diff --git a/Cargo.lock b/Cargo.lock index d5572cf..6bc9ff7 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2751,6 +2751,7 @@ name = "starhaven" version = "0.1.0" dependencies = [ "aes-gcm", + "anyhow", "atproto-identity", "atproto-oauth", "axum", @@ -2759,11 +2760,13 @@ dependencies = [ "cookie", "maud", "rand 0.9.5", + "reqwest", "serde", "serde_json", "sqlx", "thiserror 2.0.19", "tokio", + "urlencoding", ] [[package]] @@ -3188,6 +3191,12 @@ dependencies = [ "serde", ] +[[package]] +name = "urlencoding" +version = "2.1.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "daf8dba3b7eb870caf1ddeed7bc9d2a049f3cfdfae7cb521b087cc33ae4c49da" + [[package]] name = "utf8_iter" version = "1.0.4" diff --git a/Cargo.toml b/Cargo.toml index d9780b6..2b91856 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -21,4 +21,7 @@ rand = "0.9" cookie = "0.18" base64 = "0.22" chrono = { version = "0.4", features = ["serde"] } -thiserror = "2" \ No newline at end of file +thiserror = "2" +anyhow = "1" +reqwest = { version = "0.12", default-features = false, features = ["json", "rustls-tls"] } +urlencoding = "2.1" \ No newline at end of file diff --git a/src/error.rs b/src/error.rs new file mode 100644 index 0000000..9440131 --- /dev/null +++ b/src/error.rs @@ -0,0 +1,32 @@ +//! The HTTP-facing error type returned by fallible handlers. + +use axum::http::StatusCode; +use axum::response::{IntoResponse, Response}; +use thiserror::Error; + +/// Error returned by request handlers. +#[derive(Debug, Error)] +pub enum AppError { + /// An internal server error wrapping any anyhow error. + #[error("internal server error: {0}")] + Internal(#[from] anyhow::Error), + + /// The request was malformed. + #[error("bad request: {0}")] + BadRequest(String), + + /// The request was unauthenticated or the session was invalid. + #[error("unauthorized")] + Unauthorized, +} + +impl IntoResponse for AppError { + fn into_response(self) -> Response { + let status = match &self { + AppError::Internal(_) => StatusCode::INTERNAL_SERVER_ERROR, + AppError::BadRequest(_) => StatusCode::BAD_REQUEST, + AppError::Unauthorized => StatusCode::UNAUTHORIZED, + }; + (status, self.to_string()).into_response() + } +} diff --git a/src/main.rs b/src/main.rs index 5a0b6ef..d07890e 100644 --- a/src/main.rs +++ b/src/main.rs @@ -1,4 +1,5 @@ mod config; +mod error; mod oauth; mod state; @@ -29,7 +30,7 @@ async fn main() { axum::serve(listener, app).await.unwrap(); } -fn layout(body: Markup) -> Html { +pub(crate) fn layout(body: Markup) -> Html { Html( html! { (DOCTYPE) diff --git a/src/oauth/login.rs b/src/oauth/login.rs new file mode 100644 index 0000000..ff331d8 --- /dev/null +++ b/src/oauth/login.rs @@ -0,0 +1,193 @@ +//! OAuth login init: resolve the user's PDS authorization server, generate +//! PKCE + state + a per-session DPoP key, persist request state in memory, +//! run PAR, and redirect to the authorization endpoint. + +use atproto_identity::key::{generate_key, KeyType}; +use atproto_oauth::pkce; +use atproto_oauth::resources::{pds_resources, AuthorizationServer}; +use atproto_oauth::workflow::{oauth_init, OAuthClient, OAuthRequestState}; +use axum::extract::State; +use axum::response::{Html, IntoResponse, Redirect, Response}; +use axum::{routing::get, Form, Router}; +use chrono::{DateTime, Duration, Utc}; +use maud::html; +use rand::distr::{Alphanumeric, SampleString}; +use serde::Deserialize; + +use crate::error::AppError; +use crate::layout; +use crate::state::AppState; + +/// In-flight OAuth request state, persisted in memory between `login` (PAR) +/// and `callback` (code exchange). Single-use, keyed by `state`. +/// +/// Our own type, not `atproto_oauth::workflow::OAuthRequest` -- that one has +/// no `subject` field, and `callback` needs it for subject-binding. +#[derive(Debug, Clone)] +pub struct PersistedOAuthRequest { + /// The CSRF state value (map key). + pub state: String, + /// The authorization server issuer. + pub issuer: String, + /// The OAuth nonce. + pub nonce: String, + /// The PKCE code verifier. + pub pkce_verifier: String, + /// The per-session DPoP private key (`did:key:...`). + pub dpop_private_key: String, + /// The original login hint (handle/DID/URL) the user supplied. + pub login_hint: String, + /// The DID resolved at PAR time, when the login hint resolved to one. + pub subject: Option, + /// When the request was created. + pub created_at: DateTime, + /// When the request expires. + pub expires_at: DateTime, +} + +/// Form body for `POST /login`. +#[derive(Debug, Deserialize)] +pub struct LoginForm { + /// A handle, DID, or PDS URL to authenticate against. + pub handle: String, +} + +/// `GET /login` -- a minimal handle-entry form. +pub async fn login_form() -> Html { + layout(html! { + h1 { "log in" } + form method="post" action="/login" { + label for="handle" { "handle, DID, or PDS URL" } + input type="text" id="handle" name="handle" placeholder="alice.bsky.social" required; + button type="submit" { "continue" } + } + }) +} + +/// `POST /login` -- begin the OAuth authorization-code + PKCE + PAR flow. +pub async fn login( + State(state): State, + Form(form): Form, +) -> Result { + let login_hint = form.handle.trim(); + if login_hint.is_empty() { + return Err(AppError::BadRequest("handle is required".to_string())); + } + + let (authorization_server, resolved_subject) = resolve_login_hint(&state, login_hint).await?; + + let signing_key = state + .config + .oauth_private_keys + .first() + .cloned() + .ok_or_else(|| AppError::Internal(anyhow::anyhow!("no OAuth signing key configured")))?; + + let dpop_key = generate_key(KeyType::P256Private) + .map_err(|e| AppError::Internal(anyhow::anyhow!("failed to generate DPoP key: {e}")))?; + + let (pkce_verifier, code_challenge) = pkce::generate(); + let csrf_state = Alphanumeric.sample_string(&mut rand::rng(), 32); + let nonce = Alphanumeric.sample_string(&mut rand::rng(), 32); + + let oauth_client = OAuthClient { + redirect_uri: state.config.oauth_redirect_uri(), + client_id: state.config.oauth_client_id(), + private_signing_key_data: signing_key, + }; + + let oauth_request_state = OAuthRequestState { + state: csrf_state.clone(), + nonce: nonce.clone(), + code_challenge, + scope: oauth_scope(), + }; + + let par_login_hint = resolved_subject.as_deref().or(Some(login_hint)); + + let par_response = oauth_init( + &state.http_client, + &oauth_client, + &dpop_key, + par_login_hint, + &authorization_server, + &oauth_request_state, + ) + .await + .map_err(|e| AppError::Internal(anyhow::anyhow!("PAR request failed: {e}")))?; + + let now = Utc::now(); + let persisted = PersistedOAuthRequest { + state: csrf_state.clone(), + issuer: authorization_server.issuer.clone(), + nonce, + pkce_verifier, + dpop_private_key: dpop_key.to_string(), + login_hint: login_hint.to_string(), + subject: resolved_subject, + created_at: now, + expires_at: now + Duration::minutes(10), + }; + + state + .oauth_requests + .lock() + .expect("oauth_requests lock poisoned") + .insert(csrf_state, persisted); + + let authorize_url = format!( + "{}?client_id={}&request_uri={}", + authorization_server.authorization_endpoint, + urlencoding::encode(&oauth_client.client_id), + urlencoding::encode(&par_response.request_uri), + ); + + Ok(Redirect::to(&authorize_url).into_response()) +} + +/// The full OAuth scope string requested by this client. +fn oauth_scope() -> String { + "atproto transition:generic".to_string() +} + +/// Resolve a login hint (handle/DID/PDS URL) into an authorization server. +/// +/// Returns the authorization server metadata and the resolved DID, when the +/// login hint resolved to one (a bare PDS URL login resolves to `None`). +pub async fn resolve_login_hint( + state: &AppState, + login_hint: &str, +) -> Result<(AuthorizationServer, Option), AppError> { + let subject = login_hint + .trim() + .trim_start_matches("at://") + .trim_start_matches('@'); + + let (pds_endpoint, resolved_did) = if subject.starts_with("https://") { + (subject.to_string(), None) + } else { + let document = state + .identity_resolver + .resolve(subject) + .await + .map_err(|e| AppError::BadRequest(format!("could not resolve '{subject}': {e}")))?; + let pds = document + .pds_endpoints() + .first() + .map(|s| s.to_string()) + .ok_or_else(|| AppError::BadRequest(format!("no PDS endpoint for '{subject}'")))?; + (pds, Some(document.id.clone())) + }; + + let (_protected_resource, authorization_server) = + pds_resources(&state.http_client, &pds_endpoint) + .await + .map_err(|e| AppError::BadRequest(format!("OAuth resource discovery failed: {e}")))?; + + Ok((authorization_server, resolved_did)) +} + +/// Routes for the login-init step. +pub fn router() -> Router { + Router::new().route("/login", get(login_form).post(login)) +} diff --git a/src/oauth/mod.rs b/src/oauth/mod.rs index 5a5a029..c9dbdae 100644 --- a/src/oauth/mod.rs +++ b/src/oauth/mod.rs @@ -1,6 +1,7 @@ //! OAuth confidential-client flow: client metadata + JWKS discovery, login, //! callback, session refresh and logout. Owns every OAuth-related route. +pub mod login; pub mod metadata; pub mod session; @@ -14,4 +15,5 @@ pub fn router() -> Router { Router::new() .route("/client-metadata.json", get(metadata::client_metadata)) .route("/jwks.json", get(metadata::jwks)) + .merge(login::router()) } diff --git a/src/state.rs b/src/state.rs index 92138cd..1d488d6 100644 --- a/src/state.rs +++ b/src/state.rs @@ -1,9 +1,15 @@ //! Shared application state. +use std::collections::HashMap; use std::ops::Deref; -use std::sync::Arc; +use std::sync::{Arc, Mutex}; + +use atproto_identity::resolve::{ + HickoryDnsResolver, InnerIdentityResolver, SharedIdentityResolver, +}; use crate::config::Config; +use crate::oauth::login::PersistedOAuthRequest; /// Shared, cheaply-cloneable application state passed to axum handlers. #[derive(Clone)] @@ -13,12 +19,34 @@ pub struct AppState(pub Arc); pub struct Inner { /// Application configuration. pub config: Config, + /// Shared HTTP client for calling out to PDSes/authorization servers. + pub http_client: reqwest::Client, + /// DID/handle resolver. + pub identity_resolver: SharedIdentityResolver, + /// In-flight OAuth requests (PKCE verifier, CSRF state, per-flow DPoP + /// key), keyed by the CSRF `state` value, between `/login` and + /// `/callback`. In-memory and lost across restarts; an interrupted login + /// just needs to be retried. + pub oauth_requests: Mutex>, } impl AppState { /// Build application state from configuration. pub fn new(config: Config) -> Self { - AppState(Arc::new(Inner { config })) + let http_client = reqwest::Client::new(); + let dns_resolver = Arc::new(HickoryDnsResolver::create_resolver(&[])); + let identity_resolver = SharedIdentityResolver(Arc::new(InnerIdentityResolver { + dns_resolver, + http_client: http_client.clone(), + plc_hostname: "plc.directory".to_string(), + })); + + AppState(Arc::new(Inner { + config, + http_client, + identity_resolver, + oauth_requests: Mutex::new(HashMap::new()), + })) } }