diff --git a/bobbin/crates/bobbin/src/config.rs b/bobbin/crates/bobbin/src/config.rs index b831fd7a3..a993c5ef3 100644 --- a/bobbin/crates/bobbin/src/config.rs +++ b/bobbin/crates/bobbin/src/config.rs @@ -4,6 +4,7 @@ use std::path::{Path, PathBuf}; use std::str::FromStr; use anyhow::{Context, anyhow}; +use bobbin_xrpc::Cors; use confique::Config; use jacquard_common::types::did::Did; use trusted_proxies::{ProxyNetError, TrustedProxies}; @@ -17,6 +18,7 @@ const KNOWN_KEYS: &[&str] = &[ "server.shutdown_grace_secs", "server.debug_bind", "server.trusted_proxies", + "server.cors_origins", "hydrant.url", "hydrant.start_cursor", "ingest.parallelism", @@ -46,6 +48,7 @@ const KNOWN_ENVS: &[&str] = &[ "BOBBIN_SHUTDOWN_GRACE_SECS", "BOBBIN_DEBUG_BIND", "BOBBIN_TRUSTED_PROXIES", + "BOBBIN_CORS_ORIGINS", "BOBBIN_HYDRANT_URL", "BOBBIN_START_CURSOR", "BOBBIN_INGEST_PARALLELISM", @@ -153,6 +156,11 @@ pub struct ServerConfig { default = [] )] pub trusted_proxies: Vec, + + /// Exact browser origins allowed to call XRPC, or `*` for every origin. + /// Empty disables CORS. When using as an env var, comma-separated. + #[config(env = "BOBBIN_CORS_ORIGINS", parse_env = Cors::parse_env, default = [])] + pub cors_origins: Cors, } impl ServerConfig { @@ -492,6 +500,14 @@ mod tests { path } + #[test] + fn cors_wildcard_cannot_mix_with_exact_origins() { + assert!(matches!(Cors::parse_env(""), Ok(Cors::Off))); + assert!(matches!(Cors::parse_env("*"), Ok(Cors::Any))); + assert!(Cors::parse_env("*,https://x").is_err()); + assert!(Cors::parse_env("https://x,*").is_err()); + } + #[test] fn known_keys_exactly_match_template_paths() { let template = template(); diff --git a/bobbin/crates/bobbin/src/main.rs b/bobbin/crates/bobbin/src/main.rs index 8ca74da38..5381994d0 100644 --- a/bobbin/crates/bobbin/src/main.rs +++ b/bobbin/crates/bobbin/src/main.rs @@ -455,6 +455,7 @@ async fn run(cfg: BobbinConfig) -> anyhow::Result<()> { .with_mirror(mirror) .with_mirror_v2(mirror_v2) .with_proxies(trusted_proxies) + .with_cors(cfg.server.cors_origins.clone()) .with_service_did(cfg.service_auth.did.clone()) .with_settlements(settlements) .with_invites(invites) diff --git a/bobbin/crates/xrpc/Cargo.toml b/bobbin/crates/xrpc/Cargo.toml index 061dd61a4..f1a4e3865 100644 --- a/bobbin/crates/xrpc/Cargo.toml +++ b/bobbin/crates/xrpc/Cargo.toml @@ -30,7 +30,7 @@ serde = { workspace = true } serde_json = { workspace = true } thiserror = { workspace = true } tower = { workspace = true } -tower-http = { workspace = true, features = ["trace"] } +tower-http = { workspace = true, features = ["cors", "trace"] } tracing = { workspace = true } trusted-proxies = { workspace = true } url = { workspace = true } diff --git a/bobbin/crates/xrpc/src/lib.rs b/bobbin/crates/xrpc/src/lib.rs index 3af834add..4ed2bf190 100644 --- a/bobbin/crates/xrpc/src/lib.rs +++ b/bobbin/crates/xrpc/src/lib.rs @@ -11,7 +11,7 @@ use axum::{ body::{Body, Bytes}, extract::{FromRequestParts, Query, RawQuery, State, rejection::QueryRejection}, http::{ - HeaderMap, HeaderName, StatusCode, + HeaderMap, HeaderName, HeaderValue, Method, StatusCode, header::{ ACCEPT_RANGES, AUTHORIZATION, CACHE_CONTROL, CONTENT_DISPOSITION, CONTENT_ENCODING, CONTENT_LENGTH, CONTENT_RANGE, CONTENT_SECURITY_POLICY, CONTENT_TYPE, ETAG, @@ -96,7 +96,7 @@ use jacquard_common::types::string::{AtUri, Cid}; use jacquard_common::xrpc::XrpcResp; use jacquard_common::{DefaultStr, IntoStatic}; use jacquard_identity::JacquardResolver; -use serde::{Deserialize, Serialize}; +use serde::{Deserialize, Deserializer, Serialize}; use std::convert::Infallible; use std::time::Duration; use thiserror::Error; @@ -104,6 +104,7 @@ use tokio::time::Instant; use url::form_urlencoded; use tower_http::classify::ServerErrorsFailureClass; +use tower_http::cors::{AllowOrigin, CorsLayer}; use tower_http::trace::{DefaultMakeSpan, OnFailure, OnResponse, TraceLayer}; use tracing::{Level, Span}; @@ -145,6 +146,67 @@ const MAX_OUTSTANDING_OFFERS: usize = 128; pub type Directory = JacquardResolver; +#[derive(Clone, Debug, Default)] +pub enum Cors { + #[default] + Off, + Any, + Origins(Vec), +} + +#[derive(Debug, Error)] +#[error("{0}")] +pub struct CorsOriginError(String); + +impl Cors { + pub fn parse_env(raw: &str) -> Result { + Self::from_origins( + raw.split(',') + .map(str::trim) + .filter(|entry| !entry.is_empty()), + ) + } + + fn from_origins<'a>( + origins: impl IntoIterator, + ) -> Result { + let origins: Vec<&str> = origins.into_iter().collect(); + match origins.as_slice() { + [] => Ok(Self::Off), + ["*"] => Ok(Self::Any), + _ if origins.contains(&"*") => Err(CorsOriginError( + "CORS wildcard `*` cannot be combined with exact origins".into(), + )), + _ => origins + .into_iter() + .map(|origin| { + let url = url::Url::parse(origin).map_err(|e| { + CorsOriginError(format!("invalid CORS origin `{origin}`: {e}")) + })?; + if !matches!(url.scheme(), "http" | "https") + || url.origin().ascii_serialization() != origin + { + return Err(CorsOriginError(format!( + "CORS origin must be an exact http(s) origin: `{origin}`" + ))); + } + origin + .parse() + .map_err(|e| CorsOriginError(format!("invalid CORS header: {e}"))) + }) + .collect::, _>>() + .map(Self::Origins), + } + } +} + +impl<'de> Deserialize<'de> for Cors { + fn deserialize>(deserializer: D) -> Result { + let origins = Vec::::deserialize(deserializer)?; + Self::from_origins(origins.iter().map(String::as_str)).map_err(serde::de::Error::custom) + } +} + pub fn default_directory() -> Directory { JacquardResolver::new(ReqwestHttp::new(reqwest::Client::new()), Default::default()) } @@ -177,6 +239,7 @@ pub struct AppState { pub invites: Arc, /// the ingest stream's liveness, absent when nothing publishes it pub stream_health: Option>, + cors: Cors, service_auth_config: service_auth::ServiceAuthConfig>, enrich_router: Arc>, trending_cache: Arc>>, @@ -226,6 +289,7 @@ impl AppState { settlements: Arc::new(Settlements::new()), invites: Arc::new(InviteIndex::new()), stream_health: None, + cors: Cors::Off, service_auth_config: ServiceAuthConfig::new( Did::new_static("did:web:localhost").unwrap(), directory, @@ -296,6 +360,11 @@ impl AppState { self } + pub fn with_cors(mut self, cors: Cors) -> Self { + self.cors = cors; + self + } + pub fn with_service_did(mut self, did: Did) -> Self { self.service_auth_config = ServiceAuthConfig::new(did, self.directory.clone()); self @@ -337,7 +406,34 @@ impl service_auth::ServiceAuth for AppState { } pub fn router(state: AppState) -> Router { - Router::new() + let cors = match &state.cors { + Cors::Off => None, + Cors::Any => Some(AllowOrigin::any()), + Cors::Origins(origins) => Some(AllowOrigin::list(origins.clone())), + } + .map(|origins| { + CorsLayer::new() + .allow_origin(origins) + .allow_methods([Method::GET, Method::POST, Method::OPTIONS]) + .allow_headers([ + AUTHORIZATION, + CONTENT_TYPE, + HeaderName::from_static("atproto-proxy"), + HeaderName::from_static("atproto-accept-labelers"), + HeaderName::from_static("range"), + HeaderName::from_static("if-range"), + HeaderName::from_static("if-none-match"), + HeaderName::from_static("if-modified-since"), + ]) + .expose_headers([ + CONTENT_TYPE, + HeaderName::from_static("ratelimit-limit"), + HeaderName::from_static("ratelimit-remaining"), + HeaderName::from_static("ratelimit-reset"), + ]) + .max_age(Duration::from_secs(86400)) + }); + let router = Router::new() .route("/xrpc/sh.tangled.repo.getRepo", get(get_repo)) .route("/xrpc/sh.tangled.repo.getRepos", get(get_repos)) .route( @@ -606,8 +702,12 @@ pub fn router(state: AppState) -> Router { .on_request(()) .on_response(LatencyFreeTrace) .on_failure(LatencyFreeTrace), - ) - .with_state(state) + ); + let router = match cors { + Some(cors) => router.layer(cors), + None => router, + }; + router.with_state(state) } #[derive(Clone, Copy, Debug)] diff --git a/bobbin/crates/xrpc/tests/coverage.rs b/bobbin/crates/xrpc/tests/coverage.rs index 036981b9a..00e2981f8 100644 --- a/bobbin/crates/xrpc/tests/coverage.rs +++ b/bobbin/crates/xrpc/tests/coverage.rs @@ -10,7 +10,7 @@ use bobbin_resolver::RepoIdResolver; use bobbin_runtime::{RuntimeHasher, SystemClock, UnixMicros}; use bobbin_search::{DEFAULT_WRITER_HEAP_BYTES, SearchIndex, SearchReader}; use bobbin_slingshot_client::SlingshotClient; -use bobbin_xrpc::{AppState, router}; +use bobbin_xrpc::{AppState, Cors, router}; use http::{Request, StatusCode}; use serde_json::{Value, json}; use tower::ServiceExt; @@ -158,3 +158,176 @@ async fn promotion_flips_ready_field() { assert_eq!(after["ready"], json!(true)); assert_eq!(after["lastCursor"], json!(9)); } + +const WEB_ORIGIN: &str = "https://next.tangled.org"; +const COVERAGE_PATH: &str = "/xrpc/sh.tangled.bobbin.getCoverage"; + +fn cors_state(state: AppState) -> AppState { + state.with_cors(Cors::Origins(vec![WEB_ORIGIN.parse().unwrap()])) +} + +fn preflight(origin: &str) -> Request { + Request::builder() + .method("OPTIONS") + .uri(COVERAGE_PATH) + .header("origin", origin) + .header("access-control-request-method", "POST") + .header( + "access-control-request-headers", + "authorization,content-type,atproto-proxy,atproto-accept-labelers", + ) + .body(Body::empty()) + .unwrap() +} + +#[tokio::test] +async fn allowed_origin_preflight_lists_methods_and_headers() { + let h = Harness::new().await; + let response = router(cors_state(h.state)) + .oneshot(preflight(WEB_ORIGIN)) + .await + .unwrap(); + let headers = response.headers(); + assert_eq!( + headers.get("access-control-allow-origin").unwrap(), + WEB_ORIGIN + ); + assert_eq!(response.status(), StatusCode::OK); + let methods = headers + .get("access-control-allow-methods") + .unwrap() + .to_str() + .unwrap(); + for method in ["GET", "POST", "OPTIONS"] { + assert!( + methods.split(',').any(|entry| entry.trim() == method), + "{methods}" + ); + } + let allowed = headers + .get("access-control-allow-headers") + .unwrap() + .to_str() + .unwrap(); + for header in [ + "authorization", + "content-type", + "atproto-proxy", + "atproto-accept-labelers", + ] { + assert!( + allowed.split(',').any(|entry| entry.trim() == header), + "{allowed}" + ); + } + assert_eq!(headers.get("access-control-max-age").unwrap(), "86400"); +} + +#[tokio::test] +async fn disallowed_origin_has_no_acao() { + let h = Harness::new().await; + let response = router(cors_state(h.state)) + .oneshot(preflight("https://other.tangled.org")) + .await + .unwrap(); + assert!( + response + .headers() + .get("access-control-allow-origin") + .is_none() + ); +} + +#[tokio::test] +async fn unset_origins_keep_responses_without_cors_headers() { + let h = Harness::new().await; + let app = router(h.state); + let preflight = app.clone().oneshot(preflight(WEB_ORIGIN)).await.unwrap(); + let get = app + .oneshot( + Request::builder() + .uri(COVERAGE_PATH) + .header("origin", WEB_ORIGIN) + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + for response in [preflight, get] { + assert!( + !response + .headers() + .keys() + .any(|name| name.as_str().starts_with("access-control-")) + ); + } +} + +#[tokio::test] +async fn allowed_origin_get_has_acao() { + let h = Harness::new().await; + let response = router(cors_state(h.state)) + .oneshot( + Request::builder() + .uri(COVERAGE_PATH) + .header("origin", WEB_ORIGIN) + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::OK); + assert_eq!( + response + .headers() + .get("access-control-allow-origin") + .unwrap(), + WEB_ORIGIN + ); +} + +#[tokio::test] +async fn any_origin_allows_preflight_and_get() { + let h = Harness::new().await; + let app = router(h.state.with_cors(Cors::Any)); + let preflight = app + .clone() + .oneshot(preflight("https://elsewhere.example")) + .await + .unwrap(); + assert_eq!(preflight.status(), StatusCode::OK); + assert_eq!( + preflight + .headers() + .get("access-control-allow-origin") + .unwrap(), + "*" + ); + assert!( + preflight + .headers() + .get("access-control-allow-methods") + .is_some() + ); + assert!( + preflight + .headers() + .get("access-control-allow-headers") + .is_some() + ); + let get = app + .oneshot( + Request::builder() + .uri(COVERAGE_PATH) + .header("origin", "https://another.example") + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(get.status(), StatusCode::OK); + assert_eq!( + get.headers().get("access-control-allow-origin").unwrap(), + "*" + ); +} diff --git a/bobbin/example.toml b/bobbin/example.toml index e203ff367..f02c68343 100644 --- a/bobbin/example.toml +++ b/bobbin/example.toml @@ -41,6 +41,15 @@ # Default value: [] #trusted_proxies = [] +# Exact browser origins allowed to call XRPC, or `*` for every origin. +# Empty disables CORS. When using as an env var, comma-separated. +# `*` cannot be combined with exact origins. +# +# Can also be specified via environment variable `BOBBIN_CORS_ORIGINS`. +# +# Default value: [] +#cors_origins = [] + [hydrant] # Base URL of the hydrant instance - the cursor-replayable /stream lives # under this. Use `ws://` or `wss://` - `http://` and `https://`