1use std::sync::LazyLock;
9
10use axum::{
11 Form,
12 extract::State,
13 response::{Html, IntoResponse, Response},
14};
15use axum_extra::extract::Query;
16use mas_axum_utils::{
17 InternalError,
18 cookies::CookieJar,
19 csrf::{CsrfExt, ProtectedForm},
20};
21use mas_data_model::{BoxClock, BoxRng, DeviceCodeGrant, normalize_user_code};
22use mas_i18n::DataLocale;
23use mas_router::UrlBuilder;
24use mas_storage::{BoxRepository, RepositoryError};
25use mas_templates::{
26 DeviceLinkContext, DeviceLinkFormField, FieldError, FormError, FormState, TemplateContext,
27 Templates,
28};
29use opentelemetry::{Key, KeyValue, metrics::Counter};
30use serde::{Deserialize, Serialize};
31
32use crate::{Limiter, METER, PreferredLanguage, RequesterFingerprint, SiteConfig};
33
34static USER_CODE_ATTEMPT_COUNTER: LazyLock<Counter<u64>> = LazyLock::new(|| {
35 METER
36 .u64_counter("mas.oauth2.device_code_link_attempt")
37 .with_description("Number of user codes submitted on the device link page")
38 .with_unit("{attempt}")
39 .build()
40});
41const RESULT: Key = Key::from_static_str("result");
42
43#[derive(Serialize, Deserialize)]
44pub struct Params {
45 #[serde(default)]
46 code: Option<String>,
47}
48
49#[tracing::instrument(name = "handlers.oauth2.device.link.get", skip_all)]
50pub(crate) async fn get(
51 mut rng: BoxRng,
52 clock: BoxClock,
53 repo: BoxRepository,
54 PreferredLanguage(locale): PreferredLanguage,
55 State(templates): State<Templates>,
56 State(url_builder): State<UrlBuilder>,
57 State(site_config): State<SiteConfig>,
58 State(limiter): State<Limiter>,
59 requester: RequesterFingerprint,
60 cookie_jar: CookieJar,
61 Query(mut query): Query<Params>,
62) -> Result<Response, InternalError> {
63 if !site_config.device_code_grant_enabled {
64 return Err(InternalError::from_anyhow(anyhow::anyhow!(
65 "The Device Authorization Grant is disabled"
66 )));
67 }
68 if !site_config.device_code_user_code_auto_fill_enabled {
71 query.code = None;
72 }
73
74 handle_code(
75 &mut rng,
76 &clock,
77 repo,
78 &locale,
79 &templates,
80 &url_builder,
81 &limiter,
82 requester,
83 cookie_jar,
84 query,
85 )
86 .await
87}
88
89#[tracing::instrument(name = "handlers.oauth2.device.link.post", skip_all)]
90pub(crate) async fn post(
91 mut rng: BoxRng,
92 clock: BoxClock,
93 repo: BoxRepository,
94 PreferredLanguage(locale): PreferredLanguage,
95 State(templates): State<Templates>,
96 State(url_builder): State<UrlBuilder>,
97 State(site_config): State<SiteConfig>,
98 State(limiter): State<Limiter>,
99 requester: RequesterFingerprint,
100 cookie_jar: CookieJar,
101 Form(form): Form<ProtectedForm<Params>>,
102) -> Result<Response, InternalError> {
103 if !site_config.device_code_grant_enabled {
104 return Err(InternalError::from_anyhow(anyhow::anyhow!(
105 "The Device Authorization Grant is disabled"
106 )));
107 }
108
109 let form = cookie_jar.verify_form(&clock, form)?;
110
111 handle_code(
112 &mut rng,
113 &clock,
114 repo,
115 &locale,
116 &templates,
117 &url_builder,
118 &limiter,
119 requester,
120 cookie_jar,
121 form,
122 )
123 .await
124}
125
126async fn find_usable_grant(
134 repo: &mut BoxRepository,
135 clock: &BoxClock,
136 user_code: &str,
137) -> Result<Option<DeviceCodeGrant>, RepositoryError> {
138 Ok(repo
139 .oauth2_device_code_grant()
140 .find_by_user_code(user_code)
141 .await?
142 .filter(|grant| grant.is_pending())
144 .filter(|grant| grant.expires_at > clock.now()))
145}
146
147#[expect(clippy::too_many_arguments)]
148async fn handle_code(
149 rng: &mut BoxRng,
150 clock: &BoxClock,
151 mut repo: BoxRepository,
152 locale: &DataLocale,
153 templates: &Templates,
154 url_builder: &UrlBuilder,
155 limiter: &Limiter,
156 requester: RequesterFingerprint,
157 cookie_jar: CookieJar,
158 params: Params,
159) -> Result<Response, InternalError> {
160 let mut form_state = FormState::from_form(¶ms);
161
162 if let Some(code) = ¶ms.code {
164 if let Err(e) = limiter.check_device_code_link(requester) {
170 tracing::warn!(error = &e as &dyn std::error::Error, "ratelimit exceeded");
171 USER_CODE_ATTEMPT_COUNTER.add(1, &[KeyValue::new(RESULT, "rate_limited")]);
172
173 let (csrf_token, cookie_jar) = cookie_jar.csrf_token(clock, rng);
174 let ctx = DeviceLinkContext::new()
175 .with_form_state(form_state.with_error_on_form(FormError::RateLimitExceeded))
176 .with_csrf(csrf_token.form_value())
177 .with_language(*locale);
178
179 let content = templates.render_device_link(&ctx)?;
180
181 return Ok((cookie_jar, Html(content)).into_response());
182 }
183
184 let uppercased = code.to_uppercase();
193 let normalized = normalize_user_code(code);
194
195 let mut grant = find_usable_grant(&mut repo, clock, &uppercased).await?;
201
202 if grant.is_none() && normalized != uppercased {
203 grant = find_usable_grant(&mut repo, clock, &normalized).await?;
204 }
205
206 if let Some(grant) = grant {
207 USER_CODE_ATTEMPT_COUNTER.add(1, &[KeyValue::new(RESULT, "success")]);
210 let destination = url_builder.redirect(&mas_router::DeviceCodeConsent::new(grant.id));
211
212 return Ok((cookie_jar, destination).into_response());
213 }
214
215 USER_CODE_ATTEMPT_COUNTER.add(1, &[KeyValue::new(RESULT, "invalid")]);
217 form_state = form_state.with_error_on_field(DeviceLinkFormField::Code, FieldError::Invalid);
218 }
219
220 let (csrf_token, cookie_jar) = cookie_jar.csrf_token(clock, rng);
221
222 let ctx = DeviceLinkContext::new()
224 .with_form_state(form_state)
225 .with_csrf(csrf_token.form_value())
226 .with_language(*locale);
227
228 let content = templates.render_device_link(&ctx)?;
229
230 Ok((cookie_jar, Html(content)).into_response())
231}
232
233#[cfg(test)]
234mod tests {
235 use std::net::{IpAddr, Ipv4Addr};
236
237 use chrono::Duration;
238 use hyper::{
239 Request, StatusCode,
240 header::{CONTENT_TYPE, LOCATION},
241 };
242 use mas_data_model::Client;
243 use mas_router::{Route, SimpleRoute};
244 use mas_storage::oauth2::OAuth2DeviceCodeGrantParams;
245 use oauth2_types::{
246 registration::ClientRegistrationResponse, requests::DeviceAuthorizationResponse,
247 scope::OPENID,
248 };
249 use sqlx::PgPool;
250
251 use crate::test_utils::{CookieHelper, RequestBuilderExt, ResponseExt, TestState, setup};
252
253 const ALICE: IpAddr = IpAddr::V4(Ipv4Addr::new(1, 2, 3, 4));
254 const BOB: IpAddr = IpAddr::V4(Ipv4Addr::new(4, 3, 2, 1));
255
256 async fn device_client(state: &TestState) -> Client {
258 let request =
259 Request::post(mas_router::OAuth2RegistrationEndpoint::PATH).json(serde_json::json!({
260 "client_uri": "https://example.com/",
261 "token_endpoint_auth_method": "none",
262 "grant_types": ["urn:ietf:params:oauth:grant-type:device_code"],
263 "response_types": [],
264 }));
265
266 let response = state.request(request).await;
267 response.assert_status(StatusCode::CREATED);
268 let response: ClientRegistrationResponse = response.json();
269
270 let mut repo = state.repository().await.unwrap();
271 let client = repo
272 .oauth2_client()
273 .find_by_client_id(&response.client_id)
274 .await
275 .unwrap()
276 .unwrap();
277 repo.save().await.unwrap();
278 client
279 }
280
281 async fn get_user_code(state: &TestState) -> String {
283 let client = device_client(state).await;
284
285 let request = Request::post(mas_router::OAuth2DeviceAuthorizationEndpoint::PATH).form(
286 serde_json::json!({
287 "client_id": client.client_id,
288 "scope": "openid",
289 }),
290 );
291 let response = state.request(request).await;
292 response.assert_status(StatusCode::OK);
293 let response: DeviceAuthorizationResponse = response.json();
294
295 response.user_code
296 }
297
298 async fn grant_with_user_code(state: &TestState, client: &Client, user_code: &str) {
302 grant_with_user_code_expiring_in(
303 state,
304 client,
305 user_code,
306 Duration::try_minutes(20).unwrap(),
307 )
308 .await;
309 }
310
311 async fn grant_with_user_code_expiring_in(
315 state: &TestState,
316 client: &Client,
317 user_code: &str,
318 expires_in: Duration,
319 ) {
320 let mut repo = state.repository().await.unwrap();
321 repo.oauth2_device_code_grant()
322 .add(
323 &mut state.rng(),
324 &state.clock,
325 OAuth2DeviceCodeGrantParams {
326 client,
327 scope: [OPENID].into_iter().collect(),
328 device_code: format!("devicecode-{user_code}"),
331 user_code: user_code.to_owned(),
332 expires_in,
333 ip_address: None,
334 user_agent: None,
335 },
336 )
337 .await
338 .unwrap();
339 repo.save().await.unwrap();
340 }
341
342 async fn submit_code(state: &TestState, code: &str) -> bool {
347 let uri = mas_router::DeviceCodeLink::with_code(code.to_owned()).path_and_query();
348 let response = state.request(Request::get(&*uri).empty()).await;
349
350 match response.status() {
351 StatusCode::SEE_OTHER => {
352 let location = response.headers().get(LOCATION).unwrap().to_str().unwrap();
353 assert!(
354 location.contains("/device/"),
355 "expected a redirect to the consent page, got {location:?}"
356 );
357 true
358 }
359 StatusCode::OK => {
361 assert!(response.body().contains("mfa-code-input"));
362 false
363 }
364 status => panic!("unexpected status {status}"),
365 }
366 }
367
368 #[sqlx::test(migrator = "mas_storage_pg::MIGRATOR")]
369 async fn test_link_rate_limit(pool: PgPool) {
370 setup();
371 let state = TestState::from_pool(pool).await.unwrap();
372 let cookies = CookieHelper::new();
373
374 let request = Request::get(
377 mas_router::DeviceCodeLink::default()
378 .path_and_query()
379 .as_ref(),
380 )
381 .client_ip(ALICE)
382 .empty();
383 let request = cookies.with_cookies(request);
384 let response = state.request(request).await;
385 cookies.save_cookies(&response);
386 response.assert_status(StatusCode::OK);
387 response.assert_header_value(CONTENT_TYPE, "text/html; charset=utf-8");
388 let csrf_token = response
389 .body()
390 .split("name=\"csrf\" value=\"")
391 .nth(1)
392 .unwrap()
393 .split('\"')
394 .next()
395 .unwrap()
396 .to_owned();
397
398 let request = Request::post(mas_router::DeviceCodeLink::route())
399 .client_ip(ALICE)
400 .form(serde_json::json!({
401 "csrf": csrf_token,
402 "code": "AAAAAA",
403 }));
404 let request = cookies.with_cookies(request);
405
406 for _ in 0..10 {
409 let response = state.request(request.clone()).await;
410 response.assert_status(StatusCode::OK);
411 let body = response.body();
412 assert!(body.contains(r#"data-error-kind="invalid""#));
413 assert!(!body.contains("too many requests"));
414 }
415
416 let response = state.request(request.clone()).await;
418 response.assert_status(StatusCode::OK);
419 let body = response.body();
420 assert!(!body.contains(r#"data-error-kind="invalid""#));
421 assert!(body.contains("too many requests"));
422
423 let user_code = get_user_code(&state).await;
426 let request = Request::get(
427 mas_router::DeviceCodeLink::with_code(user_code.clone())
428 .path_and_query()
429 .as_ref(),
430 )
431 .client_ip(ALICE)
432 .empty();
433 let request = cookies.with_cookies(request);
434 let response = state.request(request).await;
435 response.assert_status(StatusCode::OK);
436 assert!(response.body().contains("too many requests"));
437
438 let request = Request::get(
440 mas_router::DeviceCodeLink::with_code(user_code)
441 .path_and_query()
442 .as_ref(),
443 )
444 .client_ip(BOB)
445 .empty();
446 let response = state.request(request).await;
447 response.assert_status(StatusCode::SEE_OTHER);
448 }
449
450 #[sqlx::test(migrator = "mas_storage_pg::MIGRATOR")]
451 async fn test_link_valid_code(pool: PgPool) {
452 setup();
453 let state = TestState::from_pool(pool).await.unwrap();
454 let user_code = get_user_code(&state).await;
455
456 let request = Request::get(
458 mas_router::DeviceCodeLink::with_code(user_code)
459 .path_and_query()
460 .as_ref(),
461 )
462 .client_ip(ALICE)
463 .empty();
464 let response = state.request(request).await;
465 response.assert_status(StatusCode::SEE_OTHER);
466 }
467
468 #[sqlx::test(migrator = "mas_storage_pg::MIGRATOR")]
471 async fn test_link_applies_the_decode_mapping(pool: PgPool) {
472 setup();
473 let state = TestState::from_pool(pool).await.unwrap();
474 let client = device_client(&state).await;
475
476 grant_with_user_code(&state, &client, "D0WK1B").await;
478
479 assert!(submit_code(&state, "D0WK1B").await);
481 assert!(submit_code(&state, "d0wk1b").await);
483 assert!(submit_code(&state, "DOWK1B").await);
485 assert!(submit_code(&state, "D0WKIB").await);
486 assert!(submit_code(&state, "D0WKLB").await);
487 assert!(submit_code(&state, "dowklb").await);
488 }
489
490 #[sqlx::test(migrator = "mas_storage_pg::MIGRATOR")]
494 async fn test_link_does_not_tolerate_separators(pool: PgPool) {
495 setup();
496 let state = TestState::from_pool(pool).await.unwrap();
497 let client = device_client(&state).await;
498
499 grant_with_user_code(&state, &client, "D0WK1B").await;
500
501 assert!(!submit_code(&state, "D0W-K1B").await);
502 assert!(!submit_code(&state, "D0W K1B").await);
503 assert!(!submit_code(&state, " D0WK1B ").await);
504 }
505
506 #[sqlx::test(migrator = "mas_storage_pg::MIGRATOR")]
515 async fn test_link_legacy_code_still_resolves(pool: PgPool) {
516 setup();
517 let state = TestState::from_pool(pool).await.unwrap();
518 let client = device_client(&state).await;
519
520 grant_with_user_code(&state, &client, "XIL9OZ").await;
523 grant_with_user_code(&state, &client, "QUIRKY").await;
524
525 assert!(submit_code(&state, "XIL9OZ").await);
526 assert!(submit_code(&state, "xil9oz").await);
527 assert!(submit_code(&state, "QUIRKY").await);
528 }
529
530 #[sqlx::test(migrator = "mas_storage_pg::MIGRATOR")]
534 async fn test_link_legacy_and_new_codes_do_not_collide(pool: PgPool) {
535 setup();
536 let state = TestState::from_pool(pool).await.unwrap();
537 let client = device_client(&state).await;
538
539 grant_with_user_code(&state, &client, "DOWK7B").await;
540 grant_with_user_code(&state, &client, "D0WK7B").await;
541
542 let mut repo = state.repository().await.unwrap();
543 let legacy = repo
544 .oauth2_device_code_grant()
545 .find_by_user_code("DOWK7B")
546 .await
547 .unwrap()
548 .unwrap();
549 let current = repo
550 .oauth2_device_code_grant()
551 .find_by_user_code("D0WK7B")
552 .await
553 .unwrap()
554 .unwrap();
555 repo.save().await.unwrap();
556 assert_ne!(legacy.id, current.id);
557
558 let uri = mas_router::DeviceCodeLink::with_code("DOWK7B".to_owned()).path_and_query();
561 let response = state.request(Request::get(&*uri).empty()).await;
562 response.assert_status(StatusCode::SEE_OTHER);
563 let location = response.headers().get(LOCATION).unwrap().to_str().unwrap();
564 assert!(location.contains(&legacy.id.to_string()));
565
566 let uri = mas_router::DeviceCodeLink::with_code("D0WK7B".to_owned()).path_and_query();
567 let response = state.request(Request::get(&*uri).empty()).await;
568 response.assert_status(StatusCode::SEE_OTHER);
569 let location = response.headers().get(LOCATION).unwrap().to_str().unwrap();
570 assert!(location.contains(¤t.id.to_string()));
571 }
572
573 #[sqlx::test(migrator = "mas_storage_pg::MIGRATOR")]
574 async fn test_link_invalid_code(pool: PgPool) {
575 setup();
576 let state = TestState::from_pool(pool).await.unwrap();
577 let client = device_client(&state).await;
578
579 grant_with_user_code(&state, &client, "D0WK7B").await;
580
581 assert!(!submit_code(&state, "ZZZZZZ").await);
582 assert!(!submit_code(&state, "").await);
583 assert!(!submit_code(&state, "D0WK7C").await);
585 }
586
587 #[sqlx::test(migrator = "mas_storage_pg::MIGRATOR")]
597 async fn test_link_expired_grant_does_not_shadow_the_folded_lookup(pool: PgPool) {
598 setup();
599 let state = TestState::from_pool(pool).await.unwrap();
600 let client = device_client(&state).await;
601
602 grant_with_user_code_expiring_in(
604 &state,
605 &client,
606 "DOWK1B",
607 Duration::try_minutes(-20).unwrap(),
608 )
609 .await;
610 grant_with_user_code(&state, &client, "D0WK1B").await;
612
613 let mut repo = state.repository().await.unwrap();
614 let live = repo
615 .oauth2_device_code_grant()
616 .find_by_user_code("D0WK1B")
617 .await
618 .unwrap()
619 .unwrap();
620 repo.save().await.unwrap();
621
622 let uri = mas_router::DeviceCodeLink::with_code("DOWK1B".to_owned()).path_and_query();
623 let response = state.request(Request::get(&*uri).empty()).await;
624 response.assert_status(StatusCode::SEE_OTHER);
625 let location = response.headers().get(LOCATION).unwrap().to_str().unwrap();
626 assert!(
627 location.contains(&live.id.to_string()),
628 "expected the live grant {}, got a redirect to {location:?}",
629 live.id
630 );
631 }
632
633 #[sqlx::test(migrator = "mas_storage_pg::MIGRATOR")]
642 async fn test_link_decode_mapping_is_one_directional(pool: PgPool) {
643 setup();
644 let state = TestState::from_pool(pool).await.unwrap();
645 let client = device_client(&state).await;
646
647 grant_with_user_code(&state, &client, "DOWK7B").await;
648
649 assert!(submit_code(&state, "DOWK7B").await);
650 assert!(!submit_code(&state, "D0WK7B").await);
651 }
652}