From a289f80bef106d0cd775ad72bbd702d563a28e1b Mon Sep 17 00:00:00 2001 From: Trezy Date: Sun, 12 Apr 2026 19:59:17 -0500 Subject: [PATCH] feat: add support for API clients --- .../20260412000000_create_api_clients.sql | 15 + ...412100000_auth_redirects_add_client_id.sql | 1 + ...260412200000_drop_rate_limit_allowlist.sql | 1 + .../20260412000000_create_api_clients.sql | 15 + ...412100000_auth_redirects_add_client_id.sql | 1 + ...260412200000_drop_rate_limit_allowlist.sql | 1 + src/admin/api_clients.rs | 471 +++++++++++++ src/admin/mod.rs | 16 +- src/admin/permissions.rs | 23 +- src/admin/rate_limits.rs | 126 +--- src/admin/types.rs | 71 +- src/auth/client_registry.rs | 250 +++++++ src/auth/middleware.rs | 31 +- src/auth/mod.rs | 2 + src/auth/routes.rs | 78 ++- src/lib.rs | 2 +- src/lua/atproto_api.rs | 5 +- src/lua/db_api.rs | 5 +- src/lua/execute.rs | 5 +- src/lua/http_api.rs | 5 +- src/main.rs | 44 +- src/rate_limit.rs | 263 +++++-- src/repo/session.rs | 1 + src/repo/upload_blob.rs | 8 - src/server.rs | 22 +- src/xrpc/mod.rs | 51 +- tests/common/app.rs | 5 +- tests/e2e_api_clients.rs | 643 ++++++++++++++++++ tests/lua_atproto_api.rs | 5 +- tests/lua_db_api.rs | 5 +- 30 files changed, 1847 insertions(+), 324 deletions(-) create mode 100644 migrations/postgres/20260412000000_create_api_clients.sql create mode 100644 migrations/postgres/20260412100000_auth_redirects_add_client_id.sql create mode 100644 migrations/postgres/20260412200000_drop_rate_limit_allowlist.sql create mode 100644 migrations/sqlite/20260412000000_create_api_clients.sql create mode 100644 migrations/sqlite/20260412100000_auth_redirects_add_client_id.sql create mode 100644 migrations/sqlite/20260412200000_drop_rate_limit_allowlist.sql create mode 100644 src/admin/api_clients.rs create mode 100644 src/auth/client_registry.rs create mode 100644 tests/e2e_api_clients.rs diff --git a/migrations/postgres/20260412000000_create_api_clients.sql b/migrations/postgres/20260412000000_create_api_clients.sql new file mode 100644 index 0000000..c0cfc27 --- /dev/null +++ b/migrations/postgres/20260412000000_create_api_clients.sql @@ -0,0 +1,15 @@ +CREATE TABLE IF NOT EXISTS api_clients ( + id TEXT PRIMARY KEY, + client_key TEXT NOT NULL UNIQUE, + name TEXT NOT NULL, + client_id_url TEXT NOT NULL UNIQUE, + client_uri TEXT NOT NULL, + redirect_uris TEXT NOT NULL, + scopes TEXT NOT NULL DEFAULT 'atproto', + rate_limit_capacity INTEGER, + rate_limit_refill_rate REAL, + is_active INTEGER NOT NULL DEFAULT 1, + created_by TEXT NOT NULL, + created_at TEXT NOT NULL DEFAULT '', + updated_at TEXT NOT NULL DEFAULT '' +); diff --git a/migrations/postgres/20260412100000_auth_redirects_add_client_id.sql b/migrations/postgres/20260412100000_auth_redirects_add_client_id.sql new file mode 100644 index 0000000..c31d741 --- /dev/null +++ b/migrations/postgres/20260412100000_auth_redirects_add_client_id.sql @@ -0,0 +1 @@ +ALTER TABLE auth_login_redirects ADD COLUMN IF NOT EXISTS client_id TEXT; diff --git a/migrations/postgres/20260412200000_drop_rate_limit_allowlist.sql b/migrations/postgres/20260412200000_drop_rate_limit_allowlist.sql new file mode 100644 index 0000000..aca3744 --- /dev/null +++ b/migrations/postgres/20260412200000_drop_rate_limit_allowlist.sql @@ -0,0 +1 @@ +DROP TABLE IF EXISTS rate_limit_allowlist; diff --git a/migrations/sqlite/20260412000000_create_api_clients.sql b/migrations/sqlite/20260412000000_create_api_clients.sql new file mode 100644 index 0000000..c0cfc27 --- /dev/null +++ b/migrations/sqlite/20260412000000_create_api_clients.sql @@ -0,0 +1,15 @@ +CREATE TABLE IF NOT EXISTS api_clients ( + id TEXT PRIMARY KEY, + client_key TEXT NOT NULL UNIQUE, + name TEXT NOT NULL, + client_id_url TEXT NOT NULL UNIQUE, + client_uri TEXT NOT NULL, + redirect_uris TEXT NOT NULL, + scopes TEXT NOT NULL DEFAULT 'atproto', + rate_limit_capacity INTEGER, + rate_limit_refill_rate REAL, + is_active INTEGER NOT NULL DEFAULT 1, + created_by TEXT NOT NULL, + created_at TEXT NOT NULL DEFAULT '', + updated_at TEXT NOT NULL DEFAULT '' +); diff --git a/migrations/sqlite/20260412100000_auth_redirects_add_client_id.sql b/migrations/sqlite/20260412100000_auth_redirects_add_client_id.sql new file mode 100644 index 0000000..57b3bec --- /dev/null +++ b/migrations/sqlite/20260412100000_auth_redirects_add_client_id.sql @@ -0,0 +1 @@ +ALTER TABLE auth_login_redirects ADD COLUMN client_id TEXT; diff --git a/migrations/sqlite/20260412200000_drop_rate_limit_allowlist.sql b/migrations/sqlite/20260412200000_drop_rate_limit_allowlist.sql new file mode 100644 index 0000000..aca3744 --- /dev/null +++ b/migrations/sqlite/20260412200000_drop_rate_limit_allowlist.sql @@ -0,0 +1 @@ +DROP TABLE IF EXISTS rate_limit_allowlist; diff --git a/src/admin/api_clients.rs b/src/admin/api_clients.rs new file mode 100644 index 0000000..61f95fa --- /dev/null +++ b/src/admin/api_clients.rs @@ -0,0 +1,471 @@ +use axum::Json; +use axum::extract::{Path, State}; +use axum::http::StatusCode; +use hex; +use rand::Rng; +use uuid::Uuid; + +use crate::AppState; +use crate::db::{adapt_sql, now_rfc3339}; +use crate::error::AppError; +use crate::event_log::{EventLog, Severity, log_event}; + +use super::auth::UserAuth; +use super::permissions::Permission; +use super::types::{ + ApiClientSummary, CreateApiClientBody, CreateApiClientResponse, UpdateApiClientBody, +}; + +/// POST /admin/api-clients — create a new API client. +pub(super) async fn create_api_client( + State(state): State, + auth: UserAuth, + Json(body): Json, +) -> Result<(StatusCode, Json), AppError> { + auth.require(Permission::ApiClientsCreate).await?; + + // Generate the client key: "hvc_" + 32 random hex chars. + let mut random_bytes = [0u8; 16]; + rand::rng().fill(&mut random_bytes); + let client_key = format!("hvc_{}", hex::encode(random_bytes)); + + let id = Uuid::new_v4().to_string(); + let now = now_rfc3339(); + let redirect_uris_json = + serde_json::to_string(&body.redirect_uris).unwrap_or_else(|_| "[]".to_string()); + + let insert_sql = adapt_sql( + "INSERT INTO api_clients (id, client_key, name, client_id_url, client_uri, redirect_uris, scopes, rate_limit_capacity, rate_limit_refill_rate, is_active, created_by, created_at, updated_at) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, 1, ?, ?, ?)", + state.db_backend, + ); + + sqlx::query(&insert_sql) + .bind(&id) + .bind(&client_key) + .bind(&body.name) + .bind(&body.client_id_url) + .bind(&body.client_uri) + .bind(&redirect_uris_json) + .bind(&body.scopes) + .bind(body.rate_limit_capacity) + .bind(body.rate_limit_refill_rate) + .bind(&auth.did) + .bind(&now) + .bind(&now) + .execute(&state.db) + .await + .map_err(|e| AppError::Internal(format!("failed to create api client: {e}")))?; + + // Register the new client in the OAuth registry so it's usable immediately. + let oauth_params = crate::auth::client_registry::ApiClientOAuthParams { + plc_url: state.config.plc_url.clone(), + state_store: state.oauth_state_store.clone(), + session_store_pool: state.db.clone(), + db_backend: state.db_backend, + }; + if let Err(e) = state.oauth.register_api_client( + &body.client_id_url, + &body.client_uri, + body.redirect_uris.clone(), + &body.scopes, + &oauth_params, + ) { + tracing::warn!(client_id = %body.client_id_url, error = %e, "OAuth client registration failed (DB row created)"); + } + + // Register per-client rate limit config if overrides are set. + if let (Some(capacity), Some(refill_rate)) = + (body.rate_limit_capacity, body.rate_limit_refill_rate) + { + let global = state.rate_limiter.global_config(); + state.rate_limiter.register_client_config( + client_key.clone(), + crate::rate_limit::RateLimitConfig { + capacity: capacity as u32, + refill_rate, + default_query_cost: global.default_query_cost, + default_procedure_cost: global.default_procedure_cost, + default_proxy_cost: global.default_proxy_cost, + }, + ); + } + + log_event( + &state.db, + EventLog { + event_type: "api_client.created".to_string(), + severity: Severity::Info, + actor_did: Some(auth.did.clone()), + subject: Some(body.name.clone()), + detail: serde_json::json!({ + "client_key": client_key, + "client_id_url": body.client_id_url, + }), + }, + state.db_backend, + ) + .await; + + Ok(( + StatusCode::CREATED, + Json(CreateApiClientResponse { + id, + client_key, + name: body.name, + client_id_url: body.client_id_url, + }), + )) +} + +/// GET /admin/api-clients — list all API clients. +pub(super) async fn list_api_clients( + State(state): State, + auth: UserAuth, +) -> Result>, AppError> { + auth.require(Permission::ApiClientsView).await?; + + let select_sql = adapt_sql( + "SELECT id, client_key, name, client_id_url, client_uri, redirect_uris, scopes, rate_limit_capacity, rate_limit_refill_rate, is_active, created_by, created_at, updated_at FROM api_clients ORDER BY created_at DESC", + state.db_backend, + ); + + #[allow(clippy::type_complexity)] + let rows: Vec<( + String, + String, + String, + String, + String, + String, + String, + Option, + Option, + i32, + String, + String, + String, + )> = sqlx::query_as(&select_sql) + .fetch_all(&state.db) + .await + .map_err(|e| AppError::Internal(format!("failed to list api clients: {e}")))?; + + let clients: Vec = rows + .into_iter() + .map( + |( + id, + client_key, + name, + client_id_url, + client_uri, + redirect_uris_json, + scopes, + rate_limit_capacity, + rate_limit_refill_rate, + is_active, + created_by, + created_at, + updated_at, + )| { + let redirect_uris: Vec = + serde_json::from_str(&redirect_uris_json).unwrap_or_default(); + ApiClientSummary { + id, + client_key, + name, + client_id_url, + client_uri, + redirect_uris, + scopes, + rate_limit_capacity, + rate_limit_refill_rate, + is_active: is_active != 0, + created_by, + created_at, + updated_at, + } + }, + ) + .collect(); + + Ok(Json(clients)) +} + +/// GET /admin/api-clients/:id — get a single API client. +pub(super) async fn get_api_client( + State(state): State, + auth: UserAuth, + Path(id): Path, +) -> Result, AppError> { + auth.require(Permission::ApiClientsView).await?; + + let select_sql = adapt_sql( + "SELECT id, client_key, name, client_id_url, client_uri, redirect_uris, scopes, rate_limit_capacity, rate_limit_refill_rate, is_active, created_by, created_at, updated_at FROM api_clients WHERE id = ?", + state.db_backend, + ); + + type GetRow = ( + String, + String, + String, + String, + String, + String, + String, + Option, + Option, + i32, + String, + String, + String, + ); + let row: Option = sqlx::query_as(&select_sql) + .bind(&id) + .fetch_optional(&state.db) + .await + .map_err(|e| AppError::Internal(format!("failed to get api client: {e}")))?; + + let Some(( + id, + client_key, + name, + client_id_url, + client_uri, + redirect_uris_json, + scopes, + rate_limit_capacity, + rate_limit_refill_rate, + is_active, + created_by, + created_at, + updated_at, + )) = row + else { + return Err(AppError::NotFound(format!("api client '{id}' not found"))); + }; + + let redirect_uris: Vec = serde_json::from_str(&redirect_uris_json).unwrap_or_default(); + + Ok(Json(ApiClientSummary { + id, + client_key, + name, + client_id_url, + client_uri, + redirect_uris, + scopes, + rate_limit_capacity, + rate_limit_refill_rate, + is_active: is_active != 0, + created_by, + created_at, + updated_at, + })) +} + +/// PUT /admin/api-clients/:id — update an API client. +pub(super) async fn update_api_client( + State(state): State, + auth: UserAuth, + Path(id): Path, + Json(body): Json, +) -> Result { + auth.require(Permission::ApiClientsEdit).await?; + + // Read current values + let select_sql = adapt_sql( + "SELECT client_key, name, client_id_url, client_uri, redirect_uris, scopes, rate_limit_capacity, rate_limit_refill_rate, is_active FROM api_clients WHERE id = ?", + state.db_backend, + ); + + type UpdateRow = ( + String, + String, + String, + String, + String, + String, + Option, + Option, + i32, + ); + let row: Option = sqlx::query_as(&select_sql) + .bind(&id) + .fetch_optional(&state.db) + .await + .map_err(|e| AppError::Internal(format!("failed to get api client: {e}")))?; + + let Some(( + client_key, + cur_name, + client_id_url, + cur_client_uri, + cur_redirect_uris, + cur_scopes, + cur_capacity, + cur_refill, + cur_active, + )) = row + else { + return Err(AppError::NotFound(format!("api client '{id}' not found"))); + }; + + let name = body.name.unwrap_or(cur_name); + let client_uri = body.client_uri.unwrap_or(cur_client_uri); + let redirect_uris_json = body + .redirect_uris + .map(|uris| serde_json::to_string(&uris).unwrap_or_else(|_| "[]".to_string())) + .unwrap_or(cur_redirect_uris); + let scopes = body.scopes.unwrap_or(cur_scopes); + let capacity = body.rate_limit_capacity.unwrap_or(cur_capacity); + let refill_rate = body.rate_limit_refill_rate.unwrap_or(cur_refill); + let is_active = body + .is_active + .map(|a| if a { 1i32 } else { 0i32 }) + .unwrap_or(cur_active); + let now = now_rfc3339(); + + let update_sql = adapt_sql( + "UPDATE api_clients SET name = ?, client_uri = ?, redirect_uris = ?, scopes = ?, rate_limit_capacity = ?, rate_limit_refill_rate = ?, is_active = ?, updated_at = ? WHERE id = ?", + state.db_backend, + ); + + sqlx::query(&update_sql) + .bind(&name) + .bind(&client_uri) + .bind(&redirect_uris_json) + .bind(&scopes) + .bind(capacity) + .bind(refill_rate) + .bind(is_active) + .bind(&now) + .bind(&id) + .execute(&state.db) + .await + .map_err(|e| AppError::Internal(format!("failed to update api client: {e}")))?; + + // Re-register or remove from OAuth registry based on active status. + let oauth_params = crate::auth::client_registry::ApiClientOAuthParams { + plc_url: state.config.plc_url.clone(), + state_store: state.oauth_state_store.clone(), + session_store_pool: state.db.clone(), + db_backend: state.db_backend, + }; + if is_active != 0 { + let redirect_uris: Vec = + serde_json::from_str(&redirect_uris_json).unwrap_or_default(); + if let Err(e) = state.oauth.register_api_client( + &client_id_url, + &client_uri, + redirect_uris, + &scopes, + &oauth_params, + ) { + tracing::warn!(client_id = %client_id_url, error = %e, "OAuth client re-registration failed"); + } + } else { + state.oauth.remove(&client_id_url); + } + + // Update per-client rate limit config. + if is_active != 0 { + if let (Some(cap), Some(refill)) = (capacity, refill_rate) { + let global = state.rate_limiter.global_config(); + state.rate_limiter.register_client_config( + client_key, + crate::rate_limit::RateLimitConfig { + capacity: cap as u32, + refill_rate: refill, + default_query_cost: global.default_query_cost, + default_procedure_cost: global.default_procedure_cost, + default_proxy_cost: global.default_proxy_cost, + }, + ); + } else { + // Rate limit overrides were cleared — remove per-client config. + state.rate_limiter.remove_client_config(&client_key); + } + } else { + state.rate_limiter.remove_client_config(&client_key); + } + + log_event( + &state.db, + EventLog { + event_type: "api_client.updated".to_string(), + severity: Severity::Info, + actor_did: Some(auth.did.clone()), + subject: Some(id), + detail: serde_json::json!({}), + }, + state.db_backend, + ) + .await; + + Ok(StatusCode::NO_CONTENT) +} + +/// DELETE /admin/api-clients/:id — delete an API client. +pub(super) async fn delete_api_client( + State(state): State, + auth: UserAuth, + Path(id): Path, +) -> Result { + auth.require(Permission::ApiClientsDelete).await?; + + // Look up client_id_url and client_key before deleting so we can remove from registries. + let lookup_sql = adapt_sql( + "SELECT client_id_url, client_key FROM api_clients WHERE id = ?", + state.db_backend, + ); + let client_info: Option<(String, String)> = sqlx::query_as(&lookup_sql) + .bind(&id) + .fetch_optional(&state.db) + .await + .map_err(|e| AppError::Internal(format!("failed to look up api client: {e}")))?; + + let delete_sql = adapt_sql("DELETE FROM api_clients WHERE id = ?", state.db_backend); + + let result = sqlx::query(&delete_sql) + .bind(&id) + .execute(&state.db) + .await + .map_err(|e| AppError::Internal(format!("failed to delete api client: {e}")))?; + + if result.rows_affected() == 0 { + return Err(AppError::NotFound(format!("api client '{id}' not found"))); + } + + // Remove from OAuth registry and rate limiter. + if let Some((url, key)) = client_info { + state.oauth.remove(&url); + state.rate_limiter.remove_client_config(&key); + } + + log_event( + &state.db, + EventLog { + event_type: "api_client.deleted".to_string(), + severity: Severity::Info, + actor_did: Some(auth.did.clone()), + subject: Some(id), + detail: serde_json::json!({}), + }, + state.db_backend, + ) + .await; + + Ok(StatusCode::NO_CONTENT) +} + +#[cfg(test)] +mod tests { + #[test] + fn test_client_key_prefix() { + let mut random_bytes = [0u8; 16]; + rand::Rng::fill(&mut rand::rng(), &mut random_bytes); + let key = format!("hvc_{}", hex::encode(random_bytes)); + assert!(key.starts_with("hvc_")); + assert_eq!(key.len(), 4 + 32); // "hvc_" + 32 hex chars + } +} diff --git a/src/admin/mod.rs b/src/admin/mod.rs index 9e32d34..dd4c516 100644 --- a/src/admin/mod.rs +++ b/src/admin/mod.rs @@ -1,3 +1,4 @@ +mod api_clients; mod api_keys; pub(crate) mod auth; mod backfill; @@ -74,11 +75,6 @@ pub fn admin_routes(_state: AppState) -> Router { post(rate_limits::upsert).get(rate_limits::list), ) .route("/rate-limits/enabled", put(rate_limits::set_enabled)) - .route("/rate-limits/allowlist", post(rate_limits::add_allowlist)) - .route( - "/rate-limits/allowlist/{id}", - delete(rate_limits::remove_allowlist), - ) .route("/settings", get(settings::list)) .route( "/settings/logo", @@ -96,4 +92,14 @@ pub fn admin_routes(_state: AppState) -> Router { "/plugins/{id}/secrets", get(plugins::get_secrets).put(plugins::update_secrets), ) + .route( + "/api-clients", + post(api_clients::create_api_client).get(api_clients::list_api_clients), + ) + .route( + "/api-clients/{id}", + get(api_clients::get_api_client) + .put(api_clients::update_api_client) + .delete(api_clients::delete_api_client), + ) } diff --git a/src/admin/permissions.rs b/src/admin/permissions.rs index 100b844..6bd2525 100644 --- a/src/admin/permissions.rs +++ b/src/admin/permissions.rs @@ -76,6 +76,15 @@ pub enum Permission { PluginsCreate, #[serde(rename = "plugins:delete")] PluginsDelete, + + #[serde(rename = "api-clients:view")] + ApiClientsView, + #[serde(rename = "api-clients:create")] + ApiClientsCreate, + #[serde(rename = "api-clients:edit")] + ApiClientsEdit, + #[serde(rename = "api-clients:delete")] + ApiClientsDelete, } impl Permission { @@ -112,10 +121,14 @@ impl Permission { Self::PluginsRead => "plugins:read", Self::PluginsCreate => "plugins:create", Self::PluginsDelete => "plugins:delete", + Self::ApiClientsView => "api-clients:view", + Self::ApiClientsCreate => "api-clients:create", + Self::ApiClientsEdit => "api-clients:edit", + Self::ApiClientsDelete => "api-clients:delete", } } - /// All 30 permissions. + /// All permissions. pub fn all() -> HashSet { HashSet::from([ Self::LexiconsCreate, @@ -148,6 +161,10 @@ impl Permission { Self::PluginsRead, Self::PluginsCreate, Self::PluginsDelete, + Self::ApiClientsView, + Self::ApiClientsCreate, + Self::ApiClientsEdit, + Self::ApiClientsDelete, ]) } } @@ -199,6 +216,10 @@ impl Template { perms.insert(Permission::PluginsRead); perms.insert(Permission::PluginsCreate); perms.insert(Permission::PluginsDelete); + perms.insert(Permission::ApiClientsView); + perms.insert(Permission::ApiClientsCreate); + perms.insert(Permission::ApiClientsEdit); + perms.insert(Permission::ApiClientsDelete); perms } Self::FullAccess => Permission::all(), diff --git a/src/admin/rate_limits.rs b/src/admin/rate_limits.rs index 5eeca35..e8a79dc 100644 --- a/src/admin/rate_limits.rs +++ b/src/admin/rate_limits.rs @@ -1,5 +1,5 @@ use axum::Json; -use axum::extract::{Path, State}; +use axum::extract::State; use axum::http::StatusCode; use crate::AppState; @@ -9,9 +9,7 @@ use crate::event_log::{EventLog, Severity, log_event}; use super::auth::UserAuth; use super::permissions::Permission; -use super::types::{ - AddAllowlistBody, AllowlistEntry, RateLimitsResponse, SetEnabledBody, UpsertRateLimitBody, -}; +use super::types::{RateLimitsResponse, SetEnabledBody, UpsertRateLimitBody}; /// GET /admin/rate-limits — list rate limit config. pub(super) async fn list( @@ -44,25 +42,6 @@ pub(super) async fn list( let (capacity, refill_rate, default_query_cost, default_procedure_cost, default_proxy_cost) = row.unwrap_or((100, 2.0, 1, 1, 1)); - let allowlist_sql = adapt_sql( - "SELECT id, cidr, note, created_at FROM rate_limit_allowlist ORDER BY id", - backend, - ); - let allowlist_rows: Vec<(i32, String, Option, String)> = sqlx::query_as(&allowlist_sql) - .fetch_all(&state.db) - .await - .map_err(|e| AppError::Internal(format!("failed to list allowlist: {e}")))?; - - let allowlist: Vec = allowlist_rows - .into_iter() - .map(|(id, cidr, note, created_at)| AllowlistEntry { - id, - cidr, - note, - created_at, - }) - .collect(); - Ok(Json(RateLimitsResponse { enabled: enabled == "true", capacity, @@ -70,7 +49,6 @@ pub(super) async fn list( default_query_cost, default_procedure_cost, default_proxy_cost, - allowlist, })) } @@ -178,103 +156,3 @@ pub(super) async fn set_enabled( Ok(StatusCode::NO_CONTENT) } - -/// POST /admin/rate-limits/allowlist — add an IP/CIDR to the allowlist. -pub(super) async fn add_allowlist( - State(state): State, - auth: UserAuth, - Json(body): Json, -) -> Result { - auth.require(Permission::RateLimitsCreate).await?; - - // Validate CIDR syntax; if it's a bare IP, append /32 or /128 - let cidr_str = if body.cidr.contains('/') { - body.cidr.clone() - } else if let Ok(ip) = body.cidr.parse::() { - match ip { - std::net::IpAddr::V4(_) => format!("{}/32", body.cidr), - std::net::IpAddr::V6(_) => format!("{}/128", body.cidr), - } - } else { - return Err(AppError::BadRequest(format!( - "invalid IP or CIDR: {}", - body.cidr - ))); - }; - - // Validate it parses as IpNet - if cidr_str.parse::().is_err() { - return Err(AppError::BadRequest(format!("invalid CIDR: {}", cidr_str))); - } - - let backend = state.db_backend; - let now = now_rfc3339(); - let sql = adapt_sql( - "INSERT INTO rate_limit_allowlist (cidr, note, created_at) VALUES (?, ?, ?)", - backend, - ); - sqlx::query(&sql) - .bind(&cidr_str) - .bind(&body.note) - .bind(&now) - .execute(&state.db) - .await - .map_err(|e| AppError::Internal(format!("failed to add allowlist entry: {e}")))?; - - state.rate_limiter.reload_from_db(&state.db).await; - - log_event( - &state.db, - EventLog { - event_type: "rate_limit.allowlist_added".to_string(), - severity: Severity::Info, - actor_did: Some(auth.did.clone()), - subject: Some(cidr_str), - detail: serde_json::json!({ "note": body.note }), - }, - state.db_backend, - ) - .await; - - Ok(StatusCode::CREATED) -} - -/// DELETE /admin/rate-limits/allowlist/{id} — remove an allowlist entry. -pub(super) async fn remove_allowlist( - State(state): State, - auth: UserAuth, - Path(id): Path, -) -> Result { - auth.require(Permission::RateLimitsDelete).await?; - - let backend = state.db_backend; - let sql = adapt_sql("DELETE FROM rate_limit_allowlist WHERE id = ?", backend); - let result = sqlx::query(&sql) - .bind(id) - .execute(&state.db) - .await - .map_err(|e| AppError::Internal(format!("failed to delete allowlist entry: {e}")))?; - - if result.rows_affected() == 0 { - return Err(AppError::NotFound(format!( - "allowlist entry {id} not found" - ))); - } - - state.rate_limiter.reload_from_db(&state.db).await; - - log_event( - &state.db, - EventLog { - event_type: "rate_limit.allowlist_removed".to_string(), - severity: Severity::Info, - actor_did: Some(auth.did.clone()), - subject: Some(id.to_string()), - detail: serde_json::json!({}), - }, - state.db_backend, - ) - .await; - - Ok(StatusCode::NO_CONTENT) -} diff --git a/src/admin/types.rs b/src/admin/types.rs index 9fc456d..21e4de4 100644 --- a/src/admin/types.rs +++ b/src/admin/types.rs @@ -296,6 +296,62 @@ pub(super) struct UpdatePluginSecretsBody { pub(super) secrets: std::collections::HashMap, } +// --------------------------------------------------------------------------- +// API client types +// --------------------------------------------------------------------------- + +#[derive(Deserialize)] +pub(super) struct CreateApiClientBody { + pub(super) name: String, + pub(super) client_id_url: String, + pub(super) client_uri: String, + pub(super) redirect_uris: Vec, + #[serde(default = "default_scopes")] + pub(super) scopes: String, + pub(super) rate_limit_capacity: Option, + pub(super) rate_limit_refill_rate: Option, +} + +fn default_scopes() -> String { + "atproto".to_string() +} + +#[derive(Deserialize)] +pub(super) struct UpdateApiClientBody { + pub(super) name: Option, + pub(super) client_uri: Option, + pub(super) redirect_uris: Option>, + pub(super) scopes: Option, + pub(super) rate_limit_capacity: Option>, + pub(super) rate_limit_refill_rate: Option>, + pub(super) is_active: Option, +} + +#[derive(Serialize)] +pub(super) struct ApiClientSummary { + pub(super) id: String, + pub(super) client_key: String, + pub(super) name: String, + pub(super) client_id_url: String, + pub(super) client_uri: String, + pub(super) redirect_uris: Vec, + pub(super) scopes: String, + pub(super) rate_limit_capacity: Option, + pub(super) rate_limit_refill_rate: Option, + pub(super) is_active: bool, + pub(super) created_by: String, + pub(super) created_at: String, + pub(super) updated_at: String, +} + +#[derive(Serialize)] +pub(super) struct CreateApiClientResponse { + pub(super) id: String, + pub(super) client_key: String, + pub(super) name: String, + pub(super) client_id_url: String, +} + // --------------------------------------------------------------------------- // Rate limit types // --------------------------------------------------------------------------- @@ -314,12 +370,6 @@ pub(super) struct SetEnabledBody { pub(super) enabled: bool, } -#[derive(Deserialize)] -pub(super) struct AddAllowlistBody { - pub(super) cidr: String, - pub(super) note: Option, -} - #[derive(Serialize)] pub(super) struct RateLimitsResponse { pub(super) enabled: bool, @@ -328,13 +378,4 @@ pub(super) struct RateLimitsResponse { pub(super) default_query_cost: i32, pub(super) default_procedure_cost: i32, pub(super) default_proxy_cost: i32, - pub(super) allowlist: Vec, -} - -#[derive(Serialize)] -pub(super) struct AllowlistEntry { - pub(super) id: i32, - pub(super) cidr: String, - pub(super) note: Option, - pub(super) created_at: String, } diff --git a/src/auth/client_registry.rs b/src/auth/client_registry.rs new file mode 100644 index 0000000..e08f6e4 --- /dev/null +++ b/src/auth/client_registry.rs @@ -0,0 +1,250 @@ +use dashmap::DashMap; +use std::sync::Arc; + +use atrium_identity::did::{CommonDidResolver, CommonDidResolverConfig}; +use atrium_identity::handle::{AtprotoHandleResolver, AtprotoHandleResolverConfig}; +use atrium_oauth::{ + AtprotoClientMetadata, AuthMethod, DefaultHttpClient, GrantType, OAuthClientConfig, + OAuthResolverConfig, +}; + +use crate::HappyViewOAuthClient; +use crate::auth::oauth_store::{DbSessionStore, DbStateStore}; +use crate::db::{DatabaseBackend, adapt_sql}; +use crate::dns::NativeDnsResolver; + +/// Parameters needed to build an OAuth client for an API client registration. +pub struct ApiClientOAuthParams { + pub plc_url: String, + pub state_store: DbStateStore, + pub session_store_pool: sqlx::AnyPool, + pub db_backend: DatabaseBackend, +} + +/// Registry of OAuth clients, keyed by `client_id_url`. +/// +/// Each API client gets its own `OAuthClient` instance so the PDS auth screen +/// shows the correct domain. The default client is HappyView's own identity, +/// used for dashboard auth. +pub struct OAuthClientRegistry { + default_client: Arc, + clients: DashMap>, +} + +impl OAuthClientRegistry { + pub fn new(default_client: Arc) -> Self { + Self { + default_client, + clients: DashMap::new(), + } + } + + /// Register an API client's OAuth client, keyed by its `client_id_url`. + pub fn register(&self, client_id_url: String, client: Arc) { + self.clients.insert(client_id_url, client); + } + + /// Remove an API client's OAuth client. + pub fn remove(&self, client_id_url: &str) { + self.clients.remove(client_id_url); + } + + /// Look up a client by `client_id_url`. + pub fn get(&self, client_id_url: &str) -> Option> { + self.clients.get(client_id_url).map(|r| r.value().clone()) + } + + /// Look up a client by `client_id_url`, falling back to the default. + pub fn get_or_default(&self, client_id_url: Option<&str>) -> Arc { + if let Some(url) = client_id_url { + self.clients + .get(url) + .map(|r| r.value().clone()) + .unwrap_or_else(|| self.default_client.clone()) + } else { + self.default_client.clone() + } + } + + /// Get the default (HappyView dashboard) client. + pub fn default_client(&self) -> &Arc { + &self.default_client + } + + /// Build and register a single OAuth client from API client metadata. + /// Used when creating or updating an API client via the admin UI. + pub fn register_api_client( + &self, + client_id_url: &str, + client_uri: &str, + redirect_uris: Vec, + scopes_str: &str, + params: &ApiClientOAuthParams, + ) -> Result<(), String> { + let ApiClientOAuthParams { + plc_url, + state_store, + session_store_pool, + db_backend, + } = params; + let scopes = crate::auth::parse_scope_string(scopes_str); + let scopes = if scopes.is_empty() { + vec![atrium_oauth::Scope::Known( + atrium_oauth::KnownScope::Atproto, + )] + } else { + scopes + }; + + let metadata = AtprotoClientMetadata { + client_id: client_id_url.to_string(), + client_uri: Some(client_uri.to_string()), + redirect_uris, + token_endpoint_auth_method: AuthMethod::None, + grant_types: vec![GrantType::AuthorizationCode, GrantType::RefreshToken], + scopes, + jwks_uri: None, + token_endpoint_auth_signing_alg: None, + }; + + let http = Arc::new(DefaultHttpClient::default()); + let resolver = OAuthResolverConfig { + did_resolver: CommonDidResolver::new(CommonDidResolverConfig { + plc_directory_url: plc_url.to_string(), + http_client: Arc::clone(&http), + }), + handle_resolver: AtprotoHandleResolver::new(AtprotoHandleResolverConfig { + dns_txt_resolver: NativeDnsResolver::new(), + http_client: Arc::clone(&http), + }), + authorization_server_metadata: Default::default(), + protected_resource_metadata: Default::default(), + }; + + match atrium_oauth::OAuthClient::new(OAuthClientConfig { + client_metadata: metadata, + keys: None, + state_store: state_store.clone(), + session_store: DbSessionStore::new(session_store_pool.clone(), *db_backend), + resolver, + }) { + Ok(client) => { + self.register(client_id_url.to_string(), Arc::new(client)); + Ok(()) + } + Err(e) => Err(format!("failed to create OAuth client: {e}")), + } + } + + /// Load all active API clients from the database and register OAuth clients for each. + pub async fn load_from_db( + &self, + db: &sqlx::AnyPool, + db_backend: DatabaseBackend, + plc_url: &str, + state_store: DbStateStore, + session_store_pool: sqlx::AnyPool, + ) { + let sql = adapt_sql( + "SELECT client_id_url, client_uri, redirect_uris, scopes FROM api_clients WHERE is_active = 1", + db_backend, + ); + + let rows: Vec<(String, String, String, String)> = + match sqlx::query_as(&sql).fetch_all(db).await { + Ok(r) => r, + Err(e) => { + tracing::error!("Failed to load API clients from database: {e}"); + return; + } + }; + + for (client_id_url, client_uri, redirect_uris_json, scopes_str) in rows { + let redirect_uris: Vec = + serde_json::from_str(&redirect_uris_json).unwrap_or_default(); + + let scopes = crate::auth::parse_scope_string(&scopes_str); + let scopes = if scopes.is_empty() { + vec![atrium_oauth::Scope::Known( + atrium_oauth::KnownScope::Atproto, + )] + } else { + scopes + }; + + let metadata = AtprotoClientMetadata { + client_id: client_id_url.clone(), + client_uri: Some(client_uri), + redirect_uris, + token_endpoint_auth_method: AuthMethod::None, + grant_types: vec![GrantType::AuthorizationCode, GrantType::RefreshToken], + scopes, + jwks_uri: None, + token_endpoint_auth_signing_alg: None, + }; + + // Each OAuthClient needs its own resolver instances (they're not Clone) + let http = Arc::new(DefaultHttpClient::default()); + let resolver = OAuthResolverConfig { + did_resolver: CommonDidResolver::new(CommonDidResolverConfig { + plc_directory_url: plc_url.to_string(), + http_client: Arc::clone(&http), + }), + handle_resolver: AtprotoHandleResolver::new(AtprotoHandleResolverConfig { + dns_txt_resolver: NativeDnsResolver::new(), + http_client: Arc::clone(&http), + }), + authorization_server_metadata: Default::default(), + protected_resource_metadata: Default::default(), + }; + + match atrium_oauth::OAuthClient::new(OAuthClientConfig { + client_metadata: metadata, + keys: None, + state_store: state_store.clone(), + session_store: DbSessionStore::new(session_store_pool.clone(), db_backend), + resolver, + }) { + Ok(client) => { + tracing::info!(client_id = %client_id_url, "Registered API client OAuth identity"); + self.register(client_id_url, Arc::new(client)); + } + Err(e) => { + tracing::error!(client_id = %client_id_url, error = %e, "Failed to create OAuth client for API client"); + } + } + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + // Note: we can't easily construct real OAuthClient instances in unit tests + // because they require resolvers, stores, etc. The registry logic is simple + // enough that we test it via integration tests that stand up the full stack. + // These tests verify the DashMap-based lookup logic using a mock approach. + + #[test] + fn test_registry_stores_and_retrieves() { + // We can at least verify the DashMap operations work correctly + let map: DashMap = DashMap::new(); + map.insert("key1".to_string(), "val1".to_string()); + + assert!(map.get("key1").is_some()); + assert!(map.get("key2").is_none()); + + map.remove("key1"); + assert!(map.get("key1").is_none()); + } + + #[test] + fn test_registry_overwrite() { + let map: DashMap = DashMap::new(); + map.insert("key1".to_string(), "val1".to_string()); + map.insert("key1".to_string(), "val2".to_string()); + + assert_eq!(map.get("key1").unwrap().value(), "val2"); + } +} diff --git a/src/auth/middleware.rs b/src/auth/middleware.rs index 358c3e9..d64de96 100644 --- a/src/auth/middleware.rs +++ b/src/auth/middleware.rs @@ -15,18 +15,32 @@ use crate::error::AppError; #[derive(Debug, Clone)] pub struct Claims { did: String, + /// The API client key (e.g. "hvc_...") if the user authenticated via an API client. + client_key: Option, } +/// Separator used to encode `did` and `client_key` in a single cookie value. +/// Newlines cannot appear in DIDs or client keys, so this is safe. +const COOKIE_SEP: char = '\n'; + impl Claims { /// The authenticated user's DID. pub fn did(&self) -> &str { &self.did } + /// The API client key, if the user logged in via an API client. + pub fn client_key(&self) -> Option<&str> { + self.client_key.as_deref() + } + /// Test-only constructor. #[cfg(test)] pub fn new_for_test(did: String) -> Self { - Self { did } + Self { + did, + client_key: None, + } } } @@ -43,8 +57,13 @@ impl FromRequestParts for Claims { .map_err(|_| AppError::Auth("failed to read cookies".into()))?; if let Some(cookie) = jar.get(COOKIE_NAME) { - let did = cookie.value().to_string(); - return Ok(Claims { did }); + let value = cookie.value().to_string(); + let (did, client_key) = if let Some((d, k)) = value.split_once(COOKIE_SEP) { + (d.to_string(), Some(k.to_string())) + } else { + (value, None) + }; + return Ok(Claims { did, client_key }); } // Path 2: Authorization header @@ -63,13 +82,17 @@ impl FromRequestParts for Claims { // API key auth is handled by UserAuth extractor which looks up the key. // We need to extract the DID from the api_keys table. let did = resolve_api_key_did(state, token).await?; - return Ok(Claims { did }); + return Ok(Claims { + did, + client_key: None, + }); } // Otherwise, try service auth JWT let service_auth = super::service_auth::ServiceAuth::from_bearer(token, state).await?; return Ok(Claims { did: service_auth.did, + client_key: None, }); } diff --git a/src/auth/mod.rs b/src/auth/mod.rs index af65fce..538f142 100644 --- a/src/auth/mod.rs +++ b/src/auth/mod.rs @@ -1,8 +1,10 @@ +pub mod client_registry; pub mod middleware; pub mod oauth_store; pub mod routes; pub mod service_auth; +pub use client_registry::OAuthClientRegistry; pub use middleware::Claims; pub use routes::parse_scope_string; pub use service_auth::ServiceAuth; diff --git a/src/auth/routes.rs b/src/auth/routes.rs index c41ba38..f5c32db 100644 --- a/src/auth/routes.rs +++ b/src/auth/routes.rs @@ -21,6 +21,7 @@ pub struct LoginQuery { handle: String, redirect_uri: Option, scope: Option, + client_id: Option, } /// Parse a whitespace-separated OAuth scope string into typed `Scope` values. @@ -85,7 +86,10 @@ async fn login( } }; - tracing::debug!(scopes = ?scopes, "resolved oauth scopes"); + tracing::debug!(scopes = ?scopes, client_id = ?query.client_id, "resolved oauth scopes"); + + // Select the appropriate OAuth client based on client_id + let oauth_client = state.oauth.get_or_default(query.client_id.as_deref()); // Hold the authorize lock so that authorize() + take_last_state_key() are atomic. // This prevents concurrent logins from swapping each other's state keys. @@ -96,8 +100,7 @@ async fn login( ..Default::default() }; - let url = state - .oauth + let url = oauth_client .authorize(&query.handle, options) .await .map_err(|e| AppError::Internal(format!("OAuth authorize failed: {e}")))?; @@ -113,19 +116,22 @@ async fn login( // Store the redirect URI in the database, keyed by the OAuth state parameter. // This avoids third-party cookie issues when Pentaract (cross-origin) calls this endpoint. - if let Some(redirect_uri) = &query.redirect_uri { - tracing::debug!(oauth_state = ?oauth_state, redirect_uri = %redirect_uri, "storing redirect for state"); + // Store redirect URI and client_id for the callback to use + if query.redirect_uri.is_some() || query.client_id.is_some() { + let redirect_uri = query.redirect_uri.as_deref().unwrap_or(""); + tracing::debug!(oauth_state = ?oauth_state, redirect_uri = %redirect_uri, client_id = ?query.client_id, "storing redirect for state"); if let Some(oauth_state) = oauth_state { let now = now_rfc3339(); let expires_at = (chrono::Utc::now() + chrono::Duration::minutes(10)).to_rfc3339(); let sql = adapt_sql( - "INSERT INTO auth_login_redirects (state, redirect_uri, created_at, expires_at) VALUES (?, ?, ?, ?)", + "INSERT INTO auth_login_redirects (state, redirect_uri, client_id, created_at, expires_at) VALUES (?, ?, ?, ?, ?)", state.db_backend, ); let _ = sqlx::query(&sql) .bind(&oauth_state) .bind(redirect_uri) + .bind(query.client_id.as_deref()) .bind(&now) .bind(&expires_at) .execute(&state.db) @@ -145,14 +151,14 @@ async fn callback( ) -> Result<(SignedCookieJar, Redirect), AppError> { tracing::debug!(state = ?query.state, "callback received"); - // Look up the redirect URI from the database before the OAuth library consumes the state - let redirect_url = if let Some(oauth_state) = &query.state { + // Look up the redirect URI and client_id from the database before the OAuth library consumes the state + let (redirect_url, client_id) = if let Some(oauth_state) = &query.state { let sql = adapt_sql( - "SELECT redirect_uri FROM auth_login_redirects WHERE state = ? AND expires_at > ?", + "SELECT redirect_uri, client_id FROM auth_login_redirects WHERE state = ? AND expires_at > ?", state.db_backend, ); let now = now_rfc3339(); - let row: Option<(String,)> = sqlx::query_as(&sql) + let row: Option<(String, Option)> = sqlx::query_as(&sql) .bind(oauth_state) .bind(&now) .fetch_optional(&state.db) @@ -172,20 +178,28 @@ async fn callback( } tracing::debug!(found_redirect = ?row, "redirect lookup result"); - row.map(|(uri,)| uri) + match row { + Some((uri, cid)) => { + let uri = if uri.is_empty() { None } else { Some(uri) }; + (uri, cid) + } + None => (None, None), + } } else { tracing::debug!("no state in callback query"); - None + (None, None) }; + // Use the same OAuth client that was used for authorize + let oauth_client = state.oauth.get_or_default(client_id.as_deref()); + let params = atrium_oauth::CallbackParams { code: query.code, state: query.state, iss: query.iss, }; - let (session, _app_state) = state - .oauth + let (session, _app_state) = oauth_client .callback(params) .await .map_err(|e| AppError::Internal(format!("OAuth callback failed: {e}")))?; @@ -196,13 +210,37 @@ async fn callback( .await .ok_or_else(|| AppError::Internal("no DID in OAuth session".into()))?; + // Look up the client_key for the API client so we can store it in the session cookie + // for per-client rate limiting. + let client_key = if let Some(ref cid) = client_id { + let sql = adapt_sql( + "SELECT client_key FROM api_clients WHERE client_id_url = ? AND is_active = 1", + state.db_backend, + ); + let row: Option<(String,)> = sqlx::query_as(&sql) + .bind(cid) + .fetch_optional(&state.db) + .await + .unwrap_or(None); + row.map(|(k,)| k) + } else { + None + }; + // Use DB-stored redirect, or default to "/" - let redirect_url = redirect_url.unwrap_or_else(|| "/".to_string()); + let redirect_url = redirect_url.unwrap_or_else(|| "/".into()); tracing::debug!(redirect_url = %redirect_url, "redirecting after callback"); // Set the session cookie // Must use SameSite=None for cross-origin requests (e.g., Pentaract calling HappyView) - let mut session_cookie = Cookie::new(COOKIE_NAME, did.to_string()); + // Encode did and optional client_key separated by newline. + let did_str = did.as_ref(); + let cookie_value = if let Some(ref ck) = client_key { + format!("{did_str}\n{ck}") + } else { + did_str.to_string() + }; + let mut session_cookie = Cookie::new(COOKIE_NAME, cookie_value); session_cookie.set_path("/"); session_cookie.set_http_only(true); session_cookie.set_same_site(axum_extra::extract::cookie::SameSite::None); @@ -227,9 +265,10 @@ async fn logout( jar: SignedCookieJar, ) -> Result, AppError> { if let Some(cookie) = jar.get(COOKIE_NAME) { - let did_str = cookie.value().to_string(); + let raw = cookie.value().to_string(); + let did_str = raw.split('\n').next().unwrap_or(&raw).to_string(); if let Ok(did) = atrium_api::types::string::Did::new(did_str) { - let _ = state.oauth.revoke(&did).await; + let _ = state.oauth.default_client().revoke(&did).await; } } @@ -254,7 +293,8 @@ async fn me( let cookie = jar .get(COOKIE_NAME) .ok_or(AppError::Auth("not authenticated".into()))?; - let did = cookie.value().to_string(); + let raw = cookie.value().to_string(); + let did = raw.split('\n').next().unwrap_or(&raw).to_string(); let backend = state.db_backend; let user: Option<(i32,)> = diff --git a/src/lib.rs b/src/lib.rs index b69c398..a80346b 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -57,7 +57,7 @@ pub struct AppState { pub collections_tx: watch::Sender>, pub labeler_subscriptions_tx: watch::Sender<()>, pub rate_limiter: Arc, - pub oauth: Arc, + pub oauth: Arc, pub oauth_state_store: DbStateStore, pub cookie_key: axum_extra::extract::cookie::Key, pub plugin_registry: Arc, diff --git a/src/lua/atproto_api.rs b/src/lua/atproto_api.rs index 9573d89..4491830 100644 --- a/src/lua/atproto_api.rs +++ b/src/lua/atproto_api.rs @@ -346,9 +346,10 @@ mod tests { default_procedure_cost: 1, default_proxy_cost: 1, }, - vec![], ), - oauth: std::sync::Arc::new(oauth), + oauth: std::sync::Arc::new(crate::auth::OAuthClientRegistry::new(std::sync::Arc::new( + oauth, + ))), oauth_state_store: crate::auth::oauth_store::DbStateStore::new( test_db.clone(), crate::db::DatabaseBackend::Sqlite, diff --git a/src/lua/db_api.rs b/src/lua/db_api.rs index a0614c8..bbc7f8d 100644 --- a/src/lua/db_api.rs +++ b/src/lua/db_api.rs @@ -697,9 +697,10 @@ mod tests { default_procedure_cost: 1, default_proxy_cost: 1, }, - vec![], ), - oauth: std::sync::Arc::new(oauth), + oauth: std::sync::Arc::new(crate::auth::OAuthClientRegistry::new(std::sync::Arc::new( + oauth, + ))), oauth_state_store: crate::auth::oauth_store::DbStateStore::new( test_db.clone(), crate::db::DatabaseBackend::Sqlite, diff --git a/src/lua/execute.rs b/src/lua/execute.rs index 2f3f7e4..f4ac7ae 100644 --- a/src/lua/execute.rs +++ b/src/lua/execute.rs @@ -1023,9 +1023,10 @@ mod tests { default_procedure_cost: 1, default_proxy_cost: 1, }, - vec![], ), - oauth: std::sync::Arc::new(oauth), + oauth: std::sync::Arc::new(crate::auth::OAuthClientRegistry::new(std::sync::Arc::new( + oauth, + ))), oauth_state_store: crate::auth::oauth_store::DbStateStore::new( test_db.clone(), crate::db::DatabaseBackend::Sqlite, diff --git a/src/lua/http_api.rs b/src/lua/http_api.rs index af8c324..b33b80e 100644 --- a/src/lua/http_api.rs +++ b/src/lua/http_api.rs @@ -163,9 +163,10 @@ mod tests { default_procedure_cost: 1, default_proxy_cost: 1, }, - vec![], ), - oauth: std::sync::Arc::new(oauth), + oauth: std::sync::Arc::new(crate::auth::OAuthClientRegistry::new(std::sync::Arc::new( + oauth, + ))), oauth_state_store: crate::auth::oauth_store::DbStateStore::new( test_db.clone(), crate::db::DatabaseBackend::Sqlite, diff --git a/src/main.rs b/src/main.rs index 5890403..670ad1e 100644 --- a/src/main.rs +++ b/src/main.rs @@ -5,7 +5,7 @@ use happyview::config::Config; use happyview::db; use happyview::dns::NativeDnsResolver; use happyview::lexicon::{LexiconRegistry, ParsedLexicon, ProcedureAction}; -use happyview::rate_limit::RateLimiter; +use happyview::rate_limit::{RateLimitConfig, RateLimiter}; use happyview::resolve::{fetch_lexicon_from_pds, resolve_nsid_authority}; use happyview::{AppState, jetstream, labeler, server}; use tokio::sync::watch; @@ -266,9 +266,33 @@ async fn main() { // Initialize rate limiter from DB. let rl_state = RateLimiter::load_from_db(&db_pool).await; - let rate_limiter = RateLimiter::new(rl_state.enabled, rl_state.global, rl_state.allowlist); + let rate_limiter = RateLimiter::new(rl_state.enabled, rl_state.global); tokio::spawn(rate_limiter.clone().spawn_cleanup()); + // Load per-client rate limit configs from api_clients table. + { + let client_configs: Vec<(String, i32, f64)> = sqlx::query_as( + "SELECT client_key, rate_limit_capacity, rate_limit_refill_rate FROM api_clients WHERE is_active = 1 AND rate_limit_capacity IS NOT NULL AND rate_limit_refill_rate IS NOT NULL", + ) + .fetch_all(&db_pool) + .await + .unwrap_or_default(); + + let global = rate_limiter.global_config(); + for (client_key, capacity, refill_rate) in client_configs { + rate_limiter.register_client_config( + client_key, + RateLimitConfig { + capacity: capacity as u32, + refill_rate, + default_query_cost: global.default_query_cost, + default_procedure_cost: global.default_procedure_cost, + default_proxy_cost: global.default_proxy_cost, + }, + ); + } + } + // Build atrium-oauth client let dns = NativeDnsResolver::new(); let callback_url = format!("{}/auth/callback", config.public_url.trim_end_matches('/')); @@ -370,6 +394,20 @@ async fn main() { let (collections_tx, collections_rx) = watch::channel(initial_collections); let (labeler_subscriptions_tx, labeler_subscriptions_rx) = watch::channel(()); + // Build the OAuth client registry and load API clients from DB + let oauth_registry = Arc::new(happyview::auth::OAuthClientRegistry::new(Arc::new( + oauth_client, + ))); + oauth_registry + .load_from_db( + &db_pool, + db_backend, + &config.plc_url, + oauth_state_store.clone(), + db_pool.clone(), + ) + .await; + let state = AppState { config: config.clone(), http, @@ -379,7 +417,7 @@ async fn main() { collections_tx, labeler_subscriptions_tx, rate_limiter, - oauth: Arc::new(oauth_client), + oauth: oauth_registry, oauth_state_store, cookie_key, plugin_registry, diff --git a/src/rate_limit.rs b/src/rate_limit.rs index 1990ba5..b3a9e03 100644 --- a/src/rate_limit.rs +++ b/src/rate_limit.rs @@ -1,8 +1,6 @@ use arc_swap::ArcSwap; use dashmap::DashMap; -use ipnet::IpNet; use sqlx::AnyPool; -use std::net::IpAddr; use std::sync::Arc; use std::sync::atomic::{AtomicBool, Ordering}; use std::time::{Instant, SystemTime, UNIX_EPOCH}; @@ -41,13 +39,13 @@ pub struct RateLimiter { enabled: AtomicBool, buckets: DashMap, global_config: ArcSwap, - allowlist: ArcSwap>, + /// Per-client config overrides, keyed by client_key (e.g. "hvc_...") + client_configs: DashMap, } pub struct RateLimiterState { pub enabled: bool, pub global: RateLimitConfig, - pub allowlist: Vec, } fn now_unix() -> u64 { @@ -58,32 +56,27 @@ fn now_unix() -> u64 { } impl RateLimiter { - pub fn new(enabled: bool, global: RateLimitConfig, allowlist: Vec) -> Arc { + pub fn new(enabled: bool, global: RateLimitConfig) -> Arc { Arc::new(Self { enabled: AtomicBool::new(enabled), buckets: DashMap::new(), global_config: ArcSwap::new(Arc::new(global)), - allowlist: ArcSwap::new(Arc::new(allowlist)), + client_configs: DashMap::new(), }) } - pub fn check(&self, key: &str, cost: u32, client_ip: Option) -> CheckResult { + pub fn check(&self, key: &str, cost: u32) -> CheckResult { if !self.enabled.load(Ordering::Relaxed) { return CheckResult::Disabled; } - if let Some(ip) = client_ip { - let list = self.allowlist.load(); - for net in list.iter() { - if net.contains(&ip) { - return CheckResult::Disabled; - } - } - } - - let global = self.global_config.load(); - let capacity = global.capacity; - let refill_rate = global.refill_rate; + // Use per-client config if available, otherwise fall back to global + let (capacity, refill_rate) = if let Some(client_cfg) = self.client_configs.get(key) { + (client_cfg.capacity, client_cfg.refill_rate) + } else { + let global = self.global_config.load(); + (global.capacity, global.refill_rate) + }; let cost_f64 = cost as f64; let now = Instant::now(); @@ -144,6 +137,11 @@ impl RateLimiter { } } + /// Get a snapshot of the current global config. + pub fn global_config(&self) -> Arc { + self.global_config.load_full() + } + pub fn set_enabled(&self, enabled: bool) { self.enabled.store(enabled, Ordering::Relaxed); } @@ -156,8 +154,14 @@ impl RateLimiter { self.global_config.store(Arc::new(global)); } - pub fn update_allowlist(&self, entries: Vec) { - self.allowlist.store(Arc::new(entries)); + /// Register a per-client rate limit config override. + pub fn register_client_config(&self, client_key: String, config: RateLimitConfig) { + self.client_configs.insert(client_key, config); + } + + /// Remove a per-client rate limit config override. + pub fn remove_client_config(&self, client_key: &str) { + self.client_configs.remove(client_key); } pub async fn spawn_cleanup(self: Arc) { @@ -210,22 +214,7 @@ impl RateLimiter { }, }; - // Load allowlist - let cidr_rows: Vec<(String,)> = sqlx::query_as("SELECT cidr FROM rate_limit_allowlist") - .fetch_all(db) - .await - .unwrap_or_default(); - - let allowlist: Vec = cidr_rows - .into_iter() - .filter_map(|(cidr,)| cidr.parse().ok()) - .collect(); - - RateLimiterState { - enabled, - global, - allowlist, - } + RateLimiterState { enabled, global } } /// Reload all config from DB and apply to the live limiter. @@ -233,7 +222,6 @@ impl RateLimiter { let state = Self::load_from_db(db).await; self.set_enabled(state.enabled); self.update_config(state.global); - self.update_allowlist(state.allowlist); } } @@ -252,21 +240,14 @@ mod tests { default_procedure_cost: 1, default_proxy_cost: 1, }, - vec![], ); // Should allow 3 requests (bucket starts full, cost=1 each) for _ in 0..3 { - assert!(matches!( - rl.check("k", 1, None), - CheckResult::Allowed { .. } - )); + assert!(matches!(rl.check("k", 1), CheckResult::Allowed { .. })); } // 4th should be limited - assert!(matches!( - rl.check("k", 1, None), - CheckResult::Limited { .. } - )); + assert!(matches!(rl.check("k", 1), CheckResult::Limited { .. })); } #[test] @@ -280,23 +261,19 @@ mod tests { default_procedure_cost: 1, default_proxy_cost: 1, }, - vec![], ); // Cost of 5 should allow 2 requests (10 tokens total) assert!(matches!( - rl.check("k", 5, None), + rl.check("k", 5), CheckResult::Allowed { remaining: 5, .. } )); assert!(matches!( - rl.check("k", 5, None), + rl.check("k", 5), CheckResult::Allowed { remaining: 0, .. } )); // 3rd should be limited - assert!(matches!( - rl.check("k", 5, None), - CheckResult::Limited { .. } - )); + assert!(matches!(rl.check("k", 5), CheckResult::Limited { .. })); } #[test] @@ -310,15 +287,122 @@ mod tests { default_procedure_cost: 1, default_proxy_cost: 1, }, - vec![], ); - assert!(matches!(rl.check("k", 1, None), CheckResult::Disabled)); + assert!(matches!(rl.check("k", 1), CheckResult::Disabled)); + } + + #[test] + fn default_cost_for_type() { + let rl = RateLimiter::new( + true, + RateLimitConfig { + capacity: 100, + refill_rate: 10.0, + default_query_cost: 2, + default_procedure_cost: 5, + default_proxy_cost: 3, + }, + ); + + assert_eq!(rl.default_cost_for_type("query"), 2); + assert_eq!(rl.default_cost_for_type("procedure"), 5); + assert_eq!(rl.default_cost_for_type("proxy"), 3); + assert_eq!(rl.default_cost_for_type("unknown"), 1); + } + + #[test] + fn per_client_config_override() { + let rl = RateLimiter::new( + true, + RateLimitConfig { + capacity: 10, + refill_rate: 1.0, + default_query_cost: 1, + default_procedure_cost: 1, + default_proxy_cost: 1, + }, + ); + + // Register a client with lower capacity + rl.register_client_config( + "hvc_client1".to_string(), + RateLimitConfig { + capacity: 2, + refill_rate: 0.001, + default_query_cost: 1, + default_procedure_cost: 1, + default_proxy_cost: 1, + }, + ); + + // Client key should use client config (capacity=2) + assert!(matches!( + rl.check("hvc_client1", 1), + CheckResult::Allowed { .. } + )); + assert!(matches!( + rl.check("hvc_client1", 1), + CheckResult::Allowed { .. } + )); + assert!(matches!( + rl.check("hvc_client1", 1), + CheckResult::Limited { .. } + )); + + // Other key should use global config (capacity=10) + for _ in 0..10 { + assert!(matches!( + rl.check("other_key", 1), + CheckResult::Allowed { .. } + )); + } + assert!(matches!( + rl.check("other_key", 1), + CheckResult::Limited { .. } + )); + } + + #[test] + fn per_client_config_fallback_to_global() { + let rl = RateLimiter::new( + true, + RateLimitConfig { + capacity: 3, + refill_rate: 1.0, + default_query_cost: 1, + default_procedure_cost: 1, + default_proxy_cost: 1, + }, + ); + + // No client config registered — should use global (capacity=3) + for _ in 0..3 { + assert!(matches!( + rl.check("hvc_unregistered", 1), + CheckResult::Allowed { .. } + )); + } + assert!(matches!( + rl.check("hvc_unregistered", 1), + CheckResult::Limited { .. } + )); } #[test] - fn allowlisted_ip_bypasses() { + fn register_and_remove_client_config() { let rl = RateLimiter::new( true, + RateLimitConfig { + capacity: 10, + refill_rate: 1.0, + default_query_cost: 1, + default_procedure_cost: 1, + default_proxy_cost: 1, + }, + ); + + rl.register_client_config( + "hvc_temp".to_string(), RateLimitConfig { capacity: 1, refill_rate: 0.001, @@ -326,32 +410,64 @@ mod tests { default_procedure_cost: 1, default_proxy_cost: 1, }, - vec!["10.0.0.0/8".parse().unwrap()], ); - let ip: IpAddr = "10.0.0.5".parse().unwrap(); - // Even after exhausting, allowlisted IP gets Disabled - assert!(matches!(rl.check("k", 1, Some(ip)), CheckResult::Disabled)); + // Should be limited after 1 request (client config capacity=1) + assert!(matches!( + rl.check("hvc_temp", 1), + CheckResult::Allowed { .. } + )); + assert!(matches!( + rl.check("hvc_temp", 1), + CheckResult::Limited { .. } + )); + + // Remove client config — new bucket should use global (capacity=10) + rl.remove_client_config("hvc_temp"); + // Note: the old bucket still exists and is exhausted, but capacity was + // updated to global. A new bucket would get global capacity. } #[test] - fn default_cost_for_type() { + fn different_clients_get_separate_buckets() { let rl = RateLimiter::new( true, RateLimitConfig { - capacity: 100, - refill_rate: 10.0, - default_query_cost: 2, - default_procedure_cost: 5, - default_proxy_cost: 3, + capacity: 2, + refill_rate: 0.001, + default_query_cost: 1, + default_procedure_cost: 1, + default_proxy_cost: 1, }, - vec![], ); - assert_eq!(rl.default_cost_for_type("query"), 2); - assert_eq!(rl.default_cost_for_type("procedure"), 5); - assert_eq!(rl.default_cost_for_type("proxy"), 3); - assert_eq!(rl.default_cost_for_type("unknown"), 1); + // Exhaust client A + assert!(matches!( + rl.check("clientA", 1), + CheckResult::Allowed { .. } + )); + assert!(matches!( + rl.check("clientA", 1), + CheckResult::Allowed { .. } + )); + assert!(matches!( + rl.check("clientA", 1), + CheckResult::Limited { .. } + )); + + // Client B should still have tokens + assert!(matches!( + rl.check("clientB", 1), + CheckResult::Allowed { .. } + )); + assert!(matches!( + rl.check("clientB", 1), + CheckResult::Allowed { .. } + )); + assert!(matches!( + rl.check("clientB", 1), + CheckResult::Limited { .. } + )); } #[test] @@ -365,11 +481,10 @@ mod tests { default_procedure_cost: 1, default_proxy_cost: 1, }, - vec![], ); assert!(rl.is_enabled()); rl.set_enabled(false); assert!(!rl.is_enabled()); - assert!(matches!(rl.check("k", 1, None), CheckResult::Disabled)); + assert!(matches!(rl.check("k", 1), CheckResult::Disabled)); } } diff --git a/src/repo/session.rs b/src/repo/session.rs index 8dfe95e..87335c4 100644 --- a/src/repo/session.rs +++ b/src/repo/session.rs @@ -14,6 +14,7 @@ pub(crate) async fn get_oauth_session( Did::new(did.to_string()).map_err(|_| AppError::Auth(format!("invalid DID: {did}")))?; state .oauth + .default_client() .restore(&did) .await .map_err(|e| AppError::Auth(format!("no OAuth session for {}: {e}", did.as_ref()))) diff --git a/src/repo/upload_blob.rs b/src/repo/upload_blob.rs index f1ec32c..4364d35 100644 --- a/src/repo/upload_blob.rs +++ b/src/repo/upload_blob.rs @@ -2,7 +2,6 @@ use axum::body::Bytes; use axum::extract::State; use axum::http::HeaderMap; use axum::response::Response; -use std::net::IpAddr; use crate::AppState; use crate::auth::Claims; @@ -18,17 +17,10 @@ pub async fn upload_blob( headers: HeaderMap, body: Bytes, ) -> Result { - let client_ip: Option = headers - .get("x-forwarded-for") - .and_then(|v| v.to_str().ok()) - .and_then(|s| s.split(',').next()) - .and_then(|s| s.trim().parse().ok()); - let rate_key = claims.did().to_string(); let check = state.rate_limiter.check( &rate_key, state.rate_limiter.default_cost_for_type("procedure"), - client_ip, ); if let CheckResult::Limited { diff --git a/src/server.rs b/src/server.rs index c80af62..c1b5950 100644 --- a/src/server.rs +++ b/src/server.rs @@ -7,7 +7,6 @@ use axum::{Json, Router}; use bytes::Bytes; use http_body_util::Full; use std::convert::Infallible; -use std::net::IpAddr; use tower_http::cors::CorsLayer; use tower_http::services::ServeDir; use tower_http::trace::TraceLayer; @@ -106,7 +105,8 @@ async fn config_endpoint(State(state): State) -> Json) -> Json { - let mut metadata = serde_json::to_value(&state.oauth.client_metadata).unwrap_or_default(); + let mut metadata = + serde_json::to_value(&state.oauth.default_client().client_metadata).unwrap_or_default(); // The `client_id` field in the response must exactly match the URL the // authorization server fetched. @@ -161,25 +161,15 @@ async fn client_metadata(State(state): State) -> Json) -> Option { - let forwarded = value?; - let first = forwarded.split(',').next()?; - first.trim().parse::().ok() -} - async fn get_profile( State(state): State, claims: Claims, - headers: HeaderMap, + _headers: HeaderMap, ) -> Result { - let client_ip = - ip_from_forwarded_for(headers.get("x-forwarded-for").and_then(|v| v.to_str().ok())); let rate_key = claims.did().to_string(); - let check = state.rate_limiter.check( - &rate_key, - state.rate_limiter.default_cost_for_type("query"), - client_ip, - ); + let check = state + .rate_limiter + .check(&rate_key, state.rate_limiter.default_cost_for_type("query")); if let CheckResult::Limited { retry_after, diff --git a/src/xrpc/mod.rs b/src/xrpc/mod.rs index 0aa2026..55b291b 100644 --- a/src/xrpc/mod.rs +++ b/src/xrpc/mod.rs @@ -3,13 +3,12 @@ mod query; use axum::Json; use axum::body::Body; -use axum::extract::{ConnectInfo, FromRequestParts, Path, RawQuery, State}; +use axum::extract::{FromRequestParts, Path, RawQuery, State}; use axum::http::StatusCode; use axum::http::request::Parts; use axum::response::Response; use serde_json::Value; use std::collections::HashMap; -use std::net::{IpAddr, SocketAddr}; use crate::AppState; use crate::auth::Claims; @@ -155,23 +154,6 @@ async fn proxy_to_authority( .unwrap()) } -/// Extract client IP from X-Forwarded-For header or ConnectInfo. -fn extract_client_ip(parts: &Parts) -> Option { - if let Some(forwarded) = parts - .headers - .get("x-forwarded-for") - .and_then(|v| v.to_str().ok()) - && let Some(first) = forwarded.split(',').next() - && let Ok(ip) = first.trim().parse::() - { - return Some(ip); - } - parts - .extensions - .get::>() - .map(|ci| ci.0.ip()) -} - /// Apply rate limit headers to a response. fn apply_rate_limit_headers(response: &mut Response, remaining: u32, limit: u32, reset: u64) { let headers = response.headers_mut(); @@ -189,18 +171,13 @@ pub async fn xrpc_get( ) -> Result { let raw_query = raw_query.unwrap_or_default(); let mut params = parse_query_params(&raw_query); - let client_ip = extract_client_ip(&parts); let claims = Claims::from_request_parts(&mut parts, &state).await.ok(); - // Rate limit check + // Rate limit check — keyed by client_key for API client requests, "anonymous" otherwise let rate_key = claims .as_ref() - .map(|c| c.did().to_string()) - .unwrap_or_else(|| { - client_ip - .map(|ip| ip.to_string()) - .unwrap_or_else(|| "unknown".to_string()) - }); + .and_then(|c| c.client_key().map(|k| k.to_string())) + .unwrap_or_else(|| "anonymous".to_string()); let lexicon = state.lexicons.get(&method).await; @@ -214,7 +191,7 @@ pub async fn xrpc_get( state.rate_limiter.default_cost_for_type("proxy") }; - let check = state.rate_limiter.check(&rate_key, cost, client_ip); + let check = state.rate_limiter.check(&rate_key, cost); match check { CheckResult::Limited { @@ -270,27 +247,21 @@ pub async fn xrpc_get( Ok(response) } -/// Extract client IP from X-Forwarded-For header value. -fn ip_from_forwarded_for(value: Option<&str>) -> Option { - let forwarded = value?; - let first = forwarded.split(',').next()?; - first.trim().parse::().ok() -} - /// Catch-all POST handler for XRPC procedures. pub async fn xrpc_post( State(state): State, Path(method): Path, RawQuery(raw_query): RawQuery, claims: Claims, - headers: axum::http::HeaderMap, + _headers: axum::http::HeaderMap, Json(body): Json, ) -> Result { let raw_query = raw_query.unwrap_or_default(); let mut params = parse_query_params(&raw_query); - let client_ip = - ip_from_forwarded_for(headers.get("x-forwarded-for").and_then(|v| v.to_str().ok())); - let rate_key = claims.did().to_string(); + let rate_key = claims + .client_key() + .map(|k| k.to_string()) + .unwrap_or_else(|| "anonymous".to_string()); let lexicon = state.lexicons.get(&method).await; @@ -304,7 +275,7 @@ pub async fn xrpc_post( state.rate_limiter.default_cost_for_type("proxy") }; - let check = state.rate_limiter.check(&rate_key, cost, client_ip); + let check = state.rate_limiter.check(&rate_key, cost); match check { CheckResult::Limited { diff --git a/tests/common/app.rs b/tests/common/app.rs index 4033e30..da0dab8 100644 --- a/tests/common/app.rs +++ b/tests/common/app.rs @@ -123,9 +123,10 @@ impl TestApp { default_procedure_cost: 1, default_proxy_cost: 1, }, - vec![], ), - oauth: std::sync::Arc::new(oauth), + oauth: std::sync::Arc::new(happyview::auth::OAuthClientRegistry::new( + std::sync::Arc::new(oauth), + )), oauth_state_store: happyview::auth::oauth_store::DbStateStore::new( pool.clone(), backend, diff --git a/tests/e2e_api_clients.rs b/tests/e2e_api_clients.rs new file mode 100644 index 0000000..b5db164 --- /dev/null +++ b/tests/e2e_api_clients.rs @@ -0,0 +1,643 @@ +mod common; + +use axum::body::Body; +use axum::http::{Request, StatusCode}; +use http_body_util::BodyExt; +use serde_json::{Value, json}; +use serial_test::serial; +use tower::ServiceExt; + +use common::app::TestApp; + +// --------------------------------------------------------------------------- +// Helpers +// --------------------------------------------------------------------------- + +async fn json_body(resp: axum::response::Response) -> Value { + let body = resp.into_body().collect().await.unwrap().to_bytes(); + serde_json::from_slice(&body).unwrap() +} + +fn admin_get( + uri: &str, + cookie: (axum::http::HeaderName, axum::http::HeaderValue), +) -> Request { + Request::builder() + .uri(uri) + .header(cookie.0, cookie.1) + .body(Body::empty()) + .unwrap() +} + +fn admin_post( + uri: &str, + cookie: (axum::http::HeaderName, axum::http::HeaderValue), + body: &Value, +) -> Request { + Request::builder() + .method("POST") + .uri(uri) + .header(cookie.0, cookie.1) + .header("content-type", "application/json") + .body(Body::from(serde_json::to_vec(body).unwrap())) + .unwrap() +} + +fn admin_put( + uri: &str, + cookie: (axum::http::HeaderName, axum::http::HeaderValue), + body: &Value, +) -> Request { + Request::builder() + .method("PUT") + .uri(uri) + .header(cookie.0, cookie.1) + .header("content-type", "application/json") + .body(Body::from(serde_json::to_vec(body).unwrap())) + .unwrap() +} + +fn admin_delete( + uri: &str, + cookie: (axum::http::HeaderName, axum::http::HeaderValue), +) -> Request { + Request::builder() + .method("DELETE") + .uri(uri) + .header(cookie.0, cookie.1) + .body(Body::empty()) + .unwrap() +} + +fn sample_api_client_body() -> Value { + json!({ + "name": "Test App", + "client_id_url": "https://testapp.example.com/oauth-client-metadata.json", + "client_uri": "https://testapp.example.com", + "redirect_uris": ["https://happyview.example.com/auth/callback"], + "scopes": "atproto" + }) +} + +// --------------------------------------------------------------------------- +// Create +// --------------------------------------------------------------------------- + +#[tokio::test] +#[serial] +#[ignore] +async fn create_api_client_returns_201() { + let app = TestApp::new().await; + let body = sample_api_client_body(); + + let resp = app + .router + .clone() + .oneshot(admin_post("/admin/api-clients", app.admin_cookie(), &body)) + .await + .unwrap(); + + assert_eq!(resp.status(), StatusCode::CREATED); + let json = json_body(resp).await; + assert_eq!(json["name"], "Test App"); + assert_eq!( + json["client_id_url"], + "https://testapp.example.com/oauth-client-metadata.json" + ); + let key = json["client_key"].as_str().unwrap(); + assert!(key.starts_with("hvc_"), "client_key should start with hvc_"); + assert_eq!(key.len(), 36); // "hvc_" (4) + 32 hex chars + assert!(json["id"].as_str().is_some()); +} + +#[tokio::test] +#[serial] +#[ignore] +async fn create_api_client_duplicate_client_id_url_fails() { + let app = TestApp::new().await; + let body = sample_api_client_body(); + + // First create succeeds + let resp = app + .router + .clone() + .oneshot(admin_post("/admin/api-clients", app.admin_cookie(), &body)) + .await + .unwrap(); + assert_eq!(resp.status(), StatusCode::CREATED); + + // Second create with same client_id_url should fail (UNIQUE constraint) + let resp = app + .router + .clone() + .oneshot(admin_post("/admin/api-clients", app.admin_cookie(), &body)) + .await + .unwrap(); + assert_eq!(resp.status(), StatusCode::INTERNAL_SERVER_ERROR); +} + +#[tokio::test] +#[serial] +#[ignore] +async fn create_api_client_registers_in_oauth_registry() { + let app = TestApp::new().await; + let body = sample_api_client_body(); + + let resp = app + .router + .clone() + .oneshot(admin_post("/admin/api-clients", app.admin_cookie(), &body)) + .await + .unwrap(); + assert_eq!(resp.status(), StatusCode::CREATED); + + // The OAuth registry should now have this client + let client_id_url = "https://testapp.example.com/oauth-client-metadata.json"; + assert!( + app.state.oauth.get(client_id_url).is_some(), + "OAuth registry should contain the newly created client" + ); +} + +// --------------------------------------------------------------------------- +// List +// --------------------------------------------------------------------------- + +#[tokio::test] +#[serial] +#[ignore] +async fn list_api_clients_empty() { + let app = TestApp::new().await; + + let resp = app + .router + .clone() + .oneshot(admin_get("/admin/api-clients", app.admin_cookie())) + .await + .unwrap(); + + assert_eq!(resp.status(), StatusCode::OK); + let json = json_body(resp).await; + assert!(json.as_array().unwrap().is_empty()); +} + +#[tokio::test] +#[serial] +#[ignore] +async fn list_api_clients_returns_created_clients() { + let app = TestApp::new().await; + + // Create two clients + let body1 = json!({ + "name": "App One", + "client_id_url": "https://one.example.com/oauth-client-metadata.json", + "client_uri": "https://one.example.com", + "redirect_uris": ["https://happyview.example.com/auth/callback"], + "scopes": "atproto" + }); + let body2 = json!({ + "name": "App Two", + "client_id_url": "https://two.example.com/oauth-client-metadata.json", + "client_uri": "https://two.example.com", + "redirect_uris": ["https://happyview.example.com/auth/callback"], + "scopes": "atproto" + }); + + app.router + .clone() + .oneshot(admin_post("/admin/api-clients", app.admin_cookie(), &body1)) + .await + .unwrap(); + app.router + .clone() + .oneshot(admin_post("/admin/api-clients", app.admin_cookie(), &body2)) + .await + .unwrap(); + + let resp = app + .router + .clone() + .oneshot(admin_get("/admin/api-clients", app.admin_cookie())) + .await + .unwrap(); + + assert_eq!(resp.status(), StatusCode::OK); + let json = json_body(resp).await; + let arr = json.as_array().unwrap(); + assert_eq!(arr.len(), 2); + + // Verify fields are present + for client in arr { + assert!(client["id"].as_str().is_some()); + assert!(client["client_key"].as_str().is_some()); + assert!(client["name"].as_str().is_some()); + assert!(client["client_id_url"].as_str().is_some()); + assert!(client["is_active"].as_bool().is_some()); + assert!(client["created_by"].as_str().is_some()); + } +} + +// --------------------------------------------------------------------------- +// Get +// --------------------------------------------------------------------------- + +#[tokio::test] +#[serial] +#[ignore] +async fn get_api_client_returns_details() { + let app = TestApp::new().await; + let body = sample_api_client_body(); + + let create_resp = app + .router + .clone() + .oneshot(admin_post("/admin/api-clients", app.admin_cookie(), &body)) + .await + .unwrap(); + let created = json_body(create_resp).await; + let id = created["id"].as_str().unwrap(); + + let resp = app + .router + .clone() + .oneshot(admin_get( + &format!("/admin/api-clients/{id}"), + app.admin_cookie(), + )) + .await + .unwrap(); + + assert_eq!(resp.status(), StatusCode::OK); + let json = json_body(resp).await; + assert_eq!(json["name"], "Test App"); + assert_eq!( + json["client_id_url"], + "https://testapp.example.com/oauth-client-metadata.json" + ); + assert_eq!(json["client_uri"], "https://testapp.example.com"); + assert_eq!(json["scopes"], "atproto"); + assert_eq!(json["is_active"], true); + assert_eq!(json["created_by"], "did:plc:testadmin"); + assert!(json["redirect_uris"].as_array().unwrap().len() == 1); +} + +#[tokio::test] +#[serial] +#[ignore] +async fn get_api_client_not_found() { + let app = TestApp::new().await; + + let resp = app + .router + .clone() + .oneshot(admin_get( + "/admin/api-clients/00000000-0000-0000-0000-000000000000", + app.admin_cookie(), + )) + .await + .unwrap(); + + assert_eq!(resp.status(), StatusCode::NOT_FOUND); +} + +// --------------------------------------------------------------------------- +// Update +// --------------------------------------------------------------------------- + +#[tokio::test] +#[serial] +#[ignore] +async fn update_api_client_changes_fields() { + let app = TestApp::new().await; + + // Create + let create_resp = app + .router + .clone() + .oneshot(admin_post( + "/admin/api-clients", + app.admin_cookie(), + &sample_api_client_body(), + )) + .await + .unwrap(); + let created = json_body(create_resp).await; + let id = created["id"].as_str().unwrap(); + + // Update + let update_body = json!({ + "name": "Updated App", + "scopes": "atproto transition:generic" + }); + let resp = app + .router + .clone() + .oneshot(admin_put( + &format!("/admin/api-clients/{id}"), + app.admin_cookie(), + &update_body, + )) + .await + .unwrap(); + assert_eq!(resp.status(), StatusCode::NO_CONTENT); + + // Verify + let resp = app + .router + .clone() + .oneshot(admin_get( + &format!("/admin/api-clients/{id}"), + app.admin_cookie(), + )) + .await + .unwrap(); + let json = json_body(resp).await; + assert_eq!(json["name"], "Updated App"); + assert_eq!(json["scopes"], "atproto transition:generic"); + // Unchanged fields should remain + assert_eq!(json["client_uri"], "https://testapp.example.com"); +} + +#[tokio::test] +#[serial] +#[ignore] +async fn update_api_client_not_found() { + let app = TestApp::new().await; + + let resp = app + .router + .clone() + .oneshot(admin_put( + "/admin/api-clients/00000000-0000-0000-0000-000000000000", + app.admin_cookie(), + &json!({"name": "Nope"}), + )) + .await + .unwrap(); + + assert_eq!(resp.status(), StatusCode::NOT_FOUND); +} + +#[tokio::test] +#[serial] +#[ignore] +async fn update_api_client_deactivate_removes_from_registry() { + let app = TestApp::new().await; + + let create_resp = app + .router + .clone() + .oneshot(admin_post( + "/admin/api-clients", + app.admin_cookie(), + &sample_api_client_body(), + )) + .await + .unwrap(); + let created = json_body(create_resp).await; + let id = created["id"].as_str().unwrap(); + + let client_id_url = "https://testapp.example.com/oauth-client-metadata.json"; + assert!(app.state.oauth.get(client_id_url).is_some()); + + // Deactivate + let resp = app + .router + .clone() + .oneshot(admin_put( + &format!("/admin/api-clients/{id}"), + app.admin_cookie(), + &json!({"is_active": false}), + )) + .await + .unwrap(); + assert_eq!(resp.status(), StatusCode::NO_CONTENT); + + // Should be removed from registry + assert!( + app.state.oauth.get(client_id_url).is_none(), + "Deactivated client should be removed from OAuth registry" + ); +} + +// --------------------------------------------------------------------------- +// Delete +// --------------------------------------------------------------------------- + +#[tokio::test] +#[serial] +#[ignore] +async fn delete_api_client_returns_204() { + let app = TestApp::new().await; + + let create_resp = app + .router + .clone() + .oneshot(admin_post( + "/admin/api-clients", + app.admin_cookie(), + &sample_api_client_body(), + )) + .await + .unwrap(); + let created = json_body(create_resp).await; + let id = created["id"].as_str().unwrap(); + + let resp = app + .router + .clone() + .oneshot(admin_delete( + &format!("/admin/api-clients/{id}"), + app.admin_cookie(), + )) + .await + .unwrap(); + + assert_eq!(resp.status(), StatusCode::NO_CONTENT); + + // Verify gone from list + let resp = app + .router + .clone() + .oneshot(admin_get("/admin/api-clients", app.admin_cookie())) + .await + .unwrap(); + let json = json_body(resp).await; + assert!(json.as_array().unwrap().is_empty()); +} + +#[tokio::test] +#[serial] +#[ignore] +async fn delete_api_client_removes_from_oauth_registry() { + let app = TestApp::new().await; + + let create_resp = app + .router + .clone() + .oneshot(admin_post( + "/admin/api-clients", + app.admin_cookie(), + &sample_api_client_body(), + )) + .await + .unwrap(); + let created = json_body(create_resp).await; + let id = created["id"].as_str().unwrap(); + + let client_id_url = "https://testapp.example.com/oauth-client-metadata.json"; + assert!(app.state.oauth.get(client_id_url).is_some()); + + app.router + .clone() + .oneshot(admin_delete( + &format!("/admin/api-clients/{id}"), + app.admin_cookie(), + )) + .await + .unwrap(); + + assert!( + app.state.oauth.get(client_id_url).is_none(), + "Deleted client should be removed from OAuth registry" + ); +} + +#[tokio::test] +#[serial] +#[ignore] +async fn delete_api_client_not_found() { + let app = TestApp::new().await; + + let resp = app + .router + .clone() + .oneshot(admin_delete( + "/admin/api-clients/00000000-0000-0000-0000-000000000000", + app.admin_cookie(), + )) + .await + .unwrap(); + + assert_eq!(resp.status(), StatusCode::NOT_FOUND); +} + +// --------------------------------------------------------------------------- +// Permission enforcement +// --------------------------------------------------------------------------- + +#[tokio::test] +#[serial] +#[ignore] +async fn api_clients_no_auth_returns_401() { + let app = TestApp::new().await; + + let resp = app + .router + .clone() + .oneshot( + Request::builder() + .uri("/admin/api-clients") + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + + assert_eq!(resp.status(), StatusCode::UNAUTHORIZED); +} + +#[tokio::test] +#[serial] +#[ignore] +async fn api_clients_non_admin_returns_403() { + let app = TestApp::new().await; + + let resp = app + .router + .clone() + .oneshot(admin_get( + "/admin/api-clients", + common::auth::admin_cookie_header("did:plc:notadmin", &app.state.cookie_key), + )) + .await + .unwrap(); + + assert_eq!(resp.status(), StatusCode::FORBIDDEN); +} + +// --------------------------------------------------------------------------- +// OAuth registry (unit-level via AppState) +// --------------------------------------------------------------------------- + +#[tokio::test] +#[serial] +#[ignore] +async fn oauth_registry_get_or_default_returns_default_for_unknown() { + let app = TestApp::new().await; + + let client = app + .state + .oauth + .get_or_default(Some("https://unknown.example.com/metadata.json")); + let default = app.state.oauth.default_client(); + + // Should be the same Arc (default client) + assert!(std::sync::Arc::ptr_eq(&client, default)); +} + +#[tokio::test] +#[serial] +#[ignore] +async fn oauth_registry_get_or_default_returns_default_for_none() { + let app = TestApp::new().await; + + let client = app.state.oauth.get_or_default(None); + let default = app.state.oauth.default_client(); + + assert!(std::sync::Arc::ptr_eq(&client, default)); +} + +// --------------------------------------------------------------------------- +// Rate limit config on API clients +// --------------------------------------------------------------------------- + +#[tokio::test] +#[serial] +#[ignore] +async fn create_api_client_with_rate_limit_overrides() { + let app = TestApp::new().await; + let body = json!({ + "name": "Rate Limited App", + "client_id_url": "https://ratelimited.example.com/oauth-client-metadata.json", + "client_uri": "https://ratelimited.example.com", + "redirect_uris": ["https://happyview.example.com/auth/callback"], + "scopes": "atproto", + "rate_limit_capacity": 50, + "rate_limit_refill_rate": 1.5 + }); + + let resp = app + .router + .clone() + .oneshot(admin_post("/admin/api-clients", app.admin_cookie(), &body)) + .await + .unwrap(); + assert_eq!(resp.status(), StatusCode::CREATED); + let created = json_body(resp).await; + let id = created["id"].as_str().unwrap(); + + // Verify overrides persisted + let resp = app + .router + .clone() + .oneshot(admin_get( + &format!("/admin/api-clients/{id}"), + app.admin_cookie(), + )) + .await + .unwrap(); + let json = json_body(resp).await; + assert_eq!(json["rate_limit_capacity"], 50); + assert_eq!(json["rate_limit_refill_rate"], 1.5); +} diff --git a/tests/lua_atproto_api.rs b/tests/lua_atproto_api.rs index fa0e5c8..8f28964 100644 --- a/tests/lua_atproto_api.rs +++ b/tests/lua_atproto_api.rs @@ -79,9 +79,10 @@ async fn test_state_with_pool(pool: sqlx::AnyPool, backend: DatabaseBackend) -> default_procedure_cost: 1, default_proxy_cost: 1, }, - vec![], ), - oauth: std::sync::Arc::new(oauth), + oauth: std::sync::Arc::new(happyview::auth::OAuthClientRegistry::new( + std::sync::Arc::new(oauth), + )), oauth_state_store: happyview::auth::oauth_store::DbStateStore::new(pool.clone(), backend), cookie_key: axum_extra::extract::cookie::Key::derive_from(b"test-secret"), plugin_registry: std::sync::Arc::new(happyview::plugin::PluginRegistry::new()), diff --git a/tests/lua_db_api.rs b/tests/lua_db_api.rs index 98af0ff..8bb6edd 100644 --- a/tests/lua_db_api.rs +++ b/tests/lua_db_api.rs @@ -82,9 +82,10 @@ async fn test_state_with_pool(pool: sqlx::AnyPool, backend: DatabaseBackend) -> default_procedure_cost: 1, default_proxy_cost: 1, }, - vec![], ), - oauth: std::sync::Arc::new(oauth), + oauth: std::sync::Arc::new(happyview::auth::OAuthClientRegistry::new( + std::sync::Arc::new(oauth), + )), oauth_state_store: happyview::auth::oauth_store::DbStateStore::new(pool.clone(), backend), cookie_key: axum_extra::extract::cookie::Key::derive_from(b"test-secret"), plugin_registry: std::sync::Arc::new(happyview::plugin::PluginRegistry::new()), -- 2.51.2