Skip to main content

tuwunel_service/oauth/server/
auth.rs

1use std::time::{Duration, SystemTime};
2
3use base64::{Engine as _, engine::general_purpose::URL_SAFE_NO_PAD as b64};
4use ruma::OwnedUserId;
5use serde::{Deserialize, Serialize};
6use tuwunel_core::{Err, Result, err, implement, utils, utils::hash::sha256};
7use tuwunel_database::{Cbor, Deserialized};
8
9/// A pending authorization request, kept from the authorize redirect until the
10/// completion that mints its code.
11///
12/// Native sign-in binds it to one upstream provider or claims it for the local
13/// branch, and completion takes it exactly once.
14#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
15pub struct AuthRequest {
16	pub client_id: String,
17
18	pub redirect_uri: String,
19
20	pub scope: String,
21
22	pub state: Option<String>,
23
24	pub nonce: Option<String>,
25
26	pub code_challenge: Option<String>,
27
28	pub code_challenge_method: Option<String>,
29
30	/// The identity provider ID used to authenticate the user for this
31	/// authorization request. Stored so it can be propagated to the device
32	/// at token exchange time and used for UIAA SSO provider binding.
33	pub idp_id: Option<String>,
34
35	/// Whether a local login or registration has claimed this request.
36	///
37	/// A claim excludes provider selection, as `idp_id` excludes the local
38	/// branch.
39	#[serde(default)]
40	pub local_auth_selected: bool,
41
42	pub response_mode: Option<String>,
43
44	pub created_at: SystemTime,
45
46	pub expires_at: SystemTime,
47}
48
49#[derive(Clone, Debug, Deserialize, Serialize)]
50pub struct AuthCodeSession {
51	pub code: String,
52
53	pub client_id: String,
54
55	pub redirect_uri: String,
56
57	pub scope: String,
58
59	pub state: Option<String>,
60
61	pub nonce: Option<String>,
62
63	pub code_challenge: Option<String>,
64
65	pub code_challenge_method: Option<String>,
66
67	pub user_id: OwnedUserId,
68
69	/// Propagated from the originating AuthRequest; identifies which IdP
70	/// authenticated the user so the device can be tagged at token exchange.
71	pub idp_id: Option<String>,
72
73	pub created_at: SystemTime,
74
75	pub expires_at: SystemTime,
76}
77
78pub const AUTH_REQUEST_LIFETIME: Duration = Duration::from_mins(10);
79const AUTH_CODE_LIFETIME: Duration = Duration::from_mins(10);
80const AUTH_CODE_LENGTH: usize = 64;
81
82#[implement(super::Server)]
83#[must_use]
84pub fn create_auth_code(&self, auth_req: &AuthRequest, user_id: OwnedUserId) -> String {
85	let now = SystemTime::now();
86	let code = utils::random_string(AUTH_CODE_LENGTH);
87	let session = AuthCodeSession {
88		code: code.clone(),
89		client_id: auth_req.client_id.clone(),
90		redirect_uri: auth_req.redirect_uri.clone(),
91		scope: auth_req.scope.clone(),
92		state: auth_req.state.clone(),
93		nonce: auth_req.nonce.clone(),
94		code_challenge: auth_req.code_challenge.clone(),
95		code_challenge_method: auth_req.code_challenge_method.clone(),
96		user_id,
97		idp_id: auth_req.idp_id.clone(),
98		created_at: now,
99		expires_at: now.checked_add(AUTH_CODE_LIFETIME).unwrap_or(now),
100	};
101
102	self.db
103		.oidccode_authsession
104		.raw_put(&*code, Cbor(&session));
105
106	code
107}
108
109#[implement(super::Server)]
110pub fn store_auth_request(&self, req_id: &str, request: &AuthRequest) {
111	self.db
112		.oidcreqid_authrequest
113		.raw_put(req_id, Cbor(request));
114}
115
116/// Read an authorization request without consuming it.
117///
118/// A flow that pauses for a user gesture reads the request to decide what to
119/// show, then takes it with `take_auth_request` or retires it with
120/// `retire_auth_request` once the gesture arrives. An unknown or expired request
121/// is a `NotFound`, and an expired one is evicted as it is found, so no caller
122/// ever renders against a stale request.
123#[implement(super::Server)]
124pub async fn peek_auth_request(&self, req_id: &str) -> Result<AuthRequest> {
125	let request = self
126		.db
127		.oidcreqid_authrequest
128		.get(req_id)
129		.await
130		.deserialized()
131		.map(|cbor: Cbor<AuthRequest>| cbor.0)
132		.map_err(|_| err!(Request(NotFound("Unknown or expired authorization request"))))?;
133
134	if SystemTime::now() > request.expires_at {
135		self.remove_auth_request(req_id);
136
137		return Err!(Request(NotFound("Authorization request has expired")));
138	}
139
140	Ok(request)
141}
142
143/// Bind a native authorization request to one selected upstream provider.
144///
145/// A different provider or a local claim is refused, so one request cannot
146/// complete through two branches. Choosing the same provider again is
147/// harmless, so a repeated click still redirects.
148#[implement(super::Server)]
149pub async fn bind_auth_request_to_provider(&self, req_id: &str, provider_id: &str) -> Result {
150	self.update_auth_request(req_id, |request| {
151		no_other_provider(request, Some(provider_id))
152			.and_then(local_unclaimed)
153			.map(|request| AuthRequest {
154				idp_id: Some(provider_id.to_owned()),
155				..request
156			})
157	})
158	.await
159}
160
161/// Check that a pending request may still complete through the local branch.
162///
163/// A request bound to an upstream provider is refused. One the local branch
164/// already claimed passes, so a resubmitted form still completes.
165#[implement(super::Server)]
166pub async fn check_local_auth_request(&self, req_id: &str) -> Result {
167	self.peek_auth_request(req_id)
168		.await
169		.and_then(|request| no_other_provider(request, None))
170		.map(drop)
171}
172
173/// Claim the local branch of a native authorization request.
174///
175/// A request already bound to a provider is refused. Repeating the claim is
176/// harmless, so a resubmitted form still completes, and the request stays
177/// single-use when it is taken.
178#[implement(super::Server)]
179pub async fn bind_auth_request_to_local(&self, req_id: &str) -> Result {
180	self.update_auth_request(req_id, |request| {
181		no_other_provider(request, None)
182			.map(|request| AuthRequest { local_auth_selected: true, ..request })
183	})
184	.await
185}
186
187/// Take a pending request exactly once, removing it.
188///
189/// A request that changed since the caller read it is refused. The lock is the
190/// one provider selection and local claims acquire, so neither can change the
191/// request between the comparison and the removal.
192#[implement(super::Server)]
193pub async fn take_auth_request(
194	&self,
195	req_id: &str,
196	expected: &AuthRequest,
197) -> Result<AuthRequest> {
198	let _lock = self.auth_request_locks.lock(req_id).await;
199	let request = self.peek_auth_request(req_id).await?;
200
201	if &request != expected {
202		return Err!(Request(Forbidden("Authorization request changed during sign-in")));
203	}
204
205	self.remove_auth_request(req_id);
206
207	Ok(request)
208}
209
210/// Retire a pending request under its lock without reading it.
211///
212/// A refusal needs nothing from the request, and holding the lock keeps a
213/// concurrent provider selection from writing it back.
214#[implement(super::Server)]
215pub async fn retire_auth_request(&self, req_id: &str) {
216	let _lock = self.auth_request_locks.lock(req_id).await;
217
218	self.remove_auth_request(req_id);
219}
220
221/// Remove a pending authorization request without taking its lock.
222///
223/// The request is single-use, so a flow removes it before minting anything
224/// against it. Removing a key that is already gone is a no-op.
225#[implement(super::Server)]
226pub fn remove_auth_request(&self, req_id: &str) { self.db.oidcreqid_authrequest.remove(req_id); }
227
228#[implement(super::Server)]
229pub async fn exchange_auth_code(
230	&self,
231	code: &str,
232	client_id: &str,
233	redirect_uri: &str,
234	code_verifier: Option<&str>,
235	require_pkce: bool,
236) -> Result<AuthCodeSession> {
237	let session: AuthCodeSession = self
238		.db
239		.oidccode_authsession
240		.get(code)
241		.await
242		.deserialized::<Cbor<_>>()
243		.map(|cbor: Cbor<AuthCodeSession>| cbor.0)
244		.map_err(|_| err!(Request(Forbidden("Invalid or expired authorization code"))))?;
245
246	self.db.oidccode_authsession.remove(code);
247
248	if SystemTime::now() > session.expires_at {
249		return Err!(Request(Forbidden("Authorization code has expired")));
250	}
251	if session.client_id != client_id {
252		return Err!(Request(Forbidden("client_id mismatch")));
253	}
254	if session.redirect_uri != redirect_uri {
255		return Err!(Request(Forbidden("redirect_uri mismatch")));
256	}
257
258	let Some(challenge) = &session.code_challenge else {
259		// Reject a challenge-less code when PKCE is required: the knob is
260		// reloadable and codes outlive an off->on flip of it.
261		if require_pkce {
262			return Err!(Request(Forbidden(
263				"the authorization request carried no PKCE code_challenge"
264			)));
265		}
266
267		return Ok(session);
268	};
269
270	let Some(verifier) = code_verifier else {
271		return Err!(Request(Forbidden("code_verifier required for PKCE")));
272	};
273
274	validate_code_verifier(verifier)?;
275
276	let method = session
277		.code_challenge_method
278		.as_deref()
279		.unwrap_or("S256");
280
281	// Only S256 is advertised in discovery metadata; reject plain to avoid
282	// downgrade attacks (plain challenge == verifier, trivially intercepted).
283	let computed = match method {
284		| "S256" => b64.encode(sha256::hash(verifier.as_bytes())),
285		| _ => return Err!(Request(InvalidParam("Unsupported code_challenge_method"))),
286	};
287
288	if computed != *challenge {
289		return Err!(Request(Forbidden("PKCE verification failed")));
290	}
291
292	Ok(session)
293}
294
295/// Rewrite a pending request under its lock.
296///
297/// A selection decided against one read cannot then be lost to a concurrent
298/// one.
299#[implement(super::Server)]
300async fn update_auth_request<F>(&self, req_id: &str, update: F) -> Result
301where
302	F: FnOnce(AuthRequest) -> Result<AuthRequest> + Send,
303{
304	let _lock = self.auth_request_locks.lock(req_id).await;
305	let request = self
306		.peek_auth_request(req_id)
307		.await
308		.and_then(update)?;
309
310	self.store_auth_request(req_id, &request);
311
312	Ok(())
313}
314
315/// Refuse a request bound to any provider other than `provider`.
316///
317/// The local branch passes `None`, so every provider binding refuses it.
318fn no_other_provider(request: AuthRequest, provider: Option<&str>) -> Result<AuthRequest> {
319	if request
320		.idp_id
321		.as_deref()
322		.is_some_and(|bound| Some(bound) != provider)
323	{
324		return Err!(Request(Forbidden("Authorization request already selected a provider")));
325	}
326
327	Ok(request)
328}
329
330/// Refuse a request the local branch has already claimed.
331///
332/// A claim excludes provider selection but not a repeated local claim.
333fn local_unclaimed(request: AuthRequest) -> Result<AuthRequest> {
334	if request.local_auth_selected {
335		return Err!(Request(Forbidden(
336			"A local sign-in already claimed this authorization request"
337		)));
338	}
339
340	Ok(request)
341}
342
343/// Validate code_verifier per RFC 7636 Section 4.1: must be 43-128
344/// characters using only unreserved characters [A-Z] / [a-z] / [0-9] /
345/// "-" / "." / "_" / "~".
346fn validate_code_verifier(verifier: &str) -> Result {
347	if !(43..=128).contains(&verifier.len()) {
348		return Err!(Request(InvalidParam("code_verifier must be 43-128 characters")));
349	}
350
351	if !verifier
352		.bytes()
353		.all(|b| b.is_ascii_alphanumeric() || b == b'-' || b == b'.' || b == b'_' || b == b'~')
354	{
355		return Err!(Request(InvalidParam("code_verifier contains invalid characters")));
356	}
357
358	Ok(())
359}