diff --git a/migrations/20260317000000_rate_limit_token_costs.sql b/migrations/20260317000000_rate_limit_token_costs.sql new file mode 100644 index 0000000..0f3879c --- /dev/null +++ b/migrations/20260317000000_rate_limit_token_costs.sql @@ -0,0 +1,5 @@ +ALTER TABLE rate_limits ADD COLUMN default_query_cost INTEGER NOT NULL DEFAULT 1; +ALTER TABLE rate_limits ADD COLUMN default_procedure_cost INTEGER NOT NULL DEFAULT 1; +ALTER TABLE rate_limits ADD COLUMN default_proxy_cost INTEGER NOT NULL DEFAULT 1; +DELETE FROM rate_limits WHERE method IS NOT NULL; +ALTER TABLE lexicons ADD COLUMN token_cost INTEGER; diff --git a/src/admin/lexicons.rs b/src/admin/lexicons.rs index 6aea853..f9c155e 100644 --- a/src/admin/lexicons.rs +++ b/src/admin/lexicons.rs @@ -60,6 +60,7 @@ pub(super) async fn upload_lexicon( action.clone(), body.script.clone(), body.index_hook.clone(), + body.token_cost.map(|c| c as u32), ) .map_err(|e| AppError::BadRequest(format!("failed to parse lexicon: {e}")))?; @@ -79,8 +80,8 @@ pub(super) async fn upload_lexicon( // Upsert into database let row: (i32,) = sqlx::query_as( r#" - INSERT INTO lexicons (id, lexicon_json, backfill, target_collection, action, script, index_hook, source) - VALUES ($1, $2, $3, $4, $5, $6, $7, 'manual') + INSERT INTO lexicons (id, lexicon_json, backfill, target_collection, action, script, index_hook, token_cost, source) + VALUES ($1, $2, $3, $4, $5, $6, $7, $8, 'manual') ON CONFLICT (id) DO UPDATE SET lexicon_json = EXCLUDED.lexicon_json, backfill = EXCLUDED.backfill, @@ -88,6 +89,7 @@ pub(super) async fn upload_lexicon( action = EXCLUDED.action, script = EXCLUDED.script, index_hook = EXCLUDED.index_hook, + token_cost = EXCLUDED.token_cost, source = 'manual', revision = lexicons.revision + 1, updated_at = NOW() @@ -101,6 +103,7 @@ pub(super) async fn upload_lexicon( .bind(action_str) .bind(&body.script) .bind(&body.index_hook) + .bind(body.token_cost) .fetch_one(&state.db) .await .map_err(|e| AppError::Internal(format!("failed to upsert lexicon: {e}")))?; @@ -115,6 +118,7 @@ pub(super) async fn upload_lexicon( action, body.script, body.index_hook.clone(), + body.token_cost.map(|c| c as u32), ) .map_err(|e| AppError::Internal(format!("failed to re-parse lexicon: {e}")))?; let is_record = parsed.lexicon_type == LexiconType::Record; @@ -168,9 +172,9 @@ pub(super) async fn list_lexicons( ) -> Result>, AppError> { auth.require(Permission::LexiconsRead).await?; #[allow(clippy::type_complexity)] - let rows: Vec<(String, i32, Value, bool, Option, Option, Option, Option, String, Option, Option>, chrono::DateTime, chrono::DateTime)> = + let rows: Vec<(String, i32, Value, bool, Option, Option, Option, Option, String, Option, Option>, chrono::DateTime, chrono::DateTime, Option)> = sqlx::query_as( - "SELECT id, revision, lexicon_json, backfill, action, target_collection, script, index_hook, source, authority_did, last_fetched_at, created_at, updated_at FROM lexicons ORDER BY id", + "SELECT id, revision, lexicon_json, backfill, action, target_collection, script, index_hook, source, authority_did, last_fetched_at, created_at, updated_at, token_cost FROM lexicons ORDER BY id", ) .fetch_all(&state.db) .await @@ -193,9 +197,17 @@ pub(super) async fn list_lexicons( last_fetched_at, created_at, updated_at, + token_cost, )| { - let parsed = - ParsedLexicon::parse(json, revision, None, ProcedureAction::Upsert, None, None); + let parsed = ParsedLexicon::parse( + json, + revision, + None, + ProcedureAction::Upsert, + None, + None, + None, + ); let lexicon_type = parsed .as_ref() .map(|p| format!("{:?}", p.lexicon_type).to_lowercase()) @@ -220,6 +232,7 @@ pub(super) async fn list_lexicons( created_at, updated_at, record_schema, + token_cost, } }, ) @@ -236,9 +249,9 @@ pub(super) async fn get_lexicon( ) -> Result, AppError> { auth.require(Permission::LexiconsRead).await?; #[allow(clippy::type_complexity)] - let row: Option<(String, i32, Value, bool, Option, Option, Option, Option, String, Option, Option>, chrono::DateTime, chrono::DateTime)> = + let row: Option<(String, i32, Value, bool, Option, Option, Option, Option, String, Option, Option>, chrono::DateTime, chrono::DateTime, Option)> = sqlx::query_as( - "SELECT id, revision, lexicon_json, backfill, action, target_collection, script, index_hook, source, authority_did, last_fetched_at, created_at, updated_at FROM lexicons WHERE id = $1", + "SELECT id, revision, lexicon_json, backfill, action, target_collection, script, index_hook, source, authority_did, last_fetched_at, created_at, updated_at, token_cost FROM lexicons WHERE id = $1", ) .bind(&id) .fetch_optional(&state.db) @@ -259,6 +272,7 @@ pub(super) async fn get_lexicon( last_fetched_at, created_at, updated_at, + token_cost, ) = row.ok_or_else(|| AppError::NotFound(format!("lexicon '{id}' not found")))?; let lexicon_type = ParsedLexicon::parse( @@ -268,6 +282,7 @@ pub(super) async fn get_lexicon( ProcedureAction::Upsert, None, None, + None, ) .map(|p| format!("{:?}", p.lexicon_type).to_lowercase()) .unwrap_or_else(|_| "unknown".into()); @@ -291,6 +306,7 @@ pub(super) async fn get_lexicon( "last_fetched_at": last_fetched_at, "created_at": created_at, "updated_at": updated_at, + "token_cost": token_cost, }))) } diff --git a/src/admin/mod.rs b/src/admin/mod.rs index 1a9929a..2072936 100644 --- a/src/admin/mod.rs +++ b/src/admin/mod.rs @@ -73,7 +73,6 @@ pub fn admin_routes(_state: AppState) -> Router { "/rate-limits", post(rate_limits::upsert).get(rate_limits::list), ) - .route("/rate-limits/{id}", delete(rate_limits::delete)) .route("/rate-limits/enabled", put(rate_limits::set_enabled)) .route("/rate-limits/allowlist", post(rate_limits::add_allowlist)) .route( diff --git a/src/admin/network_lexicons.rs b/src/admin/network_lexicons.rs index 7bebe4a..ad8916c 100644 --- a/src/admin/network_lexicons.rs +++ b/src/admin/network_lexicons.rs @@ -44,6 +44,7 @@ pub(super) async fn add( ProcedureAction::Upsert, None, None, + None, ) .map_err(|e| AppError::BadRequest(format!("failed to parse lexicon: {e}")))?; @@ -82,6 +83,7 @@ pub(super) async fn add( ProcedureAction::Upsert, None, None, + None, ) .map_err(|e| AppError::Internal(format!("failed to re-parse lexicon: {e}")))?; state.lexicons.upsert(parsed).await; diff --git a/src/admin/rate_limits.rs b/src/admin/rate_limits.rs index 62928b1..51a30b3 100644 --- a/src/admin/rate_limits.rs +++ b/src/admin/rate_limits.rs @@ -9,8 +9,7 @@ use crate::event_log::{EventLog, Severity, log_event}; use super::auth::UserAuth; use super::permissions::Permission; use super::types::{ - AddAllowlistBody, AllowlistEntry, RateLimitSummary, RateLimitsResponse, SetEnabledBody, - UpsertRateLimitBody, + AddAllowlistBody, AllowlistEntry, RateLimitsResponse, SetEnabledBody, UpsertRateLimitBody, }; /// GET /admin/rate-limits — list rate limit config. @@ -27,12 +26,15 @@ pub(super) async fn list( .map_err(|e| AppError::Internal(format!("failed to read rate limit settings: {e}")))? .unwrap_or_else(|| "true".to_string()); - let limits: Vec = sqlx::query_as( - "SELECT id, method, capacity, refill_rate, created_at, updated_at FROM rate_limits ORDER BY id", + let row: Option<(i32, f32, i32, i32, i32)> = sqlx::query_as( + "SELECT capacity, refill_rate, default_query_cost, default_procedure_cost, default_proxy_cost FROM rate_limits WHERE method IS NULL", ) - .fetch_all(&state.db) + .fetch_optional(&state.db) .await - .map_err(|e| AppError::Internal(format!("failed to list rate limits: {e}")))?; + .map_err(|e| AppError::Internal(format!("failed to read rate limits: {e}")))?; + + let (capacity, refill_rate, default_query_cost, default_procedure_cost, default_proxy_cost) = + row.unwrap_or((100, 2.0, 1, 1, 1)); let allowlist: Vec = sqlx::query_as("SELECT id, cidr, note, created_at FROM rate_limit_allowlist ORDER BY id") @@ -42,12 +44,16 @@ pub(super) async fn list( Ok(Json(RateLimitsResponse { enabled: enabled == "true", - limits, + capacity, + refill_rate, + default_query_cost, + default_procedure_cost, + default_proxy_cost, allowlist, })) } -/// POST /admin/rate-limits — upsert a rate limit rule. +/// POST /admin/rate-limits — upsert the global rate limit config. pub(super) async fn upsert( State(state): State, auth: UserAuth, @@ -57,17 +63,22 @@ pub(super) async fn upsert( sqlx::query( r#" - INSERT INTO rate_limits (method, capacity, refill_rate) - VALUES ($1, $2, $3) + INSERT INTO rate_limits (method, capacity, refill_rate, default_query_cost, default_procedure_cost, default_proxy_cost) + VALUES (NULL, $1, $2, $3, $4, $5) ON CONFLICT (method) DO UPDATE SET capacity = EXCLUDED.capacity, refill_rate = EXCLUDED.refill_rate, + default_query_cost = EXCLUDED.default_query_cost, + default_procedure_cost = EXCLUDED.default_procedure_cost, + default_proxy_cost = EXCLUDED.default_proxy_cost, updated_at = NOW() "#, ) - .bind(&body.method) .bind(body.capacity as i32) .bind(body.refill_rate as f32) + .bind(body.default_query_cost as i32) + .bind(body.default_procedure_cost as i32) + .bind(body.default_proxy_cost as i32) .execute(&state.db) .await .map_err(|e| AppError::Internal(format!("failed to upsert rate limit: {e}")))?; @@ -80,10 +91,13 @@ pub(super) async fn upsert( event_type: "rate_limit.upserted".to_string(), severity: Severity::Info, actor_did: Some(auth.did.clone()), - subject: body.method.clone(), + subject: None, detail: serde_json::json!({ "capacity": body.capacity, "refill_rate": body.refill_rate, + "default_query_cost": body.default_query_cost, + "default_procedure_cost": body.default_procedure_cost, + "default_proxy_cost": body.default_proxy_cost, }), }, ) @@ -92,59 +106,6 @@ pub(super) async fn upsert( Ok(StatusCode::CREATED) } -/// DELETE /admin/rate-limits/{id} — delete a rate limit rule. -pub(super) async fn delete( - State(state): State, - auth: UserAuth, - Path(id): Path, -) -> Result { - auth.require(Permission::RateLimitsDelete).await?; - - // Prevent deleting the global default (method IS NULL) - let is_global: Option<(bool,)> = - sqlx::query_as("SELECT (method IS NULL) FROM rate_limits WHERE id = $1") - .bind(id) - .fetch_optional(&state.db) - .await - .map_err(|e| AppError::Internal(format!("failed to check rate limit: {e}")))?; - - match is_global { - None => { - return Err(AppError::NotFound(format!( - "rate limit rule {id} not found" - ))); - } - Some((true,)) => { - return Err(AppError::BadRequest( - "cannot delete the global default rate limit".to_string(), - )); - } - Some((false,)) => {} - } - - sqlx::query("DELETE FROM rate_limits WHERE id = $1") - .bind(id) - .execute(&state.db) - .await - .map_err(|e| AppError::Internal(format!("failed to delete rate limit: {e}")))?; - - state.rate_limiter.reload_from_db(&state.db).await; - - log_event( - &state.db, - EventLog { - event_type: "rate_limit.deleted".to_string(), - severity: Severity::Info, - actor_did: Some(auth.did.clone()), - subject: Some(id.to_string()), - detail: serde_json::json!({}), - }, - ) - .await; - - Ok(StatusCode::NO_CONTENT) -} - /// PUT /admin/rate-limits/enabled — toggle rate limiting. pub(super) async fn set_enabled( State(state): State, diff --git a/src/admin/types.rs b/src/admin/types.rs index 9d33e00..2049a10 100644 --- a/src/admin/types.rs +++ b/src/admin/types.rs @@ -23,6 +23,7 @@ pub(super) struct LexiconSummary { /// For record-type lexicons: the `properties` object from `defs.main.record`. #[serde(skip_serializing_if = "Option::is_none")] pub(super) record_schema: Option, + pub(super) token_cost: Option, } #[derive(Deserialize)] @@ -34,6 +35,7 @@ pub(super) struct UploadLexiconBody { pub(super) action: Option, pub(super) script: Option, pub(super) index_hook: Option, + pub(super) token_cost: Option, } fn default_backfill() -> bool { @@ -211,9 +213,11 @@ pub(super) struct UpdateLabelerBody { #[derive(Deserialize)] pub(super) struct UpsertRateLimitBody { - pub(super) method: Option, pub(super) capacity: u32, pub(super) refill_rate: f64, + pub(super) default_query_cost: u32, + pub(super) default_procedure_cost: u32, + pub(super) default_proxy_cost: u32, } #[derive(Deserialize)] @@ -230,18 +234,12 @@ pub(super) struct AddAllowlistBody { #[derive(Serialize)] pub(super) struct RateLimitsResponse { pub(super) enabled: bool, - pub(super) limits: Vec, - pub(super) allowlist: Vec, -} - -#[derive(Serialize, sqlx::FromRow)] -pub(super) struct RateLimitSummary { - pub(super) id: i32, - pub(super) method: Option, pub(super) capacity: i32, pub(super) refill_rate: f32, - pub(super) created_at: chrono::DateTime, - pub(super) updated_at: chrono::DateTime, + pub(super) default_query_cost: i32, + pub(super) default_procedure_cost: i32, + pub(super) default_proxy_cost: i32, + pub(super) allowlist: Vec, } #[derive(Serialize, sqlx::FromRow)] diff --git a/src/aip.rs b/src/aip.rs index 93cfdf1..e5c5a49 100644 --- a/src/aip.rs +++ b/src/aip.rs @@ -124,8 +124,10 @@ mod tests { crate::rate_limit::RateLimitConfig { capacity: 100, refill_rate: 2.0, + default_query_cost: 1, + default_procedure_cost: 1, + default_proxy_cost: 1, }, - std::collections::HashMap::new(), vec![], ), } diff --git a/src/lexicon.rs b/src/lexicon.rs index 51edd51..19b8dec 100644 --- a/src/lexicon.rs +++ b/src/lexicon.rs @@ -83,6 +83,8 @@ pub struct ParsedLexicon { pub script: Option, /// Optional Lua script that runs when a record in this collection is indexed. pub index_hook: Option, + /// Optional per-NSID token cost for rate limiting. + pub token_cost: Option, } impl ParsedLexicon { @@ -94,6 +96,7 @@ impl ParsedLexicon { action: ProcedureAction, script: Option, index_hook: Option, + token_cost: Option, ) -> Result { let id = raw .get("id") @@ -138,6 +141,7 @@ impl ParsedLexicon { action, script, index_hook, + token_cost, }) } } @@ -172,8 +176,9 @@ impl LexiconRegistry { Option, Option, Option, + Option, )> = sqlx::query_as( - "SELECT id, lexicon_json, revision, target_collection, action, script, index_hook FROM lexicons", + "SELECT id, lexicon_json, revision, target_collection, action, script, index_hook, token_cost FROM lexicons", ) .fetch_all(db) .await @@ -183,7 +188,9 @@ impl LexiconRegistry { inner.clear(); let mut loaded = 0u32; - for (id, json, revision, target_collection, action_str, script, index_hook) in rows { + for (id, json, revision, target_collection, action_str, script, index_hook, token_cost) in + rows + { let action = match ProcedureAction::from_optional_str(action_str.as_deref()) { Ok(a) => a, Err(e) => { @@ -198,6 +205,7 @@ impl LexiconRegistry { action, script, index_hook, + token_cost.map(|c| c as u32), ) { Ok(parsed) => { inner.insert(id, parsed); @@ -363,6 +371,7 @@ mod tests { ProcedureAction::Upsert, None, None, + None, ) .unwrap(); assert_eq!(parsed.id, "games.gamesgamesgamesgames.game"); @@ -382,6 +391,7 @@ mod tests { ProcedureAction::Upsert, None, None, + None, ) .unwrap(); assert_eq!(parsed.lexicon_type, LexiconType::Query); @@ -403,6 +413,7 @@ mod tests { ProcedureAction::Upsert, None, None, + None, ) .unwrap(); assert_eq!(parsed.lexicon_type, LexiconType::Procedure); @@ -419,6 +430,7 @@ mod tests { ProcedureAction::Delete, None, None, + None, ) .unwrap(); assert_eq!(parsed.action, ProcedureAction::Delete); @@ -433,6 +445,7 @@ mod tests { ProcedureAction::Upsert, None, None, + None, ) .unwrap(); assert_eq!(parsed.lexicon_type, LexiconType::Definitions); @@ -441,7 +454,7 @@ mod tests { #[test] fn parse_missing_id_returns_error() { let raw = json!({"lexicon": 1, "defs": {}}); - let result = ParsedLexicon::parse(raw, 1, None, ProcedureAction::Upsert, None, None); + let result = ParsedLexicon::parse(raw, 1, None, ProcedureAction::Upsert, None, None, None); assert!(result.is_err()); assert!(result.unwrap_err().contains("id")); } @@ -449,9 +462,16 @@ mod tests { #[test] fn parse_preserves_raw_json() { let raw = record_lexicon_json(); - let parsed = - ParsedLexicon::parse(raw.clone(), 1, None, ProcedureAction::Upsert, None, None) - .unwrap(); + let parsed = ParsedLexicon::parse( + raw.clone(), + 1, + None, + ProcedureAction::Upsert, + None, + None, + None, + ) + .unwrap(); assert_eq!(parsed.raw, raw); } @@ -464,6 +484,7 @@ mod tests { ProcedureAction::Upsert, None, None, + None, ) .unwrap(); assert_eq!(parsed.target_collection, Some("custom.collection".into())); @@ -489,6 +510,7 @@ mod tests { ProcedureAction::Upsert, None, None, + None, ) .unwrap(); reg.upsert(parsed).await; @@ -508,6 +530,7 @@ mod tests { ProcedureAction::Upsert, None, None, + None, ) .unwrap(); reg.upsert(v1).await; @@ -519,6 +542,7 @@ mod tests { ProcedureAction::Upsert, None, None, + None, ) .unwrap(); reg.upsert(v2).await; @@ -543,6 +567,7 @@ mod tests { ProcedureAction::Upsert, None, None, + None, ) .unwrap(); reg.upsert(parsed).await; @@ -574,6 +599,7 @@ mod tests { ProcedureAction::Upsert, None, None, + None, ) .unwrap(); let query = ParsedLexicon::parse( @@ -583,6 +609,7 @@ mod tests { ProcedureAction::Upsert, None, None, + None, ) .unwrap(); let procedure = ParsedLexicon::parse( @@ -592,6 +619,7 @@ mod tests { ProcedureAction::Upsert, None, None, + None, ) .unwrap(); let defs = ParsedLexicon::parse( @@ -601,6 +629,7 @@ mod tests { ProcedureAction::Upsert, None, None, + None, ) .unwrap(); @@ -684,6 +713,7 @@ mod tests { ProcedureAction::Upsert, None, Some("function handle() end".into()), + None, ) .unwrap(); assert_eq!(parsed.index_hook, Some("function handle() end".into())); @@ -698,6 +728,7 @@ mod tests { ProcedureAction::Upsert, None, None, + None, ) .unwrap(); assert!(parsed.index_hook.is_none()); @@ -713,6 +744,7 @@ mod tests { ProcedureAction::Upsert, None, Some("function handle() log('hook') end".into()), + None, ) .unwrap(); reg.upsert(parsed).await; @@ -731,6 +763,7 @@ mod tests { ProcedureAction::Upsert, None, None, + None, ) .unwrap(); reg.upsert(parsed).await; diff --git a/src/lua/atproto_api.rs b/src/lua/atproto_api.rs index 214f90b..cc4dfcf 100644 --- a/src/lua/atproto_api.rs +++ b/src/lua/atproto_api.rs @@ -204,8 +204,10 @@ mod tests { crate::rate_limit::RateLimitConfig { capacity: 100, refill_rate: 2.0, + default_query_cost: 1, + default_procedure_cost: 1, + default_proxy_cost: 1, }, - std::collections::HashMap::new(), vec![], ), } diff --git a/src/lua/db_api.rs b/src/lua/db_api.rs index a706cc4..b7aeaaf 100644 --- a/src/lua/db_api.rs +++ b/src/lua/db_api.rs @@ -597,8 +597,10 @@ mod tests { crate::rate_limit::RateLimitConfig { capacity: 100, refill_rate: 2.0, + default_query_cost: 1, + default_procedure_cost: 1, + default_proxy_cost: 1, }, - std::collections::HashMap::new(), vec![], ), } diff --git a/src/lua/execute.rs b/src/lua/execute.rs index 40082b2..09d8ee9 100644 --- a/src/lua/execute.rs +++ b/src/lua/execute.rs @@ -932,8 +932,10 @@ mod tests { crate::rate_limit::RateLimitConfig { capacity: 100, refill_rate: 2.0, + default_query_cost: 1, + default_procedure_cost: 1, + default_proxy_cost: 1, }, - std::collections::HashMap::new(), vec![], ), } diff --git a/src/lua/http_api.rs b/src/lua/http_api.rs index 45b85bd..d7f83b0 100644 --- a/src/lua/http_api.rs +++ b/src/lua/http_api.rs @@ -114,8 +114,10 @@ mod tests { crate::rate_limit::RateLimitConfig { capacity: 100, refill_rate: 2.0, + default_query_cost: 1, + default_procedure_cost: 1, + default_proxy_cost: 1, }, - std::collections::HashMap::new(), vec![], ), } diff --git a/src/main.rs b/src/main.rs index 768a10f..ca2a84a 100644 --- a/src/main.rs +++ b/src/main.rs @@ -114,6 +114,7 @@ async fn main() { ProcedureAction::Upsert, None, None, + None, ) { Ok(parsed) => { if let Err(e) = sqlx::query( @@ -157,12 +158,7 @@ async fn main() { // Initialize rate limiter from DB. let rl_state = RateLimiter::load_from_db(&db).await; - let rate_limiter = RateLimiter::new( - rl_state.enabled, - rl_state.global, - rl_state.overrides, - rl_state.allowlist, - ); + let rate_limiter = RateLimiter::new(rl_state.enabled, rl_state.global, rl_state.allowlist); tokio::spawn(rate_limiter.clone().spawn_cleanup()); let initial_collections = lexicons.get_record_collections().await; diff --git a/src/rate_limit.rs b/src/rate_limit.rs index 1d4dbda..33bec6f 100644 --- a/src/rate_limit.rs +++ b/src/rate_limit.rs @@ -2,7 +2,6 @@ use arc_swap::ArcSwap; use dashmap::DashMap; use ipnet::IpNet; use sqlx::PgPool; -use std::collections::HashMap; use std::net::IpAddr; use std::sync::Arc; use std::sync::atomic::{AtomicBool, Ordering}; @@ -11,6 +10,9 @@ use std::time::{Instant, SystemTime, UNIX_EPOCH}; pub struct RateLimitConfig { pub capacity: u32, pub refill_rate: f64, + pub default_query_cost: u32, + pub default_procedure_cost: u32, + pub default_proxy_cost: u32, } pub enum CheckResult { @@ -39,14 +41,12 @@ pub struct RateLimiter { enabled: AtomicBool, buckets: DashMap, global_config: ArcSwap, - overrides: ArcSwap>, allowlist: ArcSwap>, } pub struct RateLimiterState { pub enabled: bool, pub global: RateLimitConfig, - pub overrides: HashMap, pub allowlist: Vec, } @@ -58,22 +58,16 @@ fn now_unix() -> u64 { } impl RateLimiter { - pub fn new( - enabled: bool, - global: RateLimitConfig, - overrides: HashMap, - allowlist: Vec, - ) -> Arc { + pub fn new(enabled: bool, global: RateLimitConfig, allowlist: Vec) -> Arc { Arc::new(Self { enabled: AtomicBool::new(enabled), buckets: DashMap::new(), global_config: ArcSwap::new(Arc::new(global)), - overrides: ArcSwap::new(Arc::new(overrides)), allowlist: ArcSwap::new(Arc::new(allowlist)), }) } - pub fn check(&self, key: &str, method: Option<&str>, client_ip: Option) -> CheckResult { + pub fn check(&self, key: &str, cost: u32, client_ip: Option) -> CheckResult { if !self.enabled.load(Ordering::Relaxed) { return CheckResult::Disabled; } @@ -87,18 +81,10 @@ impl RateLimiter { } } - let overrides = self.overrides.load(); let global = self.global_config.load(); - - let (capacity, refill_rate) = if let Some(method) = method { - if let Some(cfg) = overrides.get(method) { - (cfg.capacity, cfg.refill_rate) - } else { - (global.capacity, global.refill_rate) - } - } else { - (global.capacity, global.refill_rate) - }; + let capacity = global.capacity; + let refill_rate = global.refill_rate; + let cost_f64 = cost as f64; let now = Instant::now(); @@ -130,15 +116,15 @@ impl RateLimiter { }; let reset = now_unix() + reset_secs; - if bucket.tokens >= 1.0 { - bucket.tokens -= 1.0; + if bucket.tokens >= cost_f64 { + bucket.tokens -= cost_f64; CheckResult::Allowed { remaining: bucket.tokens.floor() as u32, limit: capacity, reset, } } else { - let retry_after = ((1.0 - bucket.tokens) / refill_rate).ceil() as u64; + let retry_after = ((cost_f64 - bucket.tokens) / refill_rate).ceil() as u64; CheckResult::Limited { retry_after, limit: capacity, @@ -147,6 +133,17 @@ impl RateLimiter { } } + /// Get the default cost for a given request type. + pub fn default_cost_for_type(&self, request_type: &str) -> u32 { + let config = self.global_config.load(); + match request_type { + "query" => config.default_query_cost, + "procedure" => config.default_procedure_cost, + "proxy" => config.default_proxy_cost, + _ => 1, + } + } + pub fn set_enabled(&self, enabled: bool) { self.enabled.store(enabled, Ordering::Relaxed); } @@ -155,13 +152,8 @@ impl RateLimiter { self.enabled.load(Ordering::Relaxed) } - pub fn update_config( - &self, - global: RateLimitConfig, - overrides: HashMap, - ) { + pub fn update_config(&self, global: RateLimitConfig) { self.global_config.store(Arc::new(global)); - self.overrides.store(Arc::new(overrides)); } pub fn update_allowlist(&self, entries: Vec) { @@ -191,31 +183,32 @@ impl RateLimiter { .map(|v| v == "true") .unwrap_or(true); - // Load rate limit configs - let rows: Vec<(Option, i32, f32)> = - sqlx::query_as("SELECT method, capacity, refill_rate FROM rate_limits") - .fetch_all(db) - .await - .unwrap_or_default(); - - let mut global = RateLimitConfig { - capacity: 100, - refill_rate: 2.0, - }; - let mut overrides = HashMap::new(); - - for (method, capacity, refill_rate) in rows { - let config = RateLimitConfig { - capacity: capacity as u32, - refill_rate: refill_rate as f64, - }; - match method { - None => global = config, - Some(m) => { - overrides.insert(m, config); + // Load global rate limit config (method IS NULL row) + let row: Option<(i32, f32, i32, i32, i32)> = sqlx::query_as( + "SELECT capacity, refill_rate, default_query_cost, default_procedure_cost, default_proxy_cost FROM rate_limits WHERE method IS NULL", + ) + .fetch_optional(db) + .await + .unwrap_or(None); + + let global = match row { + Some((capacity, refill_rate, query_cost, procedure_cost, proxy_cost)) => { + RateLimitConfig { + capacity: capacity as u32, + refill_rate: refill_rate as f64, + default_query_cost: query_cost as u32, + default_procedure_cost: procedure_cost as u32, + default_proxy_cost: proxy_cost as u32, } } - } + None => RateLimitConfig { + capacity: 100, + refill_rate: 2.0, + default_query_cost: 1, + default_procedure_cost: 1, + default_proxy_cost: 1, + }, + }; // Load allowlist let cidr_rows: Vec<(String,)> = sqlx::query_as("SELECT cidr FROM rate_limit_allowlist") @@ -231,7 +224,6 @@ impl RateLimiter { RateLimiterState { enabled, global, - overrides, allowlist, } } @@ -240,7 +232,7 @@ impl RateLimiter { pub async fn reload_from_db(&self, db: &PgPool) { let state = Self::load_from_db(db).await; self.set_enabled(state.enabled); - self.update_config(state.global, state.overrides); + self.update_config(state.global); self.update_allowlist(state.allowlist); } } @@ -256,21 +248,53 @@ mod tests { RateLimitConfig { capacity: 3, refill_rate: 1.0, + default_query_cost: 1, + default_procedure_cost: 1, + default_proxy_cost: 1, }, - HashMap::new(), vec![], ); - // Should allow 3 requests (bucket starts full) + // Should allow 3 requests (bucket starts full, cost=1 each) for _ in 0..3 { assert!(matches!( - rl.check("k", None, None), + rl.check("k", 1, None), CheckResult::Allowed { .. } )); } // 4th should be limited assert!(matches!( - rl.check("k", None, None), + rl.check("k", 1, None), + CheckResult::Limited { .. } + )); + } + + #[test] + fn cost_deducts_multiple_tokens() { + let rl = RateLimiter::new( + true, + RateLimitConfig { + capacity: 10, + refill_rate: 1.0, + default_query_cost: 1, + 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), + CheckResult::Allowed { remaining: 5, .. } + )); + assert!(matches!( + rl.check("k", 5, None), + CheckResult::Allowed { remaining: 0, .. } + )); + // 3rd should be limited + assert!(matches!( + rl.check("k", 5, None), CheckResult::Limited { .. } )); } @@ -282,11 +306,13 @@ mod tests { RateLimitConfig { capacity: 1, refill_rate: 1.0, + default_query_cost: 1, + default_procedure_cost: 1, + default_proxy_cost: 1, }, - HashMap::new(), vec![], ); - assert!(matches!(rl.check("k", None, None), CheckResult::Disabled)); + assert!(matches!(rl.check("k", 1, None), CheckResult::Disabled)); } #[test] @@ -296,53 +322,36 @@ mod tests { RateLimitConfig { capacity: 1, refill_rate: 0.001, + default_query_cost: 1, + default_procedure_cost: 1, + default_proxy_cost: 1, }, - HashMap::new(), 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", None, Some(ip)), - CheckResult::Disabled - )); + assert!(matches!(rl.check("k", 1, Some(ip)), CheckResult::Disabled)); } #[test] - fn method_override_applies() { - let mut overrides = HashMap::new(); - overrides.insert( - "com.atproto.repo.uploadBlob".to_string(), - RateLimitConfig { - capacity: 2, - refill_rate: 0.001, - }, - ); - + fn default_cost_for_type() { let rl = RateLimiter::new( true, RateLimitConfig { capacity: 100, - refill_rate: 100.0, + refill_rate: 10.0, + default_query_cost: 2, + default_procedure_cost: 5, + default_proxy_cost: 3, }, - overrides, vec![], ); - // Override has capacity 2 - assert!(matches!( - rl.check("k", Some("com.atproto.repo.uploadBlob"), None), - CheckResult::Allowed { limit: 2, .. } - )); - assert!(matches!( - rl.check("k", Some("com.atproto.repo.uploadBlob"), None), - CheckResult::Allowed { limit: 2, .. } - )); - assert!(matches!( - rl.check("k", Some("com.atproto.repo.uploadBlob"), None), - CheckResult::Limited { limit: 2, .. } - )); + 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] @@ -352,13 +361,15 @@ mod tests { RateLimitConfig { capacity: 1, refill_rate: 1.0, + default_query_cost: 1, + default_procedure_cost: 1, + default_proxy_cost: 1, }, - HashMap::new(), vec![], ); assert!(rl.is_enabled()); rl.set_enabled(false); assert!(!rl.is_enabled()); - assert!(matches!(rl.check("k", None, None), CheckResult::Disabled)); + assert!(matches!(rl.check("k", 1, None), CheckResult::Disabled)); } } diff --git a/src/repo/upload_blob.rs b/src/repo/upload_blob.rs index dcd3f3d..20d4a59 100644 --- a/src/repo/upload_blob.rs +++ b/src/repo/upload_blob.rs @@ -25,9 +25,11 @@ pub async fn upload_blob( .and_then(|s| s.trim().parse().ok()); let rate_key = claims.did().to_string(); - let check = state - .rate_limiter - .check(&rate_key, Some("com.atproto.repo.uploadBlob"), client_ip); + let check = state.rate_limiter.check( + &rate_key, + state.rate_limiter.default_cost_for_type("procedure"), + client_ip, + ); if let CheckResult::Limited { retry_after, diff --git a/src/server.rs b/src/server.rs index 69b5dcc..9cfa6a6 100644 --- a/src/server.rs +++ b/src/server.rs @@ -102,9 +102,11 @@ async fn get_profile( 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, Some("app.bsky.actor.getProfile"), client_ip); + let check = state.rate_limiter.check( + &rate_key, + state.rate_limiter.default_cost_for_type("query"), + client_ip, + ); if let CheckResult::Limited { retry_after, diff --git a/src/tap.rs b/src/tap.rs index 44beb10..bf00ef4 100644 --- a/src/tap.rs +++ b/src/tap.rs @@ -708,6 +708,7 @@ async fn handle_lexicon_schema_event(state: &AppState, did: &str, record: &TapRe ProcedureAction::Upsert, None, None, + None, ) { Ok(p) => p, Err(e) => { diff --git a/src/xrpc/mod.rs b/src/xrpc/mod.rs index 26a22fe..d17f5bb 100644 --- a/src/xrpc/mod.rs +++ b/src/xrpc/mod.rs @@ -154,9 +154,19 @@ pub async fn xrpc_get( .unwrap_or_else(|| "unknown".to_string()) }); - let check = state - .rate_limiter - .check(&rate_key, Some(&method), client_ip); + let lexicon = state.lexicons.get(&method).await; + + // Determine token cost: per-NSID override → type default → 1 + let cost = if let Some(ref lex) = lexicon { + lex.token_cost.unwrap_or_else(|| { + let type_str = format!("{:?}", lex.lexicon_type).to_lowercase(); + state.rate_limiter.default_cost_for_type(&type_str) + }) + } else { + state.rate_limiter.default_cost_for_type("proxy") + }; + + let check = state.rate_limiter.check(&rate_key, cost, client_ip); match check { CheckResult::Limited { @@ -173,7 +183,7 @@ pub async fn xrpc_get( CheckResult::Allowed { .. } | CheckResult::Disabled => {} } - let lexicon = match state.lexicons.get(&method).await { + let lexicon = match lexicon { Some(l) => l, None => { let mut response = proxy_to_authority(&state, &method, &raw_query, None).await?; @@ -227,9 +237,19 @@ pub async fn xrpc_post( 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, Some(&method), client_ip); + let lexicon = state.lexicons.get(&method).await; + + // Determine token cost: per-NSID override → type default → 1 + let cost = if let Some(ref lex) = lexicon { + lex.token_cost.unwrap_or_else(|| { + let type_str = format!("{:?}", lex.lexicon_type).to_lowercase(); + state.rate_limiter.default_cost_for_type(&type_str) + }) + } else { + state.rate_limiter.default_cost_for_type("proxy") + }; + + let check = state.rate_limiter.check(&rate_key, cost, client_ip); match check { CheckResult::Limited { @@ -246,7 +266,7 @@ pub async fn xrpc_post( CheckResult::Allowed { .. } | CheckResult::Disabled => {} } - let lexicon = match state.lexicons.get(&method).await { + let lexicon = match lexicon { Some(l) => l, None => { let mut response = proxy_to_authority(&state, &method, "", Some(&body)).await?; diff --git a/tests/common/app.rs b/tests/common/app.rs index 194f1dc..4c38c97 100644 --- a/tests/common/app.rs +++ b/tests/common/app.rs @@ -73,8 +73,10 @@ impl TestApp { happyview::rate_limit::RateLimitConfig { capacity: 100, refill_rate: 2.0, + default_query_cost: 1, + default_procedure_cost: 1, + default_proxy_cost: 1, }, - std::collections::HashMap::new(), vec![], ), }; diff --git a/tests/lua_atproto_api.rs b/tests/lua_atproto_api.rs index 560ca6a..14f25c5 100644 --- a/tests/lua_atproto_api.rs +++ b/tests/lua_atproto_api.rs @@ -37,8 +37,10 @@ async fn test_state_with_pool(pool: sqlx::PgPool) -> AppState { happyview::rate_limit::RateLimitConfig { capacity: 100, refill_rate: 2.0, + default_query_cost: 1, + default_procedure_cost: 1, + default_proxy_cost: 1, }, - std::collections::HashMap::new(), vec![], ), } diff --git a/tests/lua_db_api.rs b/tests/lua_db_api.rs index bf4b374..fffffd0 100644 --- a/tests/lua_db_api.rs +++ b/tests/lua_db_api.rs @@ -40,8 +40,10 @@ async fn test_state_with_pool(pool: sqlx::PgPool) -> AppState { happyview::rate_limit::RateLimitConfig { capacity: 100, refill_rate: 2.0, + default_query_cost: 1, + default_procedure_cost: 1, + default_proxy_cost: 1, }, - std::collections::HashMap::new(), vec![], ), } diff --git a/web/src/app/dashboard/lexicons/[id]/lexicon-detail.tsx b/web/src/app/dashboard/lexicons/[id]/lexicon-detail.tsx index 7e8a4f3..25a1386 100644 --- a/web/src/app/dashboard/lexicons/[id]/lexicon-detail.tsx +++ b/web/src/app/dashboard/lexicons/[id]/lexicon-detail.tsx @@ -22,6 +22,7 @@ import { useLuaCompletions } from "@/hooks/use-lua-completions"; import { SiteHeader } from "@/components/site-header"; import { Badge } from "@/components/ui/badge"; import { Button } from "@/components/ui/button"; +import { Input } from "@/components/ui/input"; import { Label } from "@/components/ui/label"; import { Tooltip, @@ -51,6 +52,8 @@ export default function LexiconDetailPage() { const [hookText, setHookText] = useState(""); const [originalHook, setOriginalHook] = useState(""); const [showHookEditor, setShowHookEditor] = useState(false); + const [tokenCost, setTokenCost] = useState(""); + const [originalTokenCost, setOriginalTokenCost] = useState(""); const { luaCompletions, collections } = useLuaCompletions(jsonText); const load = useCallback(() => { @@ -80,6 +83,8 @@ export default function LexiconDetailPage() { setHookText(lex.index_hook ?? ""); setOriginalHook(lex.index_hook ?? ""); setShowHookEditor(!!lex.index_hook); + setTokenCost(lex.token_cost != null ? String(lex.token_cost) : ""); + setOriginalTokenCost(lex.token_cost != null ? String(lex.token_cost) : ""); }) .catch((e) => setError(e instanceof Error ? e.message : String(e))); }, [getToken, id]); @@ -91,7 +96,8 @@ export default function LexiconDetailPage() { const isDirty = jsonText !== originalJson || luaText !== originalLua || - hookText !== originalHook; + hookText !== originalHook || + tokenCost !== originalTokenCost; async function handleSave() { if (!lexicon) return; @@ -104,6 +110,7 @@ export default function LexiconDetailPage() { backfill: lexicon.backfill, script: luaText || undefined, index_hook: hookText || undefined, + token_cost: tokenCost ? Number(tokenCost) : null, }); load(); } catch (e: unknown) { @@ -220,6 +227,24 @@ export default function LexiconDetailPage() {

)} + {(lexicon.lexicon_type === "query" || + lexicon.lexicon_type === "procedure") && ( +
+ + setTokenCost(e.target.value)} + disabled={!hasPermission("lexicons:create")} + /> +
+ )} diff --git a/web/src/app/dashboard/settings/rate-limits/page.tsx b/web/src/app/dashboard/settings/rate-limits/page.tsx index 3f8a983..e14ceeb 100644 --- a/web/src/app/dashboard/settings/rate-limits/page.tsx +++ b/web/src/app/dashboard/settings/rate-limits/page.tsx @@ -8,12 +8,11 @@ import { useCurrentUser } from "@/hooks/use-current-user"; import { getRateLimits, upsertRateLimit, - deleteRateLimit, setRateLimitEnabled, addAllowlistEntry, removeAllowlistEntry, } from "@/lib/api"; -import type { RateLimitSummary, AllowlistEntry } from "@/types/rate-limits"; +import type { AllowlistEntry } from "@/types/rate-limits"; import { SiteHeader } from "@/components/site-header"; import { Button } from "@/components/ui/button"; import { Input } from "@/components/ui/input"; @@ -42,16 +41,37 @@ export default function RateLimitsPage() { const { getToken } = useAuth(); const { hasPermission } = useCurrentUser(); const [enabled, setEnabled] = useState(false); - const [limits, setLimits] = useState([]); + const [capacity, setCapacity] = useState(""); + const [refillRate, setRefillRate] = useState(""); + const [defaultQueryCost, setDefaultQueryCost] = useState(""); + const [defaultProcedureCost, setDefaultProcedureCost] = useState(""); + const [defaultProxyCost, setDefaultProxyCost] = useState(""); const [allowlist, setAllowlist] = useState([]); const [error, setError] = useState(null); const [toggling, setToggling] = useState(false); + const [saving, setSaving] = useState(false); + + // Track original values for dirty detection + const [origCapacity, setOrigCapacity] = useState(""); + const [origRefillRate, setOrigRefillRate] = useState(""); + const [origQueryCost, setOrigQueryCost] = useState(""); + const [origProcedureCost, setOrigProcedureCost] = useState(""); + const [origProxyCost, setOrigProxyCost] = useState(""); const load = useCallback(() => { getRateLimits(getToken) .then((data) => { setEnabled(data.enabled); - setLimits(data.limits); + setCapacity(String(data.capacity)); + setRefillRate(String(data.refill_rate)); + setDefaultQueryCost(String(data.default_query_cost)); + setDefaultProcedureCost(String(data.default_procedure_cost)); + setDefaultProxyCost(String(data.default_proxy_cost)); + setOrigCapacity(String(data.capacity)); + setOrigRefillRate(String(data.refill_rate)); + setOrigQueryCost(String(data.default_query_cost)); + setOrigProcedureCost(String(data.default_procedure_cost)); + setOrigProxyCost(String(data.default_proxy_cost)); setAllowlist(data.allowlist); }) .catch((e) => setError(e.message)); @@ -61,6 +81,13 @@ export default function RateLimitsPage() { load(); }, [load]); + const isDirty = + capacity !== origCapacity || + refillRate !== origRefillRate || + defaultQueryCost !== origQueryCost || + defaultProcedureCost !== origProcedureCost || + defaultProxyCost !== origProxyCost; + async function handleToggleEnabled(checked: boolean) { setToggling(true); try { @@ -73,12 +100,35 @@ export default function RateLimitsPage() { } } - async function handleDeleteLimit(id: number) { + async function handleSave() { + setError(null); + const cap = Number(capacity); + const rate = Number(refillRate); + const qc = Number(defaultQueryCost); + const pc = Number(defaultProcedureCost); + const xc = Number(defaultProxyCost); + if (!cap || cap <= 0 || !rate || rate <= 0) { + setError("Capacity and refill rate must be positive numbers."); + return; + } + if (qc < 0 || pc < 0 || xc < 0) { + setError("Default costs must be non-negative."); + return; + } + setSaving(true); try { - await deleteRateLimit(getToken, id); + await upsertRateLimit(getToken, { + capacity: cap, + refill_rate: rate, + default_query_cost: qc || 1, + default_procedure_cost: pc || 1, + default_proxy_cost: xc || 1, + }); load(); } catch (e: unknown) { setError(e instanceof Error ? e.message : String(e)); + } finally { + setSaving(false); } } @@ -91,6 +141,8 @@ export default function RateLimitsPage() { } } + const canEdit = hasPermission("rate-limits:create"); + return ( <> @@ -102,7 +154,7 @@ export default function RateLimitsPage() { - {/* Rate limit rules */} + {/* Global bucket + Default costs */}
-
-
-

Rate Limit Rules

-

- Configure global defaults and per-method overrides. -

-
- {hasPermission("rate-limits:create") && ( - - )} +
+

Global Bucket & Default Costs

+

+ Configure the shared token bucket and default costs by request type. +

-
- - - - Method - Capacity - Refill Rate - Updated - - - - - {limits.length === 0 && ( - - - No rate limit rules yet. - - - )} - {limits.map((limit) => ( - - - {limit.method ?? ( - - Global default - - )} - - {limit.capacity} - {limit.refill_rate} tokens/sec - - {new Date(limit.updated_at).toLocaleString()} - - -
- {hasPermission("rate-limits:create") && ( - - )} - {hasPermission("rate-limits:delete") && limit.method !== null && ( - handleDeleteLimit(limit.id)} - /> - )} -
-
-
- ))} -
-
+
+
+ + setCapacity(e.target.value)} + disabled={!canEdit} + /> +
+
+ + setRefillRate(e.target.value)} + disabled={!canEdit} + /> +
+
+ + setDefaultQueryCost(e.target.value)} + disabled={!canEdit} + /> +
+
+ + setDefaultProcedureCost(e.target.value)} + disabled={!canEdit} + /> +
+
+ + setDefaultProxyCost(e.target.value)} + disabled={!canEdit} + /> +
+ + {canEdit && ( +
+ +
+ )}
{/* IP allowlist */} @@ -194,7 +248,7 @@ export default function RateLimitsPage() { IPs or CIDRs that bypass rate limiting.

- {hasPermission("rate-limits:create") && ( + {canEdit && ( )}
@@ -226,7 +280,7 @@ export default function RateLimitsPage() { {entry.cidr} - {entry.note ?? "—"} + {entry.note ?? "\u2014"} {new Date(entry.created_at).toLocaleString()} @@ -251,139 +305,6 @@ export default function RateLimitsPage() { ); } -function UpsertRuleDiag({ - getToken, - onSuccess, - existing, -}: { - getToken: () => Promise; - onSuccess: () => void; - existing?: RateLimitSummary; -}) { - const [method, setMethod] = useState(existing?.method ?? ""); - const [capacity, setCapacity] = useState(String(existing?.capacity ?? "")); - const [refillRate, setRefillRate] = useState( - String(existing?.refill_rate ?? "") - ); - const [error, setError] = useState(null); - const [open, setOpen] = useState(false); - - const isEdit = !!existing; - - async function handleSubmit() { - setError(null); - const cap = Number(capacity); - const rate = Number(refillRate); - if (!cap || cap <= 0 || !rate || rate <= 0) { - setError("Capacity and refill rate must be positive numbers."); - return; - } - try { - const body: { method?: string; capacity: number; refill_rate: number } = { - capacity: cap, - refill_rate: rate, - }; - if (isEdit && existing.method !== null) { - body.method = existing.method; - } else if (!isEdit && method.trim()) { - body.method = method.trim(); - } - await upsertRateLimit(getToken, body); - setOpen(false); - if (!isEdit) { - setMethod(""); - setCapacity(""); - setRefillRate(""); - } - onSuccess(); - } catch (e: unknown) { - setError(e instanceof Error ? e.message : String(e)); - } - } - - return ( - { - setOpen(o); - if (o) { - setMethod(existing?.method ?? ""); - setCapacity(String(existing?.capacity ?? "")); - setRefillRate(String(existing?.refill_rate ?? "")); - setError(null); - } - }} - > - - {isEdit ? ( - - ) : ( - - )} - - - - - {isEdit ? "Edit Rule" : "Add Rule"} - - - {isEdit - ? `Update rate limit for ${existing.method ?? "global default"}.` - : "Add a rate limit rule. Leave method empty for global default."} - - -
- {error &&

{error}

} - {!isEdit && ( -
- - setMethod(e.target.value)} - placeholder="com.atproto.sync.getRepo" - className="font-mono" - /> -
- )} -
- - setCapacity(e.target.value)} - placeholder="100" - /> -
-
- - setRefillRate(e.target.value)} - placeholder="10" - /> -
-
- - - - - - -
-
- ); -} - function AddAllowlistDialog({ getToken, onSuccess, diff --git a/web/src/components/app-sidebar.tsx b/web/src/components/app-sidebar.tsx index 58768e7..19f5a05 100644 --- a/web/src/components/app-sidebar.tsx +++ b/web/src/components/app-sidebar.tsx @@ -132,7 +132,7 @@ export function AppSidebar({ diff --git a/web/src/lib/api.ts b/web/src/lib/api.ts index 2143296..6736817 100644 --- a/web/src/lib/api.ts +++ b/web/src/lib/api.ts @@ -25,7 +25,7 @@ export type { EventLogEntry, EventsListResponse } from "@/types/events" export type { ScriptVariableSummary } from "@/types/script-variables" export type { LabelerSummary } from "@/types/labelers" export type { RecordLabel } from "@/types/records" -export type { RateLimitSummary, AllowlistEntry, RateLimitsResponse } from "@/types/rate-limits" +export type { AllowlistEntry, RateLimitsResponse } from "@/types/rate-limits" // The DPoP proof for admin API calls must target AIP's userinfo URL, // because the backend forwards the proof to AIP for token validation. @@ -122,6 +122,7 @@ export function uploadLexicon( action?: string script?: string index_hook?: string + token_cost?: number | null } ) { return apiFetch<{ id: string; revision: number }>("/admin/lexicons", getToken, { @@ -370,7 +371,13 @@ export function getRateLimits(getToken: () => Promise) { export function upsertRateLimit( getToken: () => Promise, - body: { method?: string; capacity: number; refill_rate: number } + body: { + capacity: number + refill_rate: number + default_query_cost: number + default_procedure_cost: number + default_proxy_cost: number + } ) { return apiFetch("/admin/rate-limits", getToken, { method: "POST", @@ -378,15 +385,6 @@ export function upsertRateLimit( }) } -export function deleteRateLimit( - getToken: () => Promise, - id: number -) { - return apiFetch(`/admin/rate-limits/${encodeURIComponent(id)}`, getToken, { - method: "DELETE", - }) -} - export function setRateLimitEnabled( getToken: () => Promise, body: { enabled: boolean } diff --git a/web/src/types/lexicons.ts b/web/src/types/lexicons.ts index 2e6394d..c7e1fb9 100644 --- a/web/src/types/lexicons.ts +++ b/web/src/types/lexicons.ts @@ -12,6 +12,7 @@ export interface LexiconSummary { last_fetched_at: string | null created_at: string updated_at: string + token_cost: number | null } export interface LexiconDetail extends LexiconSummary { diff --git a/web/src/types/rate-limits.ts b/web/src/types/rate-limits.ts index 8babcdf..93d9d14 100644 --- a/web/src/types/rate-limits.ts +++ b/web/src/types/rate-limits.ts @@ -1,12 +1,3 @@ -export interface RateLimitSummary { - id: number - method: string | null - capacity: number - refill_rate: number - created_at: string - updated_at: string -} - export interface AllowlistEntry { id: number cidr: string @@ -16,6 +7,10 @@ export interface AllowlistEntry { export interface RateLimitsResponse { enabled: boolean - limits: RateLimitSummary[] + capacity: number + refill_rate: number + default_query_cost: number + default_procedure_cost: number + default_proxy_cost: number allowlist: AllowlistEntry[] }