diff --git a/crates/tranquil-api/src/admin/status.rs b/crates/tranquil-api/src/admin/status.rs index 86f3b45..9b6c439 100644 --- a/crates/tranquil-api/src/admin/status.rs +++ b/crates/tranquil-api/src/admin/status.rs @@ -199,6 +199,19 @@ pub async fn update_subject_status( let result = if deactivated.applied { state.repos.user.deactivate_account(&did, None).await } else { + let status = state + .repos + .user + .get_status_by_did(&did) + .await + .map_err(|e| { + error!("Failed to read account status for {}: {:?}", did, e); + ApiError::InternalError(None) + })? + .ok_or(ApiError::AccountNotFound)?; + if status.takedown_ref.is_some() { + return Err(ApiError::AccountNotFound); + } state.repos.user.activate_account(&did).await }; result.map_err(|e| { @@ -289,19 +302,24 @@ pub async fn update_subject_status( Some("com.atproto.repo.strongRef") => { let uri_str = input.subject.get("uri").and_then(Value::as_str); if let Some(uri_str) = uri_str { - let cid: CidLink = uri_str + let cid_str = input.subject.get("cid").and_then(Value::as_str).ok_or_else(|| { + ApiError::InvalidRequest("Record subject must include a CID".into()) + })?; + let cid: CidLink = cid_str .parse() .map_err(|_| ApiError::InvalidRequest("Invalid CID format".into()))?; + let mut applied_ref = None; if let Some(takedown) = &input.takedown { - let takedown_ref = if takedown.applied { - takedown.r#ref.as_deref() - } else { - None - }; + applied_ref = takedown.applied.then(|| { + takedown + .r#ref + .clone() + .unwrap_or_else(|| Utc::now().to_rfc3339()) + }); state .repos .repo - .set_record_takedown(&cid, takedown_ref) + .set_record_takedown(&cid, applied_ref.as_deref()) .await .map_err(|e| { error!( @@ -315,7 +333,7 @@ pub async fn update_subject_status( "subject": input.subject, "takedown": input.takedown.as_ref().map(|t| json!({ "applied": t.applied, - "ref": t.r#ref + "ref": applied_ref })) }))); } @@ -326,16 +344,18 @@ pub async fn update_subject_status( let cid: CidLink = cid_str .parse() .map_err(|_| ApiError::InvalidRequest("Invalid CID format".into()))?; + let mut applied_ref = None; if let Some(takedown) = &input.takedown { - let takedown_ref = if takedown.applied { - takedown.r#ref.as_deref() - } else { - None - }; + applied_ref = takedown.applied.then(|| { + takedown + .r#ref + .clone() + .unwrap_or_else(|| Utc::now().to_rfc3339()) + }); state .repos .blob - .update_blob_takedown(&cid, takedown_ref) + .update_blob_takedown(&cid, applied_ref.as_deref()) .await .map_err(|e| { error!( @@ -349,7 +369,7 @@ pub async fn update_subject_status( "subject": input.subject, "takedown": input.takedown.as_ref().map(|t| json!({ "applied": t.applied, - "ref": t.r#ref + "ref": applied_ref })) }))); } diff --git a/crates/tranquil-api/src/moderation/mod.rs b/crates/tranquil-api/src/moderation/mod.rs index 88045f1..aff41ad 100644 --- a/crates/tranquil-api/src/moderation/mod.rs +++ b/crates/tranquil-api/src/moderation/mod.rs @@ -113,7 +113,7 @@ async fn proxy_to_report_service( service_did: &Did, input: &CreateReportInput, ) -> Response { - if let Err(e) = is_ssrf_safe(service_url) { + if let Err(e) = is_ssrf_safe(service_url).await { error!("Report service URL failed SSRF check: {:?}", e); return ApiError::InternalError(Some("Invalid report service configuration".into())) .into_response(); diff --git a/crates/tranquil-api/src/repo/record/read.rs b/crates/tranquil-api/src/repo/record/read.rs index 76dfb7c..c273cf9 100644 --- a/crates/tranquil-api/src/repo/record/read.rs +++ b/crates/tranquil-api/src/repo/record/read.rs @@ -8,6 +8,7 @@ use axum::{ }; use base64::Engine; use cid::Cid; +use futures::{StreamExt, stream}; use ipld_core::ipld::Ipld; use jacquard_repo::storage::BlockStore; use serde::{Deserialize, Serialize}; @@ -85,6 +86,19 @@ pub async fn get_record( return ApiError::RecordNotFound.into_response(); } }; + match state + .repos + .repo + .get_record_by_cid(&record_cid_link) + .await + { + Ok(Some(record)) if record.takedown_ref.is_none() => {} + Ok(_) => return ApiError::RecordNotFound.into_response(), + Err(e) => { + error!("Error checking record moderation status: {:?}", e); + return ApiError::InternalError(None).into_response(); + } + } let record_cid_str = record_cid_link.to_string(); if let Some(expected_cid) = &input.cid && &record_cid_str != expected_cid @@ -179,6 +193,27 @@ pub async fn list_records( } }; let last_rkey = rows.last().map(|r| r.rkey.to_string()); + let checked_rows = stream::iter(rows.into_iter().map(|row| { + let repo = state.repos.repo.as_ref(); + async move { + let moderation = repo.get_record_by_cid(&row.record_cid).await; + (row, moderation) + } + })) + .buffered(8) + .collect::>() + .await; + let mut rows = Vec::with_capacity(checked_rows.len()); + for (row, moderation) in checked_rows { + match moderation { + Ok(Some(record)) if record.takedown_ref.is_none() => rows.push(row), + Ok(_) => {} + Err(e) => { + error!("Error checking record moderation status: {:?}", e); + return ApiError::InternalError(None).into_response(); + } + } + } let parsed_rows: Vec<(Cid, String, String)> = rows .iter() .filter_map(|row| { diff --git a/crates/tranquil-api/src/server/account_status.rs b/crates/tranquil-api/src/server/account_status.rs index 2c2e806..de83896 100644 --- a/crates/tranquil-api/src/server/account_status.rs +++ b/crates/tranquil-api/src/server/account_status.rs @@ -55,14 +55,7 @@ pub async fn check_account_status( .await .log_db_err("fetching user ID for account status")? .ok_or(ApiError::InternalError(None))?; - let is_active = state - .repos - .user - .is_account_active_by_did(did) - .await - .ok() - .flatten() - .unwrap_or(false); + let is_active = auth.status.is_active(); let repo_info = state.repos.repo.get_repo(user_id).await.ok().flatten(); let (repo_commit, repo_rev_from_db) = repo_info .map(|r| (r.repo_root_cid.to_string(), r.repo_rev)) @@ -317,7 +310,7 @@ async fn assert_valid_did_document_for_service( pub async fn activate_account( State(state): State, - auth: Auth, + auth: Auth, ) -> Result, ApiError> { info!("[MIGRATION] activateAccount called"); info!( @@ -372,6 +365,25 @@ pub async fn activate_account( let result = state.repos.user.activate_account(&did).await; match result { Ok(_) => { + let final_status = state + .repos + .user + .get_status_by_did(&did) + .await + .map_err(|e| { + error!("Failed to read status after activating {}: {:?}", did, e); + ApiError::InternalError(None) + })? + .map(|status| { + tranquil_db_traits::AccountStatus::from_db_fields( + status.takedown_ref.as_deref(), + status.deactivated_at, + ) + }) + .ok_or(ApiError::AccountNotFound)?; + if !final_status.is_active() { + return Err(ApiError::AccountNotFound); + } info!( "[MIGRATION] activateAccount: DB update success for did={}", did @@ -390,6 +402,7 @@ pub async fn activate_account( .cache .delete(&tranquil_pds::cache_keys::plc_data_key(&did)) .await; + tranquil_pds::auth::invalidate_auth_cache(state.cache.as_ref(), &did).await; if state.did_resolver.refresh_did(&did).await.is_err() { warn!( "[MIGRATION] activateAccount: Failed to refresh DID cache for {}", @@ -411,7 +424,7 @@ pub async fn activate_account( if let Err(e) = tranquil_pds::repo_ops::sequence_account_event( &state, &did, - tranquil_db_traits::AccountStatus::Active, + final_status, ) .await { @@ -530,12 +543,29 @@ pub async fn deactivate_account( match result { Ok(true) => { + let final_status = state + .repos + .user + .get_status_by_did(&did) + .await + .map_err(|e| { + error!("Failed to read status after deactivating {}: {:?}", did, e); + ApiError::InternalError(None) + })? + .map(|status| { + tranquil_db_traits::AccountStatus::from_db_fields( + status.takedown_ref.as_deref(), + status.deactivated_at, + ) + }) + .ok_or(ApiError::AccountNotFound)?; if let Some(ref h) = handle { let _ = state .cache .delete(&tranquil_pds::cache_keys::handle_key(h)) .await; } + tranquil_pds::auth::invalidate_auth_cache(state.cache.as_ref(), &did).await; if let Err(e) = state .repos .repo @@ -547,7 +577,7 @@ pub async fn deactivate_account( if let Err(e) = tranquil_pds::repo_ops::sequence_account_event( &state, &did, - tranquil_db_traits::AccountStatus::Deactivated, + final_status, ) .await { @@ -555,7 +585,7 @@ pub async fn deactivate_account( } Ok(Json(EmptyResponse {})) } - Ok(false) => Ok(Json(EmptyResponse {})), + Ok(false) => Err(ApiError::AccountNotFound), Err(e) => { error!("DB error deactivating account: {:?}", e); Err(ApiError::InternalError(None)) diff --git a/crates/tranquil-oauth-server/src/endpoints/authorize/mod.rs b/crates/tranquil-oauth-server/src/endpoints/authorize/mod.rs index 2b37337..4a5e286 100644 --- a/crates/tranquil-oauth-server/src/endpoints/authorize/mod.rs +++ b/crates/tranquil-oauth-server/src/endpoints/authorize/mod.rs @@ -7,8 +7,11 @@ use axum::{ }, response::{IntoResponse, Response}, }; +use base64::{Engine as _, engine::general_purpose::URL_SAFE_NO_PAD}; use chrono::Utc; +use hmac::{Hmac, Mac}; use serde::{Deserialize, Serialize}; +use sha2::Sha256; use subtle::ConstantTimeEq; use tranquil_db_traits::{ScopePreference, WebauthnChallengeType}; use tranquil_pds::auth::{BareLoginIdentifier, NormalizedLoginIdentifier}; @@ -205,18 +208,56 @@ fn build_intermediate_redirect_url( if let Some(rm) = response_mode { url.push_str(&format!("&response_mode={}", url_encode(rm))); } + let signature = sign_redirect_parameters(redirect_uri, code, state, response_mode); + url.push_str(&format!("&signature={}", url_encode(&signature))); url } +fn sign_redirect_parameters( + redirect_uri: &str, + code: &str, + state: Option<&str>, + response_mode: Option<&str>, +) -> String { + let payload = serde_json::to_vec(&(redirect_uri, code, state, response_mode)) + .expect("redirect parameters are serializable"); + let mut mac = Hmac::::new_from_slice( + tranquil_pds::config::AuthConfig::get().dpop_secret().as_bytes(), + ) + .expect("HMAC accepts keys of any size"); + mac.update(&payload); + URL_SAFE_NO_PAD.encode(mac.finalize().into_bytes()) +} + #[derive(Debug, Deserialize)] pub struct AuthorizeRedirectParams { redirect_uri: String, code: String, state: Option, response_mode: Option, + signature: String, } pub async fn authorize_redirect(Query(params): Query) -> Response { + let expected_signature = sign_redirect_parameters( + ¶ms.redirect_uri, + ¶ms.code, + params.state.as_deref(), + params.response_mode.as_deref(), + ); + if params + .signature + .as_bytes() + .ct_eq(expected_signature.as_bytes()) + .unwrap_u8() + != 1 + { + return json_error( + StatusCode::BAD_REQUEST, + "invalid_request", + "Invalid authorization redirect", + ); + } let final_url = build_success_redirect( ¶ms.redirect_uri, ¶ms.code, diff --git a/crates/tranquil-oauth-server/src/endpoints/token/grants.rs b/crates/tranquil-oauth-server/src/endpoints/token/grants.rs index 543f71f..860ea97 100644 --- a/crates/tranquil-oauth-server/src/endpoints/token/grants.rs +++ b/crates/tranquil-oauth-server/src/endpoints/token/grants.rs @@ -19,6 +19,53 @@ const ACCESS_TOKEN_EXPIRY_SECONDS: u64 = 300; const REFRESH_TOKEN_EXPIRY_DAYS_CONFIDENTIAL: i64 = 60; const REFRESH_TOKEN_EXPIRY_DAYS_PUBLIC: i64 = 14; +async fn verify_request_client_auth( + expected_client_id: &tranquil_types::ClientId, + request_auth: &RequestClientAuth, +) -> Result<(), OAuthError> { + let request_client_id = request_auth.client_id().ok_or_else(|| { + OAuthError::InvalidClient("client_id is required".to_string()) + })?; + if request_client_id != expected_client_id.as_str() { + return Err(OAuthError::InvalidClient("client_id mismatch".to_string())); + } + + let client_auth = match request_auth { + RequestClientAuth::PrivateKeyJwt { + assertion, + assertion_type, + .. + } => { + if assertion_type != "urn:ietf:params:oauth:client-assertion-type:jwt-bearer" { + return Err(OAuthError::InvalidClient( + "Unsupported client_assertion_type".to_string(), + )); + } + ClientAuth::PrivateKeyJwt { + client_assertion: assertion.clone(), + } + } + RequestClientAuth::SecretPost { client_secret, .. } => ClientAuth::SecretPost { + client_secret: client_secret.clone(), + }, + RequestClientAuth::None { .. } => ClientAuth::None, + }; + + let client_metadata_cache = ClientMetadataCache::new(3600); + let client_metadata = client_metadata_cache.get(expected_client_id).await?; + let token_endpoint = format!( + "https://{}/oauth/token", + tranquil_config::get().server.hostname + ); + verify_client_auth( + &client_metadata_cache, + &client_metadata, + &client_auth, + &token_endpoint, + ) + .await +} + pub async fn handle_authorization_code_grant( state: AppState, _headers: HeaderMap, @@ -57,39 +104,15 @@ pub async fn handle_authorization_code_grant( .require_authorized() .map_err(|_| OAuthError::InvalidGrant("Authorization not completed".to_string()))?; - if let Some(request_client_id) = request.client_auth.client_id() - && request_client_id != authorized.client_id - { - return Err(OAuthError::InvalidGrant("client_id mismatch".to_string())); - } + verify_request_client_auth(&authorized.client_id, &request.client_auth).await?; let did = authorized.did.clone(); - let client_metadata_cache = ClientMetadataCache::new(3600); - let client_metadata = client_metadata_cache.get(&authorized.client_id).await?; - let client_auth = match &request.client_auth { - RequestClientAuth::PrivateKeyJwt { - assertion, - assertion_type, - .. - } => { - if assertion_type != "urn:ietf:params:oauth:client-assertion-type:jwt-bearer" { - return Err(OAuthError::InvalidClient( - "Unsupported client_assertion_type".to_string(), - )); - } - ClientAuth::PrivateKeyJwt { - client_assertion: assertion.clone(), - } - } - RequestClientAuth::SecretPost { client_secret, .. } => ClientAuth::SecretPost { - client_secret: client_secret.clone(), - }, - RequestClientAuth::None { .. } => ClientAuth::None, - }; - 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 - { + let redirect_uri = redirect_uri.ok_or_else(|| { + OAuthError::InvalidRequest( + "redirect_uri is required for authorization_code grant".to_string(), + ) + })?; + if redirect_uri != authorized.parameters.redirect_uri { return Err(OAuthError::InvalidGrant( "redirect_uri mismatch".to_string(), )); @@ -288,6 +311,34 @@ async fn recompute_resolved_scope( Ok(effective.permitted) } +async fn reject_takendown_token_family( + state: &AppState, + db_id: tranquil_db_traits::TokenFamilyId, + token_data: &TokenData, +) -> Result<(), OAuthError> { + let status = state + .repos + .user + .get_status_by_did(&token_data.did) + .await + .map_err(tranquil_pds::oauth::db_err_to_oauth)?; + if status + .as_ref() + .is_none_or(|status| status.takedown_ref.is_some()) + { + state + .repos + .oauth + .delete_token_family(db_id) + .await + .map_err(tranquil_pds::oauth::db_err_to_oauth)?; + return Err(OAuthError::InvalidGrant( + "Invalid refresh token".to_string(), + )); + } + Ok(()) +} + pub async fn handle_refresh_token_grant( state: AppState, _headers: HeaderMap, @@ -317,10 +368,12 @@ pub async fn handle_refresh_token_grant( let (db_id, token_data) = match lookup { RefreshTokenLookup::Valid { db_id, token_data } => (db_id, token_data), RefreshTokenLookup::InGracePeriod { - db_id: _, + db_id, token_data, rotated_at, } => { + reject_takendown_token_family(&state, db_id, &token_data).await?; + verify_request_client_auth(&token_data.client_id, &request.client_auth).await?; tracing::info!( refresh_token_prefix = %token_prefix, rotated_at = %rotated_at, @@ -392,6 +445,8 @@ pub async fn handle_refresh_token_grant( )); } }; + reject_takendown_token_family(&state, db_id, &token_data).await?; + verify_request_client_auth(&token_data.client_id, &request.client_auth).await?; let dpop_jkt = if let Some(proof) = &dpop_proof { let config = AuthConfig::get(); let verifier = DPoPVerifier::new(config.dpop_secret().as_bytes()); diff --git a/crates/tranquil-oauth/src/client.rs b/crates/tranquil-oauth/src/client.rs index cb4b588..7212b62 100644 --- a/crates/tranquil-oauth/src/client.rs +++ b/crates/tranquil-oauth/src/client.rs @@ -1,6 +1,7 @@ use reqwest::Client; use serde::{Deserialize, Serialize}; use std::collections::HashMap; +use std::net::{IpAddr, Ipv4Addr, Ipv6Addr}; use std::sync::Arc; use tokio::sync::RwLock; @@ -74,6 +75,79 @@ struct CachedJwks { cached_at: std::time::Instant, } +async fn validate_outbound_url(url: &str, field: &str) -> Result<(), OAuthError> { + let parsed = reqwest::Url::parse(url) + .map_err(|_| OAuthError::InvalidClient(format!("{field} must be a valid URL")))?; + if parsed.scheme() != "https" { + return Err(OAuthError::InvalidClient(format!( + "{field} must use https" + ))); + } + if !parsed.username().is_empty() || parsed.password().is_some() { + return Err(OAuthError::InvalidClient(format!( + "{field} must not contain credentials" + ))); + } + let host = parsed + .host_str() + .ok_or_else(|| OAuthError::InvalidClient(format!("{field} must include a host")))?; + if let Ok(ip) = host.parse::() { + return is_public_ip(&ip) + .then_some(()) + .ok_or_else(|| OAuthError::InvalidClient(format!("{field} host is not public"))); + } + let port = parsed.port().unwrap_or(443); + let addresses: Vec = tokio::net::lookup_host((host, port)) + .await + .map_err(|_| OAuthError::InvalidClient(format!("{field} host could not be resolved")))? + .collect(); + if addresses.is_empty() || addresses.iter().any(|address| !is_public_ip(&address.ip())) { + return Err(OAuthError::InvalidClient(format!( + "{field} host is not public" + ))); + } + Ok(()) +} + +fn is_public_ip(ip: &IpAddr) -> bool { + match ip { + IpAddr::V4(ip) => is_public_ipv4(ip), + IpAddr::V6(ip) => is_public_ipv6(ip), + } +} + +fn is_public_ipv4(ip: &Ipv4Addr) -> bool { + let [a, b, c, _] = ip.octets(); + !(a == 0 + || a == 10 + || (a == 100 && (64..=127).contains(&b)) + || a == 127 + || (a == 169 && b == 254) + || (a == 172 && (16..=31).contains(&b)) + || (a == 192 && b == 0 && c == 0) + || (a == 192 && b == 0 && c == 2) + || (a == 192 && b == 88 && c == 99) + || (a == 192 && b == 168) + || (a == 198 && (b == 18 || b == 19)) + || (a == 198 && b == 51 && c == 100) + || (a == 203 && b == 0 && c == 113) + || a >= 224) +} + +fn is_public_ipv6(ip: &Ipv6Addr) -> bool { + let segments = ip.segments(); + if segments[..6] == [0, 0, 0, 0, 0, 0xffff] { + let high = segments[6].to_be_bytes(); + let low = segments[7].to_be_bytes(); + return is_public_ipv4(&Ipv4Addr::new(high[0], high[1], low[0], low[1])); + } + (segments[0] & 0xe000) == 0x2000 + && !(segments[0] == 0x2001 && segments[1] < 0x0200) + && !(segments[0] == 0x2001 && segments[1] == 0x0db8) + && segments[0] != 0x2002 + && !(segments[0] == 0x3fff && segments[1] < 0x1000) +} + impl ClientMetadataCache { pub fn new(cache_ttl_secs: u64) -> Self { Self { @@ -85,6 +159,7 @@ impl ClientMetadataCache { .connect_timeout(std::time::Duration::from_secs(10)) .pool_max_idle_per_host(10) .pool_idle_timeout(std::time::Duration::from_secs(90)) + .redirect(reqwest::redirect::Policy::none()) .user_agent(concat!( "Tranquil-PDS/", env!("CARGO_PKG_VERSION"), @@ -207,14 +282,7 @@ impl ClientMetadataCache { } 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(), - )); - } + validate_outbound_url(jwks_uri, "jwks_uri").await?; let response = self .http_client .get(jwks_uri) @@ -243,19 +311,7 @@ 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") - { - return Err(OAuthError::InvalidClient( - "Non-localhost client_id must use https".to_string(), - )); - } + validate_outbound_url(client_id.as_str(), "client_id").await?; let response = self .http_client .get(client_id.as_str()) @@ -393,6 +449,7 @@ pub async fn verify_client_auth( cache: &ClientMetadataCache, metadata: &ClientMetadata, client_auth: &ClientAuth, + expected_audience: &str, ) -> Result<(), OAuthError> { let expected_method = metadata.auth_method(); match (expected_method, client_auth) { @@ -401,7 +458,7 @@ pub async fn verify_client_auth( "Client is configured for no authentication, but credentials were provided".to_string(), )), ("private_key_jwt", ClientAuth::PrivateKeyJwt { client_assertion }) => { - verify_private_key_jwt_async(cache, metadata, client_assertion).await + verify_private_key_jwt_async(cache, metadata, client_assertion, expected_audience).await } ("private_key_jwt", _) => Err(OAuthError::InvalidClient( "Client requires private_key_jwt authentication".to_string(), @@ -423,6 +480,7 @@ async fn verify_private_key_jwt_async( cache: &ClientMetadataCache, metadata: &ClientMetadata, client_assertion: &str, + expected_audience: &str, ) -> Result<(), OAuthError> { use base64::{ Engine as _, @@ -481,30 +539,32 @@ async fn verify_private_key_jwt_async( "client_assertion sub does not match client_id".to_string(), )); } + let audience_matches = match payload.get("aud") { + Some(serde_json::Value::String(audience)) => audience == expected_audience, + Some(serde_json::Value::Array(audiences)) => audiences + .iter() + .any(|audience| audience.as_str() == Some(expected_audience)), + _ => false, + }; + if !audience_matches { + return Err(OAuthError::InvalidClient( + "client_assertion aud does not match token endpoint".to_string(), + )); + } let now = chrono::Utc::now().timestamp(); - let exp = payload.get("exp").and_then(|e| e.as_i64()); + let exp = payload + .get("exp") + .and_then(|e| e.as_i64()) + .ok_or_else(|| OAuthError::InvalidClient("Missing exp in client_assertion".to_string()))?; let iat = payload.get("iat").and_then(|i| i.as_i64()); - if let Some(exp) = exp { - if exp < now { - return Err(OAuthError::InvalidClient( - "client_assertion has expired".to_string(), - )); - } - } else if let Some(iat) = iat { - let max_age_secs = 300; - if now - iat > max_age_secs { - tracing::warn!( - iat = iat, - now = now, - "client_assertion too old (no exp, using iat)" - ); - return Err(OAuthError::InvalidClient( - "client_assertion is too old".to_string(), - )); - } - } else { + if exp < now { + return Err(OAuthError::InvalidClient( + "client_assertion has expired".to_string(), + )); + } + if exp > now + 300 { return Err(OAuthError::InvalidClient( - "client_assertion must have exp or iat claim".to_string(), + "client_assertion lifetime is too long".to_string(), )); } if let Some(iat) = iat diff --git a/crates/tranquil-pds/src/api/proxy.rs b/crates/tranquil-pds/src/api/proxy.rs index 5e2af72..3853d3d 100644 --- a/crates/tranquil-pds/src/api/proxy.rs +++ b/crates/tranquil-pds/src/api/proxy.rs @@ -3,7 +3,7 @@ use std::convert::Infallible; use std::sync::LazyLock; use crate::api::error::ApiError; -use crate::api::proxy_client::proxy_client; +use crate::api::proxy_client::{is_ssrf_safe, proxy_client}; use crate::state::AppState; use crate::types::{Did, Nsid}; use crate::util::get_header_str; @@ -129,9 +129,11 @@ async fn resolve_feed_generator_did(appview_url: &str, query: Option<&str>) -> O let repo = at_uri.did()?; let collection = at_uri.collection()?; let rkey = at_uri.rkey()?; + let target_url = format!("{appview_url}/xrpc/com.atproto.repo.getRecord"); + is_ssrf_safe(&target_url).await.ok()?; let resp = proxy_client() - .get(format!("{appview_url}/xrpc/com.atproto.repo.getRecord")) + .get(target_url) .query(&[("repo", repo), ("collection", collection), ("rkey", rkey)]) .send() .await @@ -263,11 +265,6 @@ async fn proxy_handler( Some(q) => format!("{}/xrpc/{}?{}", resolved.url, method, q), None => format!("{}/xrpc/{}", resolved.url, method), }; - info!("Proxying {} request to {}", method_verb, target_url); - - let client = proxy_client(); - let mut request_builder = client.request(method_verb.clone(), &target_url); - let mut auth_header_val = headers.get(http::header::AUTHORIZATION).cloned(); if let Some(extracted) = crate::auth::extract_auth_token_from_header( crate::util::get_header_str(&headers, http::header::AUTHORIZATION), @@ -389,6 +386,14 @@ async fn proxy_handler( } } + if let Err(e) = is_ssrf_safe(&target_url).await { + warn!(did = %did, service_id = %service_id, error = %e, "Refusing unsafe proxy target"); + return ApiError::UpstreamFailure.into_response(); + } + info!("Proxying {} request to {}", method_verb, target_url); + + let client = proxy_client(); + let mut request_builder = client.request(method_verb.clone(), &target_url); if let Some(val) = auth_header_val { request_builder = request_builder.header(http::header::AUTHORIZATION, val); } diff --git a/crates/tranquil-pds/src/api/proxy_client.rs b/crates/tranquil-pds/src/api/proxy_client.rs index 0bef89e..a429dbe 100644 --- a/crates/tranquil-pds/src/api/proxy_client.rs +++ b/crates/tranquil-pds/src/api/proxy_client.rs @@ -1,6 +1,6 @@ use axum::http::HeaderName; use reqwest::{Client, ClientBuilder, Url}; -use std::net::{IpAddr, SocketAddr, ToSocketAddrs}; +use std::net::{IpAddr, Ipv4Addr, Ipv6Addr}; use std::sync::{LazyLock, OnceLock}; use std::time::Duration; use tracing::warn; @@ -37,6 +37,7 @@ pub fn did_resolution_client() -> &'static Client { .connect_timeout(DEFAULT_CONNECT_TIMEOUT) .pool_max_idle_per_host(10) .pool_idle_timeout(Duration::from_secs(90)) + .redirect(reqwest::redirect::Policy::none()) .build() .expect( "Failed to build DID resolution client - this indicates a TLS or system configuration issue", @@ -59,26 +60,18 @@ pub fn handle_resolution_client() -> &'static Client { }) } -pub fn is_ssrf_safe(url: &str) -> Result<(), SsrfError> { +pub async fn is_ssrf_safe(url: &str) -> Result<(), SsrfError> { let parsed = Url::parse(url).map_err(|_| SsrfError::InvalidUrl)?; let scheme = parsed.scheme(); if scheme != "https" { - let allow_http = tranquil_config::try_get().is_some_and(|c| c.server.allow_http_proxy) - || url.starts_with("http://127.0.0.1") - || url.starts_with("http://localhost"); - if !allow_http { - return Err(SsrfError::InsecureProtocol(scheme.to_string())); - } + return Err(SsrfError::InsecureProtocol(scheme.to_string())); } - let host = parsed.host_str().ok_or(SsrfError::NoHost)?; - if host == "localhost" { - return Ok(()); + if !parsed.username().is_empty() || parsed.password().is_some() { + return Err(SsrfError::CredentialsNotAllowed); } + let host = parsed.host_str().ok_or(SsrfError::NoHost)?; if let Ok(ip) = host.parse::() { - if ip.is_loopback() { - return Ok(()); - } - if !is_unicast_ip(&ip) { + if !is_public_ip(&ip) { return Err(SsrfError::NonUnicastIp(ip.to_string())); } return Ok(()); @@ -86,13 +79,16 @@ pub fn is_ssrf_safe(url: &str) -> Result<(), SsrfError> { let port = parsed .port() .unwrap_or(if scheme == "https" { 443 } else { 80 }); - let socket_addrs: Vec = match (host, port).to_socket_addrs() { + let socket_addrs: Vec = match tokio::net::lookup_host((host, port)).await { Ok(addrs) => addrs.collect(), Err(_) => return Err(SsrfError::DnsResolutionFailed(host.to_string())), }; - if let Some(addr) = socket_addrs.iter().find(|addr| !is_unicast_ip(&addr.ip())) { + if socket_addrs.is_empty() { + return Err(SsrfError::DnsResolutionFailed(host.to_string())); + } + if let Some(addr) = socket_addrs.iter().find(|addr| !is_public_ip(&addr.ip())) { warn!( - "DNS resolution for {} returned non-unicast IP: {}", + "DNS resolution for {} returned non-public IP: {}", host, addr.ip() ); @@ -101,32 +97,51 @@ pub fn is_ssrf_safe(url: &str) -> Result<(), SsrfError> { Ok(()) } -fn is_unicast_ip(ip: &IpAddr) -> bool { +fn is_public_ip(ip: &IpAddr) -> bool { match ip { - IpAddr::V4(v4) => { - !v4.is_loopback() - && !v4.is_broadcast() - && !v4.is_multicast() - && !v4.is_unspecified() - && !v4.is_link_local() - && !is_private_v4(v4) - } - IpAddr::V6(v6) => !v6.is_loopback() && !v6.is_multicast() && !v6.is_unspecified(), + IpAddr::V4(v4) => is_public_ipv4(v4), + IpAddr::V6(v6) => is_public_ipv6(v6), } } -fn is_private_v4(ip: &std::net::Ipv4Addr) -> bool { - let octets = ip.octets(); - octets[0] == 10 - || (octets[0] == 172 && (16..=31).contains(&octets[1])) - || (octets[0] == 192 && octets[1] == 168) - || (octets[0] == 169 && octets[1] == 254) +fn is_public_ipv4(ip: &Ipv4Addr) -> bool { + let [a, b, c, _] = ip.octets(); + !(a == 0 + || a == 10 + || (a == 100 && (64..=127).contains(&b)) + || a == 127 + || (a == 169 && b == 254) + || (a == 172 && (16..=31).contains(&b)) + || (a == 192 && b == 0 && c == 0) + || (a == 192 && b == 0 && c == 2) + || (a == 192 && b == 88 && c == 99) + || (a == 192 && b == 168) + || (a == 198 && (b == 18 || b == 19)) + || (a == 198 && b == 51 && c == 100) + || (a == 203 && b == 0 && c == 113) + || a >= 224) +} + +fn is_public_ipv6(ip: &Ipv6Addr) -> bool { + let segments = ip.segments(); + if segments[..6] == [0, 0, 0, 0, 0, 0xffff] { + let high = segments[6].to_be_bytes(); + let low = segments[7].to_be_bytes(); + return is_public_ipv4(&Ipv4Addr::new(high[0], high[1], low[0], low[1])); + } + + (segments[0] & 0xe000) == 0x2000 + && !(segments[0] == 0x2001 && segments[1] < 0x0200) + && !(segments[0] == 0x2001 && segments[1] == 0x0db8) + && segments[0] != 0x2002 + && !(segments[0] == 0x3fff && segments[1] < 0x1000) } #[derive(Debug, Clone)] pub enum SsrfError { InvalidUrl, InsecureProtocol(String), + CredentialsNotAllowed, NoHost, NonUnicastIp(String), DnsResolutionFailed(String), @@ -137,8 +152,9 @@ impl std::fmt::Display for SsrfError { match self { SsrfError::InvalidUrl => write!(f, "Invalid URL"), SsrfError::InsecureProtocol(p) => write!(f, "Insecure protocol: {}", p), + SsrfError::CredentialsNotAllowed => write!(f, "URL credentials are not allowed"), SsrfError::NoHost => write!(f, "No host in URL"), - SsrfError::NonUnicastIp(ip) => write!(f, "Non-unicast IP address: {}", ip), + SsrfError::NonUnicastIp(ip) => write!(f, "Non-public IP address: {}", ip), SsrfError::DnsResolutionFailed(host) => { write!(f, "DNS resolution failed for: {}", host) } @@ -223,35 +239,88 @@ pub fn validate_limit(limit: Option, default: u32, max: u32) -> u32 { #[cfg(test)] mod tests { use super::*; - #[test] - fn test_ssrf_safe_https() { - assert!(is_ssrf_safe("https://1.1.1.1/xrpc/test").is_ok()); + + #[tokio::test] + async fn ssrf_allows_public_https() { + assert!(is_ssrf_safe("https://1.1.1.1/xrpc/test").await.is_ok()); } - #[test] - fn test_ssrf_blocks_http_by_default() { - let result = is_ssrf_safe("http://93.184.216.34/xrpc/test"); + + #[tokio::test] + async fn ssrf_blocks_http_by_default() { + let result = is_ssrf_safe("http://93.184.216.34/xrpc/test").await; assert!(matches!(result, Err(SsrfError::InsecureProtocol(_)))); } - #[test] - fn test_ssrf_allows_localhost_http() { - assert!(is_ssrf_safe("http://127.0.0.1:8080/test").is_ok()); - assert!(is_ssrf_safe("http://localhost:8080/test").is_ok()); - } - #[test] - fn test_ssrf_blocks_non_unicast_ip() { + + #[tokio::test] + async fn ssrf_blocks_non_http_schemes() { assert!(matches!( - is_ssrf_safe("https://0.0.0.0/test"), - Err(SsrfError::NonUnicastIp(_)) + is_ssrf_safe("file:///etc/passwd").await, + Err(SsrfError::InsecureProtocol(_)) )); assert!(matches!( - is_ssrf_safe("https://224.0.0.1/test"), - Err(SsrfError::NonUnicastIp(_)) + is_ssrf_safe("ftp://1.1.1.1/test").await, + Err(SsrfError::InsecureProtocol(_)) )); + } + + #[tokio::test] + async fn ssrf_blocks_url_credentials() { assert!(matches!( - is_ssrf_safe("https://255.255.255.255/test"), - Err(SsrfError::NonUnicastIp(_)) + is_ssrf_safe("https://user:password@1.1.1.1/test").await, + Err(SsrfError::CredentialsNotAllowed) )); } + + #[tokio::test] + async fn ssrf_blocks_non_public_ipv4_addresses() { + for ip in [ + "0.0.0.0", + "10.0.0.1", + "100.64.0.1", + "127.0.0.1", + "169.254.169.254", + "172.16.0.1", + "192.168.0.1", + "192.0.2.1", + "198.18.0.1", + "198.51.100.1", + "203.0.113.1", + "224.0.0.1", + "240.0.0.1", + "255.255.255.255", + ] { + assert!( + matches!( + is_ssrf_safe(&format!("https://{ip}/test")).await, + Err(SsrfError::NonUnicastIp(_)) + ), + "expected {ip} to be rejected" + ); + } + } + + #[tokio::test] + async fn ssrf_blocks_non_public_ipv6_addresses() { + for ip in [ + "::", + "::1", + "fe80::1", + "fc00::1", + "fd00::1", + "ff02::1", + "2001:db8::1", + "::ffff:10.0.0.1", + "::ffff:169.254.169.254", + ] { + assert!( + matches!( + is_ssrf_safe(&format!("https://[{ip}]/test")).await, + Err(SsrfError::NonUnicastIp(_)) + ), + "expected {ip} to be rejected" + ); + } + } #[test] fn test_validate_at_uri() { let result = validate_at_uri("at://did:plc:test/app.bsky.feed.post/abc123"); diff --git a/crates/tranquil-pds/src/auth/extractor.rs b/crates/tranquil-pds/src/auth/extractor.rs index ab3eaa5..9d3a0d6 100644 --- a/crates/tranquil-pds/src/auth/extractor.rs +++ b/crates/tranquil-pds/src/auth/extractor.rs @@ -208,6 +208,7 @@ async fn verify_oauth_token_and_build_user( ) -> Result { match crate::oauth::verify::verify_oauth_access_token( state.repos.oauth.as_ref(), + state.repos.user.as_ref(), token, dpop_proof, method, @@ -334,7 +335,7 @@ impl Auth

{ } pub fn needs_scope_check(&self) -> bool { - self.0.is_oauth() + super::scope_check::requires_scope_check(&self.0.auth_source, self.0.scope.as_deref()) } pub fn permissions(&self) -> ScopePermissions { @@ -384,7 +385,7 @@ impl AsRef for Auth

{ impl VerifyScope for Auth

{ fn needs_scope_check(&self) -> bool { - self.0.is_oauth() + super::scope_check::requires_scope_check(&self.0.auth_source, self.0.scope.as_deref()) } fn permissions(&self) -> ScopePermissions { @@ -585,6 +586,22 @@ fn extract_bearer_token(auth_header: &str) -> Result<&str, AuthError> { mod tests { use super::*; + fn auth_with(scope: Option<&str>, auth_source: AuthSource) -> Auth { + Auth( + AuthenticatedUser { + did: "did:plc:scopechecktest".parse().unwrap(), + key_bytes: None, + is_admin: false, + status: AccountStatus::Active, + scope: scope.map(str::to_string), + controller_did: None, + session_id: "test-session".to_string(), + auth_source, + }, + PhantomData, + ) + } + #[test] fn test_extract_bearer_token() { assert_eq!(extract_bearer_token("Bearer abc123").unwrap(), "abc123"); @@ -599,4 +616,13 @@ mod tests { assert!(extract_bearer_token("abc123").is_err()); assert!(extract_bearer_token("").is_err()); } + + #[test] + fn scoped_sessions_require_permission_checks() { + assert!(auth_with(Some("repo:app.bsky.feed.post?action=create"), AuthSource::Session) + .needs_scope_check()); + assert!(!auth_with(Some(crate::auth::TokenScope::Access.as_str()), AuthSource::Session) + .needs_scope_check()); + assert!(auth_with(None, AuthSource::OAuth).needs_scope_check()); + } } diff --git a/crates/tranquil-pds/src/auth/mod.rs b/crates/tranquil-pds/src/auth/mod.rs index a728b7d..9685e65 100644 --- a/crates/tranquil-pds/src/auth/mod.rs +++ b/crates/tranquil-pds/src/auth/mod.rs @@ -591,6 +591,7 @@ pub async fn validate_token_with_dpop( }; match crate::oauth::verify::verify_oauth_access_token( oauth_repo, + user_repo, token, dpop_proof, http_method, diff --git a/crates/tranquil-pds/src/auth/scope_check.rs b/crates/tranquil-pds/src/auth/scope_check.rs index fb8f1c2..e8273bb 100644 --- a/crates/tranquil-pds/src/auth/scope_check.rs +++ b/crates/tranquil-pds/src/auth/scope_check.rs @@ -6,7 +6,7 @@ use crate::types::Nsid; use super::{AuthSource, TokenScope}; -fn requires_scope_check(auth_source: &AuthSource, scope: Option<&str>) -> bool { +pub(crate) fn requires_scope_check(auth_source: &AuthSource, scope: Option<&str>) -> bool { match auth_source { AuthSource::OAuth => true, _ => match scope { diff --git a/crates/tranquil-pds/src/did.rs b/crates/tranquil-pds/src/did.rs index 9d8731b..10b00ce 100644 --- a/crates/tranquil-pds/src/did.rs +++ b/crates/tranquil-pds/src/did.rs @@ -1,3 +1,4 @@ +use crate::api::proxy_client::is_ssrf_safe; use crate::types::Did; use reqwest::Client; use serde::{Deserialize, Serialize}; @@ -101,8 +102,11 @@ impl DidResolver { .timeout(Duration::from_secs(10)) .connect_timeout(Duration::from_secs(5)) .pool_max_idle_per_host(10) + .redirect(reqwest::redirect::Policy::none()) .build() - .unwrap_or_else(|_| Client::new()); + .expect( + "Failed to build DID resolver client - this indicates a TLS or system configuration issue", + ); info!("DID resolver initialized"); @@ -210,6 +214,9 @@ impl DidResolver { debug!("Resolving did:web {} via {}", did, url); + is_ssrf_safe(&url) + .await + .map_err(|e| DidResolutionError::HttpFailed(format!("Unsafe did:web URL: {e}")))?; let resp = self .client .get(&url) @@ -309,6 +316,9 @@ impl DidResolver { ) -> Result { let url = build_did_web_url(did)?; + is_ssrf_safe(&url) + .await + .map_err(|e| DidResolutionError::HttpFailed(format!("Unsafe did:web URL: {e}")))?; let resp = self .client .get(&url) @@ -379,46 +389,22 @@ pub fn create_did_resolver() -> Arc { } fn build_did_web_url(did: &Did) -> Result { - let host = did + let method_specific_id = did .strip_prefix("did:web:") .ok_or(DidResolutionError::InvalidDidWeb)?; - - let (host, path) = if host.contains(':') { - let decoded = host.replace("%3A", ":"); - let parts: Vec<&str> = decoded.splitn(2, '/').collect(); - if parts.len() > 1 { - (parts[0].to_string(), format!("/{}", parts[1])) - } else { - (decoded, String::new()) - } - } else { - let parts: Vec<&str> = host.splitn(2, ':').collect(); - if parts.len() > 1 && parts[1].contains('/') { - let path_parts: Vec<&str> = parts[1].splitn(2, '/').collect(); - if path_parts.len() > 1 { - ( - format!("{}:{}", parts[0], path_parts[0]), - format!("/{}", path_parts[1]), - ) - } else { - (host.to_string(), String::new()) - } - } else { - (host.to_string(), String::new()) - } - }; - - let scheme = - if host.starts_with("localhost") || host.starts_with("127.0.0.1") || host.contains(':') { - "http" - } else { - "https" - }; + let mut segments = method_specific_id.split(':'); + let host = segments + .next() + .filter(|host| !host.is_empty()) + .ok_or(DidResolutionError::InvalidDidWeb)? + .replace("%3A", ":") + .replace("%3a", ":"); + let path: Vec<&str> = segments.collect(); let url = if path.is_empty() { - format!("{}://{}/.well-known/did.json", scheme, host) + format!("https://{}/.well-known/did.json", host) } else { - format!("{}://{}{}/did.json", scheme, host, path) + format!("https://{}/{}/did.json", host, path.join("/")) }; Ok(url) @@ -428,6 +414,31 @@ fn build_did_web_url(did: &Did) -> Result { mod tests { use super::*; + #[test] + fn did_web_urls_use_https_even_for_loopback_and_ports() { + for (did, expected) in [ + ( + "did:web:localhost", + "https://localhost/.well-known/did.json", + ), + ( + "did:web:127.0.0.1", + "https://127.0.0.1/.well-known/did.json", + ), + ( + "did:web:example.com%3A8443", + "https://example.com:8443/.well-known/did.json", + ), + ( + "did:web:example.com:users:alice", + "https://example.com/users/alice/did.json", + ), + ] { + let did = did.parse::().unwrap(); + assert_eq!(build_did_web_url(&did).unwrap(), expected); + } + } + #[test] fn bounded_cache_evicts_the_oldest_entry() { let mut cache = HashMap::new(); diff --git a/crates/tranquil-pds/src/oauth/db/scope_preference.rs b/crates/tranquil-pds/src/oauth/db/scope_preference.rs index ebc85a3..a027213 100644 --- a/crates/tranquil-pds/src/oauth/db/scope_preference.rs +++ b/crates/tranquil-pds/src/oauth/db/scope_preference.rs @@ -22,10 +22,13 @@ pub async fn should_show_consent( return Ok(true); } - let stored_scopes: std::collections::HashSet<&str> = - stored_prefs.iter().map(|p| p.scope.as_str()).collect(); + let granted_scopes: std::collections::HashSet<&str> = stored_prefs + .iter() + .filter(|preference| preference.granted) + .map(|preference| preference.scope.as_str()) + .collect(); Ok(requested_scopes .iter() - .any(|scope| !stored_scopes.contains(scope.as_str()))) + .any(|scope| !granted_scopes.contains(scope.as_str()))) } diff --git a/crates/tranquil-pds/src/oauth/verify.rs b/crates/tranquil-pds/src/oauth/verify.rs index 0d54d83..cd02193 100644 --- a/crates/tranquil-pds/src/oauth/verify.rs +++ b/crates/tranquil-pds/src/oauth/verify.rs @@ -36,6 +36,7 @@ pub struct VerifyResult { pub async fn verify_oauth_access_token( oauth_repo: &dyn OAuthRepository, + user_repo: &dyn UserRepository, access_token: &str, dpop_proof: Option<&str>, http_method: &str, @@ -62,6 +63,22 @@ pub async fn verify_oauth_access_token( "Token session has expired".to_string(), )); } + let user_status = user_repo + .get_status_by_did(&token_data.did) + .await + .map_err(crate::oauth::db_err_to_oauth)?; + if user_status + .as_ref() + .is_none_or(|status| status.takedown_ref.is_some()) + { + oauth_repo + .delete_token(&token_id) + .await + .map_err(crate::oauth::db_err_to_oauth)?; + return Err(OAuthError::InvalidToken( + "Token not found or revoked".to_string(), + )); + } if let Some(expected_jkt) = &token_data.parameters.dpop_jkt { tracing::debug!(expected_jkt = %expected_jkt, "Token requires DPoP"); let proof = dpop_proof.ok_or_else(|| { @@ -299,6 +316,7 @@ impl FromRequestParts for OAuthUser { let http_uri = crate::util::build_full_url(&parts.uri.to_string()); match verify_oauth_access_token( state.repos.oauth.as_ref(), + state.repos.user.as_ref(), token, dpop_proof, http_method, diff --git a/crates/tranquil-store/src/metastore/mod.rs b/crates/tranquil-store/src/metastore/mod.rs index a5bdb0c..84bba6f 100644 --- a/crates/tranquil-store/src/metastore/mod.rs +++ b/crates/tranquil-store/src/metastore/mod.rs @@ -481,7 +481,7 @@ mod tests { } } - fn create_test_user(ms: &Metastore, did: &str, handle: &str) { + fn create_test_user(ms: &Metastore, did: &str, handle: &str) -> uuid::Uuid { let input = tranquil_db_traits::CreatePasswordAccountInput { handle: tranquil_types::Handle::new(handle.to_string()).unwrap(), email: None, @@ -502,7 +502,10 @@ mod tests { invite_code: None, birthdate_pref: None, }; - ms.user_ops().create_password_account(&input).unwrap(); + ms.user_ops() + .create_password_account(&input) + .unwrap() + .user_id } fn legacy_refresh_data( @@ -842,6 +845,261 @@ mod tests { assert_eq!(clients[1].preferences, preferences); } + #[test] + fn deleting_account_removes_all_user_oauth_state() { + use super::oauth_schema::{ + AccountDeviceValue, AuthorizedClientValue, DeviceTrustValue, OAuthRequestValue, + OAuthDeviceValue, OAuthTokenValue, ScopePrefsValue, TokenIndexValue, + TwoFactorChallengeValue, UsedRefreshValue, oauth_2fa_by_request_key, + oauth_2fa_challenge_key, oauth_account_device_key, oauth_auth_by_code_key, + oauth_auth_client_key, oauth_auth_request_key, oauth_device_key, oauth_device_trust_key, + oauth_scope_prefs_key, oauth_token_by_family_key, oauth_token_by_id_key, + oauth_token_by_prev_refresh_key, oauth_token_by_refresh_key, oauth_token_key, + oauth_used_refresh_key, + }; + + let (_dir, ms) = open_fresh(); + let did = Did::new("did:plc:oauth-delete").unwrap(); + let user_id = create_test_user(&ms, did.as_str(), "oauth-delete.test"); + let user_hash = keys::UserHash::from_did(did.as_str()); + let auth = ms.partition(Partition::Auth); + let expires_at_ms = chrono::Utc::now().timestamp_millis() + 60_000; + let family_id = 42; + let token = OAuthTokenValue { + family_id, + did: did.to_string(), + client_id: "https://client.example".to_owned(), + token_id: "token-id".to_owned(), + refresh_token: "current-refresh".to_owned(), + previous_refresh_token: Some("previous-refresh".to_owned()), + scope: "atproto".to_owned(), + expires_at_ms, + created_at_ms: expires_at_ms - 1_000, + updated_at_ms: expires_at_ms - 1_000, + parameters_json: "{}".to_owned(), + controller_did: None, + }; + let token_index = TokenIndexValue { + user_hash: user_hash.raw(), + family_id, + }; + + auth.insert( + oauth_token_key(user_hash, family_id).as_slice(), + token.serialize_with_ttl(), + ) + .unwrap(); + for key in [ + oauth_token_by_id_key(&token.token_id), + oauth_token_by_refresh_key(&token.refresh_token), + oauth_token_by_prev_refresh_key(token.previous_refresh_token.as_deref().unwrap()), + oauth_token_by_family_key(family_id), + ] { + auth.insert( + key.as_slice(), + token_index.serialize_with_ttl(expires_at_ms), + ) + .unwrap(); + } + auth.insert( + oauth_used_refresh_key("used-refresh").as_slice(), + UsedRefreshValue { family_id }.serialize_with_ttl(expires_at_ms), + ) + .unwrap(); + + let request_id = "request-id"; + let authorization_code = "authorization-code"; + auth.insert( + oauth_auth_request_key(request_id).as_slice(), + OAuthRequestValue { + client_id: token.client_id.clone(), + client_auth_json: None, + parameters_json: "{}".to_owned(), + expires_at_ms, + did: Some(did.to_string()), + device_id: None, + code: Some(authorization_code.to_owned()), + controller_did: None, + } + .serialize_with_ttl(), + ) + .unwrap(); + auth.insert( + oauth_auth_by_code_key(authorization_code).as_slice(), + request_id.as_bytes(), + ) + .unwrap(); + + let challenge_id = *uuid::Uuid::new_v4().as_bytes(); + let challenge_request = "challenge-request"; + auth.insert( + oauth_2fa_challenge_key(&challenge_id).as_slice(), + TwoFactorChallengeValue { + id: challenge_id, + did: did.to_string(), + request_uri: challenge_request.to_owned(), + code: "123456".to_owned(), + attempts: 0, + created_at_ms: expires_at_ms - 1_000, + expires_at_ms, + } + .serialize_with_ttl(), + ) + .unwrap(); + auth.insert( + oauth_2fa_by_request_key(challenge_request).as_slice(), + challenge_id, + ) + .unwrap(); + + let device_id = "device-id"; + auth.insert( + oauth_account_device_key(user_hash, device_id).as_slice(), + AccountDeviceValue { + last_used_at_ms: expires_at_ms - 1_000, + } + .serialize_with_ttl(), + ) + .unwrap(); + auth.insert( + oauth_scope_prefs_key(user_hash, &token.client_id).as_slice(), + ScopePrefsValue { + prefs_json: "[]".to_owned(), + } + .serialize(), + ) + .unwrap(); + auth.insert( + oauth_auth_client_key(user_hash, &token.client_id).as_slice(), + AuthorizedClientValue { + data_json: "{}".to_owned(), + } + .serialize(), + ) + .unwrap(); + auth.insert( + oauth_device_trust_key(user_hash, device_id).as_slice(), + DeviceTrustValue { + device_id: device_id.to_owned(), + did: did.to_string(), + user_agent: None, + friendly_name: None, + trusted_at_ms: Some(expires_at_ms - 1_000), + trusted_until_ms: Some(expires_at_ms), + last_seen_at_ms: expires_at_ms - 1_000, + } + .serialize(), + ) + .unwrap(); + + let other_did = Did::new("did:plc:oauth-delete-other").unwrap(); + let other_user_hash = keys::UserHash::from_did(other_did.as_str()); + let other_challenge_id = *uuid::Uuid::new_v4().as_bytes(); + auth.insert( + oauth_used_refresh_key("other-used-refresh").as_slice(), + UsedRefreshValue { family_id: 99 }.serialize_with_ttl(expires_at_ms), + ) + .unwrap(); + auth.insert( + oauth_auth_request_key("other-request").as_slice(), + OAuthRequestValue { + client_id: token.client_id.clone(), + client_auth_json: None, + parameters_json: "{}".to_owned(), + expires_at_ms, + did: Some(other_did.to_string()), + device_id: None, + code: Some("other-code".to_owned()), + controller_did: None, + } + .serialize_with_ttl(), + ) + .unwrap(); + auth.insert( + oauth_auth_by_code_key("other-code").as_slice(), + b"other-request", + ) + .unwrap(); + auth.insert( + oauth_2fa_challenge_key(&other_challenge_id).as_slice(), + TwoFactorChallengeValue { + id: other_challenge_id, + did: other_did.to_string(), + request_uri: "other-challenge-request".to_owned(), + code: "654321".to_owned(), + attempts: 0, + created_at_ms: expires_at_ms - 1_000, + expires_at_ms, + } + .serialize_with_ttl(), + ) + .unwrap(); + auth.insert( + oauth_2fa_by_request_key("other-challenge-request").as_slice(), + other_challenge_id, + ) + .unwrap(); + auth.insert( + oauth_auth_client_key(other_user_hash, &token.client_id).as_slice(), + AuthorizedClientValue { + data_json: "{}".to_owned(), + } + .serialize(), + ) + .unwrap(); + auth.insert( + oauth_device_key(device_id).as_slice(), + OAuthDeviceValue { + session_id: "shared-session".to_owned(), + user_agent: None, + ip_address: "127.0.0.1".to_owned(), + last_seen_at_ms: expires_at_ms - 1_000, + created_at_ms: expires_at_ms - 1_000, + } + .serialize(), + ) + .unwrap(); + + let deleted_keys = [ + oauth_token_key(user_hash, family_id), + oauth_token_by_id_key(&token.token_id), + oauth_token_by_refresh_key(&token.refresh_token), + oauth_token_by_prev_refresh_key(token.previous_refresh_token.as_deref().unwrap()), + oauth_token_by_family_key(family_id), + oauth_used_refresh_key("used-refresh"), + oauth_auth_request_key(request_id), + oauth_auth_by_code_key(authorization_code), + oauth_2fa_challenge_key(&challenge_id), + oauth_2fa_by_request_key(challenge_request), + oauth_account_device_key(user_hash, device_id), + oauth_scope_prefs_key(user_hash, &token.client_id), + oauth_auth_client_key(user_hash, &token.client_id), + oauth_device_trust_key(user_hash, device_id), + ]; + deleted_keys.iter().for_each(|key| { + assert!(auth.get(key.as_slice()).unwrap().is_some()); + }); + + ms.user_ops() + .delete_account_complete(user_id, &did) + .unwrap(); + + deleted_keys.iter().for_each(|key| { + assert!(auth.get(key.as_slice()).unwrap().is_none()); + }); + for key in [ + oauth_used_refresh_key("other-used-refresh"), + oauth_auth_request_key("other-request"), + oauth_auth_by_code_key("other-code"), + oauth_2fa_challenge_key(&other_challenge_id), + oauth_2fa_by_request_key("other-challenge-request"), + oauth_auth_client_key(other_user_hash, &token.client_id), + oauth_device_key(device_id), + ] { + assert!(auth.get(key.as_slice()).unwrap().is_some()); + } + } + fn stamp_format_version(dir: &std::path::Path, version: u64) { let ms = Metastore::open(dir, test_config()).unwrap(); let repo_data = ms.partition(Partition::RepoData); diff --git a/crates/tranquil-store/src/metastore/oauth_ops.rs b/crates/tranquil-store/src/metastore/oauth_ops.rs index b8de202..71ec739 100644 --- a/crates/tranquil-store/src/metastore/oauth_ops.rs +++ b/crates/tranquil-store/src/metastore/oauth_ops.rs @@ -11,14 +11,15 @@ use super::oauth_schema::{ OAuthRequestValue, OAuthTokenValue, ScopePrefsValue, TokenIndexValue, TwoFactorChallengeValue, UsedRefreshValue, deserialize_family_counter, oauth_2fa_by_request_key, oauth_2fa_challenge_key, oauth_2fa_challenge_prefix, oauth_account_device_key, - oauth_auth_by_code_key, oauth_auth_client_key, oauth_auth_request_key, - oauth_auth_request_prefix, oauth_device_key, oauth_device_trust_key, oauth_dpop_jti_key, - oauth_dpop_jti_prefix, oauth_scope_prefs_key, oauth_scope_prefs_prefix, - oauth_token_by_family_key, oauth_token_by_id_key, oauth_token_by_prev_refresh_key, - oauth_token_by_refresh_key, oauth_token_family_counter_key, oauth_token_key, - oauth_token_user_prefix, oauth_used_refresh_key, serialize_family_counter, + oauth_account_device_prefix, oauth_auth_by_code_key, oauth_auth_client_key, + oauth_auth_client_prefix, oauth_auth_request_key, oauth_auth_request_prefix, oauth_device_key, + oauth_device_trust_key, oauth_device_trust_prefix, oauth_dpop_jti_key, oauth_dpop_jti_prefix, + oauth_scope_prefs_key, oauth_scope_prefs_prefix, oauth_token_by_family_key, + oauth_token_by_id_key, oauth_token_by_prev_refresh_key, oauth_token_by_refresh_key, + oauth_token_family_counter_key, oauth_token_key, oauth_token_user_prefix, + oauth_used_refresh_key, oauth_used_refresh_prefix, serialize_family_counter, }; -use super::scan::point_lookup; +use super::scan::{delete_all_by_prefix, point_lookup}; use super::users::UserValue; use tranquil_db_traits::{ @@ -31,6 +32,116 @@ use tranquil_types::{ TokenId, }; +fn stage_delete_token_indexes( + auth: &Keyspace, + batch: &mut fjall::OwnedWriteBatch, + token: &OAuthTokenValue, + user_hash: UserHash, +) { + batch.remove( + auth, + oauth_token_key(user_hash, token.family_id).as_slice(), + ); + batch.remove(auth, oauth_token_by_id_key(&token.token_id).as_slice()); + batch.remove( + auth, + oauth_token_by_refresh_key(&token.refresh_token).as_slice(), + ); + if let Some(prev) = &token.previous_refresh_token { + batch.remove(auth, oauth_token_by_prev_refresh_key(prev).as_slice()); + } + batch.remove( + auth, + oauth_token_by_family_key(token.family_id).as_slice(), + ); +} + +fn collect_tokens_for_user( + auth: &Keyspace, + user_hash: UserHash, +) -> Result, MetastoreError> { + let prefix = oauth_token_user_prefix(user_hash); + auth.prefix(prefix.as_slice()) + .try_fold(Vec::new(), |mut tokens, guard| { + let (_, value_bytes) = guard.into_inner().map_err(MetastoreError::Fjall)?; + if let Some(token) = OAuthTokenValue::deserialize(&value_bytes) { + tokens.push(token); + } + Ok::<_, MetastoreError>(tokens) + }) +} + +pub(super) fn stage_delete_user_oauth_data( + auth: &Keyspace, + batch: &mut fjall::OwnedWriteBatch, + user_hash: UserHash, + did: &str, +) -> Result<(), MetastoreError> { + let tokens = collect_tokens_for_user(auth, user_hash)?; + let family_ids: Vec = tokens.iter().map(|token| token.family_id).collect(); + tokens.iter().for_each(|token| { + stage_delete_token_indexes(auth, batch, token, user_hash); + }); + delete_all_by_prefix( + auth, + batch, + oauth_token_user_prefix(user_hash).as_slice(), + )?; + + let used_prefix = oauth_used_refresh_prefix(); + auth.prefix(used_prefix.as_slice()).try_for_each(|guard| { + let (key_bytes, value_bytes) = guard.into_inner().map_err(MetastoreError::Fjall)?; + let used = UsedRefreshValue::deserialize(&value_bytes) + .ok_or(MetastoreError::CorruptData("corrupt oauth used refresh"))?; + if family_ids.contains(&used.family_id) { + batch.remove(auth, key_bytes.as_ref()); + } + Ok::<(), MetastoreError>(()) + })?; + + let request_prefix = oauth_auth_request_prefix(); + auth.prefix(request_prefix.as_slice()) + .try_for_each(|guard| { + let (key_bytes, value_bytes) = guard.into_inner().map_err(MetastoreError::Fjall)?; + let request = OAuthRequestValue::deserialize(&value_bytes) + .ok_or(MetastoreError::CorruptData("corrupt oauth auth request"))?; + if request.did.as_deref() == Some(did) { + batch.remove(auth, key_bytes.as_ref()); + if let Some(code) = request.code { + batch.remove(auth, oauth_auth_by_code_key(&code).as_slice()); + } + } + Ok::<(), MetastoreError>(()) + })?; + + let challenge_prefix = oauth_2fa_challenge_prefix(); + auth.prefix(challenge_prefix.as_slice()) + .try_for_each(|guard| { + let (key_bytes, value_bytes) = guard.into_inner().map_err(MetastoreError::Fjall)?; + let challenge = TwoFactorChallengeValue::deserialize(&value_bytes) + .ok_or(MetastoreError::CorruptData("corrupt oauth 2fa challenge"))?; + if challenge.did == did { + batch.remove(auth, key_bytes.as_ref()); + batch.remove( + auth, + oauth_2fa_by_request_key(&challenge.request_uri).as_slice(), + ); + } + Ok::<(), MetastoreError>(()) + })?; + + for prefix in [ + oauth_account_device_prefix(user_hash), + oauth_scope_prefs_prefix(user_hash), + oauth_auth_client_prefix(user_hash), + oauth_device_trust_prefix(user_hash), + ] { + delete_all_by_prefix(auth, batch, prefix.as_slice())?; + } + + Ok(()) +} + pub struct OAuthOps { db: Database, auth: Keyspace, @@ -138,44 +249,14 @@ impl OAuthOps { token: &OAuthTokenValue, user_hash: UserHash, ) { - batch.remove( - &self.auth, - oauth_token_key(user_hash, token.family_id).as_slice(), - ); - batch.remove( - &self.auth, - oauth_token_by_id_key(&token.token_id).as_slice(), - ); - batch.remove( - &self.auth, - oauth_token_by_refresh_key(&token.refresh_token).as_slice(), - ); - if let Some(prev) = &token.previous_refresh_token { - batch.remove(&self.auth, oauth_token_by_prev_refresh_key(prev).as_slice()); - } - batch.remove( - &self.auth, - oauth_token_by_family_key(token.family_id).as_slice(), - ); + stage_delete_token_indexes(&self.auth, batch, token, user_hash); } fn collect_tokens_for_did( &self, user_hash: UserHash, ) -> Result, MetastoreError> { - let prefix = oauth_token_user_prefix(user_hash); - self.auth - .prefix(prefix.as_slice()) - .try_fold(Vec::new(), |mut acc, guard| { - let (_, val_bytes) = guard.into_inner().map_err(MetastoreError::Fjall)?; - match OAuthTokenValue::deserialize(&val_bytes) { - Some(v) => { - acc.push(v); - Ok::<_, MetastoreError>(acc) - } - None => Ok(acc), - } - }) + collect_tokens_for_user(&self.auth, user_hash) } fn request_value_to_data(&self, v: &OAuthRequestValue) -> Result { diff --git a/crates/tranquil-store/src/metastore/oauth_schema.rs b/crates/tranquil-store/src/metastore/oauth_schema.rs index 13b1130..352871e 100644 --- a/crates/tranquil-store/src/metastore/oauth_schema.rs +++ b/crates/tranquil-store/src/metastore/oauth_schema.rs @@ -393,6 +393,10 @@ pub fn oauth_used_refresh_key(refresh_token: &str) -> SmallVec<[u8; 128]> { .build() } +pub fn oauth_used_refresh_prefix() -> SmallVec<[u8; 128]> { + KeyBuilder::new().tag(KeyTag::OAUTH_USED_REFRESH).build() +} + pub fn oauth_auth_request_key(request_id: &str) -> SmallVec<[u8; 128]> { KeyBuilder::new() .tag(KeyTag::OAUTH_AUTH_REQUEST) @@ -485,6 +489,13 @@ pub fn oauth_auth_client_key(user_hash: UserHash, client_id: &str) -> SmallVec<[ .build() } +pub fn oauth_auth_client_prefix(user_hash: UserHash) -> SmallVec<[u8; 128]> { + KeyBuilder::new() + .tag(KeyTag::OAUTH_AUTH_CLIENT) + .u64(user_hash.raw()) + .build() +} + pub fn oauth_device_trust_key(user_hash: UserHash, device_id: &str) -> SmallVec<[u8; 128]> { KeyBuilder::new() .tag(KeyTag::OAUTH_DEVICE_TRUST) diff --git a/crates/tranquil-store/src/metastore/user_ops.rs b/crates/tranquil-store/src/metastore/user_ops.rs index 7ca6b03..44e0ad3 100644 --- a/crates/tranquil-store/src/metastore/user_ops.rs +++ b/crates/tranquil-store/src/metastore/user_ops.rs @@ -7,6 +7,7 @@ use uuid::Uuid; use super::MetastoreError; use super::infra_schema::{channel_to_u8, u8_to_channel}; use super::keys::UserHash; +use super::oauth_ops::stage_delete_user_oauth_data; use super::repo_meta::{RepoMetaValue, RepoStatus, handle_key, repo_meta_key}; use super::repo_ops::{cid_link_to_bytes, stage_full_repo_data_removal}; use super::scan::{count_prefix, delete_all_by_prefix, point_lookup}; @@ -2348,7 +2349,7 @@ impl UserOps { batch.remove(&self.users, recovery_token_key(user_hash).as_slice()); batch.remove(&self.users, did_web_overrides_key(user_hash).as_slice()); - self.delete_auth_data_for_user(batch, user_hash)?; + self.delete_auth_data_for_user(batch, user_hash, &user.did)?; Ok(()) } @@ -2357,6 +2358,7 @@ impl UserOps { &self, batch: &mut fjall::OwnedWriteBatch, user_hash: UserHash, + did: &str, ) -> Result<(), MetastoreError> { use super::sessions::{ SessionTokenValue, session_app_password_prefix, session_by_access_key, @@ -2412,6 +2414,8 @@ impl UserOps { batch.remove(&self.auth, webauthn_challenge_key(user_hash, 0).as_slice()); batch.remove(&self.auth, webauthn_challenge_key(user_hash, 1).as_slice()); + stage_delete_user_oauth_data(&self.auth, batch, user_hash, did)?; + Ok(()) } diff --git a/crates/tranquil-sync/src/blob.rs b/crates/tranquil-sync/src/blob.rs index 0e1f43a..02b299c 100644 --- a/crates/tranquil-sync/src/blob.rs +++ b/crates/tranquil-sync/src/blob.rs @@ -35,6 +35,15 @@ pub async fn get_blob( Err(e) => return e.into_response(), }; + match state.repos.blob.get_blob_with_takedown(&cid).await { + Ok(Some(blob)) if blob.takedown_ref.is_none() => {} + Ok(_) => return ApiError::BlobNotFound(Some("Blob not found".into())).into_response(), + Err(e) => { + error!("DB error checking blob moderation status: {:?}", e); + return ApiError::InternalError(None).into_response(); + } + } + let blob_result = state.repos.blob.get_blob_metadata(&cid).await; match blob_result { Ok(Some(metadata)) => match state.blob_store.get(&metadata.storage_key).await { @@ -42,6 +51,10 @@ pub async fn get_blob( .status(StatusCode::OK) .header(header::CONTENT_TYPE, &metadata.mime_type) .header(header::CONTENT_LENGTH, metadata.size_bytes.to_string()) + .header( + header::CONTENT_DISPOSITION, + format!("attachment; filename=\"{}\"", cid), + ) .header("x-content-type-options", "nosniff") .header("content-security-policy", "default-src 'none'; sandbox") .body(Body::from(data)) diff --git a/frontend/index.html b/frontend/index.html index 52c287f..87fc750 100644 --- a/frontend/index.html +++ b/frontend/index.html @@ -3,6 +3,7 @@ + Tranquil PDS