From 8d0b6f8322c9e964a75321474f33bfb394f35bf3 Mon Sep 17 00:00:00 2001 From: Lewis Date: Sun, 26 Jul 2026 13:57:07 +0300 Subject: [PATCH] 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 | 269 ++++++++++-------- crates/tranquil-pds/src/state.rs | 41 ++- example.toml | 2 +- 13 files changed, 380 insertions(+), 429 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://") { + 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( - "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") - { - 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()))?; + tracing::debug!( + provider = %self.provider_type.as_str(), + url = %discovery_url, + "Fetching OIDC discovery document" + ); - 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).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, - }; - - Ok(OidcDiscoveryCache { config, jwks }) - }) + 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() + ))); + } + + 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(), + status = %jwks_resp.status(), + "JWKS fetch returned non-success status" + ); + None + } + } + None => None, + }; + + 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)))?; + 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 + } - 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))) - }) + 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))) } 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`. #