From c4df5a19358a6cafcb1fdb6b20f5319310cbcabd Mon Sep 17 00:00:00 2001 From: "@permadeath.com" Date: Sat, 1 Aug 2026 16:06:41 -0400 Subject: [PATCH] M2: ATProto OAuth login atrium wrapped in one module; sqlite state/session stores; signed headquarters_session cookie; bidirectional handle verification; loopback and confidential client shapes selected by PUBLIC_URL. The exchange's todo!() in atrium 0.1.7 is contained by running the callback in its own task. Co-Authored-By: Claude Fable 5 --- Cargo.lock | 2 + Cargo.toml | 4 + services/api/Cargo.toml | 2 + services/api/src/atproto/client.rs | 151 ++++++++++ services/api/src/atproto/mod.rs | 244 ++++++++++++++++ services/api/src/atproto/store.rs | 243 ++++++++++++++++ services/api/src/atproto/verify.rs | 51 ++++ services/api/src/config.rs | 217 ++++++++++++++- services/api/src/db.rs | 228 +++++++++++++++ services/api/src/main.rs | 65 +++-- services/api/src/routes.rs | 430 +++++++++++++++++++++++++++++ services/api/src/session.rs | 183 ++++++++++++ 12 files changed, 1788 insertions(+), 32 deletions(-) create mode 100644 services/api/src/atproto/client.rs create mode 100644 services/api/src/atproto/mod.rs create mode 100644 services/api/src/atproto/store.rs create mode 100644 services/api/src/atproto/verify.rs create mode 100644 services/api/src/db.rs create mode 100644 services/api/src/routes.rs create mode 100644 services/api/src/session.rs diff --git a/Cargo.lock b/Cargo.lock index 5d91c3a..05834fd 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -864,7 +864,9 @@ dependencies = [ "axum", "axum-extra", "base64", + "cookie", "hmac", + "http-body-util", "jose-jwk", "rand 0.8.7", "reqwest", diff --git a/Cargo.toml b/Cargo.toml index e8379f4..deba0e2 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -39,6 +39,9 @@ rusqlite = { version = "0.40", features = ["bundled"] } tower-http = { version = "0.6", features = ["cors"] } axum-extra = { version = "0.10", features = ["cookie"] } +# The same cookie crate axum-extra uses; direct only for time::Duration on +# Max-Age. +cookie = "0.18" hmac = "0.12" sha2 = "0.10" @@ -49,3 +52,4 @@ serde = { version = "1", features = ["derive"] } serde_json = "1" tempfile = "3" +http-body-util = "0.1" diff --git a/services/api/Cargo.toml b/services/api/Cargo.toml index 8a25b78..a0ebe6b 100644 --- a/services/api/Cargo.toml +++ b/services/api/Cargo.toml @@ -22,6 +22,7 @@ rusqlite.workspace = true tower-http.workspace = true axum-extra.workspace = true +cookie.workspace = true hmac.workspace = true sha2.workspace = true @@ -34,3 +35,4 @@ serde_json.workspace = true [dev-dependencies] tower.workspace = true tempfile.workspace = true +http-body-util.workspace = true diff --git a/services/api/src/atproto/client.rs b/services/api/src/atproto/client.rs new file mode 100644 index 0000000..160fc69 --- /dev/null +++ b/services/api/src/atproto/client.rs @@ -0,0 +1,151 @@ +//! Construction of the one concrete OAuthClient, in both shapes. +//! +//! The shapes differ only in the metadata fed to `OAuthClient::new`: +//! loopback (development) encodes its metadata in the client_id itself; +//! confidential (production) serves metadata and a JWKS from PUBLIC_URL and +//! authenticates to the token endpoint with an ES256 key. + +use std::sync::Arc; + +use atrium_identity::did::{CommonDidResolver, CommonDidResolverConfig, DEFAULT_PLC_DIRECTORY_URL}; +use atrium_identity::handle::{ + AtprotoHandleResolver, AtprotoHandleResolverConfig, DohDnsTxtResolver, DohDnsTxtResolverConfig, +}; +use atrium_oauth::store::session::SessionStore; +use atrium_oauth::{ + AtprotoClientMetadata, AtprotoLocalhostClientMetadata, AuthMethod, GrantType, KnownScope, + OAuthClient, OAuthClientConfig, OAuthResolverConfig, Scope, +}; +use jose_jwk::Jwk; + +use super::store::{SqliteSessionStore, SqliteStateStore}; +use crate::config::Config; +use crate::db::Db; + +/// atrium's own DefaultHttpClient means reqwest on native-tls means a system +/// OpenSSL; this is the same shim on the rustls build of reqwest instead. +#[derive(Clone, Default)] +pub struct HttpClient { + client: reqwest::Client, +} + +impl atrium_xrpc::HttpClient for HttpClient { + async fn send_http( + &self, + request: atrium_xrpc::http::Request>, + ) -> Result< + atrium_xrpc::http::Response>, + Box, + > { + let response = self.client.execute(request.try_into()?).await?; + let mut builder = atrium_xrpc::http::Response::builder().status(response.status()); + for (name, value) in response.headers() { + builder = builder.header(name, value); + } + builder + .body(response.bytes().await?.to_vec()) + .map_err(Into::into) + } +} + +// Handle lookups go to Cloudflare's DoH endpoint: atrium's only +// batteries-included DnsTxtResolver is DNS-over-HTTPS, and it behaves the +// same in a container as on a laptop. Swapping in a system-DNS resolver later +// is a change confined to this module. +const DOH_SERVICE_URL: &str = "https://mozilla.cloudflare-dns.com/dns-query"; + +pub type DidResolver = CommonDidResolver; +pub type HandleResolver = AtprotoHandleResolver, HttpClient>; +pub type Client = + OAuthClient; + +pub fn did_resolver(http: &HttpClient) -> DidResolver { + CommonDidResolver::new(CommonDidResolverConfig { + plc_directory_url: DEFAULT_PLC_DIRECTORY_URL.to_owned(), + http_client: Arc::new(http.clone()), + }) +} + +pub fn handle_resolver(http: &HttpClient) -> HandleResolver { + AtprotoHandleResolver::new(AtprotoHandleResolverConfig { + dns_txt_resolver: DohDnsTxtResolver::new(DohDnsTxtResolverConfig { + service_url: DOH_SERVICE_URL.to_owned(), + http_client: Arc::new(http.clone()), + }), + http_client: Arc::new(http.clone()), + }) +} + +/// `PRIVATE_KEY_JWK` accepts one JWK object or an array, so a rotation can +/// publish a new key in the JWKS before the old one stops signing (TODO.md). +/// Every key must be an EC private key with a `kid` — atrium's Keyset rejects +/// anything else at construction. +fn private_keys(jwk: &str) -> Result, String> { + if let Ok(one) = serde_json::from_str::(jwk) { + return Ok(vec![one]); + } + serde_json::from_str::>(jwk) + .map_err(|_| "PRIVATE_KEY_JWK is neither a JWK object nor an array of JWKs".to_owned()) +} + +pub fn build(config: &Config, db: Db, http: &HttpClient) -> Result { + let scopes = vec![Scope::Known(KnownScope::Atproto)]; + let resolver = OAuthResolverConfig { + did_resolver: did_resolver(http), + handle_resolver: handle_resolver(http), + authorization_server_metadata: Default::default(), + protected_resource_metadata: Default::default(), + }; + let state_store = SqliteStateStore(db.clone()); + let session_store = SqliteSessionStore(db); + + let client = match &config.public_url { + None => OAuthClient::new(OAuthClientConfig { + client_metadata: AtprotoLocalhostClientMetadata { + redirect_uris: Some(vec![config.redirect_uri()]), + scopes: Some(scopes), + }, + keys: None, + state_store, + session_store, + resolver, + http_client: http.clone(), + }), + Some(url) => { + let keys = private_keys( + config + .private_key_jwk + .as_deref() + .expect("config validated: key iff public"), + )?; + OAuthClient::new(OAuthClientConfig { + client_metadata: AtprotoClientMetadata { + client_id: format!("{url}/oauth/client-metadata.json"), + client_uri: None, + redirect_uris: vec![config.redirect_uri()], + token_endpoint_auth_method: AuthMethod::PrivateKeyJwt, + grant_types: vec![GrantType::AuthorizationCode, GrantType::RefreshToken], + scopes, + jwks_uri: Some(format!("{url}/.well-known/jwks.json")), + token_endpoint_auth_signing_alg: Some("ES256".to_owned()), + }, + keys: Some(keys), + state_store, + session_store, + resolver, + http_client: http.clone(), + }) + } + }; + + // The error's Display can only name what was wrong with our own + // configuration; no runtime secret flows through construction messages. + client.map_err(|e| format!("oauth client construction: {e}")) +} + +// Satisfy the SessionStore bound's error requirement explicitly, so a future +// atrium bump that changes the bound fails here with a readable error. +fn _assert_bounds() { + fn requires() {} + requires::(); +} diff --git a/services/api/src/atproto/mod.rs b/services/api/src/atproto/mod.rs new file mode 100644 index 0000000..ee8eabc --- /dev/null +++ b/services/api/src/atproto/mod.rs @@ -0,0 +1,244 @@ +//! The one wrapper around atrium. Everything else in this crate talks to +//! ATProto through the `Atproto` struct below; only this module and its +//! submodules import `atrium_*` or `jose_jwk` — the crates are 0.x and churn, +//! and this is the wall the churn stops at (TODO.md). + +mod client; +mod store; +mod verify; + +use std::sync::Arc; + +use atrium_api::agent::SessionManager; +use atrium_api::types::string::{Did, Handle}; +use atrium_oauth::{AuthorizeOptions, CallbackParams}; + +use crate::config::Config; +use crate::db::Db; + +pub enum LoginError { + /// The input is not a handle. + InvalidHandle, + /// The account's own infrastructure could not be reached or refused. + Upstream, +} + +/// The exchange failed after consent. Denial never reaches the wrapper: the +/// authorization server signals it with `error=` and no code, which the +/// route answers by itself. +pub struct CallbackFailed; + +pub struct Atproto { + client: Arc, + did_resolver: client::DidResolver, + handle_resolver: client::HandleResolver, + confidential: bool, +} + +impl Atproto { + pub fn new(config: &Config, db: Db) -> Result { + let http = client::HttpClient::default(); + Ok(Atproto { + client: Arc::new(client::build(config, db, &http)?), + // A second pair: the first is owned by the OAuthClient. + did_resolver: client::did_resolver(&http), + handle_resolver: client::handle_resolver(&http), + confidential: config.public_url.is_some(), + }) + } + + /// Starts the flow for a handle: resolves it, pushes the authorization + /// request, returns the URL to send the browser to — the consent screen + /// on the account's own server. + pub async fn begin_login(&self, handle: &str) -> Result { + let handle = handle.trim(); + if Handle::new(handle.to_owned()).is_err() { + return Err(LoginError::InvalidHandle); + } + self.client + .authorize(handle, AuthorizeOptions::default()) + .await + .map_err(|_| { + // atrium's error values can embed request URIs and server + // responses; log the fact, not the value. + tracing::warn!("login: authorization could not be started"); + LoginError::Upstream + }) + } + + /// Completes the flow from the callback query: code for tokens, tokens + /// into the store, handle verified both directions. Returns the DID and + /// the verified handle. + pub async fn complete_login( + &self, + code: String, + state: Option, + iss: Option, + ) -> Result<(String, Option), CallbackFailed> { + let client = Arc::clone(&self.client); + let params = CallbackParams { code, state, iss }; + // atrium 0.1.7 hits a todo!() when the token exchange itself fails — + // a panic, not an Err. Run the callback in its own task so that + // surfaces as a JoinError and the user gets the failure redirect + // instead of a dead connection. + let session = match tokio::spawn(async move { client.callback(params).await }).await { + Ok(Ok((session, _app_state))) => session, + Ok(Err(_)) => { + tracing::warn!("login: callback failed"); + return Err(CallbackFailed); + } + Err(_) => { + tracing::warn!("login: token exchange failed"); + return Err(CallbackFailed); + } + }; + + let Some(did) = session.did().await else { + tracing::warn!("login: session has no DID"); + return Err(CallbackFailed); + }; + let handle = verify::verify_handle(&self.did_resolver, &self.handle_resolver, &did).await; + Ok((did.as_str().to_owned(), handle)) + } + + /// Best-effort server-side revocation. The stored session row is deleted + /// by atrium on success; the caller deletes it unconditionally as + /// belt-and-braces either way. + pub async fn logout(&self, did: &str) { + let Ok(did) = Did::new(did.to_owned()) else { + return; + }; + if self.client.revoke(&did).await.is_err() { + tracing::info!("logout: server-side revocation failed; session row deleted anyway"); + } + } + + pub fn is_confidential(&self) -> bool { + self.confidential + } + + /// Served at /oauth/client-metadata.json in confidential mode. + pub fn client_metadata(&self) -> serde_json::Value { + serde_json::to_value(&self.client.client_metadata).expect("metadata serializes") + } + + /// Served at /.well-known/jwks.json in confidential mode. atrium strips + /// the private members before handing this out. + pub fn jwks(&self) -> serde_json::Value { + serde_json::to_value(self.client.jwks()).expect("jwks serializes") + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::collections::HashMap; + + fn temp_db() -> (tempfile::TempDir, Db) { + let dir = tempfile::tempdir().unwrap(); + let db = Db::open(&dir.path().join("test.sqlite")).unwrap(); + (dir, db) + } + + fn config(vars: &[(&str, &str)]) -> Config { + let map: HashMap = vars + .iter() + .map(|(k, v)| (k.to_string(), v.to_string())) + .collect(); + Config::from_lookup(|key| map.get(key).cloned()).unwrap() + } + + // The RFC 7515 A.3 example key with a kid added: a published test + // vector, useful precisely because it is not a secret. Nothing may ever + // deploy it. + const TEST_JWK: &str = r#"{ + "kty": "EC", "crv": "P-256", "kid": "test-1", + "x": "f83OJ3D2xF1Bg8vub9tLe1gHMzV76e8Tus9uPHvRVEU", + "y": "x_FEzRu9m36HLN_tue659LNpXW6pCyStikYjKIWI5a0", + "d": "jpsQnnGQmL-YBIffH1136cspYG6-0iY7X1fCE9-E9LI" + }"#; + + #[tokio::test] + async fn loopback_shape() { + let (_dir, db) = temp_db(); + let atproto = Atproto::new(&config(&[]), db).unwrap(); + + assert!(!atproto.is_confidential()); + let metadata = atproto.client_metadata(); + let client_id = metadata["client_id"].as_str().unwrap(); + assert!( + client_id.starts_with("http://localhost?"), + "got {client_id}" + ); + assert!( + client_id.contains("127.0.0.1"), + "redirect must be loopback IP, got {client_id}" + ); + } + + #[tokio::test] + async fn confidential_shape() { + let (_dir, db) = temp_db(); + let atproto = Atproto::new( + &config(&[ + ("PUBLIC_URL", "https://api.lance.blue"), + ("WEB_ORIGIN", "https://lance.blue"), + ("SESSION_SECRET", "0123456789abcdef0123456789abcdef"), + ("PRIVATE_KEY_JWK", TEST_JWK), + ]), + db, + ) + .unwrap(); + + assert!(atproto.is_confidential()); + let metadata = atproto.client_metadata(); + assert_eq!( + metadata["client_id"], + "https://api.lance.blue/oauth/client-metadata.json" + ); + assert_eq!(metadata["token_endpoint_auth_method"], "private_key_jwt"); + assert_eq!( + metadata["jwks_uri"], + "https://api.lance.blue/.well-known/jwks.json" + ); + assert_eq!( + metadata["redirect_uris"][0], + "https://api.lance.blue/oauth/callback" + ); + assert_eq!(metadata["scope"], "atproto"); + + let jwks = atproto.jwks(); + let key = &jwks["keys"][0]; + assert_eq!(key["kid"], "test-1"); + assert!(key.get("d").is_none(), "private member served: {jwks}"); + } + + #[tokio::test] + async fn confidential_requires_kid() { + let (_dir, db) = temp_db(); + let no_kid = TEST_JWK.replace(r#""kid": "test-1","#, ""); + let result = Atproto::new( + &config(&[ + ("PUBLIC_URL", "https://api.lance.blue"), + ("SESSION_SECRET", "0123456789abcdef0123456789abcdef"), + ("PRIVATE_KEY_JWK", &no_kid), + ]), + db, + ); + assert!(result.is_err(), "a key without a kid must be rejected"); + } + + #[tokio::test] + async fn begin_login_rejects_non_handles() { + let (_dir, db) = temp_db(); + let atproto = Atproto::new(&config(&[]), db).unwrap(); + assert!(matches!( + atproto.begin_login("not a handle").await, + Err(LoginError::InvalidHandle) + )); + assert!(matches!( + atproto.begin_login("").await, + Err(LoginError::InvalidHandle) + )); + } +} diff --git a/services/api/src/atproto/store.rs b/services/api/src/atproto/store.rs new file mode 100644 index 0000000..da49b34 --- /dev/null +++ b/services/api/src/atproto/store.rs @@ -0,0 +1,243 @@ +//! atrium's state and session stores over the sqlite file. +//! +//! The stored values are atrium's own serde types, kept as JSON in the +//! `value` column. The session rows hold token sets — they are never logged +//! and the file is 0600 (db.rs). + +use atrium_api::types::string::Did; +use atrium_common::store::Store; +use atrium_oauth::store::session::{Session, SessionStore}; +use atrium_oauth::store::state::{InternalStateData, StateStore}; +use rusqlite::OptionalExtension; + +use crate::db::{Db, DbError}; + +#[derive(Debug)] +pub enum StoreError { + Db(DbError), + Json(serde_json::Error), +} + +impl std::fmt::Display for StoreError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + StoreError::Db(e) => write!(f, "store: {e}"), + StoreError::Json(e) => write!(f, "store serialization: {e}"), + } + } +} + +impl std::error::Error for StoreError {} + +impl From for StoreError { + fn from(e: DbError) -> StoreError { + StoreError::Db(e) + } +} + +pub struct SqliteStateStore(pub(super) Db); + +impl Store for SqliteStateStore { + type Error = StoreError; + + async fn get(&self, key: &String) -> Result, StoreError> { + let key = key.clone(); + let row: Option = self + .0 + .call(move |conn| { + conn.query_row( + "SELECT value FROM oauth_state WHERE key = ?1", + [key], + |row| row.get(0), + ) + .optional() + }) + .await?; + row.map(|json| serde_json::from_str(&json)) + .transpose() + .map_err(StoreError::Json) + } + + async fn set(&self, key: String, value: InternalStateData) -> Result<(), StoreError> { + let json = serde_json::to_string(&value).map_err(StoreError::Json)?; + self.0 + .call(move |conn| { + conn.execute( + "INSERT INTO oauth_state (key, value) VALUES (?1, ?2) + ON CONFLICT(key) DO UPDATE SET value = excluded.value", + (key, json), + ) + .map(|_| ()) + }) + .await?; + Ok(()) + } + + async fn del(&self, key: &String) -> Result<(), StoreError> { + let key = key.clone(); + self.0 + .call(move |conn| { + conn.execute("DELETE FROM oauth_state WHERE key = ?1", [key]) + .map(|_| ()) + }) + .await?; + Ok(()) + } + + async fn clear(&self) -> Result<(), StoreError> { + self.0 + .call(|conn| conn.execute("DELETE FROM oauth_state", []).map(|_| ())) + .await?; + Ok(()) + } +} + +impl StateStore for SqliteStateStore {} + +pub struct SqliteSessionStore(pub(super) Db); + +impl Store for SqliteSessionStore { + type Error = StoreError; + + async fn get(&self, key: &Did) -> Result, StoreError> { + let did = key.as_str().to_owned(); + let row: Option = self + .0 + .call(move |conn| { + conn.query_row( + "SELECT value FROM oauth_session WHERE did = ?1", + [did], + |row| row.get(0), + ) + .optional() + }) + .await?; + row.map(|json| serde_json::from_str(&json)) + .transpose() + .map_err(StoreError::Json) + } + + async fn set(&self, key: Did, value: Session) -> Result<(), StoreError> { + let did = key.as_str().to_owned(); + let json = serde_json::to_string(&value).map_err(StoreError::Json)?; + self.0 + .call(move |conn| { + conn.execute( + "INSERT INTO oauth_session (did, value) VALUES (?1, ?2) + ON CONFLICT(did) DO UPDATE SET + value = excluded.value, + updated_at = unixepoch()", + (did, json), + ) + .map(|_| ()) + }) + .await?; + Ok(()) + } + + async fn del(&self, key: &Did) -> Result<(), StoreError> { + let did = key.as_str().to_owned(); + self.0 + .call(move |conn| { + conn.execute("DELETE FROM oauth_session WHERE did = ?1", [did]) + .map(|_| ()) + }) + .await?; + Ok(()) + } + + async fn clear(&self) -> Result<(), StoreError> { + self.0 + .call(|conn| conn.execute("DELETE FROM oauth_session", []).map(|_| ())) + .await?; + Ok(()) + } +} + +impl SessionStore for SqliteSessionStore {} + +#[cfg(test)] +mod tests { + use super::*; + + // The RFC 7515 A.3 example key: a published test vector, useful precisely + // because it is not a secret. Nothing may ever deploy it. + fn test_dpop_key() -> serde_json::Value { + serde_json::json!({ + "kty": "EC", + "crv": "P-256", + "x": "f83OJ3D2xF1Bg8vub9tLe1gHMzV76e8Tus9uPHvRVEU", + "y": "x_FEzRu9m36HLN_tue659LNpXW6pCyStikYjKIWI5a0", + "d": "jpsQnnGQmL-YBIffH1136cspYG6-0iY7X1fCE9-E9LI" + }) + } + + fn state_value() -> InternalStateData { + serde_json::from_value(serde_json::json!({ + "iss": "https://pds.example", + "dpop_key": test_dpop_key(), + "verifier": "not-a-real-verifier", + "app_state": null, + })) + .unwrap() + } + + fn session_value() -> Session { + serde_json::from_value(serde_json::json!({ + "dpop_key": test_dpop_key(), + "token_set": { + "iss": "https://pds.example", + "sub": "did:plc:222222222222222222222222", + "aud": "https://pds.example", + "scope": "atproto", + "refresh_token": "not-a-real-refresh-token", + "access_token": "not-a-real-access-token", + "token_type": "DPoP", + "expires_at": null, + }, + })) + .unwrap() + } + + fn temp_db() -> (tempfile::TempDir, Db) { + let dir = tempfile::tempdir().unwrap(); + let db = Db::open(&dir.path().join("test.sqlite")).unwrap(); + (dir, db) + } + + #[tokio::test] + async fn state_round_trip() { + let (_dir, db) = temp_db(); + let store = SqliteStateStore(db); + + assert!(store.get(&"nonce".to_owned()).await.unwrap().is_none()); + + store.set("nonce".to_owned(), state_value()).await.unwrap(); + let loaded = store.get(&"nonce".to_owned()).await.unwrap().unwrap(); + assert_eq!(loaded.iss, "https://pds.example"); + assert_eq!(loaded.verifier, "not-a-real-verifier"); + + store.del(&"nonce".to_owned()).await.unwrap(); + assert!(store.get(&"nonce".to_owned()).await.unwrap().is_none()); + } + + #[tokio::test] + async fn session_round_trip_survives_reopen() { + let dir = tempfile::tempdir().unwrap(); + let path = dir.path().join("test.sqlite"); + let did = Did::new("did:plc:222222222222222222222222".to_owned()).unwrap(); + + let store = SqliteSessionStore(Db::open(&path).unwrap()); + store.set(did.clone(), session_value()).await.unwrap(); + drop(store); + + // Persistence, not just cache: a fresh handle on the same file still + // has the session. + let store = SqliteSessionStore(Db::open(&path).unwrap()); + let loaded = store.get(&did).await.unwrap().unwrap(); + assert_eq!(loaded.token_set.access_token, "not-a-real-access-token"); + + store.clear().await.unwrap(); + assert!(store.get(&did).await.unwrap().is_none()); + } +} diff --git a/services/api/src/atproto/verify.rs b/services/api/src/atproto/verify.rs new file mode 100644 index 0000000..b1c0b74 --- /dev/null +++ b/services/api/src/atproto/verify.rs @@ -0,0 +1,51 @@ +//! Bidirectional handle verification. +//! +//! OAuth returns a DID. The DID document advertises a handle in +//! `alsoKnownAs`, but the account controls that document, so read alone it is +//! a claim anybody can make about any name. The binding only holds if the +//! handle also resolves *back* to the DID through the mechanism the domain +//! owner controls (docs/identity-and-sessions.md). atrium's handle resolver +//! is exactly that forward resolution; this module adds only the comparison. + +use atrium_api::types::string::{Did, Handle}; +use atrium_common::resolver::Resolver; + +use super::client::{DidResolver, HandleResolver}; + +/// The verified handle, or None for unverified — which the caller stores as +/// such and the UI shows next to the word "unverified". Any resolution +/// failure is unverified rather than an error: for some hosting setups that +/// is the honest answer, not a fault (docs/local-dev.md). +pub async fn verify_handle( + did_resolver: &DidResolver, + handle_resolver: &HandleResolver, + did: &Did, +) -> Option { + let document = match did_resolver.resolve(did).await { + Ok(document) => document, + Err(_) => { + tracing::warn!("handle verification: DID document fetch failed"); + return None; + } + }; + + let claimed = document + .also_known_as + .as_ref()? + .iter() + .find_map(|aka| aka.strip_prefix("at://"))? + .to_owned(); + let handle = Handle::new(claimed).ok()?; + + match handle_resolver.resolve(&handle).await { + Ok(resolved) if resolved == *did => Some(handle.as_str().to_owned()), + Ok(_) => { + tracing::warn!("handle verification: handle resolves to a different DID"); + None + } + Err(_) => { + tracing::warn!("handle verification: handle does not resolve"); + None + } + } +} diff --git a/services/api/src/config.rs b/services/api/src/config.rs index b1ae14d..0de003c 100644 --- a/services/api/src/config.rs +++ b/services/api/src/config.rs @@ -1,23 +1,222 @@ use std::net::SocketAddr; +use std::path::PathBuf; pub struct Config { pub bind_addr: SocketAddr, + /// The API's own public origin, e.g. `https://api.lance.blue`. Set, it + /// selects the confidential OAuth client and `Secure` cookies; unset, + /// development runs the loopback client over plain HTTP. + pub public_url: Option, + /// The site's origin — the one origin CORS admits and where the OAuth + /// callback sends the browser back to. + pub web_origin: String, + /// HMAC key for the session cookie. + pub session_secret: Vec, + /// ES256 private JWK (object or array), kept opaque here and parsed by + /// the atproto module. Required in public mode, rejected outside it. + pub private_key_jwk: Option, + pub db_path: PathBuf, +} + +// The default Debug would print the session secret and the private key. +impl std::fmt::Debug for Config { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("Config") + .field("bind_addr", &self.bind_addr) + .field("public_url", &self.public_url) + .field("web_origin", &self.web_origin) + .field("db_path", &self.db_path) + .finish_non_exhaustive() + } } impl Config { /// Reads configuration from the environment. - /// - /// `BIND_ADDR` defaults to loopback: on a laptop nothing else should be - /// able to reach a development instance (docs/local-dev.md). Inside a - /// container it must be `0.0.0.0:3000` — there, compose's port publish is - /// what keeps it on loopback from the outside. pub fn from_env() -> Result { - let bind_addr = match std::env::var("BIND_ADDR") { - Ok(s) => s + Self::from_lookup(|key| std::env::var(key).ok()) + } + + /// The environment is process-global, so tests use this with a map. + pub fn from_lookup(lookup: impl Fn(&str) -> Option) -> Result { + // `BIND_ADDR` defaults to loopback: on a laptop nothing else should + // be able to reach a development instance (docs/local-dev.md). Inside + // a container it must be `0.0.0.0:3000` — there, compose's port + // publish is what keeps it on loopback from the outside. + let bind_addr = match lookup("BIND_ADDR") { + Some(s) => s .parse() .map_err(|e| format!("BIND_ADDR {s:?} is not a socket address: {e}"))?, - Err(_) => SocketAddr::from(([127, 0, 0, 1], 3000)), + None => SocketAddr::from(([127, 0, 0, 1], 3000)), }; - Ok(Config { bind_addr }) + + let public_url = match lookup("PUBLIC_URL") { + Some(s) => { + let s = s.trim_end_matches('/').to_owned(); + if !s.starts_with("https://") || s.len() <= "https://".len() { + return Err(format!("PUBLIC_URL {s:?} must be an https:// URL")); + } + Some(s) + } + None => None, + }; + + let web_origin = match lookup("WEB_ORIGIN") { + Some(s) => { + let valid = (s.starts_with("https://") || s.starts_with("http://")) + && !s.ends_with('/') + && !s + .splitn(3, '/') + .nth(2) + .is_some_and(|rest| rest.contains('/')); + if !valid { + return Err(format!( + "WEB_ORIGIN {s:?} must be a bare origin - scheme://host[:port], no path or trailing slash" + )); + } + s + } + None => "http://127.0.0.1:5173".to_owned(), + }; + + let session_secret = match lookup("SESSION_SECRET") { + Some(s) => { + if s.len() < 32 { + return Err("SESSION_SECRET must be at least 32 bytes".to_owned()); + } + s.into_bytes() + } + None => { + if public_url.is_some() { + return Err("SESSION_SECRET is required when PUBLIC_URL is set".to_owned()); + } + // Development convenience only: sessions die with the process. + use rand::RngCore; + let mut secret = vec![0u8; 32]; + rand::rngs::OsRng.fill_bytes(&mut secret); + tracing::warn!("SESSION_SECRET not set; sessions will not survive a restart"); + secret + } + }; + + let private_key_jwk = lookup("PRIVATE_KEY_JWK"); + match (&public_url, &private_key_jwk) { + (Some(_), None) => { + return Err("PRIVATE_KEY_JWK is required when PUBLIC_URL is set".to_owned()); + } + (None, Some(_)) => { + return Err( + "PRIVATE_KEY_JWK is set but PUBLIC_URL is not; refusing a half-configured confidential client" + .to_owned(), + ); + } + _ => {} + } + + let db_path = lookup("DB_PATH") + .map(PathBuf::from) + .unwrap_or_else(|| PathBuf::from("./data/headquarters.sqlite")); + + Ok(Config { + bind_addr, + public_url, + web_origin, + session_secret, + private_key_jwk, + db_path, + }) + } + + /// The OAuth redirect URI. Loopback development derives it from the bind + /// port: the container binds 0.0.0.0:3000 but publishes on + /// 127.0.0.1:3000, so the port carries over. + pub fn redirect_uri(&self) -> String { + match &self.public_url { + Some(url) => format!("{url}/oauth/callback"), + None => format!("http://127.0.0.1:{}/oauth/callback", self.bind_addr.port()), + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::collections::HashMap; + + fn config(vars: &[(&str, &str)]) -> Result { + let map: HashMap = vars + .iter() + .map(|(k, v)| (k.to_string(), v.to_string())) + .collect(); + Config::from_lookup(|key| map.get(key).cloned()) + } + + const SECRET: &str = "0123456789abcdef0123456789abcdef"; + + #[test] + fn defaults_are_loopback_dev() { + let c = config(&[]).unwrap(); + assert_eq!(c.bind_addr, SocketAddr::from(([127, 0, 0, 1], 3000))); + assert_eq!(c.public_url, None); + assert_eq!(c.web_origin, "http://127.0.0.1:5173"); + assert_eq!(c.session_secret.len(), 32); + assert_eq!(c.redirect_uri(), "http://127.0.0.1:3000/oauth/callback"); + } + + #[test] + fn public_mode_requires_secret_and_key() { + let e = config(&[("PUBLIC_URL", "https://api.lance.blue")]).unwrap_err(); + assert!(e.contains("SESSION_SECRET")); + + let e = config(&[ + ("PUBLIC_URL", "https://api.lance.blue"), + ("SESSION_SECRET", SECRET), + ]) + .unwrap_err(); + assert!(e.contains("PRIVATE_KEY_JWK")); + } + + #[test] + fn public_mode_config_resolves() { + let c = config(&[ + ("PUBLIC_URL", "https://api.lance.blue/"), + ("SESSION_SECRET", SECRET), + ("PRIVATE_KEY_JWK", "{}"), + ("WEB_ORIGIN", "https://lance.blue"), + ]) + .unwrap(); + assert_eq!(c.public_url.as_deref(), Some("https://api.lance.blue")); + assert_eq!(c.redirect_uri(), "https://api.lance.blue/oauth/callback"); + assert_eq!(c.web_origin, "https://lance.blue"); + } + + #[test] + fn public_url_must_be_https() { + assert!(config(&[("PUBLIC_URL", "http://api.lance.blue")]).is_err()); + assert!(config(&[("PUBLIC_URL", "https://")]).is_err()); + } + + #[test] + fn key_without_public_url_is_refused() { + assert!(config(&[("PRIVATE_KEY_JWK", "{}")]).is_err()); + } + + #[test] + fn short_secret_is_refused() { + assert!(config(&[("SESSION_SECRET", "short")]).is_err()); + } + + #[test] + fn web_origin_must_be_bare() { + assert!(config(&[("WEB_ORIGIN", "https://lance.blue/app")]).is_err()); + assert!(config(&[("WEB_ORIGIN", "https://lance.blue/")]).is_err()); + assert!(config(&[("WEB_ORIGIN", "lance.blue")]).is_err()); + assert!(config(&[("WEB_ORIGIN", "http://127.0.0.1:5173")]).is_ok()); + } + + #[test] + fn debug_hides_secrets() { + let c = config(&[("SESSION_SECRET", SECRET)]).unwrap(); + let s = format!("{c:?}"); + assert!(!s.contains(SECRET)); } } diff --git a/services/api/src/db.rs b/services/api/src/db.rs new file mode 100644 index 0000000..ab17a8d --- /dev/null +++ b/services/api/src/db.rs @@ -0,0 +1,228 @@ +//! The sqlite file behind everything server-side: OAuth state, OAuth +//! sessions (which hold tokens — this file is a credential store first and a +//! database second), and the account rows `/api/session` answers from. +//! +//! rusqlite is synchronous. Every operation runs the closure on a blocking +//! thread over one mutex-serialized connection: the futures come out `Send` +//! (the atrium store traits require it), no runtime worker parks on file +//! I/O, and a single process writing a WAL database needs nothing cleverer. + +use std::path::Path; +use std::sync::{Arc, Mutex}; + +use rusqlite::Connection; + +#[derive(Debug)] +pub enum DbError { + Sqlite(rusqlite::Error), + /// The blocking task was cancelled — only happens at shutdown. + Cancelled, +} + +impl std::fmt::Display for DbError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + DbError::Sqlite(e) => write!(f, "sqlite: {e}"), + DbError::Cancelled => write!(f, "database task cancelled"), + } + } +} + +impl std::error::Error for DbError {} + +#[derive(Clone)] +pub struct Db(Arc>); + +impl Db { + pub fn open(path: &Path) -> Result { + if let Some(parent) = path.parent() + && !parent.as_os_str().is_empty() + { + std::fs::create_dir_all(parent) + .map_err(|e| format!("cannot create {}: {e}", parent.display()))?; + } + + let conn = + Connection::open(path).map_err(|e| format!("cannot open {}: {e}", path.display()))?; + + // Tokens at rest: nobody but the service user gets to read them. + let mut permissions = std::fs::metadata(path) + .map_err(|e| format!("cannot stat {}: {e}", path.display()))? + .permissions(); + std::os::unix::fs::PermissionsExt::set_mode(&mut permissions, 0o600); + std::fs::set_permissions(path, permissions) + .map_err(|e| format!("cannot chmod {}: {e}", path.display()))?; + + conn.execute_batch( + "PRAGMA journal_mode = WAL; + PRAGMA busy_timeout = 5000; + PRAGMA foreign_keys = ON; + + CREATE TABLE IF NOT EXISTS oauth_state ( + key TEXT PRIMARY KEY, + value TEXT NOT NULL, + created_at INTEGER NOT NULL DEFAULT (unixepoch()) + ); + CREATE TABLE IF NOT EXISTS oauth_session ( + did TEXT PRIMARY KEY, + value TEXT NOT NULL, + created_at INTEGER NOT NULL DEFAULT (unixepoch()), + updated_at INTEGER NOT NULL DEFAULT (unixepoch()) + ); + CREATE TABLE IF NOT EXISTS account ( + did TEXT PRIMARY KEY, + handle TEXT, + verified_at INTEGER + );", + ) + .map_err(|e| format!("cannot migrate {}: {e}", path.display()))?; + + Ok(Db(Arc::new(Mutex::new(conn)))) + } + + /// Runs `f` against the connection on a blocking thread. + pub async fn call(&self, f: F) -> Result + where + T: Send + 'static, + F: FnOnce(&Connection) -> Result + Send + 'static, + { + let db = self.0.clone(); + tokio::task::spawn_blocking(move || { + let conn = db.lock().expect("db mutex poisoned"); + f(&conn).map_err(DbError::Sqlite) + }) + .await + .map_err(|_| DbError::Cancelled)? + } + + /// Records who signed in. `handle` is the *verified* handle — None means + /// verification failed and the UI shows the DID as unverified. + pub async fn upsert_account(&self, did: &str, handle: Option<&str>) -> Result<(), DbError> { + let did = did.to_owned(); + let handle = handle.map(str::to_owned); + self.call(move |conn| { + conn.execute( + "INSERT INTO account (did, handle, verified_at) + VALUES (?1, ?2, CASE WHEN ?2 IS NULL THEN NULL ELSE unixepoch() END) + ON CONFLICT(did) DO UPDATE SET + handle = excluded.handle, + verified_at = excluded.verified_at", + (did, handle), + ) + .map(|_| ()) + }) + .await + } + + /// `None` — no such account; `Some(None)` — account with no verified + /// handle. + pub async fn account_handle(&self, did: &str) -> Result>, DbError> { + let did = did.to_owned(); + self.call(move |conn| { + use rusqlite::OptionalExtension; + conn.query_row("SELECT handle FROM account WHERE did = ?1", [did], |row| { + row.get::<_, Option>(0) + }) + .optional() + }) + .await + } + + /// Whether a stored OAuth session still exists for the DID. The cookie + /// alone is not a session: this row going away is how sign-out and a + /// wiped database actually sign someone out. + pub async fn has_oauth_session(&self, did: &str) -> Result { + let did = did.to_owned(); + self.call(move |conn| { + use rusqlite::OptionalExtension; + conn.query_row("SELECT 1 FROM oauth_session WHERE did = ?1", [did], |_| { + Ok(()) + }) + .optional() + .map(|row| row.is_some()) + }) + .await + } + + pub async fn delete_oauth_session(&self, did: &str) -> Result<(), DbError> { + let did = did.to_owned(); + self.call(move |conn| { + conn.execute("DELETE FROM oauth_session WHERE did = ?1", [did]) + .map(|_| ()) + }) + .await + } +} + +#[cfg(test)] +mod tests { + use super::*; + + async fn temp_db() -> (tempfile::TempDir, Db) { + let dir = tempfile::tempdir().unwrap(); + let db = Db::open(&dir.path().join("test.sqlite")).unwrap(); + (dir, db) + } + + #[tokio::test] + async fn account_round_trip() { + let (_dir, db) = temp_db().await; + + assert_eq!(db.account_handle("did:plc:abc").await.unwrap(), None); + + db.upsert_account("did:plc:abc", Some("someone.example")) + .await + .unwrap(); + assert_eq!( + db.account_handle("did:plc:abc").await.unwrap(), + Some(Some("someone.example".to_owned())) + ); + + // A later sign-in that fails verification demotes the handle. + db.upsert_account("did:plc:abc", None).await.unwrap(); + assert_eq!(db.account_handle("did:plc:abc").await.unwrap(), Some(None)); + } + + #[tokio::test] + async fn oauth_session_presence() { + let (_dir, db) = temp_db().await; + + assert!(!db.has_oauth_session("did:plc:abc").await.unwrap()); + db.call(|conn| { + conn.execute( + "INSERT INTO oauth_session (did, value) VALUES ('did:plc:abc', '{}')", + [], + ) + .map(|_| ()) + }) + .await + .unwrap(); + assert!(db.has_oauth_session("did:plc:abc").await.unwrap()); + + db.delete_oauth_session("did:plc:abc").await.unwrap(); + assert!(!db.has_oauth_session("did:plc:abc").await.unwrap()); + } + + #[tokio::test] + async fn file_is_private_and_survives_reopen() { + let dir = tempfile::tempdir().unwrap(); + let path = dir.path().join("test.sqlite"); + + let db = Db::open(&path).unwrap(); + db.upsert_account("did:plc:abc", Some("someone.example")) + .await + .unwrap(); + drop(db); + + let mode = std::os::unix::fs::PermissionsExt::mode( + &std::fs::metadata(&path).unwrap().permissions(), + ); + assert_eq!(mode & 0o777, 0o600); + + let db = Db::open(&path).unwrap(); + assert_eq!( + db.account_handle("did:plc:abc").await.unwrap(), + Some(Some("someone.example".to_owned())) + ); + } +} diff --git a/services/api/src/main.rs b/services/api/src/main.rs index 1ac692f..7cd9807 100644 --- a/services/api/src/main.rs +++ b/services/api/src/main.rs @@ -1,12 +1,13 @@ +mod atproto; mod config; +mod db; +mod routes; +mod session; -use axum::{Router, routing::get}; +use std::sync::Arc; use crate::config::Config; - -fn app() -> Router { - Router::new().route("/healthz", get(|| async { "ok" })) -} +use crate::routes::AppState; #[tokio::main] async fn main() { @@ -27,12 +28,47 @@ async fn main() { } }; + let db = match db::Db::open(&config.db_path) { + Ok(db) => db, + Err(e) => { + tracing::error!("database: {e}"); + std::process::exit(1); + } + }; + + let atproto = match atproto::Atproto::new(&config, db.clone()) { + Ok(atproto) => atproto, + Err(e) => { + tracing::error!("atproto: {e}"); + std::process::exit(1); + } + }; + + tracing::info!( + "oauth client: {}", + if atproto.is_confidential() { + "confidential" + } else { + "loopback (development)" + } + ); + + let state = AppState { + atproto: Arc::new(atproto), + db, + signer: Arc::new(session::SessionSigner::new( + config.session_secret.clone(), + config.public_url.is_some(), + )), + web_origin: config.web_origin.clone(), + }; + let listener = tokio::net::TcpListener::bind(config.bind_addr) .await .unwrap_or_else(|e| panic!("cannot bind {}: {e}", config.bind_addr)); tracing::info!("listening on {}", config.bind_addr); - axum::serve(listener, app()) + axum::serve(listener, routes::app(state)) .with_graceful_shutdown(shutdown_signal()) .await .expect("server error"); @@ -49,20 +85,3 @@ async fn shutdown_signal() { } tracing::info!("shutting down"); } - -#[cfg(test)] -mod tests { - use super::*; - use axum::body::Body; - use axum::http::{Request, StatusCode}; - use tower::ServiceExt; - - #[tokio::test] - async fn healthz_responds_ok() { - let response = app() - .oneshot(Request::get("/healthz").body(Body::empty()).unwrap()) - .await - .unwrap(); - assert_eq!(response.status(), StatusCode::OK); - } -} diff --git a/services/api/src/routes.rs b/services/api/src/routes.rs new file mode 100644 index 0000000..c8d4fae --- /dev/null +++ b/services/api/src/routes.rs @@ -0,0 +1,430 @@ +//! The HTTP surface. The /api/* contract is fixed by web/src/api.ts: +//! errors are `{message}`, login success is `{redirectUrl}`, session is +//! `{did, handle}` — change either side and change the other. + +use std::sync::Arc; + +use axum::extract::{Query, State}; +use axum::http::{HeaderValue, Method, StatusCode, header}; +use axum::response::{IntoResponse, Redirect, Response}; +use axum::routing::{get, post}; +use axum::{Json, Router}; +use axum_extra::extract::cookie::CookieJar; +use tower_http::cors::CorsLayer; + +use crate::atproto::{Atproto, LoginError}; +use crate::db::Db; +use crate::session::SessionSigner; + +#[derive(Clone)] +pub struct AppState { + pub atproto: Arc, + pub db: Db, + pub signer: Arc, + pub web_origin: String, +} + +pub fn app(state: AppState) -> Router { + // Exact origin plus credentials: the allow-list is the site and nothing + // else, and a JSON POST from it forces a preflight — which is currently + // the whole CSRF defence (TODO.md). + let cors = CorsLayer::new() + .allow_origin( + state + .web_origin + .parse::() + .expect("config validated origin"), + ) + .allow_methods([Method::GET, Method::POST]) + .allow_headers([header::CONTENT_TYPE]) + .allow_credentials(true); + + let mut router = Router::new() + .route("/healthz", get(|| async { "ok" })) + .route("/api/login", post(login)) + .route("/api/session", get(session)) + .route("/api/logout", post(logout)) + .route("/oauth/callback", get(callback)); + + if state.atproto.is_confidential() { + router = router + .route("/oauth/client-metadata.json", get(client_metadata)) + .route("/.well-known/jwks.json", get(jwks)); + } + + router.layer(cors).with_state(state) +} + +fn message(status: StatusCode, text: &str) -> Response { + (status, Json(serde_json::json!({ "message": text }))).into_response() +} + +#[derive(serde::Deserialize)] +struct LoginBody { + handle: String, +} + +async fn login(State(state): State, body: Option>) -> Response { + let Some(Json(LoginBody { handle })) = body else { + return message(StatusCode::BAD_REQUEST, "Send a JSON body with a handle."); + }; + match state.atproto.begin_login(&handle).await { + Ok(url) => Json(serde_json::json!({ "redirectUrl": url })).into_response(), + Err(LoginError::InvalidHandle) => { + message(StatusCode::BAD_REQUEST, "That doesn't look like a handle.") + } + Err(LoginError::Upstream) => message( + StatusCode::BAD_GATEWAY, + "Your identity provider could not be reached.", + ), + } +} + +/// Permissive on purpose: the authorization server may come back with +/// `error=` and no code (the user said no), and atrium's own params type +/// requires `code`. +#[derive(serde::Deserialize)] +struct CallbackQuery { + code: Option, + state: Option, + iss: Option, + error: Option, +} + +async fn callback(State(state): State, Query(query): Query) -> Response { + let denied = format!("{}/?error=denied", state.web_origin); + let failed = format!("{}/?error=login_failed", state.web_origin); + + let Some(code) = query.code else { + // Not worth logging the error value: it is attacker-controlled query + // text on an unauthenticated route. + return Redirect::to(if query.error.is_some() { + &denied + } else { + &failed + }) + .into_response(); + }; + + match state + .atproto + .complete_login(code, query.state, query.iss) + .await + { + Ok((did, handle)) => { + if state + .db + .upsert_account(&did, handle.as_deref()) + .await + .is_err() + { + tracing::error!("callback: account row write failed"); + return Redirect::to(&failed).into_response(); + } + let jar = CookieJar::new().add(state.signer.issue(&did)); + // SameSite=Lax lets the cookie ride this top-level cross-site + // redirect — the reason Lax was chosen over Strict. + (jar, Redirect::to(&format!("{}/", state.web_origin))).into_response() + } + Err(crate::atproto::CallbackFailed) => Redirect::to(&failed).into_response(), + } +} + +async fn session(State(state): State, jar: CookieJar) -> Response { + let did = jar + .get(crate::session::COOKIE_NAME) + .and_then(|cookie| state.signer.verify(cookie.value())); + let Some(did) = did else { + return StatusCode::UNAUTHORIZED.into_response(); + }; + + // The cookie alone is not a session: the stored OAuth session is. Gone — + // signed out elsewhere, database wiped — means signed out, and the stale + // cookie gets cleared on the way past. + match state.db.has_oauth_session(&did).await { + Ok(true) => {} + Ok(false) => { + let jar = CookieJar::new().add(state.signer.clear()); + return (StatusCode::UNAUTHORIZED, jar).into_response(); + } + Err(_) => return StatusCode::INTERNAL_SERVER_ERROR.into_response(), + } + + match state.db.account_handle(&did).await { + Ok(handle) => { + Json(serde_json::json!({ "did": did, "handle": handle.flatten() })).into_response() + } + Err(_) => StatusCode::INTERNAL_SERVER_ERROR.into_response(), + } +} + +async fn logout(State(state): State, jar: CookieJar) -> Response { + if let Some(did) = jar + .get(crate::session::COOKIE_NAME) + .and_then(|cookie| state.signer.verify(cookie.value())) + { + // Best-effort revocation at the server, unconditional deletion here: + // sign-out must work even when the PDS is down. + state.atproto.logout(&did).await; + if state.db.delete_oauth_session(&did).await.is_err() { + tracing::error!("logout: session row delete failed"); + } + } + let jar = CookieJar::new().add(state.signer.clear()); + (StatusCode::NO_CONTENT, jar).into_response() +} + +async fn client_metadata(State(state): State) -> Response { + Json(state.atproto.client_metadata()).into_response() +} + +async fn jwks(State(state): State) -> Response { + Json(state.atproto.jwks()).into_response() +} + +#[cfg(test)] +mod tests { + use super::*; + use axum::body::Body; + use axum::http::Request; + use http_body_util::BodyExt; + use tower::ServiceExt; + + const SECRET: &[u8] = b"0123456789abcdef0123456789abcdef"; + + async fn loopback_state() -> (tempfile::TempDir, AppState) { + let dir = tempfile::tempdir().unwrap(); + let db = crate::db::Db::open(&dir.path().join("test.sqlite")).unwrap(); + let config = crate::config::Config::from_lookup(|_| None).unwrap(); + let state = AppState { + atproto: Arc::new(Atproto::new(&config, db.clone()).unwrap()), + db, + signer: Arc::new(SessionSigner::new(SECRET.to_vec(), false)), + web_origin: config.web_origin, + }; + (dir, state) + } + + async fn body_json(response: Response) -> serde_json::Value { + let bytes = response.into_body().collect().await.unwrap().to_bytes(); + serde_json::from_slice(&bytes).unwrap() + } + + #[tokio::test] + async fn healthz_still_responds() { + let (_dir, state) = loopback_state().await; + let response = app(state) + .oneshot(Request::get("/healthz").body(Body::empty()).unwrap()) + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::OK); + } + + #[tokio::test] + async fn session_without_cookie_is_401() { + let (_dir, state) = loopback_state().await; + let response = app(state) + .oneshot(Request::get("/api/session").body(Body::empty()).unwrap()) + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::UNAUTHORIZED); + } + + #[tokio::test] + async fn session_with_forged_cookie_is_401() { + let (_dir, state) = loopback_state().await; + let forged = SessionSigner::new(b"wrong-secret-wrong-secret-wrong!".to_vec(), false) + .issue("did:plc:abc"); + let response = app(state) + .oneshot( + Request::get("/api/session") + .header( + header::COOKIE, + format!("{}={}", crate::session::COOKIE_NAME, forged.value()), + ) + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::UNAUTHORIZED); + } + + #[tokio::test] + async fn session_happy_path() { + let (_dir, state) = loopback_state().await; + state + .db + .upsert_account("did:plc:abc", Some("someone.example")) + .await + .unwrap(); + state + .db + .call(|conn| { + conn.execute( + "INSERT INTO oauth_session (did, value) VALUES ('did:plc:abc', '{}')", + [], + ) + .map(|_| ()) + }) + .await + .unwrap(); + let cookie = state.signer.issue("did:plc:abc"); + + let response = app(state) + .oneshot( + Request::get("/api/session") + .header( + header::COOKIE, + format!("{}={}", crate::session::COOKIE_NAME, cookie.value()), + ) + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::OK); + assert_eq!( + body_json(response).await, + serde_json::json!({ "did": "did:plc:abc", "handle": "someone.example" }) + ); + } + + #[tokio::test] + async fn session_with_cookie_but_no_stored_session_clears_it() { + let (_dir, state) = loopback_state().await; + let cookie = state.signer.issue("did:plc:abc"); + + let response = app(state) + .oneshot( + Request::get("/api/session") + .header( + header::COOKIE, + format!("{}={}", crate::session::COOKIE_NAME, cookie.value()), + ) + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::UNAUTHORIZED); + let set_cookie = response + .headers() + .get(header::SET_COOKIE) + .unwrap() + .to_str() + .unwrap(); + assert!( + set_cookie.contains("Max-Age=0"), + "stale cookie must be cleared: {set_cookie}" + ); + } + + #[tokio::test] + async fn login_wants_a_handle() { + let (_dir, state) = loopback_state().await; + let app = app(state); + + let response = app + .clone() + .oneshot( + Request::post("/api/login") + .header(header::CONTENT_TYPE, "application/json") + .body(Body::from(r#"{"handle": ""}"#)) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::BAD_REQUEST); + assert!(body_json(response).await["message"].is_string()); + + let response = app + .oneshot(Request::post("/api/login").body(Body::empty()).unwrap()) + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::BAD_REQUEST); + } + + #[tokio::test] + async fn logout_always_clears() { + let (_dir, state) = loopback_state().await; + let response = app(state) + .oneshot(Request::post("/api/logout").body(Body::empty()).unwrap()) + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::NO_CONTENT); + let set_cookie = response + .headers() + .get(header::SET_COOKIE) + .unwrap() + .to_str() + .unwrap(); + assert!(set_cookie.contains("Max-Age=0")); + } + + #[tokio::test] + async fn callback_without_code_redirects_back() { + let (_dir, state) = loopback_state().await; + let web_origin = state.web_origin.clone(); + let response = app(state) + .oneshot( + Request::get("/oauth/callback?error=access_denied") + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::SEE_OTHER); + assert_eq!( + response.headers().get(header::LOCATION).unwrap(), + &format!("{web_origin}/?error=denied") + ); + } + + #[tokio::test] + async fn metadata_routes_absent_in_loopback_mode() { + let (_dir, state) = loopback_state().await; + let app = app(state); + for path in ["/oauth/client-metadata.json", "/.well-known/jwks.json"] { + let response = app + .clone() + .oneshot(Request::get(path).body(Body::empty()).unwrap()) + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::NOT_FOUND, "{path}"); + } + } + + #[tokio::test] + async fn cors_preflight_echoes_the_one_origin() { + let (_dir, state) = loopback_state().await; + let web_origin = state.web_origin.clone(); + let response = app(state) + .oneshot( + Request::builder() + .method(Method::OPTIONS) + .uri("/api/login") + .header(header::ORIGIN, &web_origin) + .header(header::ACCESS_CONTROL_REQUEST_METHOD, "POST") + .header(header::ACCESS_CONTROL_REQUEST_HEADERS, "content-type") + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!( + response + .headers() + .get(header::ACCESS_CONTROL_ALLOW_ORIGIN) + .unwrap(), + &web_origin + ); + assert_eq!( + response + .headers() + .get(header::ACCESS_CONTROL_ALLOW_CREDENTIALS) + .unwrap(), + "true" + ); + } +} diff --git a/services/api/src/session.rs b/services/api/src/session.rs new file mode 100644 index 0000000..c1911a5 --- /dev/null +++ b/services/api/src/session.rs @@ -0,0 +1,183 @@ +//! The `headquarters_session` cookie: HMAC-signed `{did, exp}`. +//! +//! HttpOnly keeps page JavaScript from *reading* a cookie; it says nothing +//! about who can *write* one, and the DID in this cookie is the entire +//! authorization decision — unsigned, it is a text field where you type who +//! you are (docs/identity-and-sessions.md). + +use axum_extra::extract::cookie::{Cookie, SameSite}; +use base64::Engine; +use base64::engine::general_purpose::URL_SAFE_NO_PAD; +use hmac::{Hmac, Mac}; +use sha2::Sha256; + +pub const COOKIE_NAME: &str = "headquarters_session"; +const SESSION_TTL_SECONDS: i64 = 30 * 24 * 60 * 60; + +#[derive(serde::Serialize, serde::Deserialize)] +struct Claims { + did: String, + exp: i64, +} + +pub struct SessionSigner { + secret: Vec, + secure: bool, +} + +impl SessionSigner { + pub fn new(secret: Vec, secure: bool) -> SessionSigner { + SessionSigner { secret, secure } + } + + fn mac(&self, payload: &[u8]) -> Hmac { + let mut mac = + Hmac::::new_from_slice(&self.secret).expect("hmac accepts any key length"); + mac.update(payload); + mac + } + + // SameSite=Lax, not Strict: the OAuth callback is a top-level navigation + // from a different site, and Strict would drop the cookie on the redirect + // immediately after it is set. No Domain attribute: host-only, so the + // cookie goes to the API and to nothing else under the domain. + fn attributes(&self, cookie: Cookie<'static>) -> Cookie<'static> { + let mut cookie = cookie; + cookie.set_http_only(true); + cookie.set_same_site(SameSite::Lax); + cookie.set_path("/"); + cookie.set_secure(self.secure); + cookie + } + + pub fn issue(&self, did: &str) -> Cookie<'static> { + let now = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .expect("clock before 1970") + .as_secs() as i64; + let payload = serde_json::to_vec(&Claims { + did: did.to_owned(), + exp: now + SESSION_TTL_SECONDS, + }) + .expect("claims serialize"); + let tag = self.mac(&payload).finalize().into_bytes(); + let value = format!( + "{}.{}", + URL_SAFE_NO_PAD.encode(&payload), + URL_SAFE_NO_PAD.encode(tag) + ); + + let mut cookie = self.attributes(Cookie::new(COOKIE_NAME, value)); + cookie.set_max_age(cookie::time::Duration::seconds(SESSION_TTL_SECONDS)); + cookie + } + + pub fn clear(&self) -> Cookie<'static> { + let mut cookie = self.attributes(Cookie::new(COOKIE_NAME, "")); + cookie.set_max_age(cookie::time::Duration::ZERO); + cookie + } + + /// The DID, iff the tag verifies and the session has not expired. + pub fn verify(&self, value: &str) -> Option { + let (payload_b64, tag_b64) = value.split_once('.')?; + let payload = URL_SAFE_NO_PAD.decode(payload_b64).ok()?; + let tag = URL_SAFE_NO_PAD.decode(tag_b64).ok()?; + + // Constant-time comparison; a straight == would leak the tag byte by + // byte. + self.mac(&payload).verify_slice(&tag).ok()?; + + let claims: Claims = serde_json::from_slice(&payload).ok()?; + let now = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .expect("clock before 1970") + .as_secs() as i64; + (claims.exp > now).then_some(claims.did) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn signer() -> SessionSigner { + SessionSigner::new(b"0123456789abcdef0123456789abcdef".to_vec(), false) + } + + #[test] + fn round_trip() { + let s = signer(); + let cookie = s.issue("did:plc:abc"); + assert_eq!(s.verify(cookie.value()), Some("did:plc:abc".to_owned())); + } + + #[test] + fn tampering_fails() { + let s = signer(); + let value = s.issue("did:plc:abc").value().to_owned(); + + let (payload, tag) = value.split_once('.').unwrap(); + let other = URL_SAFE_NO_PAD.encode( + serde_json::to_vec(&Claims { + did: "did:plc:evil".into(), + exp: i64::MAX, + }) + .unwrap(), + ); + assert_eq!(s.verify(&format!("{other}.{tag}")), None, "forged payload"); + + let mut broken_tag = tag.to_owned(); + broken_tag.replace_range(0..1, if &tag[0..1] == "A" { "B" } else { "A" }); + assert_eq!( + s.verify(&format!("{payload}.{broken_tag}")), + None, + "broken tag" + ); + + assert_eq!(s.verify(payload), None, "no tag at all"); + assert_eq!(s.verify(""), None, "empty"); + } + + #[test] + fn expiry_is_enforced() { + let s = signer(); + let payload = serde_json::to_vec(&Claims { + did: "did:plc:abc".into(), + exp: 0, + }) + .unwrap(); + let tag = s.mac(&payload).finalize().into_bytes(); + let value = format!( + "{}.{}", + URL_SAFE_NO_PAD.encode(&payload), + URL_SAFE_NO_PAD.encode(tag) + ); + assert_eq!(s.verify(&value), None); + } + + #[test] + fn different_secrets_reject_each_other() { + let a = signer(); + let b = SessionSigner::new(b"another-secret-another-secret-32".to_vec(), false); + let cookie = a.issue("did:plc:abc"); + assert_eq!(b.verify(cookie.value()), None); + } + + #[test] + fn secure_follows_mode() { + let dev = signer().issue("did:plc:abc"); + assert_ne!(dev.secure(), Some(true)); + + let public = SessionSigner::new(b"0123456789abcdef0123456789abcdef".to_vec(), true) + .issue("did:plc:abc"); + assert_eq!(public.secure(), Some(true)); + } + + #[test] + fn clear_expires_immediately() { + let cookie = signer().clear(); + assert_eq!(cookie.value(), ""); + assert_eq!(cookie.max_age(), Some(cookie::time::Duration::ZERO)); + } +} -- 2.51.2