1use anyhow::{bail, Context as _, Result};
20use base64::engine::general_purpose::URL_SAFE_NO_PAD;
21use base64::Engine;
22use ring::digest::{digest, SHA256};
23use serde_json::Value;
24use subtle::ConstantTimeEq;
25
26const VERIFIER_BYTES: usize = 32;
29
30fn random_token(bytes: usize) -> String {
32 let mut buf = vec![0u8; bytes];
33 getrandom::fill(&mut buf).expect("OS CSPRNG unavailable; refusing to mint an OAuth secret");
34 URL_SAFE_NO_PAD.encode(&buf)
35}
36
37pub fn new_pkce_verifier() -> String {
42 random_token(VERIFIER_BYTES)
43}
44
45pub fn pkce_challenge(verifier: &str) -> String {
50 URL_SAFE_NO_PAD.encode(digest(&SHA256, verifier.as_bytes()).as_ref())
51}
52
53pub fn new_state() -> String {
59 random_token(VERIFIER_BYTES)
60}
61
62pub fn new_binding_token() -> String {
64 random_token(VERIFIER_BYTES)
65}
66
67pub fn binding_hash(token: &str) -> String {
69 URL_SAFE_NO_PAD.encode(digest(&SHA256, token.as_bytes()).as_ref())
70}
71
72pub fn binding_matches(stored_hash: &str, presented: Option<&str>) -> bool {
80 let Some(presented) = presented.filter(|t| !t.is_empty()) else {
81 return false;
82 };
83 if stored_hash.is_empty() {
84 return false;
85 }
86 let computed = binding_hash(presented);
87 computed.len() == stored_hash.len()
88 && bool::from(computed.as_bytes().ct_eq(stored_hash.as_bytes()))
89}
90
91pub struct ParRequest<'a> {
93 pub redirect_uri: &'a str,
101 pub scope: &'a str,
102 pub state: &'a str,
103 pub code_challenge: &'a str,
104 pub login_hint: Option<&'a str>,
107}
108
109pub fn par_params(request: &ParRequest<'_>) -> Vec<(&'static str, String)> {
113 let mut params = vec![
114 ("response_type", "code".to_string()),
115 ("code_challenge", request.code_challenge.to_string()),
116 ("code_challenge_method", "S256".to_string()),
118 ("state", request.state.to_string()),
119 ("redirect_uri", request.redirect_uri.to_string()),
120 ("scope", request.scope.to_string()),
121 ("response_mode", "query".to_string()),
126 ];
127 if let Some(hint) = request.login_hint {
128 params.push(("login_hint", hint.to_string()));
129 }
130 params
131}
132
133pub struct ParResponse {
135 pub request_uri: String,
136 pub expires_in: i64,
138}
139
140pub fn parse_par_response(body: &Value) -> Result<ParResponse> {
146 let request_uri = body
147 .get("request_uri")
148 .and_then(Value::as_str)
149 .context("PAR response has no `request_uri`")?;
150 if request_uri.is_empty() {
151 bail!("PAR response `request_uri` is empty");
152 }
153 let expires_in = body
154 .get("expires_in")
155 .and_then(Value::as_i64)
156 .context("PAR response has no integer `expires_in`")?;
157 if expires_in <= 0 {
158 bail!("PAR response `expires_in` must be positive, got {expires_in}");
159 }
160 Ok(ParResponse {
161 request_uri: request_uri.to_string(),
162 expires_in,
163 })
164}
165
166pub fn authorize_url(
173 authorization_endpoint: &str,
174 client_id: &str,
175 request_uri: &str,
176) -> Result<String> {
177 let mut url = url::Url::parse(authorization_endpoint).with_context(|| {
178 format!("authorization_endpoint {authorization_endpoint:?} is not a URL")
179 })?;
180 if url.scheme() != "https" {
181 bail!("authorization_endpoint must be https, got {authorization_endpoint:?}");
182 }
183 url.query_pairs_mut()
184 .clear()
185 .append_pair("client_id", client_id)
186 .append_pair("request_uri", request_uri);
187 Ok(url.to_string())
188}
189
190const KNOWN_ERRORS: [&str; 11] = [
193 "access_denied",
194 "consent_required",
195 "interaction_required",
196 "invalid_grant",
197 "invalid_request",
198 "invalid_scope",
199 "login_required",
200 "server_error",
201 "temporarily_unavailable",
202 "unauthorized_client",
203 "unsupported_response_type",
204];
205
206pub(crate) fn known_error_slug(raw: &str) -> &'static str {
214 KNOWN_ERRORS
215 .iter()
216 .find(|known| **known == raw)
217 .copied()
218 .unwrap_or("unrecognized_error")
219}
220
221#[derive(Debug, Default, Clone)]
223pub struct CallbackParams {
224 pub code: Option<String>,
225 pub state: Option<String>,
226 pub iss: Option<String>,
227 pub error: Option<String>,
228 pub error_description: Option<String>,
229 pub response: Option<String>,
231}
232
233pub fn verify_callback(params: &CallbackParams, expected_issuer: &str) -> Result<String> {
245 if params.response.is_some() {
246 bail!("authorization response uses JARM, which is not supported");
247 }
248 let state = params.state.as_deref().unwrap_or_default();
251 if state.is_empty() {
252 bail!("authorization response carries no `state`");
253 }
254 let iss = params
259 .iss
260 .as_deref()
261 .context("authorization response carries no `iss` (RFC 9207); refusing it")?;
262 if iss != expected_issuer {
263 bail!("authorization response `iss` is {iss:?}, expected {expected_issuer:?}");
264 }
265
266 if let Some(error) = params.error.as_deref() {
272 bail!(
273 "authorization server returned error {:?}",
274 known_error_slug(error)
275 );
276 }
277
278 let code = params.code.as_deref().unwrap_or_default();
279 if code.is_empty() {
280 bail!("authorization response carries no `code`");
281 }
282 Ok(code.to_string())
283}
284
285pub async fn complete_callback(
299 pool: &sqlx::SqlitePool,
300 codec: &super::crypto::Codec,
301 params: &CallbackParams,
302 presented_cookie: Option<&str>,
303 now: i64,
304) -> Result<(super::store::PendingAuth, String)> {
305 if params.response.is_some() {
306 bail!("authorization response uses JARM, which is not supported");
307 }
308 let state = params.state.as_deref().unwrap_or_default();
309 if state.is_empty() {
310 bail!("authorization response carries no `state`");
311 }
312
313 let pending = super::store::take_pending(pool, codec, state, now)
314 .await?
315 .context("no pending login for that `state` (unknown, expired, or already used)")?;
316
317 if !binding_matches(&pending.browser_binding_hash, presented_cookie) {
318 bail!(
319 "the callback did not present the browser-binding cookie for this login; \
320 refusing to complete a flow this browser did not start"
321 );
322 }
323
324 let code = verify_callback(params, &pending.issuer)?;
325 Ok((pending, code))
326}
327
328#[cfg(test)]
329mod tests {
330 use super::*;
331
332 const ISSUER: &str = "https://auth.example.com";
333
334 #[test]
338 fn a_pkce_verifier_is_within_the_unreserved_alphabet_and_length() {
339 for _ in 0..32 {
340 let verifier = new_pkce_verifier();
341 assert!(
342 (43..=128).contains(&verifier.len()),
343 "length {} out of range",
344 verifier.len()
345 );
346 assert!(
347 verifier
348 .bytes()
349 .all(|b| b.is_ascii_alphanumeric() || matches!(b, b'-' | b'.' | b'_' | b'~')),
350 "verifier outside the unreserved set: {verifier}"
351 );
352 }
353 }
354
355 #[test]
358 fn every_pkce_verifier_is_fresh() {
359 let mut seen = std::collections::HashSet::new();
360 for _ in 0..64 {
361 assert!(seen.insert(new_pkce_verifier()), "verifier repeated");
362 }
363 }
364
365 #[test]
367 fn the_challenge_is_the_s256_of_the_verifier() {
368 let verifier = "dBjftJeZ4CVP-mB92K27uhbUJU1p1r_wW1gFWFOEjXk";
370 assert_eq!(
371 pkce_challenge(verifier),
372 "E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw-cM"
373 );
374 }
375
376 #[test]
377 fn the_challenge_is_unpadded_base64url() {
378 let challenge = pkce_challenge(&new_pkce_verifier());
379 assert_eq!(URL_SAFE_NO_PAD.decode(&challenge).unwrap().len(), 32);
380 assert!(!challenge.contains('=') && !challenge.contains('+') && !challenge.contains('/'));
381 }
382
383 #[test]
389 fn a_binding_token_is_fresh_and_unguessable() {
390 let mut seen = std::collections::HashSet::new();
391 for _ in 0..64 {
392 let token = new_binding_token();
393 assert_eq!(token.len(), 43, "binding token is not 32 bytes: {token}");
394 assert!(seen.insert(token), "binding token repeated");
395 }
396 }
397
398 #[test]
399 fn a_binding_token_matches_only_its_own_hash() {
400 let token = new_binding_token();
401 let hash = binding_hash(&token);
402 assert!(binding_matches(&hash, Some(&token)));
403 assert!(!binding_matches(&hash, Some(&new_binding_token())));
404 assert!(!binding_matches(&hash, Some("")));
405 }
406
407 #[test]
410 fn a_truncated_stored_hash_does_not_match() {
411 let token = new_binding_token();
412 let hash = binding_hash(&token);
413 for len in [1, 8, hash.len() - 1] {
414 assert!(
415 !binding_matches(&hash[..len], Some(&token)),
416 "a {len}-character prefix matched"
417 );
418 }
419 }
420
421 #[test]
426 fn a_raw_token_stored_as_the_hash_does_not_match() {
427 let token = new_binding_token();
428 assert!(!binding_matches(&token, Some(&token)));
429 }
430
431 #[test]
434 fn every_state_is_fresh_and_full_length() {
435 let mut seen = std::collections::HashSet::new();
436 for _ in 0..64 {
437 let state = new_state();
438 assert_eq!(state.len(), 43, "state is not 32 bytes of base64url");
439 assert!(seen.insert(state), "state repeated");
440 }
441 }
442
443 #[test]
447 fn a_missing_cookie_never_matches() {
448 let hash = binding_hash(&new_binding_token());
449 assert!(!binding_matches(&hash, None));
450 assert!(!binding_matches("", None));
452 assert!(!binding_matches("", Some("anything")));
453 }
454
455 fn par_input() -> ParRequest<'static> {
458 ParRequest {
459 redirect_uri: "https://feather-reader.com/oauth/callback",
460 scope: "atproto transition:generic",
461 state: "state-value",
462 code_challenge: "challenge-value",
463 login_hint: Some("alice.bsky.social"),
464 }
465 }
466
467 #[test]
468 fn the_par_request_carries_the_required_parameters() {
469 let params = par_params(&par_input());
470 let get = |k: &str| {
471 params
472 .iter()
473 .find(|(name, _)| *name == k)
474 .map(|(_, v)| v.clone())
475 };
476 assert_eq!(get("response_type").as_deref(), Some("code"));
477 assert_eq!(get("code_challenge_method").as_deref(), Some("S256"));
478 assert_eq!(get("code_challenge").as_deref(), Some("challenge-value"));
479 assert_eq!(get("state").as_deref(), Some("state-value"));
480 assert_eq!(
481 get("redirect_uri").as_deref(),
482 Some("https://feather-reader.com/oauth/callback")
483 );
484 assert_eq!(get("scope").as_deref(), Some("atproto transition:generic"));
485 assert_eq!(get("login_hint").as_deref(), Some("alice.bsky.social"));
486 }
487
488 #[test]
492 fn the_challenge_method_is_always_s256() {
493 let params = par_params(&par_input());
494 let method = params
495 .iter()
496 .find(|(name, _)| *name == "code_challenge_method")
497 .map(|(_, v)| v.as_str());
498 assert_eq!(method, Some("S256"));
499 }
500
501 #[test]
508 fn the_response_mode_is_pinned_to_query() {
509 let params = par_params(&par_input());
510 assert_eq!(
511 params
512 .iter()
513 .find(|(name, _)| *name == "response_mode")
514 .map(|(_, v)| v.as_str()),
515 Some("query")
516 );
517 }
518
519 #[test]
520 fn login_hint_is_omitted_when_absent_rather_than_sent_empty() {
521 let mut input = par_input();
522 input.login_hint = None;
523 let params = par_params(&input);
524 assert!(!params.iter().any(|(name, _)| *name == "login_hint"));
525 }
526
527 #[test]
530 fn a_valid_par_response_yields_its_request_uri_and_lifetime() {
531 let body = serde_json::json!({
532 "request_uri": "urn:ietf:params:oauth:request_uri:abc",
533 "expires_in": 60
534 });
535 let parsed = parse_par_response(&body).unwrap();
536 assert_eq!(parsed.request_uri, "urn:ietf:params:oauth:request_uri:abc");
537 assert_eq!(parsed.expires_in, 60);
538 }
539
540 #[test]
543 fn a_malformed_par_response_is_rejected() {
544 for body in [
545 serde_json::json!({}),
546 serde_json::json!({ "expires_in": 60 }),
547 serde_json::json!({ "request_uri": "urn:x" }),
548 serde_json::json!({ "request_uri": 42, "expires_in": 60 }),
549 serde_json::json!({ "request_uri": "urn:x", "expires_in": 0 }),
550 serde_json::json!({ "request_uri": "urn:x", "expires_in": -1 }),
551 serde_json::json!({ "request_uri": "", "expires_in": 60 }),
552 ] {
553 assert!(parse_par_response(&body).is_err(), "accepted {body}");
554 }
555 }
556
557 #[test]
560 fn the_authorize_url_carries_only_client_id_and_request_uri() {
561 let url = authorize_url(
562 "https://auth.example.com/authorize",
563 "https://feather-reader.com/oauth/client-metadata.json",
564 "urn:ietf:params:oauth:request_uri:abc",
565 )
566 .unwrap();
567 let parsed = url::Url::parse(&url).unwrap();
568 let names: Vec<String> = parsed.query_pairs().map(|(k, _)| k.into_owned()).collect();
569 assert_eq!(names, vec!["client_id", "request_uri"]);
570 }
571
572 #[test]
575 fn the_authorize_url_percent_encodes_its_parameters() {
576 let url = authorize_url(
577 "https://auth.example.com/authorize",
578 "https://x.example/m.json?a=1&b=2",
579 "urn:oauth:uri:with spaces",
580 )
581 .unwrap();
582 assert!(
583 url.contains("%3A%2F%2F"),
584 "scheme separator unencoded: {url}"
585 );
586 assert!(!url.contains("with spaces"));
587 let parsed = url::Url::parse(&url).unwrap();
588 assert_eq!(
589 parsed.query_pairs().count(),
590 2,
591 "extra parameters spliced in"
592 );
593 }
594
595 #[test]
596 fn a_non_https_authorization_endpoint_is_refused() {
597 assert!(authorize_url("http://auth.example.com/authorize", "cid", "uri").is_err());
598 assert!(authorize_url("not a url", "cid", "uri").is_err());
599 }
600
601 fn ok_params() -> CallbackParams {
604 CallbackParams {
605 code: Some("the-code".into()),
606 state: Some("state-value".into()),
607 iss: Some(ISSUER.into()),
608 error: None,
609 error_description: None,
610 response: None,
611 }
612 }
613
614 #[test]
615 fn a_well_formed_callback_is_accepted() {
616 assert_eq!(verify_callback(&ok_params(), ISSUER).unwrap(), "the-code");
617 }
618
619 #[test]
625 fn a_callback_without_iss_is_rejected() {
626 let mut params = ok_params();
627 params.iss = None;
628 assert!(verify_callback(¶ms, ISSUER).is_err());
629 }
630
631 #[test]
632 fn a_callback_from_the_wrong_issuer_is_rejected() {
633 let mut params = ok_params();
634 params.iss = Some("https://evil.example.com".into());
635 assert!(verify_callback(¶ms, ISSUER).is_err());
636 }
637
638 #[test]
641 fn an_error_response_is_surfaced_as_such() {
642 let mut params = ok_params();
643 params.code = None;
644 params.error = Some("access_denied".into());
645 let err = verify_callback(¶ms, ISSUER).unwrap_err();
646 assert!(format!("{err:#}").contains("access_denied"));
647 }
648
649 #[test]
655 fn an_error_from_the_wrong_issuer_is_reported_as_a_mismatch() {
656 let mut params = ok_params();
657 params.code = None;
658 params.error = Some("access_denied".into());
659 params.iss = Some("https://evil.example.com".into());
660 let rendered = format!("{:#}", verify_callback(¶ms, ISSUER).unwrap_err());
661 assert!(
662 rendered.contains("iss"),
663 "reported the error before checking who sent it: {rendered}"
664 );
665 assert!(!rendered.contains("access_denied"));
666 }
667
668 #[test]
673 fn server_supplied_error_text_is_not_passed_through() {
674 let mut params = ok_params();
675 params.code = None;
676 params.error = Some("access_denied".into());
677 params.error_description = Some("<img src=x onerror=alert(1)>".into());
678 let rendered = format!("{:#}", verify_callback(¶ms, ISSUER).unwrap_err());
679 assert!(
680 !rendered.contains("<img"),
681 "raw description leaked: {rendered}"
682 );
683 assert!(!rendered.contains("onerror"));
684 }
685
686 #[test]
689 fn an_unknown_error_code_is_reduced_to_a_slug() {
690 let mut params = ok_params();
691 params.code = None;
692 params.error = Some("<script>alert(1)</script>".into());
693 let rendered = format!("{:#}", verify_callback(¶ms, ISSUER).unwrap_err());
694 assert!(
695 !rendered.contains("<script>"),
696 "raw code leaked: {rendered}"
697 );
698 }
699
700 #[test]
703 fn an_error_takes_precedence_over_a_code() {
704 let mut params = ok_params();
705 params.error = Some("invalid_request".into());
706 assert!(verify_callback(¶ms, ISSUER).is_err());
707 }
708
709 #[test]
712 fn a_jarm_response_is_refused() {
713 let mut params = ok_params();
714 params.response = Some("signed-jwt".into());
715 assert!(verify_callback(¶ms, ISSUER).is_err());
716 }
717
718 #[test]
719 fn a_callback_without_a_code_is_rejected() {
720 let mut params = ok_params();
721 params.code = None;
722 assert!(verify_callback(¶ms, ISSUER).is_err());
723 params.code = Some(String::new());
724 assert!(verify_callback(¶ms, ISSUER).is_err());
725 }
726
727 #[test]
728 fn a_callback_without_a_state_is_rejected() {
729 let mut params = ok_params();
730 params.state = None;
731 assert!(verify_callback(¶ms, ISSUER).is_err());
732 }
733
734 async fn pending_db(cookie_hash: &str) -> (sqlx::SqlitePool, crate::oauth::crypto::Codec) {
737 let pool = crate::store::init_url("sqlite::memory:").await.unwrap();
738 crate::oauth::store::init_schema(&pool).await.unwrap();
739 let codec = crate::oauth::crypto::Codec::new(Some("a".repeat(43).as_str())).unwrap();
740 let pending = crate::oauth::store::PendingAuth {
741 state: "state-value".into(),
742 browser_binding_hash: cookie_hash.into(),
743 pkce_verifier: "verifier".into(),
744 dpop_key_jwk: "{}".into(),
745 issuer: ISSUER.into(),
746 pds_url: "https://pds.example.com".into(),
747 did: "did:plc:ewvi7nxzyoun6zhxrhs64oiz".into(),
748 auth_method: "private_key_jwt".into(),
749 auth_kid: None,
750 redirect_uri: "https://feather-reader.com/oauth/callback".into(),
751 requested_scope: "atproto".into(),
752 request_uri: "urn:x".into(),
753 app_return_to: None,
754 expires_at: 2_000_000_000,
755 };
756 crate::oauth::store::put_pending(&pool, &codec, &pending)
757 .await
758 .unwrap();
759 (pool, codec)
760 }
761
762 #[tokio::test]
763 async fn complete_callback_returns_the_code_for_the_right_browser() {
764 let cookie = new_binding_token();
765 let (pool, codec) = pending_db(&binding_hash(&cookie)).await;
766 let (pending, code) =
767 complete_callback(&pool, &codec, &ok_params(), Some(&cookie), 1_700_000_000)
768 .await
769 .unwrap();
770 assert_eq!(code, "the-code");
771 assert_eq!(pending.pkce_verifier, "verifier");
772 }
773
774 #[tokio::test]
778 async fn complete_callback_refuses_a_browser_that_did_not_start_the_flow() {
779 let cookie = new_binding_token();
780 for presented in [None, Some(new_binding_token())] {
781 let (pool, codec) = pending_db(&binding_hash(&cookie)).await;
782 assert!(complete_callback(
783 &pool,
784 &codec,
785 &ok_params(),
786 presented.as_deref(),
787 1_700_000_000
788 )
789 .await
790 .is_err());
791
792 assert!(
794 complete_callback(&pool, &codec, &ok_params(), Some(&cookie), 1_700_000_000)
795 .await
796 .is_err(),
797 "the pending row survived a failed binding check"
798 );
799 }
800 }
801
802 #[tokio::test]
811 async fn complete_callback_validates_iss_against_the_stored_issuer() {
812 let cookie = new_binding_token();
813 let (pool, codec) = pending_db(&binding_hash(&cookie)).await;
814
815 let impostor = CallbackParams {
817 iss: Some("https://evil.example".into()),
818 ..ok_params()
819 };
820 let err = complete_callback(&pool, &codec, &impostor, Some(&cookie), 1_700_000_000)
821 .await
822 .expect_err("an `iss` from another authorization server must be refused");
823 let rendered = format!("{err:#}");
824 assert!(
825 rendered.contains("evil.example") || rendered.contains("iss"),
826 "failed for the wrong reason: {rendered}"
827 );
828 }
829
830 #[tokio::test]
831 async fn complete_callback_rejects_an_unknown_or_expired_state() {
832 let cookie = new_binding_token();
833 let (pool, codec) = pending_db(&binding_hash(&cookie)).await;
834
835 let mut unknown = ok_params();
836 unknown.state = Some("never-existed".into());
837 assert!(
838 complete_callback(&pool, &codec, &unknown, Some(&cookie), 1_700_000_000)
839 .await
840 .is_err()
841 );
842
843 assert!(
845 complete_callback(&pool, &codec, &ok_params(), Some(&cookie), 2_000_000_001)
846 .await
847 .is_err()
848 );
849 }
850
851 #[tokio::test]
854 async fn complete_callback_checks_the_binding_before_anything_else_about_the_response() {
855 let cookie = new_binding_token();
856 let (pool, codec) = pending_db(&binding_hash(&cookie)).await;
857
858 let mut denied = ok_params();
859 denied.code = None;
860 denied.error = Some("access_denied".into());
861 let rendered = format!(
862 "{:#}",
863 complete_callback(&pool, &codec, &denied, None, 1_700_000_000)
864 .await
865 .unwrap_err()
866 );
867 assert!(
868 rendered.contains("browser-binding"),
869 "reported the response before checking the browser: {rendered}"
870 );
871 }
872}