tuwunel_service/registration_tokens/
mod.rs1mod data;
8
9use std::{collections::HashSet, fmt::Display, sync::Arc};
10
11use data::Data;
12pub use data::{DatabaseTokenInfo, TokenExpires};
16use futures::{Stream, StreamExt, pin_mut};
17use tuwunel_core::{
18 Err, Result, error,
19 utils::{IterStream, random_string},
20};
21
22const RANDOM_TOKEN_LENGTH: usize = 16;
23
24pub struct Service {
29 db: Data,
30 services: Arc<crate::services::OnceServices>,
31}
32
33#[derive(Debug)]
38pub struct ValidToken {
39 pub token: String,
41
42 pub info: TokenInfo,
44}
45
46impl Display for ValidToken {
47 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
48 write!(f, "`{}` --- {}", self.token, self.info)
49 }
50}
51
52impl PartialEq<str> for ValidToken {
53 fn eq(&self, other: &str) -> bool { self.token == other }
54}
55
56#[derive(Clone, Copy, Debug)]
61pub enum TokenInfo {
62 Config,
64
65 Database(DatabaseTokenInfo),
67}
68
69impl Display for TokenInfo {
70 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
71 match self {
72 | Self::Config => write!(f, "Token defined in config file"),
73 | Self::Database(info) => info.fmt(f),
74 }
75 }
76}
77
78impl crate::Service for Service {
79 fn build(args: &crate::Args<'_>) -> Result<Arc<Self>> {
80 Ok(Arc::new(Self {
81 db: Data::new(args.db),
82 services: args.services.clone(),
83 }))
84 }
85
86 fn name(&self) -> &str { crate::service::make_name(std::module_path!()) }
87}
88
89impl Service {
90 pub async fn create_token(
95 &self,
96 token: Option<&str>,
97 length: Option<usize>,
98 expires: TokenExpires,
99 ) -> Result<(String, DatabaseTokenInfo)> {
100 let token = token.map(ToOwned::to_owned).unwrap_or_else(|| {
101 let length = length.unwrap_or(RANDOM_TOKEN_LENGTH);
102
103 random_string(length)
104 });
105
106 let info = self.db.save_token(&token, expires).await?;
107
108 Ok((token, info))
109 }
110
111 pub async fn get_token_info(&self, token: &str) -> Result<TokenInfo> {
116 if self.get_config_tokens().await.contains(token) {
117 return Ok(TokenInfo::Config);
118 }
119
120 self.db
121 .get_token_info(token)
122 .await
123 .map(TokenInfo::Database)
124 }
125
126 pub async fn update_token(
131 &self,
132 token: &str,
133 expires: TokenExpires,
134 ) -> Result<DatabaseTokenInfo> {
135 if self.get_config_tokens().await.contains(token) {
136 return Err!(Request(Forbidden(
137 "The token set in the config file cannot be updated"
138 )));
139 }
140
141 self.db.update_token(token, expires).await
142 }
143
144 pub async fn is_enabled(&self) -> bool {
149 let stream = self.iterate_tokens().await;
150
151 pin_mut!(stream);
152
153 stream.next().await.is_some()
154 }
155
156 pub async fn get_config_tokens(&self) -> HashSet<String> {
161 let mut tokens = HashSet::new();
162
163 if let Some(file) = &self.services.config.registration_token_file {
164 match tokio::fs::read_to_string(file).await {
165 | Err(e) => error!("Failed to read the registration token file: {e}"),
166 | Ok(text) => tokens.extend(
167 text.split_ascii_whitespace()
168 .map(ToOwned::to_owned),
169 ),
170 }
171 }
172
173 if let Some(token) = &self.services.config.registration_token {
174 tokens.insert(token.to_owned());
175 }
176
177 tokens
178 }
179
180 pub async fn is_token_valid(&self, token: &str) -> Result { self.check(token, false).await }
185
186 pub async fn try_consume(&self, token: &str) -> Result { self.check(token, true).await }
192
193 async fn check(&self, token: &str, consume: bool) -> Result {
194 if self.get_config_tokens().await.contains(token)
195 || self.db.check_token(token, consume).await
196 {
197 return Ok(());
198 }
199
200 Err!(Request(Forbidden("Registration token not valid")))
201 }
202
203 pub async fn revoke_token(&self, token: &str) -> Result {
208 if self.get_config_tokens().await.contains(token) {
209 return Err!(Request(Forbidden(
210 "The token set in the config file cannot be revoked. Edit the config file to \
211 change it."
212 )));
213 }
214
215 self.db.revoke_token(token).await
216 }
217
218 pub async fn iterate_tokens(&self) -> impl Stream<Item = ValidToken> + Send + '_ {
223 let config_tokens = self
224 .get_config_tokens()
225 .await
226 .into_iter()
227 .map(|token| ValidToken { token, info: TokenInfo::Config })
228 .stream();
229
230 let db_tokens = self
231 .db
232 .iterate_and_clean_tokens()
233 .map(|(token, info)| ValidToken {
234 token: token.to_owned(),
235 info: TokenInfo::Database(info),
236 });
237
238 config_tokens.chain(db_tokens)
239 }
240}