tuwunel_service/globals/
mod.rs1mod data;
8
9use std::{ops::Range, sync::Arc};
10
11pub use data::Data;
16use ruma::{OwnedUserId, RoomAliasId, ServerName, UserId};
17use tuwunel_core::{
18 Result, Server, err,
19 utils::{Secret, resolve_secret},
20};
21
22use crate::service;
23
24pub struct Service {
29 pub db: Data,
31 server: Arc<Server>,
32
33 pub server_user: OwnedUserId,
35}
36
37impl crate::Service for Service {
38 fn build(args: &crate::Args<'_>) -> Result<Arc<Self>> {
39 let db = Data::new(args);
40
41 let server_user =
42 server_user(&args.server.config.server_user_localpart, &args.server.name)?;
43
44 Ok(Arc::new(Self {
45 db,
46 server: args.server.clone(),
47 server_user,
48 }))
49 }
50
51 fn name(&self) -> &str { service::make_name(std::module_path!()) }
52}
53
54fn server_user(localpart: &str, server_name: &ServerName) -> Result<OwnedUserId> {
59 UserId::parse_with_server_name(localpart, server_name)
60 .map_err(|e| err!("Invalid server_user_localpart configuration: {e}"))
61 .and_then(|user| {
62 user.localpart()
63 .eq(localpart)
64 .then_some(user)
65 .ok_or_else(|| {
66 err!("server_user_localpart must be a localpart, not a full user ID")
67 })
68 })
69}
70
71impl Service {
72 #[tracing::instrument(
77 level = "trace",
78 skip_all,
79 ret,
80 fields(pending = ?self.pending_count()),
81 )]
82 pub async fn wait_pending(&self) -> Result<u64> { self.db.wait_pending().await }
83
84 #[tracing::instrument(
89 level = "trace",
90 skip_all,
91 ret,
92 fields(pending = ?self.pending_count()),
93 )]
94 pub async fn wait_count(&self, count: &u64) -> Result<u64> { self.db.wait_count(count).await }
95
96 #[tracing::instrument(
105 level = "debug",
106 skip_all,
107 fields(pending = ?self.pending_count()),
108 )]
109 #[must_use]
110 pub fn next_count(&self) -> data::Permit { self.db.next_count() }
111
112 #[must_use]
116 pub fn current_count(&self) -> u64 { self.db.current_count() }
117
118 #[must_use]
123 pub fn pending_count(&self) -> Range<u64> { self.db.pending_count() }
124
125 #[inline]
129 #[must_use]
130 pub fn server_name(&self) -> &ServerName { self.server.name.as_ref() }
131
132 #[inline]
136 #[must_use]
137 pub fn user_is_local(&self, user_id: &UserId) -> bool {
138 self.server_is_ours(user_id.server_name())
139 }
140
141 #[inline]
146 #[must_use]
147 pub fn alias_is_local(&self, alias: &RoomAliasId) -> bool {
148 self.server_is_ours(alias.server_name())
149 }
150
151 #[inline]
156 #[must_use]
157 pub fn server_is_ours(&self, server_name: &ServerName) -> bool {
158 server_name == self.server_name()
159 }
160
161 #[inline]
165 #[must_use]
166 pub fn is_read_only(&self) -> bool { self.db.db.is_read_only() }
167
168 #[must_use]
174 pub fn turn_secret(&self) -> Option<Secret> {
175 let config = &self.server.config;
176
177 resolve_secret(
178 config.turn_secret_file.as_deref(),
179 config.turn_secret.as_deref(),
180 "TURN secret",
181 )
182 }
183
184 pub fn init_rustls_provider(&self) -> Result {
189 if rustls::crypto::CryptoProvider::get_default().is_none() {
190 rustls::crypto::aws_lc_rs::default_provider()
191 .install_default()
192 .map_err(|_provider| {
193 err!(error!("Error initialising aws_lc_rs rustls crypto backend"))
194 })
195 } else {
196 Ok(())
197 }
198 }
199}
200
201#[cfg(test)]
202mod tests {
203 use ruma::server_name;
204
205 use super::server_user;
206
207 #[test]
208 fn configured_server_user_must_be_valid() {
209 let server_name = server_name!("example.org");
210
211 assert_eq!(server_user("conduit", server_name).unwrap(), "@conduit:example.org");
212 assert_eq!(server_user("_server", server_name).unwrap(), "@_server:example.org");
213 server_user("bad:localpart", server_name).unwrap_err();
214 server_user("@conduit:example.org", server_name).unwrap_err();
215 server_user("@conduit:elsewhere.org", server_name).unwrap_err();
216 }
217}