diff --git a/Cargo.lock b/Cargo.lock index 4369eb6..3885f0e 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -32,6 +32,16 @@ version = "1.0.101" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5f0e0fee31ef5ed1ba1316088939cea399010ed7731dba877ed44aeb407a75ea" +[[package]] +name = "assert-json-diff" +version = "2.0.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "47e4f2b81832e72834d7518d8487a0396a28cc408186a2e8854c0f98011faf12" +dependencies = [ + "serde", + "serde_json", +] + [[package]] name = "atoi" version = "2.0.0" @@ -287,6 +297,24 @@ version = "2.10.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d7a1e2f27636f116493b8b860f5546edb47c8d8f8ea73e1d2a20be88e28d1fea" +[[package]] +name = "deadpool" +version = "0.12.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0be2b1d1d6ec8d846f05e137292d0b89133caf95ef33695424c09568bdd39b1b" +dependencies = [ + "deadpool-runtime", + "lazy_static", + "num_cpus", + "tokio", +] + +[[package]] +name = "deadpool-runtime" +version = "0.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "092966b41edc516079bdf31ec78a2e0588d1d0c08f78b91d8307215928642b2b" + [[package]] name = "der" version = "0.7.10" @@ -495,6 +523,21 @@ dependencies = [ "percent-encoding", ] +[[package]] +name = "futures" +version = "0.3.31" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "65bc07b1a8bc7c85c5f2e110c476c7389b4554ba72af57d8445ea63a576b0876" +dependencies = [ + "futures-channel", + "futures-core", + "futures-executor", + "futures-io", + "futures-sink", + "futures-task", + "futures-util", +] + [[package]] name = "futures-channel" version = "0.3.31" @@ -568,6 +611,7 @@ version = "0.3.31" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9fa08315bb612088cc391249efdc3bc77536f16c91f6cf495e6fbe85b20a4a81" dependencies = [ + "futures-channel", "futures-core", "futures-io", "futures-macro", @@ -669,19 +713,23 @@ dependencies = [ "dotenvy", "futures-util", "hex", + "http-body-util", "jsonwebtoken", "p256", "reqwest", "serde", "serde_json", + "serial_test", "sha2", "sqlx", "tokio", "tokio-tungstenite", + "tower", "tower-http", "tracing", "tracing-subscriber", "uuid", + "wiremock", ] [[package]] @@ -716,6 +764,12 @@ version = "0.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "2304e00983f87ffb38b55b444b5e3b60a884b5d30c0fca7d82fe33449bbe55ea" +[[package]] +name = "hermit-abi" +version = "0.5.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fc0fef456e4baa96da950455cd02c081ca953b141298e41db3fc7e36b1da849c" + [[package]] name = "hex" version = "0.4.3" @@ -1276,6 +1330,16 @@ dependencies = [ "libm", ] +[[package]] +name = "num_cpus" +version = "1.17.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "91df4bbde75afed763b708b7eee1e8e7651e02d97f6d5dd763e89367e957b23b" +dependencies = [ + "hermit-abi", + "libc", +] + [[package]] name = "once_cell" version = "1.21.3" @@ -1575,6 +1639,18 @@ dependencies = [ "bitflags", ] +[[package]] +name = "regex" +version = "1.12.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e10754a14b9137dd7b1e3e5b0493cc9171fdd105e0ab477f51b72e7f3ac0e276" +dependencies = [ + "aho-corasick", + "memchr", + "regex-automata", + "regex-syntax", +] + [[package]] name = "regex-automata" version = "0.4.14" @@ -1735,6 +1811,15 @@ version = "1.0.23" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9774ba4a74de5f7b1c1451ed6cd5285a32eddb5cccb8cc655a4e50009e06477f" +[[package]] +name = "scc" +version = "2.4.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "46e6f046b7fef48e2660c57ed794263155d713de679057f2d0c169bfc6e756cc" +dependencies = [ + "sdd", +] + [[package]] name = "schannel" version = "0.1.28" @@ -1750,6 +1835,12 @@ version = "1.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "94143f37725109f92c262ed2cf5e59bce7498c01bcc1502d7b9afe439a4e9f49" +[[package]] +name = "sdd" +version = "3.0.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "490dcfcbfef26be6800d11870ff2df8774fa6e86d047e3e8c8a76b25655e41ca" + [[package]] name = "sec1" version = "0.7.3" @@ -1859,6 +1950,32 @@ dependencies = [ "serde", ] +[[package]] +name = "serial_test" +version = "3.3.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0d0b343e184fc3b7bb44dff0705fffcf4b3756ba6aff420dddd8b24ca145e555" +dependencies = [ + "futures-executor", + "futures-util", + "log", + "once_cell", + "parking_lot", + "scc", + "serial_test_derive", +] + +[[package]] +name = "serial_test_derive" +version = "3.3.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6f50427f258fb77356e4cd4aa0e87e2bd2c66dbcee41dc405282cae2bfc26c83" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + [[package]] name = "sha1" version = "0.10.6" @@ -3134,6 +3251,29 @@ version = "0.53.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d6bbff5f0aada427a1e5a6da5f1f98158182f26556f345ac9e04d36d0ebed650" +[[package]] +name = "wiremock" +version = "0.6.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "08db1edfb05d9b3c1542e521aea074442088292f00b5f28e435c714a98f85031" +dependencies = [ + "assert-json-diff", + "base64", + "deadpool", + "futures", + "http", + "http-body-util", + "hyper", + "hyper-util", + "log", + "once_cell", + "regex", + "serde", + "serde_json", + "tokio", + "url", +] + [[package]] name = "wit-bindgen" version = "0.51.0" diff --git a/Cargo.toml b/Cargo.toml index b19251d..3d32383 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -24,3 +24,9 @@ tower-http = { version = "0.6", features = ["cors", "trace"] } tracing = "0.1" tracing-subscriber = { version = "0.3", features = ["env-filter", "json"] } hex = "0.4.3" + +[dev-dependencies] +wiremock = "0.6" +tower = { version = "0.5", features = ["util"] } +http-body-util = "0.1" +serial_test = "3" diff --git a/src/admin.rs b/src/admin.rs index df06eb8..a5b0568 100644 --- a/src/admin.rs +++ b/src/admin.rs @@ -15,7 +15,7 @@ use crate::AppState; // --------------------------------------------------------------------------- /// SHA-256 hash a plaintext API key for storage/comparison. -fn hash_api_key(key: &str) -> String { +pub(crate) fn hash_api_key(key: &str) -> String { let hash = Sha256::digest(key.as_bytes()); hex::encode(hash) } @@ -545,3 +545,34 @@ async fn delete_admin( Ok(StatusCode::NO_CONTENT) } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn hash_api_key_produces_deterministic_sha256_hex() { + let h1 = hash_api_key("test-key"); + let h2 = hash_api_key("test-key"); + assert_eq!(h1, h2); + assert_eq!(h1.len(), 64); + assert!(h1.chars().all(|c| c.is_ascii_hexdigit())); + } + + #[test] + fn hash_api_key_different_inputs_differ() { + let h1 = hash_api_key("key-a"); + let h2 = hash_api_key("key-b"); + assert_ne!(h1, h2); + } + + #[test] + fn hash_api_key_known_value() { + // SHA-256 of "hello" is well-known + let hash = hash_api_key("hello"); + assert_eq!( + hash, + "2cf24dba5fb0a30e26e83b2ac5b9e29e1b161e5c1fa7425e73043362938b9824" + ); + } +} diff --git a/src/auth/middleware.rs b/src/auth/middleware.rs index 2315dd3..002de17 100644 --- a/src/auth/middleware.rs +++ b/src/auth/middleware.rs @@ -22,6 +22,12 @@ impl Claims { pub fn token(&self) -> &str { &self.token } + + /// Test-only constructor. + #[cfg(test)] + pub fn new_for_test(did: String, token: String) -> Self { + Self { did, token } + } } #[derive(Deserialize)] diff --git a/src/backfill.rs b/src/backfill.rs index a2b8ccc..0b1bf45 100644 --- a/src/backfill.rs +++ b/src/backfill.rs @@ -150,6 +150,7 @@ async fn run_job( db: &PgPool, http: &reqwest::Client, relay_url: &str, + plc_url: &str, job_id: &str, ) -> Result<(), String> { // Fetch the job @@ -232,9 +233,10 @@ async fn run_job( let db = db.clone(); let collection = collection.clone(); + let plc_url = plc_url.to_string(); let task = tokio::spawn(async move { let _permit = permit; - backfill_repo(&db, &http, &did, &collection).await + backfill_repo(&db, &http, &plc_url, &did, &collection).await }); tasks.push(task); } @@ -286,11 +288,12 @@ async fn run_job( async fn backfill_repo( db: &PgPool, http: &reqwest::Client, + plc_url: &str, did: &str, collection: &str, ) -> Result { // Resolve PDS - let pds = profile::resolve_pds_endpoint(http, did) + let pds = profile::resolve_pds_endpoint(http, plc_url, did) .await .map_err(|e| format!("PDS resolution failed for {did}: {e}"))?; @@ -330,7 +333,7 @@ async fn backfill_repo( // --------------------------------------------------------------------------- /// Spawn a background task that polls for pending backfill jobs and runs them. -pub fn spawn_worker(db: PgPool, http: reqwest::Client, relay_url: String) { +pub fn spawn_worker(db: PgPool, http: reqwest::Client, relay_url: String, plc_url: String) { tokio::spawn(async move { info!("backfill worker started"); loop { @@ -344,7 +347,7 @@ pub fn spawn_worker(db: PgPool, http: reqwest::Client, relay_url: String) { if let Some((job_id,)) = job { info!(job = %job_id, "picked up backfill job"); - if let Err(e) = run_job(&db, &http, &relay_url, &job_id).await { + if let Err(e) = run_job(&db, &http, &relay_url, &plc_url, &job_id).await { error!(job = %job_id, error = %e, "backfill job failed"); let _ = sqlx::query( "UPDATE backfill_jobs SET status = 'failed', completed_at = NOW(), error = $2 WHERE id::text = $1", diff --git a/src/config.rs b/src/config.rs index 3b2d4e0..583edd3 100644 --- a/src/config.rs +++ b/src/config.rs @@ -10,6 +10,7 @@ pub struct Config { pub jetstream_url: String, pub admin_secret: Option, pub relay_url: String, + pub plc_url: String, } impl Config { @@ -29,6 +30,8 @@ impl Config { admin_secret: env::var("ADMIN_SECRET").ok(), relay_url: env::var("RELAY_URL") .unwrap_or_else(|_| "https://bsky.network".into()), + plc_url: env::var("PLC_URL") + .unwrap_or_else(|_| "https://plc.directory".into()), } } @@ -38,3 +41,113 @@ impl Config { .expect("invalid HOST/PORT") } } + +#[cfg(test)] +mod tests { + use super::*; + use serial_test::serial; + + unsafe fn clear_env() { + for key in [ + "HOST", "PORT", "DATABASE_URL", "AIP_URL", "JETSTREAM_URL", + "ADMIN_SECRET", "RELAY_URL", "PLC_URL", + ] { + unsafe { env::remove_var(key); } + } + } + + unsafe fn set_required_env() { + unsafe { + env::set_var("DATABASE_URL", "postgres://localhost/test"); + env::set_var("AIP_URL", "http://localhost:4000"); + } + } + + #[test] + fn listen_addr_combines_host_and_port() { + let config = Config { + host: "127.0.0.1".into(), + port: 8080, + database_url: String::new(), + aip_url: String::new(), + jetstream_url: String::new(), + admin_secret: None, + relay_url: String::new(), + plc_url: String::new(), + }; + assert_eq!( + config.listen_addr(), + "127.0.0.1:8080".parse::().unwrap() + ); + } + + #[test] + #[serial] + fn from_env_reads_required_vars() { + unsafe { + clear_env(); + set_required_env(); + } + let config = Config::from_env(); + assert_eq!(config.database_url, "postgres://localhost/test"); + assert_eq!(config.aip_url, "http://localhost:4000"); + } + + #[test] + #[serial] + fn from_env_applies_defaults() { + unsafe { + clear_env(); + set_required_env(); + } + let config = Config::from_env(); + assert_eq!(config.host, "0.0.0.0"); + assert_eq!(config.port, 3000); + assert!(config.jetstream_url.contains("jetstream")); + assert_eq!(config.relay_url, "https://bsky.network"); + assert_eq!(config.plc_url, "https://plc.directory"); + assert!(config.admin_secret.is_none()); + } + + #[test] + #[serial] + fn from_env_reads_optional_overrides() { + unsafe { + clear_env(); + set_required_env(); + env::set_var("HOST", "10.0.0.1"); + env::set_var("PORT", "9090"); + env::set_var("ADMIN_SECRET", "s3cret"); + env::set_var("RELAY_URL", "https://relay.example.com"); + env::set_var("PLC_URL", "https://plc.example.com"); + } + let config = Config::from_env(); + assert_eq!(config.host, "10.0.0.1"); + assert_eq!(config.port, 9090); + assert_eq!(config.admin_secret, Some("s3cret".into())); + assert_eq!(config.relay_url, "https://relay.example.com"); + assert_eq!(config.plc_url, "https://plc.example.com"); + } + + #[test] + #[serial] + #[should_panic(expected = "DATABASE_URL must be set")] + fn from_env_panics_without_database_url() { + unsafe { + clear_env(); + env::set_var("AIP_URL", "http://localhost:4000"); + } + Config::from_env(); + } + + #[test] + #[serial] + #[should_panic(expected = "AIP_URL must be set")] + fn from_env_panics_without_aip_url() { + unsafe { + clear_env(); + env::set_var("DATABASE_URL", "postgres://localhost/test"); + } + Config::from_env(); + } +} diff --git a/src/error.rs b/src/error.rs index 471f942..5b26c6e 100644 --- a/src/error.rs +++ b/src/error.rs @@ -53,3 +53,79 @@ impl IntoResponse for AppError { } } } + +#[cfg(test)] +mod tests { + use super::*; + use axum::response::IntoResponse; + use http_body_util::BodyExt; + + async fn response_parts(err: AppError) -> (StatusCode, serde_json::Value) { + let resp = err.into_response(); + let status = resp.status(); + let body = resp.into_body().collect().await.unwrap().to_bytes(); + let json: serde_json::Value = serde_json::from_slice(&body).unwrap(); + (status, json) + } + + #[tokio::test] + async fn auth_error_returns_401() { + let (status, body) = response_parts(AppError::Auth("bad token".into())).await; + assert_eq!(status, StatusCode::UNAUTHORIZED); + assert_eq!(body["error"], "bad token"); + } + + #[tokio::test] + async fn bad_request_returns_400() { + let (status, body) = response_parts(AppError::BadRequest("missing field".into())).await; + assert_eq!(status, StatusCode::BAD_REQUEST); + assert_eq!(body["error"], "missing field"); + } + + #[tokio::test] + async fn internal_error_returns_500_and_hides_detail() { + let (status, body) = response_parts(AppError::Internal("secret details".into())).await; + assert_eq!(status, StatusCode::INTERNAL_SERVER_ERROR); + assert_eq!(body["error"], "internal server error"); + } + + #[tokio::test] + async fn not_found_returns_404() { + let (status, body) = response_parts(AppError::NotFound("no such thing".into())).await; + assert_eq!(status, StatusCode::NOT_FOUND); + assert_eq!(body["error"], "no such thing"); + } + + #[tokio::test] + async fn pds_error_preserves_status_and_body() { + let raw_body = Bytes::from(r#"{"error":"upstream"}"#); + let resp = AppError::PdsError(StatusCode::BAD_GATEWAY, raw_body.clone()).into_response(); + assert_eq!(resp.status(), StatusCode::BAD_GATEWAY); + let body = resp.into_body().collect().await.unwrap().to_bytes(); + assert_eq!(body, raw_body); + } + + #[test] + fn display_formats() { + assert_eq!( + AppError::Auth("x".into()).to_string(), + "auth error: x" + ); + assert_eq!( + AppError::BadRequest("y".into()).to_string(), + "bad request: y" + ); + assert_eq!( + AppError::Internal("z".into()).to_string(), + "internal error: z" + ); + assert_eq!( + AppError::NotFound("w".into()).to_string(), + "not found: w" + ); + assert_eq!( + AppError::PdsError(StatusCode::BAD_GATEWAY, Bytes::new()).to_string(), + "PDS error: 502 Bad Gateway" + ); + } +} diff --git a/src/lexicon.rs b/src/lexicon.rs index e5bf7ee..db81214 100644 --- a/src/lexicon.rs +++ b/src/lexicon.rs @@ -183,3 +183,230 @@ impl LexiconRegistry { inner.len() } } + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + + // ----------------------------------------------------------------------- + // ParsedLexicon::parse + // ----------------------------------------------------------------------- + + fn record_lexicon_json() -> Value { + json!({ + "lexicon": 1, + "id": "games.gamesgamesgamesgames.game", + "defs": { + "main": { + "type": "record", + "key": "tid", + "record": { + "type": "object", + "properties": { + "title": { "type": "string" } + } + } + } + } + }) + } + + fn query_lexicon_json() -> Value { + json!({ + "lexicon": 1, + "id": "games.gamesgamesgamesgames.listGames", + "defs": { + "main": { + "type": "query", + "parameters": { + "type": "params", + "properties": { + "limit": { "type": "integer" } + } + }, + "output": { + "encoding": "application/json" + } + } + } + }) + } + + fn procedure_lexicon_json() -> Value { + json!({ + "lexicon": 1, + "id": "games.gamesgamesgamesgames.createGame", + "defs": { + "main": { + "type": "procedure", + "input": { + "encoding": "application/json" + }, + "output": { + "encoding": "application/json" + } + } + } + }) + } + + fn definitions_lexicon_json() -> Value { + json!({ + "lexicon": 1, + "id": "games.gamesgamesgamesgames.defs", + "defs": { + "genre": { + "type": "string", + "knownValues": ["action", "rpg"] + } + } + }) + } + + #[test] + fn parse_record_lexicon() { + let parsed = ParsedLexicon::parse(record_lexicon_json(), 1, None).unwrap(); + assert_eq!(parsed.id, "games.gamesgamesgamesgames.game"); + assert_eq!(parsed.lexicon_type, LexiconType::Record); + assert_eq!(parsed.record_key, Some("tid".into())); + assert!(parsed.record_schema.is_some()); + assert!(parsed.parameters.is_none()); + assert!(parsed.input.is_none()); + } + + #[test] + fn parse_query_lexicon() { + let parsed = ParsedLexicon::parse(query_lexicon_json(), 2, Some("games.gamesgamesgamesgames.game".into())).unwrap(); + assert_eq!(parsed.lexicon_type, LexiconType::Query); + assert!(parsed.parameters.is_some()); + assert!(parsed.output.is_some()); + assert_eq!(parsed.target_collection, Some("games.gamesgamesgamesgames.game".into())); + assert_eq!(parsed.revision, 2); + } + + #[test] + fn parse_procedure_lexicon() { + let parsed = ParsedLexicon::parse(procedure_lexicon_json(), 1, None).unwrap(); + assert_eq!(parsed.lexicon_type, LexiconType::Procedure); + assert!(parsed.input.is_some()); + assert!(parsed.output.is_some()); + } + + #[test] + fn parse_definitions_lexicon() { + let parsed = ParsedLexicon::parse(definitions_lexicon_json(), 1, None).unwrap(); + assert_eq!(parsed.lexicon_type, LexiconType::Definitions); + } + + #[test] + fn parse_missing_id_returns_error() { + let raw = json!({"lexicon": 1, "defs": {}}); + let result = ParsedLexicon::parse(raw, 1, None); + assert!(result.is_err()); + assert!(result.unwrap_err().contains("id")); + } + + #[test] + fn parse_preserves_raw_json() { + let raw = record_lexicon_json(); + let parsed = ParsedLexicon::parse(raw.clone(), 1, None).unwrap(); + assert_eq!(parsed.raw, raw); + } + + #[test] + fn parse_target_collection_passthrough() { + let parsed = ParsedLexicon::parse( + query_lexicon_json(), + 1, + Some("custom.collection".into()), + ) + .unwrap(); + assert_eq!(parsed.target_collection, Some("custom.collection".into())); + } + + // ----------------------------------------------------------------------- + // LexiconRegistry + // ----------------------------------------------------------------------- + + #[tokio::test] + async fn registry_new_is_empty() { + let reg = LexiconRegistry::new(); + assert_eq!(reg.count().await, 0); + } + + #[tokio::test] + async fn registry_upsert_and_get() { + let reg = LexiconRegistry::new(); + let parsed = ParsedLexicon::parse(record_lexicon_json(), 1, None).unwrap(); + reg.upsert(parsed).await; + + let got = reg.get("games.gamesgamesgamesgames.game").await; + assert!(got.is_some()); + assert_eq!(got.unwrap().lexicon_type, LexiconType::Record); + } + + #[tokio::test] + async fn registry_upsert_replaces() { + let reg = LexiconRegistry::new(); + let v1 = ParsedLexicon::parse(record_lexicon_json(), 1, None).unwrap(); + reg.upsert(v1).await; + + let v2 = ParsedLexicon::parse(record_lexicon_json(), 5, None).unwrap(); + reg.upsert(v2).await; + + assert_eq!(reg.count().await, 1); + assert_eq!(reg.get("games.gamesgamesgamesgames.game").await.unwrap().revision, 5); + } + + #[tokio::test] + async fn registry_remove_existing() { + let reg = LexiconRegistry::new(); + let parsed = ParsedLexicon::parse(record_lexicon_json(), 1, None).unwrap(); + reg.upsert(parsed).await; + + assert!(reg.remove("games.gamesgamesgamesgames.game").await); + assert_eq!(reg.count().await, 0); + } + + #[tokio::test] + async fn registry_remove_nonexistent() { + let reg = LexiconRegistry::new(); + assert!(!reg.remove("nonexistent").await); + } + + #[tokio::test] + async fn registry_get_nonexistent() { + let reg = LexiconRegistry::new(); + assert!(reg.get("nonexistent").await.is_none()); + } + + #[tokio::test] + async fn registry_type_filtered_collections() { + let reg = LexiconRegistry::new(); + + let record = ParsedLexicon::parse(record_lexicon_json(), 1, None).unwrap(); + let query = ParsedLexicon::parse(query_lexicon_json(), 1, None).unwrap(); + let procedure = ParsedLexicon::parse(procedure_lexicon_json(), 1, None).unwrap(); + let defs = ParsedLexicon::parse(definitions_lexicon_json(), 1, None).unwrap(); + + reg.upsert(record).await; + reg.upsert(query).await; + reg.upsert(procedure).await; + reg.upsert(defs).await; + + assert_eq!(reg.count().await, 4); + + let records = reg.get_record_collections().await; + assert_eq!(records.len(), 1); + assert!(records.contains(&"games.gamesgamesgamesgames.game".to_string())); + + let queries = reg.get_queries().await; + assert_eq!(queries.len(), 1); + assert!(queries.contains(&"games.gamesgamesgamesgames.listGames".to_string())); + + let procedures = reg.get_procedures().await; + assert_eq!(procedures.len(), 1); + assert!(procedures.contains(&"games.gamesgamesgamesgames.createGame".to_string())); + } +} diff --git a/src/lib.rs b/src/lib.rs new file mode 100644 index 0000000..64a05e6 --- /dev/null +++ b/src/lib.rs @@ -0,0 +1,24 @@ +pub mod admin; +pub mod auth; +pub mod backfill; +pub mod config; +pub mod error; +pub mod jetstream; +pub mod lexicon; +pub mod profile; +pub mod repo; +pub mod server; +pub mod xrpc; + +use config::Config; +use lexicon::LexiconRegistry; +use tokio::sync::watch; + +#[derive(Clone)] +pub struct AppState { + pub config: Config, + pub http: reqwest::Client, + pub db: sqlx::PgPool, + pub lexicons: LexiconRegistry, + pub collections_tx: watch::Sender>, +} diff --git a/src/main.rs b/src/main.rs index c8778ec..ed5768c 100644 --- a/src/main.rs +++ b/src/main.rs @@ -1,29 +1,9 @@ -mod admin; -mod auth; -mod backfill; -mod config; -mod error; -mod jetstream; -mod lexicon; -mod profile; -mod repo; -mod server; -mod xrpc; - -use config::Config; -use lexicon::LexiconRegistry; +use happyview::config::Config; +use happyview::lexicon::LexiconRegistry; +use happyview::{admin, backfill, jetstream, server, AppState}; use tokio::sync::watch; use tracing::info; -#[derive(Clone)] -pub struct AppState { - pub config: Config, - pub http: reqwest::Client, - pub db: sqlx::PgPool, - pub lexicons: LexiconRegistry, - pub collections_tx: watch::Sender>, -} - #[tokio::main] async fn main() { dotenvy::dotenv().ok(); @@ -69,7 +49,7 @@ async fn main() { }; jetstream::spawn(state.db.clone(), config.jetstream_url.clone(), collections_rx); - backfill::spawn_worker(state.db.clone(), state.http.clone(), config.relay_url.clone()); + backfill::spawn_worker(state.db.clone(), state.http.clone(), config.relay_url.clone(), config.plc_url.clone()); let app = server::router(state); let addr = config.listen_addr(); diff --git a/src/profile.rs b/src/profile.rs index dfe86d9..722c28e 100644 --- a/src/profile.rs +++ b/src/profile.rs @@ -36,8 +36,8 @@ struct GetRecordResponse { } /// Resolve a full profile for the given DID: DID document -> handle + PDS -> profile record. -pub async fn resolve_profile(http: &reqwest::Client, did: &str) -> Result { - let did_doc = resolve_did_document(http, did).await?; +pub async fn resolve_profile(http: &reqwest::Client, plc_url: &str, did: &str) -> Result { + let did_doc = resolve_did_document(http, plc_url, did).await?; let handle = did_doc .also_known_as @@ -67,8 +67,8 @@ pub async fn resolve_profile(http: &reqwest::Client, did: &str) -> Result Result { - let did_doc = resolve_did_document(http, did).await?; +pub async fn resolve_pds_endpoint(http: &reqwest::Client, plc_url: &str, did: &str) -> Result { + let did_doc = resolve_did_document(http, plc_url, did).await?; did_doc .service @@ -80,8 +80,8 @@ pub async fn resolve_pds_endpoint(http: &reqwest::Client, did: &str) -> Result Result { - let url = format!("https://plc.directory/{did}"); +async fn resolve_did_document(http: &reqwest::Client, plc_url: &str, did: &str) -> Result { + let url = format!("{}/{did}", plc_url.trim_end_matches('/')); let resp = http .get(&url) diff --git a/src/repo.rs b/src/repo.rs index 7aa9c33..6952134 100644 --- a/src/repo.rs +++ b/src/repo.rs @@ -358,3 +358,221 @@ pub(crate) fn enrich_media_blobs(record: &mut Value, pds: &str, did: &str) { } } } + +#[cfg(test)] +mod tests { + use super::*; + + // ----------------------------------------------------------------------- + // parse_did_from_at_uri + // ----------------------------------------------------------------------- + + #[test] + fn parse_did_from_valid_at_uri() { + let did = parse_did_from_at_uri("at://did:plc:abc123/app.bsky.feed.post/3k2bqxyz").unwrap(); + assert_eq!(did, "did:plc:abc123"); + } + + #[test] + fn parse_did_from_uri_with_no_rkey() { + let did = parse_did_from_at_uri("at://did:plc:abc123/collection").unwrap(); + assert_eq!(did, "did:plc:abc123"); + } + + #[test] + fn parse_did_from_did_web_uri() { + let did = parse_did_from_at_uri("at://did:web:example.com/collection/rkey").unwrap(); + assert_eq!(did, "did:web:example.com"); + } + + #[test] + fn parse_did_from_uri_missing_prefix() { + let result = parse_did_from_at_uri("did:plc:abc123/collection/rkey"); + assert!(result.is_err()); + } + + // ----------------------------------------------------------------------- + // enrich_media_blobs + // ----------------------------------------------------------------------- + + #[test] + fn enrich_media_adds_url() { + let mut record = json!({ + "media": [{ + "blob": { + "ref": { "$link": "bafyreiabc" }, + "mimeType": "image/jpeg", + "size": 1024 + } + }] + }); + + enrich_media_blobs(&mut record, "https://pds.example.com", "did:plc:test"); + + let url = record["media"][0]["blob"]["url"].as_str().unwrap(); + assert_eq!( + url, + "https://pds.example.com/xrpc/com.atproto.sync.getBlob?did=did:plc:test&cid=bafyreiabc" + ); + } + + #[test] + fn enrich_media_noop_without_media() { + let mut record = json!({"title": "test"}); + enrich_media_blobs(&mut record, "https://pds.example.com", "did:plc:test"); + assert!(record.get("media").is_none()); + } + + #[test] + fn enrich_media_skips_items_without_ref() { + let mut record = json!({ + "media": [{ + "blob": { "mimeType": "image/png" } + }] + }); + + enrich_media_blobs(&mut record, "https://pds.example.com", "did:plc:test"); + assert!(record["media"][0]["blob"].get("url").is_none()); + } + + #[test] + fn enrich_media_handles_multiple_items() { + let mut record = json!({ + "media": [ + { "blob": { "ref": { "$link": "cid1" } } }, + { "blob": { "ref": { "$link": "cid2" } } } + ] + }); + + enrich_media_blobs(&mut record, "https://pds.example.com/", "did:plc:x"); + + let url1 = record["media"][0]["blob"]["url"].as_str().unwrap(); + let url2 = record["media"][1]["blob"]["url"].as_str().unwrap(); + assert!(url1.contains("cid1")); + assert!(url2.contains("cid2")); + } + + #[test] + fn enrich_media_trims_trailing_slash() { + let mut record = json!({ + "media": [{ + "blob": { "ref": { "$link": "bafytest" } } + }] + }); + + enrich_media_blobs(&mut record, "https://pds.example.com/", "did:plc:test"); + + let url = record["media"][0]["blob"]["url"].as_str().unwrap(); + assert!(url.starts_with("https://pds.example.com/xrpc/")); + assert!(!url.contains("//xrpc")); + } + + // ----------------------------------------------------------------------- + // generate_dpop_proof + // ----------------------------------------------------------------------- + + fn test_dpop_jwk() -> DpopJwk { + use p256::elliptic_curve::rand_core::OsRng; + use p256::elliptic_curve::sec1::ToEncodedPoint; + // Generate a valid P-256 key for testing + let secret = p256::SecretKey::random(&mut OsRng); + let public = secret.public_key(); + let point = public.to_encoded_point(false); + + DpopJwk { + x: base64::engine::general_purpose::URL_SAFE_NO_PAD + .encode(point.x().unwrap()), + y: base64::engine::general_purpose::URL_SAFE_NO_PAD + .encode(point.y().unwrap()), + d: base64::engine::general_purpose::URL_SAFE_NO_PAD + .encode(secret.to_bytes()), + } + } + + #[test] + fn dpop_proof_produces_valid_jwt_structure() { + let jwk = test_dpop_jwk(); + let token = generate_dpop_proof("POST", "https://pds.example.com/xrpc/test", &jwk, "access-tok", None).unwrap(); + + let parts: Vec<&str> = token.split('.').collect(); + assert_eq!(parts.len(), 3, "JWT should have 3 parts"); + } + + #[test] + fn dpop_proof_header_has_correct_fields() { + let jwk = test_dpop_jwk(); + let token = generate_dpop_proof("POST", "https://pds.example.com/xrpc/test", &jwk, "access-tok", None).unwrap(); + + let header_b64 = token.split('.').next().unwrap(); + let header_bytes = base64::engine::general_purpose::URL_SAFE_NO_PAD + .decode(header_b64) + .unwrap(); + let header: serde_json::Value = serde_json::from_slice(&header_bytes).unwrap(); + + assert_eq!(header["typ"], "dpop+jwt"); + assert_eq!(header["alg"], "ES256"); + assert!(header.get("jwk").is_some()); + } + + #[test] + fn dpop_proof_claims_have_correct_fields() { + let jwk = test_dpop_jwk(); + let token = generate_dpop_proof("GET", "https://pds.example.com/xrpc/test", &jwk, "my-access-token", None).unwrap(); + + let payload_b64 = token.split('.').nth(1).unwrap(); + let payload_bytes = base64::engine::general_purpose::URL_SAFE_NO_PAD + .decode(payload_b64) + .unwrap(); + let claims: serde_json::Value = serde_json::from_slice(&payload_bytes).unwrap(); + + assert_eq!(claims["htm"], "GET"); + assert_eq!(claims["htu"], "https://pds.example.com/xrpc/test"); + assert!(claims.get("jti").is_some()); + assert!(claims.get("iat").is_some()); + assert!(claims.get("exp").is_some()); + assert!(claims.get("ath").is_some()); + assert!(claims.get("nonce").is_none()); + } + + #[test] + fn dpop_proof_includes_nonce_when_provided() { + let jwk = test_dpop_jwk(); + let token = generate_dpop_proof("POST", "https://pds.example.com/xrpc/test", &jwk, "tok", Some("abc123")).unwrap(); + + let payload_b64 = token.split('.').nth(1).unwrap(); + let payload_bytes = base64::engine::general_purpose::URL_SAFE_NO_PAD + .decode(payload_b64) + .unwrap(); + let claims: serde_json::Value = serde_json::from_slice(&payload_bytes).unwrap(); + + assert_eq!(claims["nonce"], "abc123"); + } + + #[test] + fn dpop_proof_ath_is_sha256_of_access_token() { + let jwk = test_dpop_jwk(); + let access_token = "test-access-token"; + let token = generate_dpop_proof("POST", "https://example.com", &jwk, access_token, None).unwrap(); + + let payload_b64 = token.split('.').nth(1).unwrap(); + let payload_bytes = base64::engine::general_purpose::URL_SAFE_NO_PAD + .decode(payload_b64) + .unwrap(); + let claims: serde_json::Value = serde_json::from_slice(&payload_bytes).unwrap(); + + let expected_hash = Sha256::digest(access_token.as_bytes()); + let expected_ath = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(expected_hash); + assert_eq!(claims["ath"], expected_ath); + } + + #[test] + fn dpop_proof_invalid_key_returns_error() { + let jwk = DpopJwk { + x: "invalid".into(), + y: "invalid".into(), + d: "invalid".into(), + }; + let result = generate_dpop_proof("POST", "https://example.com", &jwk, "tok", None); + assert!(result.is_err()); + } +} diff --git a/src/server.rs b/src/server.rs index 57ea5d7..70339b6 100644 --- a/src/server.rs +++ b/src/server.rs @@ -36,6 +36,6 @@ async fn get_profile( State(state): State, claims: Claims, ) -> Result, AppError> { - let profile = profile::resolve_profile(&state.http, claims.did()).await?; + let profile = profile::resolve_profile(&state.http, &state.config.plc_url, claims.did()).await?; Ok(Json(profile)) } diff --git a/src/xrpc.rs b/src/xrpc.rs index da73ee6..2cb6f41 100644 --- a/src/xrpc.rs +++ b/src/xrpc.rs @@ -126,7 +126,7 @@ async fn handle_query( let unique_dids: HashSet<&str> = rows.iter().map(|(_, did, _)| did.as_str()).collect(); let mut pds_map: HashMap = HashMap::new(); for did in unique_dids { - if let Ok(pds) = profile::resolve_pds_endpoint(&state.http, did).await { + if let Ok(pds) = profile::resolve_pds_endpoint(&state.http, &state.config.plc_url, did).await { pds_map.insert(did.to_string(), pds); } } @@ -169,7 +169,7 @@ async fn handle_get_record(state: &AppState, uri: &str) -> Result