From 2c704a763e3e249c2bf09f6d8918ec9164b0d6f3 Mon Sep 17 00:00:00 2001 From: Alex Bates Date: Tue, 18 Aug 2026 12:24:31 +0100 Subject: [PATCH] clean up route code using FromRequestParts --- src/atproto/lexicon/mod.rs | 86 ++++++++++++++++ src/atproto/lexicon/mod_listing.rs | 153 ++++++++++++++++------------- src/atproto/mod.rs | 6 +- src/atproto/tap.rs | 6 +- src/config.rs | 2 +- src/mods.rs | 133 ++++++++++--------------- src/oauth/session.rs | 31 ++++++ 7 files changed, 255 insertions(+), 162 deletions(-) diff --git a/src/atproto/lexicon/mod.rs b/src/atproto/lexicon/mod.rs index 29eca64..2897d7a 100644 --- a/src/atproto/lexicon/mod.rs +++ b/src/atproto/lexicon/mod.rs @@ -13,6 +13,92 @@ pub mod mod_listing; pub use mod_listing::ModListing; use anyhow::{bail, Result}; +use serde_json::Value; +use sqlx::SqlitePool; + +use crate::atproto::id::{Did, Handle, Rkey}; +use crate::atproto::pds; +use crate::error::AppError; +use crate::oauth::session::SessionCookie; +use crate::state::AppState; + +/// A lexicon record this appview writes: its own NSID, wire shape, and local +/// cache table, plus the create/update orchestration every lexicon shares. +pub trait Record: Sized { + /// The lexicon NSID, and the record's collection name in a repo. + const NSID: &'static str; + + /// Check the constraints the lexicon declares but JSON cannot express. + fn validate(&self) -> Result<()>; + /// Serialize for the wire, with `$type` set. + fn to_value(&self) -> Result; + /// Parse a record as it appears on the wire. Does not validate - callers + /// that need a trustworthy record (as opposed to e.g. just inspecting + /// one) should call `validate` themselves. + fn from_value(value: &Value) -> Result; + + /// Mirror this record into the local cache, keyed on `(did, rkey)`. + async fn save_local(&self, db: &SqlitePool, did: &Did, rkey: &Rkey) -> Result<()>; + /// Remove this record's local mirror. + async fn delete_local(db: &SqlitePool, did: &Did, rkey: &Rkey) -> Result<()>; + + /// Validate, write to the author's PDS, then mirror into the local + /// cache. Mints a fresh rkey. + async fn create(&self, state: &AppState, session: &SessionCookie) -> Result { + self.validate().map_err(WriteError::Invalid)?; + let rkey = Rkey::new(); + let value = self.to_value().map_err(WriteError::Failed)?; + pds::put(state, session, Self::NSID, rkey.as_str(), value) + .await + .map_err(WriteError::Failed)?; + self.save_local(&state.db, &session.did, &rkey) + .await + .map_err(WriteError::Failed)?; + Ok(rkey) + } + + /// Validate, write to the author's PDS, then mirror into the local + /// cache, replacing whatever was previously stored at `rkey`. + async fn update( + &self, + state: &AppState, + session: &SessionCookie, + rkey: &Rkey, + ) -> Result<(), WriteError> { + self.validate().map_err(WriteError::Invalid)?; + let value = self.to_value().map_err(WriteError::Failed)?; + pds::put(state, session, Self::NSID, rkey.as_str(), value) + .await + .map_err(WriteError::Failed)?; + self.save_local(&state.db, &session.did, rkey) + .await + .map_err(WriteError::Failed) + } +} + +/// Either the record was invalid (the caller's fault) or writing it failed +/// (ours) - kept distinct so a route can tell a 400 from a 500 apart. +pub enum WriteError { + Invalid(anyhow::Error), + Failed(anyhow::Error), +} + +impl From for AppError { + fn from(err: WriteError) -> Self { + match err { + WriteError::Invalid(e) => AppError::BadRequest(e.to_string()), + WriteError::Failed(e) => AppError::Internal(e), + } + } +} + +/// A record resolved from a route, along with the identity of its owner. +pub struct Found { + pub did: Did, + pub handle: Handle, + pub rkey: Rkey, + pub record: T, +} /// Lexicon `maxLength` counts UTF-8 bytes. fn max_len(field: &str, value: &str, max: usize) -> Result<()> { diff --git a/src/atproto/lexicon/mod_listing.rs b/src/atproto/lexicon/mod_listing.rs index c4a8749..b13d6e7 100644 --- a/src/atproto/lexicon/mod_listing.rs +++ b/src/atproto/lexicon/mod_listing.rs @@ -13,7 +13,7 @@ use serde::{Deserialize, Serialize}; use serde_json::Value; use sqlx::{Row, SqlitePool}; -use super::{max_graphemes, max_len}; +use super::{max_graphemes, max_len, Record}; use crate::atproto::id::{Did, Rkey}; /// A game a mod can target, each implying exactly one platform (see @@ -161,16 +161,63 @@ pub struct MediaItem { } impl ModListing { - /// The lexicon NSID, which is also the collection name in a repo. - pub const NSID: &'static str = "dev.starhaven.mod.listing"; - /// The console this mod runs on, implied by `game`. pub fn platform(&self) -> Platform { self.game.platform() } + /// Load a listing back by its record key. + #[cfg(test)] + async fn load_local(db: &SqlitePool, did: &Did, rkey: &Rkey) -> Result> { + let row = sqlx::query( + "SELECT title, slug, description, details, game, category, + license, tags, media, created_at + FROM mod_listing WHERE author_did = ? AND rkey = ?", + ) + .bind(did.as_str()) + .bind(rkey.as_str()) + .fetch_optional(db) + .await?; + + row.map(row_to_listing).transpose() + } + + /// Resolve a pretty URL (`/@handle/mods/:slug`) to the listing it names. + /// + /// Slugs are not unique - two listings can legitimately share one - so + /// this resolves to the lowest rkey among matches, the same rule every + /// appview instance can apply to the same records and agree on. + pub async fn find_by_slug( + db: &SqlitePool, + did: &Did, + slug: &str, + ) -> Result> { + let row = sqlx::query( + "SELECT rkey, title, slug, description, details, game, + category, license, tags, media, created_at + FROM mod_listing WHERE author_did = ? AND slug = ? ORDER BY rkey LIMIT 1", + ) + .bind(did.as_str()) + .bind(slug) + .fetch_optional(db) + .await?; + + let Some(row) = row else { + return Ok(None); + }; + let rkey: Rkey = row.get("rkey"); + let listing = row_to_listing(row)?; + + Ok(Some((rkey, listing))) + } +} + +impl Record for ModListing { + /// The lexicon NSID, which is also the collection name in a repo. + const NSID: &'static str = "dev.starhaven.mod.listing"; + /// Check the constraints the lexicon declares but JSON cannot express. - pub fn validate(&self) -> Result<()> { + fn validate(&self) -> Result<()> { max_graphemes("title", &self.title, 300)?; max_graphemes("description", &self.description, 1000)?; max_graphemes("details", &self.details, 30000)?; @@ -204,17 +251,13 @@ impl ModListing { Ok(()) } - /// Parse and validate a record as it appears on the wire. - pub fn from_value(value: &Value) -> Result { - let record: Self = serde_json::from_value(value.clone())?; - record.validate()?; - Ok(record) + /// Parse a record as it appears on the wire. Does not validate. + fn from_value(value: &Value) -> Result { + Ok(serde_json::from_value(value.clone())?) } - // TODO(pds): once handlers write to the author's PDS, call this to build - // the `com.atproto.repo.createRecord` body before saving locally. - /// Serialize for `com.atproto.repo.createRecord`, which wants `$type`. - pub fn to_value(&self) -> Result { + /// Serialize for `com.atproto.repo.putRecord`, which wants `$type`. + fn to_value(&self) -> Result { let mut value = serde_json::to_value(self)?; match value.as_object_mut() { Some(object) => { @@ -231,7 +274,7 @@ impl ModListing { /// once (on create) and remember (on edit), so an edit updates the same /// row instead of colliding with, or losing to, another listing that /// happens to share a slug. - pub async fn save(&self, db: &SqlitePool, did: &Did, rkey: &Rkey) -> Result<()> { + async fn save_local(&self, db: &SqlitePool, did: &Did, rkey: &Rkey) -> Result<()> { sqlx::query( "INSERT INTO mod_listing (author_did, rkey, title, slug, description, details, game, @@ -267,52 +310,8 @@ impl ModListing { Ok(()) } - /// Load a listing back by its record key, e.g. before editing it. - pub async fn load(db: &SqlitePool, did: &Did, rkey: &Rkey) -> Result> { - let row = sqlx::query( - "SELECT title, slug, description, details, game, category, - license, tags, media, created_at - FROM mod_listing WHERE author_did = ? AND rkey = ?", - ) - .bind(did.as_str()) - .bind(rkey.as_str()) - .fetch_optional(db) - .await?; - - row.map(row_to_listing).transpose() - } - - /// Resolve a pretty URL (`/@handle/mods/:slug`) to the listing it names. - /// - /// Slugs are not unique - two listings can legitimately share one - so - /// this resolves to the lowest rkey among matches, the same rule every - /// appview instance can apply to the same records and agree on. - pub async fn find_by_slug( - db: &SqlitePool, - did: &Did, - slug: &str, - ) -> Result> { - let row = sqlx::query( - "SELECT rkey, title, slug, description, details, game, - category, license, tags, media, created_at - FROM mod_listing WHERE author_did = ? AND slug = ? ORDER BY rkey LIMIT 1", - ) - .bind(did.as_str()) - .bind(slug) - .fetch_optional(db) - .await?; - - let Some(row) = row else { - return Ok(None); - }; - let rkey: Rkey = row.get("rkey"); - let listing = row_to_listing(row)?; - - Ok(Some((rkey, listing))) - } - /// Remove a listing. - pub async fn delete(db: &SqlitePool, did: &Did, rkey: &Rkey) -> Result<()> { + async fn delete_local(db: &SqlitePool, did: &Did, rkey: &Rkey) -> Result<()> { sqlx::query("DELETE FROM mod_listing WHERE author_did = ? AND rkey = ?") .bind(did.as_str()) .bind(rkey.as_str()) @@ -374,16 +373,22 @@ mod tests { let rkey = Rkey::new(); let original = listing(); - original.save(&db, &did, &rkey).await.unwrap(); + original.save_local(&db, &did, &rkey).await.unwrap(); - let loaded = ModListing::load(&db, &did, &rkey).await.unwrap().unwrap(); + let loaded = ModListing::load_local(&db, &did, &rkey) + .await + .unwrap() + .unwrap(); assert_eq!(loaded.title, original.title); // Not carried by any bespoke column mapping - tags round-trips // because `load` reads the same JSON column `save` wrote. assert_eq!(loaded.tags, original.tags); - ModListing::delete(&db, &did, &rkey).await.unwrap(); - assert!(ModListing::load(&db, &did, &rkey).await.unwrap().is_none()); + ModListing::delete_local(&db, &did, &rkey).await.unwrap(); + assert!(ModListing::load_local(&db, &did, &rkey) + .await + .unwrap() + .is_none()); } #[tokio::test] @@ -392,12 +397,15 @@ mod tests { let did = Did::new("did:plc:example"); let rkey = Rkey::new(); - listing().save(&db, &did, &rkey).await.unwrap(); + listing().save_local(&db, &did, &rkey).await.unwrap(); let mut edited = listing(); edited.title = "Master Quest 2".to_string(); - edited.save(&db, &did, &rkey).await.unwrap(); + edited.save_local(&db, &did, &rkey).await.unwrap(); - let loaded = ModListing::load(&db, &did, &rkey).await.unwrap().unwrap(); + let loaded = ModListing::load_local(&db, &did, &rkey) + .await + .unwrap() + .unwrap(); assert_eq!(loaded.title, "Master Quest 2"); } @@ -409,12 +417,12 @@ mod tests { let did = Did::new("did:plc:example"); let first = Rkey::new(); - listing().save(&db, &did, &first).await.unwrap(); + listing().save_local(&db, &did, &first).await.unwrap(); let second = Rkey::new(); let mut other = listing(); other.title = "A Different Master Quest".to_string(); - other.save(&db, &did, &second).await.unwrap(); + other.save_local(&db, &did, &second).await.unwrap(); let (winner, _) = ModListing::find_by_slug(&db, &did, "master-quest") .await @@ -428,8 +436,11 @@ mod tests { assert_eq!(&winner, expected); // The loser is still there, reachable by its own rkey. - assert!(ModListing::load(&db, &did, &first).await.unwrap().is_some()); - assert!(ModListing::load(&db, &did, &second) + assert!(ModListing::load_local(&db, &did, &first) + .await + .unwrap() + .is_some()); + assert!(ModListing::load_local(&db, &did, &second) .await .unwrap() .is_some()); diff --git a/src/atproto/mod.rs b/src/atproto/mod.rs index 0f92b4c..13920f0 100644 --- a/src/atproto/mod.rs +++ b/src/atproto/mod.rs @@ -1,12 +1,8 @@ //! Everything that speaks atproto: identifiers, our lexicons, and the handle //! cache. -//! -//! The site's own concerns - routes, templates, moderation - live outside this -//! module and deal in [`id::Did`] and each lexicon's own storage methods -//! rather than raw records. pub mod actor; pub mod id; pub mod lexicon; -pub mod pds; +mod pds; pub mod tap; diff --git a/src/atproto/tap.rs b/src/atproto/tap.rs index ae6c320..df01861 100644 --- a/src/atproto/tap.rs +++ b/src/atproto/tap.rs @@ -5,7 +5,7 @@ use tokio_stream::StreamExt; use crate::atproto::actor; use crate::atproto::id::{Did, Handle, Rkey}; -use crate::atproto::lexicon::ModListing; +use crate::atproto::lexicon::{ModListing, Record}; use crate::state::AppState; /// Spawn the indexer. Runs until the process exits; reconnects on its own. @@ -50,14 +50,14 @@ async fn handle_record(state: &AppState, record: &RecordEvent) -> anyhow::Result let rkey = Rkey::from_verified(record.rkey.to_string()); match record.action { - RecordAction::Delete => ModListing::delete(&state.db, &did, &rkey).await, + RecordAction::Delete => ModListing::delete_local(&state.db, &did, &rkey).await, RecordAction::Create | RecordAction::Update => { let value = record .record_value() .ok_or_else(|| anyhow::anyhow!("{} action with no record", record.action))?; let listing = ModListing::from_value(value)?; listing.validate()?; - listing.save(&state.db, &did, &rkey).await + listing.save_local(&state.db, &did, &rkey).await } } } diff --git a/src/config.rs b/src/config.rs index 0a540b5..35d0daf 100644 --- a/src/config.rs +++ b/src/config.rs @@ -6,7 +6,7 @@ use std::sync::LazyLock; use clap::Parser; -use crate::atproto::lexicon::ModListing; +use crate::atproto::lexicon::{ModListing, Record}; // https://atproto.com/guides/permission-requests pub static OAUTH_SCOPE: LazyLock = diff --git a/src/mods.rs b/src/mods.rs index 061786a..0ff293a 100644 --- a/src/mods.rs +++ b/src/mods.rs @@ -3,8 +3,8 @@ //! No `@` sigil or `mods` segment: a handle is a domain name (a dot //! required), so it can never collide with a literal route like `/new`. -use axum::extract::{Form, Path, State}; -use axum::http::HeaderMap; +use axum::extract::{Form, FromRequestParts, Path, State}; +use axum::http::request::Parts; use axum::response::{IntoResponse, Redirect, Response}; use axum::routing::get; use axum::Router; @@ -12,12 +12,12 @@ use chrono::Utc; use maud::{html, Markup}; use serde::Deserialize; -use crate::atproto::id::{Handle, Rkey}; +use crate::atproto::id::Handle; use crate::atproto::lexicon::mod_listing::Game; -use crate::atproto::lexicon::ModListing; +use crate::atproto::lexicon::{Found, ModListing, Record}; use crate::error::AppError; use crate::layout; -use crate::oauth::session::get_session_from_headers; +use crate::oauth::session::SessionCookie; use crate::state::AppState; pub fn router() -> Router { @@ -26,10 +26,34 @@ pub fn router() -> Router { .route("/{handle}/{slug}", get(show).post(save)) } -async fn new_form(State(state): State, headers: HeaderMap) -> Result { - get_session_from_headers(&state.secrets.cookie_secret, &headers) - .ok_or(AppError::Unauthorized)?; - Ok(layout(create_form()).into_response()) +/// Resolves `/{handle}/{slug}` to the listing it names. +impl FromRequestParts for Found { + type Rejection = AppError; + + async fn from_request_parts(parts: &mut Parts, state: &AppState) -> Result { + let Path((handle, slug)) = Path::<(String, String)>::from_request_parts(parts, state) + .await + .map_err(|_| AppError::BadRequest("missing handle/slug".to_string()))?; + + let handle = Handle::new(&handle).map_err(|e| AppError::BadRequest(e.to_string()))?; + let did = handle.resolve_did(state).await?.ok_or(AppError::NotFound)?; + + let (rkey, record) = ModListing::find_by_slug(&state.db, &did, &slug) + .await + .map_err(AppError::Internal)? + .ok_or(AppError::NotFound)?; + + Ok(Self { + did, + handle, + rkey, + record, + }) + } +} + +async fn new_form(_session: SessionCookie) -> Response { + layout(create_form()).into_response() } fn create_form() -> Markup { @@ -102,12 +126,9 @@ struct NewListingForm { async fn create( State(state): State, - headers: HeaderMap, + session: SessionCookie, Form(form): Form, ) -> Result { - let session = get_session_from_headers(&state.secrets.cookie_secret, &headers) - .ok_or(AppError::Unauthorized)?; - let game: Game = form .game .parse() @@ -125,19 +146,7 @@ async fn create( media: vec![], created_at: Utc::now(), }; - listing - .validate() - .map_err(|e| AppError::BadRequest(e.to_string()))?; - - let rkey = Rkey::new(); - let record = listing.to_value().map_err(AppError::Internal)?; - crate::atproto::pds::put(&state, &session, ModListing::NSID, rkey.as_str(), record) - .await - .map_err(AppError::Internal)?; - listing - .save(&state.db, &session.did, &rkey) - .await - .map_err(AppError::Internal)?; + listing.create(&state, &session).await?; // We just wrote as this DID, so its handle - if it has a verified one - // is what names the listing's URL. @@ -149,28 +158,13 @@ async fn create( Ok(Redirect::to(&format!("/{handle}/{}", listing.slug)).into_response()) } -async fn show( - State(state): State, - Path((handle, slug)): Path<(String, String)>, - headers: HeaderMap, -) -> Result { - let handle = Handle::new(&handle).map_err(|e| AppError::BadRequest(e.to_string()))?; - let did = handle - .resolve_did(&state) - .await? - .ok_or(AppError::NotFound)?; - - let (_rkey, listing) = ModListing::find_by_slug(&state.db, &did, &slug) - .await - .map_err(AppError::Internal)? - .ok_or(AppError::NotFound)?; +async fn show(existing: Found, session: Option) -> Response { + let is_owner = session.is_some_and(|session| session.did == existing.did); + let listing = &existing.record; - let is_owner = get_session_from_headers(&state.secrets.cookie_secret, &headers) - .is_some_and(|session| session.did == did); - - Ok(layout(html! { + layout(html! { h1 { (listing.title) } - p { "by " a href={ "/" (handle) } { (handle) } } + p { "by " a href={ "/" (existing.handle) } { (existing.handle) } } p { (listing.description) } @if !listing.details.is_empty() { p { (listing.details) } } dl { @@ -181,16 +175,16 @@ async fn show( dt { "created" } dd { (listing.created_at.to_rfc3339()) } } @if is_owner { - (edit_form(&handle, &slug, &listing)) + (edit_form(&existing.handle, listing)) } }) - .into_response()) + .into_response() } // TODO -fn edit_form(handle: &Handle, slug: &str, listing: &ModListing) -> Markup { +fn edit_form(handle: &Handle, listing: &ModListing) -> Markup { html! { - div.centered { form.bg-pink method="post" action={ "/" (handle) "/" (slug) } { + div.centered { form.bg-pink method="post" action={ "/" (handle) "/" (listing.slug) } { div.form-item { label for="title" { "Title" } input type="text" id="title" name="title" value=(listing.title) required; @@ -241,35 +235,20 @@ struct EditListingForm { async fn save( State(state): State, - Path((handle, slug)): Path<(String, String)>, - headers: HeaderMap, + existing: Found, + session: SessionCookie, Form(form): Form, ) -> Result { - let handle = Handle::new(&handle).map_err(|e| AppError::BadRequest(e.to_string()))?; - // Same resolution `show` does, so the two agree on what this URL names. - let did = handle - .resolve_did(&state) - .await? - .ok_or(AppError::NotFound)?; - - let session = get_session_from_headers(&state.secrets.cookie_secret, &headers) - .ok_or(AppError::Unauthorized)?; - if session.did != did { + if session.did != existing.did { return Err(AppError::Forbidden); } - // Editing only - creation happens at `/new`, so there is always an - // existing row to update. - let (rkey, mut listing) = ModListing::find_by_slug(&state.db, &did, &slug) - .await - .map_err(AppError::Internal)? - .ok_or(AppError::NotFound)?; - let game: Game = form .game .parse() .map_err(|e: anyhow::Error| AppError::BadRequest(e.to_string()))?; + let mut listing = existing.record; listing.title = form.title; listing.description = form.description; listing.details = form.details; @@ -277,18 +256,8 @@ async fn save( listing.category = form.category; listing.license = form.license; - listing - .validate() - .map_err(|e| AppError::BadRequest(e.to_string()))?; - - let record = listing.to_value().map_err(AppError::Internal)?; - crate::atproto::pds::put(&state, &session, ModListing::NSID, rkey.as_str(), record) - .await - .map_err(AppError::Internal)?; - listing - .save(&state.db, &did, &rkey) - .await - .map_err(AppError::Internal)?; + let redirect = format!("/{}/{}", existing.handle, listing.slug); + listing.update(&state, &session, &existing.rkey).await?; - Ok(Redirect::to(&format!("/{handle}/{slug}")).into_response()) + Ok(Redirect::to(&redirect).into_response()) } diff --git a/src/oauth/session.rs b/src/oauth/session.rs index 53dbb7f..69fdcc4 100644 --- a/src/oauth/session.rs +++ b/src/oauth/session.rs @@ -6,6 +6,8 @@ use aes_gcm::aead::{Aead, KeyInit}; use aes_gcm::{Aes256Gcm, Nonce}; +use axum::extract::{FromRequestParts, OptionalFromRequestParts}; +use axum::http::request::Parts; use axum::http::{header, HeaderMap, HeaderValue}; use base64::engine::general_purpose::URL_SAFE_NO_PAD as BASE64; use base64::Engine as _; @@ -15,6 +17,8 @@ use rand::RngCore as _; use serde::{Deserialize, Serialize}; use crate::atproto::id::{Did, Handle}; +use crate::error::AppError; +use crate::state::AppState; /// Name of the encrypted session cookie. pub const SESSION_COOKIE_NAME: &str = "session"; @@ -38,6 +42,33 @@ pub struct SessionCookie { pub pds_endpoint: String, } +/// Extracts the session cookie, rejecting with 401 if there isn't one. +impl FromRequestParts for SessionCookie { + type Rejection = AppError; + + async fn from_request_parts(parts: &mut Parts, state: &AppState) -> Result { + get_session_from_headers(&state.secrets.cookie_secret, &parts.headers) + .ok_or(AppError::Unauthorized) + } +} + +/// `session: Option` for routes where being logged in is +/// optional - axum 0.8 needs this as a separate opt-in, not a blanket impl +/// over `FromRequestParts`. +impl OptionalFromRequestParts for SessionCookie { + type Rejection = AppError; + + async fn from_request_parts( + parts: &mut Parts, + state: &AppState, + ) -> Result, AppError> { + Ok(get_session_from_headers( + &state.secrets.cookie_secret, + &parts.headers, + )) + } +} + impl SessionCookie { /// Whether the access token expires within `duration`. pub fn expires_within(&self, duration: Duration) -> bool { -- 2.51.2