diff --git a/crates/tranquil-api/src/lib.rs b/crates/tranquil-api/src/lib.rs index 6852e73..9419f2f 100644 --- a/crates/tranquil-api/src/lib.rs +++ b/crates/tranquil-api/src/lib.rs @@ -489,9 +489,15 @@ pub fn webhook_routes() -> axum::Router { pub fn misc_routes() -> axum::Router { use axum::routing::get; - axum::Router::new() + let router = axum::Router::new() .route("/health", get(server::health)) .route("/robots.txt", get(server::robots_txt)) .route("/favicon.ico", get(server::get_logo)) - .route("/u/{handle}/did.json", get(identity::user_did_doc)) + .route("/u/{handle}/did.json", get(identity::user_did_doc)); + + if tranquil_config::get().server.rfc_moo_compliance { + router.route("/cow.txt", get(server::cow_txt)) + } else { + router + } } diff --git a/crates/tranquil-api/src/server/cow.txt b/crates/tranquil-api/src/server/cow.txt new file mode 100644 index 0000000..add2f00 --- /dev/null +++ b/crates/tranquil-api/src/server/cow.txt @@ -0,0 +1,57 @@ + + + + + + + .......................... + ....*o|||||||8#@@@@@@@@@@@@@@@@@@@@@@@###&|o:_.. + ..*:o|||&8##@###8888888######@#@###########################|*... + .:o|||8#####8888|:::**. *&########################@@################&o_ + .*o&8###@#8&o*_. :###@##############@########################@@##&o_ + .*o8########& :##@#@##############@############################@###|_ + .*o|8##########8o .#######################################################&o_ + *&##|_ ..*&##8&o*|88888|_ _#######################################@##################|. + *#####& *&######&o_..*o|o:_ .&##o _###########################################################&_ + _##8*##8 .|88|:::|#######8###8|*:_ .&#@@8 _##@@@########################################################&_ + _#@8_##8_ *8#8|*_ _:|#####&&####8 .&##############################################################| + _#@8.|##8_ _::o###8&##8 .|##@############################8###########################@@#|_ + *###o.|88o ..*&####|..##& _|##########################8|_ .|#############################8 + *|###|_ ._&####8|*_ _*_ _::&8888888888888888|::*_ .|##@####@@@##################| + *&###|_ _:_ .&88###8|*_ ..... .|#####@@@##################8 + .##@#& _##& .|##o _#@@#@#| .|#######&:_ _|###@####################8 + .:8##8o _o:*&##| *##8_.&@@##@#| _::o8#8|::|#####|_ _|#################88###8 + .&##&*_ *###o_###| .|##8*&##|*###o _###8####8|_ _:|###|_ .*o|||o:_ _:::&8888888888|_ _##8 + .###|_. _###o *###|*&#######8 *##8 .##8_ _:|###|_ _|###|_ .&########o _##& + _|####8|&##8:_ _|#########88o .##8 *##& _|###o .|###o .#########| .o##o + o#8|*:#@@###o _:::*__*_ _##8_ _##8_ _&##|_ *##8_ *8#####8|_ .*oo:_ o##| + *###o.&#####& _oo* .8##& .8##8_ .|##& o##& _::::_.*o|8######|_ .##8. + _###&o&##8_:*_ .###& .###|_&#####| _##8 :###o *ooo&#########@#@#& ....:##& + .|8||###&. _**_ .###88##|*&###|*._&##& *|##8_ *o&####@@####@@######& .*o||||||&#######8_ + *&###o _|88##8_ _:8######|*:###|_ _##################88|_ *&#################& + *#####o *&8o *##& _:::*_.&##|_ _#@##############8_ :##################8* + .###&##8_.|88o *&8o _@@& .###| .&####@#########|_ .####@###@@########8* + _##&.|###|_.... .|88o _##8* *&###|_ *###o _:&88######8|_ .*o|||##################o + _##8_ _|########|_ .*o8####&#@@#@##o *#@8 _*:*. .&###@###################| + .|##& _:::::&##& .*&##############@#8_ .###o _#@####################|_ + .&##|_ .&##8_ *o&####################8_ *##& .&#####@@############8* + .|##8_.&###&####8_ _########################8**##& _##################|_ + .&###&##888888##8_ .|88######@###########|*######o _|8###############|_ + .&#####o *###|_ _::::::*o##8**o##8 .|###8o .&##@#############&* + .|####o _|###|_ _##8.*&##& _*_ .#################| + _*_ _|###|_.. .|#####8|_ *&#@#########8###& + *#######&|o:_... ..*:::*. ......._:o&8####888&o:#####8_ + _###|&888#####@#####&|||o:_........................._:o||||8##@@@@####8|:*_ _:::*_ + .|#@###o _:::o#@#888######@@@@@@@@@@@@@@@@@@@@@@@@#####888|::::::**_ + _::*_ :##& *&8|_:::::::::::::::::::::::::**_ + .###8||&##8o + _|888888|_ + + + + + + + + + diff --git a/crates/tranquil-api/src/server/meta.rs b/crates/tranquil-api/src/server/meta.rs index bc7f2ad..0f07c4a 100644 --- a/crates/tranquil-api/src/server/meta.rs +++ b/crates/tranquil-api/src/server/meta.rs @@ -29,6 +29,10 @@ pub async fn robots_txt() -> impl IntoResponse { "# Hello!\n\n# Crawling the public API is allowed\nUser-agent: *\nAllow: /\n", ) } + +pub async fn cow_txt() -> &'static str { + include_str!("cow.txt") +} pub fn is_self_hosted_did_web_enabled() -> bool { tranquil_config::get().server.enable_pds_hosted_did_web } diff --git a/crates/tranquil-api/src/server/mod.rs b/crates/tranquil-api/src/server/mod.rs index 4c01aae..0bf8b55 100644 --- a/crates/tranquil-api/src/server/mod.rs +++ b/crates/tranquil-api/src/server/mod.rs @@ -28,7 +28,7 @@ pub use email::{ }; pub use invite::{create_invite_code, create_invite_codes, get_account_invite_codes}; pub use logo::get_logo; -pub use meta::{describe_server, health, robots_txt}; +pub use meta::{cow_txt, describe_server, health, robots_txt}; pub use migration::{get_did_document, update_did_document}; pub use passkey_account::{ complete_passkey_setup, create_passkey_account, recover_passkey_account, diff --git a/crates/tranquil-config/src/lib.rs b/crates/tranquil-config/src/lib.rs index e1a5ce3..3a0185e 100644 --- a/crates/tranquil-config/src/lib.rs +++ b/crates/tranquil-config/src/lib.rs @@ -445,6 +445,10 @@ pub struct ServerConfig { #[config(env = "ENABLE_PDS_HOSTED_DID_WEB", default = false)] pub enable_pds_hosted_did_web: bool, + /// iykyk! + #[config(env = "RFC_MOO_COMPLIANCE", default = false)] + pub rfc_moo_compliance: bool, + /// When set to true, skip age-assurance birthday prompt for all accounts. #[config(env = "PDS_AGE_ASSURANCE_OVERRIDE", default = false)] pub age_assurance_override: bool, diff --git a/example.toml b/example.toml index aeceade..22b97b0 100644 --- a/example.toml +++ b/example.toml @@ -34,6 +34,13 @@ # Default value: false #enable_pds_hosted_did_web = false +# iykyk! +# +# Can also be specified via environment variable `RFC_MOO_COMPLIANCE`. +# +# Default value: false +#rfc_moo_compliance = false + # When set to true, skip age-assurance birthday prompt for all accounts. # # Can also be specified via environment variable `PDS_AGE_ASSURANCE_OVERRIDE`. @@ -416,7 +423,9 @@ # List of relay / crawler notification URLs. # # Can also be specified via environment variable `CRAWLERS`. -#crawlers = +# +# Default value: ["https://relay.fire.hose.cam", "https://relay3.fr.hose.cam", "https://bsky.network", "https://northamerica.firehose.network", "https://europe.firehose.network", "https://asia.firehose.network", "https://atproto.africa", "https://relay.upcloud.world"] +#crawlers = ["https://relay.fire.hose.cam", "https://relay3.fr.hose.cam", "https://bsky.network", "https://northamerica.firehose.network", "https://europe.firehose.network", "https://asia.firehose.network", "https://atproto.africa", "https://relay.upcloud.world"] [email] # Sender email address. When unset, email sending is disabled. -- 2.51.2 From 0b8787d1dea67d8eb41951dc6318c2398233dd56 Mon Sep 17 00:00:00 2001 From: Lewis Date: Sun, 26 Jul 2026 13:57:07 +0300 Subject: [PATCH 02/22] pds: compile bsky-specific proxy, CORS, & validation out under bsky features Lewis: May this revision serve well! --- crates/tranquil-api/src/identity/provision.rs | 7 ++++ crates/tranquil-pds/src/api/proxy.rs | 42 ++++++++++--------- crates/tranquil-pds/src/lib.rs | 27 ++++++------ crates/tranquil-pds/src/types.rs | 1 + crates/tranquil-pds/src/util.rs | 15 ++++--- crates/tranquil-pds/src/validation/mod.rs | 5 +++ 6 files changed, 59 insertions(+), 38 deletions(-) diff --git a/crates/tranquil-api/src/identity/provision.rs b/crates/tranquil-api/src/identity/provision.rs index 207d5c7..5d4d387 100644 --- a/crates/tranquil-api/src/identity/provision.rs +++ b/crates/tranquil-api/src/identity/provision.rs @@ -164,6 +164,13 @@ pub async fn resolve_signing_key( } } +#[cfg_attr( + not(feature = "bsky"), + expect( + unused_variables, + reason = "only the bsky block writes display_name into the default profile record" + ) +)] pub async fn sequence_new_account( state: &AppState, did: &Did, diff --git a/crates/tranquil-pds/src/api/proxy.rs b/crates/tranquil-pds/src/api/proxy.rs index 5e2af72..c51a2e3 100644 --- a/crates/tranquil-pds/src/api/proxy.rs +++ b/crates/tranquil-pds/src/api/proxy.rs @@ -335,27 +335,29 @@ async fn proxy_handler( }; // BSKY: getFeed must be audienced to the feed generator, not the AppView. - let (token_aud, token_lxm) = - if cfg!(feature = "bsky-support") && method == "app.bsky.feed.getFeed" { - match resolve_feed_generator_did(&resolved.url, query.as_deref()).await { - Some(feed_did) => ( - feed_did, - "app.bsky.feed.getFeedSkeleton" - .parse::() - .expect("getFeedSkeleton is a valid NSID"), - ), - None => { - warn!( - "getFeed proxy: could not resolve feed generator DID; refusing \ - to mint an AppView-audienced token" - ); - return ApiError::InvalidRequest("Could not resolve feed".into()) - .into_response(); - } + #[cfg(feature = "bsky-support")] + let (token_aud, token_lxm) = if method == "app.bsky.feed.getFeed" { + match resolve_feed_generator_did(&resolved.url, query.as_deref()).await { + Some(feed_did) => ( + feed_did, + "app.bsky.feed.getFeedSkeleton" + .parse::() + .expect("getFeedSkeleton is a valid NSID"), + ), + None => { + warn!( + "getFeed proxy refuses to mint an AppView-audienced token \ + because feed generator DID resolution failed" + ); + return ApiError::InvalidRequest("Couldn't resolve feed".into()) + .into_response(); } - } else { - (resolved.did.clone(), method_nsid.clone()) - }; + } + } else { + (resolved.did.clone(), method_nsid.clone()) + }; + #[cfg(not(feature = "bsky-support"))] + let (token_aud, token_lxm) = (resolved.did.clone(), method_nsid.clone()); match crate::auth::create_service_token( &auth_user.did, diff --git a/crates/tranquil-pds/src/lib.rs b/crates/tranquil-pds/src/lib.rs index 8bf5e31..a3b2bd7 100644 --- a/crates/tranquil-pds/src/lib.rs +++ b/crates/tranquil-pds/src/lib.rs @@ -35,7 +35,7 @@ use serde_json::json; use state::AppState; use tower::ServiceBuilder; use tower_http::{ - cors::{Any, CorsLayer}, + cors::{AllowHeaders, Any, CorsLayer}, services::{ServeDir, ServeFile}, }; pub use tranquil_db_traits::AccountStatus; @@ -106,17 +106,20 @@ pub fn app_with_routes(state: AppState, external: ExternalRoutes) -> Router { CorsLayer::new() .allow_origin(Any) .allow_methods([Method::GET, Method::POST, Method::OPTIONS]) - .allow_headers([ - http::header::AUTHORIZATION, - http::header::CONTENT_TYPE, - http::header::CONTENT_ENCODING, - http::header::ACCEPT_ENCODING, - http::header::USER_AGENT, - util::HEADER_DPOP, - util::HEADER_ATPROTO_PROXY, - util::HEADER_ATPROTO_ACCEPT_LABELERS, - util::HEADER_X_BSKY_TOPICS, - ]) + .allow_headers(AllowHeaders::list( + [ + http::header::AUTHORIZATION, + http::header::CONTENT_TYPE, + http::header::CONTENT_ENCODING, + http::header::ACCEPT_ENCODING, + http::header::USER_AGENT, + util::HEADER_DPOP, + util::HEADER_ATPROTO_PROXY, + util::HEADER_ATPROTO_ACCEPT_LABELERS, + ] + .into_iter() + .chain(util::CORS_BSKY_ALLOW_HEADERS), + )) .expose_headers([ http::header::WWW_AUTHENTICATE, util::HEADER_DPOP_NONCE, diff --git a/crates/tranquil-pds/src/types.rs b/crates/tranquil-pds/src/types.rs index 9a184f2..5d1d0ed 100644 --- a/crates/tranquil-pds/src/types.rs +++ b/crates/tranquil-pds/src/types.rs @@ -1,5 +1,6 @@ pub use tranquil_types::*; +#[cfg(feature = "bsky")] use std::sync::LazyLock; #[cfg(feature = "bsky")] diff --git a/crates/tranquil-pds/src/util.rs b/crates/tranquil-pds/src/util.rs index 2b1a977..4492f30 100644 --- a/crates/tranquil-pds/src/util.rs +++ b/crates/tranquil-pds/src/util.rs @@ -89,6 +89,10 @@ pub const HEADER_ATPROTO_CONTENT_LABELERS: HeaderName = HeaderName::from_static("atproto-content-labelers"); #[cfg(feature = "bsky-support")] pub const HEADER_X_BSKY_TOPICS: HeaderName = HeaderName::from_static("x-bsky-topics"); +#[cfg(feature = "bsky-support")] +pub const CORS_BSKY_ALLOW_HEADERS: [HeaderName; 1] = [HEADER_X_BSKY_TOPICS]; +#[cfg(not(feature = "bsky-support"))] +pub const CORS_BSKY_ALLOW_HEADERS: [HeaderName; 0] = []; pub fn get_header_str( headers: &HeaderMap, @@ -250,11 +254,7 @@ pub fn build_full_url(path: &str) -> String { && (path.starts_with("/com.atproto.") // BSKY: Bluesky requires that the PDS implement some app.bsky.* endpoints so we need to deal with those here too. // TODO: surely we can figure out a way to do this more generically? - || (if cfg!(feature = "bsky-support") { - path.starts_with("/app.bsky.") - } else { - true - }) + || (cfg!(feature = "bsky-support") && path.starts_with("/app.bsky.")) || path.starts_with("/_")) { format!("/xrpc{path}") @@ -798,7 +798,10 @@ mod tests { ); assert_eq!( build_full_url("/app.bsky.feed.getTimeline"), - "https://example.com/xrpc/app.bsky.feed.getTimeline" + match cfg!(feature = "bsky-support") { + true => "https://example.com/xrpc/app.bsky.feed.getTimeline", + false => "https://example.com/app.bsky.feed.getTimeline", + } ); assert_eq!( build_full_url("/_health"), diff --git a/crates/tranquil-pds/src/validation/mod.rs b/crates/tranquil-pds/src/validation/mod.rs index b94ab83..fca4766 100644 --- a/crates/tranquil-pds/src/validation/mod.rs +++ b/crates/tranquil-pds/src/validation/mod.rs @@ -132,6 +132,10 @@ fn validate_preamble<'a>( Ok((record_type, obj)) } +#[cfg_attr( + not(feature = "bsky"), + expect(unused_variables, reason = "only bsky record checks read obj and rkey") +)] fn check_banned_content( record_type: &str, obj: &serde_json::Map, @@ -211,6 +215,7 @@ fn check_post_banned_content(obj: &serde_json::Map) -> Result<(), Ok(()) } +#[cfg(feature = "bsky")] fn check_string_field( obj: &serde_json::Map, field: &str, -- 2.51.2 From 135912194d52c9874f51d99915cce07453b4a238 Mon Sep 17 00:00:00 2001 From: Lewis Date: Sun, 26 Jul 2026 13:57:07 +0300 Subject: [PATCH 03/22] types: HttpUrl newtypes, shared cache key/JSON helpers Lewis: May this revision serve well! --- Cargo.lock | 12 +- Cargo.toml | 1 + crates/tranquil-api/src/delegation.rs | 23 +- crates/tranquil-cache/Cargo.toml | 2 +- crates/tranquil-cache/src/lib.rs | 7 +- crates/tranquil-infra/Cargo.toml | 8 + crates/tranquil-infra/src/cache_keys.rs | 103 +++ crates/tranquil-infra/src/lib.rs | 45 ++ crates/tranquil-infra/src/memory_cache.rs | 74 ++ .../src/endpoints/delegation.rs | 19 +- crates/tranquil-oauth/Cargo.toml | 1 + crates/tranquil-oauth/src/types.rs | 10 +- crates/tranquil-pds/Cargo.toml | 3 +- crates/tranquil-pds/src/cache/mod.rs | 4 +- crates/tranquil-pds/src/cache_keys.rs | 49 +- crates/tranquil-pds/src/delegation/mod.rs | 39 +- crates/tranquil-pds/src/oauth/client.rs | 135 ++-- crates/tranquil-signal/Cargo.toml | 2 +- crates/tranquil-types/Cargo.toml | 7 + crates/tranquil-types/src/lib.rs | 634 +++++++++++++++++- 20 files changed, 1011 insertions(+), 167 deletions(-) create mode 100644 crates/tranquil-infra/src/cache_keys.rs create mode 100644 crates/tranquil-infra/src/memory_cache.rs diff --git a/Cargo.lock b/Cargo.lock index 4679a5c..d65d177 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -7838,7 +7838,10 @@ dependencies = [ "async-trait", "bytes", "futures", + "serde", + "serde_json", "thiserror 2.0.18", + "tranquil-types", ] [[package]] @@ -7855,9 +7858,9 @@ dependencies = [ "thiserror 2.0.18", "tokio", "tracing", + "tranquil-infra", "tranquil-types", "unicode-segmentation", - "urlencoding", "wiremock", ] @@ -7880,6 +7883,7 @@ dependencies = [ "sqlx", "tokio", "tracing", + "tranquil-infra", "tranquil-types", "uuid", ] @@ -7911,6 +7915,7 @@ dependencies = [ "tranquil-config", "tranquil-crypto", "tranquil-db-traits", + "tranquil-infra", "tranquil-pds", "tranquil-scopes", "tranquil-types", @@ -7991,6 +7996,7 @@ dependencies = [ "tranquil-config", "tranquil-db", "tranquil-db-traits", + "tranquil-infra", "tranquil-lexicon", "tranquil-oauth", "tranquil-oauth-server", @@ -8221,10 +8227,14 @@ dependencies = [ "cid", "jacquard-common", "rand 0.8.5", + "reqwest", "serde", "serde_json", "sqlx", "thiserror 2.0.18", + "tokio", + "tracing", + "url", "uuid", ] diff --git a/Cargo.toml b/Cargo.toml index 85c38c5..cdbf10a 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -137,6 +137,7 @@ tower-layer = "0.3" tracing = "0.1" tracing-subscriber = "0.3" urlencoding = "2.1" +url = "2.5" uuid = { version = "1.19", features = ["v4", "v5", "v7", "fast-rng", "serde"] } webauthn-rs = { version = "0.5", features = ["danger-allow-state-serialisation", "danger-user-presence-only-security-keys", "conditional-ui"] } webauthn-rs-proto = "0.5" diff --git a/crates/tranquil-api/src/delegation.rs b/crates/tranquil-api/src/delegation.rs index 2e52499..7262bf4 100644 --- a/crates/tranquil-api/src/delegation.rs +++ b/crates/tranquil-api/src/delegation.rs @@ -12,8 +12,8 @@ use tranquil_pds::api::{ }; use tranquil_pds::auth::{Active, Auth}; use tranquil_pds::delegation::{ - DelegationActionType, SCOPE_PRESETS, ValidatedDelegationScope, verify_can_add_controllers, - verify_can_control_accounts, + DelegationActionType, IdentityResolutionError, SCOPE_PRESETS, ValidatedDelegationScope, + verify_can_add_controllers, verify_can_control_accounts, }; use tranquil_pds::rate_limit::{AccountCreationLimit, RateLimited}; use tranquil_pds::state::AppState; @@ -65,16 +65,16 @@ pub async fn add_controller( ) -> Result, ApiError> { let resolved = tranquil_pds::delegation::resolve_identity(&state, &input.controller_did) .await - .map_err(|_| ApiError::ControllerNotFound)?; + .map_err(|e| match e { + IdentityResolutionError::PdsEndpoint(_) => ApiError::InvalidDelegation( + "Controller PDS endpoint isn't a usable https URL".into(), + ), + IdentityResolutionError::DidResolution(_) => ApiError::ControllerNotFound, + })?; if !resolved.is_local && let Some(ref pds_url) = resolved.pds_url { - if !pds_url.starts_with("https://") { - return Err(ApiError::InvalidDelegation( - "Controller PDS must use HTTPS".into(), - )); - } match state .cross_pds_oauth .check_remote_is_delegated(pds_url, &input.controller_did) @@ -477,7 +477,12 @@ pub async fn resolve_controller( let resolved = tranquil_pds::delegation::resolve_identity(&state, &did) .await - .map_err(|_| ApiError::ControllerNotFound)?; + .map_err(|e| match e { + IdentityResolutionError::PdsEndpoint(_) => ApiError::InvalidDelegation( + "Controller PDS endpoint isn't a usable https URL".into(), + ), + IdentityResolutionError::DidResolution(_) => ApiError::ControllerNotFound, + })?; Ok(Json(resolved)) } diff --git a/crates/tranquil-cache/Cargo.toml b/crates/tranquil-cache/Cargo.toml index de971ff..bbc805e 100644 --- a/crates/tranquil-cache/Cargo.toml +++ b/crates/tranquil-cache/Cargo.toml @@ -9,7 +9,7 @@ valkey = ["dep:redis"] [dependencies] tranquil-config = { workspace = true } -tranquil-infra = { workspace = true } +tranquil-infra = { workspace = true, features = ["cache-keys"] } tranquil-ripple = { workspace = true } async-trait = { workspace = true } diff --git a/crates/tranquil-cache/src/lib.rs b/crates/tranquil-cache/src/lib.rs index 2551cf8..67804cb 100644 --- a/crates/tranquil-cache/src/lib.rs +++ b/crates/tranquil-cache/src/lib.rs @@ -1,4 +1,6 @@ -pub use tranquil_infra::{Cache, CacheError, DistributedRateLimiter}; +pub use tranquil_infra::{ + Cache, CacheError, DistributedRateLimiter, cache_keys, cached_json, read_json, write_json, +}; use async_trait::async_trait; use std::sync::Arc; @@ -173,11 +175,10 @@ pub async fn create_cache( ) -> Result<(Arc, Arc), CacheInitError> { let cache_cfg = tranquil_config::try_get().map(|c| &c.cache); let backend = cache_cfg.map(|c| c.backend.as_str()).unwrap_or("ripple"); - let valkey_url = cache_cfg.and_then(|c| c.valkey_url.as_deref()); #[cfg(feature = "valkey")] if backend == "valkey" { - if let Some(url) = valkey_url { + if let Some(url) = cache_cfg.and_then(|c| c.valkey_url.as_deref()) { match ValkeyCache::new(url).await { Ok(cache) => { tracing::info!("using valkey cache at {url}"); diff --git a/crates/tranquil-infra/Cargo.toml b/crates/tranquil-infra/Cargo.toml index 4c7436e..acef127 100644 --- a/crates/tranquil-infra/Cargo.toml +++ b/crates/tranquil-infra/Cargo.toml @@ -4,8 +4,16 @@ version.workspace = true edition.workspace = true license.workspace = true +[features] +testing = [] +cache-keys = ["dep:tranquil-types"] + [dependencies] +tranquil-types = { workspace = true, optional = true } + async-trait = { workspace = true } bytes = { workspace = true } futures = { workspace = true } +serde = { workspace = true } +serde_json = { workspace = true } thiserror = { workspace = true } diff --git a/crates/tranquil-infra/src/cache_keys.rs b/crates/tranquil-infra/src/cache_keys.rs new file mode 100644 index 0000000..c01ac62 --- /dev/null +++ b/crates/tranquil-infra/src/cache_keys.rs @@ -0,0 +1,103 @@ +use tranquil_types::{ + CidLink, ClientId, CrossPdsState, Did, EmailTokenPurpose, Handle, Jti, JwksUri, Nsid, PdsUrl, + SsoIssuer, SsoJwksUri, +}; + +pub fn session_key(did: &Did, jti: &Jti) -> String { + format!("auth:session:{}:{}", did, jti) +} + +pub fn signing_key_key(did: &Did) -> String { + format!("auth:key:{}", did) +} + +pub fn user_status_key(did: &Did) -> String { + format!("auth:status:{}", did) +} + +pub fn handle_key(handle: &Handle) -> String { + format!("handle:{}", handle) +} + +pub fn reauth_key(did: &Did) -> String { + format!("reauth:{}", did) +} + +pub fn plc_doc_key(did: &Did) -> String { + format!("plc:doc:{}", did) +} + +pub fn plc_data_key(did: &Did) -> String { + format!("plc:data:{}", did) +} + +pub fn did_web_doc_key(did: &Did) -> String { + format!("did:web:doc:{}", did) +} + +pub fn email_update_key(did: &Did) -> String { + format!("email_update:{}", did) +} + +pub fn email_token_key(did: &Did, purpose: EmailTokenPurpose) -> String { + format!("email_token:{}:{}", purpose, did) +} + +pub fn legacy_2fa_challenge_key(did: &Did) -> String { + format!("legacy_2fa:{}", did) +} + +pub fn legacy_2fa_cooldown_key(did: &Did) -> String { + format!("legacy_2fa_cooldown:{}", did) +} + +pub fn scope_ref_key(cid: &CidLink) -> String { + format!("scope_ref:{}", cid) +} + +pub fn auto_verify_sent_key(did: &Did) -> String { + format!("auto_verify_sent:{}", did) +} + +pub fn permission_set_key(nsid: &Nsid, aud: Option<&str>) -> String { + match aud { + Some(a) => format!("permset:{}:{}", nsid, a), + None => format!("permset:{}", nsid), + } +} + +pub fn oauth_client_meta_key(client_id: &ClientId) -> String { + format!("oauth:client_meta:{}", client_id) +} + +pub fn oauth_client_jwks_key(jwks_uri: &JwksUri) -> String { + format!("oauth:jwks:{}", jwks_uri.canonical()) +} + +pub fn oauth_client_jwks_cooldown_key(jwks_uri: &JwksUri) -> String { + format!("oauth:jwks_cooldown:{}", jwks_uri.canonical()) +} + +pub fn sso_jwks_key(jwks_uri: &SsoJwksUri) -> String { + format!("sso:jwks:{}", jwks_uri.canonical()) +} + +pub fn oidc_discovery_key(issuer: &SsoIssuer) -> String { + format!("oidc:discovery:{}", issuer.canonical()) +} + +pub fn cross_pds_state_key(state: &CrossPdsState) -> String { + format!("cross_pds_state:{}", state) +} + +pub fn cross_pds_oauth_meta_key(pds_url: &PdsUrl) -> String { + format!("cross_pds_oauth_meta:v2:{}", pds_url.canonical()) +} + +pub fn lexicon_doc_key(nsid: &Nsid) -> String { + format!("lexicon:doc:{}", nsid) +} + +pub fn lexicon_negative_key(nsid: &Nsid) -> String { + format!("lexicon:neg:{}", nsid) +} diff --git a/crates/tranquil-infra/src/lib.rs b/crates/tranquil-infra/src/lib.rs index 9928884..76a0783 100644 --- a/crates/tranquil-infra/src/lib.rs +++ b/crates/tranquil-infra/src/lib.rs @@ -1,6 +1,15 @@ +#[cfg(feature = "cache-keys")] +pub mod cache_keys; + +#[cfg(feature = "testing")] +mod memory_cache; +#[cfg(feature = "testing")] +pub use memory_cache::MemoryCache; + use async_trait::async_trait; use bytes::Bytes; use futures::Stream; +use std::future::Future; use std::pin::Pin; use std::time::Duration; @@ -57,6 +66,42 @@ pub trait Cache: Send + Sync { } } +pub async fn read_json(cache: &dyn Cache, key: &str) -> Option { + let json = cache.get(key).await?; + serde_json::from_str(&json).ok() +} + +pub async fn write_json( + cache: &dyn Cache, + key: &str, + value: &T, + ttl: Duration, +) { + if let Ok(json) = serde_json::to_string(value) { + let _ = cache.set(key, &json, ttl).await; + } +} + +pub async fn cached_json( + cache: &dyn Cache, + key: &str, + ttl: Duration, + fetch: impl FnOnce() -> Fut, +) -> Result +where + T: serde::Serialize + serde::de::DeserializeOwned, + Fut: Future>, +{ + match read_json(cache, key).await { + Some(value) => Ok(value), + None => { + let value = fetch().await?; + write_json(cache, key, &value, ttl).await; + Ok(value) + } + } +} + #[async_trait] pub trait DistributedRateLimiter: Send + Sync { async fn check_rate_limit(&self, key: &str, limit: u32, window_ms: u64) -> bool; diff --git a/crates/tranquil-infra/src/memory_cache.rs b/crates/tranquil-infra/src/memory_cache.rs new file mode 100644 index 0000000..eca7692 --- /dev/null +++ b/crates/tranquil-infra/src/memory_cache.rs @@ -0,0 +1,74 @@ +use crate::{Cache, CacheError}; +use async_trait::async_trait; +use std::collections::HashMap; +use std::sync::Mutex; +use std::time::{Duration, Instant}; + +struct Entry { + value: Vec, + expires_at: Instant, +} + +#[derive(Default)] +pub struct MemoryCache { + entries: Mutex>, +} + +impl MemoryCache { + pub fn new() -> Self { + Self::default() + } + + fn read(&self, key: &str) -> Option> { + let now = Instant::now(); + let mut entries = self.entries.lock().unwrap_or_else(|e| e.into_inner()); + match entries.get(key) { + Some(entry) if entry.expires_at > now => Some(entry.value.clone()), + Some(_) => { + entries.remove(key); + None + } + None => None, + } + } + + fn write(&self, key: &str, value: Vec, ttl: Duration) { + let entry = Entry { + value, + expires_at: Instant::now() + ttl, + }; + self.entries + .lock() + .unwrap_or_else(|e| e.into_inner()) + .insert(key.to_string(), entry); + } +} + +#[async_trait] +impl Cache for MemoryCache { + async fn get(&self, key: &str) -> Option { + self.read(key).and_then(|v| String::from_utf8(v).ok()) + } + + async fn set(&self, key: &str, value: &str, ttl: Duration) -> Result<(), CacheError> { + self.write(key, value.as_bytes().to_vec(), ttl); + Ok(()) + } + + async fn delete(&self, key: &str) -> Result<(), CacheError> { + self.entries + .lock() + .unwrap_or_else(|e| e.into_inner()) + .remove(key); + Ok(()) + } + + async fn get_bytes(&self, key: &str) -> Option> { + self.read(key) + } + + async fn set_bytes(&self, key: &str, value: &[u8], ttl: Duration) -> Result<(), CacheError> { + self.write(key, value.to_vec(), ttl); + Ok(()) + } +} diff --git a/crates/tranquil-oauth-server/src/endpoints/delegation.rs b/crates/tranquil-oauth-server/src/endpoints/delegation.rs index 748f90a..9f4dbba 100644 --- a/crates/tranquil-oauth-server/src/endpoints/delegation.rs +++ b/crates/tranquil-oauth-server/src/endpoints/delegation.rs @@ -13,7 +13,8 @@ use tranquil_pds::rate_limit::{LoginLimit, OAuthRateLimited, TotpVerifyLimit}; use tranquil_pds::state::AppState; use tranquil_pds::types::PlainPassword; use tranquil_pds::util::ClientIp; -use tranquil_types::did_doc::{extract_handle, extract_pds_endpoint}; +use tranquil_types::did_doc::{PdsEndpointError, extract_handle, extract_pds_endpoint}; +use tranquil_types::url_kind; use tranquil_types::{Did, RequestId}; #[allow(clippy::result_large_err)] @@ -231,11 +232,17 @@ pub async fn delegation_auth( } }; - let pds_url = match extract_pds_endpoint(&did_doc) { - Some(url) => url, - None => { + let pds_url = match extract_pds_endpoint::(&did_doc) { + Ok(url) => url, + Err(PdsEndpointError::Missing) => { return DelegationAuthResponse::err("Controller has no PDS endpoint"); } + Err(PdsEndpointError::Invalid(e)) => { + tracing::warn!(controller = %controller_did, error = %e, "Controller PDS endpoint rejected"); + return DelegationAuthResponse::err( + "Controller PDS endpoint isn't a usable https URL", + ); + } }; let hostname = &tranquil_config::get().server.hostname; @@ -447,7 +454,7 @@ pub async fn delegation_auth_token( #[derive(Debug, Deserialize)] pub struct CrossPdsCallbackParams { pub code: tranquil_types::AuthorizationCode, - pub state: String, + pub state: tranquil_types::CrossPdsState, pub iss: Option, } @@ -474,7 +481,7 @@ pub async fn delegation_callback( if let Some(ref expected_issuer) = auth_state.expected_issuer { match ¶ms.iss { - Some(iss) if iss != expected_issuer => { + Some(iss) if iss.as_str() != expected_issuer.as_str() => { tracing::error!( "Cross-PDS issuer mismatch: expected {}, got {}", expected_issuer, diff --git a/crates/tranquil-oauth/Cargo.toml b/crates/tranquil-oauth/Cargo.toml index 5194281..b7b1d5f 100644 --- a/crates/tranquil-oauth/Cargo.toml +++ b/crates/tranquil-oauth/Cargo.toml @@ -6,6 +6,7 @@ license.workspace = true [dependencies] tranquil-types = { workspace = true } +tranquil-infra = { workspace = true, features = ["cache-keys"] } anyhow = { workspace = true } sqlx = { workspace = true } diff --git a/crates/tranquil-oauth/src/types.rs b/crates/tranquil-oauth/src/types.rs index 806268f..5f8f01b 100644 --- a/crates/tranquil-oauth/src/types.rs +++ b/crates/tranquil-oauth/src/types.rs @@ -1,7 +1,7 @@ use chrono::{DateTime, Utc}; use serde::{Deserialize, Serialize}; use serde_json::Value as JsonValue; -use tranquil_types::{ClientId, Did}; +use tranquil_types::{AuthServerEndpoint, ClientId, Did, Issuer}; pub use tranquil_types::{AuthorizationCode, DeviceId, RefreshToken, RequestId, TokenId}; @@ -195,9 +195,9 @@ pub struct ProtectedResourceMetadata { #[derive(Debug, Clone, Serialize, Deserialize)] pub struct AuthorizationServerMetadata { - pub issuer: String, - pub authorization_endpoint: String, - pub token_endpoint: String, + pub issuer: Issuer, + pub authorization_endpoint: AuthServerEndpoint, + pub token_endpoint: AuthServerEndpoint, pub jwks_uri: String, pub registration_endpoint: Option, pub scopes_supported: Option>, @@ -206,7 +206,7 @@ pub struct AuthorizationServerMetadata { pub grant_types_supported: Option>, pub token_endpoint_auth_methods_supported: Option>, pub code_challenge_methods_supported: Option>, - pub pushed_authorization_request_endpoint: Option, + pub pushed_authorization_request_endpoint: Option, pub require_pushed_authorization_requests: Option, pub dpop_signing_alg_values_supported: Option>, pub authorization_response_iss_parameter_supported: Option, diff --git a/crates/tranquil-pds/Cargo.toml b/crates/tranquil-pds/Cargo.toml index 6e90db3..59fb81a 100644 --- a/crates/tranquil-pds/Cargo.toml +++ b/crates/tranquil-pds/Cargo.toml @@ -15,7 +15,7 @@ tranquil-auth = { workspace = true } tranquil-oauth = { workspace = true } tranquil-comms = { workspace = true } tranquil-signal = { workspace = true } -tranquil-db = { workspace = true } +tranquil-db = { workspace = true, features = ["postgres"] } tranquil-db-traits = { workspace = true } tranquil-store = { workspace = true } tranquil-lexicon = { workspace = true, features = ["resolve"] } @@ -86,6 +86,7 @@ frontend = [] native-tls-roots = ["tranquil-oauth/native-tls-roots"] [dev-dependencies] +tranquil-infra = { workspace = true, features = ["testing"] } tempfile = "3" ciborium = { workspace = true } ctor = { workspace = true } diff --git a/crates/tranquil-pds/src/cache/mod.rs b/crates/tranquil-pds/src/cache/mod.rs index 52ddb33..1386b86 100644 --- a/crates/tranquil-pds/src/cache/mod.rs +++ b/crates/tranquil-pds/src/cache/mod.rs @@ -1,4 +1,6 @@ -pub use tranquil_cache::{Cache, CacheError, DistributedRateLimiter, NoOpCache, create_cache}; +pub use tranquil_cache::{ + Cache, CacheError, DistributedRateLimiter, NoOpCache, cached_json, create_cache, +}; #[cfg(feature = "valkey")] pub use tranquil_cache::{RedisRateLimiter, ValkeyCache}; diff --git a/crates/tranquil-pds/src/cache_keys.rs b/crates/tranquil-pds/src/cache_keys.rs index 3a148ce..b2f2840 100644 --- a/crates/tranquil-pds/src/cache_keys.rs +++ b/crates/tranquil-pds/src/cache_keys.rs @@ -1,48 +1 @@ -use crate::types::{CidLink, Did, Handle, Jti}; - -pub fn session_key(did: &Did, jti: &Jti) -> String { - format!("auth:session:{}:{}", did, jti) -} - -pub fn signing_key_key(did: &Did) -> String { - format!("auth:key:{}", did) -} - -pub fn user_status_key(did: &Did) -> String { - format!("auth:status:{}", did) -} - -pub fn handle_key(handle: &Handle) -> String { - format!("handle:{}", handle) -} - -pub fn reauth_key(did: &Did) -> String { - format!("reauth:{}", did) -} - -pub fn plc_doc_key(did: &Did) -> String { - format!("plc:doc:{}", did) -} - -pub fn plc_data_key(did: &Did) -> String { - format!("plc:data:{}", did) -} - -pub fn email_update_key(did: &Did) -> String { - format!("email_update:{}", did) -} - -pub fn scope_ref_key(cid: &CidLink) -> String { - format!("scope_ref:{}", cid) -} - -pub fn auto_verify_sent_key(did: &Did) -> String { - format!("auto_verify_sent:{}", did) -} - -pub fn permission_set_key(nsid: &tranquil_types::Nsid, aud: Option<&str>) -> String { - match aud { - Some(a) => format!("permset:{}:{}", nsid, a), - None => format!("permset:{}", nsid), - } -} +pub use tranquil_cache::cache_keys::*; diff --git a/crates/tranquil-pds/src/delegation/mod.rs b/crates/tranquil-pds/src/delegation/mod.rs index e793f8d..3bf6e5f 100644 --- a/crates/tranquil-pds/src/delegation/mod.rs +++ b/crates/tranquil-pds/src/delegation/mod.rs @@ -13,6 +13,16 @@ pub use tranquil_db_traits::DelegationActionType; use crate::did::DidResolutionError; use crate::state::AppState; use crate::types::{Did, Handle}; +use tranquil_types::did_doc::{PdsEndpointError, extract_handle, extract_pds_endpoint}; +use tranquil_types::{InvalidHttpUrl, PdsUrl}; + +#[derive(Debug, thiserror::Error)] +pub enum IdentityResolutionError { + #[error(transparent)] + DidResolution(#[from] DidResolutionError), + #[error("remote PDS endpoint is unusable: {0}")] + PdsEndpoint(InvalidHttpUrl), +} #[derive(serde::Serialize)] #[serde(rename_all = "camelCase")] @@ -21,14 +31,14 @@ pub struct ResolvedIdentity { #[serde(skip_serializing_if = "Option::is_none")] pub handle: Option, #[serde(skip_serializing_if = "Option::is_none")] - pub pds_url: Option, + pub pds_url: Option, pub is_local: bool, } pub async fn resolve_identity( state: &AppState, did: &Did, -) -> Result { +) -> Result { let is_local = state .repos .user @@ -38,26 +48,23 @@ pub async fn resolve_identity( .flatten() .is_some(); - let did_doc = state.did_resolver.resolve_did(did).await?; + let did_doc = state.did_resolver.fetch_did_document(did).await?; - let pds_url = did_doc.services.iter().find_map(|svc| { - if (svc.id == "#atproto_pds" || svc.id.ends_with("#atproto_pds")) - && svc.service_type == "AtprotoPersonalDataServer" - { - Some(svc.service_endpoint.clone()) - } else { + let pds_url = match (extract_pds_endpoint(&did_doc), is_local) { + (Ok(url), _) => Some(url), + (Err(PdsEndpointError::Missing), _) => None, + (Err(PdsEndpointError::Invalid(e)), true) => { + tracing::debug!(did = %did, error = %e, "local account has an unusable PDS endpoint"); None } - }); - let handle = did_doc - .also_known_as - .iter() - .find_map(|alias| alias.strip_prefix("at://")) - .and_then(|s| Handle::new(s).ok()); + (Err(PdsEndpointError::Invalid(e)), false) => { + return Err(IdentityResolutionError::PdsEndpoint(e)); + } + }; Ok(ResolvedIdentity { did: did.clone(), - handle, + handle: extract_handle(&did_doc), pds_url, is_local, }) diff --git a/crates/tranquil-pds/src/oauth/client.rs b/crates/tranquil-pds/src/oauth/client.rs index 2f37144..fc65eb7 100644 --- a/crates/tranquil-pds/src/oauth/client.rs +++ b/crates/tranquil-pds/src/oauth/client.rs @@ -10,10 +10,12 @@ use tranquil_oauth::{ AuthorizationServerMetadata, ClientMetadata, compute_es256_jkt, compute_pkce_challenge, create_dpop_proof, }; -use tranquil_types::{AuthorizationCode, ClientId, Did}; +use tranquil_types::{AuthorizationCode, ClientId, CrossPdsState, Did, Issuer, PdsUrl}; use crate::cache::Cache; +const SERVER_METADATA_TTL: Duration = Duration::from_secs(300); + #[derive(Error, Debug)] pub enum CrossPdsError { #[error("failed to fetch OAuth metadata: {0}")] @@ -32,11 +34,11 @@ pub enum CrossPdsError { pub struct CrossPdsAuthState { pub original_request_uri: String, pub controller_did: Did, - pub controller_pds_url: String, + pub controller_pds_url: PdsUrl, pub code_verifier: String, pub dpop_private_key_der: String, pub delegated_did: Did, - pub expected_issuer: Option, + pub expected_issuer: Option, } #[derive(Debug, Clone, Serialize, Deserialize)] @@ -70,17 +72,23 @@ impl CrossPdsOAuthClient { let http = Client::builder() .timeout(Duration::from_secs(15)) .connect_timeout(Duration::from_secs(5)) + .redirect(tranquil_types::redirect_policy( + tranquil_types::ReachPolicy::GlobalOnly, + )) + .dns_resolver(tranquil_types::dns_guard( + tranquil_types::ReachPolicy::GlobalOnly, + )) .build() - .unwrap_or_else(|_| Client::new()); + .expect("failed to build cross-PDS OAuth HTTP client"); Self { http, cache } } pub async fn store_auth_state( &self, - state_key: &str, + state_key: &CrossPdsState, auth_state: &CrossPdsAuthState, ) -> Result<(), CrossPdsError> { - let cache_key = format!("cross_pds_state:{}", state_key); + let cache_key = crate::cache_keys::cross_pds_state_key(state_key); let json_bytes = serde_json::to_vec(auth_state) .map_err(|e| CrossPdsError::ParFailed(format!("serialize auth state: {}", e)))?; let encrypted = crate::config::encrypt_key(&json_bytes) @@ -93,9 +101,9 @@ impl CrossPdsOAuthClient { pub async fn retrieve_auth_state( &self, - state_key: &str, + state_key: &CrossPdsState, ) -> Result { - let cache_key = format!("cross_pds_state:{}", state_key); + let cache_key = crate::cache_keys::cross_pds_state_key(state_key); let encrypted_bytes = self.cache.get_bytes(&cache_key).await.ok_or_else(|| { CrossPdsError::TokenExchangeFailed("auth state expired or not found".into()) })?; @@ -110,13 +118,11 @@ impl CrossPdsOAuthClient { }) } - pub async fn check_remote_is_delegated(&self, pds_url: &str, did: &Did) -> Option { - let url = format!( - "{}/oauth/security-status?identifier={}", - pds_url.trim_end_matches('/'), - urlencoding::encode(did.as_str()) - ); - let resp = self.http.get(&url).send().await.ok()?; + pub async fn check_remote_is_delegated(&self, pds_url: &PdsUrl, did: &Did) -> Option { + let mut url = pds_url.endpoint("oauth/security-status"); + url.query_pairs_mut() + .append_pair("identifier", did.as_str()); + let resp = self.http.get(url).send().await.ok()?; if !resp.status().is_success() { return None; } @@ -176,24 +182,12 @@ impl CrossPdsOAuthClient { Ok(resp) } - fn require_https(url: &str, label: &str) -> Result<(), CrossPdsError> { - if !url.starts_with("https://") { - return Err(CrossPdsError::MetadataFetch(format!( - "{} must use HTTPS, got: {}", - label, url - ))); - } - Ok(()) - } - - async fn resolve_authorization_server(&self, pds_url: &str) -> Result { - Self::require_https(pds_url, "PDS URL")?; - - let resource_url = format!( - "{}/.well-known/oauth-protected-resource", - pds_url.trim_end_matches('/') - ); - if let Ok(resp) = self.http.get(&resource_url).send().await + async fn resolve_authorization_server( + &self, + pds_url: &PdsUrl, + ) -> Result { + let resource_url = pds_url.endpoint(".well-known/oauth-protected-resource"); + if let Ok(resp) = self.http.get(resource_url).send().await && resp.status().is_success() { #[derive(Deserialize)] @@ -203,30 +197,36 @@ impl CrossPdsOAuthClient { if let Ok(pr) = resp.json::().await && let Some(server) = pr.authorization_servers.and_then(|s| s.into_iter().next()) { - Self::require_https(&server, "Authorization server")?; - return Ok(server); + return Issuer::new(server) + .map_err(|e| CrossPdsError::MetadataFetch(e.to_string())); } } - Ok(pds_url.trim_end_matches('/').to_string()) + Issuer::new(pds_url.as_str()).map_err(|e| CrossPdsError::MetadataFetch(e.to_string())) } pub async fn fetch_server_metadata( &self, - pds_url: &str, + pds_url: &PdsUrl, ) -> Result { - let cache_key = format!("cross_pds_oauth_meta:{}", pds_url); - if let Some(cached) = self.cache.get(&cache_key).await - && let Ok(meta) = serde_json::from_str(&cached) - { - return Ok(meta); - } + crate::cache::cached_json( + self.cache.as_ref(), + &crate::cache_keys::cross_pds_oauth_meta_key(pds_url), + SERVER_METADATA_TTL, + || self.fetch_verified_server_metadata(pds_url), + ) + .await + } + async fn fetch_verified_server_metadata( + &self, + pds_url: &PdsUrl, + ) -> Result { let auth_server = self.resolve_authorization_server(pds_url).await?; - let url = format!("{}/.well-known/oauth-authorization-server", auth_server); + let url = auth_server.endpoint(".well-known/oauth-authorization-server"); let resp = self .http - .get(&url) + .get(url.clone()) .send() .await .map_err(|e| CrossPdsError::MetadataFetch(e.to_string()))?; @@ -244,11 +244,11 @@ impl CrossPdsOAuthClient { .await .map_err(|e| CrossPdsError::MetadataFetch(e.to_string()))?; - if let Ok(json_str) = serde_json::to_string(&meta) { - let _ = self - .cache - .set(&cache_key, &json_str, Duration::from_secs(300)) - .await; + if meta.issuer != auth_server { + return Err(CrossPdsError::MetadataFetch(format!( + "issuer mismatch: {} serves metadata for {}", + auth_server, meta.issuer + ))); } Ok(meta) @@ -256,22 +256,22 @@ impl CrossPdsOAuthClient { pub async fn initiate_par( &self, - pds_url: &str, + pds_url: &PdsUrl, urls: &DelegationOAuthUrls, login_hint: Option<&str>, original_request_uri: &str, controller_did: &Did, delegated_did: &Did, - ) -> Result<(ParResult, CrossPdsAuthState, String), CrossPdsError> { + ) -> Result<(ParResult, CrossPdsAuthState, CrossPdsState), CrossPdsError> { let meta = self.fetch_server_metadata(pds_url).await?; let par_endpoint = meta .pushed_authorization_request_endpoint - .as_deref() + .as_ref() .ok_or(CrossPdsError::NoParEndpoint)?; let code_verifier = crate::util::generate_random_token(); let code_challenge = compute_pkce_challenge(&code_verifier); - let state = crate::util::generate_random_token(); + let state = CrossPdsState::new(crate::util::generate_random_token()); let signing_key = SigningKey::random(&mut OsRng); let dpop_key_der = URL_SAFE_NO_PAD.encode(signing_key.to_bytes()); @@ -284,7 +284,7 @@ impl CrossPdsOAuthClient { ("client_id", urls.client_id.to_string()), ("redirect_uri", urls.redirect_uri.clone()), ("scope", "atproto".to_string()), - ("state", state.clone()), + ("state", state.to_string()), ("code_challenge", code_challenge), ("code_challenge_method", "S256".to_string()), ("dpop_jkt", dpop_jkt), @@ -294,7 +294,7 @@ impl CrossPdsOAuthClient { } let resp = self - .send_with_dpop_retry(&signing_key, "POST", par_endpoint, ¶ms, None) + .send_with_dpop_retry(&signing_key, "POST", par_endpoint.as_str(), ¶ms, None) .await .map_err(|e| CrossPdsError::ParFailed(e.to_string()))?; @@ -313,17 +313,16 @@ impl CrossPdsOAuthClient { .await .map_err(|e| CrossPdsError::ParFailed(e.to_string()))?; - let authorize_url = format!( - "{}?request_uri={}&client_id={}", - meta.authorization_endpoint, - urlencoding::encode(&par_resp.request_uri), - urlencoding::encode(&urls.client_id) - ); + let mut authorize_url = meta.authorization_endpoint.url().clone(); + authorize_url + .query_pairs_mut() + .append_pair("request_uri", &par_resp.request_uri) + .append_pair("client_id", &urls.client_id); let auth_state = CrossPdsAuthState { original_request_uri: original_request_uri.to_string(), controller_did: controller_did.clone(), - controller_pds_url: pds_url.to_string(), + controller_pds_url: pds_url.clone(), code_verifier, dpop_private_key_der: dpop_key_der, delegated_did: delegated_did.clone(), @@ -333,7 +332,7 @@ impl CrossPdsOAuthClient { Ok(( ParResult { request_uri: par_resp.request_uri, - authorize_url, + authorize_url: authorize_url.into(), }, auth_state, state, @@ -366,7 +365,13 @@ impl CrossPdsOAuthClient { ]; let resp = self - .send_with_dpop_retry(&signing_key, "POST", &meta.token_endpoint, ¶ms, None) + .send_with_dpop_retry( + &signing_key, + "POST", + meta.token_endpoint.as_str(), + ¶ms, + None, + ) .await .map_err(CrossPdsError::TokenExchangeFailed)?; diff --git a/crates/tranquil-signal/Cargo.toml b/crates/tranquil-signal/Cargo.toml index ed8d49a..3db515e 100644 --- a/crates/tranquil-signal/Cargo.toml +++ b/crates/tranquil-signal/Cargo.toml @@ -18,7 +18,7 @@ tokio = { workspace = true } tokio-util = { workspace = true } futures = { workspace = true } serde_json = { workspace = true } -url = "2.5" +url = { workspace = true } uuid = { workspace = true } thiserror = { workspace = true } diff --git a/crates/tranquil-types/Cargo.toml b/crates/tranquil-types/Cargo.toml index 4055d2b..f0c808f 100644 --- a/crates/tranquil-types/Cargo.toml +++ b/crates/tranquil-types/Cargo.toml @@ -10,8 +10,15 @@ chrono = { workspace = true } cid = { workspace = true } jacquard-common = { workspace = true } rand = { workspace = true } +reqwest = { workspace = true } serde = { workspace = true } serde_json = { workspace = true } sqlx = { workspace = true } thiserror = { workspace = true } +tokio = { workspace = true, features = ["net", "rt"] } +tracing = { workspace = true } +url = { workspace = true } uuid = { workspace = true } + +[dev-dependencies] +tokio = { workspace = true } diff --git a/crates/tranquil-types/src/lib.rs b/crates/tranquil-types/src/lib.rs index 039fdcd..6fe70a8 100644 --- a/crates/tranquil-types/src/lib.rs +++ b/crates/tranquil-types/src/lib.rs @@ -1,6 +1,8 @@ use serde::{Deserialize, Serialize}; use std::borrow::Cow; use std::fmt; +use std::hash::Hash; +use std::marker::PhantomData; use std::ops::Deref; use std::str::FromStr; @@ -813,6 +815,10 @@ simple_string_newtype! { pub struct Jti; } +simple_string_newtype_no_sqlx! { + pub struct CrossPdsState; +} + simple_string_newtype! { pub struct AuthorizationCode; } @@ -881,6 +887,425 @@ impl fmt::Display for CommsChannel { } } +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum HostReach { + Global, + Loopback, + Private, +} + +fn ipv4_reach(ip: std::net::Ipv4Addr) -> HostReach { + let [a, b, c, _] = ip.octets(); + match ip { + _ if ip.is_loopback() => HostReach::Loopback, + _ if ip.is_private() + || ip.is_link_local() + || ip.is_multicast() + || ip.is_documentation() + || a == 0 + || a == 100 && (64..128).contains(&b) + || a == 192 && b == 0 && c == 0 + || a == 192 && b == 88 && c == 99 + || a == 198 && (18..20).contains(&b) + || a & 0xf0 == 240 => + { + HostReach::Private + } + _ => HostReach::Global, + } +} + +fn ipv6_reach(ip: std::net::Ipv6Addr) -> HostReach { + let seg = ip.segments(); + let embedded_ipv4 = + |hi: u16, lo: u16| std::net::Ipv4Addr::from((u32::from(hi) << 16) | u32::from(lo)); + match ip.to_ipv4_mapped() { + Some(mapped) => ipv4_reach(mapped), + None => match ip { + _ if ip.is_loopback() => HostReach::Loopback, + _ if seg[..6] == [0, 0, 0, 0, 0, 0] => ipv4_reach(embedded_ipv4(seg[6], seg[7])), + _ if seg[..2] == [0x2001, 0] => ipv4_reach(embedded_ipv4(!seg[6], !seg[7])), + _ if seg[0] == 0x2002 => ipv4_reach(embedded_ipv4(seg[1], seg[2])), + _ if seg[..6] == [0x64, 0xff9b, 0, 0, 0, 0] => { + ipv4_reach(embedded_ipv4(seg[6], seg[7])) + } + _ if seg[..3] == [0x64, 0xff9b, 1] => HostReach::Private, + _ if ip.is_unspecified() + || ip.is_multicast() + || seg[0] & 0xfe00 == 0xfc00 + || seg[0] & 0xffc0 == 0xfe80 + || seg[..2] == [0x2001, 0x0db8] => + { + HostReach::Private + } + _ => HostReach::Global, + }, + } +} + +fn host_reach(host: url::Host<&str>) -> HostReach { + match host { + url::Host::Ipv4(ip) => ipv4_reach(ip), + url::Host::Ipv6(ip) => ipv6_reach(ip), + url::Host::Domain(name) => { + let name = name.trim_end_matches('.').to_ascii_lowercase(); + match name.as_str() { + "localhost" => HostReach::Loopback, + _ if name.ends_with(".localhost") => HostReach::Loopback, + _ if name.ends_with(".local") + || name.ends_with(".internal") + || name.ends_with(".home.arpa") + || name == "home.arpa" => + { + HostReach::Private + } + _ => HostReach::Global, + } + } + } +} + +pub fn url_reach(url: &url::Url) -> Option { + url.host().map(host_reach) +} + +pub fn ip_reach(ip: std::net::IpAddr) -> HostReach { + match ip { + std::net::IpAddr::V4(v4) => ipv4_reach(v4), + std::net::IpAddr::V6(v6) => ipv6_reach(v6), + } +} + +pub fn reach_permits(reach: HostReach, policy: ReachPolicy) -> bool { + matches!( + (reach, policy), + (HostReach::Global, _) + | ( + HostReach::Loopback, + ReachPolicy::AllowLoopback | ReachPolicy::AllowPrivate, + ) + | (HostReach::Private, ReachPolicy::AllowPrivate) + ) +} + +pub fn url_reach_permits(url: &url::Url, policy: ReachPolicy) -> bool { + let Some(reach) = url_reach(url) else { + return false; + }; + let scheme_permits = matches!( + (url.scheme(), reach), + ("https", _) | ("http", HostReach::Loopback | HostReach::Private) + ); + scheme_permits && reach_permits(reach, policy) +} + +fn parse_http_url(s: &str, policy: ReachPolicy, allow_query: bool) -> Option { + let parsed = url::Url::parse(s).ok()?; + let rejected = (parsed.query().is_some() && !allow_query) + || parsed.fragment().is_some() + || !parsed.username().is_empty() + || parsed.password().is_some() + || !url_reach_permits(&parsed, policy); + match rejected { + true => None, + false => Some(parsed), + } +} + +const REDIRECT_HOP_LIMIT: usize = 5; + +pub fn redirect_policy(policy: ReachPolicy) -> reqwest::redirect::Policy { + reqwest::redirect::Policy::custom(move |attempt| { + let over_limit = attempt.previous().len() > REDIRECT_HOP_LIMIT; + let permitted = url_reach_permits(attempt.url(), policy); + let target = attempt.url().clone(); + match (over_limit, permitted) { + (true, _) => attempt.error(format!("more than {} redirect hops", REDIRECT_HOP_LIMIT)), + (false, false) => attempt.error(format!( + "redirect target {} is outside the allowed host reach", + target + )), + (false, true) => attempt.follow(), + } + }) +} + +pub struct ReachGuardedDns(ReachPolicy); + +impl reqwest::dns::Resolve for ReachGuardedDns { + fn resolve(&self, name: reqwest::dns::Name) -> reqwest::dns::Resolving { + let policy = self.0; + Box::pin(async move { + let host = name.as_str().to_owned(); + let permitted: Vec = tokio::net::lookup_host((host.as_str(), 0)) + .await? + .filter(|addr| reach_permits(ip_reach(addr.ip()), policy)) + .collect(); + match permitted.is_empty() { + true => Err(format!( + "no resolved address for {} is inside the allowed host reach", + host + ) + .into()), + false => Ok(Box::new(permitted.into_iter()) as reqwest::dns::Addrs), + } + }) + } +} + +pub fn dns_guard(policy: ReachPolicy) -> std::sync::Arc { + std::sync::Arc::new(ReachGuardedDns(policy)) +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum ReachPolicy { + AllowLoopback, + AllowPrivate, + GlobalOnly, +} + +impl ReachPolicy { + #[cfg(debug_assertions)] + pub const DEBUG_LOOPBACK: ReachPolicy = ReachPolicy::AllowLoopback; + #[cfg(not(debug_assertions))] + pub const DEBUG_LOOPBACK: ReachPolicy = ReachPolicy::GlobalOnly; +} + +pub trait UrlKind { + const LABEL: &'static str; + const REACH_POLICY: ReachPolicy; + const ALLOW_QUERY: bool; +} + +pub mod url_kind { + use super::{ReachPolicy, UrlKind}; + + pub struct AuthServerEndpoint; + impl UrlKind for AuthServerEndpoint { + const LABEL: &'static str = "authorization server endpoint"; + const REACH_POLICY: ReachPolicy = ReachPolicy::GlobalOnly; + const ALLOW_QUERY: bool = true; + } + + pub struct Issuer; + impl UrlKind for Issuer { + const LABEL: &'static str = "issuer"; + const REACH_POLICY: ReachPolicy = ReachPolicy::GlobalOnly; + const ALLOW_QUERY: bool = false; + } + + pub struct Jwks; + impl UrlKind for Jwks { + const LABEL: &'static str = "JWKS URI"; + const REACH_POLICY: ReachPolicy = ReachPolicy::DEBUG_LOOPBACK; + const ALLOW_QUERY: bool = true; + } + + pub struct Pds; + impl UrlKind for Pds { + const LABEL: &'static str = "PDS URL"; + const REACH_POLICY: ReachPolicy = ReachPolicy::GlobalOnly; + const ALLOW_QUERY: bool = false; + } + + pub struct SchemaHost; + impl UrlKind for SchemaHost { + const LABEL: &'static str = "schema host URL"; + const REACH_POLICY: ReachPolicy = ReachPolicy::DEBUG_LOOPBACK; + const ALLOW_QUERY: bool = false; + } + + pub struct SsoIssuer; + impl UrlKind for SsoIssuer { + const LABEL: &'static str = "SSO issuer"; + const REACH_POLICY: ReachPolicy = ReachPolicy::AllowPrivate; + const ALLOW_QUERY: bool = false; + } + + pub struct SsoJwks; + impl UrlKind for SsoJwks { + const LABEL: &'static str = "SSO JWKS URI"; + const REACH_POLICY: ReachPolicy = ReachPolicy::AllowPrivate; + const ALLOW_QUERY: bool = true; + } +} + +pub struct HttpUrl { + raw: String, + parsed: url::Url, + kind: PhantomData K>, +} + +pub type AuthServerEndpoint = HttpUrl; +pub type Issuer = HttpUrl; +pub type JwksUri = HttpUrl; +pub type PdsUrl = HttpUrl; +pub type SchemaHostUrl = HttpUrl; +pub type SsoIssuer = HttpUrl; +pub type SsoJwksUri = HttpUrl; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct InvalidHttpUrl { + pub kind: &'static str, + pub value: String, +} + +impl fmt::Display for InvalidHttpUrl { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + write!(f, "invalid {}: {}", self.kind, self.value) + } +} + +impl std::error::Error for InvalidHttpUrl {} + +impl HttpUrl { + pub fn new(s: impl Into) -> Result { + let raw = s.into(); + match parse_http_url(&raw, K::REACH_POLICY, K::ALLOW_QUERY) { + Some(parsed) => Ok(Self { + raw, + parsed, + kind: PhantomData, + }), + None => Err(InvalidHttpUrl { + kind: K::LABEL, + value: raw, + }), + } + } + + /// The URL as given. + /// OIDC & OAuth define issuer comparison as an + /// exact string match, + /// so anything sent to or compared against a peer uses this. + /// Give it to us raw & wriggling!! + pub fn as_str(&self) -> &str { + &self.raw + } + + /// The parsed form: lowercased scheme and host, with `/` for a bare authority. + /// Cache keys use this so `https://oyster.cafe` and `https://oyster.cafe/` share one entry. + pub fn canonical(&self) -> &str { + self.parsed.as_str() + } + + pub fn url(&self) -> &url::Url { + &self.parsed + } + + pub fn endpoint(&self, path: &str) -> url::Url { + let mut url = self.parsed.clone(); + let base = url.path().trim_end_matches('/').to_owned(); + url.set_path(&format!("{}/{}", base, path.trim_start_matches('/'))); + url + } +} + +pub mod http_url { + use super::{HttpUrl, UrlKind}; + use serde::Deserialize; + + pub fn deserialize_optional<'de, D, K>(deserializer: D) -> Result>, D::Error> + where + D: serde::Deserializer<'de>, + K: UrlKind, + { + Ok(Option::::deserialize(deserializer)?.and_then(|s| { + HttpUrl::new(s) + .inspect_err(|e| tracing::warn!(error = %e, "discarding unusable URL field")) + .ok() + })) + } +} + +impl FromStr for HttpUrl { + type Err = InvalidHttpUrl; + + fn from_str(s: &str) -> Result { + Self::new(s) + } +} + +impl fmt::Debug for HttpUrl { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + write!(f, "{}({})", K::LABEL, self.raw) + } +} + +impl fmt::Display for HttpUrl { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str(&self.raw) + } +} + +impl Clone for HttpUrl { + fn clone(&self) -> Self { + Self { + raw: self.raw.clone(), + parsed: self.parsed.clone(), + kind: PhantomData, + } + } +} + +impl PartialEq for HttpUrl { + fn eq(&self, other: &Self) -> bool { + self.parsed == other.parsed + } +} + +impl Eq for HttpUrl {} + +impl Hash for HttpUrl { + fn hash(&self, state: &mut H) { + self.parsed.hash(state); + } +} + +impl Serialize for HttpUrl { + fn serialize(&self, serializer: S) -> Result { + serializer.serialize_str(&self.raw) + } +} + +impl<'de, K: UrlKind> Deserialize<'de> for HttpUrl { + fn deserialize(deserializer: D) -> Result + where + D: serde::Deserializer<'de>, + { + let s = String::deserialize(deserializer)?; + Self::new(s).map_err(|e| serde::de::Error::custom(e.to_string())) + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum EmailTokenPurpose { + UpdateEmail, + ConfirmEmail, + DeleteAccount, + ResetPassword, + PlcOperation, +} + +impl EmailTokenPurpose { + pub fn as_str(&self) -> &'static str { + match self { + Self::UpdateEmail => "update_email", + Self::ConfirmEmail => "confirm_email", + Self::DeleteAccount => "delete_account", + Self::ResetPassword => "reset_password", + Self::PlcOperation => "plc_operation", + } + } +} + +impl fmt::Display for EmailTokenPurpose { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + write!(f, "{}", self.as_str()) + } +} + #[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize, sqlx::Type)] #[serde(rename_all = "snake_case")] #[sqlx(type_name = "comms_type", rename_all = "snake_case")] @@ -894,24 +1319,32 @@ pub enum CommsType { } pub mod did_doc { - pub fn extract_pds_endpoint(doc: &serde_json::Value) -> Option { + use crate::{HttpUrl, InvalidHttpUrl, UrlKind}; + + #[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)] + pub enum PdsEndpointError { + #[error("DID document has no atproto PDS service entry")] + Missing, + #[error(transparent)] + Invalid(#[from] InvalidHttpUrl), + } + + pub fn extract_pds_endpoint( + doc: &serde_json::Value, + ) -> Result, PdsEndpointError> { doc.get("service") .and_then(|s| s.as_array()) .and_then(|services| { services.iter().find_map(|svc| { let id = svc.get("id").and_then(|v| v.as_str()).unwrap_or_default(); let svc_type = svc.get("type").and_then(|v| v.as_str()).unwrap_or_default(); - if (id == "#atproto_pds" || id.ends_with("#atproto_pds")) - && svc_type == "AtprotoPersonalDataServer" - { - svc.get("serviceEndpoint") - .and_then(|v| v.as_str()) - .map(|s| s.to_string()) - } else { - None - } + ((id == "#atproto_pds" || id.ends_with("#atproto_pds")) + && svc_type == "AtprotoPersonalDataServer") + .then(|| svc.get("serviceEndpoint").and_then(|v| v.as_str()))? }) }) + .ok_or(PdsEndpointError::Missing) + .and_then(|endpoint| HttpUrl::new(endpoint).map_err(PdsEndpointError::Invalid)) } pub fn extract_handle(doc: &serde_json::Value) -> Option { @@ -928,6 +1361,187 @@ pub mod did_doc { } } +#[cfg(test)] +mod http_url_tests { + use super::did_doc::{PdsEndpointError, extract_pds_endpoint}; + use super::{ + AuthServerEndpoint, Issuer, JwksUri, PdsUrl, SchemaHostUrl, SsoIssuer, SsoJwksUri, + }; + + #[test] + fn extract_pds_endpoint_selects_the_pds_service_and_reports_missing_or_invalid() { + let labeler = serde_json::json!({ + "id": "#atproto_labeler", + "type": "AtprotoLabeler", + "serviceEndpoint": "https://labeler.nel.pet" + }); + let pds = |endpoint: &str| { + serde_json::json!({ + "id": "#atproto_pds", + "type": "AtprotoPersonalDataServer", + "serviceEndpoint": endpoint + }) + }; + let both = serde_json::json!({ "service": [labeler.clone(), pds("https://oyster.cafe")] }); + assert_eq!( + extract_pds_endpoint::(&both) + .unwrap() + .as_str(), + "https://oyster.cafe" + ); + [ + serde_json::json!({ "service": [labeler] }), + serde_json::json!({}), + ] + .iter() + .for_each(|doc| { + assert_eq!( + extract_pds_endpoint::(doc).unwrap_err(), + PdsEndpointError::Missing + ); + }); + let plain_http = serde_json::json!({ "service": [pds("http://oyster.cafe")] }); + assert!(matches!( + extract_pds_endpoint::(&plain_http), + Err(PdsEndpointError::Invalid(_)) + )); + } + + #[test] + fn pds_and_jwks_kinds_reject_private_and_reserved_addresses() { + [ + "https://10.0.0.1", + "https://192.168.1.1", + "https://172.16.0.1", + "https://169.254.169.254/latest/meta-data", + "https://100.64.0.1", + "https://0.1.2.3", + "https://192.0.0.8", + "https://192.88.99.1", + "https://[fd00::1]", + "https://[fe80::1]", + "https://[ff02::1]", + "https://[::ffff:10.0.0.1]", + "https://[64:ff9b::a00:1]", + "https://[64:ff9b:1::1]", + "https://[2002:a00:1::]", + "https://[::10.0.0.1]", + "https://[2001:0:0:0:0:0:f5ff:fffe]", + "https://kelp.internal", + "https://whelk.local", + "https://limpet.home.arpa", + ] + .iter() + .for_each(|url| { + assert!(PdsUrl::new(*url).is_err(), "PdsUrl must reject {url}"); + assert!(JwksUri::new(*url).is_err(), "JwksUri must reject {url}"); + }); + [ + "https://oyster.cafe", + "https://[64:ff9b::808:808]", + "https://[::8.8.8.8]", + "https://[2001::f7f7:f7f7]", + ] + .iter() + .for_each(|url| assert!(PdsUrl::new(*url).is_ok(), "PdsUrl must accept {url}")); + } + + #[test] + fn each_kind_applies_its_own_local_host_policy() { + assert!(PdsUrl::new("http://127.0.0.1:2583").is_err()); + assert!(PdsUrl::new("https://localhost").is_err()); + assert!(Issuer::new("http://localhost:8080").is_err()); + assert_eq!( + JwksUri::new("http://localhost:8080/keys").is_ok(), + cfg!(debug_assertions) + ); + assert_eq!( + SchemaHostUrl::new("http://127.0.0.1:2583").is_ok(), + cfg!(debug_assertions) + ); + assert!(SsoJwksUri::new("http://127.0.0.1:8080/keys").is_ok()); + assert!(SsoJwksUri::new("http://[::1]:8080/keys").is_ok()); + assert!(SsoJwksUri::new("http://squid.localhost:8080/keys").is_ok()); + assert!(SsoJwksUri::new("https://keycloak.internal/keys?client=squid").is_ok()); + assert!(SsoJwksUri::new("http://oyster.cafe/keys").is_err()); + assert!(SsoIssuer::new("https://keycloak.internal/realms/uni").is_ok()); + assert!(SsoIssuer::new("http://10.0.0.5:8080").is_ok()); + assert!(SsoIssuer::new("http://localhost:8080").is_ok()); + assert!(SsoIssuer::new("http://oyster.cafe").is_err()); + assert!(AuthServerEndpoint::new("https://oyster.cafe/oauth/par?tenant=uni").is_ok()); + [ + "https://169.254.169.254/oauth/par", + "https://[fd00::1]/oauth/par", + "http://127.0.0.1:2583/oauth/par", + "http://oyster.cafe/oauth/par", + ] + .iter() + .for_each(|url| { + assert!( + AuthServerEndpoint::new(*url).is_err(), + "AuthServerEndpoint must reject {url}" + ); + }); + } + + #[test] + fn canonicalization_keeps_identity_and_rejects_query_fragment_and_userinfo() { + assert_eq!( + PdsUrl::new("HTTPS://oyster.cafe").unwrap().canonical(), + "https://oyster.cafe/" + ); + let bare = PdsUrl::new("https://oyster.cafe").unwrap(); + let slashed = PdsUrl::new("https://oyster.cafe/").unwrap(); + assert_eq!(bare, slashed); + assert_eq!(bare.canonical(), slashed.canonical()); + let issuer = Issuer::new("https://accounts.google.com").unwrap(); + assert_eq!(issuer.as_str(), "https://accounts.google.com"); + assert_eq!(issuer.canonical(), "https://accounts.google.com/"); + assert_eq!( + PdsUrl::new("https://oyster.cafe/pds/") + .unwrap() + .endpoint(".well-known/oauth-protected-resource") + .as_str(), + "https://oyster.cafe/pds/.well-known/oauth-protected-resource" + ); + assert_eq!( + JwksUri::new("https://oyster.cafe/keys?appid=abc") + .expect("JwksUri keeps the query") + .canonical(), + "https://oyster.cafe/keys?appid=abc" + ); + assert!(PdsUrl::new("https://oyster.cafe/?x=1").is_err()); + assert!(PdsUrl::new("https://oyster.cafe/#frag").is_err()); + assert!(PdsUrl::new("https://nel:pw@oyster.cafe").is_err()); + assert!(Issuer::new("https://oyster.cafe/?x=1").is_err()); + assert!(JwksUri::new("https://oyster.cafe/keys#frag").is_err()); + } +} + +#[cfg(test)] +mod dns_guard_tests { + use super::{ReachPolicy, dns_guard}; + use reqwest::dns::Resolve; + + #[tokio::test] + async fn the_policy_gates_loopback_resolution() { + let name = |host: &str| host.parse::().expect("valid hostname"); + assert!( + dns_guard(ReachPolicy::GlobalOnly) + .resolve(name("localhost")) + .await + .is_err() + ); + let addrs: Vec<_> = dns_guard(ReachPolicy::AllowLoopback) + .resolve(name("localhost")) + .await + .expect("localhost resolves") + .collect(); + assert!(!addrs.is_empty()); + assert!(addrs.iter().all(|a| a.ip().is_loopback())); + } +} + #[cfg(test)] mod validated_newtype_tests { use super::*; -- 2.51.2 From 0274f19d7554a8a4bbc037a27f9ac62502eaa75b Mon Sep 17 00:00:00 2001 From: Lewis Date: Sun, 26 Jul 2026 13:57:07 +0300 Subject: [PATCH 04/22] auth: EmailTokenPurpose from tranquil-types, shared cache key fns, MemoryCache in tests Lewis: May this revision serve well! --- crates/tranquil-oauth-server/Cargo.toml | 1 + .../endpoints/authorize/scope_resolution.rs | 40 ++---- crates/tranquil-pds/src/auth/email_token.rs | 106 +++------------- crates/tranquil-pds/src/auth/legacy_2fa.rs | 119 +++++------------- .../src/oauth/permission_set_resolver.rs | 57 +++------ 5 files changed, 68 insertions(+), 255 deletions(-) diff --git a/crates/tranquil-oauth-server/Cargo.toml b/crates/tranquil-oauth-server/Cargo.toml index 3cf9ce3..aca1519 100644 --- a/crates/tranquil-oauth-server/Cargo.toml +++ b/crates/tranquil-oauth-server/Cargo.toml @@ -37,6 +37,7 @@ webauthn-rs = { workspace = true } [dev-dependencies] async-trait = { workspace = true } +tranquil-infra = { workspace = true, features = ["testing"] } [features] bsky = [] diff --git a/crates/tranquil-oauth-server/src/endpoints/authorize/scope_resolution.rs b/crates/tranquil-oauth-server/src/endpoints/authorize/scope_resolution.rs index ac12090..d2246cf 100644 --- a/crates/tranquil-oauth-server/src/endpoints/authorize/scope_resolution.rs +++ b/crates/tranquil-oauth-server/src/endpoints/authorize/scope_resolution.rs @@ -33,36 +33,12 @@ pub async fn resolve_effective_scopes( #[cfg(test)] mod tests { use super::*; - use std::collections::HashMap; - use std::sync::Mutex; use std::time::Duration; - use tranquil_pds::cache::{Cache, CacheError}; + use tranquil_infra::MemoryCache; + use tranquil_pds::cache::Cache; - #[derive(Default)] - struct MapCache(Mutex>); - #[async_trait::async_trait] - impl Cache for MapCache { - async fn get(&self, k: &str) -> Option { - self.0.lock().unwrap().get(k).cloned() - } - async fn set(&self, k: &str, v: &str, _t: Duration) -> Result<(), CacheError> { - self.0.lock().unwrap().insert(k.into(), v.into()); - Ok(()) - } - async fn delete(&self, k: &str) -> Result<(), CacheError> { - self.0.lock().unwrap().remove(k); - Ok(()) - } - async fn get_bytes(&self, _k: &str) -> Option> { - None - } - async fn set_bytes(&self, _k: &str, _v: &[u8], _t: Duration) -> Result<(), CacheError> { - Ok(()) - } - } - - fn cache_with(nsid: &str, scopes: &str) -> MapCache { - let c = MapCache::default(); + async fn cache_with(nsid: &str, scopes: &str) -> MemoryCache { + let c = MemoryCache::new(); let key = tranquil_pds::cache_keys::permission_set_key( &tranquil_types::Nsid::new(nsid).unwrap(), None, @@ -74,7 +50,7 @@ mod tests { "refreshed_at": chrono::Utc::now().timestamp(), }) .to_string(); - c.0.lock().unwrap().insert(key, json); + let _ = c.set(&key, &json, Duration::from_secs(3600)).await; c } @@ -83,7 +59,8 @@ mod tests { let c = cache_with( "io.atcr.authFullApp", "repo:io.atcr.manifest?action=create identity:*", - ); + ) + .await; let eff = resolve_effective_scopes( &c, "atproto include:io.atcr.authFullApp", @@ -104,7 +81,8 @@ mod tests { let c = cache_with( "io.atcr.authFullApp", "repo:io.atcr.manifest?action=create identity:*", - ); + ) + .await; let granted = DbScope::new("atproto repo:* blob:*/* account:*?action=manage").unwrap(); let eff = resolve_effective_scopes( &c, diff --git a/crates/tranquil-pds/src/auth/email_token.rs b/crates/tranquil-pds/src/auth/email_token.rs index 2441a29..c597ed1 100644 --- a/crates/tranquil-pds/src/auth/email_token.rs +++ b/crates/tranquil-pds/src/auth/email_token.rs @@ -2,31 +2,13 @@ use serde::{Deserialize, Serialize}; use std::time::Duration; use crate::cache::Cache; +use crate::cache_keys::email_token_key; use crate::types::Did; use crate::util::{generate_token_code, normalize_token_code}; -const TOKEN_TTL_SECS: u64 = 900; - -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub enum EmailTokenPurpose { - UpdateEmail, - ConfirmEmail, - DeleteAccount, - ResetPassword, - PlcOperation, -} +pub use tranquil_types::EmailTokenPurpose; -impl EmailTokenPurpose { - fn as_str(&self) -> &'static str { - match self { - Self::UpdateEmail => "update_email", - Self::ConfirmEmail => "confirm_email", - Self::DeleteAccount => "delete_account", - Self::ResetPassword => "reset_password", - Self::PlcOperation => "plc_operation", - } - } -} +const TOKEN_TTL_SECS: u64 = 900; #[derive(Debug, Clone, Serialize, Deserialize)] struct TokenData { @@ -42,10 +24,6 @@ pub enum TokenError { ExpiredToken, } -fn cache_key(did: &Did, purpose: EmailTokenPurpose) -> String { - format!("email_token:{}:{}", purpose.as_str(), did) -} - fn current_timestamp() -> u64 { u64::try_from(chrono::Utc::now().timestamp()).unwrap_or(0) } @@ -69,7 +47,7 @@ pub async fn create_email_token( cache .set( - &cache_key(did, purpose), + &email_token_key(did, purpose), &json, Duration::from_secs(TOKEN_TTL_SECS), ) @@ -89,7 +67,7 @@ pub async fn validate_email_token( return Err(TokenError::CacheUnavailable); } - let key = cache_key(did, purpose); + let key = email_token_key(did, purpose); let json = cache.get(&key).await.ok_or(TokenError::InvalidToken)?; let data: TokenData = serde_json::from_str(&json).map_err(|_| TokenError::InvalidToken)?; @@ -112,7 +90,7 @@ pub async fn validate_email_token( } pub async fn delete_email_token(cache: &dyn Cache, did: &Did, purpose: EmailTokenPurpose) { - let _ = cache.delete(&cache_key(did, purpose)).await; + let _ = cache.delete(&email_token_key(did, purpose)).await; } fn constant_time_eq(a: &[u8], b: &[u8]) -> bool { @@ -128,67 +106,11 @@ fn constant_time_eq(a: &[u8], b: &[u8]) -> bool { #[cfg(test)] mod tests { use super::*; - use crate::cache::CacheError; - use async_trait::async_trait; - use std::collections::HashMap; - use std::sync::Mutex; - - struct MockCache { - data: Mutex>, - } - - impl MockCache { - fn new() -> Self { - Self { - data: Mutex::new(HashMap::new()), - } - } - } - - #[async_trait] - impl Cache for MockCache { - async fn get(&self, key: &str) -> Option { - let data = self.data.lock().unwrap(); - let now = current_timestamp(); - data.get(key) - .filter(|(_, exp)| *exp > now) - .map(|(v, _)| v.clone()) - } - - async fn set(&self, key: &str, value: &str, ttl: Duration) -> Result<(), CacheError> { - let mut data = self.data.lock().unwrap(); - let expires = current_timestamp() + ttl.as_secs(); - data.insert(key.to_string(), (value.to_string(), expires)); - Ok(()) - } - - async fn delete(&self, key: &str) -> Result<(), CacheError> { - let mut data = self.data.lock().unwrap(); - data.remove(key); - Ok(()) - } - - async fn get_bytes(&self, _key: &str) -> Option> { - None - } - - async fn set_bytes( - &self, - _key: &str, - _value: &[u8], - _ttl: Duration, - ) -> Result<(), CacheError> { - Ok(()) - } - - fn is_available(&self) -> bool { - true - } - } + use tranquil_infra::MemoryCache; #[tokio::test] async fn test_create_and_validate_token() { - let cache = MockCache::new(); + let cache = MemoryCache::new(); let did = Did::new("did:plc:teq").expect("valid DID"); let token = create_email_token(&cache, &did, EmailTokenPurpose::UpdateEmail) @@ -205,7 +127,7 @@ mod tests { #[tokio::test] async fn test_token_consumed_after_use() { - let cache = MockCache::new(); + let cache = MemoryCache::new(); let did = Did::new("did:plc:teq").expect("valid DID"); let token = create_email_token(&cache, &did, EmailTokenPurpose::UpdateEmail) @@ -223,7 +145,7 @@ mod tests { #[tokio::test] async fn test_invalid_token_rejected() { - let cache = MockCache::new(); + let cache = MemoryCache::new(); let did = Did::new("did:plc:teq").expect("valid DID"); let _token = create_email_token(&cache, &did, EmailTokenPurpose::UpdateEmail) @@ -237,7 +159,7 @@ mod tests { #[tokio::test] async fn test_wrong_purpose_rejected() { - let cache = MockCache::new(); + let cache = MemoryCache::new(); let did = Did::new("did:plc:teq").expect("valid DID"); let token = create_email_token(&cache, &did, EmailTokenPurpose::UpdateEmail) @@ -252,7 +174,7 @@ mod tests { #[tokio::test] async fn test_token_format() { // The emitted token is the display form: uppercase `XXXXX-XXXXX`. - let cache = MockCache::new(); + let cache = MemoryCache::new(); let did = Did::new("did:plc:teq").expect("valid DID"); (0..50).for_each(|_| { let token = futures::executor::block_on(create_email_token( @@ -269,7 +191,7 @@ mod tests { #[tokio::test] async fn test_case_insensitive_validation() { - let cache = MockCache::new(); + let cache = MemoryCache::new(); let did = Did::new("did:plc:teq").expect("valid DID"); let token = create_email_token(&cache, &did, EmailTokenPurpose::UpdateEmail) @@ -284,7 +206,7 @@ mod tests { #[tokio::test] async fn test_hyphen_insensitive_validation() { - let cache = MockCache::new(); + let cache = MemoryCache::new(); let did = Did::new("did:plc:teq").expect("valid DID"); let token = create_email_token(&cache, &did, EmailTokenPurpose::UpdateEmail) diff --git a/crates/tranquil-pds/src/auth/legacy_2fa.rs b/crates/tranquil-pds/src/auth/legacy_2fa.rs index e8d3aba..b85aca6 100644 --- a/crates/tranquil-pds/src/auth/legacy_2fa.rs +++ b/crates/tranquil-pds/src/auth/legacy_2fa.rs @@ -3,6 +3,7 @@ use serde::{Deserialize, Serialize}; use std::time::Duration; use crate::cache::Cache; +use crate::cache_keys::{legacy_2fa_challenge_key, legacy_2fa_cooldown_key}; use crate::types::Did; use crate::util::{generate_token_code, normalize_token_code}; @@ -58,8 +59,8 @@ pub async fn create_challenge( } pub async fn clear_challenge(cache: &dyn Cache, did: &Did) { - let _ = cache.delete(&challenge_key(did)).await; - let _ = cache.delete(&cooldown_key(did)).await; + let _ = cache.delete(&legacy_2fa_challenge_key(did)).await; + let _ = cache.delete(&legacy_2fa_cooldown_key(did)).await; } async fn validate_challenge_internal( @@ -71,7 +72,7 @@ async fn validate_challenge_internal( return Err(ValidationError::CacheUnavailable); } - let challenge_k = challenge_key(did); + let challenge_k = legacy_2fa_challenge_key(did); let json = cache .get(&challenge_k) @@ -114,19 +115,11 @@ async fn validate_challenge_internal( } let _ = cache.delete(&challenge_k).await; - let _ = cache.delete(&cooldown_key(did)).await; + let _ = cache.delete(&legacy_2fa_cooldown_key(did)).await; Ok(()) } -fn challenge_key(did: &Did) -> String { - format!("legacy_2fa:{}", did) -} - -fn cooldown_key(did: &Did) -> String { - format!("legacy_2fa_cooldown:{}", did) -} - fn current_timestamp() -> u64 { u64::try_from(Utc::now().timestamp()).unwrap_or(0) } @@ -226,7 +219,7 @@ async fn create_challenge_code( return Err(ChallengeError::CacheUnavailable); } - let cooldown = cooldown_key(did); + let cooldown = legacy_2fa_cooldown_key(did); if cache.get(&cooldown).await.is_some() { return Err(ChallengeError::RateLimited); } @@ -244,7 +237,7 @@ async fn create_challenge_code( cache .set( - &challenge_key(did), + &legacy_2fa_challenge_key(did), &json, Duration::from_secs(CHALLENGE_TTL_SECS), ) @@ -280,67 +273,11 @@ impl From for Legacy2faFlowError { #[cfg(test)] mod tests { use super::*; - use crate::cache::CacheError; - use async_trait::async_trait; - use std::collections::HashMap; - use std::sync::Mutex; - - struct MockCache { - data: Mutex>, - } - - impl MockCache { - fn new() -> Self { - Self { - data: Mutex::new(HashMap::new()), - } - } - } - - #[async_trait] - impl Cache for MockCache { - async fn get(&self, key: &str) -> Option { - let data = self.data.lock().unwrap(); - let now = current_timestamp(); - data.get(key) - .filter(|(_, exp)| *exp > now) - .map(|(v, _)| v.clone()) - } - - async fn set(&self, key: &str, value: &str, ttl: Duration) -> Result<(), CacheError> { - let mut data = self.data.lock().unwrap(); - let expires = current_timestamp() + ttl.as_secs(); - data.insert(key.to_string(), (value.to_string(), expires)); - Ok(()) - } - - async fn delete(&self, key: &str) -> Result<(), CacheError> { - let mut data = self.data.lock().unwrap(); - data.remove(key); - Ok(()) - } - - async fn get_bytes(&self, _key: &str) -> Option> { - None - } - - async fn set_bytes( - &self, - _key: &str, - _value: &[u8], - _ttl: Duration, - ) -> Result<(), CacheError> { - Ok(()) - } - - fn is_available(&self) -> bool { - true - } - } + use tranquil_infra::MemoryCache; #[tokio::test] async fn test_create_and_validate_challenge() { - let cache = MockCache::new(); + let cache = MemoryCache::new(); let did = Did::new("did:plc:test123".to_string()).unwrap(); let code = create_challenge(&cache, &did).await.unwrap(); @@ -352,7 +289,7 @@ mod tests { #[tokio::test] async fn test_challenge_code_format() { - let cache = MockCache::new(); + let cache = MemoryCache::new(); let did = Did::new("did:plc:test123".to_string()).unwrap(); let code = create_challenge(&cache, &did).await.unwrap(); @@ -364,7 +301,7 @@ mod tests { #[tokio::test] async fn test_case_insensitive_validation() { - let cache = MockCache::new(); + let cache = MemoryCache::new(); let did = Did::new("did:plc:test123".to_string()).unwrap(); let code = create_challenge(&cache, &did).await.unwrap(); @@ -375,7 +312,7 @@ mod tests { #[tokio::test] async fn test_hyphen_insensitive_validation() { - let cache = MockCache::new(); + let cache = MemoryCache::new(); let did = Did::new("did:plc:test123".to_string()).unwrap(); let code = create_challenge(&cache, &did).await.unwrap(); @@ -386,7 +323,7 @@ mod tests { #[tokio::test] async fn test_invalid_code_rejected() { - let cache = MockCache::new(); + let cache = MemoryCache::new(); let did = Did::new("did:plc:test123".to_string()).unwrap(); let _code = create_challenge(&cache, &did).await.unwrap(); @@ -396,7 +333,7 @@ mod tests { #[tokio::test] async fn test_challenge_consumed_on_success() { - let cache = MockCache::new(); + let cache = MemoryCache::new(); let did = Did::new("did:plc:test123".to_string()).unwrap(); let code = create_challenge(&cache, &did).await.unwrap(); @@ -410,7 +347,7 @@ mod tests { #[tokio::test] async fn test_max_attempts_exceeded() { - let cache = MockCache::new(); + let cache = MemoryCache::new(); let did = Did::new("did:plc:test123".to_string()).unwrap(); let _code = create_challenge(&cache, &did).await.unwrap(); @@ -425,7 +362,7 @@ mod tests { #[tokio::test] async fn test_rate_limiting() { - let cache = MockCache::new(); + let cache = MemoryCache::new(); let did = Did::new("did:plc:test123".to_string()).unwrap(); let _first = create_challenge(&cache, &did).await.unwrap(); @@ -453,7 +390,7 @@ mod tests { #[tokio::test] async fn test_process_flow_not_required() { - let cache = MockCache::new(); + let cache = MemoryCache::new(); let did = Did::new("did:plc:test".to_string()).unwrap(); let ctx = Legacy2faContext { is_app_password: false, @@ -470,7 +407,7 @@ mod tests { #[tokio::test] async fn test_process_flow_not_required_because_app_password() { - let cache = MockCache::new(); + let cache = MemoryCache::new(); let did = Did::new("did:plc:test".to_string()).unwrap(); let ctx = Legacy2faContext { is_app_password: true, @@ -487,7 +424,7 @@ mod tests { #[tokio::test] async fn test_process_flow_blocked() { - let cache = MockCache::new(); + let cache = MemoryCache::new(); let did = Did::new("did:plc:test".to_string()).unwrap(); let ctx = Legacy2faContext { is_app_password: false, @@ -504,7 +441,7 @@ mod tests { #[tokio::test] async fn test_process_flow_challenge_sent_totp() { - let cache = MockCache::new(); + let cache = MemoryCache::new(); let did = Did::new("did:plc:test".to_string()).unwrap(); let ctx = Legacy2faContext { is_app_password: false, @@ -521,7 +458,7 @@ mod tests { #[tokio::test] async fn test_process_flow_challenge_sent_email_2fa_enabled() { - let cache = MockCache::new(); + let cache = MemoryCache::new(); let did = Did::new("did:plc:test2".to_string()).unwrap(); let ctx = Legacy2faContext { is_app_password: false, @@ -538,7 +475,7 @@ mod tests { #[tokio::test] async fn test_process_flow_verified() { - let cache = MockCache::new(); + let cache = MemoryCache::new(); let did = Did::new("did:plc:test".to_string()).unwrap(); let ctx = Legacy2faContext { is_app_password: false, @@ -557,7 +494,7 @@ mod tests { #[tokio::test] async fn test_attempts_persist_across_failures() { - let cache = MockCache::new(); + let cache = MemoryCache::new(); let did = Did::new("did:plc:test123".to_string()).unwrap(); let code = create_challenge(&cache, &did).await.unwrap(); @@ -590,7 +527,7 @@ mod tests { #[tokio::test] async fn test_totp_shaped_token_accepted_via_verifier() { - let cache = MockCache::new(); + let cache = MemoryCache::new(); let did = Did::new("did:plc:totp1".to_string()).unwrap(); let ctx = Legacy2faContext { is_app_password: false, @@ -607,7 +544,7 @@ mod tests { #[tokio::test] async fn test_totp_shaped_token_rejected_does_not_touch_email_challenge() { - let cache = MockCache::new(); + let cache = MemoryCache::new(); let did = Did::new("did:plc:totp2".to_string()).unwrap(); let ctx = Legacy2faContext { is_app_password: false, @@ -641,7 +578,7 @@ mod tests { #[tokio::test] async fn test_email_shaped_token_routes_to_email_path_when_totp_present() { - let cache = MockCache::new(); + let cache = MemoryCache::new(); let did = Did::new("did:plc:totp3".to_string()).unwrap(); let ctx = Legacy2faContext { is_app_password: false, @@ -662,7 +599,7 @@ mod tests { #[tokio::test] async fn test_backup_code_shaped_token_routes_to_verifier() { - let cache = MockCache::new(); + let cache = MemoryCache::new(); let did = Did::new("did:plc:totp4".to_string()).unwrap(); let ctx = Legacy2faContext { is_app_password: false, @@ -681,7 +618,7 @@ mod tests { #[tokio::test] async fn test_totp_shaped_token_ignored_when_no_totp() { - let cache = MockCache::new(); + let cache = MemoryCache::new(); let did = Did::new("did:plc:totp5".to_string()).unwrap(); let ctx = Legacy2faContext { is_app_password: false, diff --git a/crates/tranquil-pds/src/oauth/permission_set_resolver.rs b/crates/tranquil-pds/src/oauth/permission_set_resolver.rs index a9420bc..defe864 100644 --- a/crates/tranquil-pds/src/oauth/permission_set_resolver.rs +++ b/crates/tranquil-pds/src/oauth/permission_set_resolver.rs @@ -137,39 +137,12 @@ fn map_err(e: &ScopeExpansionError) -> ResolveFailure { #[cfg(test)] mod tests { use super::*; - use crate::cache::{Cache, CacheError}; - use std::collections::HashMap; - use std::sync::Mutex; use std::time::Duration; + use tranquil_infra::MemoryCache; - #[derive(Default)] - struct MapCache(Mutex>); + const SEED_TTL: Duration = Duration::from_secs(3600); - #[async_trait::async_trait] - impl Cache for MapCache { - async fn get(&self, key: &str) -> Option { - self.0.lock().unwrap().get(key).cloned() - } - async fn set(&self, key: &str, value: &str, _ttl: Duration) -> Result<(), CacheError> { - self.0 - .lock() - .unwrap() - .insert(key.to_string(), value.to_string()); - Ok(()) - } - async fn delete(&self, key: &str) -> Result<(), CacheError> { - self.0.lock().unwrap().remove(key); - Ok(()) - } - async fn get_bytes(&self, _key: &str) -> Option> { - None - } - async fn set_bytes(&self, _k: &str, _v: &[u8], _t: Duration) -> Result<(), CacheError> { - Ok(()) - } - } - - fn seed_at(cache: &MapCache, nsid: &str, scope: &str, refreshed_at: i64) { + async fn seed_at(cache: &MemoryCache, nsid: &str, scope: &str, refreshed_at: i64) { let key = crate::cache_keys::permission_set_key(&tranquil_types::Nsid::new(nsid).unwrap(), None); let val = serde_json::to_string(&CachedPermissionSet { @@ -179,21 +152,22 @@ mod tests { refreshed_at, }) .unwrap(); - cache.0.lock().unwrap().insert(key, val); + let _ = cache.set(&key, &val, SEED_TTL).await; } - fn seed(cache: &MapCache, nsid: &str, scope: &str) { - seed_at(cache, nsid, scope, now_secs()); + async fn seed(cache: &MemoryCache, nsid: &str, scope: &str) { + seed_at(cache, nsid, scope, now_secs()).await; } #[tokio::test] async fn cache_hit_expands_without_network() { - let cache = MapCache::default(); + let cache = MemoryCache::new(); seed( &cache, "io.atcr.authFullApp", "repo:io.atcr.manifest?action=create identity:*", - ); + ) + .await; let out = expand_scopes(&cache, "atproto include:io.atcr.authFullApp").await; assert!(out.failures.is_empty()); assert_eq!(out.passthrough, vec!["atproto".to_string()]); @@ -208,13 +182,14 @@ mod tests { #[tokio::test] async fn stale_entry_is_served_when_refresh_fails() { - let cache = MapCache::default(); + let cache = MemoryCache::new(); seed_at( &cache, "nonexistent.fake.permissionSet", "repo:nonexistent.fake.record?action=create", now_secs() - STALE_AFTER_SECS - 1, - ); + ) + .await; let out = expand_scopes(&cache, "include:nonexistent.fake.permissionSet").await; assert!( out.failures.is_empty(), @@ -230,7 +205,7 @@ mod tests { #[tokio::test] async fn entry_without_refreshed_at_is_treated_as_stale_but_usable() { - let cache = MapCache::default(); + let cache = MemoryCache::new(); let key = crate::cache_keys::permission_set_key( &tranquil_types::Nsid::new("nonexistent.fake.permissionSet").unwrap(), None, @@ -238,7 +213,7 @@ mod tests { // Shape written before `refreshed_at` existed. let legacy = r#"{"scope":"repo:nonexistent.fake.record?action=create","title":null,"detail":null}"#; - cache.0.lock().unwrap().insert(key, legacy.to_string()); + let _ = cache.set(&key, legacy, SEED_TTL).await; let out = expand_scopes(&cache, "include:nonexistent.fake.permissionSet").await; assert!(out.failures.is_empty()); assert_eq!(out.sets.len(), 1); @@ -246,7 +221,7 @@ mod tests { #[tokio::test] async fn passthrough_scopes_untouched() { - let cache = MapCache::default(); + let cache = MemoryCache::new(); let out = expand_scopes(&cache, "atproto repo:app.bsky.feed.post?action=create").await; assert!(out.failures.is_empty()); assert!(out.sets.is_empty()); @@ -255,7 +230,7 @@ mod tests { #[tokio::test] async fn cache_miss_unresolvable_is_a_failure() { - let cache = MapCache::default(); + let cache = MemoryCache::new(); let out = expand_scopes(&cache, "include:nonexistent.fake.permissionSet").await; assert_eq!(out.sets.len(), 0); assert_eq!(out.failures.len(), 1); -- 2.51.2 From 52d5236e892834cdfdb4ef8ba31f4567d8ee9725 Mon Sep 17 00:00:00 2001 From: Lewis Date: Sun, 26 Jul 2026 13:57:07 +0300 Subject: [PATCH 05/22] plc: dedup fetch paths, cache TTL from config Lewis: May this revision serve well! --- crates/tranquil-pds/src/plc/mod.rs | 135 +++++++++-------------------- 1 file changed, 42 insertions(+), 93 deletions(-) diff --git a/crates/tranquil-pds/src/plc/mod.rs b/crates/tranquil-pds/src/plc/mod.rs index 4561e00..ea723f1 100644 --- a/crates/tranquil-pds/src/plc/mod.rs +++ b/crates/tranquil-pds/src/plc/mod.rs @@ -165,12 +165,11 @@ impl PlcOpOrTombstone { } } -const PLC_CACHE_TTL_SECS: u64 = 300; - pub struct PlcClient { base_url: String, client: Client, cache: Option>, + cache_ttl: Duration, } impl PlcClient { @@ -193,12 +192,19 @@ impl PlcClient { .connect_timeout(Duration::from_secs(connect_timeout_secs)) .pool_max_idle_per_host(5) .pool_idle_timeout(Duration::from_secs(90)) + .redirect(tranquil_types::redirect_policy( + tranquil_types::ReachPolicy::DEBUG_LOOPBACK, + )) + .dns_resolver(tranquil_types::dns_guard( + tranquil_types::ReachPolicy::DEBUG_LOOPBACK, + )) .build() - .unwrap_or_else(|_| Client::new()); + .expect("failed to build PLC directory HTTP client"); Self { base_url, client, cache, + cache_ttl: Duration::from_secs(cfg.map_or(300, |c| c.plc.did_cache_ttl_secs)), } } @@ -206,15 +212,7 @@ impl PlcClient { urlencoding::encode(did.as_str()).to_string() } - pub async fn get_document(&self, did: &Did) -> Result { - let cache_key = crate::cache_keys::plc_doc_key(did); - if let Some(ref cache) = self.cache - && let Some(cached) = cache.get(&cache_key).await - && let Ok(value) = serde_json::from_str(&cached) - { - return Ok(value); - } - let url = format!("{}/{}", self.base_url, Self::encode_did(did)); + async fn fetch_json(&self, url: String) -> Result { let response = self.client.get(&url).send().await?; if response.status() == reqwest::StatusCode::NOT_FOUND { return Err(PlcError::NotFound); @@ -227,101 +225,52 @@ impl PlcClient { status, body ))); } - let value: Value = response + response .json() .await - .map_err(|e| PlcError::InvalidResponse(e.to_string()))?; - if let Some(ref cache) = self.cache - && let Ok(json_str) = serde_json::to_string(&value) - { - let _ = cache - .set( - &cache_key, - &json_str, - Duration::from_secs(PLC_CACHE_TTL_SECS), - ) - .await; + .map_err(|e| PlcError::InvalidResponse(e.to_string())) + } + + async fn cached_fetch(&self, cache_key: &str, url: String) -> Result { + match &self.cache { + Some(cache) => { + crate::cache::cached_json(cache.as_ref(), cache_key, self.cache_ttl, || { + self.fetch_json(url) + }) + .await + } + None => self.fetch_json(url).await, } - Ok(value) + } + + pub async fn get_document(&self, did: &Did) -> Result { + let url = format!("{}/{}", self.base_url, Self::encode_did(did)); + self.cached_fetch(&crate::cache_keys::plc_doc_key(did), url) + .await } pub async fn get_document_data(&self, did: &Did) -> Result { - let cache_key = crate::cache_keys::plc_data_key(did); - if let Some(ref cache) = self.cache - && let Some(cached) = cache.get(&cache_key).await - && let Ok(value) = serde_json::from_str(&cached) - { - return Ok(value); - } let url = format!("{}/{}/data", self.base_url, Self::encode_did(did)); - let response = self.client.get(&url).send().await?; - if response.status() == reqwest::StatusCode::NOT_FOUND { - return Err(PlcError::NotFound); - } - if !response.status().is_success() { - let status = response.status(); - let body = response.text().await.unwrap_or_default(); - return Err(PlcError::InvalidResponse(format!( - "HTTP {}: {}", - status, body - ))); - } - let value: Value = response - .json() + self.cached_fetch(&crate::cache_keys::plc_data_key(did), url) .await - .map_err(|e| PlcError::InvalidResponse(e.to_string()))?; - if let Some(ref cache) = self.cache - && let Ok(json_str) = serde_json::to_string(&value) - { - let _ = cache - .set( - &cache_key, - &json_str, - Duration::from_secs(PLC_CACHE_TTL_SECS), - ) - .await; - } - Ok(value) } pub async fn get_last_op(&self, did: &Did) -> Result { - let url = format!("{}/{}/log/last", self.base_url, Self::encode_did(did)); - let response = self.client.get(&url).send().await?; - if response.status() == reqwest::StatusCode::NOT_FOUND { - return Err(PlcError::NotFound); - } - if !response.status().is_success() { - let status = response.status(); - let body = response.text().await.unwrap_or_default(); - return Err(PlcError::InvalidResponse(format!( - "HTTP {}: {}", - status, body - ))); - } - response - .json() - .await - .map_err(|e| PlcError::InvalidResponse(e.to_string())) + self.fetch_json(format!( + "{}/{}/log/last", + self.base_url, + Self::encode_did(did) + )) + .await } pub async fn get_audit_log(&self, did: &Did) -> Result, PlcError> { - let url = format!("{}/{}/log/audit", self.base_url, Self::encode_did(did)); - let response = self.client.get(&url).send().await?; - if response.status() == reqwest::StatusCode::NOT_FOUND { - return Err(PlcError::NotFound); - } - if !response.status().is_success() { - let status = response.status(); - let body = response.text().await.unwrap_or_default(); - return Err(PlcError::InvalidResponse(format!( - "HTTP {}: {}", - status, body - ))); - } - response - .json() - .await - .map_err(|e| PlcError::InvalidResponse(e.to_string())) + self.fetch_json(format!( + "{}/{}/log/audit", + self.base_url, + Self::encode_did(did) + )) + .await } pub async fn send_operation(&self, did: &Did, operation: &Value) -> Result<(), PlcError> { -- 2.51.2 From 0fc577316ea9870f06a4b004874d7fb32b787351 Mon Sep 17 00:00:00 2001 From: Lewis Date: Sun, 26 Jul 2026 13:57:07 +0300 Subject: [PATCH 06/22] lexicon: schema docs & negative results via cluster cache Lewis: May this revision serve well! --- crates/tranquil-lexicon/Cargo.toml | 7 +- crates/tranquil-lexicon/src/dynamic.rs | 246 +++++++++++++++--- crates/tranquil-lexicon/src/registry.rs | 5 + crates/tranquil-lexicon/src/resolve.rs | 177 +++++++------ crates/tranquil-lexicon/src/schema.rs | 30 +-- .../tests/resolve_integration.rs | 17 +- 6 files changed, 345 insertions(+), 137 deletions(-) diff --git a/crates/tranquil-lexicon/Cargo.toml b/crates/tranquil-lexicon/Cargo.toml index 99e7eb0..2974aac 100644 --- a/crates/tranquil-lexicon/Cargo.toml +++ b/crates/tranquil-lexicon/Cargo.toml @@ -5,10 +5,11 @@ edition.workspace = true license.workspace = true [features] -resolve = ["dep:reqwest", "dep:hickory-resolver", "dep:tokio", "dep:parking_lot", "dep:tracing", "dep:urlencoding"] +resolve = ["dep:reqwest", "dep:hickory-resolver", "dep:tokio", "dep:parking_lot", "dep:tracing", "dep:tranquil-infra"] [dependencies] -tranquil-types = { path = "../tranquil-types", default-features = false } +tranquil-types = { workspace = true } +tranquil-infra = { workspace = true, optional = true, features = ["cache-keys"] } serde = { workspace = true } serde_json = { workspace = true } thiserror = { workspace = true } @@ -19,9 +20,9 @@ hickory-resolver = { workspace = true, optional = true } tokio = { workspace = true, optional = true } parking_lot = { workspace = true, optional = true } tracing = { workspace = true, optional = true } -urlencoding = { workspace = true, optional = true } [dev-dependencies] wiremock = { workspace = true } tokio = { workspace = true } futures = { workspace = true } +tranquil-infra = { workspace = true, features = ["testing", "cache-keys"] } diff --git a/crates/tranquil-lexicon/src/dynamic.rs b/crates/tranquil-lexicon/src/dynamic.rs index 96581e2..5d5db99 100644 --- a/crates/tranquil-lexicon/src/dynamic.rs +++ b/crates/tranquil-lexicon/src/dynamic.rs @@ -6,9 +6,11 @@ use std::sync::Arc; use std::sync::atomic::{AtomicBool, Ordering}; use std::time::{Duration, Instant}; use tokio::sync::Notify; +use tranquil_infra::cache_keys::{lexicon_doc_key, lexicon_negative_key}; +use tranquil_infra::{Cache, read_json, write_json}; use tranquil_types::Nsid; -const NEGATIVE_CACHE_TTL: Duration = Duration::from_secs(24 * 60 * 60); +const NEGATIVE_CACHE_TTL: Duration = Duration::from_secs(60 * 60); const POSITIVE_CACHE_TTL: Duration = Duration::from_secs(24 * 60 * 60); const REFRESH_FAILURE_BACKOFF: Duration = Duration::from_secs(60); const MAX_DYNAMIC_SCHEMAS: usize = 1024; @@ -17,6 +19,13 @@ struct NegativeEntry { expires_at: Instant, } +fn negative_ttl_for(error: &ResolveError) -> Duration { + match error.is_definitive() { + true => NEGATIVE_CACHE_TTL, + false => REFRESH_FAILURE_BACKOFF, + } +} + struct PositiveEntry { doc: Arc, expires_at: Instant, @@ -44,6 +53,7 @@ pub struct DynamicRegistry { negative_cache: RwLock>, in_flight: RwLock>>, network_disabled: AtomicBool, + shared: RwLock>>, } struct InFlightGuard<'a> { @@ -70,9 +80,18 @@ impl DynamicRegistry { negative_cache: RwLock::new(HashMap::new()), in_flight: RwLock::new(HashMap::new()), network_disabled: AtomicBool::new(false), + shared: RwLock::new(None), } } + pub fn set_shared_cache(&self, cache: Arc) { + *self.shared.write() = Some(cache); + } + + fn shared_cache(&self) -> Option> { + self.shared.read().clone() + } + pub fn from_env() -> Self { let registry = Self::new(); let disabled = @@ -105,13 +124,17 @@ impl DynamicRegistry { } pub fn is_negative_cached(&self, nsid: &Nsid) -> bool { - let cache = self.negative_cache.read(); - cache + self.negative_remaining(nsid).is_some() + } + + fn negative_remaining(&self, nsid: &Nsid) -> Option { + self.negative_cache + .read() .get(nsid) - .is_some_and(|entry| entry.expires_at > Instant::now()) + .and_then(|entry| entry.expires_at.checked_duration_since(Instant::now())) } - fn insert_negative(&self, nsid: &Nsid) { + fn insert_negative(&self, nsid: &Nsid, ttl: Duration) { let mut cache = self.negative_cache.write(); if cache.len() >= MAX_DYNAMIC_SCHEMAS { let now = Instant::now(); @@ -120,7 +143,7 @@ impl DynamicRegistry { cache.insert( nsid.clone(), NegativeEntry { - expires_at: Instant::now() + NEGATIVE_CACHE_TTL, + expires_at: Instant::now() + ttl, }, ); } @@ -159,6 +182,44 @@ impl DynamicRegistry { arc } + async fn shared_get(&self, nsid: &Nsid) -> Option> { + let cache = self.shared_cache()?; + let doc = read_json::(cache.as_ref(), &lexicon_doc_key(nsid)).await?; + Some(self.insert_schema(doc)) + } + + async fn shared_put(&self, doc: &LexiconDoc) { + let Some(cache) = self.shared_cache() else { + return; + }; + write_json( + cache.as_ref(), + &lexicon_doc_key(&doc.id), + doc, + POSITIVE_CACHE_TTL, + ) + .await; + let _ = cache.delete(&lexicon_negative_key(&doc.id)).await; + } + + async fn shared_is_negative(&self, nsid: &Nsid) -> bool { + match self.shared_cache() { + Some(cache) => cache.get(&lexicon_negative_key(nsid)).await.is_some(), + None => false, + } + } + + async fn shared_put_negative(&self, nsid: &Nsid, error: &ResolveError) { + if !error.is_definitive() { + return; + } + if let Some(cache) = self.shared_cache() { + let _ = cache + .set(&lexicon_negative_key(nsid), "1", NEGATIVE_CACHE_TTL) + .await; + } + } + fn bump_expiry(&self, nsid: &Nsid, duration: Duration) { let mut store = self.store.write(); if let Some(entry) = store.schemas.get_mut(nsid) { @@ -203,15 +264,23 @@ impl DynamicRegistry { match self.acquire_leadership(nsid) { Some(_guard) => match resolver(nsid.clone()).await { - Ok(doc) => Ok(self.insert_schema(doc)), + Ok(doc) => { + self.shared_put(&doc).await; + Ok(self.insert_schema(doc)) + } Err(e) => { + let (doc, source) = match self.shared_get(nsid).await { + Some(doc) => (doc, "shared"), + None => (stale, "local"), + }; self.bump_expiry(nsid, REFRESH_FAILURE_BACKOFF); tracing::warn!( nsid = %nsid, error = %e, - "lexicon refresh failed, serving stale cached entry" + source, + "lexicon refresh failed, serving cached entry" ); - Ok(stale) + Ok(doc) } }, None => { @@ -230,34 +299,59 @@ impl DynamicRegistry { F: FnOnce(Nsid) -> Fut, Fut: std::future::Future>, { - if self.network_disabled.load(Ordering::Relaxed) { - return Err(ResolveError::NetworkDisabled); + if let Some(doc) = self.shared_get(nsid).await { + return Ok(doc); } - if self.is_negative_cached(nsid) { + + if let Some(remaining) = self.negative_remaining(nsid) { + return Err(ResolveError::NegativelyCached { + nsid: nsid.clone(), + ttl_secs: remaining.as_secs(), + }); + } + + if self.shared_is_negative(nsid).await { + // Cache reports 0 remaining TTL for shared negative hit, + // so we mirror for the backoff rather than a full `NEGATIVE_CACHE_TTL`. + self.insert_negative(nsid, REFRESH_FAILURE_BACKOFF); return Err(ResolveError::NegativelyCached { nsid: nsid.clone(), - ttl_secs: NEGATIVE_CACHE_TTL.as_secs(), + ttl_secs: REFRESH_FAILURE_BACKOFF.as_secs(), }); } + if self.network_disabled.load(Ordering::Relaxed) { + return Err(ResolveError::NetworkDisabled); + } + match self.acquire_leadership(nsid) { Some(_guard) => match resolver(nsid.clone()).await { - Ok(doc) => Ok(self.insert_schema(doc)), + Ok(doc) => { + self.shared_put(&doc).await; + Ok(self.insert_schema(doc)) + } Err(e) => { - self.insert_negative(nsid); - tracing::debug!(nsid = %nsid, error = %e, "caching negative resolution result"); + let ttl = negative_ttl_for(&e); + self.insert_negative(nsid, ttl); + self.shared_put_negative(nsid, &e).await; + tracing::debug!( + nsid = %nsid, + error = %e, + ttl_secs = ttl.as_secs(), + "caching negative resolution result" + ); Err(e) } }, None => { self.wait_for_leader(nsid).await; - match self.get_cached(nsid) { - Some(doc) => Ok(doc), - None if self.is_negative_cached(nsid) => Err(ResolveError::NegativelyCached { + match (self.get_cached(nsid), self.negative_remaining(nsid)) { + (Some(doc), _) => Ok(doc), + (None, Some(remaining)) => Err(ResolveError::NegativelyCached { nsid: nsid.clone(), - ttl_secs: NEGATIVE_CACHE_TTL.as_secs(), + ttl_secs: remaining.as_secs(), }), - None => Err(ResolveError::LeaderAborted { nsid: nsid.clone() }), + (None, None) => Err(ResolveError::LeaderAborted { nsid: nsid.clone() }), } } } @@ -316,6 +410,7 @@ impl Default for DynamicRegistry { #[cfg(test)] mod tests { use super::*; + use tranquil_infra::MemoryCache; fn nsid(s: &str) -> Nsid { s.parse().unwrap() @@ -324,19 +419,19 @@ mod tests { #[test] fn test_negative_cache() { let registry = DynamicRegistry::new(); - assert!(!registry.is_negative_cached(&nsid("com.example.test"))); + assert!(!registry.is_negative_cached(&nsid("pet.nel.negative"))); - registry.insert_negative(&nsid("com.example.test")); - assert!(registry.is_negative_cached(&nsid("com.example.test"))); + registry.insert_negative(&nsid("pet.nel.negative"), NEGATIVE_CACHE_TTL); + assert!(registry.is_negative_cached(&nsid("pet.nel.negative"))); } #[tokio::test] async fn test_negative_cache_returns_appropriate_error_variant() { let registry = DynamicRegistry::new(); - registry.insert_negative(&nsid("com.example.cached")); + registry.insert_negative(&nsid("pet.nel.cached"), NEGATIVE_CACHE_TTL); let err = registry - .resolve_and_cache(&nsid("com.example.cached")) + .resolve_and_cache(&nsid("pet.nel.cached")) .await .unwrap_err(); @@ -383,17 +478,17 @@ mod tests { fn test_negative_cache_cleared_on_insert() { let registry = DynamicRegistry::new(); - registry.insert_negative(&nsid("com.example.test")); - assert!(registry.is_negative_cached(&nsid("com.example.test"))); + registry.insert_negative(&nsid("pet.nel.cleared"), NEGATIVE_CACHE_TTL); + assert!(registry.is_negative_cached(&nsid("pet.nel.cleared"))); let doc = LexiconDoc { lexicon: 1, - id: nsid("com.example.test"), + id: nsid("pet.nel.cleared"), defs: HashMap::new(), }; registry.insert_schema(doc); - assert!(!registry.is_negative_cached(&nsid("com.example.test"))); + assert!(!registry.is_negative_cached(&nsid("pet.nel.cleared"))); } #[test] @@ -692,4 +787,95 @@ mod tests { "evicted Arc should be freed when no external references remain" ); } + + #[tokio::test] + async fn test_shared_positive_hit_skips_resolver() { + let registry = DynamicRegistry::new(); + let cache = Arc::new(MemoryCache::new()); + registry.set_shared_cache(cache.clone()); + let doc = LexiconDoc { + lexicon: 1, + id: nsid("pet.nel.sharedDoc"), + defs: HashMap::new(), + }; + cache + .set( + &lexicon_doc_key(&nsid("pet.nel.sharedDoc")), + &serde_json::to_string(&doc).unwrap(), + POSITIVE_CACHE_TTL, + ) + .await + .unwrap(); + + let resolved = registry + .resolve_and_cache_with(&nsid("pet.nel.sharedDoc"), |_| async move { + panic!("resolver mustn't run on a shared positive hit") + }) + .await + .unwrap(); + + assert_eq!(resolved.id, "pet.nel.sharedDoc"); + assert!(registry.get_cached(&nsid("pet.nel.sharedDoc")).is_some()); + } + + #[tokio::test] + async fn test_definitive_failure_writes_shared_negative_and_peers_mirror_it() { + let cache = Arc::new(MemoryCache::new()); + let registry = DynamicRegistry::new(); + registry.set_shared_cache(cache.clone()); + + let _ = registry + .resolve_and_cache_with(&nsid("pet.nel.gone"), |n| async move { + Err::(ResolveError::SchemaNotFound { + nsid: n, + url: "https://oyster.cafe".to_string(), + }) + }) + .await; + assert!( + cache + .get(&lexicon_negative_key(&nsid("pet.nel.gone"))) + .await + .is_some(), + "definitive failure must write the shared negative key" + ); + + let _ = registry + .resolve_and_cache_with(&nsid("pet.nel.transient"), |n| async move { + Err::(ResolveError::DnsLookup { + domain: n.into_inner(), + reason: "simulated".to_string(), + }) + }) + .await; + assert!( + cache + .get(&lexicon_negative_key(&nsid("pet.nel.transient"))) + .await + .is_none(), + "transient failure must stay out of the shared negative key" + ); + + let peer = DynamicRegistry::new(); + peer.set_shared_cache(cache); + let err = peer + .resolve_and_cache_with(&nsid("pet.nel.gone"), |_| async move { + panic!("resolver mustn't run on a shared negative hit") + }) + .await + .unwrap_err(); + match err { + ResolveError::NegativelyCached { ttl_secs, .. } => assert!( + ttl_secs <= REFRESH_FAILURE_BACKOFF.as_secs(), + "local mirror must use the backoff TTL, got {}s", + ttl_secs + ), + other => panic!("expected NegativelyCached, got: {}", other), + } + assert!( + peer.negative_remaining(&nsid("pet.nel.gone")) + .expect("local mirror exists") + <= REFRESH_FAILURE_BACKOFF + ); + } } diff --git a/crates/tranquil-lexicon/src/registry.rs b/crates/tranquil-lexicon/src/registry.rs index 2fa209d..65a14f8 100644 --- a/crates/tranquil-lexicon/src/registry.rs +++ b/crates/tranquil-lexicon/src/registry.rs @@ -125,6 +125,11 @@ impl LexiconRegistry { pub fn is_negative_cached(&self, nsid: &Nsid) -> bool { self.dynamic.is_negative_cached(nsid) } + + #[cfg(feature = "resolve")] + pub fn set_shared_cache(&self, cache: Arc) { + self.dynamic.set_shared_cache(cache); + } } pub struct ResolvedRef { diff --git a/crates/tranquil-lexicon/src/resolve.rs b/crates/tranquil-lexicon/src/resolve.rs index 4662223..c66f69b 100644 --- a/crates/tranquil-lexicon/src/resolve.rs +++ b/crates/tranquil-lexicon/src/resolve.rs @@ -4,7 +4,10 @@ use hickory_resolver::config::{ResolverConfig, ResolverOpts}; use reqwest::Client; use std::sync::OnceLock; use std::time::Duration; -use tranquil_types::{Did, Nsid}; +use tranquil_types::did_doc::extract_pds_endpoint; +use tranquil_types::{ + Did, Nsid, SchemaHostUrl, UrlKind, dns_guard, redirect_policy, url_kind, url_reach_permits, +}; static RESOLVER_CLIENT: OnceLock = OnceLock::new(); @@ -17,7 +20,8 @@ fn client() -> &'static Client { .connect_timeout(Duration::from_secs(5)) .pool_max_idle_per_host(4) .pool_idle_timeout(Duration::from_secs(60)) - .redirect(reqwest::redirect::Policy::limited(3)) + .redirect(redirect_policy(url_kind::SchemaHost::REACH_POLICY)) + .dns_resolver(dns_guard(url_kind::SchemaHost::REACH_POLICY)) .build() .expect("failed to build lexicon resolver HTTP client") }) @@ -63,6 +67,8 @@ pub enum ResolveError { NoPdsEndpoint { did: Did }, #[error("schema fetch failed from {url}: {reason}")] SchemaFetch { url: String, reason: String }, + #[error("no schema record for {nsid} at {url}")] + SchemaNotFound { nsid: Nsid, url: String }, #[error("schema deserialization failed: {0}")] InvalidSchema(String), #[error("schema resolution recently failed for {nsid}, cached for {ttl_secs}s")] @@ -73,6 +79,23 @@ pub enum ResolveError { LeaderAborted { nsid: Nsid }, } +impl ResolveError { + pub fn is_definitive(&self) -> bool { + match self { + Self::NoDid { .. } + | Self::NoPdsEndpoint { .. } + | Self::InvalidSchema(_) + | Self::SchemaNotFound { .. } => true, + Self::DnsLookup { .. } + | Self::DidResolution { .. } + | Self::SchemaFetch { .. } + | Self::NegativelyCached { .. } + | Self::NetworkDisabled + | Self::LeaderAborted { .. } => false, + } + } +} + pub fn nsid_to_authority(nsid: &Nsid) -> String { let mut segments: Vec<&str> = nsid.split('.').collect(); segments.pop(); @@ -123,7 +146,7 @@ pub async fn resolve_did_from_dns(authority: &str) -> Result pub async fn resolve_pds_endpoint( did: &Did, plc_directory_url: Option<&str>, -) -> Result { +) -> Result { let plc_base = plc_directory_url.unwrap_or(DEFAULT_PLC_DIRECTORY); let url = match did @@ -131,7 +154,20 @@ pub async fn resolve_pds_endpoint( .and_then(|(_, rest)| rest.split_once(':')) { Some(("plc", _)) => format!("{}/{}", plc_base.trim_end_matches('/'), did), - Some(("web", domain)) => format!("https://{}/.well-known/did.json", domain), + Some(("web", domain)) => { + let url = format!("https://{}/.well-known/did.json", domain); + let permitted = reqwest::Url::parse(&url) + .is_ok_and(|u| url_reach_permits(&u, url_kind::SchemaHost::REACH_POLICY)); + match permitted { + true => url, + false => { + return Err(ResolveError::DidResolution { + did: did.clone(), + reason: "did:web host is outside the allowed host reach".to_string(), + }); + } + } + } _ => { return Err(ResolveError::DidResolution { did: did.clone(), @@ -162,39 +198,29 @@ pub async fn resolve_pds_endpoint( reason: e.to_string(), })?; - extract_pds_endpoint(&doc).ok_or_else(|| ResolveError::NoPdsEndpoint { did: did.clone() }) + extract_pds_endpoint(&doc).map_err(|_| ResolveError::NoPdsEndpoint { did: did.clone() }) } -fn extract_pds_endpoint(doc: &serde_json::Value) -> Option { - doc.get("service") - .and_then(|s| s.as_array()) - .and_then(|services| { - services.iter().find_map(|svc| { - let is_pds = svc - .get("type") - .and_then(|t| t.as_str()) - .is_some_and(|t| t == "AtprotoPersonalDataServer"); - is_pds - .then(|| svc.get("serviceEndpoint").and_then(|ep| ep.as_str()))? - .map(|s| s.to_string()) - }) - }) +fn is_record_absent(xrpc_error: &str, xrpc_message: &str) -> bool { + xrpc_error == "RecordNotFound" + || xrpc_error == "InvalidRequest" && xrpc_message.starts_with("Could not locate record") } pub async fn fetch_schema_from_pds( - pds_endpoint: &str, + pds_endpoint: &SchemaHostUrl, did: &Did, nsid: &Nsid, ) -> Result { - let url = format!( - "{}/xrpc/com.atproto.repo.getRecord?repo={}&collection=com.atproto.lexicon.schema&rkey={}", - pds_endpoint.trim_end_matches('/'), - urlencoding::encode(did.as_str()), - urlencoding::encode(nsid.as_str()) - ); + let mut request_url = pds_endpoint.endpoint("xrpc/com.atproto.repo.getRecord"); + request_url + .query_pairs_mut() + .append_pair("repo", did.as_str()) + .append_pair("collection", "com.atproto.lexicon.schema") + .append_pair("rkey", nsid.as_str()); + let url = request_url.to_string(); let resp = client() - .get(&url) + .get(request_url) .send() .await .map_err(|e| ResolveError::SchemaFetch { @@ -204,10 +230,27 @@ pub async fn fetch_schema_from_pds( let status = resp.status(); if !status.is_success() { - return Err(ResolveError::SchemaFetch { - url, - reason: format!("HTTP {}", status), - }); + let body = read_body_limited(resp, MAX_RESPONSE_BYTES) + .await + .ok() + .and_then(|bytes| serde_json::from_slice::(&bytes).ok()) + .unwrap_or(serde_json::Value::Null); + let field = |name: &str| { + body.get(name) + .and_then(|v| v.as_str()) + .unwrap_or_default() + .to_string() + }; + return match is_record_absent(&field("error"), &field("message")) { + true => Err(ResolveError::SchemaNotFound { + nsid: nsid.clone(), + url, + }), + false => Err(ResolveError::SchemaFetch { + url, + reason: format!("HTTP {}", status), + }), + }; } let body = read_body_limited(resp, MAX_RESPONSE_BYTES) @@ -292,6 +335,27 @@ mod tests { s.parse().unwrap() } + #[test] + fn is_record_absent_recognizes_only_the_reference_pds_absence_shapes() { + assert!(is_record_absent( + "RecordNotFound", + "Could not locate record: at://did:plc:nel/com.atproto.lexicon.schema/x" + )); + assert!(is_record_absent("RecordNotFound", "")); + assert!(is_record_absent( + "InvalidRequest", + "Could not locate record" + )); + assert!(!is_record_absent( + "InvalidRequest", + "Error: rkey must be a valid record key" + )); + assert!(!is_record_absent("InvalidRequest", "")); + assert!(!is_record_absent("InternalServerError", "")); + assert!(!is_record_absent("RateLimitExceeded", "")); + assert!(!is_record_absent("", "")); + } + #[test] fn test_nsid_to_authority() { assert_eq!( @@ -316,57 +380,6 @@ mod tests { ); } - #[test] - fn test_extract_pds_endpoint_valid() { - let doc = serde_json::json!({ - "service": [{ - "type": "AtprotoPersonalDataServer", - "serviceEndpoint": "https://pds.example.com" - }] - }); - assert_eq!( - extract_pds_endpoint(&doc), - Some("https://pds.example.com".to_string()) - ); - } - - #[test] - fn test_extract_pds_endpoint_multiple_services() { - let doc = serde_json::json!({ - "service": [ - { - "type": "AtprotoLabeler", - "serviceEndpoint": "https://labeler.example.com" - }, - { - "type": "AtprotoPersonalDataServer", - "serviceEndpoint": "https://pds.example.com" - } - ] - }); - assert_eq!( - extract_pds_endpoint(&doc), - Some("https://pds.example.com".to_string()) - ); - } - - #[test] - fn test_extract_pds_endpoint_missing() { - let doc = serde_json::json!({ - "service": [{ - "type": "AtprotoLabeler", - "serviceEndpoint": "https://labeler.example.com" - }] - }); - assert_eq!(extract_pds_endpoint(&doc), None); - } - - #[test] - fn test_extract_pds_endpoint_no_services() { - let doc = serde_json::json!({}); - assert_eq!(extract_pds_endpoint(&doc), None); - } - #[test] fn test_validate_fetched_schema_ok() { let doc = LexiconDoc { diff --git a/crates/tranquil-lexicon/src/schema.rs b/crates/tranquil-lexicon/src/schema.rs index d71b12f..261e766 100644 --- a/crates/tranquil-lexicon/src/schema.rs +++ b/crates/tranquil-lexicon/src/schema.rs @@ -1,8 +1,8 @@ -use serde::Deserialize; +use serde::{Deserialize, Serialize}; use std::collections::HashMap; use tranquil_types::Nsid; -#[derive(Debug, Deserialize)] +#[derive(Debug, Serialize, Deserialize)] pub struct LexiconDoc { pub lexicon: u32, pub id: Nsid, @@ -10,7 +10,7 @@ pub struct LexiconDoc { pub defs: HashMap, } -#[derive(Debug, Deserialize)] +#[derive(Debug, Serialize, Deserialize)] #[serde(tag = "type")] pub enum LexDef { #[serde(rename = "record")] @@ -35,14 +35,14 @@ pub enum LexDef { PermissionSet {}, } -#[derive(Debug, Deserialize)] +#[derive(Debug, Serialize, Deserialize)] pub struct LexRecord { #[serde(default)] pub key: Option, pub record: LexObject, } -#[derive(Debug, Deserialize)] +#[derive(Debug, Serialize, Deserialize)] pub struct LexObject { #[serde(default)] pub required: Vec, @@ -52,7 +52,7 @@ pub struct LexObject { pub properties: HashMap, } -#[derive(Debug, Deserialize)] +#[derive(Debug, Serialize, Deserialize)] #[serde(tag = "type")] pub enum LexProperty { #[serde(rename = "string")] @@ -79,7 +79,7 @@ pub enum LexProperty { Object(LexObject), } -#[derive(Debug, Deserialize)] +#[derive(Debug, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] pub struct LexString { #[serde(default)] @@ -102,7 +102,7 @@ pub struct LexString { pub default: Option, } -#[derive(Debug, Deserialize)] +#[derive(Debug, Serialize, Deserialize)] pub struct LexInteger { #[serde(default)] pub minimum: Option, @@ -116,7 +116,7 @@ pub struct LexInteger { pub const_value: Option, } -#[derive(Debug, Deserialize)] +#[derive(Debug, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] pub struct LexBytes { #[serde(default)] @@ -125,7 +125,7 @@ pub struct LexBytes { pub min_length: Option, } -#[derive(Debug, Deserialize)] +#[derive(Debug, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] pub struct LexBlob { #[serde(default)] @@ -134,7 +134,7 @@ pub struct LexBlob { pub max_size: Option, } -#[derive(Debug, Deserialize)] +#[derive(Debug, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] pub struct LexArray { pub items: Box, @@ -144,7 +144,7 @@ pub struct LexArray { pub max_length: Option, } -#[derive(Debug, Deserialize)] +#[derive(Debug, Serialize, Deserialize)] pub struct LexUnion { #[serde(default)] pub refs: Vec, @@ -152,14 +152,14 @@ pub struct LexUnion { pub closed: bool, } -#[derive(Debug, Deserialize)] +#[derive(Debug, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] pub struct LexRef { #[serde(rename = "ref")] pub reference: String, } -#[derive(Debug, Clone, Deserialize)] +#[derive(Debug, Clone, Serialize, Deserialize)] pub enum StringFormat { #[serde(rename = "did")] Did, @@ -204,6 +204,6 @@ pub fn parse_ref(reference: &str) -> ParsedRef<'_> { } } -#[derive(Debug, Deserialize)] +#[derive(Debug, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] pub struct LexStringDef {} diff --git a/crates/tranquil-lexicon/tests/resolve_integration.rs b/crates/tranquil-lexicon/tests/resolve_integration.rs index 68158ee..c0f2c62 100644 --- a/crates/tranquil-lexicon/tests/resolve_integration.rs +++ b/crates/tranquil-lexicon/tests/resolve_integration.rs @@ -74,7 +74,7 @@ async fn test_resolve_pds_endpoint_from_plc() { let endpoint = resolve_pds_endpoint(&did.parse().unwrap(), Some(&plc_server.uri())) .await .unwrap(); - assert_eq!(endpoint, "https://pds.example.com"); + assert_eq!(endpoint.as_str(), "https://pds.example.com"); } #[tokio::test] @@ -130,14 +130,17 @@ async fn test_resolve_pds_endpoint_multiple_services_picks_pds() { "id": did, "service": [ { + "id": "#atproto_labeler", "type": "AtprotoLabeler", "serviceEndpoint": "https://labeler.example.com" }, { + "id": "#bsky_notif", "type": "BskyNotificationService", "serviceEndpoint": "https://notify.example.com" }, { + "id": "#atproto_pds", "type": "AtprotoPersonalDataServer", "serviceEndpoint": "https://pds.example.com" } @@ -149,7 +152,7 @@ async fn test_resolve_pds_endpoint_multiple_services_picks_pds() { let endpoint = resolve_pds_endpoint(&did.parse().unwrap(), Some(&plc_server.uri())) .await .unwrap(); - assert_eq!(endpoint, "https://pds.example.com"); + assert_eq!(endpoint.as_str(), "https://pds.example.com"); } #[tokio::test] @@ -168,7 +171,7 @@ async fn test_fetch_schema_from_pds_success() { .await; let doc = fetch_schema_from_pds( - &pds_server.uri(), + &pds_server.uri().parse().unwrap(), &did.parse().unwrap(), &nsid.parse().unwrap(), ) @@ -195,7 +198,7 @@ async fn test_fetch_schema_missing_value_field() { .await; let result = fetch_schema_from_pds( - &pds_server.uri(), + &pds_server.uri().parse().unwrap(), &did.parse().unwrap(), &nsid.parse().unwrap(), ) @@ -222,7 +225,7 @@ async fn test_fetch_schema_invalid_lexicon_json() { .await; let result = fetch_schema_from_pds( - &pds_server.uri(), + &pds_server.uri().parse().unwrap(), &did.parse().unwrap(), &nsid.parse().unwrap(), ) @@ -352,7 +355,7 @@ async fn test_pds_trailing_slash_handled() { let pds_url_with_slash = format!("{}/", pds_server.uri()); let doc = fetch_schema_from_pds( - &pds_url_with_slash, + &pds_url_with_slash.parse().unwrap(), &did.parse().unwrap(), &nsid.parse().unwrap(), ) @@ -377,7 +380,7 @@ async fn test_fetch_schema_error_status_gives_meaningful_error() { .await; let result = fetch_schema_from_pds( - &pds_server.uri(), + &pds_server.uri().parse().unwrap(), &did.parse().unwrap(), &nsid.parse().unwrap(), ) -- 2.51.2 From 8d0b6f8322c9e964a75321474f33bfb394f35bf3 Mon Sep 17 00:00:00 2001 From: Lewis Date: Sun, 26 Jul 2026 13:57:07 +0300 Subject: [PATCH 07/22] cache: DID, SSO, & OAuth client metadata caches onto shared cache Lewis: May this revision serve well! --- crates/tranquil-api/src/identity/account.rs | 9 +- crates/tranquil-api/src/server/session.rs | 6 +- crates/tranquil-config/src/lib.rs | 2 +- .../src/endpoints/authorize/consent.rs | 2 +- .../src/endpoints/authorize/login.rs | 2 +- .../src/endpoints/authorize/mod.rs | 3 +- .../src/endpoints/par.rs | 6 +- .../src/endpoints/token/grants.rs | 7 +- crates/tranquil-oauth/src/client.rs | 194 +++++++------ crates/tranquil-pds/src/did.rs | 266 +++++------------- crates/tranquil-pds/src/sso/providers.rs | 261 ++++++++++------- crates/tranquil-pds/src/state.rs | 41 ++- example.toml | 2 +- 13 files changed, 376 insertions(+), 425 deletions(-) diff --git a/crates/tranquil-api/src/identity/account.rs b/crates/tranquil-api/src/identity/account.rs index 2b2a14a..3575c2d 100644 --- a/crates/tranquil-api/src/identity/account.rs +++ b/crates/tranquil-api/src/identity/account.rs @@ -147,12 +147,7 @@ async fn try_reactivate_migration( Json(CreateAccountOutput { handle: handle.clone(), did: did.clone(), - did_doc: state - .did_resolver - .fetch_did_document(did) - .await - .ok() - .map(|f| (*f).clone()), + did_doc: state.did_resolver.fetch_did_document(did).await.ok(), access_jwt: access_meta.token, refresh_jwt: refresh_meta.token, verification_required, @@ -568,7 +563,7 @@ pub async fn create_account( Json(CreateAccountOutput { handle: handle.clone(), did, - did_doc: did_doc.map(|f| (*f).clone()), + did_doc, access_jwt: session.access_jwt, refresh_jwt: session.refresh_jwt, verification_required: !is_migration, diff --git a/crates/tranquil-api/src/server/session.rs b/crates/tranquil-api/src/server/session.rs index e772f40..4cc2d40 100644 --- a/crates/tranquil-api/src/server/session.rs +++ b/crates/tranquil-api/src/server/session.rs @@ -351,7 +351,7 @@ pub async fn create_session( refresh_jwt: refresh_meta.token, handle, did: row.did, - did_doc: did_doc.ok().map(|f| (*f).clone()), + did_doc: did_doc.ok(), email: row.email, email_confirmed: Some(row.channel_verification.email), email_auth_factor: email_auth_factor_out, @@ -444,7 +444,7 @@ pub async fn get_session( status: account_state.status_for_session().map(String::from), migrated_to_pds, migrated_at, - did_doc: did_doc.ok().map(|f| (*f).clone()), + did_doc: did_doc.ok(), })) } Ok(None) => Err(ApiError::AuthenticationFailed(None)), @@ -800,7 +800,7 @@ async fn build_refresh_session_output( preferred_locale: u.preferred_locale, is_admin: u.is_admin, active: account_state.is_active(), - did_doc: did_doc.ok().map(|f| (*f).clone()), + did_doc: did_doc.ok(), status: account_state.status_for_session().map(String::from), })) } diff --git a/crates/tranquil-config/src/lib.rs b/crates/tranquil-config/src/lib.rs index 3a0185e..f71e7e7 100644 --- a/crates/tranquil-config/src/lib.rs +++ b/crates/tranquil-config/src/lib.rs @@ -835,7 +835,7 @@ pub struct PlcConfig { #[config(env = "PLC_CONNECT_TIMEOUT_SECS", default = 5)] pub connect_timeout_secs: u64, - /// Seconds to cache DID documents in memory. + /// Seconds to cache DID documents. #[config(env = "DID_CACHE_TTL_SECS", default = 300)] pub did_cache_ttl_secs: u64, } diff --git a/crates/tranquil-oauth-server/src/endpoints/authorize/consent.rs b/crates/tranquil-oauth-server/src/endpoints/authorize/consent.rs index 5743817..d6c1e40 100644 --- a/crates/tranquil-oauth-server/src/endpoints/authorize/consent.rs +++ b/crates/tranquil-oauth-server/src/endpoints/authorize/consent.rs @@ -120,7 +120,7 @@ pub async fn consent_get( }; let did = flow_with_user.did().clone(); - let client_cache = ClientMetadataCache::new(3600); + let client_cache = &state.client_metadata_cache; let client_metadata = client_cache .get(&request_data.parameters.client_id) .await diff --git a/crates/tranquil-oauth-server/src/endpoints/authorize/login.rs b/crates/tranquil-oauth-server/src/endpoints/authorize/login.rs index 0e3bd8d..577d4a6 100644 --- a/crates/tranquil-oauth-server/src/endpoints/authorize/login.rs +++ b/crates/tranquil-oauth-server/src/endpoints/authorize/login.rs @@ -80,7 +80,7 @@ pub async fn authorize_get( "Authorization request has expired. Please start a new request.", ); } - let client_cache = ClientMetadataCache::new(3600); + let client_cache = &state.client_metadata_cache; let client_name = client_cache .get(&request_data.parameters.client_id) .await diff --git a/crates/tranquil-oauth-server/src/endpoints/authorize/mod.rs b/crates/tranquil-oauth-server/src/endpoints/authorize/mod.rs index 2b37337..d9283cd 100644 --- a/crates/tranquil-oauth-server/src/endpoints/authorize/mod.rs +++ b/crates/tranquil-oauth-server/src/endpoints/authorize/mod.rs @@ -14,8 +14,7 @@ use tranquil_db_traits::{ScopePreference, WebauthnChallengeType}; use tranquil_pds::auth::{BareLoginIdentifier, NormalizedLoginIdentifier}; use tranquil_pds::comms::comms_repo::enqueue_2fa_code; use tranquil_pds::oauth::{ - AuthFlow, ClientMetadataCache, DeviceData, DeviceId, OAuthError, Prompt, SessionId, - db::should_show_consent, + AuthFlow, DeviceData, DeviceId, OAuthError, Prompt, SessionId, db::should_show_consent, }; use tranquil_pds::rate_limit::{ OAuthAuthorizeLimit, OAuthRateLimited, OAuthRegisterCompleteLimit, TotpVerifyLimit, diff --git a/crates/tranquil-oauth-server/src/endpoints/par.rs b/crates/tranquil-oauth-server/src/endpoints/par.rs index dcaa075..c122161 100644 --- a/crates/tranquil-oauth-server/src/endpoints/par.rs +++ b/crates/tranquil-oauth-server/src/endpoints/par.rs @@ -3,8 +3,8 @@ use axum::{Json, extract::State, http::HeaderMap}; use chrono::{Duration, Utc}; use serde::{Deserialize, Serialize}; use tranquil_pds::oauth::{ - AuthorizationRequestParameters, ClientAuth, ClientMetadataCache, CodeChallengeMethod, - OAuthError, Prompt, RequestData, RequestId, ResponseMode, ResponseType, + AuthorizationRequestParameters, ClientAuth, CodeChallengeMethod, OAuthError, Prompt, + RequestData, RequestId, ResponseMode, ResponseType, scopes::{ParsedScope, parse_scope}, }; use tranquil_pds::rate_limit::{OAuthParLimit, OAuthRateLimited}; @@ -80,7 +80,7 @@ pub async fn pushed_authorization_request( .ok_or_else(|| OAuthError::InvalidRequest("code_challenge is required".to_string()))?; let code_challenge_method = parse_code_challenge_method(request.code_challenge_method.as_deref())?; - let client_cache = ClientMetadataCache::new(3600); + let client_cache = &state.client_metadata_cache; let client_metadata = client_cache.get(&request.client_id).await?; client_cache.validate_redirect_uri(&client_metadata, &request.redirect_uri)?; let client_auth = determine_client_auth(&request)?; diff --git a/crates/tranquil-oauth-server/src/endpoints/token/grants.rs b/crates/tranquil-oauth-server/src/endpoints/token/grants.rs index b9698d5..981d5a9 100644 --- a/crates/tranquil-oauth-server/src/endpoints/token/grants.rs +++ b/crates/tranquil-oauth-server/src/endpoints/token/grants.rs @@ -8,8 +8,7 @@ use chrono::{Duration, Utc}; use tranquil_db_traits::RefreshTokenLookup; use tranquil_pds::config::AuthConfig; use tranquil_pds::oauth::{ - AuthFlow, ClientAuth, ClientMetadataCache, DPoPVerifier, OAuthError, RefreshToken, TokenData, - TokenId, + AuthFlow, ClientAuth, DPoPVerifier, OAuthError, RefreshToken, TokenData, TokenId, db::{enforce_token_limit_for_user, lookup_refresh_token}, verify_client_auth, }; @@ -63,7 +62,7 @@ pub async fn handle_authorization_code_grant( return Err(OAuthError::InvalidGrant("client_id mismatch".to_string())); } let did = authorized.did.clone(); - let client_metadata_cache = ClientMetadataCache::new(3600); + let client_metadata_cache = &state.client_metadata_cache; let client_metadata = client_metadata_cache.get(&authorized.client_id).await?; let client_auth = match &request.client_auth { RequestClientAuth::PrivateKeyJwt { @@ -85,7 +84,7 @@ pub async fn handle_authorization_code_grant( }, RequestClientAuth::None { .. } => ClientAuth::None, }; - verify_client_auth(&client_metadata_cache, &client_metadata, &client_auth).await?; + verify_client_auth(client_metadata_cache, &client_metadata, &client_auth).await?; verify_pkce(&authorized.parameters.code_challenge, &code_verifier)?; if let Some(req_redirect_uri) = &redirect_uri && req_redirect_uri != &authorized.parameters.redirect_uri diff --git a/crates/tranquil-oauth/src/client.rs b/crates/tranquil-oauth/src/client.rs index cb4b588..8b126a3 100644 --- a/crates/tranquil-oauth/src/client.rs +++ b/crates/tranquil-oauth/src/client.rs @@ -1,12 +1,19 @@ use reqwest::Client; use serde::{Deserialize, Serialize}; -use std::collections::HashMap; use std::sync::Arc; -use tokio::sync::RwLock; +use std::time::Duration; use crate::OAuthError; use crate::types::ClientAuth; -use tranquil_types::ClientId; +use tranquil_infra::cache_keys::{ + oauth_client_jwks_cooldown_key, oauth_client_jwks_key, oauth_client_meta_key, +}; +use tranquil_infra::{Cache, cached_json, write_json}; +use tranquil_types::{ + ClientId, JwksUri, ReachPolicy, dns_guard, redirect_policy, url_reach_permits, +}; + +const JWKS_REFRESH_COOLDOWN: Duration = Duration::from_secs(60); #[derive(Debug, Clone, Serialize, Deserialize)] pub struct ClientMetadata { @@ -30,8 +37,12 @@ pub struct ClientMetadata { pub dpop_bound_access_tokens: Option, #[serde(skip_serializing_if = "Option::is_none")] pub jwks: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub jwks_uri: Option, + #[serde( + default, + skip_serializing_if = "Option::is_none", + deserialize_with = "tranquil_types::http_url::deserialize_optional" + )] + pub jwks_uri: Option, #[serde(skip_serializing_if = "Option::is_none")] pub application_type: Option, } @@ -58,33 +69,23 @@ impl Default for ClientMetadata { #[derive(Clone)] pub struct ClientMetadataCache { - cache: Arc>>, - jwks_cache: Arc>>, + cache: Arc, http_client: Client, - cache_ttl_secs: u64, -} - -struct CachedMetadata { - metadata: ClientMetadata, - cached_at: std::time::Instant, -} - -struct CachedJwks { - jwks: serde_json::Value, - cached_at: std::time::Instant, + cache_ttl: Duration, } impl ClientMetadataCache { - pub fn new(cache_ttl_secs: u64) -> Self { + pub fn new(cache: Arc, cache_ttl: Duration) -> Self { Self { - cache: Arc::new(RwLock::new(HashMap::new())), - jwks_cache: Arc::new(RwLock::new(HashMap::new())), + cache, http_client: { let builder = Client::builder() .timeout(std::time::Duration::from_secs(30)) .connect_timeout(std::time::Duration::from_secs(10)) .pool_max_idle_per_host(10) .pool_idle_timeout(std::time::Duration::from_secs(90)) + .redirect(redirect_policy(ReachPolicy::DEBUG_LOOPBACK)) + .dns_resolver(dns_guard(ReachPolicy::DEBUG_LOOPBACK)) .user_agent(concat!( "Tranquil-PDS/", env!("CARGO_PKG_VERSION"), @@ -92,9 +93,11 @@ impl ClientMetadataCache { )); #[cfg(feature = "native-tls-roots")] let builder = builder.danger_accept_invalid_certs(true); - builder.build().unwrap_or_else(|_| Client::new()) + builder + .build() + .expect("failed to build client metadata HTTP client") }, - cache_ttl_secs, + cache_ttl, } } @@ -150,26 +153,13 @@ impl ClientMetadataCache { if Self::is_loopback_client(client_id) { return Self::build_loopback_metadata(client_id); } - { - let cache = self.cache.read().await; - if let Some(cached) = cache.get(client_id.as_str()) - && cached.cached_at.elapsed().as_secs() < self.cache_ttl_secs - { - return Ok(cached.metadata.clone()); - } - } - let metadata = self.fetch_metadata(client_id).await?; - { - let mut cache = self.cache.write().await; - cache.insert( - client_id.to_string(), - CachedMetadata { - metadata: metadata.clone(), - cached_at: std::time::Instant::now(), - }, - ); - } - Ok(metadata) + cached_json( + self.cache.as_ref(), + &oauth_client_meta_key(client_id), + self.cache_ttl, + || self.fetch_metadata(client_id), + ) + .await } pub async fn get_jwks( @@ -181,43 +171,57 @@ impl ClientMetadataCache { } let jwks_uri = metadata.jwks_uri.as_ref().ok_or_else(|| { OAuthError::InvalidClient( - "Client using private_key_jwt must have jwks or jwks_uri".to_string(), + "Client using private_key_jwt must have jwks or a usable jwks_uri".to_string(), ) })?; - { - let cache = self.jwks_cache.read().await; - if let Some(cached) = cache.get(jwks_uri) - && cached.cached_at.elapsed().as_secs() < self.cache_ttl_secs - { - return Ok(cached.jwks.clone()); + cached_json( + self.cache.as_ref(), + &oauth_client_jwks_key(jwks_uri), + self.cache_ttl, + || self.fetch_jwks(jwks_uri), + ) + .await + } + + async fn refresh_jwks( + &self, + metadata: &ClientMetadata, + ) -> Result, OAuthError> { + match (&metadata.jwks, &metadata.jwks_uri) { + (None, Some(jwks_uri)) => { + let cooldown_key = oauth_client_jwks_cooldown_key(jwks_uri); + if self.cache.get(&cooldown_key).await.is_some() { + return Ok(None); + } + let _ = self + .cache + .set(&cooldown_key, "1", JWKS_REFRESH_COOLDOWN) + .await; + self.fetch_and_store_jwks(jwks_uri).await.map(Some) } + _ => Ok(None), } + } + + async fn fetch_and_store_jwks( + &self, + jwks_uri: &JwksUri, + ) -> Result { let jwks = self.fetch_jwks(jwks_uri).await?; - { - let mut cache = self.jwks_cache.write().await; - cache.insert( - jwks_uri.clone(), - CachedJwks { - jwks: jwks.clone(), - cached_at: std::time::Instant::now(), - }, - ); - } + write_json( + self.cache.as_ref(), + &oauth_client_jwks_key(jwks_uri), + &jwks, + self.cache_ttl, + ) + .await; Ok(jwks) } - async fn fetch_jwks(&self, jwks_uri: &str) -> Result { - if !jwks_uri.starts_with("https://") - && (!jwks_uri.starts_with("http://") - || (!jwks_uri.contains("localhost") && !jwks_uri.contains("127.0.0.1"))) - { - return Err(OAuthError::InvalidClient( - "jwks_uri must use https (except for localhost)".to_string(), - )); - } + async fn fetch_jwks(&self, jwks_uri: &JwksUri) -> Result { let response = self .http_client - .get(jwks_uri) + .get(jwks_uri.as_str()) .header("Accept", "application/json") .send() .await @@ -243,22 +247,16 @@ impl ClientMetadataCache { } async fn fetch_metadata(&self, client_id: &ClientId) -> Result { - if !client_id.starts_with("http://") && !client_id.starts_with("https://") { - return Err(OAuthError::InvalidClient( - "client_id must be a URL".to_string(), - )); - } - if client_id.starts_with("http://") - && !client_id.contains("localhost") - && !client_id.contains("127.0.0.1") - { + let url = reqwest::Url::parse(client_id) + .map_err(|_| OAuthError::InvalidClient("client_id must be a URL".to_string()))?; + if !url_reach_permits(&url, ReachPolicy::DEBUG_LOOPBACK) { return Err(OAuthError::InvalidClient( - "Non-localhost client_id must use https".to_string(), + "client_id must be an https URL inside the allowed host reach".to_string(), )); } let response = self .http_client - .get(client_id.as_str()) + .get(url) .header("Accept", "application/json") .send() .await @@ -514,7 +512,29 @@ async fn verify_private_key_jwt_async( "client_assertion iat is in the future".to_string(), )); } + let signing_input = format!("{}.{}", parts[0], parts[1]); + let signature_bytes = URL_SAFE_NO_PAD + .decode(parts[2]) + .map_err(|_| OAuthError::InvalidClient("Invalid signature encoding".to_string()))?; let jwks = cache.get_jwks(metadata).await?; + match verify_assertion_signature(&jwks, kid, alg, &signing_input, &signature_bytes) { + Ok(()) => Ok(()), + Err(cached_failure) => match cache.refresh_jwks(metadata).await { + Ok(Some(fresh)) => { + verify_assertion_signature(&fresh, kid, alg, &signing_input, &signature_bytes) + } + Ok(None) | Err(_) => Err(cached_failure), + }, + } +} + +fn verify_assertion_signature( + jwks: &serde_json::Value, + kid: Option<&str>, + alg: &str, + signing_input: &str, + signature: &[u8], +) -> Result<(), OAuthError> { let keys = jwks .get("keys") .and_then(|k| k.as_array()) @@ -531,10 +551,6 @@ async fn verify_private_key_jwt_async( "No matching key found in client JWKS".to_string(), )); } - let signing_input = format!("{}.{}", parts[0], parts[1]); - let signature_bytes = URL_SAFE_NO_PAD - .decode(parts[2]) - .map_err(|_| OAuthError::InvalidClient("Invalid signature encoding".to_string()))?; matching_keys .into_iter() .filter(|key| { @@ -544,12 +560,12 @@ async fn verify_private_key_jwt_async( .find_map(|key| { let kty = key.get("kty").and_then(|k| k.as_str()).unwrap_or(""); match (alg, kty) { - ("ES256", "EC") => verify_es256(key, &signing_input, &signature_bytes).ok(), - ("ES384", "EC") => verify_es384(key, &signing_input, &signature_bytes).ok(), + ("ES256", "EC") => verify_es256(key, signing_input, signature).ok(), + ("ES384", "EC") => verify_es384(key, signing_input, signature).ok(), ("RS256" | "RS384" | "RS512", "RSA") => { - verify_rsa(alg, key, &signing_input, &signature_bytes).ok() + verify_rsa(alg, key, signing_input, signature).ok() } - ("EdDSA", "OKP") => verify_eddsa(key, &signing_input, &signature_bytes).ok(), + ("EdDSA", "OKP") => verify_eddsa(key, signing_input, signature).ok(), _ => None, } }) diff --git a/crates/tranquil-pds/src/did.rs b/crates/tranquil-pds/src/did.rs index ac41376..0a066d4 100644 --- a/crates/tranquil-pds/src/did.rs +++ b/crates/tranquil-pds/src/did.rs @@ -1,10 +1,9 @@ +use crate::cache::Cache; use crate::types::Did; use reqwest::Client; use serde::{Deserialize, Serialize}; -use std::collections::HashMap; use std::sync::Arc; -use std::time::{Duration, Instant}; -use tokio::sync::RwLock; +use std::time::Duration; use tracing::{debug, info, warn}; #[derive(Debug, thiserror::Error)] @@ -13,6 +12,8 @@ pub enum DidResolutionError { UnsupportedDidMethod(String), #[error("Invalid did:web format")] InvalidDidWeb, + #[error("did:web host {0} is outside the allowed host reach")] + DidWebHostRejected(String), #[error("HTTP request failed: {0}")] HttpFailed(String), #[error("Invalid DID document: {0}")] @@ -53,43 +54,50 @@ pub struct DidService { pub struct ResolvedService { pub url: String, pub did: Did, - pub service_id: String, } -type TimedCache = RwLock, (Instant, Arc)>>; - pub struct DidResolver { - did_doc_cache: TimedCache, - parsed_did_doc_cache: TimedCache, - service_cache: TimedCache, + cache: Arc, client: Client, cache_ttl: Duration, plc_directory_url: String, } impl DidResolver { - pub fn new() -> Self { + pub fn new(cache: Arc) -> Self { let cfg = tranquil_config::get(); - let cache_ttl_secs = cfg.plc.did_cache_ttl_secs; - - let plc_directory_url = cfg.plc.directory_url.clone(); let client = Client::builder() .timeout(Duration::from_secs(10)) .connect_timeout(Duration::from_secs(5)) .pool_max_idle_per_host(10) + .redirect(tranquil_types::redirect_policy( + tranquil_types::ReachPolicy::DEBUG_LOOPBACK, + )) + .dns_resolver(tranquil_types::dns_guard( + tranquil_types::ReachPolicy::DEBUG_LOOPBACK, + )) .build() - .unwrap_or_else(|_| Client::new()); + .expect("failed to build DID resolver HTTP client"); info!("DID resolver initialized"); Self { - did_doc_cache: RwLock::new(HashMap::new()), - parsed_did_doc_cache: RwLock::new(HashMap::new()), - service_cache: RwLock::new(HashMap::new()), + cache, client, - cache_ttl: Duration::from_secs(cache_ttl_secs), - plc_directory_url, + cache_ttl: Duration::from_secs(cfg.plc.did_cache_ttl_secs), + plc_directory_url: cfg.plc.directory_url.clone(), + } + } + + fn doc_cache_key(did: &Did) -> Result { + match (did.is_plc(), did.is_web()) { + (true, _) => Ok(crate::cache_keys::plc_doc_key(did)), + (_, true) => Ok(crate::cache_keys::did_web_doc_key(did)), + _ => { + warn!("Unsupported DID method: {}", did); + Err(DidResolutionError::UnsupportedDidMethod(did.to_string())) + } } } @@ -97,175 +105,50 @@ impl DidResolver { &self, did: &Did, service_id: &str, - ) -> Result, ServiceResolutionError> { - { - let cache = self.service_cache.read().await; - if let Some(cached) = cache.get(&*format!("{did}#{service_id}")) - && cached.0.elapsed() < self.cache_ttl - { - return Ok(cached.1.clone()); - } - } - + ) -> Result { let did_doc = self.resolve_did(did).await?; - let Some(service) = did_doc + let suffix = format!("#{service_id}"); + did_doc .services .iter() - .find(|s| s.id.ends_with(&format!("#{service_id}"))) - else { - return Err(ServiceResolutionError::ServiceIdNotFound(service_id.into())); - }; - - let resolved = Arc::new(ResolvedService { - url: service.service_endpoint.clone(), - did: did.clone(), - service_id: service_id.into(), - }); - - { - let mut cache = self.service_cache.write().await; - cache.insert( - format!("{did}#{service_id}").into(), - (Instant::now(), resolved.clone()), - ); - } - - Ok(resolved) + .find(|s| s.id.ends_with(&suffix)) + .map(|service| ResolvedService { + url: service.service_endpoint.clone(), + did: did.clone(), + }) + .ok_or_else(|| ServiceResolutionError::ServiceIdNotFound(service_id.into())) } - pub async fn resolve_did(&self, did: &Did) -> Result, DidResolutionError> { - { - let cache = self.parsed_did_doc_cache.read().await; - if let Some(cached) = cache.get(did.as_str()) - && cached.0.elapsed() < self.cache_ttl - { - return Ok(cached.1.clone()); - } - } - - let resolved = Arc::new(self.resolve_did_uncached(did).await?); - - { - let mut cache = self.parsed_did_doc_cache.write().await; - cache.insert(did.as_str().into(), (Instant::now(), resolved.clone())); - } - - Ok(resolved) + pub async fn resolve_did(&self, did: &Did) -> Result { + self.cached_did_document(did).await } - pub async fn refresh_did(&self, did: &Did) -> Result, DidResolutionError> { - { - let mut cache = self.parsed_did_doc_cache.write().await; - cache.remove(did.as_str()); - let mut cache = self.service_cache.write().await; - cache.retain(|k, _| !k.starts_with(did.as_str())); - } + pub async fn refresh_did(&self, did: &Did) -> Result { + let _ = self.cache.delete(&Self::doc_cache_key(did)?).await; self.resolve_did(did).await } - async fn resolve_did_uncached(&self, did: &Did) -> Result { - if did.is_web() { - self.resolve_did_web(did).await - } else if did.is_plc() { - self.resolve_did_plc(did).await - } else { - warn!("Unsupported DID method: {}", did); - Err(DidResolutionError::UnsupportedDidMethod(did.to_string())) - } - } - - async fn resolve_did_web(&self, did: &Did) -> Result { - let url = build_did_web_url(did)?; - - debug!("Resolving did:web {} via {}", did, url); - - let resp = self - .client - .get(&url) - .send() - .await - .map_err(|e| DidResolutionError::HttpFailed(e.to_string()))?; - - if !resp.status().is_success() { - return Err(DidResolutionError::HttpFailed(format!( - "HTTP {}", - resp.status() - ))); - } - - resp.json::() - .await - .map_err(|e| DidResolutionError::InvalidDocument(e.to_string())) - } - - async fn resolve_did_plc(&self, did: &Did) -> Result { - let url = format!( - "{}/{}", - self.plc_directory_url, - urlencoding::encode(did.as_str()) - ); - - debug!("Resolving did:plc {} via {}", did, url); - - let resp = self - .client - .get(&url) - .send() - .await - .map_err(|e| DidResolutionError::HttpFailed(e.to_string()))?; - - if resp.status() == reqwest::StatusCode::NOT_FOUND { - return Err(DidResolutionError::NotFound); - } - - if !resp.status().is_success() { - return Err(DidResolutionError::HttpFailed(format!( - "HTTP {}", - resp.status() - ))); - } - - resp.json::() - .await - .map_err(|e| DidResolutionError::InvalidDocument(e.to_string())) - } - pub async fn fetch_did_document( &self, did: &Did, - ) -> Result, DidResolutionError> { - { - let cache = self.did_doc_cache.read().await; - if let Some(cached) = cache.get(did.as_str()) - && cached.0.elapsed() < self.cache_ttl - { - return Ok(cached.1.clone()); - } - } - - let resolved = Arc::new(self.fetch_did_document_uncached(did).await?); - - { - let mut cache = self.did_doc_cache.write().await; - cache.insert(did.as_str().into(), (Instant::now(), resolved.clone())); - } - - Ok(resolved) + ) -> Result { + self.cached_did_document(did).await } - // TODO: make cached version - async fn fetch_did_document_uncached( + async fn cached_did_document( &self, did: &Did, - ) -> Result { - if did.is_web() { - self.fetch_did_document_web(did).await - } else if did.is_plc() { - self.fetch_did_document_plc(did).await - } else { - warn!("Unsupported DID method: {}", did); - Err(DidResolutionError::UnsupportedDidMethod(did.to_string())) - } + ) -> Result { + let cache_key = Self::doc_cache_key(did)?; + let doc = + crate::cache::cached_json(self.cache.as_ref(), &cache_key, self.cache_ttl, || async { + match did.is_plc() { + true => self.fetch_did_document_plc(did).await, + false => self.fetch_did_document_web(did).await, + } + }) + .await?; + serde_json::from_value(doc).map_err(|e| DidResolutionError::InvalidDocument(e.to_string())) } async fn fetch_did_document_web( @@ -274,6 +157,8 @@ impl DidResolver { ) -> Result { let url = build_did_web_url(did)?; + debug!("Resolving did:web {} via {}", did, url); + let resp = self .client .get(&url) @@ -303,6 +188,8 @@ impl DidResolver { urlencoding::encode(did.as_str()) ); + debug!("Resolving did:plc {} via {}", did, url); + let resp = self .client .get(&url) @@ -325,21 +212,6 @@ impl DidResolver { .await .map_err(|e| DidResolutionError::InvalidDocument(e.to_string())) } - - pub async fn invalidate_cache(&self, did: &Did) { - let mut doc_cache = self.parsed_did_doc_cache.write().await; - doc_cache.remove(did.as_str()); - } -} - -impl Default for DidResolver { - fn default() -> Self { - Self::new() - } -} - -pub fn create_did_resolver() -> Arc { - Arc::new(DidResolver::new()) } fn build_did_web_url(did: &Did) -> Result { @@ -372,18 +244,18 @@ fn build_did_web_url(did: &Did) -> Result { } }; - let scheme = - if host.starts_with("localhost") || host.starts_with("127.0.0.1") || host.contains(':') { - "http" - } else { - "https" - }; - - let url = if path.is_empty() { - format!("{}://{}/.well-known/did.json", scheme, host) + let https = if path.is_empty() { + format!("https://{}/.well-known/did.json", host) } else { - format!("{}://{}{}/did.json", scheme, host, path) + format!("https://{}{}/did.json", host, path) }; - Ok(url) + let mut url = reqwest::Url::parse(&https).map_err(|_| DidResolutionError::InvalidDidWeb)?; + if tranquil_types::url_reach(&url) == Some(tranquil_types::HostReach::Loopback) { + let _ = url.set_scheme("http"); + } + match tranquil_types::url_reach_permits(&url, tranquil_types::ReachPolicy::DEBUG_LOOPBACK) { + true => Ok(url.to_string()), + false => Err(DidResolutionError::DidWebHostRejected(host)), + } } diff --git a/crates/tranquil-pds/src/sso/providers.rs b/crates/tranquil-pds/src/sso/providers.rs index 2ced382..0b2e4be 100644 --- a/crates/tranquil-pds/src/sso/providers.rs +++ b/crates/tranquil-pds/src/sso/providers.rs @@ -4,15 +4,23 @@ use jsonwebtoken::{Algorithm, DecodingKey, EncodingKey, Header, Validation, jwk: use reqwest::Client; use serde::{Deserialize, Serialize}; use std::collections::HashMap; -use std::sync::Arc; +use std::sync::{Arc, LazyLock}; use std::time::{Duration, SystemTime, UNIX_EPOCH}; use thiserror::Error; -use tokio::sync::{OnceCell, RwLock}; +use tokio::sync::RwLock; use tranquil_db_traits::SsoProviderType; +use tranquil_types::{SsoIssuer, SsoJwksUri}; use super::config::{AppleProviderConfig, ProviderConfig, SsoConfig}; +use crate::cache::{Cache, cached_json}; +use crate::cache_keys::{oidc_discovery_key, sso_jwks_key}; const SSO_HTTP_TIMEOUT: Duration = Duration::from_secs(15); +const SSO_DISCOVERY_TTL: Duration = Duration::from_secs(3600); +static APPLE_JWKS_URI: LazyLock = LazyLock::new(|| { + SsoJwksUri::new("https://appleid.apple.com/auth/keys") + .expect("Apple JWKS URI is a valid https URL") +}); struct PkceChallenge { code_verifier: String, @@ -28,6 +36,12 @@ fn create_http_client() -> Client { Client::builder() .timeout(SSO_HTTP_TIMEOUT) .connect_timeout(Duration::from_secs(5)) + .redirect(tranquil_types::redirect_policy( + tranquil_types::ReachPolicy::AllowPrivate, + )) + .dns_resolver(tranquil_types::dns_guard( + tranquil_types::ReachPolicy::AllowPrivate, + )) .build() .expect("Failed to create HTTP client") } @@ -367,16 +381,21 @@ impl SsoProvider for DiscordProvider { } } -#[derive(Debug, Clone, Deserialize)] +#[derive(Debug, Clone, Serialize, Deserialize)] pub struct OidcDiscoveryConfig { - pub issuer: String, + pub issuer: SsoIssuer, pub authorization_endpoint: String, pub token_endpoint: String, pub userinfo_endpoint: Option, - pub jwks_uri: Option, + #[serde( + default, + deserialize_with = "tranquil_types::http_url::deserialize_optional" + )] + pub jwks_uri: Option, } -struct OidcDiscoveryCache { +#[derive(Serialize, Deserialize)] +struct OidcDiscovery { config: OidcDiscoveryConfig, jwks: Option, } @@ -385,10 +404,10 @@ pub struct OidcProvider { provider_type: SsoProviderType, client_id: String, client_secret: String, - issuer: String, + issuer: SsoIssuer, display_name: String, http_client: Client, - discovery_cache: OnceCell, + cache: Arc, } impl OidcProvider { @@ -397,11 +416,25 @@ impl OidcProvider { config: &ProviderConfig, default_issuer: Option<&str>, default_name: &str, + cache: Arc, ) -> Option { - let issuer = config + let issuer = match config .issuer .clone() - .or_else(|| default_issuer.map(String::from))?; + .or_else(|| default_issuer.map(String::from)) + .map(SsoIssuer::new) + { + Some(Ok(issuer)) => issuer, + Some(Err(e)) => { + tracing::error!( + provider = %provider_type.as_str(), + error = %e, + "SSO provider disabled because its issuer isn't a usable http or https URL" + ); + return None; + } + None => return None, + }; Some(Self { provider_type, @@ -413,74 +446,80 @@ impl OidcProvider { .clone() .unwrap_or_else(|| default_name.to_string()), http_client: create_http_client(), - discovery_cache: OnceCell::new(), + cache, }) } - async fn get_discovery(&self) -> Result<&OidcDiscoveryCache, SsoError> { - self.discovery_cache - .get_or_try_init(|| async { - let discovery_url = format!( - "{}/.well-known/openid-configuration", - self.issuer.trim_end_matches('/') - ); + async fn get_discovery(&self) -> Result { + cached_json( + self.cache.as_ref(), + &oidc_discovery_key(&self.issuer), + SSO_DISCOVERY_TTL, + || self.fetch_discovery(), + ) + .await + } - tracing::debug!( - provider = %self.provider_type.as_str(), - url = %discovery_url, - "Fetching OIDC discovery document" - ); + async fn fetch_discovery(&self) -> Result { + let discovery_url = self.issuer.endpoint(".well-known/openid-configuration"); - let resp = self - .http_client - .get(&discovery_url) - .send() - .await - .map_err(|e| SsoError::Discovery(e.to_string()))?; - - if !resp.status().is_success() { - return Err(SsoError::Discovery(format!( - "Discovery endpoint returned {}", - resp.status() - ))); - } + tracing::debug!( + provider = %self.provider_type.as_str(), + url = %discovery_url, + "Fetching OIDC discovery document" + ); - let config: OidcDiscoveryConfig = resp - .json() - .await - .map_err(|e| SsoError::Discovery(e.to_string()))?; + let resp = self + .http_client + .get(discovery_url) + .send() + .await + .map_err(|e| SsoError::Discovery(e.to_string()))?; - let jwks = match &config.jwks_uri { - Some(jwks_uri) => { - tracing::debug!( + if !resp.status().is_success() { + return Err(SsoError::Discovery(format!( + "Discovery endpoint returned {}", + resp.status() + ))); + } + + let config: OidcDiscoveryConfig = resp + .json() + .await + .map_err(|e| SsoError::Discovery(e.to_string()))?; + + let jwks = + match &config.jwks_uri { + Some(jwks_uri) => { + tracing::debug!( + provider = %self.provider_type.as_str(), + url = %jwks_uri, + "Fetching JWKS" + ); + let jwks_resp = self + .http_client + .get(jwks_uri.as_str()) + .send() + .await + .map_err(|e| SsoError::Discovery(format!("JWKS fetch failed: {}", e)))?; + + if jwks_resp.status().is_success() { + Some(jwks_resp.json::().await.map_err(|e| { + SsoError::Discovery(format!("JWKS parse failed: {}", e)) + })?) + } else { + tracing::warn!( provider = %self.provider_type.as_str(), - url = %jwks_uri, - "Fetching JWKS" + status = %jwks_resp.status(), + "JWKS fetch returned non-success status" ); - let jwks_resp = - self.http_client.get(jwks_uri).send().await.map_err(|e| { - SsoError::Discovery(format!("JWKS fetch failed: {}", e)) - })?; - - if jwks_resp.status().is_success() { - Some(jwks_resp.json::().await.map_err(|e| { - SsoError::Discovery(format!("JWKS parse failed: {}", e)) - })?) - } else { - tracing::warn!( - provider = %self.provider_type.as_str(), - status = %jwks_resp.status(), - "JWKS fetch returned non-success status" - ); - None - } + None } - None => None, - }; + } + None => None, + }; - Ok(OidcDiscoveryCache { config, jwks }) - }) - .await + Ok(OidcDiscovery { config, jwks }) } fn generate_pkce() -> PkceChallenge { @@ -602,9 +641,7 @@ impl SsoProvider for OidcProvider { let auth_endpoint = match self.provider_type { SsoProviderType::Google => "https://accounts.google.com/o/oauth2/v2/auth".to_string(), - SsoProviderType::Gitlab => { - format!("{}/oauth/authorize", self.issuer.trim_end_matches('/')) - } + SsoProviderType::Gitlab => self.issuer.endpoint("oauth/authorize").to_string(), _ => { let discovery = self.get_discovery().await?; discovery.config.authorization_endpoint.clone() @@ -638,7 +675,7 @@ impl SsoProvider for OidcProvider { ) -> Result { let token_endpoint = match self.provider_type { SsoProviderType::Google => "https://oauth2.googleapis.com/token".to_string(), - SsoProviderType::Gitlab => format!("{}/oauth/token", self.issuer.trim_end_matches('/')), + SsoProviderType::Gitlab => self.issuer.endpoint("oauth/token").to_string(), _ => { let discovery = self.get_discovery().await?; discovery.config.token_endpoint.clone() @@ -721,9 +758,7 @@ impl SsoProvider for OidcProvider { SsoProviderType::Google => { "https://openidconnect.googleapis.com/v1/userinfo".to_string() } - SsoProviderType::Gitlab => { - format!("{}/oauth/userinfo", self.issuer.trim_end_matches('/')) - } + SsoProviderType::Gitlab => self.issuer.endpoint("oauth/userinfo").to_string(), _ => { let discovery = self.get_discovery().await?; discovery @@ -777,11 +812,11 @@ pub struct AppleProvider { private_key_pem: String, http_client: Client, client_secret_cache: RwLock>, - jwks_cache: OnceCell, + cache: Arc, } impl AppleProvider { - pub fn new(config: &AppleProviderConfig) -> Result { + pub fn new(config: &AppleProviderConfig, cache: Arc) -> Result { let key_pem = config.private_key_pem.replace("\\n", "\n"); jsonwebtoken::EncodingKey::from_ec_pem(key_pem.as_bytes()) @@ -794,7 +829,7 @@ impl AppleProvider { private_key_pem: key_pem, http_client: create_http_client(), client_secret_cache: RwLock::new(None), - jwks_cache: OnceCell::new(), + cache, }) } @@ -868,29 +903,35 @@ impl AppleProvider { Ok(generated.secret) } - async fn get_jwks(&self) -> Result<&JwkSet, SsoError> { - self.jwks_cache - .get_or_try_init(|| async { - tracing::debug!("Fetching Apple JWKS"); - let resp = self - .http_client - .get("https://appleid.apple.com/auth/keys") - .send() - .await - .map_err(|e| SsoError::Discovery(format!("Apple JWKS fetch failed: {}", e)))?; - - if !resp.status().is_success() { - return Err(SsoError::Discovery(format!( - "Apple JWKS returned {}", - resp.status() - ))); - } + async fn get_jwks(&self) -> Result { + cached_json( + self.cache.as_ref(), + &sso_jwks_key(&APPLE_JWKS_URI), + SSO_DISCOVERY_TTL, + || self.fetch_jwks(), + ) + .await + } + + async fn fetch_jwks(&self) -> Result { + tracing::debug!("Fetching Apple JWKS"); + let resp = self + .http_client + .get(APPLE_JWKS_URI.as_str()) + .send() + .await + .map_err(|e| SsoError::Discovery(format!("Apple JWKS fetch failed: {}", e)))?; + + if !resp.status().is_success() { + return Err(SsoError::Discovery(format!( + "Apple JWKS returned {}", + resp.status() + ))); + } - resp.json::() - .await - .map_err(|e| SsoError::Discovery(format!("Apple JWKS parse failed: {}", e))) - }) + resp.json() .await + .map_err(|e| SsoError::Discovery(format!("Apple JWKS parse failed: {}", e))) } fn validate_id_token( @@ -1043,7 +1084,7 @@ impl SsoProvider for AppleProvider { })?; let jwks = self.get_jwks().await?; - let claims = self.validate_id_token(id_token, jwks, expected_nonce)?; + let claims = self.validate_id_token(id_token, &jwks, expected_nonce)?; tracing::debug!( sub = %claims.sub, @@ -1063,10 +1104,11 @@ impl SsoProvider for AppleProvider { #[derive(Clone)] pub struct SsoManager { providers: HashMap>, + config: &'static SsoConfig, } impl SsoManager { - pub fn from_config(config: &SsoConfig) -> Self { + pub fn from_config(config: &'static SsoConfig, cache: Arc) -> Self { let mut providers: HashMap> = HashMap::new(); if let Some(ref cfg) = config.github { @@ -1086,13 +1128,15 @@ impl SsoManager { cfg, Some("https://accounts.google.com"), "Google", + cache.clone(), ) { providers.insert(SsoProviderType::Google, Arc::new(provider)); } if let Some(ref cfg) = config.gitlab - && let Some(provider) = OidcProvider::new(SsoProviderType::Gitlab, cfg, None, "GitLab") + && let Some(provider) = + OidcProvider::new(SsoProviderType::Gitlab, cfg, None, "GitLab", cache.clone()) { providers.insert(SsoProviderType::Gitlab, Arc::new(provider)); } @@ -1103,13 +1147,14 @@ impl SsoManager { cfg, None, cfg.display_name.as_deref().unwrap_or("SSO"), + cache.clone(), ) { providers.insert(SsoProviderType::Oidc, Arc::new(provider)); } if let Some(ref cfg) = config.apple { - match AppleProvider::new(cfg) { + match AppleProvider::new(cfg, cache.clone()) { Ok(provider) => { providers.insert(SsoProviderType::Apple, Arc::new(provider)); } @@ -1119,7 +1164,11 @@ impl SsoManager { } } - Self { providers } + Self { providers, config } + } + + pub fn config(&self) -> &'static SsoConfig { + self.config } pub fn get_provider(&self, provider_type: SsoProviderType) -> Option> { @@ -1137,9 +1186,3 @@ impl SsoManager { !self.providers.is_empty() } } - -impl Default for SsoManager { - fn default() -> Self { - Self::from_config(SsoConfig::get()) - } -} diff --git a/crates/tranquil-pds/src/state.rs b/crates/tranquil-pds/src/state.rs index 66f5fd1..19da343 100644 --- a/crates/tranquil-pds/src/state.rs +++ b/crates/tranquil-pds/src/state.rs @@ -15,10 +15,12 @@ use std::error::Error; use std::path::PathBuf; use std::sync::Arc; use std::sync::atomic::{AtomicBool, Ordering}; +use std::time::Duration; use tokio::sync::broadcast; use tokio_util::sync::CancellationToken; use tranquil_db::PostgresRepositories; use tranquil_db_traits::SequencedEvent; +use tranquil_oauth::ClientMetadataCache; static RATE_LIMITING_DISABLED: AtomicBool = AtomicBool::new(false); @@ -49,6 +51,7 @@ pub struct AppState { pub sso_manager: SsoManager, pub webauthn_config: Arc, pub cross_pds_oauth: Arc, + pub client_metadata_cache: ClientMetadataCache, pub shutdown: CancellationToken, pub bootstrap_invite_code: Option, pub signal_sender: Option>, @@ -210,6 +213,27 @@ impl RateLimitKind { } } +const CLIENT_METADATA_TTL: Duration = Duration::from_secs(3600); + +struct CacheBound { + did_resolver: Arc, + cross_pds_oauth: Arc, + client_metadata_cache: ClientMetadataCache, + sso_manager: SsoManager, +} + +impl CacheBound { + fn new(cache: &Arc, sso_config: &'static SsoConfig) -> Self { + tranquil_lexicon::LexiconRegistry::global().set_shared_cache(cache.clone()); + Self { + did_resolver: Arc::new(DidResolver::new(cache.clone())), + cross_pds_oauth: Arc::new(CrossPdsOAuthClient::new(cache.clone())), + client_metadata_cache: ClientMetadataCache::new(cache.clone(), CLIENT_METADATA_TTL), + sso_manager: SsoManager::from_config(sso_config, cache.clone()), + } + } +} + impl AppState { pub fn plc_client(&self) -> PlcClient { PlcClient::with_cache(None, Some(self.cache.clone())) @@ -366,10 +390,7 @@ impl AppState { let (cache, distributed_rate_limiter) = create_cache(shutdown.clone()) .await .expect("Failed to initialize cache and distributed rate limiter at startup"); - let did_resolver = Arc::new(DidResolver::new()); - let cross_pds_oauth = Arc::new(CrossPdsOAuthClient::new(cache.clone())); - let sso_config = SsoConfig::init(); - let sso_manager = SsoManager::from_config(sso_config); + let bound = CacheBound::new(&cache, SsoConfig::init()); let webauthn_config = Arc::new( WebAuthnConfig::new(&cfg.server.hostname) .expect("Failed to create WebAuthn config at startup"), @@ -385,9 +406,10 @@ impl AppState { circuit_breakers, cache, distributed_rate_limiter, - did_resolver, - cross_pds_oauth, - sso_manager, + did_resolver: bound.did_resolver, + cross_pds_oauth: bound.cross_pds_oauth, + client_metadata_cache: bound.client_metadata_cache, + sso_manager: bound.sso_manager, webauthn_config, shutdown, bootstrap_invite_code: None, @@ -410,6 +432,11 @@ impl AppState { cache: Arc, distributed_rate_limiter: Arc, ) -> Self { + let bound = CacheBound::new(&cache, self.sso_manager.config()); + self.did_resolver = bound.did_resolver; + self.cross_pds_oauth = bound.cross_pds_oauth; + self.client_metadata_cache = bound.client_metadata_cache; + self.sso_manager = bound.sso_manager; self.cache = cache; self.distributed_rate_limiter = distributed_rate_limiter; self diff --git a/example.toml b/example.toml index 22b97b0..26c9cd2 100644 --- a/example.toml +++ b/example.toml @@ -390,7 +390,7 @@ # Default value: 5 #connect_timeout_secs = 5 -# Seconds to cache DID documents in memory. +# Seconds to cache DID documents. # # Can also be specified via environment variable `DID_CACHE_TTL_SECS`. # -- 2.51.2 From ed3d129594ff9d25f5258b6678752ad288eef393 Mon Sep 17 00:00:00 2001 From: Lewis Date: Tue, 11 Aug 2026 20:14:26 +0300 Subject: [PATCH 08/22] just: clippy over all targets, lint the bsky-off build Lewis: May this revision serve well! --- justfile | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/justfile b/justfile index 0cb9050..ebc8618 100644 --- a/justfile +++ b/justfile @@ -16,12 +16,15 @@ build-release: check: cargo check clippy: - cargo clippy -- -D warnings + cargo clippy --all-targets -- -D warnings +lint-no-bsky: + cargo clippy -p tranquil-server --no-default-features --features frontend,postgres,s3,valkey --all-targets -- -D warnings + cargo clippy -p tranquil-pds --no-default-features --all-targets -- -D warnings fmt: cargo fmt fmt-check: cargo fmt -- --check -lint: fmt-check clippy +lint: fmt-check clippy lint-no-bsky test-store: SQLX_OFFLINE=true cargo nextest run -p tranquil-store --features tranquil-store/test-harness -- 2.51.2 From 1b5a2b319c2687fc33a0486a8550becc70b50002 Mon Sep 17 00:00:00 2001 From: Louis Escher Date: Wed, 19 Aug 2026 14:28:52 +0200 Subject: [PATCH 09/22] fix: getServiceAuth aud parsing --- crates/tranquil-api/src/moderation/mod.rs | 4 +- .../tranquil-api/src/server/service_auth.rs | 6 +- crates/tranquil-auth/src/token.rs | 4 +- crates/tranquil-pds/src/api/proxy.rs | 4 +- crates/tranquil-pds/tests/jwt_security.rs | 4 +- crates/tranquil-pds/tests/oauth.rs | 46 ++++++ crates/tranquil-pds/tests/scope_edge_cases.rs | 32 ++++ crates/tranquil-pds/tests/server.rs | 40 +++++ crates/tranquil-types/src/lib.rs | 139 ++++++++++++++++++ 9 files changed, 269 insertions(+), 10 deletions(-) diff --git a/crates/tranquil-api/src/moderation/mod.rs b/crates/tranquil-api/src/moderation/mod.rs index 88045f1..27119bb 100644 --- a/crates/tranquil-api/src/moderation/mod.rs +++ b/crates/tranquil-api/src/moderation/mod.rs @@ -12,7 +12,7 @@ use tranquil_pds::api::ApiError; use tranquil_pds::api::proxy_client::{is_ssrf_safe, proxy_client}; use tranquil_pds::auth::{AnyUser, Auth}; use tranquil_pds::state::AppState; -use tranquil_pds::types::{Did, Nsid}; +use tranquil_pds::types::{Did, DidRef, Nsid}; static CREATE_REPORT_NSID: LazyLock = LazyLock::new(|| "com.atproto.moderation.createReport".parse().unwrap()); @@ -151,7 +151,7 @@ async fn proxy_to_report_service( let service_token = match tranquil_pds::auth::create_service_token( &auth_user.did, - service_did, + &DidRef::from(service_did), Some(&CREATE_REPORT_NSID), &key_bytes, ) { diff --git a/crates/tranquil-api/src/server/service_auth.rs b/crates/tranquil-api/src/server/service_auth.rs index da08672..5b66b27 100644 --- a/crates/tranquil-api/src/server/service_auth.rs +++ b/crates/tranquil-api/src/server/service_auth.rs @@ -11,7 +11,7 @@ use tracing::{error, info, warn}; use tranquil_pds::api::error::ApiError; use tranquil_pds::auth::extractor::{Auth, Permissive}; use tranquil_pds::state::AppState; -use tranquil_pds::types::Did; +use tranquil_pds::types::DidRef; use tranquil_types::Nsid; static CREATE_ACCOUNT_NSID: LazyLock = @@ -45,7 +45,7 @@ static PROTECTED_METHODS: LazyLock> = LazyLock::new(|| { #[derive(Deserialize)] pub struct GetServiceAuthParams { - pub aud: Did, + pub aud: DidRef, pub lxm: Option, pub exp: Option, } @@ -146,6 +146,8 @@ pub async fn get_service_auth( .into_response(); } + // NOTE: exp is validated here but never reaches create_service_token, which hardcodes a 60 + // second lifetime, so a client asking for longer silently gets 60 seconds if let Some(exp) = params.exp { let now = chrono::Utc::now().timestamp(); let diff = exp - now; diff --git a/crates/tranquil-auth/src/token.rs b/crates/tranquil-auth/src/token.rs index 9fee5fb..8a851cc 100644 --- a/crates/tranquil-auth/src/token.rs +++ b/crates/tranquil-auth/src/token.rs @@ -10,7 +10,7 @@ use chrono::{DateTime, Duration, Utc}; use hmac::{Hmac, Mac}; use k256::ecdsa::{Signature, SigningKey, signature::Signer}; use sha2::Sha256; -use tranquil_types::{Did, Jti, Nsid}; +use tranquil_types::{Did, DidRef, Jti, Nsid}; type HmacSha256 = Hmac; @@ -127,7 +127,7 @@ pub fn create_refresh_token_with_jti( pub fn create_service_token( did: &Did, - aud: &Did, + aud: &DidRef, lxm: Option<&Nsid>, key_bytes: &[u8], ) -> Result { diff --git a/crates/tranquil-pds/src/api/proxy.rs b/crates/tranquil-pds/src/api/proxy.rs index c51a2e3..5d5c638 100644 --- a/crates/tranquil-pds/src/api/proxy.rs +++ b/crates/tranquil-pds/src/api/proxy.rs @@ -5,7 +5,7 @@ use std::sync::LazyLock; use crate::api::error::ApiError; use crate::api::proxy_client::proxy_client; use crate::state::AppState; -use crate::types::{Did, Nsid}; +use crate::types::{Did, DidRef, Nsid}; use crate::util::get_header_str; use axum::{ body::Bytes, @@ -361,7 +361,7 @@ async fn proxy_handler( match crate::auth::create_service_token( &auth_user.did, - &token_aud, + &DidRef::from(&token_aud), Some(&token_lxm), &key_bytes, ) { diff --git a/crates/tranquil-pds/tests/jwt_security.rs b/crates/tranquil-pds/tests/jwt_security.rs index e462df3..c899b95 100644 --- a/crates/tranquil-pds/tests/jwt_security.rs +++ b/crates/tranquil-pds/tests/jwt_security.rs @@ -14,7 +14,7 @@ use tranquil_pds::auth::{ get_did_from_token, get_jti_from_token, verify_access_token, verify_refresh_token, verify_token, }; -use tranquil_types::{Did, Nsid}; +use tranquil_types::{Did, DidRef, Nsid}; fn generate_user_key() -> Vec { let secret_key = SecretKey::random(&mut OsRng); @@ -169,7 +169,7 @@ fn test_token_type_confusion() { let service_token = create_service_token( &did, - &Did::new("did:web:nel.pet").expect("valid DID"), + &DidRef::new("did:web:nel.pet").expect("valid DID reference"), Some(&Nsid::new("cafe.oyster.method").expect("valid NSID")), &key_bytes, ) diff --git a/crates/tranquil-pds/tests/oauth.rs b/crates/tranquil-pds/tests/oauth.rs index 284ac81..2a05cd3 100644 --- a/crates/tranquil-pds/tests/oauth.rs +++ b/crates/tranquil-pds/tests/oauth.rs @@ -1270,6 +1270,52 @@ async fn test_granular_scope_rpc_specific_method() { ); } +#[tokio::test] +async fn test_granular_scope_rpc_aud_with_service_id() { + let url = base_url().await; + let http_client = client(); + let (token, _, _) = + get_oauth_token_with_scope("rpc:app.bsky.feed.getTimeline?aud=did:web:api.bsky.app").await; + let allowed_res = http_client + .get(format!("{}/xrpc/com.atproto.server.getServiceAuth", url)) + .bearer_auth(&token) + .query(&[ + ("aud", "did:web:api.bsky.app#bsky_appview"), + ("lxm", "app.bsky.feed.getTimeline"), + ]) + .send() + .await + .unwrap(); + assert_eq!( + allowed_res.status(), + StatusCode::OK, + "A scope granted for a service must cover a request naming one of its service ids" + ); + let body: Value = allowed_res.json().await.unwrap(); + let service_token = body["token"].as_str().unwrap(); + let payload = service_token.split('.').nth(1).unwrap(); + let claims: Value = serde_json::from_slice(&URL_SAFE_NO_PAD.decode(payload).unwrap()).unwrap(); + assert_eq!( + claims["aud"], "did:web:api.bsky.app#bsky_appview", + "the service id must reach the signed claim even on the granular scope path" + ); + let blocked_res = http_client + .get(format!("{}/xrpc/com.atproto.server.getServiceAuth", url)) + .bearer_auth(&token) + .query(&[ + ("aud", "did:web:other.example#bsky_appview"), + ("lxm", "app.bsky.feed.getTimeline"), + ]) + .send() + .await + .unwrap(); + assert_eq!( + blocked_res.status(), + StatusCode::FORBIDDEN, + "A service id must not smuggle in a different audience" + ); +} + #[tokio::test] async fn test_oauth_metadata_includes_prompt_values_supported() { let url = base_url().await; diff --git a/crates/tranquil-pds/tests/scope_edge_cases.rs b/crates/tranquil-pds/tests/scope_edge_cases.rs index bdca3bb..9cffed4 100644 --- a/crates/tranquil-pds/tests/scope_edge_cases.rs +++ b/crates/tranquil-pds/tests/scope_edge_cases.rs @@ -181,6 +181,38 @@ fn test_permissions_rpc_lxm_wildcard_prefix() { assert!(!perms.allows_rpc("did:web:api.bsky.app", &c("app.bsky.actor.getProfile"))); } +#[test] +fn test_permissions_rpc_aud_service_id_is_normalized() { + let perms = + ScopePermissions::from_scope_string(Some("rpc:app.bsky.feed.*?aud=did:web:api.bsky.app")); + assert!( + perms.allows_rpc( + "did:web:api.bsky.app#bsky_appview", + &c("app.bsky.feed.getTimeline") + ), + "a scope granted for a service must cover a request naming one of its service ids" + ); + assert!( + perms.allows_rpc("did:web:api.bsky.app", &c("app.bsky.feed.getTimeline")), + "the bare form must keep working" + ); + assert!( + !perms.allows_rpc( + "did:web:other.example#bsky_appview", + &c("app.bsky.feed.getTimeline") + ), + "a service id must not smuggle in a different audience" + ); + + let fragment_scope = ScopePermissions::from_scope_string(Some( + "rpc:app.bsky.feed.*?aud=did:web:api.bsky.app%23bsky_appview", + )); + assert!( + fragment_scope.allows_rpc("did:web:api.bsky.app", &c("app.bsky.feed.getTimeline")), + "a scope granted with a service id must still cover the bare audience" + ); +} + #[test] fn test_delegation_intersect_mismatched_params_empty() { let result = intersect_scopes("repo:*?action=create", "repo:*?action=delete"); diff --git a/crates/tranquil-pds/tests/server.rs b/crates/tranquil-pds/tests/server.rs index 6cd861a..4951705 100644 --- a/crates/tranquil-pds/tests/server.rs +++ b/crates/tranquil-pds/tests/server.rs @@ -132,6 +132,46 @@ async fn test_service_auth() { let lxm_payload = URL_SAFE_NO_PAD.decode(lxm_parts[1]).unwrap(); let lxm_claims: Value = serde_json::from_slice(&lxm_payload).unwrap(); assert_eq!(lxm_claims["lxm"], "com.atproto.repo.getRecord"); + let fragment_res = client + .get(format!("{}/xrpc/com.atproto.server.getServiceAuth", base)) + .bearer_auth(&access_jwt) + .query(&[ + ("aud", "did:web:example.com#colibri_appview"), + ("lxm", "com.atproto.repo.getRecord"), + ]) + .send() + .await + .unwrap(); + assert_eq!(fragment_res.status(), StatusCode::OK); + let fragment_body: Value = fragment_res.json().await.unwrap(); + let fragment_token = fragment_body["token"].as_str().unwrap(); + let fragment_parts: Vec<&str> = fragment_token.split('.').collect(); + let fragment_payload = URL_SAFE_NO_PAD.decode(fragment_parts[1]).unwrap(); + let fragment_claims: Value = serde_json::from_slice(&fragment_payload).unwrap(); + assert_eq!( + fragment_claims["aud"], "did:web:example.com#colibri_appview", + "the service id must survive into the signed claim so the receiver can match it \ + against its own DID document" + ); + + let empty_fragment = client + .get(format!("{}/xrpc/com.atproto.server.getServiceAuth", base)) + .bearer_auth(&access_jwt) + .query(&[("aud", "did:web:example.com#")]) + .send() + .await + .unwrap(); + assert_eq!(empty_fragment.status(), StatusCode::BAD_REQUEST); + + let double_fragment = client + .get(format!("{}/xrpc/com.atproto.server.getServiceAuth", base)) + .bearer_auth(&access_jwt) + .query(&[("aud", "did:web:example.com#a#b")]) + .send() + .await + .unwrap(); + assert_eq!(double_fragment.status(), StatusCode::BAD_REQUEST); + let unauth = client .get(format!("{}/xrpc/com.atproto.server.getServiceAuth", base)) .query(&[("aud", "did:web:example.com")]) diff --git a/crates/tranquil-types/src/lib.rs b/crates/tranquil-types/src/lib.rs index 6fe70a8..eea83df 100644 --- a/crates/tranquil-types/src/lib.rs +++ b/crates/tranquil-types/src/lib.rs @@ -225,6 +225,54 @@ impl Did { } } +const DID_REF_MAX_LEN: usize = 2048; + +fn is_service_id(s: &str) -> bool { + !s.is_empty() + && !s + .chars() + .any(|c| c.is_whitespace() || c.is_control() || matches!(c, '#' | '/' | '?')) +} + +validated_string_newtype! { + pub struct DidRef; + error = DidRefError; + label = "DID reference"; + validator = |s| { + if s.len() > DID_REF_MAX_LEN { + return Err(()); + } + match s.split_once('#') { + None => jacquard_common::types::string::Did::new(s) + .map(|v| v.as_str().to_owned()) + .map_err(|_| ()), + Some((did, service_id)) => { + if !is_service_id(service_id) { + return Err(()); + } + let base = jacquard_common::types::string::Did::new(did).map_err(|_| ())?; + Ok(format!("{}#{}", base.as_str(), service_id)) + } + } + }; +} + +impl DidRef { + pub fn did(&self) -> &str { + self.0.split('#').next().unwrap_or(&self.0) + } + + pub fn service_id(&self) -> Option<&str> { + self.0.split_once('#').map(|(_, service_id)| service_id) + } +} + +impl From<&Did> for DidRef { + fn from(did: &Did) -> Self { + Self(did.0.clone()) + } +} + #[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, sqlx::Type)] #[serde(transparent)] #[sqlx(transparent)] @@ -1588,6 +1636,97 @@ mod validated_newtype_tests { ); } + #[test] + fn a_bare_did_ref_names_no_service() { + let aud = DidRef::new("did:plc:abc").unwrap(); + assert_eq!(aud.as_str(), "did:plc:abc"); + assert_eq!(aud.did(), "did:plc:abc"); + assert_eq!( + aud.service_id(), + None, + "an absent fragment is not the same as an empty one" + ); + } + + #[test] + fn a_did_ref_keeps_the_service_id_it_was_given() { + let aud = DidRef::new("did:web:api.colibri.social#colibri_appview").unwrap(); + assert_eq!( + aud.as_str(), + "did:web:api.colibri.social#colibri_appview", + "the fragment is what tells the receiver which of its services was audienced, \ + so it must survive entirely" + ); + assert_eq!(aud.did(), "did:web:api.colibri.social"); + assert_eq!(aud.service_id(), Some("colibri_appview")); + } + + #[test] + fn a_did_ref_normalizes_its_did_half_the_way_a_did_does() { + assert_eq!( + DidRef::new("at://did:plc:abc#colibri_appview") + .unwrap() + .as_str(), + "did:plc:abc#colibri_appview" + ); + assert_eq!( + DidRef::new("did:plc:def").unwrap().as_str(), + Did::new("did:plc:def").unwrap().as_str(), + "a fragmentless DidRef must be byte-identical to the Did it replaces" + ); + } + + #[test] + fn a_did_ref_rejects_anything_that_cannot_name_one_service() { + for bad in [ + "did:web:oyster.cafe#", + "did:web:oyster.cafe#a#b", + "did:web:oyster.cafe# whelk", + "did:web:oyster.cafe#a/b", + "did:web:oyster.cafe#a?b", + "not-a-did#colibri_appview", + "#colibri_appview", + ] { + assert!( + DidRef::new(bad).is_err(), + "{bad} should not parse as a DID reference" + ); + } + } + + #[test] + fn a_did_ref_does_not_second_guess_the_service_ids_it_has_not_seen() { + for good in [ + "did:web:oyster.cafe#atproto_pds", + "did:web:oyster.cafe#atproto_labeler", + "did:plc:abc#bsky_chat", + "did:web:oyster.cafe#whelk.v2", + ] { + assert!( + DidRef::new(good).is_ok(), + "{good} names a service the receiver resolves in its own DID document, \ + so rejecting it here would recreate the bug this type exists to fix" + ); + } + } + + #[test] + fn an_over_long_did_ref_is_rejected() { + let long = format!("did:web:{}#def", "a".repeat(2048)); + assert!( + DidRef::new(&long).is_err(), + "the lexicon bounds aud at 2048 bytes" + ); + } + + #[test] + fn a_did_ref_built_from_a_did_names_no_service() { + let did = Did::new("did:plc:def").unwrap(); + let aud = DidRef::from(&did); + assert_eq!(aud.as_str(), did.as_str()); + assert_eq!(aud.service_id(), None); + } + #[test] fn the_earliest_tid_sorts_below_every_generated_tid() { let earliest = Tid::earliest(); -- 2.51.2 From 32c58b1d0b487efc57ebd106e077684b7cae0f2c Mon Sep 17 00:00:00 2001 From: Louis Escher Date: Wed, 19 Aug 2026 16:29:04 +0200 Subject: [PATCH 10/22] fix: make thingy allow list --- crates/tranquil-types/src/lib.rs | 51 +++++++++++++++++++++++++++++--- 1 file changed, 47 insertions(+), 4 deletions(-) diff --git a/crates/tranquil-types/src/lib.rs b/crates/tranquil-types/src/lib.rs index eea83df..3703bff 100644 --- a/crates/tranquil-types/src/lib.rs +++ b/crates/tranquil-types/src/lib.rs @@ -226,12 +226,33 @@ impl Did { } const DID_REF_MAX_LEN: usize = 2048; +const SERVICE_ID_MAX_LEN: usize = 128; + +const fn is_pchar(b: u8) -> bool { + b.is_ascii_alphanumeric() + || matches!( + b, + b'-' | b'.' + | b'_' + | b'~' + | b'!' + | b'$' + | b'&' + | b'\'' + | b'(' + | b')' + | b'*' + | b'+' + | b',' + | b';' + | b'=' + | b':' + | b'@' + ) +} fn is_service_id(s: &str) -> bool { - !s.is_empty() - && !s - .chars() - .any(|c| c.is_whitespace() || c.is_control() || matches!(c, '#' | '/' | '?')) + !s.is_empty() && s.len() <= SERVICE_ID_MAX_LEN && s.bytes().all(is_pchar) } validated_string_newtype! { @@ -1684,6 +1705,13 @@ mod validated_newtype_tests { "did:web:oyster.cafe# whelk", "did:web:oyster.cafe#a/b", "did:web:oyster.cafe#a?b", + "did:web:oyster.cafe#