diff --git a/.env.example b/.env.example index 648713bd..de417fc8 100644 --- a/.env.example +++ b/.env.example @@ -31,6 +31,8 @@ DB_PASSWORD=__SECRET_HEX_16__ REDIS_PASSWORD=__SECRET_HEX_16__ # dev: REDIS_URL=redis://:${REDIS_PASSWORD}@localhost:6379 # prod: REDIS_URL=redis://:${REDIS_PASSWORD}@redis:6379 +# Redis Cluster: REDIS_URL=redis-cluster://:${REDIS_PASSWORD}@redis-0:6379?node=redis-1:6379 +# (see deploy/helm/think-watch/README.md, "Redis Cluster") # --- Application --- JWT_SECRET=__SECRET_HEX_32__ diff --git a/CHANGELOG.md b/CHANGELOG.md index dd1dbfb9..ee98abb2 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -11,6 +11,60 @@ target. ## [Unreleased] +### Read before upgrading + +- **Rate-limit windows start empty.** Rate-limit counters move to new Redis + keys (one hash per counter, tagged so that Redis Cluster can run them), and + the counts from before the upgrade are not carried over: every window starts + empty and fills from the first request after the upgrade. The old keys + expire by themselves within two window lengths. Budget counters are kept. +- **Route health starts fresh.** A route's samples, circuit breaker and + lifetime request count move to new keys for the same reason, so every route + starts closed with nothing counted. The old lifetime counters never expire; + `redis-cli --scan --pattern 'route_health:[0-9a-f]*' | xargs redis-cli del` + removes them (the new keys start `route_health:{`). +- **A key's limits no longer replace its owner's.** A rate limit or budget set + on an API key used to take the place of the owner's limit for the same + window or period. Both now apply, each on its own counter: a key's limits can + narrow what its owner may do through that key, never widen it. A key given a + higher limit than its owner to give it more room needs the owner's limit + raised instead. + +### Fixed + +- **Token limits refuse requests.** A `tokens` rate limit never refused + anything, and stopped counting once a request would have taken it past its + limit. A request is now refused once the window's recorded usage reaches the + limit, and every request's tokens are recorded after it, even past the limit + — a window can overshoot by what was in flight when it filled. +- **Several request limits at once.** With two or more `requests` limits on a + user (per minute and per hour, say), every request that passed was counted + twice, and only one of the limits could refuse; with + `security.rate_limit_fail_closed` on, every request was refused as + `rate_limiter_unavailable`. +- **An API key's limits count on that key.** They were counted on its owner's + counter, which every key of the owner shared, and the usage the console + reads for a key (`/api/admin/limits/api_key/{id}/usage`) was always 0. Each + key now has counters of its own, rate limits and budgets alike, for the + gateway and the MCP gateway, and its usage shows what it used. +- **A refused request counts against nothing.** A request refused by a spent + budget, or by one rate limit after another had passed, was still counted + against the request limits. Budgets are now checked first and every rate + limit in one step, so a refused request leaves every counter as it was. +- **`Retry-After` says when to retry.** A `429` from the gateway's own limits + said `Retry-After: 30` whatever the limit. It now gives the seconds until the + window has room for another request, or until a spent budget's period ends + (the next midnight, Monday or 1st of the month, UTC). A spent budget also + sends `x-should-retry: false`, so the OpenAI and Anthropic SDKs don't retry it + by themselves. The body stays in the caller's API format. +- **Redis Cluster.** The rate-limit, route-health and quota scripts touched + keys of several hash slots, which a Redis Cluster refuses: on a cluster, rate + limits silently stopped applying (or refused every request with + `security.rate_limit_fail_closed`), circuit breakers never tripped, and cache + invalidation reached one node only. Every key a script touches now shares a + hash tag, pattern deletes scan every node, and the Helm chart's README + describes a `redis-cluster://` URL. + ## [3.1.0] — 2026-10-05 The thinkwatch-core crates move from v0.59.0 to v0.62.0. Two changes reach diff --git a/Cargo.lock b/Cargo.lock index 93be0587..1f5b01eb 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -4077,6 +4077,7 @@ dependencies = [ "clickhouse", "dotenvy", "fred", + "futures", "hex", "hmac 0.13.0", "http 1.4.0", diff --git a/README.md b/README.md index 6e8773fa..70e29e53 100644 --- a/README.md +++ b/README.md @@ -92,7 +92,9 @@ The gateway (port `3000`) is the only part that clients need to reach. The conso - A model's maximum output tokens, set on the Models page, caps `max_tokens` on every request to that model; it replaces the old output length guardrail. **Limits and budgets** -- Request-count limits are checked before the request; token limits and budgets are counted after the response, so one request can cross a budget before the next is refused. +- Every limit and budget is checked before the request, against what earlier requests used; tokens are counted after the response, so one request can cross a token limit or a budget before the next is refused. A refused request counts against nothing, and its `429` says in `Retry-After` when the limit frees. +- A limit set on an API key applies on top of its owner's, on a counter of its own. +- Redis can be a single node or a Redis Cluster. - If Redis is unavailable, limits fail open by default. Setting `security.rate_limit_fail_closed` refuses requests instead. - Budget alerts fire once per period at 50%, 80%, 95% and 100%. diff --git a/README.zh-CN.md b/README.zh-CN.md index a740754b..2d98bc6d 100644 --- a/README.zh-CN.md +++ b/README.zh-CN.md @@ -92,7 +92,9 @@ cd web && pnpm install && pnpm dev - 模型的「最大输出 token」在模型页设置,限制发往该模型的每个请求的 `max_tokens`,取代原来的输出长度护栏。 **限流与预算** -- 请求数限制在请求发出前检查;Token 限制与预算在响应返回后计入,因此单个请求可能越过预算,此后的请求才会被拒绝。 +- 所有限流与预算都在请求发出前按此前的用量检查;Token 在响应返回后计入,因此单个请求可能越过 Token 限制或预算,此后的请求才会被拒绝。被拒绝的请求不计入任何限制,其 `429` 响应以 `Retry-After` 说明限制何时解除。 +- 设置在 API Key 上的限制叠加在其所属用户的限制之上,单独计数。 +- Redis 可以是单节点,也可以是 Redis Cluster。 - Redis 不可用时,限流默认放行;设置 `security.rate_limit_fail_closed` 后改为拒绝请求。 - 预算提醒在每个周期内于 50%、80%、95% 和 100% 各触发一次。 diff --git a/crates/auth/src/rbac.rs b/crates/auth/src/rbac.rs index 4adac13e..9867d1c6 100644 --- a/crates/auth/src/rbac.rs +++ b/crates/auth/src/rbac.rs @@ -373,61 +373,26 @@ pub async fn compute_user_surface_constraints( Ok(apply_user_overrides(role_merged, override_constraints)) } -/// Like [`compute_user_surface_constraints`] but also folds in any -/// active rate-limit / budget overrides keyed on a specific -/// `api_key_id`. Per-key overrides REPLACE the user-derived value -/// in the matching `(surface, metric, window)` / `(surface, period)` -/// slot — the same merge semantics as user-scope overrides. -/// -/// Use this in the gateway hot path (where the auth middleware -/// already knows the api_key id). Other callers (analytics -/// dashboards, admin "what does this user see today" views) keep -/// using the user-only variant since they have no key context. -pub async fn compute_effective_surface_constraints( +/// The limits attached to one API key — its lineage's active +/// `rate_limit_rules` / `budget_caps` rows — on their own. They are not +/// merged into the owner's: the gateway counts them on the lineage's +/// counters and checks them on top of the owner's limits, so a key's +/// limits can narrow what its owner may do through it but never widen +/// it. Keyed on the lineage so they survive rotation. +pub async fn compute_key_surface_constraints( pool: &PgPool, - user_id: Uuid, - api_key_id: Uuid, + lineage_id: Uuid, ) -> Result { use think_watch_common::limits::{ - self, BudgetSubject, RateLimitSubject, apply_user_overrides, - list_enabled_caps_for_subjects, list_enabled_rules_for_subjects, side_table_as_constraints, + BudgetSubject, RateLimitSubject, list_enabled_caps_for_subjects, + list_enabled_rules_for_subjects, side_table_as_constraints, }; - - let user_merged = compute_user_surface_constraints(pool, user_id).await?; - - // Per-key overrides are stored against the key's `lineage_id` - // (subject_kind = 'api_key_lineage') so they survive rotation. - // Resolve api_key_id → lineage_id once and bind every lookup on - // the lineage. A non-existent api_key_id (never happens at - // runtime — the auth middleware just authenticated this id — - // but defend anyway) maps to an empty override set. - let lineage_id: Option = - sqlx::query_scalar("SELECT lineage_id FROM api_keys WHERE id = $1") - .bind(api_key_id) - .fetch_optional(pool) - .await?; - let Some(lineage_id) = lineage_id else { - return Ok(user_merged); - }; - - // The api_key-side overrides are loaded separately so the - // existing `compute_user_*` helper stays a pure function of - // user_id (used by analytics + admin views). Loading both kinds - // here is two queries instead of one — that's bounded and the - // hot path already does enough DB round-trips that it doesn't - // dominate latency. - let key_rules = + let rules = list_enabled_rules_for_subjects(pool, &[(RateLimitSubject::ApiKeyLineage, lineage_id)]) .await?; - let key_caps = + let caps = list_enabled_caps_for_subjects(pool, &[(BudgetSubject::ApiKeyLineage, lineage_id)]).await?; - let key_overrides = side_table_as_constraints(&key_rules, &key_caps); - - if key_overrides == limits::SurfaceConstraints::default() { - return Ok(user_merged); - } - - Ok(apply_user_overrides(user_merged, key_overrides)) + Ok(side_table_as_constraints(&rules, &caps)) } /// Compute the set of permissions that are explicitly denied to `user_id` diff --git a/crates/common/Cargo.toml b/crates/common/Cargo.toml index 8826a6ee..b50cf489 100644 --- a/crates/common/Cargo.toml +++ b/crates/common/Cargo.toml @@ -31,6 +31,7 @@ hmac = { workspace = true } metrics = { workspace = true } clickhouse = { workspace = true } bytes = { workspace = true } +futures = { workspace = true } url = { workspace = true } # S3-compatible body offload (matches the same SigV4 + reqwest pattern # the Bedrock provider uses — no aws-sdk-s3 dependency, so the build diff --git a/crates/common/src/lib.rs b/crates/common/src/lib.rs index b3d1d944..5ca12adc 100644 --- a/crates/common/src/lib.rs +++ b/crates/common/src/lib.rs @@ -49,5 +49,6 @@ pub mod crypto; // AES-256-GCM envelope for secrets at rest pub mod fixed_window; pub mod json_secret; // `{"$enc": ...}` — a secret nested inside a JSONB column pub mod pii; // BlobRedactor — at-rest body redaction shared by gateway + mcp-gateway +pub mod redis_keys; // pattern deletes that work on one node and on a Redis Cluster pub mod tasks; // supervised_spawn — panic-isolated background tasks pub mod validation; diff --git a/crates/common/src/lifecycle/mod.rs b/crates/common/src/lifecycle/mod.rs index 7723c8b3..d6ae2cd0 100644 --- a/crates/common/src/lifecycle/mod.rs +++ b/crates/common/src/lifecycle/mod.rs @@ -8,8 +8,8 @@ //! //! ```text //! Raw -//! → check_limits (rate-limit gate) -//! → check_budget (pre-call budget peek) +//! → check_budget (pre-call budget peek; charges nothing) +//! → check_limits (rate-limit gate; charges only when it passes) //! → check_access (allowed_models / allowed_tools) //! → surface-specific (cache lookup, breaker, credential resolution, //! invoke_upstream → Invocation) diff --git a/crates/common/src/lifecycle/stages/check_budget.rs b/crates/common/src/lifecycle/stages/check_budget.rs index fbd5e66f..48e8c609 100644 --- a/crates/common/src/lifecycle/stages/check_budget.rs +++ b/crates/common/src/lifecycle/stages/check_budget.rs @@ -16,26 +16,27 @@ //! (request allowed) unless the caller passes `fail_closed = true` //! — matching the [`super::check_limits`] semantics. //! -//! Ordering note: this stage runs AFTER `check_limits`. A request -//! that hits the per-window requests counter via `check_limits` and -//! then gets rejected here for being over-budget will have its -//! `requests` counter incremented anyway — rate-limit counters -//! measure "attempts", not "successes". Operators querying the -//! rate-limit metric will see budget-rejected requests reflected -//! there, by design. +//! Ordering note: this stage runs BEFORE `check_limits`, on the +//! [`Raw`] state. `check_limits` charges the `requests` counters, so +//! running it first would charge a request this stage then refuses. +//! Since this stage charges nothing, a request refused by either +//! charges nothing. +use chrono::Utc; use fred::clients::Client; use crate::audit::AuditLogger; use crate::limits::{BudgetCap, budget}; use super::super::Surface; -use super::super::state::LimitsChecked; +use super::super::state::Raw; /// Run a read-only spend check against every supplied cap. On /// allow, returns the input state unchanged (passthrough — no /// extra type narrowing). On deny, emits a `"budget_exceeded"` -/// audit row + short-circuits with `S::budget_exceeded_response`. +/// audit row + short-circuits with `S::budget_exceeded_response`, +/// naming the spent cap whose period ends last — the request can't +/// succeed before then — and the seconds until it does. /// /// `fail_closed` controls behaviour on a Redis read error: /// - `false` (default): bumps `lifecycle_budget_fail_open_total` @@ -47,12 +48,12 @@ use super::super::state::LimitsChecked; fields(trace_id = %state.trace_id, cap_count = caps.len()), )] pub async fn check_budget( - state: LimitsChecked, + state: Raw, caps: &[BudgetCap], redis: &Client, fail_closed: bool, audit: &AuditLogger, -) -> Result, S::Response> { +) -> Result, S::Response> { if caps.is_empty() { return Ok(state); } @@ -84,27 +85,34 @@ pub async fn check_budget( // Walk caps + statuses in lockstep — `current_spend` preserves // input order so the pairing is positional. - for (cap, status) in caps.iter().zip(statuses.iter()) { - if status.current >= status.limit { - let label = budget_label(cap); - metrics::counter!("lifecycle_budget_exceeded_total").increment(1); - tracing::warn!( - trace_id = %state.trace_id, - cap = %label, - current = status.current, - limit = status.limit, - "budget cap exceeded" - ); - let entry = S::audit_entry(&state.identity, "budget_exceeded") - .trace_id(state.trace_id.clone()) - .detail(serde_json::json!({ - "limit": label, - "current": status.current, - "max": status.limit, - })); - audit.log(entry); - return Err(S::budget_exceeded_response(&label)); - } + let now = Utc::now(); + let spent = caps + .iter() + .zip(statuses.iter()) + .filter(|(_, status)| status.current >= status.limit) + .map(|(cap, status)| (cap, status, budget::secs_until_period_end(cap.period, now))) + .max_by_key(|(_, _, wait)| *wait); + if let Some((cap, status, retry_after_secs)) = spent { + let label = budget_label(cap); + metrics::counter!("lifecycle_budget_exceeded_total").increment(1); + tracing::warn!( + trace_id = %state.trace_id, + cap = %label, + current = status.current, + limit = status.limit, + retry_after_secs, + "budget cap exceeded" + ); + let entry = S::audit_entry(&state.identity, "budget_exceeded") + .trace_id(state.trace_id.clone()) + .detail(serde_json::json!({ + "limit": label, + "current": status.current, + "max": status.limit, + "retry_after_secs": retry_after_secs, + })); + audit.log(entry); + return Err(S::budget_exceeded_response(&label, retry_after_secs)); } Ok(state) @@ -125,7 +133,6 @@ fn budget_label(cap: &BudgetCap) -> String { #[cfg(test)] mod tests { - use super::super::super::state::{LimitCheckRecord, LimitsChecked}; use super::super::super::test_surface::{TestResponse, TestSurface, make_raw}; use super::*; use crate::limits::{BudgetPeriod, BudgetSubject}; @@ -142,26 +149,13 @@ mod tests { crate::audit::AuditLogger::test_drain() } - fn make_limits_checked(user_id: Uuid) -> LimitsChecked { - let raw = make_raw(user_id); - LimitsChecked { - identity: raw.identity, - trace_id: raw.trace_id, - started_at: raw.started_at, - client_ip: raw.client_ip, - limit_check: LimitCheckRecord { - currents: Vec::new(), - }, - } - } - /// No caps configured → trivially pass through without touching /// Redis. The disconnected dummy client confirms the function /// never reached out. #[tokio::test] async fn passes_through_with_no_caps() { let user_id = Uuid::new_v4(); - let state = make_limits_checked(user_id); + let state = make_raw(user_id); let trace_id = state.trace_id.clone(); let started_at = state.started_at; @@ -194,7 +188,7 @@ mod tests { /// short-circuit Response type to be propagated via `?`-style /// match — same shape `check_limits` uses. #[allow(dead_code)] - async fn type_state_compiles(state: LimitsChecked) { + async fn type_state_compiles(state: Raw) { let _ = check_budget::(state, &[], &dummy_redis(), false, &dummy_audit()) .await .map_err(|r| match r { diff --git a/crates/common/src/lifecycle/stages/check_limits.rs b/crates/common/src/lifecycle/stages/check_limits.rs index 31d86f6e..210fb2ae 100644 --- a/crates/common/src/lifecycle/stages/check_limits.rs +++ b/crates/common/src/lifecycle/stages/check_limits.rs @@ -1,11 +1,17 @@ -//! `check_limits` — the first stage of the request lifecycle. -//! Charges the user's `requests` rate-limit counters and either: +//! `check_limits` — the rate-limit gate of the request lifecycle. +//! Checks every rate-limit rule the request is held to, in one atomic +//! step, and either: //! -//! * passes through to [`LimitsChecked`] (carrying the post-INCR -//! currents for downstream audit), or -//! * short-circuits with `S::rate_limited_response(label)` / -//! `S::rate_limiter_unavailable_response()` after emitting the -//! audit row for the deny. +//! * passes through to [`LimitsChecked`] (carrying the post-charge +//! currents for downstream audit), having charged 1 to every +//! `requests` rule, or +//! * short-circuits with `S::rate_limited_response(label, retry_after)` +//! / `S::rate_limiter_unavailable_response()` after emitting the +//! audit row for the deny, having charged nothing. +//! +//! A `tokens` rule refuses the request once its window's recorded +//! usage has reached the limit; the tokens the call uses are added +//! after it, by the surface's `record_usage` hook. //! //! No surface-specific logic lives here — the stage takes //! pre-built `RateLimitRule`s. The surface's pipeline runner is @@ -13,17 +19,20 @@ //! before calling. use fred::clients::Client; +use uuid::Uuid; use crate::audit::AuditLogger; -use crate::limits::{RateLimitRule, RateMetric, sliding}; +use crate::limits::{RateLimitRule, sliding}; use super::super::Surface; use super::super::state::{LimitCheckRecord, LimitsChecked, Raw}; -/// Resolve and charge `requests` counters. See module-level docs. +/// Check the rules and charge the request's `requests` counters. See +/// module-level docs. `owner` is the user the request runs as: every +/// counter carries their hash tag (`sliding::counter_key`). /// /// `fail_closed` controls behaviour on Redis errors: -/// - `false` (default): bumps `lifecycle_rate_limiter_fail_open_total` +/// - `false` (default): bumps `gateway_rate_limiter_fail_open_total` /// and lets the request through. /// - `true`: emits the audit row and short-circuits with /// `S::rate_limiter_unavailable_response()`. Wired up from the @@ -35,16 +44,15 @@ use super::super::state::{LimitCheckRecord, LimitsChecked, Raw}; pub async fn check_limits( state: Raw, rules: &[RateLimitRule], + owner: Uuid, redis: &Client, fail_closed: bool, audit: &AuditLogger, ) -> Result, S::Response> { - let resolved = sliding::resolve_rules(rules, RateMetric::Requests); - // No rules configured ⇒ trivially pass. Skip Redis entirely so // a misconfigured surface (no rules attached) doesn't pay a // round-trip per request. - if resolved.is_empty() { + if rules.is_empty() { return Ok(LimitsChecked { identity: state.identity, trace_id: state.trace_id, @@ -56,63 +64,45 @@ pub async fn check_limits( }); } - let outcome = match sliding::check_and_record(redis, &resolved, 1, !fail_closed).await { + let outcome = match sliding::admit(redis, rules, owner, !fail_closed).await { Ok(o) => o, Err(e) => { - if fail_closed { - metrics::counter!("lifecycle_rate_limiter_unavailable_total").increment(1); - tracing::warn!( - error = %e, - trace_id = %state.trace_id, - "rate limiter unavailable; failing closed" - ); - let entry = S::audit_entry(&state.identity, "rate_limiter_unavailable") - .trace_id(state.trace_id.clone()); - audit.log(entry); - return Err(S::rate_limiter_unavailable_response()); - } - // Fail-open path: bump the counter, log at warn, and - // synthesise an "allowed" outcome with empty currents - // so downstream stages see an unobjectionable result. - metrics::counter!("lifecycle_rate_limiter_fail_open_total").increment(1); + // `admit` only errs when failing closed; it fails open by + // itself. + metrics::counter!("lifecycle_rate_limiter_unavailable_total").increment(1); tracing::warn!( error = %e, trace_id = %state.trace_id, - "rate limiter unavailable; failing open" + "rate limiter unavailable; failing closed" ); - sliding::CheckOutcome { - allowed: true, - exceeded_index: -1, - currents: Vec::new(), - } + let entry = S::audit_entry(&state.identity, "rate_limiter_unavailable") + .trace_id(state.trace_id.clone()); + audit.log(entry); + return Err(S::rate_limiter_unavailable_response()); } }; if !outcome.allowed { - // `exceeded_index` is into the `requests`-filtered slice; - // re-filter to find the rule that tripped so we can label - // the response with `subject:metric/window`. - let label = (outcome.exceeded_index >= 0) - .then(|| { - rules - .iter() - .filter(|r| r.metric == RateMetric::Requests) - .nth(outcome.exceeded_index as usize) - .map(sliding::rate_label) - }) - .flatten() + let label = outcome + .exceeded_index + .and_then(|i| rules.get(i)) + .map(sliding::rate_label) .unwrap_or_else(|| "rate limit".to_string()); metrics::counter!("lifecycle_rate_limited_total").increment(1); tracing::warn!( trace_id = %state.trace_id, limit = %label, + retry_after_secs = outcome.retry_after_secs, "rate limited" ); let entry = S::audit_entry(&state.identity, "rate_limited") .trace_id(state.trace_id.clone()) - .detail(serde_json::json!({ "limit": label })); + .detail(serde_json::json!({ + "limit": label, + "retry_after_secs": outcome.retry_after_secs, + })); audit.log(entry); - return Err(S::rate_limited_response(&label)); + return Err(S::rate_limited_response(&label, outcome.retry_after_secs)); } Ok(LimitsChecked { @@ -162,6 +152,7 @@ mod tests { let result = check_limits::( raw, &[], + user_id, &dummy_redis(), true, // fail_closed — irrelevant when no rules &dummy_audit(), diff --git a/crates/common/src/lifecycle/stages/mod.rs b/crates/common/src/lifecycle/stages/mod.rs index 0109e0c5..d0dd2ffb 100644 --- a/crates/common/src/lifecycle/stages/mod.rs +++ b/crates/common/src/lifecycle/stages/mod.rs @@ -1,7 +1,7 @@ //! Pipeline stages — each surface-agnostic stage is a plain //! `async fn` here. See [`super`] for the full pipeline shape. //! -//! Short-circuit stages (`check_limits`, `check_budget`, +//! Short-circuit stages (`check_budget`, `check_limits`, //! `check_access`) live as standalone fns that take the previous //! state struct and return either the next state or //! `Err(S::Response)`. The four post-invoke stages diff --git a/crates/common/src/lifecycle/surface.rs b/crates/common/src/lifecycle/surface.rs index 1fe02cb0..1d5a6420 100644 --- a/crates/common/src/lifecycle/surface.rs +++ b/crates/common/src/lifecycle/surface.rs @@ -70,8 +70,9 @@ pub trait Surface: Sized + Send + Sync + 'static { /// Render the surface's "you got rate limited" response. The /// `label` is a human string like `"user:requests/1m"` that /// the limit engine produces; the surface decides whether to - /// surface it verbatim or wrap it. - fn rate_limited_response(label: &str) -> Self::Response; + /// surface it verbatim or wrap it. `retry_after_secs` is how long + /// until that window has room again (HTTP `Retry-After`). + fn rate_limited_response(label: &str, retry_after_secs: u64) -> Self::Response; /// Render the surface's "rate-limiter unavailable, fail-closed" /// response. Distinct from `rate_limited_response` because the @@ -99,11 +100,10 @@ pub trait Surface: Sized + Send + Sync + 'static { /// Render the surface's "budget cap exhausted" response. The /// label is `":budget/"` (e.g. /// `"user:budget/monthly"`) so clients can tell which cap fired - /// without parsing prose. Maps to 429 on the wire so existing - /// rate-limit retry semantics apply — `GatewayError::LocalRateLimited`'s - /// docstring explicitly covers both rate and budget under that - /// status family. - fn budget_exceeded_response(label: &str) -> Self::Response; + /// without parsing prose. `retry_after_secs` is how long until the + /// cap's period ends — a budget does not free before that, so the + /// AI gateway also tells SDKs not to retry on their own. + fn budget_exceeded_response(label: &str, retry_after_secs: u64) -> Self::Response; /// Render the surface's "budget read backend unavailable" /// response (Redis outage + `fail_closed` enabled). Same wire diff --git a/crates/common/src/lifecycle/test_surface.rs b/crates/common/src/lifecycle/test_surface.rs index 6cc9758f..129e3374 100644 --- a/crates/common/src/lifecycle/test_surface.rs +++ b/crates/common/src/lifecycle/test_surface.rs @@ -30,6 +30,7 @@ pub enum TestResponse { Ok, RateLimited { label: String, + retry_after_secs: u64, }, RateLimiterUnavailable, AccessDenied { @@ -37,6 +38,7 @@ pub enum TestResponse { }, BudgetExceeded { label: String, + retry_after_secs: u64, }, BudgetUnavailable, } @@ -133,9 +135,10 @@ impl Surface for TestSurface { .audit(action) } - fn rate_limited_response(label: &str) -> Self::Response { + fn rate_limited_response(label: &str, retry_after_secs: u64) -> Self::Response { TestResponse::RateLimited { label: label.to_owned(), + retry_after_secs, } } @@ -156,9 +159,10 @@ impl Surface for TestSurface { } } - fn budget_exceeded_response(label: &str) -> Self::Response { + fn budget_exceeded_response(label: &str, retry_after_secs: u64) -> Self::Response { TestResponse::BudgetExceeded { label: label.to_owned(), + retry_after_secs, } } diff --git a/crates/common/src/limits/budget.rs b/crates/common/src/limits/budget.rs index 0e1d9346..d88e8673 100644 --- a/crates/common/src/limits/budget.rs +++ b/crates/common/src/limits/budget.rs @@ -24,18 +24,21 @@ // Why a separate module from `sliding`: // - Period semantics are different from window semantics — a 5h // "today" budget would be confusing. -// - The check is post-hoc only (responses, not requests) so the -// all-or-nothing Lua dance is unnecessary; plain INCRBY + -// value-read is enough. +// - The pre-call check only reads (`current_spend`) and the debit +// comes after the response, so the all-or-nothing Lua dance is +// unnecessary; plain INCRBY + value-read is enough. // - The Redis key namespace is intentionally separate to avoid // collisions if a future migration changes one shape. +// +// Redis Cluster: every command here touches one key, so the counters +// need no hash tag. // ============================================================================ -use chrono::{DateTime, Datelike, Utc}; +use chrono::{DateTime, Datelike, Days, Months, NaiveTime, Utc}; use fred::clients::Client; use fred::interfaces::KeysInterface; -use super::{BudgetCap, BudgetSubject}; +use super::{BudgetCap, BudgetPeriod, BudgetSubject}; // ---------------------------------------------------------------------------- // Threshold alerting @@ -107,6 +110,29 @@ pub fn bucket_id(period: &str, now: DateTime) -> String { } } +/// When the period that contains `now` ends — and its counter stops +/// counting: the next midnight (daily), the next Monday 00:00 (weekly, +/// ISO weeks), the 1st of the next month 00:00 (monthly), all UTC like +/// [`bucket_id`]. +pub fn period_end(period: BudgetPeriod, now: DateTime) -> DateTime { + let today = now.date_naive(); + let first_day_after = match period { + BudgetPeriod::Daily => today + Days::new(1), + BudgetPeriod::Weekly => { + today + Days::new(7 - u64::from(today.weekday().num_days_from_monday())) + } + BudgetPeriod::Monthly => today.with_day(1).expect("every month has a 1st") + Months::new(1), + }; + first_day_after.and_time(NaiveTime::MIN).and_utc() +} + +/// Whole seconds from `now` until [`period_end`], rounded up and at least +/// one: what a refused request's `Retry-After` says. +pub fn secs_until_period_end(period: BudgetPeriod, now: DateTime) -> u64 { + let ms = (period_end(period, now) - now).num_milliseconds(); + super::sliding::retry_after_secs(ms) +} + /// TTL (in seconds) to set on the period counter when we INCRBY it. /// Picked at 2 × the period so a slow process can still find the /// key on the day after, but it's well gone before any chance of @@ -300,6 +326,53 @@ mod tests { assert!(w.starts_with("2026-W"), "got {w}"); } + #[test] + fn a_period_ends_where_its_bucket_id_changes() { + // Wednesday 2026-04-08, mid-afternoon. + let t = Utc.with_ymd_and_hms(2026, 4, 8, 15, 30, 0).unwrap(); + let at = |y, m, d| Utc.with_ymd_and_hms(y, m, d, 0, 0, 0).unwrap(); + assert_eq!(period_end(BudgetPeriod::Daily, t), at(2026, 4, 9)); + assert_eq!(period_end(BudgetPeriod::Weekly, t), at(2026, 4, 13)); + assert_eq!(period_end(BudgetPeriod::Monthly, t), at(2026, 5, 1)); + // Boundaries: a Monday at midnight is the start of a week, not + // its end; December rolls into the next year. + assert_eq!( + period_end(BudgetPeriod::Weekly, at(2026, 4, 13)), + at(2026, 4, 20) + ); + assert_eq!( + period_end(BudgetPeriod::Monthly, at(2026, 12, 31)), + at(2027, 1, 1) + ); + for p in [ + BudgetPeriod::Daily, + BudgetPeriod::Weekly, + BudgetPeriod::Monthly, + ] { + let end = period_end(p, t); + assert_ne!( + bucket_id(p.as_str(), end - chrono::Duration::seconds(1)), + bucket_id(p.as_str(), end), + "{p:?}" + ); + assert_eq!( + bucket_id(p.as_str(), end - chrono::Duration::seconds(1)), + bucket_id(p.as_str(), t), + "{p:?}" + ); + } + } + + #[test] + fn a_spent_budget_waits_until_its_period_ends() { + let t = Utc.with_ymd_and_hms(2026, 4, 30, 23, 59, 59).unwrap() + + chrono::Duration::milliseconds(500); + assert_eq!(secs_until_period_end(BudgetPeriod::Daily, t), 1); + assert_eq!(secs_until_period_end(BudgetPeriod::Monthly, t), 1); + let t = Utc.with_ymd_and_hms(2026, 4, 30, 0, 0, 0).unwrap(); + assert_eq!(secs_until_period_end(BudgetPeriod::Daily, t), 86_400); + } + #[test] fn unknown_period_falls_back_to_monthly() { let t = Utc.with_ymd_and_hms(2026, 4, 8, 0, 0, 0).unwrap(); diff --git a/crates/common/src/limits/mod.rs b/crates/common/src/limits/mod.rs index 9e748422..3acdc827 100644 --- a/crates/common/src/limits/mod.rs +++ b/crates/common/src/limits/mod.rs @@ -31,14 +31,22 @@ // Role- and team-level constraints are NOT their own subjects — they live // in `rbac_roles.policy_document` (and, if ever added, the analogous field // on teams) and fold into each member's merged policy at request time, -// materializing as `subject = User` rules. Redis counters therefore stay -// user-scoped and grouping membership never becomes a shared pool. +// materializing as `subject = User` rules. Grouping membership therefore +// never becomes a shared pool. // -// At request time the proxy resolves which subjects apply (user + -// api_key, plus the merged role/team constraint set attributed to that -// same user) and runs every matching enabled rule through -// `sliding::check_and_record`. Any single failure rejects the request — -// Lua handles the all-or-nothing INCR. +// At request time the proxy builds a [`RequestLimits`]: the user's rules +// and caps (role defaults with the user's overrides) on the user's +// counters, and the calling key's own rules and caps on the key +// lineage's counters. Both sets apply — a key's limits narrow what its +// owner may do through that key, they never widen the owner's. Every +// rule goes through one `sliding::admit` call, so a request refused by +// any of them charges none. +// +// Redis Cluster: every rate-limit counter of a request carries the hash +// tag `{user:}` (see `sliding::counter_key`), so the user's and +// the key's counters share a slot and one script can check them all. +// Budget counters are read and written one key per command and need no +// tag. // // See `plan.md` (limits chapter) for the full design. // ============================================================================ @@ -1138,6 +1146,104 @@ fn apply_block_overrides(target: &mut Option, overrides: Option, + pub caps: Vec, +} + +impl RequestLimits { + /// `user` is the owner's constraints — role defaults with the user's + /// overrides — and counts on the user's counters. `key`, when the + /// request came with an API key that has limits of its own, is the + /// key's lineage id and those limits; they count on the lineage's + /// counters. Both apply: a key's limits narrow its owner's, never + /// widen them. + pub fn for_request( + surface: Surface, + owner: Uuid, + user: &SurfaceConstraints, + key: Option<(Uuid, &SurfaceConstraints)>, + ) -> Self { + let mut out = Self { + owner, + ..Self::default() + }; + out.add(surface, RateLimitSubject::User, owner, user); + if let Some((lineage, constraints)) = key { + out.add( + surface, + RateLimitSubject::ApiKeyLineage, + lineage, + constraints, + ); + } + out + } + + fn add( + &mut self, + surface: Surface, + subject_kind: RateLimitSubject, + subject_id: Uuid, + constraints: &SurfaceConstraints, + ) { + let Some(block) = constraints.block(surface) else { + return; + }; + let budget_kind = match subject_kind { + RateLimitSubject::User => BudgetSubject::User, + RateLimitSubject::ApiKeyLineage => BudgetSubject::ApiKeyLineage, + }; + self.rules.extend( + block + .rules + .iter() + .filter(|r| r.enabled) + .map(|r| RateLimitRule { + id: Uuid::nil(), + subject_kind, + subject_id, + surface, + metric: r.metric, + window_secs: r.window_secs, + max_count: r.max_count, + enabled: true, + expires_at: None, + reason: None, + created_by: None, + }), + ); + self.caps.extend( + block + .budgets + .iter() + .filter(|b| b.enabled) + .map(|b| BudgetCap { + id: Uuid::nil(), + subject_kind: budget_kind, + subject_id, + period: b.period, + limit_tokens: b.limit_tokens, + enabled: true, + expires_at: None, + reason: None, + created_by: None, + }), + ); + } +} + // ---------------------------------------------------------------------------- // Cache-invalidation pubsub // @@ -1584,6 +1690,67 @@ mod tests { ); } + #[test] + fn a_request_is_held_to_its_users_limits_and_its_keys_each_on_their_own_counter() { + let user = Uuid::new_v4(); + let lineage = Uuid::new_v4(); + let block = |max: i64, budget: i64| SurfaceConstraints { + ai_gateway: Some(SurfaceBlock { + rules: vec![SurfaceRule { + metric: RateMetric::Requests, + window_secs: 60, + max_count: max, + enabled: true, + }], + budgets: vec![SurfaceBudget { + period: BudgetPeriod::Daily, + limit_tokens: budget, + enabled: true, + }], + }), + mcp_gateway: None, + }; + let limits = RequestLimits::for_request( + Surface::AiGateway, + user, + &block(1, 100), + Some((lineage, &block(5, 500))), + ); + assert_eq!(limits.owner, user); + // The key's rule of the same window does not replace the user's: + // both are checked. + let rules: Vec<_> = limits + .rules + .iter() + .map(|r| (r.subject_kind, r.subject_id, r.max_count)) + .collect(); + assert_eq!( + rules, + vec![ + (RateLimitSubject::User, user, 1), + (RateLimitSubject::ApiKeyLineage, lineage, 5), + ] + ); + let caps: Vec<_> = limits + .caps + .iter() + .map(|c| (c.subject_kind, c.subject_id, c.limit_tokens)) + .collect(); + assert_eq!( + caps, + vec![ + (BudgetSubject::User, user, 100), + (BudgetSubject::ApiKeyLineage, lineage, 500), + ] + ); + // Another surface's block is not this request's. + assert!( + RequestLimits::for_request(Surface::McpGateway, user, &block(1, 100), None) + .rules + .is_empty() + ); + } + #[test] fn window_to_secs_roundtrip() { for (s, expected) in [ diff --git a/crates/common/src/limits/sliding.rs b/crates/common/src/limits/sliding.rs index ef5cd682..fb340f13 100644 --- a/crates/common/src/limits/sliding.rs +++ b/crates/common/src/limits/sliding.rs @@ -4,9 +4,9 @@ // A pure-sliding window over a 1-week timespan would need ~50k members // in a single Redis ZSET (one per request) and would dominate memory // at any meaningful traffic. We approximate with **fixed-bucket -// sliding**: each window is split into 60 buckets, each bucket is one -// INCR counter, the "current value" is the sum of the last 60 buckets. -// Precision is ~1.6%, more than enough for rate limiting. +// sliding**: each window is split into 60 buckets, the "current value" +// is the sum of the last 60 buckets. Precision is ~1.6%, more than +// enough for rate limiting. // // Bucket sizing per window: // @@ -17,28 +17,43 @@ // 1d → 24m × 60 // 1w → 168m × 60 // -// All N rules for a request go through ONE Lua script invocation. The -// script: +// Storage: one Redis hash per counter, its fields the bucket ids +// (`now_secs / bucket_secs`), each holding what that bucket counted: // -// 1. For each rule, computes its `current` = sum(last 60 bucket counters). -// 2. Checks `current + cost <= max_count` for every rule. -// 3. If any rule would be exceeded → returns `{0, exceeded_rule_index}` -// and writes nothing (atomic all-or-nothing). -// 4. Otherwise INCRs every rule's current bucket by `cost` and refreshes -// the bucket TTL to `2 * window_secs` (so old buckets self-evict). -// 5. Returns `{1, ...}` on success. +// ratelimit:{user:}::::: // -// One round trip, one atomic decision. The script is loaded with -// SCRIPT LOAD on first use and reused via EVALSHA — see `script_sha`. +// `` is the user the request runs as — for a key-lineage counter, +// the key's owner. The braces are a Redis Cluster hash tag: every counter +// one request touches carries the same tag, so all of them sit in one +// slot and one script can read and write them together. Each counter is +// one declared key; the script never builds a key name of its own. +// +// A request goes through two scripts: +// +// * `admit`, before the call: refuses when any rule's window is full +// (`current >= max_count`) and otherwise charges 1 to every +// `requests` rule. `tokens` rules are only read — a request can't +// know its tokens yet. All or nothing: a refused request charges no +// counter at all. A refusal also says how long until the limiting +// window has room again. +// * `record`, after the call: adds the tokens it used to every +// `tokens` rule, unconditionally. Usage that overshoots a limit is +// still recorded — the overshoot is bounded by what was in flight +// when the window filled, and the next request is refused. +// +// Old buckets are pruned by `admit` as it reads them; the hash itself +// expires 2 × window after its last write. // ============================================================================ +use std::collections::HashMap; use std::sync::OnceLock; use fred::clients::Client; use fred::interfaces::LuaInterface; use sha1::{Digest, Sha1}; +use uuid::Uuid; -use super::{RateLimitRule, RateMetric}; +use super::{RateLimitRule, RateLimitSubject, RateMetric, Surface}; /// 60 buckets per window — fixed across every supported window size. /// Larger N tightens precision but balloons Redis memory linearly. @@ -53,80 +68,138 @@ pub fn bucket_secs(window_secs: i32) -> i32 { } // ---------------------------------------------------------------------------- -// The Lua script — see file header for the algorithm +// The Lua scripts — see file header for the algorithm // ---------------------------------------------------------------------------- -const LUA_CHECK_AND_RECORD: &str = r#" --- KEYS: one base key per rule (the prefix; bucket suffix is appended in-script) --- ARGV: [now_ms, cost, n_rules, b1_secs, m1_max, b2_secs, m2_max, ...] --- Returns: {ok, idx_of_first_exceeded_rule_or_-1, [debug_current_per_rule...]} - -local now_ms = tonumber(ARGV[1]) -local cost = tonumber(ARGV[2]) -local n_rules = tonumber(ARGV[3]) -local now_secs = math.floor(now_ms / 1000) - --- Phase 1: read every rule's current sum, decide pass/fail. -local currents = {} -for i = 1, n_rules do - local bucket_secs = tonumber(ARGV[3 + (i - 1) * 2 + 1]) - local max_count = tonumber(ARGV[3 + (i - 1) * 2 + 2]) - local base_key = KEYS[i] - - -- Window starts `bucket_secs * 60` ago. Sum the 60 buckets in - -- [now - 60*bucket_secs, now], one MGET-style loop. - local sum = 0 - local current_bucket = math.floor(now_secs / bucket_secs) - for b = 0, 59 do - local bucket_id = current_bucket - b - local v = redis.call("GET", base_key .. ":" .. bucket_id) - if v then sum = sum + tonumber(v) end +const LUA_ADMIT: &str = r#" +-- KEYS: one counter hash per rule +-- ARGV: [now_ms, then (bucket_secs, max_count, charge) per rule] +-- Returns: {allowed, limiting rule (1-based, 0 when allowed), +-- ms until it has room (0 when allowed), current per rule...} + +local now_ms = tonumber(ARGV[1]) +local now_secs = math.floor(now_ms / 1000) +local n = #KEYS + +local sums = {} +local limiting, wait_ms = 0, 0 +for i = 1, n do + local bucket_secs = tonumber(ARGV[2 + (i - 1) * 3]) + local max_count = tonumber(ARGV[3 + (i - 1) * 3]) + local current = math.floor(now_secs / bucket_secs) + local oldest = current - 59 + + local counts, stale, sum = {}, {}, 0 + local fields = redis.call('HGETALL', KEYS[i]) + for j = 1, #fields, 2 do + local b = tonumber(fields[j]) + if b == nil or b < oldest then + stale[#stale + 1] = fields[j] + elseif b <= current then + local v = tonumber(fields[j + 1]) or 0 + counts[b] = v + sum = sum + v + end + end + if #stale > 0 then + redis.call('HDEL', KEYS[i], unpack(stale)) end - currents[i] = sum + sums[i] = sum - if sum + cost > max_count then - return {0, i, currents} + if sum >= max_count then + -- The window has room for one more once enough of its oldest + -- buckets have left it. Bucket b leaves at (b + 60) * bucket_secs. + local need, freed, b = sum - max_count + 1, 0, oldest + while b < current do + freed = freed + (counts[b] or 0) + if freed >= need then break end + b = b + 1 + end + local w = (b + 60) * bucket_secs * 1000 - now_ms + if w > wait_ms then + wait_ms = w + limiting = i + end end end --- Phase 2: every rule passed → INCR each rule's current bucket. -for i = 1, n_rules do - local bucket_secs = tonumber(ARGV[3 + (i - 1) * 2 + 1]) - local base_key = KEYS[i] - local current_bucket = math.floor(now_secs / bucket_secs) - local k = base_key .. ":" .. current_bucket - redis.call("INCRBY", k, cost) - -- TTL = 2 × window so old buckets vanish before they're reused. - redis.call("EXPIRE", k, bucket_secs * 60 * 2) +if limiting > 0 then + local out = {0, limiting, wait_ms} + for i = 1, n do out[#out + 1] = sums[i] end + return out end -return {1, -1, currents} +for i = 1, n do + local charge = ARGV[4 + (i - 1) * 3] + if tonumber(charge) > 0 then + local bucket_secs = tonumber(ARGV[2 + (i - 1) * 3]) + redis.call('HINCRBY', KEYS[i], math.floor(now_secs / bucket_secs), charge) + redis.call('EXPIRE', KEYS[i], bucket_secs * 120) + sums[i] = sums[i] + tonumber(charge) + end +end +local out = {1, 0, 0} +for i = 1, n do out[#out + 1] = sums[i] end +return out "#; -fn script_sha() -> &'static String { - static SHA: OnceLock = OnceLock::new(); - SHA.get_or_init(|| { +const LUA_RECORD: &str = r#" +-- KEYS: one counter hash per rule +-- ARGV: [now_ms, amount, then bucket_secs per rule] +local now_secs = math.floor(tonumber(ARGV[1]) / 1000) +for i = 1, #KEYS do + local bucket_secs = tonumber(ARGV[2 + i]) + redis.call('HINCRBY', KEYS[i], math.floor(now_secs / bucket_secs), ARGV[2]) + redis.call('EXPIRE', KEYS[i], bucket_secs * 120) +end +return #KEYS +"#; + +fn sha(script: &'static str, cell: &'static OnceLock) -> &'static str { + cell.get_or_init(|| { let mut h = Sha1::new(); - h.update(LUA_CHECK_AND_RECORD.as_bytes()); + h.update(script.as_bytes()); hex::encode(h.finalize()) }) } +/// EVALSHA, falling back to EVAL when the node doesn't have the script +/// cached yet (a restart, a failover, a node new to the cluster). +async fn run( + redis: &Client, + script: &'static str, + cell: &'static OnceLock, + keys: Vec, + args: Vec, +) -> Result, fred::error::Error> { + match redis + .evalsha::, _, _, _>(sha(script, cell), keys.clone(), args.clone()) + .await + { + Ok(v) => Ok(v), + Err(e) if e.details().starts_with("NOSCRIPT") => redis.eval(script, keys, args).await, + Err(e) => Err(e), + } +} + +static ADMIT_SHA: OnceLock = OnceLock::new(); +static RECORD_SHA: OnceLock = OnceLock::new(); + // ---------------------------------------------------------------------------- // Public API // ---------------------------------------------------------------------------- /// Render a rule's identity as the canonical /// `:/` label (e.g. -/// `api_key:tokens/1h`, `user:requests/1m`). Used by: -/// - the `Retry-After`-style HTTP body for rate-limited responses, +/// `api_key_lineage:tokens/1h`, `user:requests/1m`). Used by: +/// - the HTTP body for rate-limited responses, /// - the `gateway_logs` / `mcp_logs` audit row's `limits.label`, /// - log scrapers that group by `{subject_kind, metric, window}`. /// /// Lives next to the engine because every surface that runs -/// `check_and_record` also needs to render the label that caused a -/// deny — keeping them apart bred two identical copies (one each in -/// the AI gateway and MCP gateway). One implementation, one format. +/// [`admit`] also needs to render the label that caused a deny — +/// keeping them apart bred two identical copies (one each in the AI +/// gateway and MCP gateway). One implementation, one format. pub fn rate_label(rule: &super::RateLimitRule) -> String { let window = match rule.window_secs { 60 => "1m".to_string(), @@ -145,202 +218,264 @@ pub fn rate_label(rule: &super::RateLimitRule) -> String { ) } -/// One rule resolved against the current request's metric, ready to -/// hand to `check_and_record`. The `cost` parameter on the helper -/// applies to all rules in the slice equally; callers MUST pre-filter -/// the slice to a single metric before calling. +/// The Redis Cluster hash tag every limit counter of `owner`'s requests +/// carries, so one script can touch all of them. +pub fn hash_tag(owner: Uuid) -> String { + format!("{{user:{owner}}}") +} + +/// The Redis key of one counter. `owner` is the user the requests run +/// as: the subject itself for a user rule, the key's owner for a key +/// lineage rule. +pub fn counter_key( + owner: Uuid, + surface: Surface, + subject_kind: RateLimitSubject, + subject_id: Uuid, + metric: RateMetric, + window_secs: i32, +) -> String { + format!( + "ratelimit:{}:{}:{}:{subject_id}:{}:{window_secs}", + hash_tag(owner), + surface.as_str(), + subject_kind.as_str(), + metric.as_str() + ) +} + +/// One rule's counter, ready for the scripts. #[derive(Debug, Clone)] pub struct ResolvedRule { - pub id: uuid::Uuid, - pub base_key: String, + pub id: Uuid, + pub key: String, pub bucket_secs: i32, pub max_count: i64, } +impl ResolvedRule { + pub fn new(rule: &RateLimitRule, owner: Uuid) -> Self { + Self { + id: rule.id, + key: counter_key( + owner, + rule.surface, + rule.subject_kind, + rule.subject_id, + rule.metric, + rule.window_secs, + ), + bucket_secs: bucket_secs(rule.window_secs), + max_count: rule.max_count, + } + } +} + #[derive(Debug, Clone)] pub struct CheckOutcome { pub allowed: bool, - /// Index into the `rules` slice of the first rule that would be - /// exceeded. -1 when allowed. - pub exceeded_index: i32, - /// Current sum (pre-INCR) for every rule in the slice, in input - /// order. Useful for surfacing "X / Y used" in the response or - /// for `gateway_rate_limit_remaining` headers. + /// Index into the rules handed to [`admit`] of the rule that + /// refused the request — of those whose windows were full, the one + /// that frees last. `None` when allowed. + pub exceeded_index: Option, + /// Seconds until that rule's window has room for the request again, + /// rounded up. 0 when allowed. + pub retry_after_secs: u64, + /// Each rule's count in its window, after this request's charge + /// when it was allowed, in input order. Empty when Redis failed + /// open. pub currents: Vec, } -/// Build a Redis key prefix for a rule. The bucket id is appended by -/// the Lua script. Format: -/// -/// ratelimit::::: -/// -/// Splitting on the metric AND window means a (subject, surface) pair -/// can hold multiple rules without colliding their counters. -pub fn build_base_key( - surface: &str, - subject_kind: &str, - subject_id: uuid::Uuid, - metric: RateMetric, - window_secs: i32, -) -> String { - format!( - "ratelimit:{surface}:{subject_kind}:{subject_id}:{}:{window_secs}", - metric.as_str() - ) -} - -/// Convert a slice of `RateLimitRule` rows into the engine-facing -/// `ResolvedRule` shape, dropping any rules whose metric doesn't match -/// `metric_filter`. Use this to split a "load every rule" result into -/// the requests-pass and tokens-pass batches. -pub fn resolve_rules(rules: &[RateLimitRule], metric_filter: RateMetric) -> Vec { - rules - .iter() - .filter(|r| r.metric == metric_filter) - .map(|r| ResolvedRule { - id: r.id, - base_key: build_base_key( - r.surface.as_str(), - r.subject_kind.as_str(), - r.subject_id, - r.metric, - r.window_secs, - ), - bucket_secs: bucket_secs(r.window_secs), - max_count: r.max_count, - }) - .collect() +impl CheckOutcome { + fn open() -> Self { + Self { + allowed: true, + exceeded_index: None, + retry_after_secs: 0, + currents: Vec::new(), + } + } } -/// Atomic "check then INCR-by-cost" for an entire batch of rules. +/// Before the call: refuse the request if any rule's window is full, +/// otherwise charge 1 to every `requests` rule. `tokens` rules are only +/// read; [`record`] charges them after the call. Atomic across every +/// rule — a refused request charges nothing. /// -/// `cost` is in metric units: `1` for `requests`, weighted-token count -/// for `tokens`. All rules in the slice MUST share a single metric and -/// must already be filtered through `resolve_rules`. +/// `owner` is the user the request runs as (see [`counter_key`]). /// /// On Redis error the policy is controlled by `fail_open`: /// /// * `fail_open = true` (default): returns `Ok(allowed = true)` with /// empty `currents`, bumps `gateway_rate_limiter_fail_open_total`, -/// and lets the request through. This is the historical behavior -/// and matches what most operators want — Redis should not be a -/// single point of failure for the AI control plane. +/// and lets the request through. Redis should not be a single point +/// of failure for the AI control plane. /// * `fail_open = false`: returns the underlying `fred` error so /// callers can refuse the request. Wired up via the -/// `security.rate_limit_fail_closed` system setting; flip it on -/// when the limits engine is the only thing standing between you -/// and a cost incident. -pub async fn check_and_record( +/// `security.rate_limit_fail_closed` system setting. +pub async fn admit( redis: &Client, - rules: &[ResolvedRule], - cost: i64, + rules: &[RateLimitRule], + owner: Uuid, fail_open: bool, ) -> Result { - if rules.is_empty() || cost <= 0 { - return Ok(CheckOutcome { - allowed: true, - exceeded_index: -1, - currents: Vec::new(), - }); - } - - let now_ms = chrono::Utc::now().timestamp_millis(); - let n_rules = rules.len(); - - // Build KEYS[] in the same order as the loop below appends ARGV - // pairs so the script's index math lines up. - let keys: Vec = rules.iter().map(|r| r.base_key.clone()).collect(); + admit_at( + redis, + rules, + owner, + chrono::Utc::now().timestamp_millis(), + fail_open, + ) + .await +} - // ARGV layout: [now_ms, cost, n_rules, then (bucket_secs, max_count) × n] - let mut args: Vec = Vec::with_capacity(3 + n_rules * 2); +/// [`admit`] at a given time, in Unix milliseconds — for tests. +pub async fn admit_at( + redis: &Client, + rules: &[RateLimitRule], + owner: Uuid, + now_ms: i64, + fail_open: bool, +) -> Result { + if rules.is_empty() { + return Ok(CheckOutcome::open()); + } + let mut keys = Vec::with_capacity(rules.len()); + let mut args = Vec::with_capacity(1 + rules.len() * 3); args.push(now_ms.to_string()); - args.push(cost.to_string()); - args.push(n_rules.to_string()); - for r in rules { + for rule in rules { + let r = ResolvedRule::new(rule, owner); + keys.push(r.key); args.push(r.bucket_secs.to_string()); args.push(r.max_count.to_string()); + args.push(match rule.metric { + RateMetric::Requests => "1".to_string(), + RateMetric::Tokens => "0".to_string(), + }); } - // Try EVALSHA first; on NOSCRIPT load and retry. fred's evalsha - // surfaces script-not-found as an error so we fall back manually. - let raw: Result, _> = redis - .evalsha(script_sha().as_str(), keys.clone(), args.clone()) - .await; - let result: Vec = match raw { + let reply = match run(redis, LUA_ADMIT, &ADMIT_SHA, keys, args).await { Ok(v) => v, + Err(e) if fail_open => { + tracing::warn!("rate-limit check failed: {e}; failing open"); + metrics::counter!("gateway_rate_limiter_fail_open_total").increment(1); + return Ok(CheckOutcome::open()); + } Err(e) => { - // Either NOSCRIPT or a transient redis problem. Try a - // direct EVAL once — Redis caches the loaded script for - // subsequent EVALSHA calls. - tracing::debug!("rate-limit evalsha failed ({e}); falling back to EVAL"); - match redis - .eval::, _, _, _>(LUA_CHECK_AND_RECORD, keys, args) - .await - { - Ok(v) => v, - Err(e) => { - if fail_open { - tracing::warn!("rate-limit EVAL failed: {e}; failing open"); - metrics::counter!("gateway_rate_limiter_fail_open_total").increment(1); - return Ok(CheckOutcome { - allowed: true, - exceeded_index: -1, - currents: Vec::new(), - }); - } else { - tracing::error!( - "rate-limit EVAL failed: {e}; failing closed per security.rate_limit_fail_closed" - ); - metrics::counter!("gateway_rate_limiter_fail_closed_total").increment(1); - return Err(e); - } - } - } + tracing::error!( + "rate-limit check failed: {e}; failing closed per security.rate_limit_fail_closed" + ); + metrics::counter!("gateway_rate_limiter_fail_closed_total").increment(1); + return Err(e); } }; + Ok(parse_admit_reply(&reply)) +} - // Lua return shape: [ok, exceeded_index, currents...] - let ok = result.first().copied().unwrap_or(1); - let exceeded_index = result.get(1).copied().unwrap_or(-1) as i32; - let currents = result.iter().skip(2).copied().collect(); - - Ok(CheckOutcome { - allowed: ok == 1, - // Lua is 1-indexed; convert to 0-indexed for Rust callers. - exceeded_index: if exceeded_index >= 1 { - exceeded_index - 1 +/// Reply shape: `[allowed, limiting rule (1-based, 0 = none), wait_ms, +/// currents...]`. +fn parse_admit_reply(reply: &[i64]) -> CheckOutcome { + let allowed = reply.first().copied().unwrap_or(1) == 1; + let limiting = reply.get(1).copied().unwrap_or(0); + let wait_ms = reply.get(2).copied().unwrap_or(0).max(0); + CheckOutcome { + allowed, + exceeded_index: (!allowed && limiting >= 1).then(|| (limiting - 1) as usize), + retry_after_secs: if allowed { + 0 } else { - -1 + retry_after_secs(wait_ms) }, - currents, - }) + currents: reply.iter().skip(3).copied().collect(), + } +} + +/// Whole seconds a client should wait for `wait_ms` to pass: rounded up, +/// and at least one — `Retry-After: 0` reads as "now". +pub fn retry_after_secs(wait_ms: i64) -> u64 { + (wait_ms.max(0) as u64).div_ceil(1000).max(1) +} + +/// After the call: add `amount` to every rule of `metric`, whether or +/// not that takes it past its limit — the call has been made, and +/// hiding what it used would only let the next one through too. +pub async fn record( + redis: &Client, + rules: &[RateLimitRule], + owner: Uuid, + metric: RateMetric, + amount: i64, +) -> Result<(), fred::error::Error> { + record_at( + redis, + rules, + owner, + metric, + amount, + chrono::Utc::now().timestamp_millis(), + ) + .await +} + +/// [`record`] at a given time, in Unix milliseconds — for tests. +pub async fn record_at( + redis: &Client, + rules: &[RateLimitRule], + owner: Uuid, + metric: RateMetric, + amount: i64, + now_ms: i64, +) -> Result<(), fred::error::Error> { + let resolved: Vec = rules + .iter() + .filter(|r| r.metric == metric) + .map(|r| ResolvedRule::new(r, owner)) + .collect(); + if resolved.is_empty() || amount <= 0 { + return Ok(()); + } + let mut args = Vec::with_capacity(2 + resolved.len()); + args.push(now_ms.to_string()); + args.push(amount.to_string()); + args.extend(resolved.iter().map(|r| r.bucket_secs.to_string())); + let keys = resolved.into_iter().map(|r| r.key).collect(); + run(redis, LUA_RECORD, &RECORD_SHA, keys, args) + .await + .map(|_| ()) } /// Read-only "what's the current sum for this rule" helper. Used by -/// the limits CRUD's usage endpoint to render "X / Y used" without -/// any side effects. Walks the same 60 buckets the Lua script does -/// but issues plain GETs instead of an INCR. +/// the console's usage views to render "X / Y used" without any side +/// effects. /// /// Returns 0 on Redis error so the UI can fall back to "no data" /// rather than 500. Real failures are logged. pub async fn current_count(redis: &Client, rule: &ResolvedRule) -> i64 { - use fred::interfaces::KeysInterface; - let now_secs = chrono::Utc::now().timestamp(); - let bucket_secs = rule.bucket_secs as i64; + use fred::interfaces::HashesInterface; + match redis.hgetall::, _>(&rule.key).await { + Ok(buckets) => window_sum(&buckets, chrono::Utc::now().timestamp(), rule.bucket_secs), + Err(e) => { + tracing::warn!(key = %rule.key, "rate-limit usage read failed: {e}"); + 0 + } + } +} + +/// The sum of the buckets inside the window that ends at `now_secs`. +fn window_sum(buckets: &HashMap, now_secs: i64, bucket_secs: i32) -> i64 { + let bucket_secs = i64::from(bucket_secs); if bucket_secs <= 0 { return 0; } - let current_bucket = now_secs / bucket_secs; - let mut sum: i64 = 0; - for b in 0..BUCKETS_PER_WINDOW { - let bucket_id = current_bucket - b; - let key = format!("{}:{}", rule.base_key, bucket_id); - let v: Option = redis.get(&key).await.ok().flatten(); - if let Some(n) = v { - sum += n; - } - } - sum + let current = now_secs.div_euclid(bucket_secs); + let oldest = current - (BUCKETS_PER_WINDOW - 1); + buckets + .iter() + .filter_map(|(b, v)| b.parse::().ok().map(|b| (b, *v))) + .filter(|(b, _)| (oldest..=current).contains(b)) + .map(|(_, v)| v) + .sum() } #[cfg(test)] @@ -364,4 +499,108 @@ mod tests { assert_eq!(bucket_secs(86_400), 1_440); assert_eq!(bucket_secs(604_800), 10_080); } + + fn rule(kind: RateLimitSubject, subject: Uuid, metric: RateMetric) -> RateLimitRule { + RateLimitRule { + id: Uuid::nil(), + subject_kind: kind, + subject_id: subject, + surface: Surface::AiGateway, + metric, + window_secs: 60, + max_count: 1, + enabled: true, + expires_at: None, + reason: None, + created_by: None, + } + } + + /// Redis Cluster runs a script only when every key it declares hashes + /// to one slot. A request's counters — the user's and the key + /// lineage's, requests and tokens — must all land together. + #[test] + fn every_counter_of_a_request_hashes_to_one_slot() { + use fred::util::redis_keyslot; + let user = Uuid::new_v4(); + let lineage = Uuid::new_v4(); + let keys: Vec = [ + rule(RateLimitSubject::User, user, RateMetric::Requests), + rule(RateLimitSubject::User, user, RateMetric::Tokens), + rule( + RateLimitSubject::ApiKeyLineage, + lineage, + RateMetric::Requests, + ), + rule(RateLimitSubject::ApiKeyLineage, lineage, RateMetric::Tokens), + ] + .iter() + .map(|r| ResolvedRule::new(r, user).key) + .collect(); + let slot = redis_keyslot(keys[0].as_bytes()); + for k in &keys { + assert_eq!(redis_keyslot(k.as_bytes()), slot, "{k}"); + } + // The tag is what decides: the slot is the tag's own. + assert_eq!(slot, redis_keyslot(format!("user:{user}").as_bytes())); + } + + #[test] + fn a_key_lineage_counter_is_named_after_the_lineage_under_its_owners_tag() { + let user = Uuid::new_v4(); + let lineage = Uuid::new_v4(); + let key = ResolvedRule::new( + &rule(RateLimitSubject::ApiKeyLineage, lineage, RateMetric::Tokens), + user, + ) + .key; + assert_eq!( + key, + format!("ratelimit:{{user:{user}}}:ai_gateway:api_key_lineage:{lineage}:tokens:60") + ); + } + + #[test] + fn the_window_sum_counts_the_last_sixty_buckets_only() { + let now = 10_000; + let b = |id: i64, v: i64| (id.to_string(), v); + let buckets: HashMap = [ + b(now, 1), // current bucket + b(now - 59, 2), // oldest still inside + b(now - 60, 4), // just left + b(now + 1, 8), // a pod whose clock runs ahead + ] + .into_iter() + .chain([("garbage".to_string(), 16)]) + .collect(); + assert_eq!(window_sum(&buckets, now, 1), 3); + // Five-second buckets: bucket id = now / 5. + let five: HashMap = [b(now / 5, 7)].into_iter().collect(); + assert_eq!(window_sum(&five, now, 5), 7); + } + + #[test] + fn retry_after_rounds_up_and_never_says_now() { + assert_eq!(retry_after_secs(1), 1); + assert_eq!(retry_after_secs(1_000), 1); + assert_eq!(retry_after_secs(1_001), 2); + assert_eq!(retry_after_secs(59_999), 60); + assert_eq!(retry_after_secs(0), 1); + assert_eq!(retry_after_secs(-5), 1); + } + + #[test] + fn a_refusal_names_the_limiting_rule_and_its_wait() { + let refused = parse_admit_reply(&[0, 2, 4_500, 3, 10]); + assert!(!refused.allowed); + assert_eq!(refused.exceeded_index, Some(1)); + assert_eq!(refused.retry_after_secs, 5); + assert_eq!(refused.currents, vec![3, 10]); + + let allowed = parse_admit_reply(&[1, 0, 0, 4]); + assert!(allowed.allowed); + assert_eq!(allowed.exceeded_index, None); + assert_eq!(allowed.retry_after_secs, 0); + assert_eq!(allowed.currents, vec![4]); + } } diff --git a/crates/common/src/redis_keys.rs b/crates/common/src/redis_keys.rs new file mode 100644 index 00000000..a750c611 --- /dev/null +++ b/crates/common/src/redis_keys.rs @@ -0,0 +1,100 @@ +//! Key-space helpers that work the same on one Redis node and on a +//! Redis Cluster. +//! +//! On a cluster, `SCAN` walks only the node it is sent to, and a `DEL` +//! naming keys of different hash slots is refused (`CROSSSLOT`). A +//! pattern delete therefore has to scan every primary and delete slot +//! by slot. + +use std::collections::BTreeMap; + +use fred::clients::Client; +use fred::interfaces::{ClientLike, KeysInterface}; +use fred::types::Key; +use fred::types::scan::ScanType; +use futures::StreamExt; + +/// Keys deleted per `DEL` on a single node. +const BATCH: usize = 256; + +/// Delete every key matching `pattern` (a `SCAN MATCH` glob), of +/// `kind` when given. Returns how many were deleted. Keys written while +/// the scan runs may survive it, as with any `SCAN`. +pub async fn delete_matching( + redis: &Client, + pattern: &str, + kind: Option, +) -> Result { + let clustered = redis.is_clustered(); + let mut keys = if clustered { + redis + .scan_cluster_buffered(pattern.to_owned(), Some(BATCH as u32), kind) + .boxed() + } else { + redis + .scan_buffered(pattern.to_owned(), Some(BATCH as u32), kind) + .boxed() + }; + let mut batch: Vec = Vec::with_capacity(BATCH); + let mut deleted = 0; + while let Some(key) = keys.next().await { + batch.push(key?); + if batch.len() == BATCH { + deleted += delete(redis, std::mem::take(&mut batch), clustered).await?; + } + } + deleted += delete(redis, batch, clustered).await?; + Ok(deleted) +} + +async fn delete( + redis: &Client, + keys: Vec, + clustered: bool, +) -> Result { + if keys.is_empty() { + return Ok(0); + } + if !clustered { + return Ok(redis.del::(keys).await? as usize); + } + let mut deleted = 0; + for (_, keys) in by_slot(keys) { + deleted += redis.del::(keys).await? as usize; + } + Ok(deleted) +} + +/// Group keys by Redis Cluster hash slot: one `DEL` per group is one a +/// cluster accepts. +fn by_slot(keys: Vec) -> BTreeMap> { + let mut out: BTreeMap> = BTreeMap::new(); + for key in keys { + out.entry(fred::util::redis_keyslot(key.as_bytes())) + .or_default() + .push(key); + } + out +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn keys_of_one_slot_are_deleted_together() { + let keys: Vec = ["{a}1", "{a}2", "{b}1", "c"] + .into_iter() + .map(Key::from) + .collect(); + let groups = by_slot(keys); + assert_eq!(groups.len(), 3); + let a = fred::util::redis_keyslot(b"a"); + assert_eq!(groups[&a].len(), 2); + for (slot, keys) in &groups { + for k in keys { + assert_eq!(fred::util::redis_keyslot(k.as_bytes()), *slot); + } + } + } +} diff --git a/crates/gateway/src/cache.rs b/crates/gateway/src/cache.rs index 68202ce0..bdc86333 100644 --- a/crates/gateway/src/cache.rs +++ b/crates/gateway/src/cache.rs @@ -147,34 +147,16 @@ impl ResponseCache { }) } - /// Invalidate all cached responses by deleting keys matching the cache prefix. - /// Uses Lua script for atomic pattern deletion. + /// Invalidate all cached responses by deleting keys matching the + /// cache prefix — on every node, when Redis is a cluster. pub async fn invalidate_all(&self) { - use fred::interfaces::LuaInterface; - // Use Lua EVAL to scan and delete in batches server-side - const LUA_INVALIDATE: &str = r#" -local cursor = '0' -local total = 0 -repeat - local result = redis.call('SCAN', cursor, 'MATCH', ARGV[1], 'COUNT', 100) - cursor = result[1] - local keys = result[2] - if #keys > 0 then - redis.call('DEL', unpack(keys)) - total = total + #keys - end -until cursor == '0' -return total -"#; - let deleted: i64 = self - .redis - .eval( - LUA_INVALIDATE, - Vec::::new(), - vec!["llm_cache:*".to_string()], - ) - .await - .unwrap_or(0); + let deleted = + think_watch_common::redis_keys::delete_matching(&self.redis, "llm_cache:*", None) + .await + .unwrap_or_else(|e| { + tracing::warn!("Cache invalidation failed: {e}"); + 0 + }); metrics::counter!("gateway_cache_invalidations_total").increment(1); tracing::info!(deleted, "Cache invalidated"); } diff --git a/crates/gateway/src/error.rs b/crates/gateway/src/error.rs index 38a48020..56d88a1a 100644 --- a/crates/gateway/src/error.rs +++ b/crates/gateway/src/error.rs @@ -46,12 +46,22 @@ pub enum GatewayError { /// names the account and the IAM principal. #[error("Authentication failed with upstream")] UpstreamAuthError { status: u16, message: String }, - /// Local rate limit / budget cap was hit. The String is the rule - /// label so the response body can tell the caller WHICH limit - /// fired (e.g. "user requests/5h", "api_key tokens/1d", - /// "monthly budget"). Maps to 429 in `IntoResponse`. - #[error("Rate limited: {0}")] - LocalRateLimited(String), + /// Local rate limit / budget cap was hit. `label` names the limit + /// so the response body can tell the caller WHICH one fired (e.g. + /// `user:requests/5h`, `api_key_lineage:tokens/1d`, + /// `user:budget/monthly`). Maps to 429 in `IntoResponse`, with + /// `Retry-After: retry_after_secs`. `retry` is false for a spent + /// budget: it frees only when its period ends, so SDKs that retry a + /// 429 by themselves are told not to (`x-should-retry: false`). + /// Build it with [`GatewayError::rate_limited`], + /// [`GatewayError::budget_exhausted`] or + /// [`GatewayError::limiter_unavailable`]. + #[error("Rate limited: {label}")] + LocalRateLimited { + label: String, + retry_after_secs: u32, + retry: bool, + }, /// Refused by the gateway's own policy — a content filter rule set to /// refuse matched what the caller sent, or a tool call the upstream /// returned matched a rule set to cut it. Not a malformed request (not @@ -77,7 +87,7 @@ impl GatewayError { GatewayError::ProviderInvalidResponse(_) => 502, GatewayError::TransformError(_) => 400, GatewayError::NetworkError(_) => 502, - GatewayError::UpstreamRateLimited { .. } | GatewayError::LocalRateLimited(_) => 429, + GatewayError::UpstreamRateLimited { .. } | GatewayError::LocalRateLimited { .. } => 429, GatewayError::UpstreamAuthError { .. } => 401, GatewayError::PolicyBlocked(_) => 403, } @@ -95,29 +105,75 @@ impl GatewayError { GatewayError::TransformError(_) => "TransformError", GatewayError::NetworkError(_) => "NetworkError", GatewayError::UpstreamRateLimited { .. } => "UpstreamRateLimited", - GatewayError::LocalRateLimited(_) => "LocalRateLimited", + GatewayError::LocalRateLimited { .. } => "LocalRateLimited", GatewayError::UpstreamAuthError { .. } => "UpstreamAuthError", GatewayError::PolicyBlocked(_) => "PolicyBlocked", } } + /// A rate-limit window is full. It has room again in + /// `retry_after_secs`, by itself. + pub fn rate_limited(label: impl Into, retry_after_secs: u64) -> Self { + GatewayError::LocalRateLimited { + label: label.into(), + retry_after_secs: clamp_secs(retry_after_secs), + retry: true, + } + } + + /// A budget is spent. It frees when its period ends, in + /// `retry_after_secs` — far too long for an SDK's automatic retries. + pub fn budget_exhausted(label: impl Into, retry_after_secs: u64) -> Self { + GatewayError::LocalRateLimited { + label: label.into(), + retry_after_secs: clamp_secs(retry_after_secs), + retry: false, + } + } + + /// The limit counters can't be read and the gateway fails closed. + /// Nothing says when they will be back; a short wait is a guess. + pub fn limiter_unavailable(label: impl Into) -> Self { + const UNAVAILABLE_RETRY_SECS: u32 = 30; + GatewayError::LocalRateLimited { + label: label.into(), + retry_after_secs: UNAVAILABLE_RETRY_SECS, + retry: true, + } + } + /// Hint, in seconds, for `Retry-After` on a 429 response. For - /// upstream limits we echo the upstream's own header when present; - /// for local limits we fall back to a conservative 30s so naive - /// clients don't spin into a tight retry loop while the bucket is - /// still refilling. Capped at one hour to keep the header sane - /// even when an upstream returns an absurd value. + /// upstream limits we echo the upstream's own header when present, + /// capped at one hour to keep the header sane even when an upstream + /// returns an absurd value. For local limits it is when the limit + /// lets a request through again — for a budget, the end of its + /// period, which can be weeks away. pub fn retry_after_secs(&self) -> Option { const HARD_CAP_SECS: u32 = 3600; - const LOCAL_DEFAULT_SECS: u32 = 30; match self { GatewayError::UpstreamRateLimited { retry_after_secs } => { retry_after_secs.map(|s| s.min(HARD_CAP_SECS)) } - GatewayError::LocalRateLimited(_) => Some(LOCAL_DEFAULT_SECS), + GatewayError::LocalRateLimited { + retry_after_secs, .. + } => Some(*retry_after_secs), _ => None, } } + + /// `Some(false)` when the client must not retry on its own: sent as + /// `x-should-retry`, which the OpenAI and Anthropic SDKs read before + /// retrying a 429. + pub fn should_retry(&self) -> Option { + match self { + GatewayError::LocalRateLimited { retry: false, .. } => Some(false), + _ => None, + } + } +} + +fn clamp_secs(secs: u64) -> u32 { + u32::try_from(secs).unwrap_or(u32::MAX).max(1) } /// Parse RFC 7231 `Retry-After` (delta-seconds form). HTTP-date is @@ -141,4 +197,25 @@ mod tests { assert_eq!(e.error_tag(), "PolicyBlocked"); assert_eq!(e.retry_after_secs(), None); } + + #[test] + fn a_full_window_says_when_it_frees_and_a_spent_budget_says_not_to_retry() { + let window = GatewayError::rate_limited("user:requests/1m", 17); + assert_eq!(window.status_code(), 429); + assert_eq!(window.error_tag(), "LocalRateLimited"); + assert_eq!(window.retry_after_secs(), Some(17)); + assert_eq!(window.should_retry(), None); + assert_eq!(window.to_string(), "Rate limited: user:requests/1m"); + + // A month away is not capped like an upstream's hint. + let budget = GatewayError::budget_exhausted("user:budget/monthly", 2_000_000); + assert_eq!(budget.status_code(), 429); + assert_eq!(budget.error_tag(), "LocalRateLimited"); + assert_eq!(budget.retry_after_secs(), Some(2_000_000)); + assert_eq!(budget.should_retry(), Some(false)); + + let down = GatewayError::limiter_unavailable("rate_limiter_unavailable"); + assert_eq!(down.retry_after_secs(), Some(30)); + assert_eq!(down.should_retry(), None); + } } diff --git a/crates/gateway/src/health.rs b/crates/gateway/src/health.rs index c26797c8..83a77503 100644 --- a/crates/gateway/src/health.rs +++ b/crates/gateway/src/health.rs @@ -27,6 +27,10 @@ //! express — operators tuning weights need to know whether a //! route has actually carried any requests at all. //! +//! The braces are literal: `{}` is a Redis Cluster hash tag, +//! so a route's three keys share a slot and one script (or one `DEL`) +//! can touch them together. +//! //! ### One round trip, two when the state changes //! //! Recording a completion is one Lua call: insert the sample, drop what @@ -183,11 +187,14 @@ impl CircuitBreakerConfig { } } +/// A route's samples, state and counters keys. `{}` is a Redis +/// Cluster hash tag: the scripts below declare two or three of them at +/// once, which a cluster accepts only when they share a slot. fn keys(route_id: Uuid) -> (String, String, String) { ( - format!("route_health:{route_id}:samples"), - format!("route_health:{route_id}:state"), - format!("route_health:{route_id}:counters"), + format!("route_health:{{{route_id}}}:samples"), + format!("route_health:{{{route_id}}}:state"), + format!("route_health:{{{route_id}}}:counters"), ) } @@ -401,9 +408,7 @@ impl HealthTracker { /// orphan keys after route churn. Best-effort: a Redis hiccup /// here is not worth failing the delete over. pub async fn forget(&self, route_id: Uuid) { - let samples_key = format!("route_health:{route_id}:samples"); - let state_key = format!("route_health:{route_id}:state"); - let counters_key = format!("route_health:{route_id}:counters"); + let (samples_key, state_key, counters_key) = keys(route_id); if let Err(e) = self .redis .del::(vec![samples_key, state_key, counters_key]) @@ -418,6 +423,14 @@ impl HealthTracker { mod tests { use super::*; + #[test] + fn a_routes_keys_share_one_cluster_slot() { + let (samples, state, counters) = keys(Uuid::new_v4()); + let slot = fred::util::redis_keyslot(samples.as_bytes()); + assert_eq!(fred::util::redis_keyslot(state.as_bytes()), slot); + assert_eq!(fred::util::redis_keyslot(counters.as_bytes()), slot); + } + fn cfg() -> CircuitBreakerConfig { CircuitBreakerConfig { enabled: true, diff --git a/crates/gateway/src/lifecycle/mod.rs b/crates/gateway/src/lifecycle/mod.rs index 35061314..0b72fb79 100644 --- a/crates/gateway/src/lifecycle/mod.rs +++ b/crates/gateway/src/lifecycle/mod.rs @@ -28,7 +28,7 @@ use rust_decimal::Decimal; use think_watch_common::audit::{AuditActor, AuditEntry, GatewayActor}; use think_watch_common::lifecycle::Surface; use think_watch_common::lifecycle::state::{CapturedView, Invoked, LimitCheckRecord}; -use think_watch_common::limits::{BudgetCap, RateLimitRule}; +use think_watch_common::limits::RequestLimits; use tw_dialect::ir::Dialect; use crate::guards::Guards; @@ -108,8 +108,7 @@ pub(crate) struct ChatRequestSnapshot { /// Pre-flight rule + cap lists, reused by the post-flight debit. pub(crate) struct ChatPreflightLists { - pub request_rules: Vec, - pub budget_caps: Vec, + pub limits: RequestLimits, } /// The route that actually served the request. @@ -472,20 +471,19 @@ impl Surface for ChatCompletionSurface { .audit(action) } - fn rate_limited_response(label: &str) -> Self::Response { + fn rate_limited_response(label: &str, retry_after_secs: u64) -> Self::Response { // `LocalRateLimited` so `status_code() == 429` on the wire. - // The label (`":/"`) is what the - // pre-migration `preflight_request_limits` already produced; - // keeping it intact lets clients diff exhausted windows. - ChatCompletionOutcome::ShortCircuit(GatewayError::LocalRateLimited(label.to_owned())) + // The label (`":/"`) lets clients + // diff exhausted windows; `Retry-After` says when this one + // has room again. + ChatCompletionOutcome::ShortCircuit(GatewayError::rate_limited(label, retry_after_secs)) } fn rate_limiter_unavailable_response() -> Self::Response { - // Matches the pre-migration `preflight_request_limits` - // fail-closed path — same `LocalRateLimited` variant with - // the sentinel label dashboards already filter on. - ChatCompletionOutcome::ShortCircuit(GatewayError::LocalRateLimited( - "rate_limiter_unavailable".to_owned(), + // Same `LocalRateLimited` variant with the sentinel label + // dashboards already filter on. + ChatCompletionOutcome::ShortCircuit(GatewayError::limiter_unavailable( + "rate_limiter_unavailable", )) } @@ -507,18 +505,16 @@ impl Surface for ChatCompletionSurface { ))) } - fn budget_exceeded_response(label: &str) -> Self::Response { - // `LocalRateLimited` per its docstring's explicit budget - // coverage — wire status 429, label carries which cap fired - // (e.g. `"user:budget/monthly"`) so dashboards can split - // budget exhaustion from rate-limit hits. - ChatCompletionOutcome::ShortCircuit(GatewayError::LocalRateLimited(label.to_owned())) + fn budget_exceeded_response(label: &str, retry_after_secs: u64) -> Self::Response { + // Wire status 429; the label carries which cap fired (e.g. + // `"user:budget/monthly"`) so dashboards can split budget + // exhaustion from rate-limit hits. `Retry-After` is the end of + // the cap's period, and SDKs are told not to retry by themselves. + ChatCompletionOutcome::ShortCircuit(GatewayError::budget_exhausted(label, retry_after_secs)) } fn budget_unavailable_response() -> Self::Response { - ChatCompletionOutcome::ShortCircuit(GatewayError::LocalRateLimited( - "budget_unavailable".to_owned(), - )) + ChatCompletionOutcome::ShortCircuit(GatewayError::limiter_unavailable("budget_unavailable")) } async fn record_outcome(deps: &Self::PostInvokeDeps, invoked: &Invoked) { @@ -574,8 +570,7 @@ impl Surface for ChatCompletionSurface { deps.state.weight_cache.clone(), deps.request.mapped_model.clone(), priced(&extract_usage(&invoked.view)), - deps.preflight.request_rules.clone(), - deps.preflight.budget_caps.clone(), + &deps.preflight.limits, deps.request.identity.user_id.clone(), deps.request.identity.user_email.clone(), deps.request.identity.api_key_id.clone(), diff --git a/crates/gateway/src/proxy/accounting.rs b/crates/gateway/src/proxy/accounting.rs index 770372be..3d17dda9 100644 --- a/crates/gateway/src/proxy/accounting.rs +++ b/crates/gateway/src/proxy/accounting.rs @@ -1,12 +1,12 @@ //! Post-flight accounting + token resolution for streaming responses. //! -//! [`post_flight_account`] runs the token-metric sliding rules and -//! budget caps against the token counts the upstream returned — or the -//! estimate, when it returned none — with cache reads and writes -//! weighted apart from plain input. Used from BOTH the non-streaming branch (called -//! inline after the upstream future resolves) and the streaming branch -//! (called from the post-invoke pipeline after the SSE stream is -//! drained). +//! [`post_flight_account`] adds what a request used to its token-metric +//! sliding rules and budget caps, from the token counts the upstream +//! returned — or the estimate, when it returned none — with cache reads +//! and writes weighted apart from plain input. Used from BOTH the +//! non-streaming branch (called inline after the upstream future +//! resolves) and the streaming branch (called from the post-invoke +//! pipeline after the SSE stream is drained). //! //! All errors are logged and swallowed — by the time we get here the //! caller has already received their response, so refusing to account @@ -17,7 +17,7 @@ use std::sync::Arc; use sqlx::PgPool; use think_watch_common::dynamic_config::DynamicConfig; -use think_watch_common::limits::{self, BudgetCap, RateMetric, sliding, weight}; +use think_watch_common::limits::{self, RateMetric, RequestLimits, sliding, weight}; #[allow(clippy::too_many_arguments)] pub(crate) async fn post_flight_account( @@ -27,15 +27,11 @@ pub(crate) async fn post_flight_account( weight_cache: weight::WeightCache, model: String, tokens: weight::TokenCounts, - request_rules: Vec, - budget_caps: Vec, + request_limits: &RequestLimits, // Actor attribution for `budget.threshold_crossed` audit entries. - // Without these the crossing log carries only `cap_id`, and the - // cap's `subject_id` may be a team/role — operators investigating - // a 100 %-cross had to time-join against gateway_logs to find the - // user who pushed it over. Cloned in the streaming path (the - // tokio task moves owned values) and inlined in the non-streaming - // path; passing `Option` keeps both call shapes flat. + // Without these the crossing log carries only `cap_id`, and + // operators investigating a 100 %-cross had to time-join against + // gateway_logs to find the user who pushed it over. actor_user_id: Option, actor_user_email: Option, actor_api_key_id: Option, @@ -48,27 +44,29 @@ pub(crate) async fn post_flight_account( return; } - // Token-metric sliding rules — same rule set the pre-flight - // loaded, filtered to tokens. Post-flight always runs fail-open - // because the response has already been delivered: refusing to - // record the spend would just hide it from analytics without - // recovering anything. - let resolved_token_rules = sliding::resolve_rules(&request_rules, RateMetric::Tokens); - if !resolved_token_rules.is_empty() - && let Err(e) = - sliding::check_and_record(&redis, &resolved_token_rules, weighted, true).await + // Token-metric sliding rules — the user's and the key's, the same + // rules the pre-flight checked. Recorded whatever they come to: a + // window this request overshoots refuses the next one. Post-flight + // always runs fail-open because the response has already been + // delivered: refusing to record the spend would just hide it from + // analytics without recovering anything. + if let Err(e) = sliding::record( + &redis, + &request_limits.rules, + request_limits.owner, + RateMetric::Tokens, + weighted, + ) + .await { tracing::warn!("token rate-limit accounting failed: {e}"); } - // Natural-period budget caps derived from the user's merged - // role-inline constraints. `db` is unused here — kept on the - // signature so future per-user overrides can fold in cleanly. - let _ = db; - if !budget_caps.is_empty() { - let caps = budget_caps; + // Natural-period budget caps — the user's and the key's. + if !request_limits.caps.is_empty() { + let caps = &request_limits.caps; { - match limits::budget::add_weighted_tokens(&redis, &caps, weighted).await { + match limits::budget::add_weighted_tokens(&redis, caps, weighted).await { Ok((_statuses, crossings)) if !crossings.is_empty() => { // Emit one `budget.threshold_crossed` audit-log // entry per crossing. The action is namespaced diff --git a/crates/gateway/src/proxy/generate.rs b/crates/gateway/src/proxy/generate.rs index e391ece7..9e93eab6 100644 --- a/crates/gateway/src/proxy/generate.rs +++ b/crates/gateway/src/proxy/generate.rs @@ -849,8 +849,7 @@ async fn run( input_estimate, }, preflight: crate::lifecycle::ChatPreflightLists { - request_rules: preflight.request_rules.clone(), - budget_caps: preflight.budget_caps.clone(), + limits: preflight.limits.clone(), }, route: crate::lifecycle::ChatPickedRoute { provider_name: route.provider_name.clone(), diff --git a/crates/gateway/src/proxy/identity.rs b/crates/gateway/src/proxy/identity.rs index 1b152391..82747ef1 100644 --- a/crates/gateway/src/proxy/identity.rs +++ b/crates/gateway/src/proxy/identity.rs @@ -1,78 +1,87 @@ -//! Materialize merged surface constraints into rate-limit rule rows -//! and budget caps keyed by the authenticated user id. +//! Materialize the constraints the auth middleware resolved into the +//! rate-limit rules and budget caps one request is held to. use uuid::Uuid; use super::GatewayRequestIdentity; -use think_watch_common::limits::{ - BudgetCap, BudgetSubject, RateLimitRule, RateLimitSubject, Surface, -}; +use think_watch_common::limits::{RequestLimits, Surface}; -/// Materialize the merged surface constraints into `RateLimitRule` -/// rows keyed by the authenticated user id. Redis counters now live -/// at `ratelimit::user::...` — one set per user -/// regardless of how many roles they hold. Roles merged to empty -/// (no user_id, or no rules) produce an empty list. -pub(super) fn rules_for_ai_gateway(identity: &GatewayRequestIdentity) -> Vec { - let Some(user_id) = identity - .user_id - .as_deref() - .and_then(|s| Uuid::parse_str(s).ok()) - else { - return Vec::new(); +/// The user's limits (role defaults with the user's overrides) on the +/// user's counters, and the calling key's own limits on its lineage's +/// counters — both apply. A request with no user (never past the auth +/// middleware today) is held to nothing. +pub(super) fn limits_for_ai_gateway(identity: &GatewayRequestIdentity) -> RequestLimits { + let parse = |s: &Option| s.as_deref().and_then(|s| Uuid::parse_str(s).ok()); + let Some(user_id) = parse(&identity.user_id) else { + return RequestLimits::default(); }; - let Some(block) = identity.surface_constraints.block(Surface::AiGateway) else { - return Vec::new(); - }; - block - .rules - .iter() - .filter(|r| r.enabled) - .map(|r| RateLimitRule { - // Synthetic id — stable across a single request so the - // exceeded_index in `CheckOutcome` maps back to the same - // rule without needing a persistence layer. - id: Uuid::nil(), - subject_kind: RateLimitSubject::User, - subject_id: user_id, - surface: Surface::AiGateway, - metric: r.metric, - window_secs: r.window_secs, - max_count: r.max_count, - enabled: true, - // In-memory synthesis — override metadata lives on persisted rows only. - expires_at: None, - reason: None, - created_by: None, - }) - .collect() + let key = + parse(&identity.api_key_lineage_id).map(|lineage| (lineage, &identity.key_constraints)); + RequestLimits::for_request( + Surface::AiGateway, + user_id, + &identity.surface_constraints, + key, + ) } -pub(super) fn budgets_for_ai_gateway(identity: &GatewayRequestIdentity) -> Vec { - let Some(user_id) = identity - .user_id - .as_deref() - .and_then(|s| Uuid::parse_str(s).ok()) - else { - return Vec::new(); - }; - let Some(block) = identity.surface_constraints.block(Surface::AiGateway) else { - return Vec::new(); +#[cfg(test)] +mod tests { + use super::*; + use think_watch_common::limits::{ + RateLimitSubject, RateMetric, SurfaceBlock, SurfaceConstraints, SurfaceRule, }; - block - .budgets - .iter() - .filter(|b| b.enabled) - .map(|b| BudgetCap { - id: Uuid::nil(), - subject_kind: BudgetSubject::User, - subject_id: user_id, - period: b.period, - limit_tokens: b.limit_tokens, - enabled: true, - expires_at: None, - reason: None, - created_by: None, - }) - .collect() + + fn one_rule(max_count: i64) -> SurfaceConstraints { + SurfaceConstraints { + ai_gateway: Some(SurfaceBlock { + rules: vec![SurfaceRule { + metric: RateMetric::Requests, + window_secs: 60, + max_count, + enabled: true, + }], + budgets: vec![], + }), + mcp_gateway: None, + } + } + + #[test] + fn a_keys_rules_count_on_its_lineage_and_its_owners_on_the_user() { + let user = Uuid::new_v4(); + let lineage = Uuid::new_v4(); + let identity = GatewayRequestIdentity { + user_id: Some(user.to_string()), + api_key_id: Some(Uuid::new_v4().to_string()), + api_key_lineage_id: Some(lineage.to_string()), + surface_constraints: one_rule(10), + key_constraints: one_rule(2), + ..Default::default() + }; + let limits = limits_for_ai_gateway(&identity); + assert_eq!(limits.owner, user); + let subjects: Vec<_> = limits + .rules + .iter() + .map(|r| (r.subject_kind, r.subject_id, r.max_count)) + .collect(); + assert_eq!( + subjects, + vec![ + (RateLimitSubject::User, user, 10), + (RateLimitSubject::ApiKeyLineage, lineage, 2), + ] + ); + } + + #[test] + fn no_user_no_limits() { + let identity = GatewayRequestIdentity { + surface_constraints: one_rule(1), + ..Default::default() + }; + let limits = limits_for_ai_gateway(&identity); + assert!(limits.rules.is_empty() && limits.caps.is_empty()); + } } diff --git a/crates/gateway/src/proxy/mod.rs b/crates/gateway/src/proxy/mod.rs index 4bebb01d..45088b48 100644 --- a/crates/gateway/src/proxy/mod.rs +++ b/crates/gateway/src/proxy/mod.rs @@ -114,12 +114,17 @@ pub struct GatewayRequestIdentity { /// PG via `rotated_from_id`. pub api_key_lineage_id: Option, pub allowed_models: Option>, - /// Merged-across-roles inline limits (most restrictive per - /// surface+metric+window / surface+period). Computed once by the - /// auth middleware via `rbac::compute_user_surface_constraints` - /// and consumed directly here — no side-table lookups on the - /// hot path. + /// The user's limits: merged-across-roles inline limits (most + /// restrictive per surface+metric+window / surface+period) with the + /// user's own overrides on top. Computed once by the auth middleware + /// via `rbac::compute_user_surface_constraints` and consumed directly + /// here — no side-table lookups on the hot path. Counted on the + /// user's counters. pub surface_constraints: SurfaceConstraints, + /// The calling key's own limits (`rbac::compute_key_surface_constraints`), + /// counted on the key lineage's counters. They apply on top of + /// `surface_constraints`, never instead of them. + pub key_constraints: SurfaceConstraints, /// Resolved client IP (honours `client_ip_source` + `trusted_proxies`). /// Populated by the API-key middleware via `extract_client_ip` so /// every `gateway_logs` row carries it without each handler reading @@ -190,8 +195,8 @@ impl IntoResponse for GatewayErrorResponse { body, ) .into_response(); - // Echo the upstream's Retry-After (or our local default) so - // well-behaved clients back off the right amount instead of + // Echo the upstream's Retry-After (or when our own limit frees) + // so well-behaved clients back off the right amount instead of // burning quota with tight 3× retries that all hit the same // open window. if let Some(secs) = self.error.retry_after_secs() @@ -199,6 +204,14 @@ impl IntoResponse for GatewayErrorResponse { { response.headers_mut().insert(header::RETRY_AFTER, v); } + // A spent budget: the OpenAI and Anthropic SDKs retry a 429 by + // themselves unless told not to. + if self.error.should_retry() == Some(false) { + response.headers_mut().insert( + header::HeaderName::from_static("x-should-retry"), + HeaderValue::from_static("false"), + ); + } response } } @@ -239,7 +252,7 @@ mod helper_tests { }, 429, ), - (GatewayError::LocalRateLimited("rule".into()), 429), + (GatewayError::rate_limited("rule", 5), 429), ( GatewayError::UpstreamAuthError { status: 401, @@ -355,16 +368,37 @@ mod helper_tests { "no header when upstream didn't tell us — guessing would mislead clients" ); - // Local limit — we set our own conservative default so SDKs - // see a number instead of immediately retrying. - let resp = GatewayErrorResponse::from(GatewayError::LocalRateLimited("budget".into())) + // Local limit — when the window has room again, and nothing + // that stops an SDK's own retries. + let resp = GatewayErrorResponse::from(GatewayError::rate_limited("user:requests/1m", 42)) .into_response(); assert_eq!(resp.status().as_u16(), 429); - assert!( + assert_eq!( resp.headers() .get(axum::http::header::RETRY_AFTER) - .is_some(), - "local rate-limit must carry a Retry-After default" + .and_then(|v| v.to_str().ok()), + Some("42") + ); + assert!(resp.headers().get("x-should-retry").is_none()); + + // A spent budget — the end of its period, and no retries. + let resp = GatewayErrorResponse::from(GatewayError::budget_exhausted( + "user:budget/monthly", + 1_234_567, + )) + .into_response(); + assert_eq!(resp.status().as_u16(), 429); + assert_eq!( + resp.headers() + .get(axum::http::header::RETRY_AFTER) + .and_then(|v| v.to_str().ok()), + Some("1234567") + ); + assert_eq!( + resp.headers() + .get("x-should-retry") + .and_then(|v| v.to_str().ok()), + Some("false") ); // Non-429 responses must NOT carry Retry-After — would @@ -384,10 +418,11 @@ mod helper_tests { async fn the_error_body_is_in_the_callers_format() { use tw_dialect::ir::Dialect; async fn body(d: Dialect) -> serde_json::Value { - let resp = GatewayErrorResponse::from(GatewayError::LocalRateLimited("rule".into())) + let resp = GatewayErrorResponse::from(GatewayError::budget_exhausted("rule", 60)) .in_dialect(d) .into_response(); assert_eq!(resp.status().as_u16(), 429); + assert_eq!(resp.headers()["x-should-retry"], "false"); let bytes = axum::body::to_bytes(resp.into_body(), usize::MAX) .await .unwrap(); diff --git a/crates/gateway/src/proxy/pipeline.rs b/crates/gateway/src/proxy/pipeline.rs index 6eb7de88..506f3175 100644 --- a/crates/gateway/src/proxy/pipeline.rs +++ b/crates/gateway/src/proxy/pipeline.rs @@ -12,7 +12,7 @@ //! Plus [`LogCtx::new`] (in `log_ctx.rs`) builds the audit context in //! one call instead of 12-field literals at every site. -use super::identity::{budgets_for_ai_gateway, rules_for_ai_gateway}; +use super::identity::limits_for_ai_gateway; use super::{GatewayErrorResponse, GatewayRequestIdentity, GatewayState}; use super::shaper::StreamShaper; @@ -24,19 +24,20 @@ use think_watch_common::lifecycle::stages::{ check_access, check_budget, check_limits, run_post_invoke, }; use think_watch_common::lifecycle::state::{CapturedView, Invoked, LimitCheckRecord, Raw}; -use think_watch_common::limits::{BudgetCap, RateLimitRule}; +use think_watch_common::limits::RequestLimits; use tw_dialect::ir::Dialect; /// Pre-flight result threaded through to `ChatPostInvokeDeps` later /// in the handler. Computed once by [`run_preflight_stages`] so the /// caller doesn't have to re-derive the rule/cap lists from identity. pub(super) struct Preflight { - pub(super) request_rules: Vec, - pub(super) budget_caps: Vec, + pub(super) limits: RequestLimits, } -/// Run the three shared pre-flight stages — `check_limits` → -/// `check_budget` → `check_access` — that each AI surface gates on. +/// Run the three shared pre-flight stages — `check_budget` → +/// `check_limits` → `check_access` — that each AI surface gates on. +/// The budget peek goes first because it charges nothing: a request a +/// spent budget refuses must not have used up a request limit. /// /// Returns the resolved rule + cap lists so the caller can feed them /// into [`ChatPostInvokeDeps`] without re-walking the identity. @@ -50,8 +51,7 @@ pub(super) async fn run_preflight_stages( trace_id: &str, model: &str, ) -> Result { - let request_rules = rules_for_ai_gateway(identity); - let budget_caps = budgets_for_ai_gateway(identity); + let limits = limits_for_ai_gateway(identity); let fail_closed = state.dynamic_config.rate_limit_fail_closed().await; let raw = Raw::::new( @@ -59,18 +59,19 @@ pub(super) async fn run_preflight_stages( trace_id.to_string(), identity.ip_address.clone(), ); - let limits_checked = check_limits::( + let raw = check_budget::( raw, - &request_rules, + &limits.caps, &state.redis, fail_closed, &state.audit, ) .await .map_err(short_circuit_to_response)?; - let limits_checked = check_budget::( - limits_checked, - &budget_caps, + let limits_checked = check_limits::( + raw, + &limits.rules, + limits.owner, &state.redis, fail_closed, &state.audit, @@ -81,10 +82,7 @@ pub(super) async fn run_preflight_stages( .await .map_err(short_circuit_to_response)?; - Ok(Preflight { - request_rules, - budget_caps, - }) + Ok(Preflight { limits }) } fn short_circuit_to_response(outcome: ChatCompletionOutcome) -> GatewayErrorResponse { diff --git a/crates/gateway/src/quota.rs b/crates/gateway/src/quota.rs index 59f2efdc..31d2ef50 100644 --- a/crates/gateway/src/quota.rs +++ b/crates/gateway/src/quota.rs @@ -37,13 +37,16 @@ impl QuotaManager { chrono::Utc::now().format("%Y-%m").to_string() } + // `{}` is a Redis Cluster hash tag: `consume` and + // `check_and_consume` declare both keys in one script, which a + // cluster runs only when they share a slot. fn limit_key(key: &str) -> String { - format!("quota:{key}:limit") + format!("quota:{{{key}}}:limit") } fn usage_key(key: &str) -> String { let month = Self::current_month(); - format!("quota:{key}:used:{month}") + format!("quota:{{{key}}}:used:{month}") } /// Check if user/team has enough quota. Returns remaining tokens. @@ -213,3 +216,22 @@ return limit - used - tokens }) } } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn a_quotas_limit_and_usage_share_one_cluster_slot() { + for key in [ + "8a7c0c1e-0000-0000-0000-000000000000:gpt-4o", + "k:{odd}:model", + ] { + assert_eq!( + fred::util::redis_keyslot(QuotaManager::limit_key(key).as_bytes()), + fred::util::redis_keyslot(QuotaManager::usage_key(key).as_bytes()), + "{key}" + ); + } + } +} diff --git a/crates/gateway/src/rate_limiter.rs b/crates/gateway/src/rate_limiter.rs index 521df847..d43d0a80 100644 --- a/crates/gateway/src/rate_limiter.rs +++ b/crates/gateway/src/rate_limiter.rs @@ -95,7 +95,9 @@ impl RateLimiter { let now_ms = chrono::Utc::now().timestamp_millis() as f64; let window_start = now_ms - 60_000.0; let member_id = uuid::Uuid::new_v4().to_string(); - let rpm_key = format!("ratelimit:rpm:{key}"); + // `{}` is a Redis Cluster hash tag: the combined script + // declares both keys, which a cluster runs only in one slot. + let rpm_key = format!("ratelimit:rpm:{{{key}}}"); let reset_at = chrono::Utc::now().timestamp() + 60; // window resets in ~60s // When BOTH RPM and TPM are configured, evaluate them atomically @@ -106,7 +108,7 @@ impl RateLimiter { if let (Some(tpm_limit), Some(tokens)) = (tpm_limit, estimated_tokens) && tokens > 0 { - let tpm_key = format!("ratelimit:tpm:{key}"); + let tpm_key = format!("ratelimit:tpm:{{{key}}}"); let member_with_tokens = format!("{member_id}:{tokens}"); let result: Vec = self .redis diff --git a/crates/mcp-gateway/src/cache.rs b/crates/mcp-gateway/src/cache.rs index 8f7374e1..2d7936aa 100644 --- a/crates/mcp-gateway/src/cache.rs +++ b/crates/mcp-gateway/src/cache.rs @@ -254,40 +254,23 @@ impl McpResponseCache { user_id: Option, server_id: &Uuid, ) { - let mut cursor: String = "0".to_string(); - let mut deleted: usize = 0; - loop { - let page: Result<(String, Vec), _> = self - .redis - .scan_page( - cursor.clone(), - pattern.clone(), - Some(256), - Some(ScanType::String), - ) - .await; - let (next, keys) = match page { - Ok(p) => p, - Err(e) => { - tracing::warn!( - server = %server_id, user = ?user_id, scope, error = %e, - "MCP cache invalidate: SCAN failed; some stale entries may persist" - ); - return; - } - }; - if !keys.is_empty() { - let n: Result = self.redis.del(keys.clone()).await; - match n { - Ok(n) => deleted += n as usize, - Err(e) => tracing::warn!(error = %e, scope, "MCP cache invalidate: DEL failed"), - } - } - if next == "0" { - break; + // On every node, slot by slot, when Redis is a cluster. + let deleted = match think_watch_common::redis_keys::delete_matching( + &self.redis, + &pattern, + Some(ScanType::String), + ) + .await + { + Ok(n) => n, + Err(e) => { + tracing::warn!( + server = %server_id, user = ?user_id, scope, error = %e, + "MCP cache invalidate failed; some stale entries may persist" + ); + return; } - cursor = next; - } + }; if deleted > 0 { tracing::info!( server = %server_id, diff --git a/crates/mcp-gateway/src/lifecycle/mod.rs b/crates/mcp-gateway/src/lifecycle/mod.rs index 4c58c17a..095f343f 100644 --- a/crates/mcp-gateway/src/lifecycle/mod.rs +++ b/crates/mcp-gateway/src/lifecycle/mod.rs @@ -18,7 +18,7 @@ use think_watch_common::lifecycle::Surface; use think_watch_common::lifecycle::state::{CapturedView, Invoked}; use think_watch_common::lifecycle::streaming::StreamOutcome; use think_watch_common::limits::{ - RateLimitRule, RateLimitSubject, Surface as LimitSurface, SurfaceConstraints, + RateMetric, RequestLimits, Surface as LimitSurface, SurfaceConstraints, }; use crate::access_control::is_tool_allowed; @@ -76,10 +76,10 @@ impl Surface for McpSurface { .audit(action) } - fn rate_limited_response(label: &str) -> Self::Response { + fn rate_limited_response(label: &str, _retry_after_secs: u64) -> Self::Response { // Bumps the existing operator-facing metric so dashboards // built around `mcp_rate_limited_total` keep working after - // the migration. + // the migration. A JSON-RPC error has no `Retry-After`. metrics::counter!("mcp_rate_limited_total").increment(1); err_response(None, INVALID_REQUEST, format!("Rate limited: {label}")) } @@ -104,7 +104,7 @@ impl Surface for McpSurface { err_response(None, INVALID_REQUEST, "Access denied for this tool") } - fn budget_exceeded_response(label: &str) -> Self::Response { + fn budget_exceeded_response(label: &str, _retry_after_secs: u64) -> Self::Response { // MCP doesn't currently wire budget caps into // `handle_tools_call` (the AI gateway is the only consumer // of `check_budget` today). The factory exists so the @@ -298,34 +298,18 @@ fn response_for_hooks(invoked: &Invoked, deps: &McpPostInvokeDeps) - } } -/// Build the MCP-surface `requests` rate-limit rules from a -/// materialised [`SurfaceConstraints`]. Mirrors the inline -/// extraction in the previous `handle_tools_call` — same shape, -/// just lifted out so the surface stage gets a clean -/// `&[RateLimitRule]` slice without re-implementing the -/// extraction at every call site. -pub fn rate_limit_rules(constraints: &SurfaceConstraints, user_id: Uuid) -> Vec { - constraints - .block(LimitSurface::McpGateway) - .map(|block| { - block - .rules - .iter() - .filter(|r| r.enabled) - .map(|r| RateLimitRule { - id: Uuid::nil(), - subject_kind: RateLimitSubject::User, - subject_id: user_id, - surface: LimitSurface::McpGateway, - metric: r.metric, - window_secs: r.window_secs, - max_count: r.max_count, - enabled: true, - expires_at: None, - reason: None, - created_by: None, - }) - .collect() - }) - .unwrap_or_default() +/// The MCP-surface rate limits one `tools/call` is held to: the user's +/// (role defaults with the user's overrides) on the user's counters, and +/// the calling key's own on its lineage's — both apply. `requests` rules +/// only: a tool call has no tokens to count, so a `tokens` rule on this +/// surface would only ever read zero. +pub fn rate_limits( + user_id: Uuid, + user: &SurfaceConstraints, + key: Option<(Uuid, &SurfaceConstraints)>, +) -> RequestLimits { + let mut limits = RequestLimits::for_request(LimitSurface::McpGateway, user_id, user, key); + limits.rules.retain(|r| r.metric == RateMetric::Requests); + limits.caps.clear(); + limits } diff --git a/crates/mcp-gateway/src/proxy.rs b/crates/mcp-gateway/src/proxy.rs index 1e0af486..d268bdcc 100644 --- a/crates/mcp-gateway/src/proxy.rs +++ b/crates/mcp-gateway/src/proxy.rs @@ -40,7 +40,12 @@ pub struct RequestContext<'a> { pub user_id: Uuid, pub user_email: &'a str, pub client_session_id: &'a str, + /// The user's limits, counted on the user's counters. pub surface_constraints: &'a SurfaceConstraints, + /// The calling key's lineage and its own limits, counted on the + /// lineage's counters on top of the user's. + pub api_key_lineage_id: Option, + pub key_constraints: &'a SurfaceConstraints, pub allowed_mcp_tools: Option<&'a [String]>, pub trace_id: &'a str, /// Per-server MCP account override JSON from the calling API key @@ -586,7 +591,12 @@ impl McpProxy { // stays identical to the pre-migration version. We bind // the `request.id` onto the response after the fact // because the stage doesn't know the wire-level id. - let rules = crate::lifecycle::rate_limit_rules(surface_constraints, user_id); + let limits = crate::lifecycle::rate_limits( + user_id, + surface_constraints, + ctx.api_key_lineage_id + .map(|lineage| (lineage, ctx.key_constraints)), + ); let fail_closed = self.dynamic_config.rate_limit_fail_closed().await; let raw = think_watch_common::lifecycle::state::Raw::::new( crate::lifecycle::McpIdentity { @@ -601,7 +611,14 @@ impl McpProxy { ); let limits_checked = match think_watch_common::lifecycle::stages::check_limits::< crate::lifecycle::McpSurface, - >(raw, &rules, &self.redis, fail_closed, &self.audit) + >( + raw, + &limits.rules, + limits.owner, + &self.redis, + fail_closed, + &self.audit, + ) .await { Ok(s) => s, diff --git a/crates/mcp-gateway/src/transport/streamable_http.rs b/crates/mcp-gateway/src/transport/streamable_http.rs index 380a4d36..a0f86554 100644 --- a/crates/mcp-gateway/src/transport/streamable_http.rs +++ b/crates/mcp-gateway/src/transport/streamable_http.rs @@ -38,11 +38,17 @@ pub struct McpRequestIdentity { /// keys (no associated user) — those will be denied any tool /// that requires a role match. pub user_roles: Vec, - /// Effective (most-restrictive across roles) rate-limit rules and - /// budget caps for this user. Materialized by the parent crate - /// once per request so the MCP proxy can apply them without a DB - /// round-trip. + /// Effective (most-restrictive across roles, with the user's own + /// overrides) rate-limit rules and budget caps for this user. + /// Materialized by the parent crate once per request so the MCP + /// proxy can apply them without a DB round-trip. pub surface_constraints: think_watch_common::limits::SurfaceConstraints, + /// The calling key's lineage id. Limits attached to the key live + /// on it so they survive rotation. + pub api_key_lineage_id: Uuid, + /// The calling key's own limits, counted on the lineage's counters + /// on top of the user's. + pub key_constraints: think_watch_common::limits::SurfaceConstraints, /// MCP tool access patterns from role union. `None` = unrestricted. pub allowed_mcp_tools: Option>, /// Per-server account-label override map carried by the calling @@ -134,6 +140,8 @@ pub async fn handle_post( user_email: &identity.user_email, client_session_id: &session_id, surface_constraints: &identity.surface_constraints, + api_key_lineage_id: Some(identity.api_key_lineage_id), + key_constraints: &identity.key_constraints, allowed_mcp_tools: identity.allowed_mcp_tools.as_deref(), trace_id: &trace_id, mcp_account_overrides: &identity.mcp_account_overrides, diff --git a/crates/server/src/handlers/limits.rs b/crates/server/src/handlers/limits.rs index 1d49b64f..7a98f347 100644 --- a/crates/server/src/handlers/limits.rs +++ b/crates/server/src/handlers/limits.rs @@ -675,21 +675,27 @@ pub async fn get_usage( let rate_subject = parse_rate_subject(&kind)?; let storage_id = resolve_subject_id(&state.db, &kind, subject_id).await?; let rules = limits::list_rules(&state.db, rate_subject, storage_id).await?; + // The counters the gateway writes carry the requesting user's hash + // tag: for a key, its owner. A key without one has never been + // let through, so it has counted nothing. + let owner = match rate_subject { + RateLimitSubject::User => Some(subject_id), + RateLimitSubject::ApiKeyLineage => { + sqlx::query_scalar::<_, Option>("SELECT user_id FROM api_keys WHERE id = $1") + .bind(subject_id) + .fetch_optional(&state.db) + .await? + .flatten() + } + }; let mut rule_usage: Vec = Vec::with_capacity(rules.len()); for r in &rules { - let resolved = sliding::ResolvedRule { - id: r.id, - base_key: sliding::build_base_key( - r.surface.as_str(), - r.subject_kind.as_str(), - r.subject_id, - r.metric, - r.window_secs, - ), - bucket_secs: sliding::bucket_secs(r.window_secs), - max_count: r.max_count, + let current = match owner { + Some(owner) => { + sliding::current_count(&state.redis, &sliding::ResolvedRule::new(r, owner)).await + } + None => 0, }; - let current = sliding::current_count(&state.redis, &resolved).await; rule_usage.push(RuleUsage { rule_id: r.id, current, diff --git a/crates/server/src/handlers/user_limits.rs b/crates/server/src/handlers/user_limits.rs index 9d97baa8..fa9e76eb 100644 --- a/crates/server/src/handlers/user_limits.rs +++ b/crates/server/src/handlers/user_limits.rs @@ -217,9 +217,10 @@ async fn build_effective_rules( // shouldn't 500 the whole dashboard, just show 0. let resolved = sliding::ResolvedRule { id: ov.map(|o| o.id).unwrap_or(Uuid::nil()), - base_key: sliding::build_base_key( - surface.as_str(), - "user", + key: sliding::counter_key( + user_id, + surface, + RateLimitSubject::User, user_id, rule.metric, rule.window_secs, @@ -518,27 +519,22 @@ async fn reset_rule_counter( .window_secs .ok_or_else(|| AppError::BadRequest("window_secs is required for rule reset".into()))?; - let base_key = sliding::build_base_key(surface.as_str(), "user", user_id, metric, window_secs); - // Buckets are timestamp-derived (`now_secs / bucket_secs`). DEL the - // 60-window range plus a small safety margin in case a late write - // lands after we read the clock. - let bucket_secs = sliding::bucket_secs(window_secs) as i64; - if bucket_secs <= 0 { - return Ok(0); - } - let now = chrono::Utc::now().timestamp(); - let current_bucket = now / bucket_secs; - let mut deleted = 0usize; - for b in 0..(sliding::BUCKETS_PER_WINDOW + 2) { - let key = format!("{}:{}", base_key, current_bucket - b); - match state.redis.del::(&key).await { - Ok(n) => deleted += n as usize, - Err(e) => { - tracing::warn!("reset_rule_counter DEL failed for {key}: {e}"); - } + // One hash holds every bucket of the window. + let key = sliding::counter_key( + user_id, + surface, + RateLimitSubject::User, + user_id, + metric, + window_secs, + ); + match state.redis.del::(&key).await { + Ok(n) => Ok(n as usize), + Err(e) => { + tracing::warn!("reset_rule_counter DEL failed for {key}: {e}"); + Ok(0) } } - Ok(deleted) } async fn reset_cap_counter( diff --git a/crates/server/src/middleware/api_key_auth.rs b/crates/server/src/middleware/api_key_auth.rs index 55eda7c0..1f13aad4 100644 --- a/crates/server/src/middleware/api_key_auth.rs +++ b/crates/server/src/middleware/api_key_auth.rs @@ -292,24 +292,24 @@ pub fn require_api_key( // can gate per-tool access without re-querying the DB, and // the aggregated `surface_constraints` JSON so the gateway // hot path has rate limits + budgets without further lookups. - let (role_limits, user_roles, surface_constraints) = if let Some(uid) = row.user_id { + let (role_limits, user_roles, surface_constraints, key_constraints) = if let Some(uid) = row.user_id { let limits = rbac::compute_user_resource_limits(&state.db, uid) .await .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; let names = rbac::load_user_role_names(&state.db, uid) .await .unwrap_or_default(); - // Use the api_key-aware variant so per-key - // `rate_limit_rules` / `budget_caps` rows fire on the - // gateway hot path. Falling back to the user-only - // function would silently drop api_key-scope - // overrides — the schema supports them but the - // gateway would never see them. - let constraints = - rbac::compute_effective_surface_constraints(&state.db, uid, row.id) + // The user's limits and the key's own, kept apart: + // the gateway counts each on its own counters and + // checks both. + let constraints = rbac::compute_user_surface_constraints(&state.db, uid) + .await + .unwrap_or_default(); + let key_constraints = + rbac::compute_key_surface_constraints(&state.db, row.lineage_id) .await .unwrap_or_default(); - (limits, names, constraints) + (limits, names, constraints, key_constraints) } else { // A key without an owner (its user row was removed // and `user_id` set NULL) has no roles to grant @@ -320,6 +320,7 @@ pub fn require_api_key( rbac::UserResourceLimits::none(), Vec::new(), think_watch_common::limits::SurfaceConstraints::default(), + think_watch_common::limits::SurfaceConstraints::default(), ) }; @@ -382,6 +383,7 @@ pub fn require_api_key( api_key_lineage_id: Some(row.lineage_id.to_string()), allowed_models: merged_models.clone(), surface_constraints: surface_constraints.clone(), + key_constraints: key_constraints.clone(), ip_address: client_ip.clone(), }; @@ -414,6 +416,8 @@ pub fn require_api_key( user_email, user_roles, surface_constraints: surface_constraints.clone(), + api_key_lineage_id: row.lineage_id, + key_constraints, allowed_mcp_tools: merged_mcp_tools.clone(), mcp_account_overrides: row.mcp_account_overrides.clone(), ip_address: client_ip.clone(), diff --git a/crates/test-support/src/lib.rs b/crates/test-support/src/lib.rs index e7eae132..868fecd4 100644 --- a/crates/test-support/src/lib.rs +++ b/crates/test-support/src/lib.rs @@ -20,6 +20,7 @@ pub mod client; pub mod fixtures; pub mod mock_provider; pub mod pg; +pub mod redis_scripts; use std::sync::Arc; @@ -85,6 +86,10 @@ pub struct SpawnOptions { /// to exercise the offload path inject an in-memory store here /// without standing up a real S3 backend. pub blob_store: Option>, + /// Boot against this Redis instead of `TEST_REDIS_URL`'s per-slot + /// logical DB — e.g. a `redis-cluster://` URL. Nothing is flushed: + /// the test must use keys no other test touches. + pub redis_url: Option, } /// An SSRF guard that lets a `wiremock` on `127.0.0.1` through and still @@ -151,7 +156,11 @@ impl TestApp { let redis_url = std::env::var("TEST_REDIS_URL").unwrap_or_else(|_| { "redis://:225b3facaf55212ff86ad6595e6d6471@localhost:6379/1".into() }); - let redis_url = redis_url_for_slot(&redis_url)?; + let shared_redis = opts.redis_url.is_some(); + let redis_url = match &opts.redis_url { + Some(url) => url.clone(), + None => redis_url_for_slot(&redis_url)?, + }; // Per-test database with migrations applied. let db_owner = IsolatedDatabase::create(&base_url) @@ -170,7 +179,7 @@ impl TestApp { // fred 10 doesn't expose FLUSHDB directly (only FLUSHALL), // and we don't want to nuke the dev DB. Send the raw // command so we only clear the test logical DB. - { + if !shared_redis { use fred::interfaces::ClientLike; use fred::types::{ClusterHash, CustomCommand}; let cmd = CustomCommand::new("FLUSHDB", ClusterHash::FirstKey, false); diff --git a/crates/test-support/src/redis_scripts.rs b/crates/test-support/src/redis_scripts.rs new file mode 100644 index 00000000..5109a280 --- /dev/null +++ b/crates/test-support/src/redis_scripts.rs @@ -0,0 +1,79 @@ +//! The rate-limit scripts, run against whatever Redis a test hands in — +//! one node in `tests/limits.rs`, a Redis Cluster in +//! `tests/redis_cluster.rs`. + +use uuid::Uuid; + +fn rule( + kind: think_watch_common::limits::RateLimitSubject, + subject: Uuid, + metric: think_watch_common::limits::RateMetric, + window_secs: i32, + max_count: i64, +) -> think_watch_common::limits::RateLimitRule { + think_watch_common::limits::RateLimitRule { + id: Uuid::nil(), + subject_kind: kind, + subject_id: subject, + surface: think_watch_common::limits::Surface::AiGateway, + metric, + window_secs, + max_count, + enabled: true, + expires_at: None, + reason: None, + created_by: None, + } +} + +/// The admit/record sequence both the single-node and the cluster test +/// run: windows that fill, free bucket by bucket, refuse without +/// charging, and record past a limit. +pub async fn exercise_the_limit_scripts(redis: &fred::clients::Client) { + use fred::interfaces::HashesInterface; + use think_watch_common::limits::{RateLimitSubject as S, RateMetric as M, sliding}; + + let owner = Uuid::new_v4(); + let lineage = Uuid::new_v4(); + let rules = [ + rule(S::User, owner, M::Requests, 60, 3), + rule(S::ApiKeyLineage, lineage, M::Tokens, 300, 100), + ]; + // A whole minute, so every bucket boundary below is exact. + let t0: i64 = 1_800_000_000_000; + let admit = |ms: i64| sliding::admit_at(redis, &rules, owner, t0 + ms, false); + + for (i, at) in [0, 10_000, 20_000].into_iter().enumerate() { + let o = admit(at).await.unwrap(); + assert!(o.allowed, "request {i}"); + assert_eq!(o.currents, vec![i as i64 + 1, 0]); + } + // Full. The first request's bucket leaves the window at t0 + 60 s. + let o = admit(30_000).await.unwrap(); + assert!(!o.allowed); + assert_eq!(o.exceeded_index, Some(0)); + assert_eq!(o.retry_after_secs, 30); + assert_eq!(o.currents, vec![3, 0], "a refusal charges nothing"); + + // Tokens are recorded in full, past the limit. + sliding::record_at(redis, &rules, owner, M::Tokens, 150, t0 + 30_000) + .await + .unwrap(); + // The requests window has room again; the tokens one does not, and + // frees when the 5-second bucket of t0 + 30 s leaves it, 300 s on. + let o = admit(61_000).await.unwrap(); + assert!(!o.allowed); + assert_eq!(o.exceeded_index, Some(1)); + assert_eq!(o.retry_after_secs, 330 - 61); + assert_eq!(o.currents, vec![2, 150]); + + // Both windows have moved past everything: one request counted, and + // the old buckets are gone from the hash. + let o = admit(400_000).await.unwrap(); + assert!(o.allowed); + assert_eq!(o.currents, vec![1, 0]); + let requests = sliding::ResolvedRule::new(&rules[0], owner); + let tokens = sliding::ResolvedRule::new(&rules[1], owner); + assert_eq!(redis.hlen::(&requests.key).await.unwrap(), 1); + assert_eq!(redis.hlen::(&tokens.key).await.unwrap(), 0); +} diff --git a/crates/test-support/tests/limits.rs b/crates/test-support/tests/limits.rs index 5565305a..0161533c 100644 --- a/crates/test-support/tests/limits.rs +++ b/crates/test-support/tests/limits.rs @@ -249,9 +249,8 @@ async fn rate_limit_window_validation_rejects_off_grid_seconds() { async fn api_key_scope_rate_limit_isolates_from_other_keys() { // Per-key rate-limit rules MUST fire on the gateway hot path — // schema supports `subject_kind='api_key_lineage'` and the auth - // middleware passes `api_key_id` through - // `compute_effective_surface_constraints`, which resolves it to - // `lineage_id` before binding. Two keys for the same user: one + // middleware loads the key lineage's rules through + // `compute_key_surface_constraints`. Two keys for the same user: one // carries a max_count=1 rule keyed on its lineage_id, the other // carries nothing. Each key must behave independently. let app = TestApp::spawn().await; @@ -445,3 +444,497 @@ async fn api_key_rate_limit_survives_rotation_via_lineage_id() { r.status ); } + +// ---------------------------------------------------------------------------- +// Enforcement: token limits, key counters, refusals that charge nothing, +// and the Retry-After a refusal carries. +// ---------------------------------------------------------------------------- + +fn chat() -> Json { + json!({"model": "gpt-test", "messages": [{"role": "user", "content": "x"}]}) +} + +/// `current` for every rule and cap on a subject, from the console's +/// usage endpoint — what an operator sees. +async fn usage(con: &TestClient, kind: &str, id: Uuid) -> (Vec, Vec) { + let r = con + .get(&format!("/api/admin/limits/{kind}/{id}/usage")) + .await + .unwrap(); + r.assert_ok(); + let v: Json = r.json().unwrap(); + let currents = |field: &str| -> Vec { + v[field] + .as_array() + .unwrap() + .iter() + .map(|x| x["current"].as_i64().unwrap()) + .collect() + }; + (currents("rules"), currents("caps")) +} + +fn retry_after(r: &think_watch_test_support::client::TestResponse) -> u64 { + r.headers + .get("retry-after") + .and_then(|v| v.to_str().ok()) + .and_then(|v| v.parse().ok()) + .unwrap_or_else(|| panic!("a 429 carries Retry-After, headers: {:?}", r.headers)) +} + +#[ignore = "integration test — run via `make test-it`"] +#[tokio::test] +async fn a_token_limit_refuses_once_the_window_is_used_up() { + // A request costs ~10 weighted tokens against the mock. A limit of 5 + // lets the first request through (nothing was used yet), records all + // of what it used even though that overshoots, and refuses the next. + let app = TestApp::spawn().await; + let (api_key, user_id) = seed_runtime(&app).await; + fixtures::create_rate_limit_rule(&app.db, "user", user_id, "ai_gateway", "tokens", 60, 5) + .await + .unwrap(); + let con = admin_session(&app).await; + + let gw = app.gateway_client(); + gw.set_bearer(&api_key); + gw.post("/v1/chat/completions", chat()) + .await + .unwrap() + .assert_ok(); + + let (rules, _) = usage(&con, "user", user_id).await; + assert!( + rules[0] > 5, + "the whole use is recorded, past the limit; got {rules:?}" + ); + + let r = gw.post("/v1/chat/completions", chat()).await.unwrap(); + assert_eq!(r.status.as_u16(), 429, "body={}", r.text()); + let secs = retry_after(&r); + assert!((1..=60).contains(&secs), "Retry-After {secs}"); + assert_eq!( + usage(&con, "user", user_id).await.0, + rules, + "a refused request records nothing" + ); +} + +#[ignore = "integration test — run via `make test-it`"] +#[tokio::test] +async fn every_request_limit_applies_when_there_are_several() { + let app = TestApp::spawn().await; + let (api_key, user_id) = seed_runtime(&app).await; + fixtures::create_rate_limit_rule(&app.db, "user", user_id, "ai_gateway", "requests", 60, 2) + .await + .unwrap(); + fixtures::create_rate_limit_rule(&app.db, "user", user_id, "ai_gateway", "requests", 300, 5) + .await + .unwrap(); + + let gw = app.gateway_client(); + gw.set_bearer(&api_key); + for _ in 0..2 { + gw.post("/v1/chat/completions", chat()) + .await + .unwrap() + .assert_ok(); + } + let r = gw.post("/v1/chat/completions", chat()).await.unwrap(); + assert_eq!(r.status.as_u16(), 429, "body={}", r.text()); +} + +/// A key and its owner, each with limits of their own. +async fn second_key(app: &TestApp, user_id: Uuid) -> fixtures::SeededApiKey { + fixtures::create_api_key( + &app.db, + user_id, + &unique_name("limit-key"), + &["ai_gateway"], + None, + None, + ) + .await + .unwrap() +} + +#[ignore = "integration test — run via `make test-it`"] +#[tokio::test] +async fn a_keys_limits_count_on_the_key_not_on_its_owner() { + let app = TestApp::spawn().await; + let (_, user_id) = seed_runtime(&app).await; + let key_a = second_key(&app, user_id).await; + let key_b = second_key(&app, user_id).await; + for key in [&key_a, &key_b] { + fixtures::create_rate_limit_rule( + &app.db, + "api_key_lineage", + key.row.lineage_id, + "ai_gateway", + "requests", + 60, + 1, + ) + .await + .unwrap(); + } + let con = admin_session(&app).await; + let gw = app.gateway_client(); + + gw.set_bearer(&key_a.plaintext); + gw.post("/v1/chat/completions", chat()) + .await + .unwrap() + .assert_ok(); + // Key B has a counter of its own: A's request does not use it up. + gw.set_bearer(&key_b.plaintext); + gw.post("/v1/chat/completions", chat()) + .await + .unwrap() + .assert_ok(); + gw.set_bearer(&key_a.plaintext); + let r = gw.post("/v1/chat/completions", chat()).await.unwrap(); + assert_eq!(r.status.as_u16(), 429, "body={}", r.text()); + + // The console reads the counter the gateway writes. + assert_eq!(usage(&con, "api_key", key_a.row.id).await.0, vec![1]); + assert_eq!(usage(&con, "api_key", key_b.row.id).await.0, vec![1]); +} + +#[ignore = "integration test — run via `make test-it`"] +#[tokio::test] +async fn a_key_limit_does_not_lift_its_owners() { + // The owner may make one request a minute; the key's own limit of + // five does not raise that. + let app = TestApp::spawn().await; + let (_, user_id) = seed_runtime(&app).await; + let key = second_key(&app, user_id).await; + fixtures::create_rate_limit_rule(&app.db, "user", user_id, "ai_gateway", "requests", 60, 1) + .await + .unwrap(); + fixtures::create_rate_limit_rule( + &app.db, + "api_key_lineage", + key.row.lineage_id, + "ai_gateway", + "requests", + 60, + 5, + ) + .await + .unwrap(); + + let gw = app.gateway_client(); + gw.set_bearer(&key.plaintext); + gw.post("/v1/chat/completions", chat()) + .await + .unwrap() + .assert_ok(); + let r = gw.post("/v1/chat/completions", chat()).await.unwrap(); + assert_eq!(r.status.as_u16(), 429, "body={}", r.text()); +} + +#[ignore = "integration test — run via `make test-it`"] +#[tokio::test] +async fn a_key_budget_counts_on_the_key() { + let app = TestApp::spawn().await; + let (_, user_id) = seed_runtime(&app).await; + let key = second_key(&app, user_id).await; + fixtures::create_budget_cap( + &app.db, + "api_key_lineage", + key.row.lineage_id, + "daily", + 1_000_000, + ) + .await + .unwrap(); + let con = admin_session(&app).await; + + let gw = app.gateway_client(); + gw.set_bearer(&key.plaintext); + gw.post("/v1/chat/completions", chat()) + .await + .unwrap() + .assert_ok(); + let (_, caps) = usage(&con, "api_key", key.row.id).await; + assert!( + caps[0] > 0, + "the key's budget counted the request: {caps:?}" + ); +} + +#[ignore = "integration test — run via `make test-it`"] +#[tokio::test] +async fn a_request_refused_by_a_budget_charges_no_request_limit() { + use fred::interfaces::KeysInterface; + use think_watch_common::limits::budget; + + let app = TestApp::spawn().await; + let (api_key, user_id) = seed_runtime(&app).await; + fixtures::create_rate_limit_rule(&app.db, "user", user_id, "ai_gateway", "requests", 60, 2) + .await + .unwrap(); + fixtures::create_budget_cap(&app.db, "user", user_id, "daily", 100) + .await + .unwrap(); + let spent = budget::build_key("user", user_id, "daily", chrono::Utc::now()); + let _: () = app + .state + .redis + .set(&spent, 105, None, None, false) + .await + .unwrap(); + let con = admin_session(&app).await; + + let gw = app.gateway_client(); + gw.set_bearer(&api_key); + for _ in 0..3 { + let r = gw.post("/v1/chat/completions", chat()).await.unwrap(); + assert_eq!(r.status.as_u16(), 429, "body={}", r.text()); + } + assert_eq!(usage(&con, "user", user_id).await.0, vec![0]); + + // With the budget freed, both requests the limit allows go through. + let _: () = app.state.redis.del(&spent).await.unwrap(); + for _ in 0..2 { + gw.post("/v1/chat/completions", chat()) + .await + .unwrap() + .assert_ok(); + } +} + +#[ignore = "integration test — run via `make test-it`"] +#[tokio::test] +async fn a_request_refused_by_a_token_limit_charges_no_request_limit() { + let app = TestApp::spawn().await; + let (api_key, user_id) = seed_runtime(&app).await; + let requests = + fixtures::create_rate_limit_rule(&app.db, "user", user_id, "ai_gateway", "requests", 60, 5) + .await + .unwrap(); + fixtures::create_rate_limit_rule(&app.db, "user", user_id, "ai_gateway", "tokens", 300, 1) + .await + .unwrap(); + let con = admin_session(&app).await; + + let gw = app.gateway_client(); + gw.set_bearer(&api_key); + gw.post("/v1/chat/completions", chat()) + .await + .unwrap() + .assert_ok(); + for _ in 0..3 { + let r = gw.post("/v1/chat/completions", chat()).await.unwrap(); + assert_eq!(r.status.as_u16(), 429, "body={}", r.text()); + } + + let r: Json = con + .get(&format!("/api/admin/limits/user/{user_id}/usage")) + .await + .unwrap() + .json() + .unwrap(); + let requests_used = r["rules"] + .as_array() + .unwrap() + .iter() + .find(|x| x["rule_id"] == json!(requests)) + .map(|x| x["current"].as_i64().unwrap()); + assert_eq!(requests_used, Some(1), "usage: {r}"); +} + +#[ignore = "integration test — run via `make test-it`"] +#[tokio::test] +async fn a_rate_limit_says_when_its_window_frees() { + let app = TestApp::spawn().await; + let (api_key, user_id) = seed_runtime(&app).await; + fixtures::create_rate_limit_rule(&app.db, "user", user_id, "ai_gateway", "requests", 300, 1) + .await + .unwrap(); + + let gw = app.gateway_client(); + gw.set_bearer(&api_key); + gw.post("/v1/chat/completions", chat()) + .await + .unwrap() + .assert_ok(); + let r = gw.post("/v1/chat/completions", chat()).await.unwrap(); + assert_eq!(r.status.as_u16(), 429); + // The one request leaves the five-minute window in under five + // minutes, and not before the window has nearly run its course. + let secs = retry_after(&r); + assert!((290..=300).contains(&secs), "Retry-After {secs}"); + assert!( + r.headers.get("x-should-retry").is_none(), + "a window frees by itself; SDK retries are welcome" + ); +} + +#[ignore = "integration test — run via `make test-it`"] +#[tokio::test] +async fn a_spent_budget_says_when_its_period_ends_and_not_to_retry() { + use chrono::{Datelike, TimeZone, Utc}; + use fred::interfaces::KeysInterface; + use think_watch_common::limits::budget; + + let app = TestApp::spawn().await; + let (api_key, user_id) = seed_runtime(&app).await; + fixtures::create_budget_cap(&app.db, "user", user_id, "monthly", 100) + .await + .unwrap(); + let now = Utc::now(); + let spent = budget::build_key("user", user_id, "monthly", now); + let _: () = app + .state + .redis + .set(&spent, 100, None, None, false) + .await + .unwrap(); + let (y, m) = if now.month() == 12 { + (now.year() + 1, 1) + } else { + (now.year(), now.month() + 1) + }; + let next_month = Utc.with_ymd_and_hms(y, m, 1, 0, 0, 0).unwrap(); + let expected = (next_month - now).num_seconds(); + + let gw = app.gateway_client(); + gw.set_bearer(&api_key); + // The Anthropic surface, to see the refusal in the caller's format. + let r = gw + .post( + "/v1/messages", + json!({"model": "gpt-test", "max_tokens": 16, + "messages": [{"role": "user", "content": "x"}]}), + ) + .await + .unwrap(); + assert_eq!(r.status.as_u16(), 429, "body={}", r.text()); + let secs = retry_after(&r) as i64; + assert!( + (expected - 5..=expected + 5).contains(&secs), + "Retry-After {secs}, the month ends in {expected}" + ); + assert_eq!( + r.headers + .get("x-should-retry") + .and_then(|v| v.to_str().ok()), + Some("false") + ); + let body: Json = r.json().unwrap(); + assert_eq!(body["type"], "error", "{body}"); + assert_eq!(body["error"]["type"], "rate_limit_error", "{body}"); +} + +// ---------------------------------------------------------------------------- +// The scripts themselves, on a clock the test sets. +// ---------------------------------------------------------------------------- + +#[ignore = "integration test — run via `make test-it`"] +#[tokio::test] +async fn the_limit_scripts_fill_free_and_record() { + let app = TestApp::spawn().await; + think_watch_test_support::redis_scripts::exercise_the_limit_scripts(&app.state.redis).await; +} + +// ---------------------------------------------------------------------------- +// MCP: a key's limits count on the key there too. +// ---------------------------------------------------------------------------- + +#[ignore = "integration test — run via `make test-it`"] +#[tokio::test] +async fn a_keys_mcp_limits_count_on_the_key_not_on_its_owner() { + use axum::{Json as AxumJson, Router, routing::post}; + + let app = TestApp::spawn().await; + let user = fixtures::create_random_user(&app.db).await.unwrap(); + + let upstream = Router::new().route( + "/mcp", + post(|AxumJson(req): AxumJson| async move { + AxumJson(json!({"jsonrpc": "2.0", "id": req["id"], "result": {"content": []}})) + }), + ); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + tokio::spawn(async move { + let _ = axum::serve(listener, upstream).await; + }); + let server_id = fixtures::create_mcp_server_with( + &app.db, + &unique_name("limits-mcp"), + "lim", + &format!("http://{addr}/mcp"), + fixtures::McpServerOpts::default(), + ) + .await + .unwrap(); + let row = sqlx::query_as::<_, think_watch_common::models::McpServer>( + "SELECT * FROM mcp_servers WHERE id = $1", + ) + .bind(server_id) + .fetch_one(&app.db) + .await + .unwrap(); + let registered = think_watch_server::mcp_runtime::build_registered_server( + &app.db, + &row, + &app.state.config.encryption_key, + ) + .await + .unwrap(); + app.state.mcp_registry.register(registered).await; + + let mut keys = Vec::new(); + for _ in 0..2 { + let key = fixtures::create_api_key( + &app.db, + user.user.id, + &unique_name("mcp-key"), + &["mcp_gateway"], + None, + None, + ) + .await + .unwrap(); + fixtures::create_rate_limit_rule( + &app.db, + "api_key_lineage", + key.row.lineage_id, + "mcp_gateway", + "requests", + 60, + 1, + ) + .await + .unwrap(); + keys.push(key); + } + let con = admin_session(&app).await; + + let gw = app.gateway_client(); + let call = |id: i64| { + json!({"jsonrpc": "2.0", "id": id, "method": "tools/call", + "params": {"name": "lim__anything", "arguments": {}}}) + }; + let mut answers = Vec::new(); + for (id, key) in [(1, &keys[0]), (2, &keys[1]), (3, &keys[0])] { + gw.set_bearer(&key.plaintext); + let r = gw.post("/mcp", call(id)).await.unwrap(); + r.assert_ok(); + answers.push(r.json::().unwrap()); + } + assert!(answers[0]["result"].is_object(), "{}", answers[0]); + assert!( + answers[1]["result"].is_object(), + "the second key has its own counter: {}", + answers[1] + ); + assert_eq!( + answers[2]["error"]["message"], "Rate limited: api_key_lineage:requests/1m", + "{}", + answers[2] + ); + assert_eq!(usage(&con, "api_key", keys[0].row.id).await.0, vec![1]); +} diff --git a/crates/test-support/tests/redis_cluster.rs b/crates/test-support/tests/redis_cluster.rs new file mode 100644 index 00000000..f4ab0407 --- /dev/null +++ b/crates/test-support/tests/redis_cluster.rs @@ -0,0 +1,282 @@ +//! The gateway's Lua scripts and multi-key commands against a real +//! Redis Cluster. +//! +//! A cluster runs a script only when every key it declares hashes to one +//! slot, refuses a `DEL` across slots (`CROSSSLOT`), and answers `SCAN` +//! for one node only. These tests point at a cluster given by +//! `TEST_REDIS_CLUSTER_URL` (e.g. `redis-cluster://127.0.0.1:37001`) and +//! skip without it, so CI — which has a single Redis — stays green. To +//! run them locally, start a three-primary cluster whose nodes announce +//! an address this process can reach, for example: +//! +//! ```text +//! docker run -d --rm --name tw-redis-cluster \ +//! -p 37001:37001 -p 37002:37002 -p 37003:37003 redis:8-alpine sh -c ' +//! for p in 37001 37002 37003; do +//! redis-server --port $p --cluster-enabled yes --cluster-config-file n-$p.conf \ +//! --cluster-announce-ip 127.0.0.1 --protected-mode no --save "" --daemonize yes --dir /tmp +//! done; sleep 1 +//! redis-cli --cluster create 127.0.0.1:37001 127.0.0.1:37002 127.0.0.1:37003 \ +//! --cluster-replicas 0 --cluster-yes; tail -f /dev/null' +//! TEST_REDIS_CLUSTER_URL=redis-cluster://127.0.0.1:37001 \ +//! cargo nextest run -p think-watch-test-support --test redis_cluster --run-ignored only +//! ``` +//! +//! Keys carry fresh UUIDs, so nothing is flushed and runs don't collide. + +use fred::clients::Client; +use fred::interfaces::{ClientLike, KeysInterface}; +use fred::types::Builder; +use fred::types::config::Config; +use think_watch_test_support::prelude::*; + +fn cluster_url() -> Option { + let url = std::env::var("TEST_REDIS_CLUSTER_URL").ok(); + if url.is_none() { + eprintln!("TEST_REDIS_CLUSTER_URL not set — skipping the Redis Cluster test"); + } + url +} + +async fn cluster(url: &str) -> Client { + let client = Builder::from_config(Config::from_url(url).unwrap()) + .build() + .unwrap(); + client.init().await.unwrap(); + assert!(client.is_clustered(), "{url} is not a cluster URL"); + client +} + +#[ignore = "integration test — needs TEST_REDIS_CLUSTER_URL"] +#[tokio::test] +async fn the_limit_scripts_run_on_a_cluster() { + let Some(url) = cluster_url() else { return }; + let redis = cluster(&url).await; + think_watch_test_support::redis_scripts::exercise_the_limit_scripts(&redis).await; +} + +#[ignore = "integration test — needs TEST_REDIS_CLUSTER_URL"] +#[tokio::test] +async fn budgets_route_health_and_quotas_run_on_a_cluster() { + use think_watch_common::limits::{BudgetCap, BudgetPeriod, BudgetSubject, budget}; + use think_watch_gateway::health::{CircuitBreakerConfig, HealthTracker}; + use think_watch_gateway::quota::QuotaManager; + + let Some(url) = cluster_url() else { return }; + let redis = cluster(&url).await; + + // Budgets: one key per command, any slot. + let cap = |kind, id| BudgetCap { + id: Uuid::nil(), + subject_kind: kind, + subject_id: id, + period: BudgetPeriod::Daily, + limit_tokens: 100, + enabled: true, + expires_at: None, + reason: None, + created_by: None, + }; + let caps = [ + cap(BudgetSubject::User, Uuid::new_v4()), + cap(BudgetSubject::ApiKeyLineage, Uuid::new_v4()), + ]; + budget::add_weighted_tokens(&redis, &caps, 60) + .await + .unwrap(); + let (statuses, crossings) = budget::add_weighted_tokens(&redis, &caps, 60) + .await + .unwrap(); + assert_eq!( + statuses.iter().map(|s| s.current).collect::>(), + [120, 120] + ); + assert_eq!(crossings.len(), 6, "80/95/100 % on both caps"); + let spent = budget::current_spend(&redis, &caps).await.unwrap(); + assert_eq!( + spent.iter().map(|s| s.current).collect::>(), + [120, 120] + ); + + // Route health: three keys in one script, then two, then a DEL of + // all three. + let health = HealthTracker::new(redis.clone()); + let route = Uuid::new_v4(); + let cfg = CircuitBreakerConfig { + enabled: true, + error_pct: 50, + min_samples: 2, + window_secs: 60, + open_secs: 30, + }; + health.record(route, 10, true, cfg).await; + let after = health.record(route, 10, true, cfg).await; + assert_eq!(after.total, 2, "the record script ran"); + assert_eq!(after.lifetime_requests, 2); + assert_eq!(format!("{:?}", after.state), "Open", "the state write ran"); + assert_eq!(format!("{:?}", health.state(route, cfg).await), "Open"); + health.forget(route).await; + assert_eq!(health.snapshot(route, cfg).await.lifetime_requests, 0); + + // Quotas: limit and usage in one script. + let quota = QuotaManager::new(redis.clone()); + let key = format!("{}:gpt-test", Uuid::new_v4()); + quota.set_limit(&key, 100).await.unwrap(); + assert_eq!(quota.consume(&key, 30).await.unwrap(), 70); + assert_eq!(quota.check_and_consume(&key, 20).await.unwrap(), 50); + assert!(quota.check_and_consume(&key, 60).await.is_err()); + assert_eq!(quota.get_usage(&key).await.unwrap().used, 50); +} + +#[ignore = "integration test — needs TEST_REDIS_CLUSTER_URL"] +#[tokio::test] +async fn a_pattern_delete_reaches_every_node_of_a_cluster() { + let Some(url) = cluster_url() else { return }; + let redis = cluster(&url).await; + + let run = Uuid::new_v4(); + let keys: Vec = (0..300).map(|i| format!("tw-test:{run}:{i}")).collect(); + for k in &keys { + let _: () = redis.set(k, 1, None, None, false).await.unwrap(); + } + let slots: std::collections::HashSet = keys + .iter() + .map(|k| fred::util::redis_keyslot(k.as_bytes())) + .collect(); + assert!(slots.len() > 100, "the keys spread over the cluster"); + + let deleted = + think_watch_common::redis_keys::delete_matching(&redis, &format!("tw-test:{run}:*"), None) + .await + .unwrap(); + assert_eq!(deleted, keys.len()); + for k in &keys { + assert_eq!(redis.exists::(k).await.unwrap(), 0, "{k}"); + } +} + +#[ignore = "integration test — needs TEST_REDIS_CLUSTER_URL"] +#[tokio::test] +async fn the_gateway_enforces_limits_on_a_cluster() { + let Some(url) = cluster_url() else { return }; + let app = TestApp::try_spawn_with(SpawnOptions { + redis_url: Some(url), + ..Default::default() + }) + .await + .unwrap(); + + let user = fixtures::create_random_user(&app.db).await.unwrap(); + let user_id = user.user.id; + let mock = MockProvider::openai_chat_ok("gpt-test").await; + let provider = fixtures::create_provider( + &app.db, + &unique_name("cluster-prov"), + "openai", + &mock.uri(), + None, + ) + .await + .unwrap(); + fixtures::create_model_and_route(&app.db, provider.id, "gpt-test") + .await + .unwrap(); + app.rebuild_gateway_router().await; + let key = fixtures::create_api_key( + &app.db, + user_id, + &unique_name("cluster-key"), + &["ai_gateway"], + None, + None, + ) + .await + .unwrap(); + // The user's request limit and the key's token limit and budget: one + // script checks counters of both subjects. + fixtures::create_rate_limit_rule(&app.db, "user", user_id, "ai_gateway", "requests", 60, 2) + .await + .unwrap(); + fixtures::create_rate_limit_rule( + &app.db, + "api_key_lineage", + key.row.lineage_id, + "ai_gateway", + "tokens", + 60, + 1_000_000, + ) + .await + .unwrap(); + fixtures::create_budget_cap( + &app.db, + "api_key_lineage", + key.row.lineage_id, + "daily", + 1_000_000, + ) + .await + .unwrap(); + let con = admin_session(&app).await; + + let gw = app.gateway_client(); + gw.set_bearer(&key.plaintext); + // A prompt of its own: the response cache lives in the same Redis, + // which nothing flushes, and a cached answer records no tokens. + let chat = json!({"model": "gpt-test", + "messages": [{"role": "user", "content": Uuid::new_v4().to_string()}]}); + for _ in 0..2 { + gw.post("/v1/chat/completions", chat.clone()) + .await + .unwrap() + .assert_ok(); + } + let r = gw.post("/v1/chat/completions", chat).await.unwrap(); + assert_eq!(r.status.as_u16(), 429, "body={}", r.text()); + assert!(r.headers.get("retry-after").is_some()); + + let usage: Json = con + .get(&format!("/api/admin/limits/api_key/{}/usage", key.row.id)) + .await + .unwrap() + .json() + .unwrap(); + assert!( + usage["rules"][0]["current"].as_i64().unwrap() > 0, + "the key's tokens were recorded: {usage}" + ); + assert!( + usage["caps"][0]["current"].as_i64().unwrap() > 0, + "the key's budget was debited: {usage}" + ); + let usage: Json = con + .get(&format!("/api/admin/limits/user/{user_id}/usage")) + .await + .unwrap() + .json() + .unwrap(); + assert_eq!(usage["rules"][0]["current"], 2, "{usage}"); +} + +#[ignore = "integration test — needs TEST_REDIS_CLUSTER_URL"] +#[tokio::test] +async fn config_change_notices_reach_a_subscriber_on_a_cluster() { + // What `init::spawn_config_subscriber` does with the same URL: a + // subscriber on one node hears a publish sent through another. + use fred::interfaces::{EventInterface, PubsubInterface}; + let Some(url) = cluster_url() else { return }; + let publisher = cluster(&url).await; + let subscriber = Builder::from_config(Config::from_url(&url).unwrap()) + .build_subscriber_client() + .unwrap(); + subscriber.init().await.unwrap(); + let mut rx = subscriber.message_rx(); + subscriber.subscribe("config:changed").await.unwrap(); + + think_watch_common::dynamic_config::notify_config_changed(&publisher).await; + let msg = tokio::time::timeout(std::time::Duration::from_secs(5), rx.recv()) + .await + .expect("a notice within 5 s") + .unwrap(); + assert_eq!(msg.channel, "config:changed"); +} diff --git a/db/schema.sql b/db/schema.sql index 3b21f660..0245ea58 100644 --- a/db/schema.sql +++ b/db/schema.sql @@ -90,7 +90,10 @@ CREATE TABLE IF NOT EXISTS teams ( id UUID PRIMARY KEY DEFAULT gen_random_uuid(), name VARCHAR(255) NOT NULL UNIQUE, description TEXT, - -- Budget caps live in `budget_caps` (subject_kind = 'team'). + -- Teams carry no limits or budgets of their own: `budget_caps` and + -- `rate_limit_rules` attach to users and API keys only (see the + -- note above `rate_limit_rules`). A team's members are limited + -- through the roles the team grants. created_at TIMESTAMPTZ NOT NULL DEFAULT now() ); diff --git a/deploy/helm/think-watch/README.md b/deploy/helm/think-watch/README.md index bd84dd75..06478b71 100644 --- a/deploy/helm/think-watch/README.md +++ b/deploy/helm/think-watch/README.md @@ -72,6 +72,31 @@ helm upgrade --install thinkwatch deploy/helm/think-watch \ When `bundled=false` and `externalUrl` is empty the chart fails at install-time with an explicit message — no silent broken Secret. +### Redis Cluster + +The external Redis can be a Redis Cluster. Give its URL the +`redis-cluster://` scheme and name one node or more; the server finds the +rest of the cluster from them: + +```yaml +redis: + bundled: false + externalUrl: redis-cluster://:pass@redis-0.redis:6379?node=redis-1.redis:6379&node=redis-2.redis:6379 +``` + +- Every node must be reachable from the server pods at the address it + announces to the cluster (`cluster-announce-ip` / `-port`): the server + follows the cluster's redirects to it. +- A cluster has only database 0, so the URL names no `/`. +- Nothing else is needed. Every script the server runs keeps its keys in + one hash slot — the counters of one user's request share the tag + `{user:}`, a route's health keys `{}` — and pattern + deletes (cache invalidation) scan every primary. + +Rate limits, budgets, route health, caches and the config change +notices between instances are tested against a three-primary Redis 8 +cluster (`crates/test-support/tests/redis_cluster.rs`). + ## Rotating secrets `-secrets` is kept on `helm uninstall`. To rotate passwords: