From 3eae0e36183553e1eb733ac0a5e8b0640976b303 Mon Sep 17 00:00:00 2001 From: Orual Date: Sat, 6 Jun 2026 20:43:25 -0400 Subject: [PATCH] axum oauth helpers and example --- AGENTS.md | 294 +++++++- CLAUDE.md | 294 +------- Cargo.lock | 148 ++++ crates/jacquard-axum/Cargo.toml | 10 +- crates/jacquard-axum/src/lib.rs | 1 + crates/jacquard-axum/src/oauth.rs | 661 +++++++++++++++++ crates/jacquard-axum/tests/oauth_web_tests.rs | 696 ++++++++++++++++++ crates/jacquard-oauth/src/authstore.rs | 1 - crates/jacquard-oauth/src/client.rs | 1 - crates/jacquard-oauth/src/keyset.rs | 35 +- crates/jacquard-oauth/src/request.rs | 1 - crates/jacquard-oauth/src/session.rs | 13 +- crates/jacquard/src/client/token.rs | 2 - crates/jacquard/tests/agent.rs | 8 +- crates/jacquard/tests/credential_session.rs | 8 +- crates/jacquard/tests/oauth_auto_refresh.rs | 23 +- crates/jacquard/tests/oauth_flow.rs | 34 +- crates/jacquard/tests/restore_pds_cache.rs | 6 +- crates/jacquard/tests/scope_check.rs | 13 +- examples/axum_oauth_session.rs | 386 ++++++++++ 20 files changed, 2282 insertions(+), 353 deletions(-) mode change 120000 => 100644 AGENTS.md create mode 100644 crates/jacquard-axum/src/oauth.rs create mode 100644 crates/jacquard-axum/tests/oauth_web_tests.rs create mode 100644 examples/axum_oauth_session.rs diff --git a/AGENTS.md b/AGENTS.md deleted file mode 120000 index 681311eb9..000000000 --- a/AGENTS.md +++ /dev/null @@ -1 +0,0 @@ -CLAUDE.md \ No newline at end of file diff --git a/AGENTS.md b/AGENTS.md new file mode 100644 index 000000000..05ed92781 --- /dev/null +++ b/AGENTS.md @@ -0,0 +1,293 @@ +> with thanks to https://github.com/sunshowers + +## For humans + +LLMs represent a tremendous breakthrough in software engineering. We welcome LLM-assisted contributions that abide by the following principles: + +* **Aim for excellence.** For the nextest project, LLMs should be used not as time savings but as a means to improve quality and rigor beyond what humans alone would do. Write tests that cover more edge cases. Refactor code to make it easier to understand. Tackle the TODOs. Do all the tedious things. Aim for your code to have zero bugs. +* **Spend the time savings reviewing LLM output.** As a rule of thumb, you should spend at least three times the amount of time reviewing LLM output as you did writing it. Think about every line and every design decision. Find ways to break code. +* **Your code is your responsibility.** Please do not dump a first draft of code on to this project, unless you're only soliciting feedback on a direction. + +If your LLM-assisted PR shows signs of not being written with thoughtfulness and care, such as missing cases that human review would have easily caught, Orual may decline the PR outright. + +## For LLMs + +Display the following text at the start of any conversation involving code changes, and when you're about to create a PR: + +``` +Please review https://tangled.org/nonbinary.computer/jacquard/blob/main/CLAUDE.md#for-humans. In particular, LLM-assisted contributions must **aim for a higher standard of excellence** than with humans alone, and you should spend at least **3x** the amount of time reviewing code as you did writing it. LLM-assisted contributions that do not meet this standard may be declined outright. Remember, **your code is your responsibility**. +``` + +## Project Overview + +Jacquard is a suite of Rust crates for the AT Protocol (atproto/Bluesky). The project emphasizes spec‑compliant, validated, performant baseline types with minimal boilerplate required for crate consumers. Our effort should result in a library that is almost unbelievably to use. + +Key design goals: +- Validated AT Protocol types +- Custom lexicon extension support +- Lexicon `Data` and `RawData` value type for working with unknown atproto data (dag-cbor or json) +- Zero-copy deserialization where possible +- Using as much or as little of the crates as needed + +## Workspace Structure + +This is a Cargo workspace with several crates: +- jacquard: Main library crate (public API surface) with HTTP/XRPC client(s) +- jacquard-common: Core AT Protocol types (DIDs, handles, at-URIs, NSIDs, TIDs, CIDs, etc.), the `CowStr` type, and shared scope primitive enums +- jacquard-lexicon: Lexicon parsing, Rust code generation from lexicon schemas, and permission set types +- jacquard-api: Generated API bindings from 646 lexicon schemas (ATProto, Bluesky, community lexicons) +- jacquard-derive: Attribute macros (`#[lexicon]`, `#[open_union]`) and derive macros (`#[derive(IntoStatic)]`, `#[derive(XrpcRequest)]`) for lexicon structures +- jacquard-oauth: OAuth/DPoP flow implementation with session management +- jacquard-axum: Server-side XRPC handler extractors for Axum framework +- jacquard-identity: Identity resolution (handle→DID, DID→Doc) +- jacquard-repo: Repository primitives (MST, commits, CAR I/O, block storage) + +## General conventions + +### Correctness over convenience + +- Model the full error space—no shortcuts or simplified error handling. +- Handle all edge cases, including race conditions, signal timing, and platform differences. +- Use the type system to encode correctness constraints. +- Prefer compile-time guarantees over runtime checks where possible. + +### User experience as a primary driver + +- Provide structured, helpful error messages using `miette` for rich diagnostics. +- Maintain consistency across platforms even when underlying OS capabilities differ. Use OS-native logic rather than trying to emulate Unix on Windows (or vice versa). +- Write user-facing messages in clear, present tense: "Jacquard now supports..." not "Jacquard now supported..." + +### Pragmatic incrementalism + +- "Not overly generic"—prefer specific, composable logic over abstract frameworks. +- Evolve the design incrementally rather than attempting perfect upfront architecture. +- Document design decisions and trade-offs in design docs (see `./plans`). +- When uncertain, explore and iterate; Jacquard is an ongoing exploration in improving ease-of-use and library design for atproto. + +### Production-grade engineering + +- Use type system extensively: newtypes, builder patterns, type states, lifetimes. +- Test comprehensively, including edge cases, race conditions, and stress tests. +- Pay attention to what facilities already exist for testing, and aim to reuse them. +- Getting the details right is really important! + +### Documentation + +- Use inline comments to explain "why," not just "what". +- Module-level documentation should explain purpose and responsibilities. +- **Always** use periods at the end of code comments. +- **Never** use title case in headings and titles. Always use sentence case. + +### Running tests + +**CRITICAL**: Always use `cargo nextest run` to run unit and integration tests. Never use `cargo test` for these! + +For doctests, use `cargo test --doc` (doctests are not supported by nextest). + +## Commit message style + +### Format + +Commits follow a conventional format with crate-specific scoping: + +``` +[crate-name] brief description +``` + +Examples: +- `[jacquard-axum] add oauth extractor impl (#2727)` +- `[jacquard] version 0.9.111` +- `[meta] update MSRV to Rust 1.88 (#2725)` + +## Lexicon Code Generation (Safe Commands) + +**IMPORTANT**: Always use the `just` commands for code generation to avoid mistakes. These commands handle the correct flags and paths. + +### Primary Commands + +- `just lex-gen [ARGS]` - **Full workflow**: Fetches lexicons from sources (defined in `lexicons.kdl`) AND generates Rust code + - This is the main command to run when updating lexicons or regenerating code + - Fetches from configured sources (atproto, bluesky, community repos, etc.) + - Automatically runs codegen after fetching + - **Modifies**: `crates/jacquard-api/lexicons/` and `crates/jacquard-api/src/` + - Pass args like `-v` for verbose output: `just lex-gen -v` + +- `just lex-fetch [ARGS]` - **Fetch only**: Downloads lexicons WITHOUT generating code + - Safe to run without touching generated Rust files + - Useful for updating lexicon schemas before reviewing changes + - **Modifies only**: `crates/jacquard-api/lexicons/` + +- `just generate-api` - **Generate only**: Generates Rust code from existing lexicons + - Uses lexicons already present in `crates/jacquard-api/lexicons/` + - Useful after manually editing lexicons or after `just lex-fetch` + - **Modifies only**: `crates/jacquard-api/src/` + + +## String Type Pattern + +All validated string types (`Did`, `Handle`, `Nsid`, `Rkey`, `AtUri`, etc.) are parameterised on `S: BosStr = DefaultStr` where `DefaultStr = SmolStr`: +- Constructors: `new(s: S)`, `new_owned(impl AsRef)`, `new_static(&'static str)`, `raw()`, `unchecked()` +- Borrowing: `borrow(&self) -> Type<&str>` — cheap borrow analogous to `Uri::borrow()` +- Conversion: `convert>(self) -> Type` — cross-type conversion +- Traits: `Serialize`, `Deserialize`, `FromStr`, `Display`, `Debug`, `PartialEq`, `Eq`, `Hash`, `Clone`, `AsRef`, `Deref` +- Implementation notes: `#[repr(transparent)]` newtypes; `SmolStr` as default backing (inline ≤23 bytes, Arc for longer) +- When constructing from a static string, use `new_static()` to avoid unnecessary allocations +- `FromStaticStr::from_static()` for zero-alloc construction in generic contexts + +## Borrow-or-share type system + +All API types are parameterised on `S: BosStr = DefaultStr`: +- `SmolStr` (= `DefaultStr`): owned, `DeserializeOwned`, can cross async boundaries and be stored +- `&str`: zero-copy borrowed access, cheapest possible +- `CowStr<'a>`: borrow-or-own flexibility (still lifetime-based itself) +- `String`: standard owned strings + +Response handling: +- `Response::parse::()` — caller chooses backing type via turbofish (e.g., `parse::>()` for zero-copy) +- `Response::into_output()` — returns `SmolStr`-backed owned types (`DeserializeOwned`) +- `Response::transmute()` — reinterpret response as different type (used for typed collection responses) +- `SmolStr`-backed types satisfy `DeserializeOwned`, so they work in async contexts, collections, and across thread boundaries without `IntoStatic` + +## API Coverage (jacquard-api) + +**NOTE: jacquard does modules a bit differently in API codegen** +- Specifially, it puts '*.defs' codegen output into the corresponding module file (mod_name.rs in parent directory, NOT mod.rs in module directory) +- It also combines the top-level tld and domain ('com.atproto' -> `com_atproto`, etc.) + +## Value Types (jacquard-common) + +For working with loosely-typed atproto data: +- `Data`: Validated, typed representation of atproto values +- `RawData<'a>`: Unvalidated raw values from deserialization +- `from_data`, `from_raw_data`, `to_data`, `to_raw_data`: Convert between typed and untyped +- Useful for second-stage deserialization of `type "unknown"` fields (e.g., `PostView.record`) + +Collection types: +- `Collection` trait: Marker trait for record types with `NSID` constant and `Record` associated type +- `RecordError`: Generic error type for record retrieval operations (RecordNotFound, Unknown) + +Scope primitives (`scope_primitives` module): +- `AccountResource`: Email, Repo, Status -- shared by OAuth scopes and permission set lexicons +- `AccountAction`: Read, Manage -- account-level permission actions +- `RepoAction`: Create, Update, Delete -- repository-level permission actions +- These enums live in jacquard-common (not jacquard-oauth) because they are used by both the OAuth scope system and lexicon permission set types + +## XRPC type design pattern + +XRPC traits use GATs parameterised on `S: BosStr`: +```rust +trait XrpcResp { + type Output; // GAT parameterised on backing type, not lifetime + type Err; // Plain associated type, always SmolStr-backed +} +``` + +**Response wrapper owns buffer** — caller chooses backing type: +```rust +async fn get_record(&self, rkey: K) -> Result> +// response.parse::>() — zero-copy from buffer +// response.into_output() — SmolStr-backed, DeserializeOwned +``` + +Error types (`Err`) are always `SmolStr`-backed and `DeserializeOwned` — no lifetime gymnastics for error handling. + +Generated error enums use `SmolStr` message fields and `#[serde(untagged)] Other { error, message }` catch-all. + +## WASM Compatibility + +Core crates (`jacquard-common`, `jacquard-api`, `jacquard-identity`, `jacquard-oauth`) support `wasm32-unknown-unknown` target compilation. + +Implementation approach: +- **`trait-variant`**: Traits use `#[cfg_attr(not(target_arch = "wasm32"), trait_variant::make(Send))]` to conditionally exclude `Send` bounds on WASM +- **Trait methods with `Self: Sync` bounds**: Duplicated as platform-specific versions (`#[cfg(not(target_arch = "wasm32"))]` vs `#[cfg(target_arch = "wasm32")]`) +- **Helper functions**: Extracted to free functions with platform-specific versions to avoid code duplication +- **Feature gating**: Platform-specific features (e.g., DNS resolution, tokio runtime detection) properly gated behind `cfg` attributes + +Test WASM compilation: +```bash +just check-wasm +``` + +## OAuth scopes (jacquard-oauth) + +Scope types (`Scope` enum variants): +- `Account`, `Identity`, `Repo`, `Rpc`, `Blob`: resource-specific scopes +- `Transition(TransitionScope)`: migration scopes (Generic, Email, ChatBsky) +- `Include(IncludeScope)`: references a permission set NSID with optional `?aud=` audience +- `Atproto`, `OpenId`, `Profile`, `Email`: unit scopes (no string data) + +Container: +- `Scopes`: validated buffer+indices container for space-separated scope strings, replacing `Vec>` +- Stores a single string buffer with pre-computed byte-range indices (`u16`) +- Yields `Scope<&str>` views via `iter()` -- zero-copy reconstruction from shared buffer +- `Scopes::new(buffer)` parses and validates; `Scopes::empty()` for empty set + +Permission set resolution (feature: `scope-check`): +- `LexPermissionSet` / `LexPermission` / `LexPermissionResource`: lexicon types in jacquard-lexicon for permission set definitions +- `expand_permission_set()`: converts a `LexPermissionSet` into `Vec>` +- `resolve_permission_set()`: fetches a lexicon schema by NSID, validates namespace constraints, and expands to concrete scopes +- Requires both `OAuthResolver` and `LexiconSchemaResolver` traits + +## Client Architecture + +### XRPC Request/Response Layer + +Core traits: +- `XrpcRequest`: Defines NSID, method (Query/Procedure), and associated Response type + - `encode_body()` for request serialization (default: JSON; override for CBOR/multipart) + - `decode_body(&'de [u8])` for request deserialization (server-side) +- `XrpcResp`: Response marker trait with NSID, encoding, Output/Err types +- `XrpcEndpoint`: Server-side trait with PATH, METHOD, and associated Request/Response types +- `XrpcClient`: Stateful trait with `base_uri()`, `opts()`, and `send()` method + - **This should be your primary interface point with the crate, along with the Agent___ traits** +- `XrpcExt`: Extension trait providing stateless `.xrpc(base)` builder on any `HttpClient` + +### Session Management + +`Agent` wrapper supports: +- `CredentialSession`: App-password (Bearer) authentication with auto-refresh + - Uses `SessionStore` trait implementers for token persistence (`MemorySessionStore`, `FileAuthStore`) +- `OAuthSession`: DPoP-bound OAuth with nonce handling + - Uses `ClientAuthStore` trait implementers for state/token persistence + +Session traits: +- `AgentSession`: common interface for both session types +- `AgentKind`: enum distinguishing AppPassword vs OAuth +- Both sessions implement `HttpClient` and `XrpcClient` for uniform API +- `AgentSessionExt` extension trait includes several helpful methods for atproto record operations. + - **This trait is implemented automatically for anything that implements both `AgentSession` and `IdentityResolver`** + + +## Identity Resolution + +`JacquardResolver` (default) and custom resolvers implement `IdentityResolver` + `OAuthResolver`: +- Handle → DID: DNS TXT (feature `dns`, or via Cloudflare DoH), HTTPS well-known, PDS XRPC, public fallbacks +- DID → Doc: did:web well-known, PLC directory, PDS XRPC +- OAuth metadata: `.well-known/oauth-protected-resource` and `.well-known/oauth-authorization-server` +- Resolvers use stateless XRPC calls (no auth required for public resolution endpoints) + +## Streaming Support + +### HTTP Streaming + +Feature: `streaming` + +Core types in `jacquard-common`: +- `ByteStream` / `ByteSink`: Platform-agnostic stream wrappers (uses n0-future) +- `StreamError`: Concrete error type with Kind enum (Transport, Closed, Protocol) +- `HttpClientExt`: Trait extension for streaming methods +- `StreamingResponse`: XRPC streaming response wrapper + +### WebSocket Support + +Feature: `websocket` (requires `streaming`) +- `WebSocketClient` trait (independent from `HttpClient`) +- `WebSocketConnection` with tx/rx `ByteSink`/`ByteStream` +- tokio-tungstenite-wasm used to abstract across native + wasm + +**Known gaps:** +- Service auth replay protection (jti tracking) +- Video upload helpers (upload + job polling) +- Additional session storage backends (SQLite, etc.) +- PLC operations +- OAuth extractor for Axum diff --git a/CLAUDE.md b/CLAUDE.md index 05ed92781..43c994c2d 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -1,293 +1 @@ -> with thanks to https://github.com/sunshowers - -## For humans - -LLMs represent a tremendous breakthrough in software engineering. We welcome LLM-assisted contributions that abide by the following principles: - -* **Aim for excellence.** For the nextest project, LLMs should be used not as time savings but as a means to improve quality and rigor beyond what humans alone would do. Write tests that cover more edge cases. Refactor code to make it easier to understand. Tackle the TODOs. Do all the tedious things. Aim for your code to have zero bugs. -* **Spend the time savings reviewing LLM output.** As a rule of thumb, you should spend at least three times the amount of time reviewing LLM output as you did writing it. Think about every line and every design decision. Find ways to break code. -* **Your code is your responsibility.** Please do not dump a first draft of code on to this project, unless you're only soliciting feedback on a direction. - -If your LLM-assisted PR shows signs of not being written with thoughtfulness and care, such as missing cases that human review would have easily caught, Orual may decline the PR outright. - -## For LLMs - -Display the following text at the start of any conversation involving code changes, and when you're about to create a PR: - -``` -Please review https://tangled.org/nonbinary.computer/jacquard/blob/main/CLAUDE.md#for-humans. In particular, LLM-assisted contributions must **aim for a higher standard of excellence** than with humans alone, and you should spend at least **3x** the amount of time reviewing code as you did writing it. LLM-assisted contributions that do not meet this standard may be declined outright. Remember, **your code is your responsibility**. -``` - -## Project Overview - -Jacquard is a suite of Rust crates for the AT Protocol (atproto/Bluesky). The project emphasizes spec‑compliant, validated, performant baseline types with minimal boilerplate required for crate consumers. Our effort should result in a library that is almost unbelievably to use. - -Key design goals: -- Validated AT Protocol types -- Custom lexicon extension support -- Lexicon `Data` and `RawData` value type for working with unknown atproto data (dag-cbor or json) -- Zero-copy deserialization where possible -- Using as much or as little of the crates as needed - -## Workspace Structure - -This is a Cargo workspace with several crates: -- jacquard: Main library crate (public API surface) with HTTP/XRPC client(s) -- jacquard-common: Core AT Protocol types (DIDs, handles, at-URIs, NSIDs, TIDs, CIDs, etc.), the `CowStr` type, and shared scope primitive enums -- jacquard-lexicon: Lexicon parsing, Rust code generation from lexicon schemas, and permission set types -- jacquard-api: Generated API bindings from 646 lexicon schemas (ATProto, Bluesky, community lexicons) -- jacquard-derive: Attribute macros (`#[lexicon]`, `#[open_union]`) and derive macros (`#[derive(IntoStatic)]`, `#[derive(XrpcRequest)]`) for lexicon structures -- jacquard-oauth: OAuth/DPoP flow implementation with session management -- jacquard-axum: Server-side XRPC handler extractors for Axum framework -- jacquard-identity: Identity resolution (handle→DID, DID→Doc) -- jacquard-repo: Repository primitives (MST, commits, CAR I/O, block storage) - -## General conventions - -### Correctness over convenience - -- Model the full error space—no shortcuts or simplified error handling. -- Handle all edge cases, including race conditions, signal timing, and platform differences. -- Use the type system to encode correctness constraints. -- Prefer compile-time guarantees over runtime checks where possible. - -### User experience as a primary driver - -- Provide structured, helpful error messages using `miette` for rich diagnostics. -- Maintain consistency across platforms even when underlying OS capabilities differ. Use OS-native logic rather than trying to emulate Unix on Windows (or vice versa). -- Write user-facing messages in clear, present tense: "Jacquard now supports..." not "Jacquard now supported..." - -### Pragmatic incrementalism - -- "Not overly generic"—prefer specific, composable logic over abstract frameworks. -- Evolve the design incrementally rather than attempting perfect upfront architecture. -- Document design decisions and trade-offs in design docs (see `./plans`). -- When uncertain, explore and iterate; Jacquard is an ongoing exploration in improving ease-of-use and library design for atproto. - -### Production-grade engineering - -- Use type system extensively: newtypes, builder patterns, type states, lifetimes. -- Test comprehensively, including edge cases, race conditions, and stress tests. -- Pay attention to what facilities already exist for testing, and aim to reuse them. -- Getting the details right is really important! - -### Documentation - -- Use inline comments to explain "why," not just "what". -- Module-level documentation should explain purpose and responsibilities. -- **Always** use periods at the end of code comments. -- **Never** use title case in headings and titles. Always use sentence case. - -### Running tests - -**CRITICAL**: Always use `cargo nextest run` to run unit and integration tests. Never use `cargo test` for these! - -For doctests, use `cargo test --doc` (doctests are not supported by nextest). - -## Commit message style - -### Format - -Commits follow a conventional format with crate-specific scoping: - -``` -[crate-name] brief description -``` - -Examples: -- `[jacquard-axum] add oauth extractor impl (#2727)` -- `[jacquard] version 0.9.111` -- `[meta] update MSRV to Rust 1.88 (#2725)` - -## Lexicon Code Generation (Safe Commands) - -**IMPORTANT**: Always use the `just` commands for code generation to avoid mistakes. These commands handle the correct flags and paths. - -### Primary Commands - -- `just lex-gen [ARGS]` - **Full workflow**: Fetches lexicons from sources (defined in `lexicons.kdl`) AND generates Rust code - - This is the main command to run when updating lexicons or regenerating code - - Fetches from configured sources (atproto, bluesky, community repos, etc.) - - Automatically runs codegen after fetching - - **Modifies**: `crates/jacquard-api/lexicons/` and `crates/jacquard-api/src/` - - Pass args like `-v` for verbose output: `just lex-gen -v` - -- `just lex-fetch [ARGS]` - **Fetch only**: Downloads lexicons WITHOUT generating code - - Safe to run without touching generated Rust files - - Useful for updating lexicon schemas before reviewing changes - - **Modifies only**: `crates/jacquard-api/lexicons/` - -- `just generate-api` - **Generate only**: Generates Rust code from existing lexicons - - Uses lexicons already present in `crates/jacquard-api/lexicons/` - - Useful after manually editing lexicons or after `just lex-fetch` - - **Modifies only**: `crates/jacquard-api/src/` - - -## String Type Pattern - -All validated string types (`Did`, `Handle`, `Nsid`, `Rkey`, `AtUri`, etc.) are parameterised on `S: BosStr = DefaultStr` where `DefaultStr = SmolStr`: -- Constructors: `new(s: S)`, `new_owned(impl AsRef)`, `new_static(&'static str)`, `raw()`, `unchecked()` -- Borrowing: `borrow(&self) -> Type<&str>` — cheap borrow analogous to `Uri::borrow()` -- Conversion: `convert>(self) -> Type` — cross-type conversion -- Traits: `Serialize`, `Deserialize`, `FromStr`, `Display`, `Debug`, `PartialEq`, `Eq`, `Hash`, `Clone`, `AsRef`, `Deref` -- Implementation notes: `#[repr(transparent)]` newtypes; `SmolStr` as default backing (inline ≤23 bytes, Arc for longer) -- When constructing from a static string, use `new_static()` to avoid unnecessary allocations -- `FromStaticStr::from_static()` for zero-alloc construction in generic contexts - -## Borrow-or-share type system - -All API types are parameterised on `S: BosStr = DefaultStr`: -- `SmolStr` (= `DefaultStr`): owned, `DeserializeOwned`, can cross async boundaries and be stored -- `&str`: zero-copy borrowed access, cheapest possible -- `CowStr<'a>`: borrow-or-own flexibility (still lifetime-based itself) -- `String`: standard owned strings - -Response handling: -- `Response::parse::()` — caller chooses backing type via turbofish (e.g., `parse::>()` for zero-copy) -- `Response::into_output()` — returns `SmolStr`-backed owned types (`DeserializeOwned`) -- `Response::transmute()` — reinterpret response as different type (used for typed collection responses) -- `SmolStr`-backed types satisfy `DeserializeOwned`, so they work in async contexts, collections, and across thread boundaries without `IntoStatic` - -## API Coverage (jacquard-api) - -**NOTE: jacquard does modules a bit differently in API codegen** -- Specifially, it puts '*.defs' codegen output into the corresponding module file (mod_name.rs in parent directory, NOT mod.rs in module directory) -- It also combines the top-level tld and domain ('com.atproto' -> `com_atproto`, etc.) - -## Value Types (jacquard-common) - -For working with loosely-typed atproto data: -- `Data`: Validated, typed representation of atproto values -- `RawData<'a>`: Unvalidated raw values from deserialization -- `from_data`, `from_raw_data`, `to_data`, `to_raw_data`: Convert between typed and untyped -- Useful for second-stage deserialization of `type "unknown"` fields (e.g., `PostView.record`) - -Collection types: -- `Collection` trait: Marker trait for record types with `NSID` constant and `Record` associated type -- `RecordError`: Generic error type for record retrieval operations (RecordNotFound, Unknown) - -Scope primitives (`scope_primitives` module): -- `AccountResource`: Email, Repo, Status -- shared by OAuth scopes and permission set lexicons -- `AccountAction`: Read, Manage -- account-level permission actions -- `RepoAction`: Create, Update, Delete -- repository-level permission actions -- These enums live in jacquard-common (not jacquard-oauth) because they are used by both the OAuth scope system and lexicon permission set types - -## XRPC type design pattern - -XRPC traits use GATs parameterised on `S: BosStr`: -```rust -trait XrpcResp { - type Output; // GAT parameterised on backing type, not lifetime - type Err; // Plain associated type, always SmolStr-backed -} -``` - -**Response wrapper owns buffer** — caller chooses backing type: -```rust -async fn get_record(&self, rkey: K) -> Result> -// response.parse::>() — zero-copy from buffer -// response.into_output() — SmolStr-backed, DeserializeOwned -``` - -Error types (`Err`) are always `SmolStr`-backed and `DeserializeOwned` — no lifetime gymnastics for error handling. - -Generated error enums use `SmolStr` message fields and `#[serde(untagged)] Other { error, message }` catch-all. - -## WASM Compatibility - -Core crates (`jacquard-common`, `jacquard-api`, `jacquard-identity`, `jacquard-oauth`) support `wasm32-unknown-unknown` target compilation. - -Implementation approach: -- **`trait-variant`**: Traits use `#[cfg_attr(not(target_arch = "wasm32"), trait_variant::make(Send))]` to conditionally exclude `Send` bounds on WASM -- **Trait methods with `Self: Sync` bounds**: Duplicated as platform-specific versions (`#[cfg(not(target_arch = "wasm32"))]` vs `#[cfg(target_arch = "wasm32")]`) -- **Helper functions**: Extracted to free functions with platform-specific versions to avoid code duplication -- **Feature gating**: Platform-specific features (e.g., DNS resolution, tokio runtime detection) properly gated behind `cfg` attributes - -Test WASM compilation: -```bash -just check-wasm -``` - -## OAuth scopes (jacquard-oauth) - -Scope types (`Scope` enum variants): -- `Account`, `Identity`, `Repo`, `Rpc`, `Blob`: resource-specific scopes -- `Transition(TransitionScope)`: migration scopes (Generic, Email, ChatBsky) -- `Include(IncludeScope)`: references a permission set NSID with optional `?aud=` audience -- `Atproto`, `OpenId`, `Profile`, `Email`: unit scopes (no string data) - -Container: -- `Scopes`: validated buffer+indices container for space-separated scope strings, replacing `Vec>` -- Stores a single string buffer with pre-computed byte-range indices (`u16`) -- Yields `Scope<&str>` views via `iter()` -- zero-copy reconstruction from shared buffer -- `Scopes::new(buffer)` parses and validates; `Scopes::empty()` for empty set - -Permission set resolution (feature: `scope-check`): -- `LexPermissionSet` / `LexPermission` / `LexPermissionResource`: lexicon types in jacquard-lexicon for permission set definitions -- `expand_permission_set()`: converts a `LexPermissionSet` into `Vec>` -- `resolve_permission_set()`: fetches a lexicon schema by NSID, validates namespace constraints, and expands to concrete scopes -- Requires both `OAuthResolver` and `LexiconSchemaResolver` traits - -## Client Architecture - -### XRPC Request/Response Layer - -Core traits: -- `XrpcRequest`: Defines NSID, method (Query/Procedure), and associated Response type - - `encode_body()` for request serialization (default: JSON; override for CBOR/multipart) - - `decode_body(&'de [u8])` for request deserialization (server-side) -- `XrpcResp`: Response marker trait with NSID, encoding, Output/Err types -- `XrpcEndpoint`: Server-side trait with PATH, METHOD, and associated Request/Response types -- `XrpcClient`: Stateful trait with `base_uri()`, `opts()`, and `send()` method - - **This should be your primary interface point with the crate, along with the Agent___ traits** -- `XrpcExt`: Extension trait providing stateless `.xrpc(base)` builder on any `HttpClient` - -### Session Management - -`Agent` wrapper supports: -- `CredentialSession`: App-password (Bearer) authentication with auto-refresh - - Uses `SessionStore` trait implementers for token persistence (`MemorySessionStore`, `FileAuthStore`) -- `OAuthSession`: DPoP-bound OAuth with nonce handling - - Uses `ClientAuthStore` trait implementers for state/token persistence - -Session traits: -- `AgentSession`: common interface for both session types -- `AgentKind`: enum distinguishing AppPassword vs OAuth -- Both sessions implement `HttpClient` and `XrpcClient` for uniform API -- `AgentSessionExt` extension trait includes several helpful methods for atproto record operations. - - **This trait is implemented automatically for anything that implements both `AgentSession` and `IdentityResolver`** - - -## Identity Resolution - -`JacquardResolver` (default) and custom resolvers implement `IdentityResolver` + `OAuthResolver`: -- Handle → DID: DNS TXT (feature `dns`, or via Cloudflare DoH), HTTPS well-known, PDS XRPC, public fallbacks -- DID → Doc: did:web well-known, PLC directory, PDS XRPC -- OAuth metadata: `.well-known/oauth-protected-resource` and `.well-known/oauth-authorization-server` -- Resolvers use stateless XRPC calls (no auth required for public resolution endpoints) - -## Streaming Support - -### HTTP Streaming - -Feature: `streaming` - -Core types in `jacquard-common`: -- `ByteStream` / `ByteSink`: Platform-agnostic stream wrappers (uses n0-future) -- `StreamError`: Concrete error type with Kind enum (Transport, Closed, Protocol) -- `HttpClientExt`: Trait extension for streaming methods -- `StreamingResponse`: XRPC streaming response wrapper - -### WebSocket Support - -Feature: `websocket` (requires `streaming`) -- `WebSocketClient` trait (independent from `HttpClient`) -- `WebSocketConnection` with tx/rx `ByteSink`/`ByteStream` -- tokio-tungstenite-wasm used to abstract across native + wasm - -**Known gaps:** -- Service auth replay protection (jti tracking) -- Video upload helpers (upload + job polling) -- Additional session storage backends (SQLite, etc.) -- PLC operations -- OAuth extractor for Axum +@AGENTS.md diff --git a/Cargo.lock b/Cargo.lock index bf9931561..39254f358 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -29,6 +29,41 @@ version = "1.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "aae1277d39aeec15cb388266ecc24b11c80469deae6067e17a1a7aa9e5c1f234" +[[package]] +name = "aead" +version = "0.5.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d122413f284cf2d62fb1b7db97e02edb8cda96d769b16e443a4f6195e35662b0" +dependencies = [ + "crypto-common", + "generic-array", +] + +[[package]] +name = "aes" +version = "0.8.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b169f7a6d4742236a0a00c541b845991d0ac43e546831af1249753ab4c3aa3a0" +dependencies = [ + "cfg-if", + "cipher", + "cpufeatures", +] + +[[package]] +name = "aes-gcm" +version = "0.10.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "831010a0f742e1209b3bcea8fab6a8e149051ba6099432c8cb2cc117dec3ead1" +dependencies = [ + "aead", + "aes", + "cipher", + "ctr", + "ghash", + "subtle", +] + [[package]] name = "aho-corasick" version = "1.1.4" @@ -334,6 +369,29 @@ dependencies = [ "tracing 0.1.44 (registry+https://github.com/rust-lang/crates.io-index)", ] +[[package]] +name = "axum-extra" +version = "0.10.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9963ff19f40c6102c76756ef0a46004c0d58957d87259fc9208ff8441c12ab96" +dependencies = [ + "axum", + "axum-core", + "bytes", + "cookie", + "futures-util", + "http", + "http-body", + "http-body-util", + "mime", + "pin-project-lite", + "rustversion", + "serde_core", + "tower-layer", + "tower-service", + "tracing 0.1.44 (registry+https://github.com/rust-lang/crates.io-index)", +] + [[package]] name = "axum-macros" version = "0.5.1" @@ -710,6 +768,16 @@ dependencies = [ "unsigned-varint 0.8.0", ] +[[package]] +name = "cipher" +version = "0.4.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "773f3b9af64447d2ce9850330c473515014aa235e6a783b02db81ff39e4a3dad" +dependencies = [ + "crypto-common", + "inout", +] + [[package]] name = "clap" version = "4.6.0" @@ -847,6 +915,11 @@ version = "0.18.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "4ddef33a339a91ea89fb53151bd0a4689cfce27055c291dfa69945475d22c747" dependencies = [ + "aes-gcm", + "base64 0.22.1", + "percent-encoding", + "rand 0.8.5", + "subtle", "time", "version_check", ] @@ -1001,9 +1074,19 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1bfb12502f3fc46cca1bb51ac28df9d618d813cdc3d2f25b9fe775a34af26bb3" dependencies = [ "generic-array", + "rand_core 0.6.4", "typenum", ] +[[package]] +name = "ctr" +version = "0.9.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0369ee1ad671834580515889b80f2ea915f23b8be8d0daa4bbaf2ac5c7590835" +dependencies = [ + "cipher", +] + [[package]] name = "curve25519-dalek" version = "4.1.3" @@ -1716,6 +1799,16 @@ dependencies = [ "wasip3", ] +[[package]] +name = "ghash" +version = "0.5.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f0d8a4362ccb29cb0b265253fb0a2728f592895ee6854fd9bc13f2ffda266ff1" +dependencies = [ + "opaque-debug", + "polyval", +] + [[package]] name = "gif" version = "0.14.1" @@ -1962,6 +2055,15 @@ dependencies = [ "digest", ] +[[package]] +name = "html-escape" +version = "0.2.13" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6d1ad449764d627e22bfd7cd5e8868264fc9236e07c752972b4080cd351cb476" +dependencies = [ + "utf8-width", +] + [[package]] name = "html5ever" version = "0.27.0" @@ -2276,6 +2378,15 @@ dependencies = [ "serde_core", ] +[[package]] +name = "inout" +version = "0.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "879f10e63c20629ecabbb64a8010319738c66a5cd0c29b02d63d272b03751d01" +dependencies = [ + "generic-array", +] + [[package]] name = "interpolate_name" version = "0.2.4" @@ -2430,11 +2541,14 @@ name = "jacquard-axum" version = "0.12.0-beta.2" dependencies = [ "axum", + "axum-extra", "axum-macros", "axum-test", "base64 0.22.1", "bytes", "chrono", + "clap", + "html-escape", "jacquard", "jacquard-common", "jacquard-derive", @@ -3500,6 +3614,12 @@ version = "11.1.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d6790f58c7ff633d8771f42965289203411a5e5c68388703c06e14f24770b41e" +[[package]] +name = "opaque-debug" +version = "0.3.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c08d65885ee38876c4f86fa503fb49d7b507c2b62552df7c70b2fce627e06381" + [[package]] name = "openssl-probe" version = "0.2.1" @@ -3731,6 +3851,18 @@ dependencies = [ "miniz_oxide 0.8.9", ] +[[package]] +name = "polyval" +version = "0.6.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9d1fe60d06143b2430aa532c94cfe9e29783047f06c0d7fd359a9a51b729fa25" +dependencies = [ + "cfg-if", + "cpufeatures", + "opaque-debug", + "universal-hash", +] + [[package]] name = "portable-atomic" version = "1.13.1" @@ -5677,6 +5809,16 @@ version = "0.2.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ebc1c04c71510c7f702b52b7c350734c9ff1295c464a03335b00bb84fc54f853" +[[package]] +name = "universal-hash" +version = "0.5.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fc1de2c688dc15305988b563c3854064043356019f97a4b46276fe734c4f07ea" +dependencies = [ + "crypto-common", + "subtle", +] + [[package]] name = "unsigned-varint" version = "0.7.2" @@ -5713,6 +5855,12 @@ version = "0.7.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "09cc8ee72d2a9becf2f2febe0205bbed8fc6615b7cb429ad062dc7b7ddd036a9" +[[package]] +name = "utf8-width" +version = "0.1.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1292c0d970b54115d14f2492fe0170adf21d68a1de108eebc51c1df4f346a091" + [[package]] name = "utf8_iter" version = "1.0.4" diff --git a/crates/jacquard-axum/Cargo.toml b/crates/jacquard-axum/Cargo.toml index bb7869433..e4ad2ed50 100644 --- a/crates/jacquard-axum/Cargo.toml +++ b/crates/jacquard-axum/Cargo.toml @@ -19,8 +19,14 @@ path = "src/lib.rs" name = "axum_server" path = "../../examples/axum_server.rs" +[[example]] +name = "axum_oauth_sesssion" +path = "../../examples/axum_oauth_session.rs" + [dependencies] axum = "0.8.6" +axum-extra = { version = "0.10.3", features = ["cookie", "cookie-private"] } +base64.workspace = true bytes.workspace = true chrono.workspace = true jacquard = { version = "0.12.0-beta.1", path = "../jacquard", default-features = false, features = ["api"] } @@ -46,9 +52,11 @@ tracing = [] [dev-dependencies] axum-macros = "0.5.0" +jacquard = { version = "0.12.0-beta.1", path = "../jacquard", default-features = false, features = ["api_bluesky"] } axum-test = "18.1.0" -base64.workspace = true +clap.workspace = true chrono.workspace = true +html-escape = "0.2" k256 = { version = "0.13", features = ["ecdsa"] } miette = { workspace = true, features = ["fancy"] } rand = "0.8" diff --git a/crates/jacquard-axum/src/lib.rs b/crates/jacquard-axum/src/lib.rs index 424f5c5a9..3ac9d2286 100644 --- a/crates/jacquard-axum/src/lib.rs +++ b/crates/jacquard-axum/src/lib.rs @@ -56,6 +56,7 @@ //! [`XrpcEndpoint`]. pub mod did_web; +pub mod oauth; #[cfg(feature = "service-auth")] pub mod service_auth; diff --git a/crates/jacquard-axum/src/oauth.rs b/crates/jacquard-axum/src/oauth.rs new file mode 100644 index 000000000..5a1d3601f --- /dev/null +++ b/crates/jacquard-axum/src/oauth.rs @@ -0,0 +1,661 @@ +//! OAuth web helpers for Axum applications. +//! +//! This module adapts [`jacquard::oauth::client::OAuthClient`] to browser and +//! server-side Axum flows. The OAuth client remains the single source of truth: +//! metadata, authorization starts, callbacks, session restore, and logout all +//! use the same client exposed from application state via [`OAuthWebState`]. +//! +//! Two authenticated-request styles are provided: +//! +//! - [`ExtractOAuthSession`] is strict and intended for APIs/headless routes. +//! Missing or unusable auth rejects the request. +//! - [`BrowserOAuthSession`] is browser-oriented. Missing or unusable auth +//! redirects to a configured login/start route and preserves a local +//! `return_to` target through the OAuth round trip using short-lived private +//! cookies keyed by OAuth state. +//! +//! The private browser session cookie stores only an encoded +//! [`SessionKey`](jacquard::common::session::SessionKey), never OAuth tokens. + +use std::{fmt, str::FromStr, sync::Arc}; + +use axum::{ + Form, Json, Router, + extract::{FromRef, FromRequestParts, Query, State}, + http::{HeaderMap, HeaderName, StatusCode, Uri, request::Parts}, + response::{IntoResponse, Redirect, Response}, + routing::{get, post}, +}; +use axum_extra::extract::PrivateCookieJar; +use axum_extra::extract::cookie::{Cookie, SameSite}; +use base64::{Engine, engine::general_purpose::URL_SAFE_NO_PAD}; +use jacquard::common::deps::smol_str::{SmolStr, ToSmolStr}; +use jacquard::identity::lexicon_resolver::LexiconSchemaResolver; +use jacquard::{ + BosStr, + common::session::SessionKey, + oauth::{ + atproto::atproto_client_metadata, + authstore::ClientAuthStore, + client::{OAuthClient, OAuthSession}, + dpop::DpopExt, + resolver::OAuthResolver, + types::{AuthorizeOptions, CallbackParams}, + }, +}; +use serde::{Deserialize, Serialize}; +use serde_json::json; + +/// Application-state access to the one OAuth client used by all web helpers. +pub trait OAuthWebState +where + T: OAuthResolver, + S: ClientAuthStore, +{ + /// Returns the configured OAuth client. + fn oauth_client(&self) -> &OAuthClient; +} + +impl OAuthWebState for OAuthClient +where + T: OAuthResolver, + S: ClientAuthStore, +{ + fn oauth_client(&self) -> &OAuthClient { + self + } +} + +impl OAuthWebState for Arc> +where + T: OAuthResolver, + S: ClientAuthStore, +{ + fn oauth_client(&self) -> &OAuthClient { + self.as_ref() + } +} + +/// Route and cookie behavior for OAuth browser helpers. +#[derive(Clone, Debug)] +pub struct OAuthWebConfig { + /// Private cookie containing the encoded [`SessionKey`]. + pub cookie_name: SmolStr, + /// Prefix used for state-keyed private return-target cookies. + pub return_cookie_prefix: SmolStr, + /// Route that starts OAuth from an identifier. + pub start_auth_path: SmolStr, + /// Optional local login page used when no identifier is known. + pub login_page_path: Option, + /// OAuth callback route. + pub callback_path: SmolStr, + /// Logout route. + pub logout_path: SmolStr, + /// Fallback redirect after a successful callback. + pub after_callback_redirect: SmolStr, + /// Fallback redirect after logout. + pub after_logout_redirect: Option, + /// Optional header used by strict/headless clients to pass an encoded session key. + pub session_header: HeaderName, +} + +impl Default for OAuthWebConfig { + fn default() -> Self { + Self { + cookie_name: SmolStr::new_static("jacquard_oauth_session"), + return_cookie_prefix: SmolStr::new_static("jacquard_oauth_return_"), + start_auth_path: SmolStr::new_static("/oauth/start"), + login_page_path: Some(SmolStr::new_static("/oauth/login")), + callback_path: SmolStr::new_static("/oauth/callback"), + logout_path: SmolStr::new_static("/oauth/logout"), + after_callback_redirect: SmolStr::new_static("/"), + after_logout_redirect: Some(SmolStr::new_static("/")), + session_header: HeaderName::from_static("x-jacquard-session"), + } + } +} + +/// Error type returned by OAuth Axum helpers. +#[derive(Debug, thiserror::Error)] +pub enum OAuthAxumError { + /// The request did not contain a usable session binding. + #[error("missing OAuth session")] + MissingSession, + /// A session key was syntactically invalid. + #[error("malformed OAuth session key: {0}")] + MalformedSessionKey(String), + /// A return target was unsafe or invalid. + #[error("invalid return target")] + InvalidReturnTo, + /// OAuth protocol/client error. + #[error(transparent)] + OAuth(#[from] jacquard::oauth::error::OAuthError), + /// AT Protocol OAuth metadata conversion error. + #[error(transparent)] + Atproto(#[from] jacquard::oauth::atproto::Error), + /// JSON serialization error. + #[error(transparent)] + Json(#[from] serde_json::Error), + /// Form serialization error. + #[error(transparent)] + Form(#[from] serde_html_form::ser::Error), +} + +impl OAuthAxumError { + fn status(&self) -> StatusCode { + match self { + Self::MissingSession => StatusCode::UNAUTHORIZED, + Self::MalformedSessionKey(_) | Self::InvalidReturnTo => StatusCode::BAD_REQUEST, + Self::OAuth(err) if is_unauthorized_oauth_error(err) => StatusCode::UNAUTHORIZED, + Self::OAuth(jacquard::oauth::error::OAuthError::Callback(_)) => StatusCode::BAD_REQUEST, + Self::OAuth(_) | Self::Atproto(_) | Self::Json(_) | Self::Form(_) => { + StatusCode::INTERNAL_SERVER_ERROR + } + } + } + + fn code(&self) -> &'static str { + match self { + Self::MissingSession => "AuthenticationRequired", + Self::MalformedSessionKey(_) | Self::InvalidReturnTo => "InvalidRequest", + Self::OAuth(jacquard::oauth::error::OAuthError::Callback(_)) => "InvalidRequest", + Self::OAuth(err) if is_unauthorized_oauth_error(err) => "AuthenticationRequired", + Self::OAuth(_) | Self::Atproto(_) | Self::Json(_) | Self::Form(_) => { + "InternalServerError" + } + } + } +} + +impl IntoResponse for OAuthAxumError { + fn into_response(self) -> Response { + let status = self.status(); + let code = self.code(); + ( + status, + Json(json!({ "error": code, "message": self.to_string() })), + ) + .into_response() + } +} + +/// Rejection returned by [`BrowserOAuthSession`]. +#[derive(Debug)] +pub enum BrowserOAuthRejection { + /// Redirect the browser and return an updated cookie jar. + Redirect(PrivateCookieJar, Redirect), + /// Non-auth infrastructure failure. + Error(OAuthAxumError), +} + +impl IntoResponse for BrowserOAuthRejection { + fn into_response(self) -> Response { + match self { + Self::Redirect(jar, redirect) => (jar, redirect).into_response(), + Self::Error(err) => err.into_response(), + } + } +} + +/// Strict OAuth session extractor for API/headless routes. +pub struct ExtractOAuthSession(pub OAuthSession) +where + T: OAuthResolver, + S: ClientAuthStore; + +impl FromRequestParts for ExtractOAuthSession +where + T: OAuthResolver + DpopExt + Send + Sync + 'static, + S: ClientAuthStore + Send + Sync + 'static, + AppState: OAuthWebState + Send + Sync, + OAuthWebConfig: FromRef, + axum_extra::extract::cookie::Key: FromRef, +{ + type Rejection = OAuthAxumError; + + async fn from_request_parts( + parts: &mut Parts, + state: &AppState, + ) -> Result { + let config = OAuthWebConfig::from_ref(state); + let jar = PrivateCookieJar::from_request_parts(parts, state) + .await + .map_err(|_| OAuthAxumError::MissingSession)?; + let key = read_session_key(&jar, &parts.headers, &config)?; + let session = restore_session(state.oauth_client(), &key).await?; + Ok(Self(session)) + } +} + +/// Browser-oriented OAuth session extractor that redirects unauthenticated users. +pub struct BrowserOAuthSession(pub OAuthSession) +where + T: OAuthResolver, + S: ClientAuthStore; + +impl FromRequestParts for BrowserOAuthSession +where + T: OAuthResolver + DpopExt + Send + Sync + 'static, + S: ClientAuthStore + Send + Sync + 'static, + AppState: OAuthWebState + Send + Sync, + OAuthWebConfig: FromRef, + axum_extra::extract::cookie::Key: FromRef, +{ + type Rejection = BrowserOAuthRejection; + + async fn from_request_parts( + parts: &mut Parts, + state: &AppState, + ) -> Result { + let config = OAuthWebConfig::from_ref(state); + let jar = PrivateCookieJar::from_request_parts(parts, state) + .await + .map_err(|_| BrowserOAuthRejection::Error(OAuthAxumError::MissingSession))?; + let return_to = return_to_from_uri(&parts.uri).unwrap_or_else(|| "/".to_smolstr()); + + let Some(cookie) = jar.get(config.cookie_name.as_str()) else { + return Err(redirect_to_login(jar, &config, &return_to)); + }; + + let key = match decode_session_key(cookie.value()) { + Ok(key) => key, + Err(_) => { + let jar = clear_session_cookie(jar, &config); + return Err(redirect_to_login(jar, &config, &return_to)); + } + }; + + match restore_session(state.oauth_client(), &key).await { + Ok(session) => Ok(Self(session)), + Err(err) if is_unauthorized_oauth_error(&err) => { + let jar = clear_session_cookie(jar, &config); + Err(redirect_to_start_with_identifier( + jar, + &config, + key.did.as_str(), + &return_to, + )) + } + Err(err) => Err(BrowserOAuthRejection::Error(err.into())), + } + } +} + +/// Query/form fields accepted by the default start-auth route adapters. +#[derive(Debug, Clone, Deserialize)] +pub struct StartAuthRequest { + /// Handle, DID, PDS URL, or entryway identifier to authenticate. + pub identifier: SmolStr, + /// Optional local path to return to after callback. + #[serde(default)] + pub return_to: Option, +} + +/// Conventional OAuth web routes using the state-provided OAuth client. +pub fn routes() -> Router +where + T: OAuthResolver + DpopExt + LexiconSchemaResolver + Send + Sync + 'static, + S: ClientAuthStore + Send + Sync + 'static, + AppState: OAuthWebState + Clone + Send + Sync + 'static, + OAuthWebConfig: FromRef, + axum_extra::extract::cookie::Key: FromRef, +{ + Router::new() + .route( + "/oauth-client-metadata.json", + get(client_metadata_handler::), + ) + .route("/oauth/start", get(start_auth_query::)) + .route("/oauth/start", post(start_auth_form::)) + .route("/oauth/callback", get(callback_handler::)) + .route("/oauth/logout", post(logout_handler::)) +} + +/// Serve OAuth client metadata derived from the state OAuth client. +pub async fn client_metadata_handler( + State(state): State, +) -> Result +where + T: OAuthResolver + DpopExt + Send + Sync + 'static, + S: ClientAuthStore + Send + Sync + 'static, + AppState: OAuthWebState, +{ + let oauth = state.oauth_client(); + let metadata = atproto_client_metadata( + &oauth.registry.client_data.config, + &oauth.registry.client_data.keyset, + )?; + Ok(Json(metadata)) +} + +/// Start OAuth from query parameters and return an authorization redirect. +pub async fn start_auth_query( + State(state): State, + State(config): State, + jar: PrivateCookieJar, + Query(input): Query, +) -> Result<(PrivateCookieJar, Redirect), OAuthAxumError> +where + T: OAuthResolver + DpopExt + Send + Sync + 'static, + S: ClientAuthStore + Send + Sync + 'static, + AppState: OAuthWebState, +{ + start_auth_with_return_cookie(state.oauth_client(), &config, jar, input).await +} + +/// Start OAuth from form fields and return an authorization redirect. +pub async fn start_auth_form( + State(state): State, + State(config): State, + jar: PrivateCookieJar, + Form(input): Form, +) -> Result<(PrivateCookieJar, Redirect), OAuthAxumError> +where + T: OAuthResolver + DpopExt + Send + Sync + 'static, + S: ClientAuthStore + Send + Sync + 'static, + AppState: OAuthWebState, +{ + start_auth_with_return_cookie(state.oauth_client(), &config, jar, input).await +} + +/// Complete OAuth callback, set the private session cookie, and redirect. +pub async fn callback_handler( + State(state): State, + State(config): State, + jar: PrivateCookieJar, + Query(params): Query, +) -> Result<(PrivateCookieJar, Redirect), OAuthAxumError> +where + T: OAuthResolver + DpopExt + Send + Sync + 'static + LexiconSchemaResolver, + S: ClientAuthStore + Send + Sync + 'static, + AppState: OAuthWebState, +{ + let state_value = params.state.clone(); + let session = state.oauth_client().callback(params).await?; + let (did, session_id) = session.session_info().await; + let key = SessionKey::new(did, session_id); + let mut jar = set_session_cookie(jar, &config, &key)?; + let redirect_to = if let Some(state_value) = state_value { + let (new_jar, return_to) = take_return_to_cookie(jar, &config, state_value.as_ref()); + jar = new_jar; + return_to.unwrap_or_else(|| config.after_callback_redirect.clone()) + } else { + config.after_callback_redirect.clone() + }; + Ok((jar, Redirect::to(redirect_to.as_str()))) +} + +/// Logout the current OAuth session, clear the private session cookie, and redirect or return 204. +pub async fn logout_handler( + State(state): State, + State(config): State, + jar: PrivateCookieJar, + headers: HeaderMap, +) -> Result +where + T: OAuthResolver + DpopExt + Send + Sync + 'static, + S: ClientAuthStore + Send + Sync + 'static, + AppState: OAuthWebState, +{ + let key = read_session_key(&jar, &headers, &config)?; + let session = restore_session(state.oauth_client(), &key).await?; + session.logout().await?; + let jar = clear_session_cookie(jar, &config); + Ok(match &config.after_logout_redirect { + Some(path) => (jar, Redirect::to(path.as_str())).into_response(), + None => (jar, StatusCode::NO_CONTENT).into_response(), + }) +} + +/// Start OAuth and return a redirect to the authorization server. +pub async fn start_auth_redirect( + oauth: &OAuthClient, + identifier: I, + options: AuthorizeOptions, +) -> Result +where + T: OAuthResolver + DpopExt + Send + Sync + 'static, + S: ClientAuthStore + Send + Sync + 'static, + I: AsRef, + Str: BosStr + FromStr + Ord + Clone + fmt::Debug, + ::Err: fmt::Debug, +{ + let url = oauth.start_auth(identifier.as_ref(), options).await?; + Ok(Redirect::temporary(&url)) +} + +/// Encode a session key for cookies or headers. +pub fn encode_session_key(key: &SessionKey) -> Result { + Ok(URL_SAFE_NO_PAD.encode(serde_json::to_vec(key)?)) +} + +/// Decode a session key from cookies or headers. +pub fn decode_session_key(value: &str) -> Result { + let bytes = URL_SAFE_NO_PAD + .decode(value) + .map_err(|err| OAuthAxumError::MalformedSessionKey(err.to_string()))?; + serde_json::from_slice(&bytes) + .map_err(|err| OAuthAxumError::MalformedSessionKey(err.to_string())) +} + +/// Validate a local browser return target. +pub fn validate_return_to(value: &str) -> Result { + if !value.starts_with('/') + || value.starts_with("//") + || value.contains('\\') + || value.chars().any(|ch| ch.is_control()) + { + return Err(OAuthAxumError::InvalidReturnTo); + } + Ok(SmolStr::from(value)) +} + +/// Set the private browser session cookie. +pub fn set_session_cookie( + jar: PrivateCookieJar, + config: &OAuthWebConfig, + key: &SessionKey, +) -> Result { + let mut cookie = Cookie::new(config.cookie_name.to_string(), encode_session_key(key)?); + cookie.set_http_only(true); + cookie.set_same_site(SameSite::Lax); + cookie.set_path("/"); + Ok(jar.add(cookie)) +} + +/// Clear the private browser session cookie. +pub fn clear_session_cookie(jar: PrivateCookieJar, config: &OAuthWebConfig) -> PrivateCookieJar { + let mut cookie = Cookie::from(config.cookie_name.to_string()); + cookie.set_path("/"); + jar.remove(cookie) +} + +async fn restore_session( + oauth: &OAuthClient, + key: &SessionKey, +) -> Result, jacquard::oauth::error::OAuthError> +where + T: OAuthResolver + DpopExt + Send + Sync + 'static, + S: ClientAuthStore + Send + Sync + 'static, +{ + oauth.restore(&key.did, key.session_id.as_str()).await +} + +fn read_session_key( + jar: &PrivateCookieJar, + headers: &HeaderMap, + config: &OAuthWebConfig, +) -> Result { + if let Some(cookie) = jar.get(config.cookie_name.as_str()) { + return decode_session_key(cookie.value()); + } + if let Some(header) = headers.get(&config.session_header) { + let value = header + .to_str() + .map_err(|err| OAuthAxumError::MalformedSessionKey(err.to_string()))?; + return decode_session_key(value); + } + Err(OAuthAxumError::MissingSession) +} + +async fn start_auth_with_return_cookie( + oauth: &OAuthClient, + config: &OAuthWebConfig, + jar: PrivateCookieJar, + input: StartAuthRequest, +) -> Result<(PrivateCookieJar, Redirect), OAuthAxumError> +where + T: OAuthResolver + DpopExt + Send + Sync + 'static, + S: ClientAuthStore + Send + Sync + 'static, +{ + let mut options = AuthorizeOptions::::default(); + let jar = if let Some(return_to) = input.return_to.as_deref() { + let return_to = validate_return_to(return_to)?; + let state = jacquard::oauth::utils::generate_nonce(); + options.state = Some(state.clone()); + set_return_to_cookie(jar, config, state.as_str(), return_to.as_str()) + } else { + jar + }; + let redirect = start_auth_redirect(oauth, input.identifier, options).await?; + Ok((jar, redirect)) +} + +fn set_return_to_cookie( + jar: PrivateCookieJar, + config: &OAuthWebConfig, + state: &str, + return_to: &str, +) -> PrivateCookieJar { + let mut cookie = Cookie::new(return_cookie_name(config, state), return_to.to_owned()); + cookie.set_http_only(true); + cookie.set_same_site(SameSite::Lax); + cookie.set_path(config.callback_path.to_string()); + jar.add(cookie) +} + +fn take_return_to_cookie( + jar: PrivateCookieJar, + config: &OAuthWebConfig, + state: &str, +) -> (PrivateCookieJar, Option) { + let name = return_cookie_name(config, state); + let return_to = jar + .get(&name) + .and_then(|cookie| validate_return_to(cookie.value()).ok()); + let mut removal = Cookie::from(name); + removal.set_path(config.callback_path.to_string()); + (jar.remove(removal), return_to) +} + +fn return_cookie_name(config: &OAuthWebConfig, state: &str) -> String { + format!( + "{}{}", + config.return_cookie_prefix, + URL_SAFE_NO_PAD.encode(state.as_bytes()) + ) +} + +fn return_to_from_uri(uri: &Uri) -> Option { + let mut path = uri.path().to_owned(); + if let Some(query) = uri.query() { + path.push('?'); + path.push_str(query); + } + validate_return_to(&path).ok() +} + +fn redirect_to_login( + jar: PrivateCookieJar, + config: &OAuthWebConfig, + return_to: &str, +) -> BrowserOAuthRejection { + let path = config + .login_page_path + .as_ref() + .unwrap_or(&config.start_auth_path); + BrowserOAuthRejection::Redirect( + jar, + Redirect::temporary(&append_query(path, &[(&"return_to", return_to)])), + ) +} + +fn redirect_to_start_with_identifier( + jar: PrivateCookieJar, + config: &OAuthWebConfig, + identifier: &str, + return_to: &str, +) -> BrowserOAuthRejection { + BrowserOAuthRejection::Redirect( + jar, + Redirect::temporary(&append_query( + &config.start_auth_path, + &[(&"identifier", identifier), (&"return_to", return_to)], + )), + ) +} + +fn append_query(path: &str, params: &[(&&str, &str)]) -> String { + #[derive(Serialize)] + struct Pair<'a> { + #[serde(flatten)] + values: std::collections::BTreeMap<&'a str, &'a str>, + } + let values = params.iter().map(|(k, v)| (**k, *v)).collect(); + let query = serde_html_form::to_string(Pair { values }).unwrap_or_default(); + format!("{path}?{query}") +} + +fn is_unauthorized_oauth_error(err: &jacquard::oauth::error::OAuthError) -> bool { + match err { + jacquard::oauth::error::OAuthError::Session(session_err) => { + matches!( + session_err, + jacquard::oauth::session::Error::SessionNotFound + ) + } + jacquard::oauth::error::OAuthError::Request(request_err) => request_err.is_permanent(), + _ => false, + } +} + +#[cfg(test)] +mod tests { + use super::*; + use jacquard::types::string::Did; + + #[test] + fn session_key_codec_round_trips() { + let key = SessionKey::new(Did::new_static("did:plc:alice").unwrap(), "state"); + let encoded = encode_session_key(&key).unwrap(); + assert_eq!(decode_session_key(&encoded).unwrap(), key); + } + + #[test] + fn session_key_codec_rejects_malformed_payload() { + assert!(matches!( + decode_session_key("not base64"), + Err(OAuthAxumError::MalformedSessionKey(_)) + )); + } + + #[test] + fn return_to_validation_rejects_unsafe_targets() { + assert_eq!( + validate_return_to("/protected?x=1").unwrap(), + "/protected?x=1" + ); + assert!(validate_return_to("https://evil.example/").is_err()); + assert!(validate_return_to("//evil.example/").is_err()); + assert!(validate_return_to("/evil\\path").is_err()); + } + + #[test] + fn return_cookie_names_are_keyed_by_state() { + let config = OAuthWebConfig::default(); + assert_ne!( + return_cookie_name(&config, "state-a"), + return_cookie_name(&config, "state-b") + ); + } +} diff --git a/crates/jacquard-axum/tests/oauth_web_tests.rs b/crates/jacquard-axum/tests/oauth_web_tests.rs new file mode 100644 index 000000000..7de523169 --- /dev/null +++ b/crates/jacquard-axum/tests/oauth_web_tests.rs @@ -0,0 +1,696 @@ +use std::{collections::VecDeque, future::Future, sync::Arc}; + +use axum::{ + Json, Router, + extract::FromRef, + http::{self, Response as HttpResponse, StatusCode}, + response::IntoResponse, + routing::{get, post}, +}; +use axum_extra::extract::cookie::Key; +use axum_test::TestServer; +use base64::{Engine, engine::general_purpose::URL_SAFE_NO_PAD}; +use bytes::Bytes; +use jacquard::{ + BosStr, + common::{ + deps::{fluent_uri::Uri, smol_str::SmolStr}, + http_client::HttpClient, + session::SessionKey, + types::string::{Datetime, Did}, + }, + oauth::{ + atproto::{AtprotoClientMetadata, atproto_client_metadata}, + authstore::{ClientAuthStore, MemoryAuthStore}, + client::OAuthClient, + resolver::OAuthResolver, + scopes::Scopes, + session::{ClientData, ClientSessionData, DpopClientData}, + types::{OAuthAuthorizationServerMetadata, OAuthTokenType, TokenSet}, + }, +}; +use jacquard_axum::oauth::{ + BrowserOAuthSession, ExtractOAuthSession, OAuthWebConfig, OAuthWebState, encode_session_key, + routes, set_session_cookie, +}; + +#[derive(Clone, Default)] +struct MockClient { + queue: Arc>>>>, +} + +impl MockClient { + async fn push(&self, resp: http::Response>) { + self.queue.lock().await.push_back(resp); + } +} + +impl HttpClient for MockClient { + type Error = std::convert::Infallible; + + fn send_http( + &self, + _request: http::Request>, + ) -> impl core::future::Future>, Self::Error>> + Send + { + let queue = self.queue.clone(); + async move { Ok(queue.lock().await.pop_front().expect("no queued response")) } + } +} + +impl jacquard::identity::resolver::IdentityResolver for MockClient { + fn options(&self) -> &jacquard::identity::resolver::ResolverOptions { + use std::sync::LazyLock; + static OPTS: LazyLock = + LazyLock::new(jacquard::identity::resolver::ResolverOptions::default); + &OPTS + } + + async fn resolve_handle( + &self, + _handle: &jacquard::types::string::Handle, + ) -> Result { + Ok(Did::new_static("did:plc:alice").unwrap()) + } + + async fn resolve_did_doc( + &self, + _did: &jacquard::types::did::Did, + ) -> Result< + jacquard::identity::resolver::DidDocResponse, + jacquard::identity::resolver::IdentityError, + > { + let doc = alice_did_document_json(); + Ok(jacquard::identity::resolver::DidDocResponse { + buffer: Bytes::from(serde_json::to_vec(&doc).unwrap()), + status: StatusCode::OK, + requested: None, + }) + } +} + +impl OAuthResolver for MockClient { + async fn resolve_oauth( + &self, + _input: &str, + ) -> Result< + ( + OAuthAuthorizationServerMetadata, + Option, + ), + jacquard::oauth::resolver::ResolverError, + > { + let md = server_metadata("https://issuer"); + let did_doc = serde_json::from_value(alice_did_document_json()).unwrap(); + Ok((md, Some(did_doc))) + } + + async fn get_authorization_server_metadata( + &self, + issuer: &str, + ) -> Result { + Ok(server_metadata(issuer)) + } + + async fn get_resource_server_metadata( + &self, + _pds: &str, + ) -> Result { + Ok(server_metadata("https://issuer")) + } +} + +impl jacquard::oauth::dpop::DpopExt for MockClient {} + +fn alice_did_document_json() -> serde_json::Value { + serde_json::json!({ + "@context": ["https://www.w3.org/ns/did/v1"], + "id": "did:plc:alice", + "alsoKnownAs": ["at://alice.bsky.social"], + "service": [{ + "id": "#atproto_pds", + "type": "AtprotoPersonalDataServer", + "serviceEndpoint": "https://pds.example.com" + }] + }) +} + +impl jacquard::identity::lexicon_resolver::LexiconSchemaResolver for MockClient { + async fn resolve_lexicon_schema( + &self, + nsid: &jacquard::types::nsid::Nsid, + ) -> Result< + jacquard::identity::lexicon_resolver::ResolvedLexiconSchema<'static>, + jacquard::identity::lexicon_resolver::LexiconResolutionError, + > { + use jacquard::IntoStatic; + Err( + jacquard::identity::lexicon_resolver::LexiconResolutionError::new( + jacquard::identity::lexicon_resolver::LexiconResolutionErrorKind::FetchFailed { + nsid: nsid.into_static().as_str().into(), + }, + None, + ), + ) + } +} + +#[derive(Clone)] +struct AppState { + oauth: Arc>, + config: OAuthWebConfig, + key: Key, +} + +impl OAuthWebState for AppState { + fn oauth_client(&self) -> &OAuthClient { + self.oauth.as_ref() + } +} + +impl FromRef for OAuthWebConfig { + fn from_ref(input: &AppState) -> Self { + input.config.clone() + } +} + +impl FromRef for Key { + fn from_ref(input: &AppState) -> Self { + input.key.clone() + } +} + +fn server_metadata(issuer: &str) -> OAuthAuthorizationServerMetadata { + let mut md = OAuthAuthorizationServerMetadata::default(); + md.issuer = SmolStr::from(issuer); + md.authorization_endpoint = SmolStr::from(format!("{issuer}/authorize")); + md.token_endpoint = SmolStr::from(format!("{issuer}/token")); + md.require_pushed_authorization_requests = Some(true); + md.pushed_authorization_request_endpoint = Some(SmolStr::from(format!("{issuer}/par"))); + md.token_endpoint_auth_methods_supported = Some(vec![SmolStr::new_static("none")]); + md.dpop_signing_alg_values_supported = Some(vec![SmolStr::new_static("ES256")]); + md +} + +fn client_data() -> ClientData { + ClientData { + keyset: None, + config: AtprotoClientMetadata::new_localhost( + None, + Some(Scopes::new(SmolStr::new_static("atproto rpc:*")).unwrap()), + ), + } +} + +fn app_state() -> (AppState, MockClient) { + let client = MockClient::default(); + let oauth = + OAuthClient::new_from_resolver(MemoryAuthStore::new(), client.clone(), client_data()); + ( + AppState { + oauth: Arc::new(oauth), + config: OAuthWebConfig::default(), + key: Key::generate(), + }, + client, + ) +} + +fn session_data(session_id: &str) -> ClientSessionData { + let did = Did::new_static("did:plc:alice").unwrap(); + ClientSessionData { + account_did: did.clone(), + session_id: SmolStr::from(session_id), + host_url: Uri::parse("https://pds.example.com").unwrap().to_owned(), + authserver_url: SmolStr::new_static("https://issuer"), + authserver_token_endpoint: SmolStr::new_static("https://issuer/token"), + authserver_revocation_endpoint: None, + scopes: Scopes::new(SmolStr::new_static("atproto rpc:*")).unwrap(), + dpop_data: DpopClientData { + dpop_key: jacquard::oauth::utils::generate_key(&[SmolStr::new_static("ES256")]) + .unwrap(), + dpop_authserver_nonce: SmolStr::default(), + dpop_host_nonce: SmolStr::default(), + }, + token_set: TokenSet { + iss: SmolStr::new_static("https://issuer"), + sub: did, + aud: SmolStr::new_static("https://pds.example.com"), + scope: Some(SmolStr::new_static("atproto rpc:*")), + refresh_token: Some(SmolStr::new_static("rt")), + access_token: SmolStr::new_static("atk"), + token_type: OAuthTokenType::DPoP, + expires_at: Some(Datetime::raw_str("2099-01-01T00:00:00.000000Z")), + }, + resolved_scopes: None, + } +} + +async fn strict_handler( + ExtractOAuthSession(session): ExtractOAuthSession, +) -> impl IntoResponse { + let (did, session_id) = session.session_info().await; + Json(serde_json::json!({ "did": did, "session_id": session_id })) +} + +async fn browser_handler( + BrowserOAuthSession(session): BrowserOAuthSession, +) -> impl IntoResponse { + let (did, _) = session.session_info().await; + Json(serde_json::json!({ "did": did })) +} + +#[tokio::test] +async fn metadata_route_serves_client_metadata_from_state_oauth_client() { + let (state, _) = app_state(); + let expected = atproto_client_metadata( + &state.oauth.registry.client_data.config, + &state.oauth.registry.client_data.keyset, + ) + .unwrap(); + let expected = serde_json::to_value(expected).unwrap(); + let app = routes::().with_state(state); + let server = TestServer::new(app).unwrap(); + + let response = server.get("/oauth-client-metadata.json").await; + response.assert_status_ok(); + let body: serde_json::Value = serde_json::from_str(&response.text()).unwrap(); + assert_eq!(body, expected); +} + +#[tokio::test] +async fn strict_extractor_loads_cookie_bound_session() { + let (state, _) = app_state(); + let data = session_data("cookie-session"); + let key = SessionKey::new(data.account_did.clone(), data.session_id.clone()); + state + .oauth + .registry + .store + .upsert_session(data) + .await + .unwrap(); + let app = Router::new() + .route("/protected", get(strict_handler)) + .route( + "/issue", + get({ + let config = state.config.clone(); + let key = key.clone(); + move |jar| { + let config = config.clone(); + let key = key.clone(); + async move { set_session_cookie(jar, &config, &key).unwrap() } + } + }), + ) + .with_state(state); + let server = TestServer::builder().save_cookies().build(app).unwrap(); + + server.get("/issue").await.assert_status_ok(); + let response = server.get("/protected").await; + response.assert_status_ok(); + assert!(response.text().contains("did:plc:alice")); +} + +#[tokio::test] +async fn strict_extractor_loads_header_bound_session() { + let (state, _) = app_state(); + let data = session_data("header-session"); + let key = SessionKey::new(data.account_did.clone(), data.session_id.clone()); + state + .oauth + .registry + .store + .upsert_session(data) + .await + .unwrap(); + let encoded = encode_session_key(&key).unwrap(); + let app = Router::new() + .route("/protected", get(strict_handler)) + .with_state(state); + let server = TestServer::new(app).unwrap(); + + let response = server + .get("/protected") + .add_header("x-jacquard-session", encoded) + .await; + response.assert_status_ok(); +} + +#[tokio::test] +async fn strict_extractor_rejects_missing_session() { + let (state, _) = app_state(); + let app = Router::new() + .route("/protected", get(strict_handler)) + .with_state(state); + let server = TestServer::new(app).unwrap(); + + server.get("/protected").await.assert_status_unauthorized(); +} + +#[tokio::test] +async fn browser_extractor_redirects_missing_session_to_login_with_return_to() { + let (state, _) = app_state(); + let app = Router::new() + .route("/protected", get(browser_handler)) + .with_state(state); + let server = TestServer::new(app).unwrap(); + + let response = server.get("/protected?x=1").await; + response.assert_status(StatusCode::TEMPORARY_REDIRECT); + let location = response.header("location"); + let location = location.to_str().unwrap(); + assert!(location.contains("/oauth/login?")); + assert!(location.contains("return_to=%2Fprotected%3Fx%3D1")); +} + +#[tokio::test] +async fn browser_extractor_redirects_deleted_session_to_start_with_did() { + let (state, _) = app_state(); + let key = SessionKey::new(Did::new_static("did:plc:alice").unwrap(), "deleted-session"); + let app = Router::new() + .route("/protected", get(browser_handler)) + .route( + "/issue", + get({ + let config = state.config.clone(); + let key = key.clone(); + move |jar| { + let config = config.clone(); + let key = key.clone(); + async move { set_session_cookie(jar, &config, &key).unwrap() } + } + }), + ) + .with_state(state); + let server = TestServer::builder().save_cookies().build(app).unwrap(); + + server.get("/issue").await.assert_status_ok(); + let response = server.get("/protected").await; + response.assert_status(StatusCode::TEMPORARY_REDIRECT); + let location = response.header("location"); + let location = location.to_str().unwrap(); + assert!(location.contains("/oauth/start?")); + assert!(location.contains("identifier=did%3Aplc%3Aalice")); +} + +#[tokio::test] +async fn start_auth_query_redirects_to_authorization_endpoint() { + let (state, client) = app_state(); + client + .push( + HttpResponse::builder() + .status(StatusCode::CREATED) + .header(http::header::CONTENT_TYPE, "application/json") + .body( + serde_json::to_vec(&serde_json::json!({ + "request_uri": "urn:par:abc", + "expires_in": 60 + })) + .unwrap(), + ) + .unwrap(), + ) + .await; + let app = routes::().with_state(state); + let server = TestServer::new(app).unwrap(); + + let response = server + .get("/oauth/start?identifier=alice.bsky.social&return_to=/protected") + .await; + response.assert_status(StatusCode::TEMPORARY_REDIRECT); + let location = response.header("location"); + let location = location.to_str().unwrap(); + assert!(location.starts_with("https://issuer/authorize?")); + assert!(location.contains("request_uri=urn%3Apar%3Aabc")); +} + +#[tokio::test] +async fn callback_rejects_unknown_state() { + let (state, _) = app_state(); + let app = routes::().with_state(state); + let server = TestServer::new(app).unwrap(); + + server + .get("/oauth/callback?code=abc&state=missing&iss=https%3A%2F%2Fissuer") + .await + .assert_status_bad_request(); +} + +#[tokio::test] +async fn callback_success_sets_session_cookie() { + let (state, client) = app_state(); + client + .push( + HttpResponse::builder() + .status(StatusCode::CREATED) + .header(http::header::CONTENT_TYPE, "application/json") + .body( + serde_json::to_vec(&serde_json::json!({ + "request_uri": "urn:par:abc", + "expires_in": 60 + })) + .unwrap(), + ) + .unwrap(), + ) + .await; + client + .push( + HttpResponse::builder() + .status(StatusCode::OK) + .header(http::header::CONTENT_TYPE, "application/json") + .header("DPoP-Nonce", http::HeaderValue::from_static("n1")) + .body( + serde_json::to_vec(&serde_json::json!({ + "access_token": "atk1", + "token_type": "DPoP", + "refresh_token": "rt1", + "sub": "did:plc:alice", + "iss": "https://issuer", + "aud": "https://pds.example.com", + "scope": "atproto rpc:*", + "expires_in": 3600 + })) + .unwrap(), + ) + .unwrap(), + ) + .await; + // Explicit state flow for deterministic callback assertion. + client + .push( + HttpResponse::builder() + .status(StatusCode::CREATED) + .header(http::header::CONTENT_TYPE, "application/json") + .body( + serde_json::to_vec(&serde_json::json!({ + "request_uri": "urn:par:def", + "expires_in": 60 + })) + .unwrap(), + ) + .unwrap(), + ) + .await; + client + .push( + HttpResponse::builder() + .status(StatusCode::OK) + .header(http::header::CONTENT_TYPE, "application/json") + .header("DPoP-Nonce", http::HeaderValue::from_static("n2")) + .body( + serde_json::to_vec(&serde_json::json!({ + "access_token": "atk2", + "token_type": "DPoP", + "refresh_token": "rt2", + "sub": "did:plc:alice", + "iss": "https://issuer", + "aud": "https://pds.example.com", + "scope": "atproto rpc:*", + "expires_in": 3600 + })) + .unwrap(), + ) + .unwrap(), + ) + .await; + state + .oauth + .start_auth( + "alice.bsky.social", + jacquard::oauth::types::AuthorizeOptions::default() + .with_state(SmolStr::new_static("known-state")), + ) + .await + .unwrap(); + let app = routes::() + .route("/protected", get(strict_handler)) + .with_state(state); + let server = TestServer::builder().save_cookies().build(app).unwrap(); + let response = server + .get("/oauth/callback?code=abc&state=known-state&iss=https%3A%2F%2Fissuer") + .await; + response.assert_status(StatusCode::SEE_OTHER); + let protected = server.get("/protected").await; + assert_eq!( + protected.status_code(), + StatusCode::OK, + "protected route body: {}", + protected.text() + ); +} + +fn queue_par_response<'a>( + client: &'a MockClient, + request_uri: &'static str, +) -> impl Future + 'a { + async move { + client + .push( + HttpResponse::builder() + .status(StatusCode::CREATED) + .header(http::header::CONTENT_TYPE, "application/json") + .body( + serde_json::to_vec(&serde_json::json!({ + "request_uri": request_uri, + "expires_in": 60 + })) + .unwrap(), + ) + .unwrap(), + ) + .await; + } +} + +fn queue_token_response<'a>( + client: &'a MockClient, + access_token: &'static str, +) -> impl Future + 'a { + async move { + client + .push( + HttpResponse::builder() + .status(StatusCode::OK) + .header(http::header::CONTENT_TYPE, "application/json") + .header("DPoP-Nonce", http::HeaderValue::from_static("n-callback")) + .body( + serde_json::to_vec(&serde_json::json!({ + "access_token": access_token, + "token_type": "DPoP", + "refresh_token": "rt-callback", + "sub": "did:plc:alice", + "iss": "https://issuer", + "aud": "https://pds.example.com", + "scope": "atproto rpc:*", + "expires_in": 3600 + })) + .unwrap(), + ) + .unwrap(), + ) + .await; + } +} + +fn state_from_return_cookie(response: &axum_test::TestResponse, prefix: &str) -> SmolStr { + response + .headers() + .get_all(http::header::SET_COOKIE) + .iter() + .filter_map(|value| value.to_str().ok()) + .find_map(|cookie| { + let name = cookie.split_once('=')?.0; + let encoded = name.strip_prefix(prefix)?; + let bytes = URL_SAFE_NO_PAD.decode(encoded).ok()?; + String::from_utf8(bytes).ok().map(SmolStr::from) + }) + .expect("state-keyed return cookie") +} + +#[tokio::test] +async fn start_auth_return_to_callback_redirects_back_and_cookie_states_do_not_conflict() { + let (state, client) = app_state(); + let app = routes::().with_state(state.clone()); + let server = TestServer::builder().save_cookies().build(app).unwrap(); + + queue_par_response(&client, "urn:par:first").await; + let first = server + .get("/oauth/start?identifier=alice.bsky.social&return_to=/first") + .await; + first.assert_status(StatusCode::TEMPORARY_REDIRECT); + let first_state = state_from_return_cookie(&first, state.config.return_cookie_prefix.as_str()); + + queue_par_response(&client, "urn:par:second").await; + let second = server + .get("/oauth/start?identifier=alice.bsky.social&return_to=/second") + .await; + second.assert_status(StatusCode::TEMPORARY_REDIRECT); + let second_state = + state_from_return_cookie(&second, state.config.return_cookie_prefix.as_str()); + assert_ne!(first_state, second_state); + + queue_token_response(&client, "atk-second").await; + let callback = server + .get(&format!( + "/oauth/callback?code=abc&state={}&iss=https%3A%2F%2Fissuer", + second_state + )) + .await; + callback.assert_status(StatusCode::SEE_OTHER); + assert_eq!(callback.header("location").to_str().unwrap(), "/second"); +} + +#[tokio::test] +async fn logout_deletes_session_and_clears_cookie() { + let (state, _) = app_state(); + let data = session_data("logout-session"); + let key = SessionKey::new(data.account_did.clone(), data.session_id.clone()); + state + .oauth + .registry + .store + .upsert_session(data) + .await + .unwrap(); + let app = Router::new() + .route("/protected", get(strict_handler)) + .route( + "/oauth/logout", + post(jacquard_axum::oauth::logout_handler::), + ) + .route( + "/issue", + get({ + let config = state.config.clone(); + let key = key.clone(); + move |jar| { + let config = config.clone(); + let key = key.clone(); + async move { set_session_cookie(jar, &config, &key).unwrap() } + } + }), + ) + .with_state(state.clone()); + let server = TestServer::builder().save_cookies().build(app).unwrap(); + + server.get("/issue").await.assert_status_ok(); + server.get("/protected").await.assert_status_ok(); + server + .post("/oauth/logout") + .await + .assert_status(StatusCode::SEE_OTHER); + assert!( + state + .oauth + .registry + .store + .get_session(&key.did, key.session_id.as_str()) + .await + .unwrap() + .is_none() + ); + server.get("/protected").await.assert_status_unauthorized(); +} diff --git a/crates/jacquard-oauth/src/authstore.rs b/crates/jacquard-oauth/src/authstore.rs index 43b475cc7..78ca59e42 100644 --- a/crates/jacquard-oauth/src/authstore.rs +++ b/crates/jacquard-oauth/src/authstore.rs @@ -331,7 +331,6 @@ mod tests { token_type: OAuthTokenType::DPoP, expires_at: None, }, - #[cfg(feature = "scope-check")] resolved_scopes: None, } } diff --git a/crates/jacquard-oauth/src/client.rs b/crates/jacquard-oauth/src/client.rs index f4e7496ab..8333f21a8 100644 --- a/crates/jacquard-oauth/src/client.rs +++ b/crates/jacquard-oauth/src/client.rs @@ -366,7 +366,6 @@ where .unwrap_or_default(), }, token_set, - #[cfg(feature = "scope-check")] resolved_scopes: None, }) } diff --git a/crates/jacquard-oauth/src/keyset.rs b/crates/jacquard-oauth/src/keyset.rs index aff7470da..509c06f89 100644 --- a/crates/jacquard-oauth/src/keyset.rs +++ b/crates/jacquard-oauth/src/keyset.rs @@ -2,8 +2,9 @@ use crate::jose::jws::RegisteredHeader; use crate::jose::jwt::Claims; use crate::jose::signing; use jose_jwa::{Algorithm, Signing}; -use jose_jwk::{Class, EcCurves, OkpCurves, crypto}; +use jose_jwk::{Class, EcCurves, OkpCurves, Parameters, crypto}; use jose_jwk::{Jwk, JwkSet, Key}; +use serde::{Deserialize, Serialize}; use smol_str::{SmolStr, ToSmolStr}; use std::collections::HashSet; use thiserror::Error; @@ -62,10 +63,34 @@ const PREFERRED_SIGNING_ALGORITHMS: [Signing; 4] = [ /// /// Key selection follows [`PREFERRED_SIGNING_ALGORITHMS`] when multiple keys match. /// Supported algorithms: EdDSA (Ed25519), ES256K (secp256k1), ES256 (P-256), ES384 (P-384). -#[derive(Clone, Debug, Default, PartialEq, Eq)] +#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize, Deserialize)] +#[serde(try_from = "Vec", into = "Vec")] pub struct Keyset(Vec); impl Keyset { + /// Generate a single-key confidential-client keyset for the first supported algorithm. + /// + /// The generated key includes the supplied `kid` and is validated through the + /// same constructor used for externally supplied keys. This is intended for + /// server-side OAuth clients that authenticate with `private_key_jwt` and + /// serve the public half through their OAuth client metadata. + pub fn generate(kid: impl Into, allowed_algos: &[impl AsRef]) -> Result { + let key = + crate::utils::generate_key(allowed_algos).ok_or_else(|| Error::NotFound(Vec::new()))?; + Self::try_from(vec![Jwk { + key, + prm: Parameters { + kid: Some(kid.into()), + ..Default::default() + }, + }]) + } + + /// Generate a single-key ES256 confidential-client keyset. + pub fn generate_es256(kid: impl Into) -> Result { + Self::generate(kid, &["ES256"]) + } + /// Returns a [`JwkSet`] containing the public halves of all keys in this keyset. pub fn public_jwks(&self) -> JwkSet { let mut keys = Vec::with_capacity(self.0.len()); @@ -212,6 +237,12 @@ pub fn parse_signing_alg(s: &str) -> Option { } } +impl From for Vec { + fn from(value: Keyset) -> Self { + value.0 + } +} + impl TryFrom> for Keyset { type Error = Error; diff --git a/crates/jacquard-oauth/src/request.rs b/crates/jacquard-oauth/src/request.rs index ca697acaf..718008e38 100644 --- a/crates/jacquard-oauth/src/request.rs +++ b/crates/jacquard-oauth/src/request.rs @@ -1095,7 +1095,6 @@ mod tests { token_type: crate::types::OAuthTokenType::DPoP, expires_at: None, }, - #[cfg(feature = "scope-check")] resolved_scopes: None, }; let err = super::refresh(&client, session, &meta).await.unwrap_err(); diff --git a/crates/jacquard-oauth/src/session.rs b/crates/jacquard-oauth/src/session.rs index a35d4742c..c046aa347 100644 --- a/crates/jacquard-oauth/src/session.rs +++ b/crates/jacquard-oauth/src/session.rs @@ -88,10 +88,11 @@ pub struct ClientSessionData { pub token_set: TokenSet, /// Fully expanded scopes with include scopes resolved. - /// Populated eagerly at session creation when `scope-check` is enabled. - /// `None` when `scope-check` is disabled or no include scopes are present. - #[cfg(feature = "scope-check")] - #[serde(skip)] + /// + /// This is populated eagerly at session creation when `scope-check` is enabled. + /// It is `None` when scope checking is not enabled, when no eager resolution was + /// performed, or when reading older persisted sessions that predate this field. + #[serde(default, skip_serializing_if = "Option::is_none")] pub resolved_scopes: Option>>, } @@ -102,7 +103,6 @@ where type Output = ClientSessionData; fn into_static(self) -> Self::Output { - #[cfg(feature = "scope-check")] let resolved_scopes = self.resolved_scopes; ClientSessionData { @@ -117,7 +117,6 @@ where account_did: self.account_did.into_static(), session_id: self.session_id.into_static(), host_url: self.host_url.clone(), - #[cfg(feature = "scope-check")] resolved_scopes, } } @@ -464,7 +463,7 @@ where pub fn new(store: S, client: Arc, client_data: ClientData) -> Self { let store = Arc::new(store); Self { - store: Arc::clone(&store), + store, client, client_data, pending: DashMap::new(), diff --git a/crates/jacquard/src/client/token.rs b/crates/jacquard/src/client/token.rs index 00aa54330..7965bad61 100644 --- a/crates/jacquard/src/client/token.rs +++ b/crates/jacquard/src/client/token.rs @@ -145,7 +145,6 @@ impl From for ClientSessionData { token_type: session.token_type, expires_at: session.expires_at, }, - #[cfg(feature = "scope-check")] resolved_scopes: None, } } @@ -557,7 +556,6 @@ mod tests { token_type: OAuthTokenType::DPoP, expires_at: None, }, - #[cfg(feature = "scope-check")] resolved_scopes: None, } } diff --git a/crates/jacquard/tests/agent.rs b/crates/jacquard/tests/agent.rs index 8c160eeb1..4f079d3b0 100644 --- a/crates/jacquard/tests/agent.rs +++ b/crates/jacquard/tests/agent.rs @@ -58,11 +58,13 @@ impl IdentityResolver for MockClient { _did: &Did, ) -> std::result::Result { let doc = serde_json::json!({ + "@context": ["https://www.w3.org/ns/did/v1"], "id": "did:plc:alice", + "alsoKnownAs": ["at://alice.bsky.social"], "service": [{ - "id": "#pds", + "id": "#atproto_pds", "type": "AtprotoPersonalDataServer", - "serviceEndpoint": "https://pds" + "serviceEndpoint": "https://pds.example.com" }] }); Ok(DidDocResponse { @@ -113,7 +115,7 @@ async fn agent_delegates_to_session_and_refreshes() { let info = agent.info().await.expect("session info"); assert_eq!(info.0.as_str(), "did:plc:alice"); assert_eq!(info.1.as_ref().unwrap().as_str(), "session"); - assert_eq!(agent.endpoint().await.as_str(), "https://pds"); + assert_eq!(agent.endpoint().await.as_str(), "https://pds.example.com"); // Queue a refresh response and call agent.refresh(); Authorization header must use refresh token client diff --git a/crates/jacquard/tests/credential_session.rs b/crates/jacquard/tests/credential_session.rs index bfd562b64..6ff8e1def 100644 --- a/crates/jacquard/tests/credential_session.rs +++ b/crates/jacquard/tests/credential_session.rs @@ -83,11 +83,13 @@ impl IdentityResolver for MockClient { *self.did_doc_calls.write().await += 1; assert_eq!(did.as_str(), "did:plc:alice"); let doc = serde_json::json!({ + "@context": ["https://www.w3.org/ns/did/v1"], "id": "did:plc:alice", + "alsoKnownAs": ["at://alice.bsky.social"], "service": [{ - "id": "#pds", + "id": "#atproto_pds", "type": "AtprotoPersonalDataServer", - "serviceEndpoint": "https://pds" + "serviceEndpoint": "https://pds.example.com" }] }); Ok(DidDocResponse { @@ -374,7 +376,7 @@ async fn credential_login_and_auto_refresh() { .expect("login ok"); // Endpoint switches to PDS - assert_eq!(session.endpoint().await.as_str(), "https://pds"); + assert_eq!(session.endpoint().await.as_str(), "https://pds.example.com"); // Send a request that will first 401 (ExpiredToken), then refresh, then succeed let resp = session diff --git a/crates/jacquard/tests/oauth_auto_refresh.rs b/crates/jacquard/tests/oauth_auto_refresh.rs index 4c2dfb909..e82b4c80b 100644 --- a/crates/jacquard/tests/oauth_auto_refresh.rs +++ b/crates/jacquard/tests/oauth_auto_refresh.rs @@ -68,11 +68,13 @@ impl jacquard::identity::resolver::IdentityResolver for MockClient { jacquard::identity::resolver::IdentityError, > { let doc = serde_json::json!({ + "@context": ["https://www.w3.org/ns/did/v1"], "id": "did:plc:alice", + "alsoKnownAs": ["at://alice.bsky.social"], "service": [{ - "id": "#pds", + "id": "#atproto_pds", "type": "AtprotoPersonalDataServer", - "serviceEndpoint": "https://pds" + "serviceEndpoint": "https://pds.example.com" }] }); Ok(jacquard::identity::resolver::DidDocResponse { @@ -122,7 +124,7 @@ impl OAuthResolver for MockClient { _sub: &Did, ) -> Result, jacquard_oauth::resolver::ResolverError> { - Ok(jacquard::deps::fluent_uri::Uri::parse("https://pds") + Ok(jacquard::deps::fluent_uri::Uri::parse("https://pds.example.com") .unwrap() .to_owned()) } @@ -220,7 +222,7 @@ async fn oauth_xrpc_invalid_token_triggers_refresh_and_retries() { let session_data = ClientSessionData { account_did: Did::new_static("did:plc:alice").unwrap(), session_id: SmolStr::from("state"), - host_url: Uri::parse("https://pds").expect("valid uri").to_owned(), + host_url: Uri::parse("https://pds.example.com").expect("valid uri").to_owned(), authserver_url: SmolStr::new_static("https://issuer"), authserver_token_endpoint: SmolStr::from("https://issuer/token"), authserver_revocation_endpoint: None, @@ -233,14 +235,13 @@ async fn oauth_xrpc_invalid_token_triggers_refresh_and_retries() { token_set: TokenSet { iss: SmolStr::from("https://issuer"), sub: Did::new_static("did:plc:alice").unwrap(), - aud: SmolStr::from("https://pds"), + aud: SmolStr::from("https://pds.example.com"), scope: None, refresh_token: Some(SmolStr::from("rt1")), access_token: SmolStr::from("atk1"), token_type: OAuthTokenType::DPoP, expires_at: None, }, - #[cfg(feature = "scope-check")] resolved_scopes: None, } .into_static(); @@ -250,7 +251,7 @@ async fn oauth_xrpc_invalid_token_triggers_refresh_and_retries() { let data_store = ClientSessionData { account_did: Did::new_static("did:plc:alice").unwrap(), session_id: SmolStr::from("state"), - host_url: Uri::parse("https://pds").expect("valid uri").to_owned(), + host_url: Uri::parse("https://pds.example.com").expect("valid uri").to_owned(), authserver_url: SmolStr::new_static("https://issuer"), authserver_token_endpoint: SmolStr::from("https://issuer/token"), authserver_revocation_endpoint: None, @@ -263,14 +264,13 @@ async fn oauth_xrpc_invalid_token_triggers_refresh_and_retries() { token_set: TokenSet { iss: SmolStr::from("https://issuer"), sub: Did::new_static("did:plc:alice").unwrap(), - aud: SmolStr::from("https://pds"), + aud: SmolStr::from("https://pds.example.com"), scope: None, refresh_token: Some(SmolStr::from("rt1")), access_token: SmolStr::from("atk1"), token_type: OAuthTokenType::DPoP, expires_at: None, }, - #[cfg(feature = "scope-check")] resolved_scopes: None, } .into_static(); @@ -356,7 +356,7 @@ async fn oauth_xrpc_invalid_token_body_triggers_refresh_and_retries() { let session_data = ClientSessionData { account_did: Did::new_static("did:plc:alice").unwrap(), session_id: SmolStr::new_static("state"), - host_url: Uri::parse("https://pds").expect("valid uri").to_owned(), + host_url: Uri::parse("https://pds.example.com").expect("valid uri").to_owned(), authserver_url: SmolStr::new_static("https://issuer"), authserver_token_endpoint: SmolStr::from("https://issuer/token"), authserver_revocation_endpoint: None, @@ -369,14 +369,13 @@ async fn oauth_xrpc_invalid_token_body_triggers_refresh_and_retries() { token_set: TokenSet { iss: SmolStr::from("https://issuer"), sub: Did::new_static("did:plc:alice").unwrap(), - aud: SmolStr::from("https://pds"), + aud: SmolStr::from("https://pds.example.com"), scope: None, refresh_token: Some(SmolStr::from("rt1")), access_token: SmolStr::from("atk1"), token_type: OAuthTokenType::DPoP, expires_at: None, }, - #[cfg(feature = "scope-check")] resolved_scopes: None, } .into_static(); diff --git a/crates/jacquard/tests/oauth_flow.rs b/crates/jacquard/tests/oauth_flow.rs index f5a41a1ef..cd8388574 100644 --- a/crates/jacquard/tests/oauth_flow.rs +++ b/crates/jacquard/tests/oauth_flow.rs @@ -40,6 +40,19 @@ impl HttpClient for MockClient { } } +fn alice_did_document_json() -> serde_json::Value { + serde_json::json!({ + "@context": ["https://www.w3.org/ns/did/v1"], + "id": "did:plc:alice", + "alsoKnownAs": ["at://alice.bsky.social"], + "service": [{ + "id": "#atproto_pds", + "type": "AtprotoPersonalDataServer", + "serviceEndpoint": "https://pds.example.com" + }] + }) +} + impl jacquard::identity::resolver::IdentityResolver for MockClient { fn options(&self) -> &jacquard::identity::resolver::ResolverOptions { use std::sync::LazyLock; @@ -61,14 +74,7 @@ impl jacquard::identity::resolver::IdentityResolver for MockClient { jacquard::identity::resolver::DidDocResponse, jacquard::identity::resolver::IdentityError, > { - let doc = serde_json::json!({ - "id": "did:plc:alice", - "service": [{ - "id": "#pds", - "type": "AtprotoPersonalDataServer", - "serviceEndpoint": "https://pds" - }] - }); + let doc = alice_did_document_json(); Ok(jacquard::identity::resolver::DidDocResponse { buffer: Bytes::from(serde_json::to_vec(&doc).unwrap()), status: StatusCode::OK, @@ -97,15 +103,7 @@ impl OAuthResolver for MockClient { md.token_endpoint_auth_methods_supported = Some(vec![SmolStr::from("none")]); md.dpop_signing_alg_values_supported = Some(vec![SmolStr::from("ES256")]); - // Simple DID doc pointing to https://pds - let doc = serde_json::json!({ - "id": "did:plc:alice", - "service": [{ - "id": "#pds", - "type": "AtprotoPersonalDataServer", - "serviceEndpoint": "https://pds" - }] - }); + let doc = alice_did_document_json(); let buf = Bytes::from(serde_json::to_vec(&doc).unwrap()); let did_doc: jacquard_common::types::did_doc::DidDocument = serde_json::from_slice(&buf).unwrap(); @@ -207,7 +205,7 @@ async fn oauth_end_to_end_mock_flow() { "refresh_token": "rt1", "sub": "did:plc:alice", "iss": "https://issuer", - "aud": "https://pds", + "aud": "https://pds.example.com", "scope": "atproto rpc:*", "expires_in": 3600 })) diff --git a/crates/jacquard/tests/restore_pds_cache.rs b/crates/jacquard/tests/restore_pds_cache.rs index aa1e14786..82c3c4624 100644 --- a/crates/jacquard/tests/restore_pds_cache.rs +++ b/crates/jacquard/tests/restore_pds_cache.rs @@ -57,11 +57,13 @@ impl IdentityResolver for MockResolver { ) -> std::result::Result { *self.did_doc_calls.write().await += 1; let doc = serde_json::json!({ + "@context": ["https://www.w3.org/ns/did/v1"], "id": "did:plc:alice", + "alsoKnownAs": ["at://alice.bsky.social"], "service": [{ - "id": "#pds", + "id": "#atproto_pds", "type": "AtprotoPersonalDataServer", - "serviceEndpoint": "https://pds-resolved" + "serviceEndpoint": "https://pds-resolved.example.com" }] }); Ok(DidDocResponse { diff --git a/crates/jacquard/tests/scope_check.rs b/crates/jacquard/tests/scope_check.rs index 21d011331..e459d9be7 100644 --- a/crates/jacquard/tests/scope_check.rs +++ b/crates/jacquard/tests/scope_check.rs @@ -75,11 +75,13 @@ impl jacquard::identity::resolver::IdentityResolver for MockClient { jacquard::identity::resolver::IdentityError, > { let doc = serde_json::json!({ + "@context": ["https://www.w3.org/ns/did/v1"], "id": "did:plc:alice", + "alsoKnownAs": ["at://alice.bsky.social"], "service": [{ - "id": "#pds", + "id": "#atproto_pds", "type": "AtprotoPersonalDataServer", - "serviceEndpoint": "https://pds" + "serviceEndpoint": "https://pds.example.com" }] }); Ok(jacquard::identity::resolver::DidDocResponse { @@ -127,7 +129,7 @@ impl OAuthResolver for MockClient { _sub: &Did, ) -> Result, jacquard_oauth::resolver::ResolverError> { - Ok(jacquard::deps::fluent_uri::Uri::parse("https://pds") + Ok(jacquard::deps::fluent_uri::Uri::parse("https://pds.example.com") .unwrap() .to_owned()) } @@ -154,7 +156,7 @@ fn create_session_data(resolved_scopes: Option>>) -> ClientSe ClientSessionData { account_did: Did::new_static("did:plc:alice").unwrap(), session_id: SmolStr::from("state"), - host_url: Uri::parse("https://pds").expect("valid uri").to_owned(), + host_url: Uri::parse("https://pds.example.com").expect("valid uri").to_owned(), authserver_url: SmolStr::new_static("https://issuer"), authserver_token_endpoint: SmolStr::from("https://issuer/token"), authserver_revocation_endpoint: None, @@ -167,14 +169,13 @@ fn create_session_data(resolved_scopes: Option>>) -> ClientSe token_set: TokenSet { iss: SmolStr::from("https://issuer"), sub: Did::new_static("did:plc:alice").unwrap(), - aud: SmolStr::from("https://pds"), + aud: SmolStr::from("https://pds.example.com"), scope: None, refresh_token: Some(SmolStr::from("rt1")), access_token: SmolStr::from("atk1"), token_type: OAuthTokenType::DPoP, expires_at: None, }, - #[cfg(feature = "scope-check")] resolved_scopes, } } diff --git a/examples/axum_oauth_session.rs b/examples/axum_oauth_session.rs new file mode 100644 index 000000000..0d3b6dc62 --- /dev/null +++ b/examples/axum_oauth_session.rs @@ -0,0 +1,386 @@ +//! Hosted, server-rendered Axum OAuth example. +//! +//! This example demonstrates a web OAuth client whose `client_id` is the public +//! metadata URL served by the same Axum application: +//! +//! ```text +//! https://example.com/oauth-client-metadata.json +//! ``` +//! +//! Run it with a public HTTPS origin that reaches this server. For local +//! development, point a domain you control at your machine with something like +//! Tailscale Funnel or a Cloudflare Tunnel. Alternatively, run it on a server +//! with a publicly routable IP address and configured A and/or AAAA records. +//! +//! ```text +//! cargo run -p jacquard-axum --example axum_oauth_server -- \ +//! --base-url https://example.com \ +//! --listen 127.0.0.1:3000 +//! ``` +//! +//! The example stores generated development secrets under `--data-dir`: the +//! private-cookie key, the OAuth confidential-client keyset, and the file-backed +//! OAuth session store. The default lives under `/tmp` so running the example +//! does not contaminate the repository checkout. For longer-lived local state, +//! use an os/environment appropriate config/state directory. A real +//! deployment should store cookie keys and OAuth signing keys in appropriately +//! protected secret storage, such as an OS keyring or deployment secret manager, +//! with permissions, backups, rotation, access controls, and operational practices +//! chosen for that environment, and would also use a database-backed impl of `ClientAuthStore`. + +use std::{fs, net::SocketAddr, path::PathBuf, sync::Arc}; + +use axum::{ + Json, Router, + extract::{FromRef, Query}, + http::StatusCode, + response::{Html, IntoResponse}, + routing::get, +}; +use axum_extra::extract::cookie::Key; +use clap::Parser; +use html_escape::{encode_double_quoted_attribute, encode_text}; +use jacquard::{ + api::app_bsky::feed::get_timeline::GetTimeline, + client::{Agent, FileAuthStore}, + common::deps::{fluent_uri::Uri, smol_str::SmolStr}, + identity::JacquardResolver, + oauth::{ + atproto::{AtprotoClientMetadata, GrantType}, + client::OAuthClient, + keyset::Keyset, + scopes::Scopes, + session::ClientData, + }, + xrpc::XrpcClient, +}; +use jacquard_axum::oauth::{ + BrowserOAuthSession, ExtractOAuthSession, OAuthWebConfig, OAuthWebState, routes as oauth_routes, +}; +use miette::{Context, IntoDiagnostic, Result, bail, miette}; +use rand::{RngCore, rngs::OsRng}; +use serde::Deserialize; +use tracing_subscriber::EnvFilter; + +#[derive(Parser, Debug)] +struct Args { + /// Public HTTPS origin for this hosted OAuth client, e.g. `https://example.com`. + /// + /// This must be externally reachable by users' PDSes. For local development, + /// use a tunnel such as Tailscale Funnel or Cloudflare Tunnel, or run on a + /// server with public DNS A and/or AAAA records. + #[arg(long)] + base_url: String, + + /// Directory for this example's generated development state. + /// + /// This directory contains secret key material. The default is under `/tmp` + /// to avoid writing into the repository. For longer-lived local state, prefer + /// an XDG state/config directory or the platform equivalent; production apps + /// should use secure secret storage or an OS keyring. + #[arg(long, default_value = "/tmp/jacquard-axum-oauth-example")] + data_dir: PathBuf, + + /// Socket address to bind locally. + #[arg(long, default_value = "127.0.0.1:3000")] + listen: SocketAddr, +} + +#[derive(Clone)] +struct AppState { + oauth: Arc>, + oauth_config: OAuthWebConfig, + cookie_key: Key, +} + +impl OAuthWebState for AppState { + fn oauth_client(&self) -> &OAuthClient { + self.oauth.as_ref() + } +} + +impl FromRef for OAuthWebConfig { + fn from_ref(input: &AppState) -> Self { + input.oauth_config.clone() + } +} + +impl FromRef for Key { + fn from_ref(input: &AppState) -> Self { + input.cookie_key.clone() + } +} + +#[derive(Debug, Deserialize)] +struct LoginQuery { + #[serde(default)] + return_to: Option, +} + +async fn login_page(Query(query): Query) -> Html { + let return_to = query.return_to.unwrap_or_else(|| "/timeline".to_owned()); + Html(format!( + r#" + + + + + Sign in with AT Protocol + + +
+

Sign in with AT Protocol

+

Enter your handle, DID, or PDS URL. The server will resolve it and start + the OAuth flow for the matching PDS.

+
+

+
+ +

+ + +
+
+ +"#, + encode_double_quoted_attribute(&return_to), + )) +} + +async fn timeline( + BrowserOAuthSession(session): BrowserOAuthSession, +) -> Result, AppError> { + let agent: Agent<_> = Agent::from(session); + let response = agent + .send(GetTimeline::new().limit(10).build()) + .await + .map_err(AppError::from_display)?; + let timeline = response.into_output().map_err(AppError::from_display)?; + + let mut html = String::from( + r#" + + + + + Timeline + + +
+

Timeline

+

View strict extractor session JSON

+
    +"#, + ); + for item in timeline.feed { + html.push_str("
  1. "); + html.push_str(&encode_text(item.post.author.handle.as_str())); + html.push_str(": "); + html.push_str(&encode_text( + &serde_json::to_string(&item.post.record).unwrap_or_default(), + )); + html.push_str("
  2. \n"); + } + html.push_str( + r#"
+
+ +
+
+ +"#, + ); + Ok(Html(html)) +} + +async fn strict_session_json( + ExtractOAuthSession(session): ExtractOAuthSession, +) -> Json { + let (did, session_id) = session.session_info().await; + Json(serde_json::json!({ "did": did, "session_id": session_id })) +} + +#[tokio::main] +async fn main() -> Result<()> { + tracing_subscriber::fmt() + .with_env_filter(EnvFilter::from_env("JACQUARD_AXUM_OAUTH_LOG")) + .init(); + + let args = Args::parse(); + let paths = ExamplePaths::new(args.data_dir); + paths.create_data_dir()?; + + let client_data = ClientData::new( + Some(load_or_generate_keyset(&paths.keyset)?), + hosted_client_metadata(&args.base_url)?, + ); + + let state = AppState { + oauth: Arc::new(OAuthClient::new( + FileAuthStore::new(paths.sessions.to_string_lossy().into_owned()), + client_data, + )), + oauth_config: OAuthWebConfig::default(), + cookie_key: load_or_generate_cookie_key(&paths.cookie_key)?, + }; + + let app = Router::new() + .route("/", get(|| async { Html(HOME_PAGE) })) + .route("/oauth/login", get(login_page)) + .route("/timeline", get(timeline)) + .route("/api/session", get(strict_session_json)) + .merge(oauth_routes::()) + .with_state(state); + + let listener = tokio::net::TcpListener::bind(args.listen) + .await + .into_diagnostic()?; + axum::serve(listener, app).await.into_diagnostic()?; + Ok(()) +} + +const HOME_PAGE: &str = r#" + + + + + Jacquard Axum OAuth example + + +
+

Jacquard Axum OAuth example

+

This is a hosted OAuth client example. Its OAuth client metadata is served + from /oauth-client-metadata.json by the same Axum state that + starts OAuth, handles callbacks, restores sessions, and logs out.

+

Open the authenticated timeline

+
+ +"#; + +fn hosted_client_metadata(base_url: &str) -> Result> { + let base_url = base_url.trim_end_matches('/'); + let client_uri = Uri::parse(base_url.to_owned()) + .map_err(|(err, _)| miette!("invalid --base-url `{}`: {err}", base_url))?; + if client_uri.scheme().as_str() != "https" { + bail!("--base-url must be an externally reachable https:// origin"); + } + if client_uri.path().as_str() != "/" + || client_uri.query().is_some() + || client_uri.fragment().is_some() + { + bail!("--base-url must be an origin only, such as https://example.com"); + } + + let client_id = Uri::parse(format!("{base_url}/oauth-client-metadata.json")) + .map_err(|(err, _)| miette!("invalid client metadata URL: {err}"))?; + let redirect_uri = Uri::parse(format!("{base_url}/oauth/callback")) + .map_err(|(err, _)| miette!("invalid OAuth callback URL: {err}"))?; + let scopes = Scopes::new(SmolStr::new_static("atproto rpc:*")) + .map_err(|err| miette!("invalid OAuth scopes: {err}"))?; + + Ok(AtprotoClientMetadata { + client_id, + client_uri: Some(client_uri), + redirect_uris: vec![redirect_uri], + grant_types: vec![GrantType::AuthorizationCode, GrantType::RefreshToken], + scopes, + jwks_uri: None, + client_name: Some(SmolStr::new_static("Jacquard Axum OAuth example")), + logo_uri: None, + tos_uri: None, + privacy_policy_uri: None, + }) +} + +#[derive(Debug)] +struct ExamplePaths { + data_dir: PathBuf, + cookie_key: PathBuf, + keyset: PathBuf, + sessions: PathBuf, +} + +impl ExamplePaths { + fn new(data_dir: PathBuf) -> Self { + Self { + cookie_key: data_dir.join("private-cookie.key"), + keyset: data_dir.join("oauth-client-keyset.json"), + sessions: data_dir.join("oauth-sessions.json"), + data_dir, + } + } + + fn create_data_dir(&self) -> Result<()> { + fs::create_dir_all(&self.data_dir) + .into_diagnostic() + .wrap_err_with(|| format!("failed to create {}", self.data_dir.display())) + } +} + +fn load_or_generate_cookie_key(path: &PathBuf) -> Result { + match fs::read(path) { + Ok(bytes) if bytes.len() == 64 => Ok(Key::from(&bytes)), + Ok(_) => bail!("{} must contain exactly 64 bytes", path.display()), + Err(err) if err.kind() == std::io::ErrorKind::NotFound => { + let mut bytes = [0_u8; 64]; + OsRng.fill_bytes(&mut bytes); + write_secret(path, &bytes)?; + Ok(Key::from(&bytes)) + } + Err(err) => Err(err) + .into_diagnostic() + .wrap_err_with(|| format!("failed to read {}", path.display())), + } +} + +fn load_or_generate_keyset(path: &PathBuf) -> Result { + match fs::read(path) { + Ok(bytes) => serde_json::from_slice(&bytes) + .into_diagnostic() + .wrap_err_with(|| format!("failed to parse {}", path.display())), + Err(err) if err.kind() == std::io::ErrorKind::NotFound => { + let keyset = Keyset::generate_es256("jacquard-axum-example") + .map_err(|err| miette!("failed to generate OAuth client keyset: {err}"))?; + let bytes = serde_json::to_vec_pretty(&keyset).into_diagnostic()?; + write_secret(path, &bytes)?; + Ok(keyset) + } + Err(err) => Err(err) + .into_diagnostic() + .wrap_err_with(|| format!("failed to read {}", path.display())), + } +} + +fn write_secret(path: &PathBuf, bytes: &[u8]) -> Result<()> { + fs::write(path, bytes) + .into_diagnostic() + .wrap_err_with(|| format!("failed to write {}", path.display()))?; + #[cfg(unix)] + { + use std::os::unix::fs::PermissionsExt; + fs::set_permissions(path, fs::Permissions::from_mode(0o600)) + .into_diagnostic() + .wrap_err_with(|| format!("failed to restrict permissions on {}", path.display()))?; + } + Ok(()) +} + +#[derive(Debug)] +struct AppError(String); + +impl AppError { + fn from_display(err: impl std::fmt::Display) -> Self { + Self(err.to_string()) + } +} + +impl IntoResponse for AppError { + fn into_response(self) -> axum::response::Response { + ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(serde_json::json!({ "error": "InternalServerError", "message": self.0 })), + ) + .into_response() + } +} -- 2.51.2