From 89dcec9c1eac9887716a42f216b0284cff4351cc Mon Sep 17 00:00:00 2001 From: Vishu Bhatnagar Date: Mon, 5 Oct 2026 22:12:05 +0100 Subject: [PATCH 1/3] feat(rate-limiter): add per-server user limits Signed-off-by: Vishu Bhatnagar --- Cargo.lock | 2 +- .../python-package/rate_limiter/Cargo.toml | 2 +- .../python-package/rate_limiter/README.md | 24 +- .../cpex_rate_limiter/plugin-manifest.yaml | 5 +- .../cpex_rate_limiter/rate_limiter.py | 1 + .../rate_limiter_rust/__init__.pyi | 4 +- .../python-package/rate_limiter/src/config.rs | 33 ++- .../python-package/rate_limiter/src/engine.rs | 215 +++++++++++++++++- .../python-package/rate_limiter/src/lib.rs | 1 + .../python-package/rate_limiter/src/plugin.rs | 137 +++++++++-- .../rate_limiter/test_config_integration.py | 7 + .../tests/rate_limiter/test_integration.py | 77 ++++++- .../rate_limiter/test_redis_integration.py | 61 ++++- 13 files changed, 515 insertions(+), 54 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 21400694..0ef7e5b6 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1545,7 +1545,7 @@ checksum = "63b8176103e19a2643978565ca18b50549f6101881c443590420e4dc998a3c69" [[package]] name = "rate_limiter" -version = "0.1.10" +version = "0.1.11" dependencies = [ "cpex_framework_bridge", "criterion", diff --git a/plugins/rust/python-package/rate_limiter/Cargo.toml b/plugins/rust/python-package/rate_limiter/Cargo.toml index c03203e1..f4020811 100644 --- a/plugins/rust/python-package/rate_limiter/Cargo.toml +++ b/plugins/rust/python-package/rate_limiter/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "rate_limiter" -version = "0.1.10" +version = "0.1.11" edition.workspace = true authors.workspace = true license.workspace = true diff --git a/plugins/rust/python-package/rate_limiter/README.md b/plugins/rust/python-package/rate_limiter/README.md index c667f382..90b50897 100644 --- a/plugins/rust/python-package/rate_limiter/README.md +++ b/plugins/rust/python-package/rate_limiter/README.md @@ -2,7 +2,7 @@ > Author: ContextForge Contributors -Enforces rate limits per user, tenant, and tool across `tool_pre_invoke` and `prompt_pre_fetch` hooks. Supports pluggable counting algorithms (fixed window, sliding window, token bucket), an in-process memory backend (single-instance), and a Redis backend (shared across all gateway instances). +Enforces global per-user, per-user/per-server, tenant, and tool rate limits across `tool_pre_invoke` and `prompt_pre_fetch` hooks. Supports pluggable counting algorithms (fixed window, sliding window, token bucket), an in-process memory backend (single-instance), and a Redis backend (shared across all gateway instances). ## Runtime Requirements @@ -12,8 +12,8 @@ This plugin depends on `cpex>=0.1.0,<0.2` and imports hook models from `cpex.fra | Hook | When it runs | |---|---| -| `tool_pre_invoke` | Before every tool call — checks `by_user`, `by_tenant`, `by_tool` | -| `prompt_pre_fetch` | Before every prompt fetch — checks `by_user`, `by_tenant`, `by_tool` | +| `tool_pre_invoke` | Before every tool call — checks `by_user`, `by_user_per_server`, `by_tenant`, `by_tool` | +| `prompt_pre_fetch` | Before every prompt fetch — checks `by_user`, `by_user_per_server`, `by_tenant`, `by_tool` | If any configured dimension is exceeded, the plugin returns a violation with HTTP 429. All requests include `X-RateLimit-*` headers. The most restrictive active dimension is surfaced (e.g. if both user and tenant limits are active, the one closest to exhaustion is reported). @@ -27,7 +27,8 @@ If any configured dimension is exceeded, the plugin returns a violation with HTT - tool_pre_invoke mode: enforce # enforce | permissive | disabled config: - by_user: "30/m" # per-user limit across all tools + by_user: "300/m" # per-user limit across all servers + by_user_per_server: "60/m" # per-user limit for each MCP server by_tenant: "300/m" # shared limit across all users in a tenant by_tool: # per-tool overrides (applied on top of by_user) search: "10/m" @@ -55,6 +56,7 @@ If any configured dimension is exceeded, the plugin returns a violation with HTT | Field | Type | Default | Description | |---|---|---|---| | `by_user` | string | `null` | Per-user rate limit, e.g. `"60/m"` | +| `by_user_per_server` | string | `null` | Per-user rate limit independently enforced for each MCP server, e.g. `"60/m"` | | `by_tenant` | string | `null` | Per-tenant rate limit, e.g. `"600/m"` | | `by_tool` | dict | `{}` | Per-tool overrides, e.g. `{"search": "10/m"}` | | `algorithm` | string | `"fixed_window"` | Counting algorithm: `"fixed_window"`, `"sliding_window"`, or `"token_bucket"` | @@ -75,6 +77,8 @@ If any configured dimension is exceeded, the plugin returns a violation with HTT **Omitting a dimension** (e.g. no `by_tenant`) means that dimension is unlimited — no counter is tracked for it. +`by_user` and `by_user_per_server` are independent and may be enabled together. For example, `by_user: "300/m"` plus `by_user_per_server: "60/m"` gives each user at most 60 requests per server and 300 requests across all servers. This prevents one user from multiplying their effective quota without relying on the shared `by_tenant` ceiling. When `by_user_per_server` is configured, requests without a safe non-empty `server_id` are blocked with `RATE_LIMIT_CONTEXT_MISSING` (HTTP 503) rather than bypassing the limit. + ## Response headers Every request (allowed or blocked) includes: @@ -106,7 +110,7 @@ Stores a timestamp for every request in the current window. At each check, expir ### Token bucket -Each identity (user, tenant, tool) has a bucket that holds up to `count` tokens. Tokens refill at a steady rate of `count/window`. A request consumes one token. Bursts up to the bucket capacity are allowed; sustained rate above `count/window` is rejected. Useful for APIs where short spikes are acceptable but sustained overload is not. +Each identity (user, user/server, tenant, tool) has a bucket that holds up to `count` tokens. Tokens refill at a steady rate of `count/window`. A request consumes one token. Bursts up to the bucket capacity are allowed; sustained rate above `count/window` is rejected. Useful for APIs where short spikes are acceptable but sustained overload is not. **Redis support:** `token_bucket` with `backend: redis` is fully supported. The plugin stores `{tokens, last_refill}` in a Redis hash per key and uses an atomic Lua script to refill and consume tokens in a single round-trip — the same pattern as the other two algorithms. This means `token_bucket` enforces a true cluster-wide limit in multi-instance deployments. @@ -178,11 +182,12 @@ When the plugin context carries a `tenant_id`, every dimension key is prefixed w ``` rl:{tenant_id}:user:{email}:{window_seconds} +rl:{tenant_id}:user:{email}:server:{server_id}:{window_seconds} rl:{tenant_id}:tenant:{tenant_id}:{window_seconds} rl:{tenant_id}:tool:{tool_name}:{window_seconds} ``` -When `tenant_id` is absent (single-tenant deployments), the prefix is omitted and keys revert to the pre-tenant-scoping layout (`rl:user:{email}:{window}`), so single-tenant behaviour is unchanged. +When `tenant_id` is absent (single-tenant deployments), the tenant prefix is omitted. Global user keys retain the existing layout (`rl:user:{email}:{window}`); per-server keys use `rl:user:{email}:server:{server_id}:{window}`. **Upgrade note:** the first deploy of the tenant-scoping change causes counters under `rl:user:*` / `rl:tool:*` to be orphaned while new writes land at `rl:{tenant}:user:*`. Counters effectively reset once for all in-flight windows — non-event for typical second/minute windows. @@ -192,7 +197,8 @@ When `tenant_id` is absent (single-tenant deployments), the prefix is omitted an ```yaml config: - by_user: "60/m" + by_user: "300/m" + by_user_per_server: "60/m" by_tenant: "600/m" ``` @@ -267,6 +273,8 @@ result.metadata["rate_limiter"] = { ## Migration Note +Version `0.1.11` adds the optional `by_user_per_server` dimension. Existing configurations and global `by_user` Redis keys remain unchanged. Enabling the new field creates fresh user/server counters; deploy `0.1.11` to every replica before enabling it so all replicas enforce the same dimensions. Use the stable catalog UUID carried in `GlobalContext.server_id`; display names or slugs can reset quotas when renamed. + Version `0.1.7` is a **breaking change** for any existing consumer reading rate-limit metadata: - The old flat, unconditional `result.metadata` write (the engine's `meta` dict — `limited`, `remaining`, `reset_in`, `dimensions` — written on every allowed/not-limited call regardless of trace context) is now gated on a valid `trace_id` and replaced by the namespaced `rate_limiter` key containing only `allowed`/`throttled`/`backend` (the engine's own fields are not folded in — see "Returned Metadata" above for why). @@ -291,5 +299,5 @@ Without `shutdown`, the cached Redis connection would leak across plugin re-inst | Fixed window allows up to 2× limit at window boundary | LOW | Use `sliding_window` algorithm, or use `by_user` with headroom | | `by_tool` matching is case-sensitive | LOW | Fixed — tool names are normalised with `.strip().lower()` | | Whitespace-only user identity bypasses anonymous bucket | LOW | Fixed — `_extract_user_identity` strips whitespace and falls back to `'anonymous'` | -| No per-server limits (`server_id` dimension missing) | LOW | Not implemented | +| Per-server user limits require a stable server ID | MEDIUM | Use the catalog UUID propagated as `GlobalContext.server_id` | | No config hot-reload — rate string changes require restart | LOW | Not implemented | diff --git a/plugins/rust/python-package/rate_limiter/cpex_rate_limiter/plugin-manifest.yaml b/plugins/rust/python-package/rate_limiter/cpex_rate_limiter/plugin-manifest.yaml index f9e4fcfa..c7209de4 100644 --- a/plugins/rust/python-package/rate_limiter/cpex_rate_limiter/plugin-manifest.yaml +++ b/plugins/rust/python-package/rate_limiter/cpex_rate_limiter/plugin-manifest.yaml @@ -1,11 +1,12 @@ -description: "Rate limiting by user/tenant/tool — memory (single-process) or Redis (shared across instances)" +description: "Rate limiting by user, user/server, tenant, and tool — memory or Redis backends" author: "ContextForge Contributors" -version: "0.1.10" +version: "0.1.11" kind: "cpex_rate_limiter.rate_limiter.RateLimiterPlugin" available_hooks: - "prompt_pre_fetch" - "tool_pre_invoke" default_configs: by_user: "60/m" + by_user_per_server: null by_tenant: "600/m" by_tool: {} diff --git a/plugins/rust/python-package/rate_limiter/cpex_rate_limiter/rate_limiter.py b/plugins/rust/python-package/rate_limiter/cpex_rate_limiter/rate_limiter.py index 0a209ddc..d0e2905b 100644 --- a/plugins/rust/python-package/rate_limiter/cpex_rate_limiter/rate_limiter.py +++ b/plugins/rust/python-package/rate_limiter/cpex_rate_limiter/rate_limiter.py @@ -21,6 +21,7 @@ def _parse_rate(rate: str) -> tuple[int, int]: class RateLimiterConfig: __slots__ = ( "by_user", + "by_user_per_server", "by_tenant", "by_tool", "algorithm", diff --git a/plugins/rust/python-package/rate_limiter/cpex_rate_limiter/rate_limiter_rust/__init__.pyi b/plugins/rust/python-package/rate_limiter/cpex_rate_limiter/rate_limiter_rust/__init__.pyi index ac8aa0a6..e2f15577 100644 --- a/plugins/rust/python-package/rate_limiter/cpex_rate_limiter/rate_limiter_rust/__init__.pyi +++ b/plugins/rust/python-package/rate_limiter/cpex_rate_limiter/rate_limiter_rust/__init__.pyi @@ -114,7 +114,7 @@ class RateLimiterEngine: `redis_ur` instead of `redis_url`) surface visibly instead of being silently ignored. """ - def check(self, user: builtins.str, tenant: typing.Optional[builtins.str], tool: builtins.str, now_unix: builtins.int, include_retry_after: builtins.bool, context_prefix: typing.Optional[builtins.str]) -> tuple[builtins.bool, dict, dict]: + def check(self, user: builtins.str, tenant: typing.Optional[builtins.str], tool: builtins.str, now_unix: builtins.int, include_retry_after: builtins.bool, context_prefix: typing.Optional[builtins.str], server_id: typing.Optional[builtins.str] = None) -> tuple[builtins.bool, dict, dict]: r""" High-level check: builds dimension keys internally, evaluates, and returns pre-built Python dicts for headers and metadata. @@ -135,7 +135,7 @@ class RateLimiterEngine: path is intended for the memory backend. The `debug_assert` below guards against accidental misuse. """ - def check_async(self, user: builtins.str, tenant: typing.Optional[builtins.str], tool: builtins.str, now_unix: builtins.int, include_retry_after: builtins.bool, context_prefix: typing.Optional[builtins.str]) -> typing.Any: + def check_async(self, user: builtins.str, tenant: typing.Optional[builtins.str], tool: builtins.str, now_unix: builtins.int, include_retry_after: builtins.bool, context_prefix: typing.Optional[builtins.str], server_id: typing.Optional[builtins.str] = None) -> typing.Any: r""" Async variant of `check()` for Redis-backed deployments. diff --git a/plugins/rust/python-package/rate_limiter/src/config.rs b/plugins/rust/python-package/rate_limiter/src/config.rs index 6b990b6f..ecf3920f 100644 --- a/plugins/rust/python-package/rate_limiter/src/config.rs +++ b/plugins/rust/python-package/rate_limiter/src/config.rs @@ -108,6 +108,7 @@ impl Algorithm { #[derive(Debug, Clone)] pub struct EngineConfig { pub by_user: Option, + pub by_user_per_server: Option, pub by_tenant: Option, /// Normalised key → limit. Keys are already `.trim().to_lowercase()`. pub by_tool: HashMap, @@ -119,6 +120,7 @@ impl EngineConfig { /// that are relevant to the Rust engine — strict subset per IFACE-04). pub fn new( by_user: Option<&str>, + by_user_per_server: Option<&str>, by_tenant: Option<&str>, by_tool: HashMap, algorithm: &str, @@ -131,6 +133,14 @@ impl EngineConfig { }) }) .transpose()?; + let by_user_per_server = by_user_per_server + .map(|rate| { + parse_rate(rate).map_err(|err| ConfigError::FieldError { + field: format!("by_user_per_server={rate:?}"), + message: err.to_string(), + }) + }) + .transpose()?; let by_tenant = by_tenant .map(|rate| { parse_rate(rate).map_err(|err| ConfigError::FieldError { @@ -155,6 +165,7 @@ impl EngineConfig { .ok_or_else(|| ConfigError::InvalidAlgorithm(algorithm.to_string()))?; Ok(Self { by_user, + by_user_per_server, by_tenant, by_tool, algorithm, @@ -256,9 +267,17 @@ mod tests { by_tool.insert("Search".to_string(), "10/m".to_string()); by_tool.insert(" Summarise ".to_string(), "5/m".to_string()); - let cfg = EngineConfig::new(Some("30/m"), Some("300/m"), by_tool, "fixed_window").unwrap(); + let cfg = EngineConfig::new( + Some("30/m"), + Some("15/m"), + Some("300/m"), + by_tool, + "fixed_window", + ) + .unwrap(); assert_eq!(cfg.by_user.unwrap().count, 30); + assert_eq!(cfg.by_user_per_server.unwrap().count, 15); assert_eq!(cfg.by_tenant.unwrap().count, 300); // Keys must be normalised assert!(cfg.by_tool.contains_key("search")); @@ -269,19 +288,25 @@ mod tests { #[test] fn engine_config_all_none_is_valid() { - let cfg = EngineConfig::new(None, None, HashMap::new(), "sliding_window").unwrap(); + let cfg = EngineConfig::new(None, None, None, HashMap::new(), "sliding_window").unwrap(); assert!(cfg.by_user.is_none()); + assert!(cfg.by_user_per_server.is_none()); assert!(cfg.by_tenant.is_none()); assert!(cfg.by_tool.is_empty()); } #[test] fn engine_config_invalid_rate_propagates_error() { - assert!(EngineConfig::new(Some("bad"), None, HashMap::new(), "fixed_window").is_err()); + assert!( + EngineConfig::new(Some("bad"), None, None, HashMap::new(), "fixed_window").is_err() + ); + assert!( + EngineConfig::new(None, Some("bad"), None, HashMap::new(), "fixed_window").is_err() + ); } #[test] fn engine_config_invalid_algorithm_propagates_error() { - assert!(EngineConfig::new(None, None, HashMap::new(), "leaky_bucket").is_err()); + assert!(EngineConfig::new(None, None, None, HashMap::new(), "leaky_bucket").is_err()); } } diff --git a/plugins/rust/python-package/rate_limiter/src/engine.rs b/plugins/rust/python-package/rate_limiter/src/engine.rs index b63ae460..40b50151 100644 --- a/plugins/rust/python-package/rate_limiter/src/engine.rs +++ b/plugins/rust/python-package/rate_limiter/src/engine.rs @@ -79,6 +79,35 @@ impl RateLimiterEngine { matches!(self.backend, EngineBackend::Redis(_)) } + /// Validate and normalize the server identifier only when a per-server + /// user limit is configured. Server IDs become part of backend keys, so + /// keep them bounded and delimiter-safe. + pub(crate) fn validated_server_id<'a>( + &self, + server_id: Option<&'a str>, + ) -> Result, &'static str> { + if self.config.by_user_per_server.is_none() { + return Ok(None); + } + + let server_id = server_id + .map(str::trim) + .filter(|value| !value.is_empty()) + .ok_or("server_id is required when by_user_per_server is configured")?; + if server_id.len() > 128 { + return Err("server_id exceeds the maximum length of 128 bytes"); + } + if !server_id + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_' | b'.')) + { + return Err( + "server_id may contain only ASCII letters, digits, hyphen, underscore, and period", + ); + } + Ok(Some(server_id)) + } + /// Release backend-held resources. For Redis, this drops the cached /// multiplexed connection so the server can close the socket; in-flight /// requests that already cloned the handle remain valid. Memory backend @@ -118,6 +147,9 @@ impl RateLimiterEngine { let _ = warn_on_unknown_config_keys(config); let by_user: Option = config.get_item("by_user")?.and_then(|v| v.extract().ok()); + let by_user_per_server: Option = config + .get_item("by_user_per_server")? + .and_then(|v| v.extract().ok()); let by_tenant: Option = config.get_item("by_tenant")?.and_then(|v| v.extract().ok()); let algorithm: String = config @@ -132,6 +164,7 @@ impl RateLimiterEngine { let engine_config = EngineConfig::new( by_user.as_deref(), + by_user_per_server.as_deref(), by_tenant.as_deref(), by_tool, &algorithm, @@ -211,6 +244,7 @@ impl RateLimiterEngine { /// The Python wrapper routes Redis to `check_async()` instead; this sync /// path is intended for the memory backend. The `debug_assert` below /// guards against accidental misuse. + #[pyo3(signature = (user, tenant, tool, now_unix, include_retry_after, context_prefix, server_id=None))] #[allow(clippy::too_many_arguments)] pub fn check<'py>( &self, @@ -221,13 +255,17 @@ impl RateLimiterEngine { now_unix: i64, include_retry_after: bool, context_prefix: Option<&str>, + server_id: Option<&str>, ) -> PyResult<(bool, Bound<'py, PyDict>, Bound<'py, PyDict>)> { if matches!(self.backend, EngineBackend::Redis(_)) { return Err(pyo3::exceptions::PyRuntimeError::new_err( "check() must not be called with the Redis backend — use check_async() instead", )); } - let checks = self.build_checks(user, tenant, tool, context_prefix); + let server_id = self + .validated_server_id(server_id) + .map_err(pyo3::exceptions::PyValueError::new_err)?; + let checks = self.build_checks(user, tenant, tool, context_prefix, server_id); if checks.is_empty() { let headers = PyDict::new(py); let meta = PyDict::new(py); @@ -270,6 +308,7 @@ impl RateLimiterEngine { /// Async variant of `check()` for Redis-backed deployments. /// /// Returns an awaitable that resolves to `(allowed, headers_dict, meta_dict)`. + #[pyo3(signature = (user, tenant, tool, now_unix, include_retry_after, context_prefix, server_id=None))] #[allow(clippy::too_many_arguments)] pub fn check_async<'py>( &self, @@ -280,8 +319,12 @@ impl RateLimiterEngine { now_unix: i64, include_retry_after: bool, context_prefix: Option<&str>, + server_id: Option<&str>, ) -> PyResult> { - let checks = self.build_checks(user, tenant, tool, context_prefix); + let server_id = self + .validated_server_id(server_id) + .map_err(pyo3::exceptions::PyValueError::new_err)?; + let checks = self.build_checks(user, tenant, tool, context_prefix, server_id); if checks.is_empty() { return future_into_py(py, async move { Python::attach(|py| -> PyResult> { @@ -376,8 +419,9 @@ impl RateLimiterEngine { tenant: Option<&str>, tool: &str, context_prefix: Option<&str>, + server_id: Option<&str>, ) -> Vec<(String, u64, u64)> { - let mut checks = Vec::with_capacity(3); + let mut checks = Vec::with_capacity(4); let pfx = context_prefix.unwrap_or(""); if let Some(ref rl) = self.config.by_user { let key = if pfx.is_empty() { @@ -387,6 +431,14 @@ impl RateLimiterEngine { }; checks.push((key, rl.count, rl.window_nanos)); } + if let (Some(server_id), Some(rl)) = (server_id, &self.config.by_user_per_server) { + let key = if pfx.is_empty() { + format!("user:{}:server:{}", user, server_id) + } else { + format!("{}:user:{}:server:{}", pfx, user, server_id) + }; + checks.push((key, rl.count, rl.window_nanos)); + } if let (Some(t), Some(rl)) = (tenant, &self.config.by_tenant) { let key = if pfx.is_empty() { format!("tenant:{}", t) @@ -417,6 +469,7 @@ impl RateLimiterEngine { fn warn_on_unknown_config_keys(config: &Bound<'_, PyDict>) -> Vec { const KNOWN: &[&str] = &[ "by_user", + "by_user_per_server", "by_tenant", "by_tool", "algorithm", @@ -593,6 +646,7 @@ mod tests { let mut by_tool = HashMap::new(); let cfg = EngineConfig { by_user: by_user.map(|s| crate::config::parse_rate(s).unwrap()), + by_user_per_server: None, by_tenant: None, by_tool: { by_tool.insert( @@ -614,6 +668,7 @@ mod tests { let cfg = EngineConfig::new( Some("10/s"), None, + None, { let mut m = HashMap::new(); m.insert("Search".to_string(), "5/m".to_string()); @@ -702,7 +757,7 @@ mod tests { #[test] fn build_checks_without_prefix_produces_unprefixed_keys() { let (engine, _handle) = engine_with_fake_clock(Some("10/s"), Algorithm::FixedWindow); - let checks = engine.build_checks("alice", Some("acme"), "search", None); + let checks = engine.build_checks("alice", Some("acme"), "search", None, None); let keys: Vec<&str> = checks.iter().map(|(k, _, _)| k.as_str()).collect(); assert!(keys.contains(&"user:alice")); assert!(keys.contains(&"tool:search")); @@ -711,7 +766,7 @@ mod tests { #[test] fn build_checks_with_prefix_prepends_to_all_keys() { let (engine, _handle) = engine_with_fake_clock(Some("10/s"), Algorithm::FixedWindow); - let checks = engine.build_checks("alice", Some("acme"), "search", Some("team_a")); + let checks = engine.build_checks("alice", Some("acme"), "search", Some("team_a"), None); let keys: Vec<&str> = checks.iter().map(|(k, _, _)| k.as_str()).collect(); assert!(keys.contains(&"team_a:user:alice"), "keys: {:?}", keys); assert!(keys.contains(&"team_a:tool:search"), "keys: {:?}", keys); @@ -723,22 +778,161 @@ mod tests { let (clock, _handle) = FakeClock::new(1_000_000); let cfg = EngineConfig { by_user: Some(crate::config::parse_rate("10/s").unwrap()), + by_user_per_server: None, by_tenant: Some(crate::config::parse_rate("100/s").unwrap()), by_tool: HashMap::new(), algorithm: Algorithm::FixedWindow, }; let engine = RateLimiterEngine::new_with_clock(cfg, Arc::new(clock)); - let checks = engine.build_checks("alice", Some("acme"), "search", Some("team_a")); + let checks = engine.build_checks("alice", Some("acme"), "search", Some("team_a"), None); let keys: Vec<&str> = checks.iter().map(|(k, _, _)| k.as_str()).collect(); assert!(keys.contains(&"team_a:user:alice"), "keys: {:?}", keys); assert!(keys.contains(&"team_a:tenant:acme"), "keys: {:?}", keys); } + #[test] + fn build_checks_adds_independent_global_and_server_user_dimensions() { + init_python(); + let (clock, _handle) = FakeClock::new(1_000_000); + let cfg = EngineConfig { + by_user: Some(crate::config::parse_rate("100/m").unwrap()), + by_user_per_server: Some(crate::config::parse_rate("60/m").unwrap()), + by_tenant: None, + by_tool: HashMap::new(), + algorithm: Algorithm::FixedWindow, + }; + let engine = RateLimiterEngine::new_with_clock(cfg, Arc::new(clock)); + let server_id = engine.validated_server_id(Some("server-a")).unwrap(); + let checks = engine.build_checks("alice", Some("acme"), "search", Some("acme"), server_id); + let keys: Vec<&str> = checks.iter().map(|(key, _, _)| key.as_str()).collect(); + + assert!(keys.contains(&"acme:user:alice"), "keys: {keys:?}"); + assert!( + keys.contains(&"acme:user:alice:server:server-a"), + "keys: {keys:?}" + ); + } + + #[test] + fn per_server_user_limit_requires_safe_server_id() { + init_python(); + let (clock, _handle) = FakeClock::new(1_000_000); + let cfg = EngineConfig { + by_user: None, + by_user_per_server: Some(crate::config::parse_rate("60/m").unwrap()), + by_tenant: None, + by_tool: HashMap::new(), + algorithm: Algorithm::FixedWindow, + }; + let engine = RateLimiterEngine::new_with_clock(cfg, Arc::new(clock)); + + assert!(engine.validated_server_id(None).is_err()); + assert!(engine.validated_server_id(Some(" ")).is_err()); + assert!(engine.validated_server_id(Some("server:unsafe")).is_err()); + assert!(engine.validated_server_id(Some(&"x".repeat(129))).is_err()); + assert_eq!( + engine.validated_server_id(Some(" server-a ")).unwrap(), + Some("server-a") + ); + } + + #[test] + fn per_server_user_counters_are_independent_for_every_algorithm() { + for algorithm in [ + Algorithm::FixedWindow, + Algorithm::SlidingWindow, + Algorithm::TokenBucket, + ] { + init_python(); + let (clock, _handle) = FakeClock::new(1_000_000); + let cfg = EngineConfig { + by_user: None, + by_user_per_server: Some(crate::config::parse_rate("2/s").unwrap()), + by_tenant: None, + by_tool: HashMap::new(), + algorithm, + }; + let engine = RateLimiterEngine::new_with_clock(cfg, Arc::new(clock)); + let checks_for = |server_id| { + let server_id = engine.validated_server_id(Some(server_id)).unwrap(); + engine.build_checks("alice", None, "search", None, server_id) + }; + + assert!( + engine + .evaluate_many(checks_for("server-a"), 1_000_000) + .unwrap() + .allowed + ); + assert!( + engine + .evaluate_many(checks_for("server-a"), 1_000_000) + .unwrap() + .allowed + ); + assert!( + !engine + .evaluate_many(checks_for("server-a"), 1_000_000) + .unwrap() + .allowed + ); + assert!( + engine + .evaluate_many(checks_for("server-b"), 1_000_000) + .unwrap() + .allowed + ); + } + } + + #[test] + fn global_user_counter_caps_traffic_across_servers() { + init_python(); + let (clock, _handle) = FakeClock::new(1_000_000); + let cfg = EngineConfig { + by_user: Some(crate::config::parse_rate("3/s").unwrap()), + by_user_per_server: Some(crate::config::parse_rate("2/s").unwrap()), + by_tenant: None, + by_tool: HashMap::new(), + algorithm: Algorithm::FixedWindow, + }; + let engine = RateLimiterEngine::new_with_clock(cfg, Arc::new(clock)); + let checks_for = |server_id| { + let server_id = engine.validated_server_id(Some(server_id)).unwrap(); + engine.build_checks("alice", None, "search", None, server_id) + }; + + assert!( + engine + .evaluate_many(checks_for("server-a"), 1_000_000) + .unwrap() + .allowed + ); + assert!( + engine + .evaluate_many(checks_for("server-a"), 1_000_000) + .unwrap() + .allowed + ); + assert!( + engine + .evaluate_many(checks_for("server-b"), 1_000_000) + .unwrap() + .allowed + ); + assert!( + !engine + .evaluate_many(checks_for("server-b"), 1_000_000) + .unwrap() + .allowed + ); + } + #[test] fn different_prefixes_produce_isolated_counters() { let (engine, _handle) = engine_with_fake_clock(Some("2/s"), Algorithm::FixedWindow); // Exhaust limit for team_a - let checks_a = || engine.build_checks("alice", None, "search", Some("team_a")); + let checks_a = || engine.build_checks("alice", None, "search", Some("team_a"), None); let _ = engine.evaluate_many(checks_a(), 1_000_000).unwrap(); let _ = engine.evaluate_many(checks_a(), 1_000_000).unwrap(); let result_a = engine.evaluate_many(checks_a(), 1_000_000).unwrap(); @@ -748,7 +942,7 @@ mod tests { ); // team_b should still be allowed — different prefix, different counters - let checks_b = || engine.build_checks("alice", None, "search", Some("team_b")); + let checks_b = || engine.build_checks("alice", None, "search", Some("team_b"), None); let result_b = engine.evaluate_many(checks_b(), 1_000_000).unwrap(); assert!( result_b.allowed, @@ -759,8 +953,8 @@ mod tests { #[test] fn empty_prefix_matches_no_prefix_behavior() { let (engine, _handle) = engine_with_fake_clock(Some("10/s"), Algorithm::FixedWindow); - let checks_none = engine.build_checks("alice", None, "search", None); - let checks_empty = engine.build_checks("alice", None, "search", Some("")); + let checks_none = engine.build_checks("alice", None, "search", None, None); + let checks_empty = engine.build_checks("alice", None, "search", Some(""), None); // Both should produce the same unprefixed keys assert_eq!(checks_none.len(), checks_empty.len()); for ((k1, _, _), (k2, _, _)) in checks_none.iter().zip(checks_empty.iter()) { @@ -778,6 +972,7 @@ mod tests { // Every key the engine recognises today, including the four TLS knobs. for k in [ "by_user", + "by_user_per_server", "by_tenant", "by_tool", "algorithm", diff --git a/plugins/rust/python-package/rate_limiter/src/lib.rs b/plugins/rust/python-package/rate_limiter/src/lib.rs index 89d1704e..96fd1b97 100644 --- a/plugins/rust/python-package/rate_limiter/src/lib.rs +++ b/plugins/rust/python-package/rate_limiter/src/lib.rs @@ -26,6 +26,7 @@ pub use types::{EvalDimension, EvalResult}; fn compat_default_config(py: Python<'_>) -> PyResult> { let defaults = PyDict::new(py); defaults.set_item("by_user", py.None())?; + defaults.set_item("by_user_per_server", py.None())?; defaults.set_item("by_tenant", py.None())?; defaults.set_item("by_tool", py.None())?; defaults.set_item("algorithm", "fixed_window")?; diff --git a/plugins/rust/python-package/rate_limiter/src/plugin.rs b/plugins/rust/python-package/rate_limiter/src/plugin.rs index 782fb858..8fa8c723 100644 --- a/plugins/rust/python-package/rate_limiter/src/plugin.rs +++ b/plugins/rust/python-package/rate_limiter/src/plugin.rs @@ -87,7 +87,13 @@ impl RateLimiterPluginCore { .extract::()? .trim() .to_ascii_lowercase(); - let (user, tenant) = extract_request_context(context)?; + let (user, tenant, server_id) = extract_request_context(context)?; + if let Err(message) = self.engine.validated_server_id(server_id.as_deref()) { + warn!("rate limiter: refusing prompt request: {message}"); + return Ok( + rate_limit_context_error_result(py, "PromptPrehookResult", message)?.into_bound(py), + ); + } // Use tenant_id as the context prefix so that each team's rate limit // counters are isolated in Redis. Without this, all teams share keys. let context_prefix = tenant.as_deref(); @@ -99,6 +105,7 @@ impl RateLimiterPluginCore { tenant.as_deref(), &prompt, context_prefix, + server_id.as_deref(), ) { Ok((allowed, headers, meta)) => Ok(build_prehook_result( py, @@ -130,6 +137,7 @@ impl RateLimiterPluginCore { tenant.as_deref(), &prompt, context_prefix_owned.as_deref(), + server_id.as_deref(), ) .await { @@ -167,7 +175,13 @@ impl RateLimiterPluginCore { .extract::()? .trim() .to_ascii_lowercase(); - let (user, tenant) = extract_request_context(context)?; + let (user, tenant, server_id) = extract_request_context(context)?; + if let Err(message) = self.engine.validated_server_id(server_id.as_deref()) { + warn!("rate limiter: refusing tool request: {message}"); + return Ok( + rate_limit_context_error_result(py, "ToolPreInvokeResult", message)?.into_bound(py), + ); + } let context_prefix = tenant.as_deref(); let fail_closed = self.fail_closed; if !self.use_async { @@ -177,6 +191,7 @@ impl RateLimiterPluginCore { tenant.as_deref(), &tool, context_prefix, + server_id.as_deref(), ) { Ok((allowed, headers, meta)) => Ok(build_prehook_result( py, @@ -208,6 +223,7 @@ impl RateLimiterPluginCore { tenant.as_deref(), &tool, context_prefix_owned.as_deref(), + server_id.as_deref(), ) .await { @@ -320,12 +336,63 @@ fn backend_error_result( } } +fn rate_limit_context_error_result( + py: Python<'_>, + class_name: &str, + message: &str, +) -> PyResult> { + let details = PyDict::new(py); + details.set_item("field", "server_id")?; + details.set_item("error", message)?; + let violation = build_framework_object( + py, + "PluginViolation", + [ + ( + "reason", + "Rate limit context unavailable" + .into_pyobject(py)? + .into_any() + .unbind(), + ), + ( + "description", + message.into_pyobject(py)?.into_any().unbind(), + ), + ( + "code", + "RATE_LIMIT_CONTEXT_MISSING" + .into_pyobject(py)? + .into_any() + .unbind(), + ), + ("details", details.into_any().unbind()), + ( + "http_status_code", + 503i32.into_pyobject(py)?.into_any().unbind(), + ), + ], + )?; + build_framework_object( + py, + class_name, + [ + ( + "continue_processing", + false.into_pyobject(py)?.to_owned().into_any().unbind(), + ), + ("violation", violation), + ], + ) +} + fn evaluate_sync_request( engine: &RateLimiterEngine, user: &str, tenant: Option<&str>, tool_or_prompt: &str, context_prefix: Option<&str>, + server_id: Option<&str>, ) -> PyResult<(bool, Py, Py)> { let now_unix = std::time::SystemTime::now() .duration_since(std::time::UNIX_EPOCH) @@ -341,6 +408,7 @@ fn evaluate_sync_request( now_unix, true, context_prefix, + server_id, )?; Ok((allowed, headers.unbind(), meta.unbind())) }) @@ -352,6 +420,7 @@ async fn evaluate_async_request( tenant: Option<&str>, tool_or_prompt: &str, context_prefix: Option<&str>, + server_id: Option<&str>, ) -> PyResult<(bool, Py, Py)> { let now_unix = std::time::SystemTime::now() .duration_since(std::time::UNIX_EPOCH) @@ -367,6 +436,7 @@ async fn evaluate_async_request( now_unix, true, context_prefix, + server_id, ) .map(|awaitable| awaitable.unbind()) })?; @@ -583,21 +653,27 @@ fn build_violation( ) } -fn extract_request_context(context: &Bound<'_, PyAny>) -> PyResult<(String, Option)> { +fn extract_request_context( + context: &Bound<'_, PyAny>, +) -> PyResult<(String, Option, Option)> { let global_context = context.getattr("global_context")?; let user = extract_user_identity(&global_context.getattr("user")?)?; - let tenant = match global_context.getattr("tenant_id") { + let tenant = extract_optional_string_attribute(&global_context, "tenant_id")?; + let server_id = extract_optional_string_attribute(&global_context, "server_id")?; + Ok((user, tenant, server_id)) +} + +fn extract_optional_string_attribute( + object: &Bound<'_, PyAny>, + name: &str, +) -> PyResult> { + match object.getattr(name) { Ok(value) if !value.is_none() => { let trimmed = value.extract::()?.trim().to_string(); - if trimmed.is_empty() { - None - } else { - Some(trimmed) - } + Ok((!trimmed.is_empty()).then_some(trimmed)) } - _ => None, - }; - Ok((user, tenant)) + _ => Ok(None), + } } fn extract_user_identity(user: &Bound<'_, PyAny>) -> PyResult { @@ -712,13 +788,14 @@ class PromptPayload: self.prompt_id = prompt_id class GlobalContext: - def __init__(self, user): + def __init__(self, user, server_id=None): self.user = user self.tenant_id = None + self.server_id = server_id class Context: - def __init__(self, user): - self.global_context = GlobalContext(user) + def __init__(self, user, server_id=None): + self.global_context = GlobalContext(user, server_id) "# ), pyo3::ffi::c_str!("rl_test_payloads.py"), @@ -827,6 +904,36 @@ class Context: .unwrap(); } + #[test] + fn tool_pre_invoke_blocks_missing_server_id_for_per_server_limit() { + Python::initialize(); + Python::attach(|py| -> PyResult<()> { + install_framework_module(py)?; + let module = payload_module(py)?; + + let config = PyDict::new(py); + config.set_item("by_user_per_server", "5/s")?; + config.set_item("backend", "memory")?; + let plugin = RateLimiterPluginCore::new(&config)?; + let payload = module.getattr("ToolPayload")?.call1(("search",))?; + let context = module.getattr("Context")?.call1(("alice",))?; + + let result = plugin.tool_pre_invoke(py, &payload, &context, None)?; + assert!(!result.getattr("continue_processing")?.extract::()?); + let violation = result.getattr("violation")?; + assert_eq!( + violation.getattr("code")?.extract::()?, + "RATE_LIMIT_CONTEXT_MISSING" + ); + assert_eq!( + violation.getattr("http_status_code")?.extract::()?, + 503 + ); + Ok(()) + }) + .unwrap(); + } + #[test] fn tool_pre_invoke_allowed_emits_metrics_and_headers_with_trace_id_present() { Python::initialize(); diff --git a/plugins/tests/rate_limiter/test_config_integration.py b/plugins/tests/rate_limiter/test_config_integration.py index 25a07f92..f610e146 100644 --- a/plugins/tests/rate_limiter/test_config_integration.py +++ b/plugins/tests/rate_limiter/test_config_integration.py @@ -28,6 +28,7 @@ class TestRateLimiterPluginConfig: def test_module_level_config_defaults_remain_importable(self): config = RateLimiterConfig() assert config.by_user is None + assert config.by_user_per_server is None assert config.by_tenant is None assert config.by_tool is None assert config.algorithm == "fixed_window" @@ -38,6 +39,7 @@ def test_module_level_config_defaults_remain_importable(self): def test_top_level_package_reexports_public_compatibility_names(self): config = PackageRateLimiterConfig() assert PackageRateLimiterPlugin is RateLimiterPlugin + assert config.by_user_per_server is None assert config.algorithm == "fixed_window" assert package_parse_rate("60/sec") == (60, 1) @@ -68,6 +70,7 @@ def test_defaults_construct_successfully(self): def test_all_fields_construct_successfully(self): plugin = RateLimiterPlugin(_config( by_user="60/m", + by_user_per_server="20/m", by_tenant="600/m", by_tool={"search": "10/s"}, algorithm="sliding_window", @@ -88,6 +91,10 @@ def test_invalid_by_user_rate_rejected(self): with pytest.raises(ValueError, match="by_user"): RateLimiterPlugin(_config(by_user="not-a-rate")) + def test_invalid_by_user_per_server_rate_rejected(self): + with pytest.raises(ValueError, match="by_user_per_server"): + RateLimiterPlugin(_config(by_user_per_server="not-a-rate")) + def test_invalid_by_tenant_rate_rejected(self): with pytest.raises(ValueError, match="by_tenant"): RateLimiterPlugin(_config(by_tenant="bad")) diff --git a/plugins/tests/rate_limiter/test_integration.py b/plugins/tests/rate_limiter/test_integration.py index f4d005dd..832f1346 100644 --- a/plugins/tests/rate_limiter/test_integration.py +++ b/plugins/tests/rate_limiter/test_integration.py @@ -49,9 +49,9 @@ def _make_config(**overrides) -> PluginConfig: return PluginConfig(name="rate_limiter", config=config) -def _make_context(user="testuser", tenant_id="tenant-1") -> PluginContext: +def _make_context(user="testuser", tenant_id="tenant-1", server_id=None) -> PluginContext: return PluginContext( - global_context=GlobalContext(user=user, tenant_id=tenant_id), + global_context=GlobalContext(user=user, tenant_id=tenant_id, server_id=server_id), ) @@ -128,7 +128,6 @@ async def test_blocked_over_limit(self, plugin): assert result.continue_processing is False assert result.violation is not None assert result.violation.http_status_code == 429 - assert result.violation.code == "RATE_LIMIT" async def test_different_users_independent(self, plugin): payload = ToolPreInvokePayload(name="search") @@ -139,6 +138,64 @@ async def test_different_users_independent(self, plugin): result = await plugin.tool_pre_invoke(payload, _make_context(user="userB")) assert result.continue_processing is True + @pytest.mark.parametrize("algorithm", ["fixed_window", "sliding_window", "token_bucket"]) + async def test_per_server_user_limits_are_independent(self, algorithm): + plugin = RateLimiterPlugin( + _make_config( + by_user=None, + by_user_per_server="2/s", + algorithm=algorithm, + ) + ) + payload = ToolPreInvokePayload(name="search") + server_a = _make_context(user="alice", server_id="server-a") + server_b = _make_context(user="alice", server_id="server-b") + + for _ in range(2): + result = await plugin.tool_pre_invoke(payload, server_a) + assert result.continue_processing is True + + blocked = await plugin.tool_pre_invoke(payload, server_a) + allowed = await plugin.tool_pre_invoke(payload, server_b) + other_user = await plugin.tool_pre_invoke( + payload, + _make_context(user="bob", server_id="server-a"), + ) + + assert blocked.continue_processing is False + assert blocked.violation.code == "RATE_LIMIT" + assert allowed.continue_processing is True + assert other_user.continue_processing is True + + async def test_global_user_limit_still_caps_all_servers(self): + plugin = RateLimiterPlugin( + _make_config(by_user="3/s", by_user_per_server="2/s") + ) + payload = ToolPreInvokePayload(name="search") + server_a = _make_context(user="alice", server_id="server-a") + server_b = _make_context(user="alice", server_id="server-b") + + assert (await plugin.tool_pre_invoke(payload, server_a)).continue_processing + assert (await plugin.tool_pre_invoke(payload, server_a)).continue_processing + assert (await plugin.tool_pre_invoke(payload, server_b)).continue_processing + + blocked = await plugin.tool_pre_invoke(payload, server_b) + assert blocked.continue_processing is False + assert blocked.violation.http_status_code == 429 + + async def test_per_server_limit_blocks_when_server_id_missing(self): + plugin = RateLimiterPlugin( + _make_config(by_user=None, by_user_per_server="2/s") + ) + result = await plugin.tool_pre_invoke( + ToolPreInvokePayload(name="search"), + _make_context(user="alice", server_id=None), + ) + + assert result.continue_processing is False + assert result.violation.code == "RATE_LIMIT_CONTEXT_MISSING" + assert result.violation.http_status_code == 503 + async def test_dict_user_identity_uses_email_before_other_fields(self, plugin): payload = ToolPreInvokePayload(name="search") for _ in range(5): @@ -266,6 +323,20 @@ async def test_blocked_over_limit(self, plugin): assert result.violation is not None assert result.violation.http_status_code == 429 + async def test_per_server_prompt_limits_are_independent(self): + plugin = RateLimiterPlugin( + _make_config(by_user=None, by_user_per_server="1/s") + ) + payload = PromptPrehookPayload(prompt_id="my-prompt") + server_a = _make_context(user="alice", server_id="server-a") + server_b = _make_context(user="alice", server_id="server-b") + + assert (await plugin.prompt_pre_fetch(payload, server_a)).continue_processing + blocked = await plugin.prompt_pre_fetch(payload, server_a) + assert not blocked.continue_processing + assert blocked.violation.code == "RATE_LIMIT" + assert (await plugin.prompt_pre_fetch(payload, server_b)).continue_processing + # --------------------------------------------------------------------------- # by_tenant limiting diff --git a/plugins/tests/rate_limiter/test_redis_integration.py b/plugins/tests/rate_limiter/test_redis_integration.py index da7bd76f..8196fd2f 100644 --- a/plugins/tests/rate_limiter/test_redis_integration.py +++ b/plugins/tests/rate_limiter/test_redis_integration.py @@ -1013,20 +1013,27 @@ def redis_url_for_integration(): subprocess.run(["docker", "stop", container_id], check=False) -def _make_redis_plugin(redis_url: str, algorithm: str = "fixed_window", limit: str = "3/s") -> RateLimiterPlugin: +def _make_redis_plugin( + redis_url: str, + algorithm: str = "fixed_window", + limit: str | None = "3/s", + **overrides, +) -> RateLimiterPlugin: """Create a RateLimiterPlugin backed by real Redis.""" + config = { + "by_user": limit, + "backend": "redis", + "redis_url": redis_url, + "algorithm": algorithm, + } + config.update(overrides) return RateLimiterPlugin( PluginConfig( name="RateLimiter", kind="cpex_rate_limiter.rate_limiter.RateLimiterPlugin", hooks=["tool_pre_invoke"], priority=100, - config={ - "by_user": limit, - "backend": "redis", - "redis_url": redis_url, - "algorithm": algorithm, - }, + config=config, ) ) @@ -1115,6 +1122,45 @@ async def test_redis_shared_counter_across_plugin_instances(self, redis_url_for_ result = await plugin_b.tool_pre_invoke(payload, ctx) assert result.violation is not None, "Redis backend must share counters across plugin instances — " "instance B must be blocked after instance A exhausts the limit" + @pytest.mark.asyncio + async def test_redis_per_server_user_counters_are_shared_and_isolated(self, redis_url_for_integration): + """Replicas share one user/server bucket while different servers remain independent.""" + await _flush_redis(redis_url_for_integration) + + plugin_a = _make_redis_plugin( + redis_url_for_integration, + limit=None, + by_user_per_server="2/s", + ) + plugin_b = _make_redis_plugin( + redis_url_for_integration, + limit=None, + by_user_per_server="2/s", + ) + payload = ToolPreInvokePayload(name="tool", arguments={}) + server_a = PluginContext( + global_context=GlobalContext(request_id="r1", user="alice", server_id="server-a") + ) + server_b = PluginContext( + global_context=GlobalContext(request_id="r2", user="alice", server_id="server-b") + ) + + assert (await plugin_a.tool_pre_invoke(payload, server_a)).violation is None + assert (await plugin_a.tool_pre_invoke(payload, server_a)).violation is None + assert (await plugin_b.tool_pre_invoke(payload, server_a)).violation is not None + assert (await plugin_b.tool_pre_invoke(payload, server_b)).violation is None + + keys_a = await _keys_in_redis( + redis_url_for_integration, + "rl:user:alice:server:server-a:*", + ) + keys_b = await _keys_in_redis( + redis_url_for_integration, + "rl:user:alice:server:server-b:*", + ) + assert keys_a == ["rl:user:alice:server:server-a:1"] + assert keys_b == ["rl:user:alice:server:server-b:1"] + @pytest.mark.asyncio async def test_redis_window_resets_after_ttl(self, redis_url_for_integration): """After the rate window expires, Redis TTL resets counters and requests are allowed again.""" @@ -2132,4 +2178,3 @@ async def test_rediss_handshake_succeeds_against_real_tls_redis( "expected a counter key for user 'alice' to appear in TLS Redis " f"after a successful rustls handshake; got keys={keys!r}" ) - From adfa02ba71b73974af0895fe252a4595289e67fc Mon Sep 17 00:00:00 2001 From: Vishu Bhatnagar Date: Tue, 6 Oct 2026 12:26:39 +0100 Subject: [PATCH 2/3] test(rate-limiter): cover server ID and context boundaries Signed-off-by: Vishu Bhatnagar --- .../python-package/rate_limiter/src/engine.rs | 7 +++++ .../python-package/rate_limiter/src/plugin.rs | 29 ++++++++++++++++++- 2 files changed, 35 insertions(+), 1 deletion(-) diff --git a/plugins/rust/python-package/rate_limiter/src/engine.rs b/plugins/rust/python-package/rate_limiter/src/engine.rs index 40b50151..9f47ba51 100644 --- a/plugins/rust/python-package/rate_limiter/src/engine.rs +++ b/plugins/rust/python-package/rate_limiter/src/engine.rs @@ -829,6 +829,13 @@ mod tests { assert!(engine.validated_server_id(None).is_err()); assert!(engine.validated_server_id(Some(" ")).is_err()); assert!(engine.validated_server_id(Some("server:unsafe")).is_err()); + let max_length_server_id = "x".repeat(128); + assert_eq!( + engine + .validated_server_id(Some(&max_length_server_id)) + .unwrap(), + Some(max_length_server_id.as_str()) + ); assert!(engine.validated_server_id(Some(&"x".repeat(129))).is_err()); assert_eq!( engine.validated_server_id(Some(" server-a ")).unwrap(), diff --git a/plugins/rust/python-package/rate_limiter/src/plugin.rs b/plugins/rust/python-package/rate_limiter/src/plugin.rs index 8fa8c723..2b0fadbf 100644 --- a/plugins/rust/python-package/rate_limiter/src/plugin.rs +++ b/plugins/rust/python-package/rate_limiter/src/plugin.rs @@ -719,7 +719,7 @@ fn log_exception(py: Python<'_>, message: &str) -> PyResult<()> { mod tests { use super::await_async_tuple; use super::ensure_crypto_provider; - use super::{RateLimiterPluginCore, read_trace_id}; + use super::{RateLimiterPluginCore, extract_request_context, read_trace_id}; use pyo3::prelude::*; use pyo3::types::{PyAnyMethods, PyDict, PyDictMethods, PyModule}; @@ -803,6 +803,33 @@ class Context: ) } + #[test] + fn request_context_ignores_blank_optional_identifiers() { + Python::initialize(); + Python::attach(|py| -> PyResult<()> { + let module = payload_module(py)?; + let context = module.getattr("Context")?.call1(("alice", " "))?; + context + .getattr("global_context")? + .setattr("tenant_id", " ")?; + + assert_eq!( + extract_request_context(&context)?, + ("alice".to_string(), None, None) + ); + + context + .getattr("global_context")? + .setattr("server_id", " server-a ")?; + assert_eq!( + extract_request_context(&context)?, + ("alice".to_string(), None, Some("server-a".to_string())) + ); + Ok(()) + }) + .unwrap(); + } + fn extensions_with_trace<'py>(py: Python<'py>, trace_id: &str) -> PyResult> { let ext_module = PyModule::from_code( py, From 0e9bcc516cc643133fe286769ab685d5e44a8642 Mon Sep 17 00:00:00 2001 From: Vishu Bhatnagar Date: Tue, 6 Oct 2026 14:51:16 +0100 Subject: [PATCH 3/3] test(rate-limiter): avoid token bucket refill race Signed-off-by: Vishu Bhatnagar --- plugins/tests/rate_limiter/test_redis_integration.py | 8 +++++--- 1 file changed, 5 insertions(+), 3 deletions(-) diff --git a/plugins/tests/rate_limiter/test_redis_integration.py b/plugins/tests/rate_limiter/test_redis_integration.py index 8196fd2f..79686c34 100644 --- a/plugins/tests/rate_limiter/test_redis_integration.py +++ b/plugins/tests/rate_limiter/test_redis_integration.py @@ -1280,7 +1280,8 @@ async def test_redis_token_bucket_enforces_limit(self, redis_url_for_integration """token_bucket on real Redis blocks when bucket is empty.""" await _flush_redis(redis_url_for_integration) - plugin = _make_redis_plugin(redis_url_for_integration, algorithm="token_bucket", limit="3/s") + # Keep refill out of the immediate exhaustion check on slow runners. + plugin = _make_redis_plugin(redis_url_for_integration, algorithm="token_bucket", limit="3/m") ctx = PluginContext(global_context=GlobalContext(request_id="r1", user="alice")) payload = ToolPreInvokePayload(name="tool", arguments={}) @@ -1301,8 +1302,9 @@ async def test_redis_token_bucket_shared_counter_across_instances(self, redis_ur """ await _flush_redis(redis_url_for_integration) - plugin_a = _make_redis_plugin(redis_url_for_integration, algorithm="token_bucket", limit="3/s") - plugin_b = _make_redis_plugin(redis_url_for_integration, algorithm="token_bucket", limit="3/s") + # A slow refill avoids a wall-clock race in the cross-instance assertion. + plugin_a = _make_redis_plugin(redis_url_for_integration, algorithm="token_bucket", limit="3/m") + plugin_b = _make_redis_plugin(redis_url_for_integration, algorithm="token_bucket", limit="3/m") ctx = PluginContext(global_context=GlobalContext(request_id="r1", user="alice")) payload = ToolPreInvokePayload(name="tool", arguments={})