Skip to main content

tuwunel_service/oauth/
mod.rs

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
44/// Per-client-IP token-bucket table: last-refill instant and remaining tokens.
45type 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
83/// Cap on the rate-limit table; fully refilled buckets are pruned past it.
84const RATELIMIT_MAP_CAP: usize = 1 << 16;
85
86#[cfg(test)]
87mod tests;
88
89/// Always-on throttle for the RFC 8628 device user-code endpoints. The
90/// `user_code` is low-entropy by design (§6.1), so §5.1 requires bounding
91/// guesses regardless of the optional `oidc_rc_*` knobs; the burst stays
92/// generous for the one code a real user enters.
93const DEVICE_RC_PER_SECOND: f64 = 1.0;
94/// Generous allowance for the one code a real user enters.
95const DEVICE_RC_BURST: f64 = 60.0;
96
97/// Shared per-client-IP token-bucket throttle for the OIDC endpoints. A no-op
98/// unless both `oidc_rc_per_second` and `oidc_rc_burst_count` are configured.
99#[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/// Always-on anti-brute-force throttle for the device user-code endpoints
113/// (RFC 8628 §5.1), independent of the optional `oidc_rc_*` knobs.
114#[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/// Remove all session state for a user. For debug and developer use only;
186/// deleting state can cause registration conflicts and unintended
187/// re-registrations.
188#[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/// Revoke all session tokens for a user.
201#[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/// Get user's authorizations. Lists pairs of `(Provider, Session)` for a user.
216#[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/// Network request to a Provider returning userinfo for a Session. The session
228/// must have a valid access token.
229#[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/// Network request to a Provider returning information for a Session based on
250/// its access token.
251#[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/// Network request to a Provider revoking a Session's token.
274#[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/// Network request to a Provider to obtain an access token for a Session using
302/// a provided code.
303#[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/// Send a request to a provider; this is somewhat abstract since URL's are
343/// formed prior to this call and could point at anything, however this function
344/// uses the oauth-specific http client and is configured for JSON with special
345/// casing for an `error` property in the response.
346#[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/// Generate a unique-id string determined by the combination of `Provider` and
403/// `Session` instances.
404#[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/// Generate a unique-id string determined by the combination of `Provider`
410/// instance and `sub` string.
411#[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/// Generate a unique-id string determined by the combination of `issuer_url`
417/// and `Session` instance.
418#[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
423/// Generate a unique-id string determined by the `issuer_url` and the `sub`
424/// strings directly.
425pub 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
448/// Issuer string used as input to the identity hash. Pinned per-brand for
449/// providers whose published issuer has changed under us, so existing account
450/// associations survive the change.
451fn 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}