1pub mod providers;
2pub mod server;
3pub mod sessions;
4pub mod token_response;
5pub mod user_info;
6
7use std::{
8 collections::HashMap,
9 net::IpAddr,
10 sync::{Arc, Mutex},
11 time::Instant,
12};
13
14use base64::{Engine as _, engine::general_purpose::URL_SAFE_NO_PAD as b64encode};
15use futures::{Stream, StreamExt, TryStreamExt};
16use http::StatusCode;
17use reqwest::{
18 Method,
19 header::{ACCEPT, CONTENT_TYPE},
20};
21use ruma::{
22 UserId,
23 api::error::{ErrorKind, LimitExceededErrorData},
24};
25use serde::Serialize;
26use serde_json::Value as JsonValue;
27use tuwunel_core::{
28 Err, Error, Result, err, implement,
29 utils::{hash::sha256, result::LogErr, stream::ReadyExt},
30 warn,
31};
32use url::Url;
33
34use self::{providers::Providers, sessions::Sessions};
35pub use self::{
36 providers::{Provider, ProviderId},
37 server::Server,
38 sessions::{CODE_VERIFIER_LENGTH, SESSION_ID_LENGTH, Session, SessionId},
39 token_response::TokenResponse,
40 user_info::UserInfo,
41};
42use crate::{SelfServices, client::read_response_capped};
43
44type Ratelimiter = Mutex<HashMap<IpAddr, (Instant, f64)>>;
46
47pub struct Service {
48 services: SelfServices,
49 pub providers: Arc<Providers>,
50 pub sessions: Arc<Sessions>,
51 pub server: Option<Arc<Server>>,
52 ratelimiter: Ratelimiter,
53 device_ratelimiter: Ratelimiter,
54}
55
56impl crate::Service for Service {
57 fn build(args: &crate::Args<'_>) -> Result<Arc<Self>> {
58 let providers = Arc::new(Providers::build(args));
59 let sessions = Arc::new(Sessions::build(args, providers.clone()));
60 let server = Server::build(args)?.map(Arc::new);
61
62 Ok(Arc::new(Self {
63 services: args.services.clone(),
64 sessions,
65 providers,
66 server,
67 ratelimiter: Mutex::new(HashMap::new()),
68 device_ratelimiter: Mutex::new(HashMap::new()),
69 }))
70 }
71
72 fn name(&self) -> &str { crate::service::make_name(std::module_path!()) }
73}
74
75#[implement(Service)]
76#[inline]
77pub fn get_server(&self) -> Result<&Server> {
78 self.server
79 .as_deref()
80 .ok_or_else(|| err!(Request(Unrecognized("OIDC server not configured"))))
81}
82
83const RATELIMIT_MAP_CAP: usize = 1 << 16;
85
86#[cfg(test)]
87mod tests;
88
89const DEVICE_RC_PER_SECOND: f64 = 1.0;
94const DEVICE_RC_BURST: f64 = 60.0;
96
97#[implement(Service)]
100pub fn check_rate_limit(&self, client: IpAddr) -> Result {
101 let config = &self.services.config;
102 let rate = f64::from(config.oidc_rc_per_second);
103 let burst = f64::from(config.oidc_rc_burst_count);
104
105 if rate <= 0.0 || burst <= 0.0 {
106 return Ok(());
107 }
108
109 check_bucket(&self.ratelimiter, client, rate, burst)
110}
111
112#[implement(Service)]
115pub fn check_device_rate_limit(&self, client: IpAddr) -> Result {
116 check_bucket(&self.device_ratelimiter, client, DEVICE_RC_PER_SECOND, DEVICE_RC_BURST)
117}
118
119fn check_bucket(table: &Ratelimiter, client: IpAddr, rate: f64, burst: f64) -> Result {
120 check_bucket_at(table, client, rate, burst, Instant::now(), RATELIMIT_MAP_CAP)
121}
122
123fn check_bucket_at(
124 table: &Ratelimiter,
125 client: IpAddr,
126 rate: f64,
127 burst: f64,
128 now: Instant,
129 cap: usize,
130) -> Result {
131 let mut buckets = table.lock()?;
132 debug_assert!(cap > 0, "rate-limit table cap must be positive");
133 debug_assert!(buckets.len() <= cap, "rate-limit table exceeded its cap");
134
135 if buckets.len() >= cap && !buckets.contains_key(&client) {
136 let mut oldest = None;
137
138 buckets.retain(|client, bucket| {
139 let (last_time, tokens) = *bucket;
140 let refilled = now
141 .duration_since(last_time)
142 .as_secs_f64()
143 .mul_add(rate, tokens);
144
145 let retain = refilled < burst;
146
147 if retain && oldest.is_none_or(|(_, oldest_at)| last_time < oldest_at) {
148 oldest = Some((*client, last_time));
149 }
150
151 retain
152 });
153
154 if buckets.len() >= cap
155 && let Some((oldest, _)) = oldest
156 {
157 buckets.remove(&oldest);
158 }
159 }
160
161 let (last_time, tokens) = buckets
162 .entry(client)
163 .or_insert_with(|| (now, burst));
164
165 let new_tokens = now
166 .duration_since(*last_time)
167 .as_secs_f64()
168 .mul_add(rate, *tokens)
169 .min(burst);
170
171 if new_tokens < 1.0 {
172 return Err(Error::Request(
173 ErrorKind::LimitExceeded(LimitExceededErrorData { retry_after: None }),
174 "Too many OIDC requests.".into(),
175 StatusCode::TOO_MANY_REQUESTS,
176 ));
177 }
178
179 *last_time = now;
180 *tokens = new_tokens - 1.0;
181
182 Ok(())
183}
184
185#[implement(Service)]
189#[tracing::instrument(level = "debug", skip(self))]
190pub async fn delete_user_sessions(&self, user_id: &UserId) {
191 self.user_sessions(user_id)
192 .ready_filter_map(Result::ok)
193 .ready_filter_map(|(_, session)| session.sess_id)
194 .for_each(async |sess_id| {
195 self.sessions.delete(&sess_id).await;
196 })
197 .await;
198}
199
200#[implement(Service)]
202#[tracing::instrument(level = "debug", skip(self))]
203pub async fn revoke_user_tokens(&self, user_id: &UserId) {
204 self.user_sessions(user_id)
205 .ready_filter_map(Result::ok)
206 .for_each(async |(provider, session)| {
207 self.revoke_token((&provider, &session))
208 .await
209 .log_err()
210 .ok();
211 })
212 .await;
213}
214
215#[implement(Service)]
217#[tracing::instrument(level = "debug", skip(self))]
218pub fn user_sessions(
219 &self,
220 user_id: &UserId,
221) -> impl Stream<Item = Result<(Provider, Session)>> + Send {
222 self.sessions
223 .get_by_user(user_id)
224 .and_then(async |session| Ok((self.sessions.provider(&session).await?, session)))
225}
226
227#[implement(Service)]
230#[tracing::instrument(level = "debug", skip_all, ret)]
231pub async fn request_userinfo(
232 &self,
233 (provider, session): (&Provider, &Session),
234) -> Result<UserInfo> {
235 #[derive(Debug, Serialize)]
236 struct Query;
237
238 let url = provider
239 .userinfo_url
240 .clone()
241 .ok_or_else(|| err!(Config("userinfo_url", "Missing userinfo URL in config")))?;
242
243 self.request((Some(provider), Some(session)), Method::GET, url, Option::<Query>::None)
244 .await
245 .and_then(|value| serde_json::from_value(value).map_err(Into::into))
246 .log_err()
247}
248
249#[implement(Service)]
252#[tracing::instrument(level = "debug", skip_all, ret)]
253pub async fn request_tokeninfo(
254 &self,
255 (provider, session): (&Provider, &Session),
256) -> Result<UserInfo> {
257 #[derive(Debug, Serialize)]
258 struct Query;
259
260 let url = provider
261 .introspection_url
262 .clone()
263 .ok_or_else(|| {
264 err!(Config("introspection_url", "Missing introspection URL in config"))
265 })?;
266
267 self.request((Some(provider), Some(session)), Method::GET, url, Option::<Query>::None)
268 .await
269 .and_then(|value| serde_json::from_value(value).map_err(Into::into))
270 .log_err()
271}
272
273#[implement(Service)]
275#[tracing::instrument(level = "debug", skip_all, ret)]
276pub async fn revoke_token(&self, (provider, session): (&Provider, &Session)) -> Result {
277 #[derive(Debug, Serialize)]
278 struct RevokeQuery<'a> {
279 client_id: &'a str,
280 client_secret: &'a str,
281 }
282
283 let client_secret = provider.get_client_secret().await?;
284
285 let query = RevokeQuery {
286 client_id: &provider.client_id,
287 client_secret: &client_secret,
288 };
289
290 let url = provider
291 .revocation_url
292 .clone()
293 .ok_or_else(|| err!(Config("revocation_url", "Missing revocation URL in config")))?;
294
295 self.request((Some(provider), Some(session)), Method::POST, url, Some(query))
296 .await
297 .log_err()
298 .map(|_| ())
299}
300
301#[implement(Service)]
304#[tracing::instrument(level = "debug", skip_all, ret)]
305pub async fn request_token(
306 &self,
307 (provider, session): (&Provider, &Session),
308 code: &str,
309) -> Result<TokenResponse> {
310 #[derive(Debug, Serialize)]
311 struct TokenQuery<'a> {
312 client_id: &'a str,
313 client_secret: &'a str,
314 grant_type: &'a str,
315 code: &'a str,
316 code_verifier: Option<&'a str>,
317 redirect_uri: Option<&'a str>,
318 }
319
320 let client_secret = provider.get_client_secret().await?;
321
322 let query = TokenQuery {
323 client_id: &provider.client_id,
324 client_secret: &client_secret,
325 grant_type: "authorization_code",
326 code,
327 code_verifier: session.code_verifier.as_deref(),
328 redirect_uri: provider.callback_url.as_ref().map(Url::as_str),
329 };
330
331 let url = provider
332 .token_url
333 .clone()
334 .ok_or_else(|| err!(Config("token_url", "Missing token URL in config")))?;
335
336 self.request((Some(provider), Some(session)), Method::POST, url, Some(query))
337 .await
338 .and_then(|value| serde_json::from_value(value).map_err(Into::into))
339 .log_err()
340}
341
342#[implement(Service)]
347#[tracing::instrument(
348 name = "request",
349 level = "debug",
350 ret(level = "trace"),
351 skip(self, body)
352)]
353pub async fn request<Body>(
354 &self,
355 (provider, session): (Option<&Provider>, Option<&Session>),
356 method: Method,
357 url: Url,
358 body: Option<Body>,
359) -> Result<JsonValue>
360where
361 Body: Serialize,
362{
363 let mut request = self
364 .services
365 .client
366 .oauth
367 .request(method, url)
368 .header(ACCEPT, "application/json");
369
370 if let Some(body) = body.map(serde_html_form::to_string).transpose()? {
371 request = request
372 .header(CONTENT_TYPE, "application/x-www-form-urlencoded")
373 .body(body);
374 }
375
376 if let Some(session) = session
377 && let Some(access_token) = session.access_token.clone()
378 {
379 request = request.bearer_auth(access_token);
380 }
381
382 let limit = self.services.config.max_response_size;
383 let http_response = request.send().await?.error_for_status()?;
384
385 let body = read_response_capped(http_response, limit).await?;
386 let response: JsonValue = serde_json::from_slice(&body)?;
387
388 if let Some(response) = response.as_object().as_ref()
389 && let Some(error) = response.get("error").and_then(JsonValue::as_str)
390 {
391 let description = response
392 .get("error_description")
393 .and_then(JsonValue::as_str)
394 .unwrap_or("(no description)");
395
396 return Err!(Request(Forbidden("Error from provider: {error}: {description}",)));
397 }
398
399 Ok(response)
400}
401
402#[inline]
405pub fn unique_id((provider, session): (&Provider, &Session)) -> Result<String> {
406 unique_id_parts((provider, session)).and_then(unique_id_iss_sub)
407}
408
409#[inline]
412pub fn unique_id_sub((provider, sub): (&Provider, &str)) -> Result<String> {
413 unique_id_sub_parts((provider, sub)).and_then(unique_id_iss_sub)
414}
415
416#[inline]
419pub fn unique_id_iss((iss, session): (&str, &Session)) -> Result<String> {
420 unique_id_iss_parts((iss, session)).and_then(unique_id_iss_sub)
421}
422
423pub fn unique_id_iss_sub((iss, sub): (&str, &str)) -> Result<String> {
426 let hash = sha256::delimited([iss, sub].iter());
427 let b64 = b64encode.encode(hash);
428
429 Ok(b64)
430}
431
432fn unique_id_parts<'a>(
433 (provider, session): (&'a Provider, &'a Session),
434) -> Result<(&'a str, &'a str)> {
435 identity_issuer(provider)
436 .ok_or_else(|| err!(Config("issuer_url", "issuer_url not found for this provider.")))
437 .and_then(|iss| unique_id_iss_parts((iss, session)))
438}
439
440fn unique_id_sub_parts<'a>(
441 (provider, sub): (&'a Provider, &'a str),
442) -> Result<(&'a str, &'a str)> {
443 identity_issuer(provider)
444 .ok_or_else(|| err!(Config("issuer_url", "issuer_url not found for this provider.")))
445 .map(|iss| (iss, sub))
446}
447
448fn identity_issuer(provider: &Provider) -> Option<&str> {
452 match provider.brand.as_str() {
453 | "github" => Some("https://github.com/"),
454 | _ => provider.issuer_url.as_ref().map(Url::as_str),
455 }
456}
457
458fn unique_id_iss_parts<'a>((iss, session): (&'a str, &'a Session)) -> Result<(&'a str, &'a str)> {
459 session
460 .user_info
461 .as_ref()
462 .map(|user_info| user_info.sub.as_str())
463 .ok_or_else(|| err!(Request(NotFound("user_info not found for this session."))))
464 .map(|sub| (iss, sub))
465}