From 0ad78af842b96593da47a6c81ab1d7f3810fee29 Mon Sep 17 00:00:00 2001 From: Trezy Date: Tue, 7 Jul 2026 11:14:34 -0500 Subject: [PATCH] fix: prevent CORS from reflecting credentials to arbitrary origins Signed-off-by: Trezy --- src/domain.rs | 39 ++++++++++++ src/server.rs | 125 ++++++++++++++++++++++++++++++++----- tests/common/db.rs | 12 +++- tests/cors.rs | 151 +++++++++++++++++++++++++++++++++++++++++++++ 4 files changed, 309 insertions(+), 18 deletions(-) create mode 100644 tests/cors.rs diff --git a/src/domain.rs b/src/domain.rs index 4bfe539..bd9a1e0 100644 --- a/src/domain.rs +++ b/src/domain.rs @@ -99,6 +99,23 @@ impl DomainCache { let by_host = self.by_host.read().await; by_host.values().cloned().collect() } + + /// Return `true` if `origin` (a browser `Origin` header value, e.g. + /// `https://example.com` or `http://localhost:3000`) exactly matches a + /// registered domain's URL. + /// + /// This is the trusted-origin allowlist for credentialed (cookie-bearing) + /// CORS: only first-party domains HappyView actually serves may send + /// credentials cross-origin. The match is on the full origin + /// (scheme + host + port), not just the host, so `http` and `https` or a + /// different port are treated as distinct origins. + pub async fn is_allowed_origin(&self, origin: &str) -> bool { + let target = origin.trim_end_matches('/'); + let by_host = self.by_host.read().await; + by_host + .values() + .any(|d| d.url.trim_end_matches('/') == target) + } } impl Default for DomainCache { @@ -179,6 +196,28 @@ mod tests { assert!(cache.get("example.com").await.is_none()); } + #[tokio::test] + async fn is_allowed_origin_matches_full_origin() { + let cache = DomainCache::new(); + cache + .load(vec![ + make_domain("https://example.com", true), + make_domain("http://localhost:3000", false), + ]) + .await; + + // Exact matches (trailing slash tolerated). + assert!(cache.is_allowed_origin("https://example.com").await); + assert!(cache.is_allowed_origin("https://example.com/").await); + assert!(cache.is_allowed_origin("http://localhost:3000").await); + + // Scheme, port, and host must all match. + assert!(!cache.is_allowed_origin("http://example.com").await); + assert!(!cache.is_allowed_origin("https://example.com:8443").await); + assert!(!cache.is_allowed_origin("http://localhost:3001").await); + assert!(!cache.is_allowed_origin("https://evil.example").await); + } + #[tokio::test] async fn set_primary_updates() { let cache = DomainCache::new(); diff --git a/src/server.rs b/src/server.rs index 71d655f..f6732ce 100644 --- a/src/server.rs +++ b/src/server.rs @@ -7,7 +7,6 @@ use base64::Engine; use bytes::Bytes; use http_body_util::Full; use std::convert::Infallible; -use tower_http::cors::CorsLayer; use tower_http::services::ServeDir; use tower_http::trace::TraceLayer; @@ -179,25 +178,117 @@ pub fn router(state: AppState) -> Router { outer .layer(TraceLayer::new_for_http()) - .layer( - CorsLayer::new() - .allow_origin(tower_http::cors::AllowOrigin::mirror_request()) - .allow_methods([Method::GET, Method::POST, Method::DELETE, Method::OPTIONS]) - .allow_headers([ - header::CONTENT_TYPE, - header::AUTHORIZATION, - header::COOKIE, - axum::http::HeaderName::from_static("x-client-key"), - axum::http::HeaderName::from_static("x-client-secret"), - axum::http::HeaderName::from_static("dpop"), - axum::http::HeaderName::from_static("atproto-accept-labelers"), - axum::http::HeaderName::from_static("atproto-proxy"), - ]) - .allow_credentials(true), - ) + .layer(axum::middleware::from_fn_with_state(state.clone(), cors)) .with_state(state) } +/// Allowed request methods, shared by both CORS policies. +const CORS_ALLOW_METHODS: &str = "GET, POST, DELETE, OPTIONS"; + +/// Headers a credentialed (first-party, cookie-bearing) request may send. +const CORS_ALLOW_HEADERS_CREDENTIALED: &str = "content-type, authorization, cookie, x-client-key, x-client-secret, dpop, \ + atproto-accept-labelers, atproto-proxy"; + +/// Headers a credential-less (third-party, cookieless DPoP) request may send. +/// Identical to the credentialed set minus `cookie`. +const CORS_ALLOW_HEADERS_ANON: &str = "content-type, authorization, x-client-key, x-client-secret, dpop, \ + atproto-accept-labelers, atproto-proxy"; + +/// Cross-Origin Resource Sharing policy. +/// +/// This deliberately replaces a single permissive `CorsLayer`. The old policy +/// reflected *any* `Origin` **and** allowed credentials, which let a malicious +/// page drive the admin API with the victim's cookie and read the response +/// (finding C2). Instead we apply two policies keyed on trust: +/// +/// - **Trusted first-party origins** — those in the [`DomainCache`] (the domains +/// HappyView actually serves the dashboard/admin UI on) — get their origin +/// reflected *with* `Access-Control-Allow-Credentials: true`, so the +/// cookie-authenticated dashboard works cross-origin if ever hosted on a +/// second registered domain. +/// - **Any other origin** — e.g. a third-party app or an attacker page — gets a +/// credential-*less* grant: its origin is reflected but credentials are never +/// allowed. Third-party clients authenticate with explicit DPoP + client-key +/// headers (never ambient cookies), so they keep working; an attacker page +/// can neither ride the admin cookie nor read a credentialed response. +/// +/// The one rule that must never be violated: reflecting an arbitrary origin and +/// allowing credentials at the same time. +async fn cors( + State(state): State, + req: axum::extract::Request, + next: axum::middleware::Next, +) -> Response { + let origin = req + .headers() + .get(header::ORIGIN) + .and_then(|v| v.to_str().ok()) + .map(|s| s.to_string()); + + // No `Origin` header → not a CORS request (same-origin navigation, + // server-to-server, curl). Emit no CORS headers at all. + let Some(origin) = origin else { + return next.run(req).await; + }; + + let credentialed = state.domain_cache.is_allowed_origin(&origin).await; + + let is_preflight = req.method() == Method::OPTIONS + && req + .headers() + .contains_key(header::ACCESS_CONTROL_REQUEST_METHOD); + + let mut cors_headers = header::HeaderMap::new(); + if let Ok(value) = header::HeaderValue::from_str(&origin) { + cors_headers.insert(header::ACCESS_CONTROL_ALLOW_ORIGIN, value); + } else { + // Malformed origin — refuse it entirely rather than emit a broken header. + if is_preflight { + return preflight_response(cors_headers); + } + return next.run(req).await; + } + cors_headers.insert(header::VARY, header::HeaderValue::from_static("origin")); + if credentialed { + cors_headers.insert( + header::ACCESS_CONTROL_ALLOW_CREDENTIALS, + header::HeaderValue::from_static("true"), + ); + } + + if is_preflight { + cors_headers.insert( + header::ACCESS_CONTROL_ALLOW_METHODS, + header::HeaderValue::from_static(CORS_ALLOW_METHODS), + ); + cors_headers.insert( + header::ACCESS_CONTROL_ALLOW_HEADERS, + header::HeaderValue::from_static(if credentialed { + CORS_ALLOW_HEADERS_CREDENTIALED + } else { + CORS_ALLOW_HEADERS_ANON + }), + ); + cors_headers.insert( + header::ACCESS_CONTROL_MAX_AGE, + header::HeaderValue::from_static("86400"), + ); + return preflight_response(cors_headers); + } + + let mut resp = next.run(req).await; + resp.headers_mut().extend(cors_headers); + resp +} + +/// Build a `204 No Content` preflight response carrying the given CORS headers. +fn preflight_response(cors_headers: header::HeaderMap) -> Response { + let mut resp = Response::new(axum::body::Body::empty()); + *resp.status_mut() = axum::http::StatusCode::NO_CONTENT; + resp.headers_mut().extend(cors_headers); + resp +} + async fn health() -> &'static str { "ok" } diff --git a/tests/common/db.rs b/tests/common/db.rs index 80f4968..1f5dcd0 100644 --- a/tests/common/db.rs +++ b/tests/common/db.rs @@ -48,7 +48,7 @@ pub async fn truncate_all(pool: &AnyPool) { match backend { DatabaseBackend::Postgres => { sqlx::query( - "TRUNCATE happyview_records, happyview_lexicons, happyview_backfill_jobs, happyview_users, happyview_user_permissions, happyview_api_keys, happyview_event_logs, happyview_script_variables, happyview_scripts, happyview_dead_letter_scripts, happyview_dead_letter_hooks, happyview_record_refs, happyview_labeler_subscriptions, happyview_labels, happyview_instance_settings, happyview_domains, happyview_dpop_sessions, happyview_dpop_keys, happyview_api_clients, happyview_delegated_accounts, happyview_account_delegates, happyview_service_identity, happyview_service_entries, happyview_service_entry_xrpcs, happyview_jobs RESTART IDENTITY CASCADE", + "TRUNCATE happyview_records, happyview_lexicons, happyview_backfill_jobs, happyview_users, happyview_user_permissions, happyview_api_keys, happyview_event_logs, happyview_script_variables, happyview_scripts, happyview_dead_letter_scripts, happyview_dead_letter_hooks, happyview_record_refs, happyview_labeler_subscriptions, happyview_labels, happyview_instance_settings, happyview_domains, happyview_dpop_sessions, happyview_dpop_keys, happyview_api_clients, happyview_delegated_accounts, happyview_account_delegates, happyview_service_identity, happyview_service_entries, happyview_service_entry_xrpcs, happyview_jobs, happyview_spaces, happyview_space_members, happyview_space_records, happyview_space_repo_state, happyview_space_record_oplog, happyview_space_notify_registrations, happyview_space_invites RESTART IDENTITY CASCADE", ) .execute(pool) .await @@ -56,6 +56,16 @@ pub async fn truncate_all(pool: &AnyPool) { } DatabaseBackend::Sqlite => { let tables = [ + // Spaces tables (children before parents — no cascade on SQLite). + "happyview_space_credentials", + "happyview_space_dids", + "happyview_space_invites", + "happyview_space_notify_registrations", + "happyview_space_record_oplog", + "happyview_space_repo_state", + "happyview_space_records", + "happyview_space_members", + "happyview_spaces", "happyview_service_entry_xrpcs", "happyview_service_entries", "happyview_service_identity", diff --git a/tests/cors.rs b/tests/cors.rs new file mode 100644 index 0000000..3da9772 --- /dev/null +++ b/tests/cors.rs @@ -0,0 +1,151 @@ +mod common; + +use axum::body::Body; +use axum::http::Request; +use serial_test::serial; +use tower::ServiceExt; + +/// The origin the TestApp registers in its DomainCache (a trusted first-party +/// domain — where the dashboard / admin UI is served). +const TRUSTED_ORIGIN: &str = "http://127.0.0.1:0"; +/// An arbitrary untrusted origin (e.g. a third-party app or a malicious page). +const UNTRUSTED_ORIGIN: &str = "https://evil.example"; + +fn preflight(origin: &str, request_method: &str) -> Request { + Request::builder() + .method("OPTIONS") + .uri("/xrpc/com.example.test") + .header("host", "127.0.0.1") + .header("origin", origin) + .header("access-control-request-method", request_method) + .header( + "access-control-request-headers", + "content-type, authorization", + ) + .body(Body::empty()) + .unwrap() +} + +fn get_with_origin(uri: &str, origin: &str) -> Request { + Request::builder() + .method("GET") + .uri(uri) + .header("host", "127.0.0.1") + .header("origin", origin) + .body(Body::empty()) + .unwrap() +} + +fn header<'a>(resp: &'a axum::http::Response, name: &str) -> Option<&'a str> { + resp.headers().get(name).and_then(|v| v.to_str().ok()) +} + +#[tokio::test] +#[serial] +async fn preflight_from_trusted_origin_allows_credentials() { + common::require_db!(); + let app = common::app::TestApp::new().await; + + let resp = app + .router + .clone() + .oneshot(preflight(TRUSTED_ORIGIN, "POST")) + .await + .unwrap(); + + assert!(resp.status().is_success()); + assert_eq!( + header(&resp, "access-control-allow-origin"), + Some(TRUSTED_ORIGIN) + ); + assert_eq!( + header(&resp, "access-control-allow-credentials"), + Some("true"), + "trusted first-party origins must be allowed to send credentials" + ); +} + +#[tokio::test] +#[serial] +async fn preflight_from_untrusted_origin_never_allows_credentials() { + common::require_db!(); + let app = common::app::TestApp::new().await; + + let resp = app + .router + .clone() + .oneshot(preflight(UNTRUSTED_ORIGIN, "POST")) + .await + .unwrap(); + + // The untrusted origin may still use credential-less (cookieless) CORS — + // e.g. a third-party DPoP client — but it must NEVER be granted credentials. + assert_eq!( + header(&resp, "access-control-allow-credentials"), + None, + "untrusted origins must never be allowed to send credentials" + ); +} + +#[tokio::test] +#[serial] +async fn actual_request_from_trusted_origin_allows_credentials() { + common::require_db!(); + let app = common::app::TestApp::new().await; + + let resp = app + .router + .clone() + .oneshot(get_with_origin("/health", TRUSTED_ORIGIN)) + .await + .unwrap(); + + assert_eq!( + header(&resp, "access-control-allow-origin"), + Some(TRUSTED_ORIGIN) + ); + assert_eq!( + header(&resp, "access-control-allow-credentials"), + Some("true") + ); +} + +#[tokio::test] +#[serial] +async fn actual_request_from_untrusted_origin_never_allows_credentials() { + common::require_db!(); + let app = common::app::TestApp::new().await; + + let resp = app + .router + .clone() + .oneshot(get_with_origin("/health", UNTRUSTED_ORIGIN)) + .await + .unwrap(); + + assert_eq!( + header(&resp, "access-control-allow-credentials"), + None, + "untrusted origins must never be allowed to send credentials" + ); +} + +#[tokio::test] +#[serial] +async fn request_without_origin_gets_no_cors_headers() { + common::require_db!(); + let app = common::app::TestApp::new().await; + + let req = Request::builder() + .method("GET") + .uri("/health") + .header("host", "127.0.0.1") + .body(Body::empty()) + .unwrap(); + + let resp = app.router.clone().oneshot(req).await.unwrap(); + + assert!(resp.status().is_success()); + assert_eq!(header(&resp, "access-control-allow-origin"), None); + assert_eq!(header(&resp, "access-control-allow-credentials"), None); +} -- 2.51.2