tuwunel_service/oauth/server/
auth.rs1use 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#[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 pub idp_id: Option<String>,
34
35 #[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 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#[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#[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#[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#[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#[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#[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#[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 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 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#[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
315fn 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
330fn 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
343fn 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}