diff --git a/model_gateway/src/app_context.rs b/model_gateway/src/app_context.rs index 2873848039..fd2b6a85e0 100644 --- a/model_gateway/src/app_context.rs +++ b/model_gateway/src/app_context.rs @@ -19,6 +19,7 @@ use crate::{ middleware::TokenBucket, observability::inflight_tracker::InFlightRequestTracker, policies::PolicyRegistry, + rate_limit::LocalTokenRateLimiter, routers::{ common::openai_bridge::FormatRegistry, grpc::multimodal::MultimodalConfigRegistry, openai::realtime::RealtimeRegistry, router_manager::RouterManager, @@ -51,6 +52,7 @@ pub struct AppContext { pub client: Client, pub router_config: RouterConfig, pub rate_limiter: Option>, + pub token_rate_limiter: Option>, pub tokenizer_registry: Arc, pub multimodal_config_registry: Arc, pub reasoning_parser_factory: Option, @@ -91,6 +93,7 @@ pub struct AppContextBuilder { client: Option, router_config: Option, rate_limiter: Option>, + token_rate_limiter: Option>, tokenizer_registry: Option>, reasoning_parser_factory: Option, tool_parser_factory: Option, @@ -144,6 +147,7 @@ impl AppContextBuilder { client: None, router_config: None, rate_limiter: None, + token_rate_limiter: None, tokenizer_registry: None, reasoning_parser_factory: None, tool_parser_factory: None, @@ -337,6 +341,7 @@ impl AppContextBuilder { .ok_or(AppContextBuildError::MissingField("client"))?, router_config, rate_limiter: self.rate_limiter, + token_rate_limiter: self.token_rate_limiter, tokenizer_registry: self .tokenizer_registry .ok_or(AppContextBuildError::MissingField("tokenizer_registry"))?, @@ -488,6 +493,11 @@ impl AppContextBuilder { ))) } }; + self.token_rate_limiter = config.multi_tenant_rate_limit.enabled.then(|| { + Arc::new(LocalTokenRateLimiter::new( + config.multi_tenant_rate_limit.clone(), + )) + }); self } diff --git a/model_gateway/src/config/builder.rs b/model_gateway/src/config/builder.rs index b9bf025852..913805ab0c 100644 --- a/model_gateway/src/config/builder.rs +++ b/model_gateway/src/config/builder.rs @@ -1,4 +1,4 @@ -use std::collections::HashMap; +use std::collections::{hash_map::Entry, HashMap}; use smg_mcp::McpConfig; @@ -7,7 +7,7 @@ use super::{ HistoryBackend, MetricsConfig, OracleConfig, PolicyConfig, PostgresConfig, RedisConfig, RetryConfig, RouterConfig, RoutingMode, TokenizerCacheConfig, TraceConfig, }; -use crate::worker::ConnectionMode; +use crate::{rate_limit::TenantTokenPolicy, worker::ConnectionMode}; /// Builder for RouterConfig that wraps the config itself /// This eliminates field duplication and stays in sync automatically @@ -608,6 +608,59 @@ impl RouterConfigBuilder { self } + pub fn multi_tenant_rate_limit_enabled(mut self, enabled: bool) -> Self { + self.config.multi_tenant_rate_limit.enabled = enabled; + self + } + + pub fn default_tokens_per_minute(mut self, limit: u32) -> Self { + self.config + .multi_tenant_rate_limit + .default_tokens_per_minute = limit; + self + } + + pub fn default_requests_per_minute(mut self, limit: u32) -> Self { + self.config + .multi_tenant_rate_limit + .default_requests_per_minute = limit; + self + } + + pub fn tenant_rate_limit>( + mut self, + tenant_key: S, + tokens_per_minute: u32, + requests_per_minute: u32, + ) -> Self { + let tenant_key = tenant_key.into(); + let new_policy = TenantTokenPolicy { + tokens_per_minute, + requests_per_minute, + }; + + match self + .config + .multi_tenant_rate_limit + .tenants + .entry(tenant_key.clone()) + { + Entry::Vacant(entry) => { + entry.insert(new_policy); + } + Entry::Occupied(mut entry) => { + tracing::warn!( + tenant_key = %tenant_key, + previous_policy = ?entry.get(), + new_policy = ?new_policy, + "overwriting duplicate tenant rate limit policy" + ); + entry.insert(new_policy); + } + } + self + } + pub fn maybe_model_path(mut self, path: Option>) -> Self { self.config.model_path = path.map(|p| p.into()); self @@ -897,6 +950,56 @@ mod tests { assert!(modified.trace_config.is_some()); } + #[test] + fn test_builder_multi_tenant_rate_limit_round_trip() { + let config = RouterConfigBuilder::new() + .regular_mode(vec!["http://worker1:8000".to_string()]) + .multi_tenant_rate_limit_enabled(true) + .default_tokens_per_minute(10_000) + .default_requests_per_minute(60) + .tenant_rate_limit("team-a", 50_000, 600) + .tenant_rate_limit("team-b", 100_000, 1_200) + .build() + .unwrap(); + + assert!(config.multi_tenant_rate_limit.enabled); + assert_eq!( + config.multi_tenant_rate_limit.default_tokens_per_minute, + 10_000 + ); + assert_eq!( + config.multi_tenant_rate_limit.default_requests_per_minute, + 60 + ); + let team_a = config + .multi_tenant_rate_limit + .tenants + .get("team-a") + .expect("team-a override registered"); + assert_eq!(team_a.tokens_per_minute, 50_000); + assert_eq!(team_a.requests_per_minute, 600); + assert_eq!(config.multi_tenant_rate_limit.tenants.len(), 2); + } + + #[test] + fn test_builder_duplicate_tenant_rate_limit_overwrites_latest_policy() { + let config = RouterConfigBuilder::new() + .regular_mode(vec!["http://worker1:8000".to_string()]) + .tenant_rate_limit("team-a", 50_000, 600) + .tenant_rate_limit("team-a", 100_000, 1_200) + .build() + .unwrap(); + + let team_a = config + .multi_tenant_rate_limit + .tenants + .get("team-a") + .expect("team-a override registered"); + assert_eq!(team_a.tokens_per_minute, 100_000); + assert_eq!(team_a.requests_per_minute, 1_200); + assert_eq!(config.multi_tenant_rate_limit.tenants.len(), 1); + } + /// Test complex routing mode helper method #[test] fn test_builder_prefill_decode_mode() { diff --git a/model_gateway/src/config/types.rs b/model_gateway/src/config/types.rs index 7a08ca8d0a..6750906c45 100644 --- a/model_gateway/src/config/types.rs +++ b/model_gateway/src/config/types.rs @@ -8,7 +8,10 @@ pub use smg_data_connector::{ }; use super::{validation::ConfigValidator, ConfigResult}; -use crate::{tenant::DEFAULT_TENANT_HEADER_NAME, worker::ConnectionMode}; +use crate::{ + rate_limit::MultiTenantRateLimitConfig, tenant::DEFAULT_TENANT_HEADER_NAME, + worker::ConnectionMode, +}; /// Main router configuration #[derive(Debug, Clone, Serialize, Deserialize)] @@ -54,6 +57,8 @@ pub struct RouterConfig { pub storage_context_headers: HashMap, #[serde(default)] pub tenant_resolution: TenantResolutionConfig, + #[serde(default)] + pub multi_tenant_rate_limit: MultiTenantRateLimitConfig, /// Set to -1 to disable rate limiting pub max_concurrent_requests: i32, pub queue_size: usize, @@ -680,6 +685,7 @@ impl Default for RouterConfig { request_id_headers: None, storage_context_headers: HashMap::new(), tenant_resolution: TenantResolutionConfig::default(), + multi_tenant_rate_limit: MultiTenantRateLimitConfig::default(), max_concurrent_requests: -1, queue_size: 100, queue_timeout_secs: 60, diff --git a/model_gateway/src/lib.rs b/model_gateway/src/lib.rs index 2d8661262d..cafc830198 100644 --- a/model_gateway/src/lib.rs +++ b/model_gateway/src/lib.rs @@ -5,6 +5,7 @@ pub mod mesh; pub mod middleware; pub mod observability; pub mod policies; +pub mod rate_limit; pub mod routers; pub mod server; pub mod service_discovery; diff --git a/model_gateway/src/main.rs b/model_gateway/src/main.rs index 8a4c42e4d5..f808a45d5f 100644 --- a/model_gateway/src/main.rs +++ b/model_gateway/src/main.rs @@ -422,6 +422,22 @@ struct CliArgs { #[arg(long, help_heading = "Rate Limiting")] rate_limit_tokens_per_second: Option, + /// Enable tenant-aware token rate limiting. + #[arg(long, default_value_t = false, help_heading = "Rate Limiting")] + multi_tenant_rate_limit_enabled: bool, + + /// Default token budget per minute for tenants without an explicit override. + #[arg(long, default_value_t = 0, help_heading = "Rate Limiting")] + default_tokens_per_minute: u32, + + /// Default request budget per minute for tenants without an explicit override. + #[arg(long, default_value_t = 0, help_heading = "Rate Limiting")] + default_requests_per_minute: u32, + + /// Per-tenant override in the form tenant_key:tpm:rpm, e.g. header:team-a:1000:10 + #[arg(long = "tenant-rate-limit", num_args = 0.., help_heading = "Rate Limiting")] + tenant_rate_limits: Vec, + // ==================== Retry Configuration ==================== /// Maximum number of retry attempts #[arg(long, default_value_t = 5, help_heading = "Retry Configuration")] @@ -1330,6 +1346,9 @@ impl CliArgs { .trust_tenant_header(self.trust_tenant_header) .tenant_header_name(&self.tenant_header_name) .maybe_rate_limit_tokens_per_second(self.rate_limit_tokens_per_second) + .multi_tenant_rate_limit_enabled(self.multi_tenant_rate_limit_enabled) + .default_tokens_per_minute(self.default_tokens_per_minute) + .default_requests_per_minute(self.default_requests_per_minute) .maybe_model_path(self.model_path.as_ref()) .maybe_tokenizer_path(self.tokenizer_path.as_ref()) .maybe_chat_template(self.chat_template.as_ref()) @@ -1348,6 +1367,30 @@ impl CliArgs { .dp_minimum_tokens_scheduler(self.dp_minimum_tokens_scheduler) .maybe_server_cert_and_key(self.tls_cert_path.as_ref(), self.tls_key_path.as_ref()); + let mut builder = builder; + for spec in &self.tenant_rate_limits { + let mut parts = spec.rsplitn(3, ':'); + let rpm = parts.next().and_then(|s| s.parse::().ok()); + let tpm = parts.next().and_then(|s| s.parse::().ok()); + let tenant_key = parts.next(); + if let (Some(tenant_key), Some(tpm), Some(rpm)) = (tenant_key, tpm, rpm) { + if tenant_key.is_empty() { + return Err(ConfigError::ValidationFailed { + reason: format!( + "invalid --tenant-rate-limit '{spec}'; expected tenant_key:tpm:rpm" + ), + }); + } + builder = builder.tenant_rate_limit(tenant_key, tpm, rpm); + } else { + return Err(ConfigError::ValidationFailed { + reason: format!( + "invalid --tenant-rate-limit '{spec}'; expected tenant_key:tpm:rpm" + ), + }); + } + } + builder.build() } diff --git a/model_gateway/src/rate_limit/error.rs b/model_gateway/src/rate_limit/error.rs new file mode 100644 index 0000000000..e4950309a5 --- /dev/null +++ b/model_gateway/src/rate_limit/error.rs @@ -0,0 +1,49 @@ +use axum::{ + http::{self, header::RETRY_AFTER, HeaderValue}, + response::Response, +}; + +use super::local::TERMINAL_REJECTION_RETRY_AFTER_SECS; +use crate::routers::error::create_error; + +pub fn rate_limit_exceeded_response(retry_after_secs: u64) -> Response { + if retry_after_secs == TERMINAL_REJECTION_RETRY_AFTER_SECS { + return create_error( + http::StatusCode::PAYLOAD_TOO_LARGE, + "tenant_rate_limit_exceeded", + "Request exceeds the tenant capacity limit and cannot be retried without reducing its size", + ); + } + + let mut response = create_error( + http::StatusCode::TOO_MANY_REQUESTS, + "tenant_rate_limit_exceeded", + "Tenant rate limit exceeded for this request", + ); + + if let Ok(v) = HeaderValue::from_str(&retry_after_secs.max(1).to_string()) { + response.headers_mut().insert(RETRY_AFTER, v); + } + + response +} + +#[cfg(test)] +mod tests { + use axum::http::{header::RETRY_AFTER, StatusCode}; + + use super::{rate_limit_exceeded_response, TERMINAL_REJECTION_RETRY_AFTER_SECS}; + use crate::routers::error::extract_error_code_from_response; + + #[test] + fn returns_distinct_terminal_rate_limit_rejection() { + let response = rate_limit_exceeded_response(TERMINAL_REJECTION_RETRY_AFTER_SECS); + + assert_eq!(response.status(), StatusCode::PAYLOAD_TOO_LARGE); + assert!(response.headers().get(RETRY_AFTER).is_none()); + assert_eq!( + extract_error_code_from_response(&response), + "tenant_rate_limit_exceeded" + ); + } +} diff --git a/model_gateway/src/rate_limit/local.rs b/model_gateway/src/rate_limit/local.rs new file mode 100644 index 0000000000..b773ac72e2 --- /dev/null +++ b/model_gateway/src/rate_limit/local.rs @@ -0,0 +1,370 @@ +use std::{ + sync::{ + atomic::{AtomicU64, Ordering}, + Arc, + }, + time::{Duration, Instant}, +}; + +use dashmap::DashMap; +use parking_lot::Mutex; + +use super::types::{MultiTenantRateLimitConfig, TenantTokenPolicy}; + +pub const TERMINAL_REJECTION_RETRY_AFTER_SECS: u64 = u64::MAX; +const STALE_BUCKET_TTL_MULTIPLIER: u32 = 10; +const CLEANUP_INTERVAL: Duration = Duration::from_secs(60); + +#[derive(Debug)] +struct BucketState { + available_tokens: f64, + available_requests: f64, + last_refill: Instant, +} + +#[derive(Debug)] +struct Bucket { + policy: TenantTokenPolicy, + state: Mutex, +} + +/// In-memory tenant token bucket rate limiter. +/// +/// Bucket cardinality grows with the number of distinct `tenant_key` values observed. +/// To avoid unbounded growth from untrusted tenant identifiers, pair this limiter with +/// upstream tenant validation or keep `RouterConfigBuilder::trust_tenant_header(false)`. +/// +/// Each bucket stores a cloned `TenantTokenPolicy` at allocation time. Runtime policy +/// updates do not affect already-allocated buckets until they are evicted by cleanup or +/// otherwise replaced, so policy-change flows may need explicit bucket invalidation. +#[derive(Debug)] +pub struct LocalTokenRateLimiter { + config: MultiTenantRateLimitConfig, + buckets: DashMap>, + started_at: Instant, + last_cleanup_ms: AtomicU64, +} + +impl LocalTokenRateLimiter { + #[must_use] + pub fn new(config: MultiTenantRateLimitConfig) -> Self { + Self { + config, + buckets: DashMap::new(), + started_at: Instant::now(), + last_cleanup_ms: AtomicU64::new(0), + } + } + + #[must_use] + pub fn is_enabled(&self) -> bool { + self.config.enabled + } + + pub fn check_and_consume(&self, tenant_key: &str, estimated_tokens: u32) -> Result<(), u64> { + self.maybe_cleanup_stale_buckets(); + let Some(policy) = self.config.policy_for(tenant_key) else { + return Ok(()); + }; + + if policy.tokens_per_minute == 0 && policy.requests_per_minute == 0 { + return Ok(()); + } + + let bucket = if let Some(bucket) = self.buckets.get(tenant_key) { + bucket.value().clone() + } else { + self.buckets + .entry(tenant_key.to_string()) + .or_insert_with(|| { + Arc::new(Bucket { + policy: policy.clone(), + state: Mutex::new(BucketState { + available_tokens: policy.tokens_per_minute as f64, + available_requests: policy.requests_per_minute as f64, + last_refill: Instant::now(), + }), + }) + }) + .value() + .clone() + }; + + let mut state = bucket.state.lock(); + let now = Instant::now(); + let elapsed = now.duration_since(state.last_refill).as_secs_f64(); + state.last_refill = now; + + if bucket.policy.tokens_per_minute > 0 { + let refill = elapsed * (bucket.policy.tokens_per_minute as f64 / 60.0); + state.available_tokens = + (state.available_tokens + refill).min(bucket.policy.tokens_per_minute as f64); + } + if bucket.policy.requests_per_minute > 0 { + let refill = elapsed * (bucket.policy.requests_per_minute as f64 / 60.0); + state.available_requests = + (state.available_requests + refill).min(bucket.policy.requests_per_minute as f64); + } + + let need_tokens = estimated_tokens as f64; + let exceeds_token_capacity = bucket.policy.tokens_per_minute > 0 + && need_tokens > bucket.policy.tokens_per_minute as f64; + + if exceeds_token_capacity { + return Err(TERMINAL_REJECTION_RETRY_AFTER_SECS); + } + + let token_denied = + bucket.policy.tokens_per_minute > 0 && state.available_tokens < need_tokens; + let request_denied = + bucket.policy.requests_per_minute > 0 && state.available_requests < 1.0; + + if token_denied || request_denied { + let token_retry = if token_denied && bucket.policy.tokens_per_minute > 0 { + let debt = (need_tokens - state.available_tokens).max(0.0); + (debt / (bucket.policy.tokens_per_minute as f64 / 60.0)).ceil() as u64 + } else { + 0 + }; + let request_retry = if request_denied && bucket.policy.requests_per_minute > 0 { + let debt = (1.0 - state.available_requests).max(0.0); + (debt / (bucket.policy.requests_per_minute as f64 / 60.0)).ceil() as u64 + } else { + 0 + }; + return Err(token_retry.max(request_retry).max(1)); + } + + if bucket.policy.tokens_per_minute > 0 { + state.available_tokens -= need_tokens; + } + if bucket.policy.requests_per_minute > 0 { + state.available_requests -= 1.0; + } + + Ok(()) + } + + fn maybe_cleanup_stale_buckets(&self) { + let now = Instant::now(); + let now_ms = duration_millis(now.duration_since(self.started_at)); + let cleanup_interval_ms = duration_millis(CLEANUP_INTERVAL); + let last_cleanup_ms = self.last_cleanup_ms.load(Ordering::Relaxed); + + if now_ms.saturating_sub(last_cleanup_ms) < cleanup_interval_ms { + return; + } + + if self + .last_cleanup_ms + .compare_exchange(last_cleanup_ms, now_ms, Ordering::AcqRel, Ordering::Relaxed) + .is_err() + { + return; + } + + self.cleanup_stale_buckets_at(now); + } + + fn cleanup_stale_buckets_at(&self, now: Instant) { + let stale_keys: Vec = self + .buckets + .iter() + .filter_map(|entry| { + let state = entry.value().state.lock(); + let ttl = stale_bucket_ttl(&entry.value().policy); + (now.duration_since(state.last_refill) > ttl).then(|| entry.key().clone()) + }) + .collect(); + + for key in stale_keys { + let _ = self.buckets.remove(&key); + } + } + + #[cfg(test)] + fn force_cleanup_stale_buckets(&self) { + self.cleanup_stale_buckets_at(Instant::now()); + } +} + +fn duration_millis(duration: Duration) -> u64 { + duration.as_millis().try_into().unwrap_or(u64::MAX) +} + +fn stale_bucket_ttl(policy: &TenantTokenPolicy) -> Duration { + let token_window_secs = (policy.tokens_per_minute > 0).then_some(60); + let request_window_secs = (policy.requests_per_minute > 0).then_some(60); + let refill_window_secs = token_window_secs + .into_iter() + .chain(request_window_secs) + .max() + .unwrap_or(60); + Duration::from_secs(u64::from(refill_window_secs * STALE_BUCKET_TTL_MULTIPLIER)) +} + +#[cfg(test)] +mod tests { + use std::{ + collections::HashMap, + thread, + time::{Duration, Instant}, + }; + + use super::{ + LocalTokenRateLimiter, STALE_BUCKET_TTL_MULTIPLIER, TERMINAL_REJECTION_RETRY_AFTER_SECS, + }; + use crate::rate_limit::types::{MultiTenantRateLimitConfig, TenantTokenPolicy}; + + fn limiter(policy: TenantTokenPolicy) -> LocalTokenRateLimiter { + LocalTokenRateLimiter::new(MultiTenantRateLimitConfig { + enabled: true, + default_tokens_per_minute: 0, + default_requests_per_minute: 0, + tenants: HashMap::from([("tenant-a".to_string(), policy)]), + }) + } + + #[test] + fn allows_unknown_tenant_without_policy() { + let limiter = LocalTokenRateLimiter::new(MultiTenantRateLimitConfig { + enabled: true, + default_tokens_per_minute: 0, + default_requests_per_minute: 0, + tenants: HashMap::new(), + }); + + assert!(limiter.check_and_consume("tenant-a", 10).is_ok()); + } + + #[test] + fn enforces_request_limit_per_tenant() { + let limiter = limiter(TenantTokenPolicy { + tokens_per_minute: 0, + requests_per_minute: 1, + }); + + assert!(limiter.check_and_consume("tenant-a", 0).is_ok()); + let retry_after = limiter + .check_and_consume("tenant-a", 0) + .expect_err("second request should be rejected"); + + assert_eq!(retry_after, 60); + } + + #[test] + fn enforces_token_limit_per_tenant() { + let limiter = limiter(TenantTokenPolicy { + tokens_per_minute: 5, + requests_per_minute: 0, + }); + + assert!(limiter.check_and_consume("tenant-a", 3).is_ok()); + let retry_after = limiter + .check_and_consume("tenant-a", 3) + .expect_err("second token-heavy request should be rejected"); + + assert_eq!(retry_after, 12); + } + + #[test] + fn rejects_requests_larger_than_token_bucket_capacity_as_terminal() { + let limiter = limiter(TenantTokenPolicy { + tokens_per_minute: 5, + requests_per_minute: 0, + }); + + let retry_after = limiter + .check_and_consume("tenant-a", 6) + .expect_err("request larger than full bucket capacity should be terminal"); + + assert_eq!(retry_after, TERMINAL_REJECTION_RETRY_AFTER_SECS); + } + + #[test] + fn tracks_tenants_independently() { + let limiter = LocalTokenRateLimiter::new(MultiTenantRateLimitConfig { + enabled: true, + default_tokens_per_minute: 0, + default_requests_per_minute: 0, + tenants: HashMap::from([ + ( + "tenant-a".to_string(), + TenantTokenPolicy { + tokens_per_minute: 0, + requests_per_minute: 1, + }, + ), + ( + "tenant-b".to_string(), + TenantTokenPolicy { + tokens_per_minute: 0, + requests_per_minute: 1, + }, + ), + ]), + }); + + assert!(limiter.check_and_consume("tenant-a", 0).is_ok()); + assert!(limiter.check_and_consume("tenant-b", 0).is_ok()); + assert!(limiter.check_and_consume("tenant-a", 0).is_err()); + assert!(limiter.check_and_consume("tenant-b", 0).is_err()); + } + + #[test] + fn evicts_stale_buckets_during_periodic_cleanup() { + let limiter = LocalTokenRateLimiter::new(MultiTenantRateLimitConfig { + enabled: true, + default_tokens_per_minute: 0, + default_requests_per_minute: 0, + tenants: HashMap::from([ + ( + "tenant-a".to_string(), + TenantTokenPolicy { + tokens_per_minute: 60, + requests_per_minute: 0, + }, + ), + ( + "tenant-b".to_string(), + TenantTokenPolicy { + tokens_per_minute: 60, + requests_per_minute: 0, + }, + ), + ]), + }); + + assert!(limiter.check_and_consume("tenant-a", 1).is_ok()); + { + let bucket = limiter + .buckets + .get("tenant-a") + .expect("bucket should exist"); + let mut state = bucket.state.lock(); + state.last_refill = Instant::now() + - Duration::from_secs(u64::from(60 * STALE_BUCKET_TTL_MULTIPLIER + 1)); + } + limiter.force_cleanup_stale_buckets(); + + assert!(limiter.buckets.get("tenant-a").is_none()); + } + + #[test] + fn refills_capacity_over_time() { + let limiter = limiter(TenantTokenPolicy { + tokens_per_minute: 60, + requests_per_minute: 60, + }); + + assert!(limiter.check_and_consume("tenant-a", 60).is_ok()); + let retry_after = limiter + .check_and_consume("tenant-a", 1) + .expect_err("bucket should be empty immediately after consuming full capacity"); + assert_eq!(retry_after, 1); + + thread::sleep(Duration::from_millis(1100)); + + assert!(limiter.check_and_consume("tenant-a", 1).is_ok()); + } +} diff --git a/model_gateway/src/rate_limit/mod.rs b/model_gateway/src/rate_limit/mod.rs new file mode 100644 index 0000000000..5d92c03ec6 --- /dev/null +++ b/model_gateway/src/rate_limit/mod.rs @@ -0,0 +1,7 @@ +pub mod error; +pub mod local; +pub mod types; + +pub use error::rate_limit_exceeded_response; +pub use local::LocalTokenRateLimiter; +pub use types::{MultiTenantRateLimitConfig, TenantTokenPolicy}; diff --git a/model_gateway/src/rate_limit/types.rs b/model_gateway/src/rate_limit/types.rs new file mode 100644 index 0000000000..888350d30d --- /dev/null +++ b/model_gateway/src/rate_limit/types.rs @@ -0,0 +1,112 @@ +use std::collections::HashMap; + +use serde::{Deserialize, Serialize}; + +#[derive(Debug, Clone, Serialize, Deserialize, Default)] +#[serde(default)] +pub struct MultiTenantRateLimitConfig { + pub enabled: bool, + pub default_tokens_per_minute: u32, + pub default_requests_per_minute: u32, + pub tenants: HashMap, +} + +#[derive(Debug, Clone, Serialize, Deserialize, Default)] +#[serde(default)] +pub struct TenantTokenPolicy { + pub tokens_per_minute: u32, + pub requests_per_minute: u32, +} + +impl MultiTenantRateLimitConfig { + #[must_use] + pub fn policy_for(&self, tenant_key: &str) -> Option { + if !self.enabled { + return None; + } + + self.tenants.get(tenant_key).cloned().or_else(|| { + (self.default_tokens_per_minute > 0 || self.default_requests_per_minute > 0).then_some( + TenantTokenPolicy { + tokens_per_minute: self.default_tokens_per_minute, + requests_per_minute: self.default_requests_per_minute, + }, + ) + }) + } +} + +#[cfg(test)] +mod tests { + use std::collections::HashMap; + + use super::{MultiTenantRateLimitConfig, TenantTokenPolicy}; + + #[test] + fn policy_for_returns_none_when_disabled() { + let config = MultiTenantRateLimitConfig { + enabled: false, + default_tokens_per_minute: 100, + default_requests_per_minute: 10, + tenants: HashMap::from([( + "tenant-a".to_string(), + TenantTokenPolicy { + tokens_per_minute: 50, + requests_per_minute: 5, + }, + )]), + }; + + assert!(config.policy_for("tenant-a").is_none()); + assert!(config.policy_for("missing").is_none()); + } + + #[test] + fn policy_for_prefers_tenant_override() { + let tenant_policy = TenantTokenPolicy { + tokens_per_minute: 50, + requests_per_minute: 5, + }; + let config = MultiTenantRateLimitConfig { + enabled: true, + default_tokens_per_minute: 100, + default_requests_per_minute: 10, + tenants: HashMap::from([("tenant-a".to_string(), tenant_policy.clone())]), + }; + + let policy = config.policy_for("tenant-a").expect("tenant policy"); + + assert_eq!(policy.tokens_per_minute, tenant_policy.tokens_per_minute); + assert_eq!( + policy.requests_per_minute, + tenant_policy.requests_per_minute + ); + } + + #[test] + fn policy_for_uses_default_when_tenant_missing_and_defaults_enabled() { + let config = MultiTenantRateLimitConfig { + enabled: true, + default_tokens_per_minute: 100, + default_requests_per_minute: 10, + tenants: HashMap::new(), + }; + + let policy = config.policy_for("missing").expect("default policy"); + + assert_eq!(policy.tokens_per_minute, 100); + assert_eq!(policy.requests_per_minute, 10); + } + + #[test] + fn policy_for_returns_none_when_no_matching_tenant_or_defaults() { + let config = MultiTenantRateLimitConfig { + enabled: true, + default_tokens_per_minute: 0, + default_requests_per_minute: 0, + tenants: HashMap::new(), + }; + + assert!(config.policy_for("missing").is_none()); + } +} diff --git a/model_gateway/src/routers/common/mod.rs b/model_gateway/src/routers/common/mod.rs index 1cd09ea658..53f5c9f50a 100644 --- a/model_gateway/src/routers/common/mod.rs +++ b/model_gateway/src/routers/common/mod.rs @@ -17,6 +17,8 @@ //! used by every router for transport-level retries. Has zero //! coupling to the `Worker` trait — it lived in `worker/` for //! historical reasons before this extraction. +//! - [`token_count`] — shared request token estimation for routing and +//! rate-limiting paths. //! - [`sse`] — shared SSE codec (encoder + decoder) for streaming //! responses to clients and parsing upstream SSE byte streams @@ -26,4 +28,5 @@ pub mod openai_bridge; pub mod persistence_utils; pub mod retry; pub mod sse; +pub mod token_count; pub mod worker_selection; diff --git a/model_gateway/src/routers/common/token_count.rs b/model_gateway/src/routers/common/token_count.rs new file mode 100644 index 0000000000..294391baf9 --- /dev/null +++ b/model_gateway/src/routers/common/token_count.rs @@ -0,0 +1,28 @@ +use openai_protocol::common::GenerationRequest; + +/// Estimate request token count for routing and rate limiting. +/// +/// Falls back to a conservative character-based estimate when no tokenizer is +/// registered for the model or tokenization fails. +pub fn count_tokens( + tokenizer_registry: &llm_tokenizer::registry::TokenizerRegistry, + body: &T, + model_id: &str, +) -> u32 { + let text = body.extract_text_for_routing(); + if text.is_empty() { + return 1; + } + + let fallback_estimate = || ((text.chars().count() as u32) / 4).max(1); + + let Some(tokenizer) = tokenizer_registry.get(model_id) else { + return fallback_estimate(); + }; + + tokenizer + .encode(&text, false) + .map(|encoding| encoding.token_ids().len() as u32) + .unwrap_or_else(|_| fallback_estimate()) + .max(1) +} diff --git a/model_gateway/src/routers/grpc/router.rs b/model_gateway/src/routers/grpc/router.rs index 0578faeebd..37c719fa50 100644 --- a/model_gateway/src/routers/grpc/router.rs +++ b/model_gateway/src/routers/grpc/router.rs @@ -6,9 +6,9 @@ use axum::{ response::{IntoResponse, Response}, }; use openai_protocol::{ - chat::ChatCompletionRequest, classify::ClassifyRequest, completion::CompletionRequest, - embedding::EmbeddingRequest, generate::GenerateRequest, messages::CreateMessageRequest, - responses::ResponsesRequest, + chat::ChatCompletionRequest, classify::ClassifyRequest, common::GenerationRequest, + completion::CompletionRequest, embedding::EmbeddingRequest, generate::GenerateRequest, + messages::CreateMessageRequest, responses::ResponsesRequest, }; use tracing::debug; @@ -27,8 +27,10 @@ use crate::{ config::types::RetryConfig, middleware::TenantRequestMeta, observability::metrics::{metrics_labels, Metrics}, + rate_limit::{rate_limit_exceeded_response, LocalTokenRateLimiter}, routers::{ common::retry::{is_retryable_status, RetryExecutor}, + common::token_count::count_tokens, RouterTrait, }, worker::WorkerRegistry, @@ -48,6 +50,7 @@ pub struct GrpcRouter { responses_context: ResponsesContext, harmony_responses_context: ResponsesContext, retry_config: RetryConfig, + app_token_rate_limiter: Option>, } impl GrpcRouter { @@ -170,9 +173,25 @@ impl GrpcRouter { responses_context, harmony_responses_context, retry_config: ctx.router_config.effective_retry_config(), + app_token_rate_limiter: ctx.token_rate_limiter.clone(), }) } + fn enforce_token_rate_limit( + &self, + tenant_meta: &TenantRequestMeta, + body: &T, + model_id: &str, + ) -> Option { + let limiter = self.app_token_rate_limiter.as_ref()?; + let estimated_tokens = + count_tokens(&self.shared_components.tokenizer_registry, body, model_id); + limiter + .check_and_consume(tenant_meta.tenant_key().as_str(), estimated_tokens) + .err() + .map(rate_limit_exceeded_response) + } + /// Main route_chat implementation async fn route_chat_impl( &self, @@ -181,6 +200,10 @@ impl GrpcRouter { body: &ChatCompletionRequest, model_id: &str, ) -> Response { + if let Some(response) = self.enforce_token_rate_limit(tenant_meta, body, model_id) { + return response; + } + // Choose Harmony pipeline if workers indicate Harmony (checks architectures, hf_model_type) let is_harmony = HarmonyDetector::is_harmony_model_in_registry(&self.worker_registry, &body.model); @@ -253,6 +276,10 @@ impl GrpcRouter { body: &GenerateRequest, model_id: &str, ) -> Response { + if let Some(response) = self.enforce_token_rate_limit(tenant_meta, body, model_id) { + return response; + } + debug!("Processing generate request for model: {}", model_id); // Clone values needed for retry closure @@ -315,6 +342,10 @@ impl GrpcRouter { body: &ResponsesRequest, model_id: &str, ) -> Response { + if let Some(response) = self.enforce_token_rate_limit(tenant_meta, body, model_id) { + return response; + } + // 0. Fast worker validation (fail-fast before expensive operations) if let Some(error_response) = validate_worker_availability(&self.worker_registry, model_id) { @@ -374,6 +405,10 @@ impl GrpcRouter { body: &EmbeddingRequest, model_id: &str, ) -> Response { + if let Some(response) = self.enforce_token_rate_limit(tenant_meta, body, model_id) { + return response; + } + debug!("Processing embedding request for model: {}", model_id); self.embedding_pipeline @@ -451,6 +486,10 @@ impl GrpcRouter { body: &CompletionRequest, model_id: &str, ) -> Response { + if let Some(response) = self.enforce_token_rate_limit(tenant_meta, body, model_id) { + return response; + } + debug!("Processing completion request for model: {}", model_id); let request = Arc::new(body.clone()); @@ -511,6 +550,10 @@ impl GrpcRouter { body: &ClassifyRequest, model_id: &str, ) -> Response { + if let Some(response) = self.enforce_token_rate_limit(tenant_meta, body, model_id) { + return response; + } + debug!("Processing classify request for model: {}", model_id); self.classify_pipeline diff --git a/model_gateway/src/routers/http/pd_router.rs b/model_gateway/src/routers/http/pd_router.rs index ea9e72ce91..fccc108257 100644 --- a/model_gateway/src/routers/http/pd_router.rs +++ b/model_gateway/src/routers/http/pd_router.rs @@ -11,7 +11,7 @@ use futures_util::StreamExt; use memchr::memmem; use openai_protocol::{ chat::{ChatCompletionRequest, ChatMessage, MessageContent}, - common::{InputIds, StringOrArray}, + common::{GenerationRequest, InputIds, StringOrArray}, completion::CompletionRequest, generate::GenerateRequest, rerank::RerankRequest, @@ -32,11 +32,13 @@ use crate::{ otel_trace::inject_trace_context_http, }, policies::{LoadBalancingPolicy, PolicyRegistry, SelectWorkerInfo}, + rate_limit::{rate_limit_exceeded_response, LocalTokenRateLimiter}, routers::{ common::{ header_utils, retry::{is_retryable_status, RetryExecutor}, sse::SseEncoder, + token_count::count_tokens, }, error, grpc::utils::{error_type_from_status, route_to_endpoint}, @@ -45,13 +47,31 @@ use crate::{ worker::{HashRing, Worker, WorkerLoadGuard, WorkerRegistry, WorkerType, UNKNOWN_MODEL_ID}, }; -#[derive(Debug)] pub struct PDRouter { pub worker_registry: Arc, pub policy_registry: Arc, pub client: Client, pub retry_config: RetryConfig, pub api_key: Option, + tokenizer_registry: Arc, + app_token_rate_limiter: Option>, +} + +impl std::fmt::Debug for PDRouter { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("PDRouter") + .field("worker_registry", &self.worker_registry) + .field("policy_registry", &self.policy_registry) + .field("client", &self.client) + .field("retry_config", &self.retry_config) + .field("api_key", &self.api_key.as_ref().map(|_| "configured")) + .field("tokenizer_registry", &"configured") + .field( + "app_token_rate_limiter", + &self.app_token_rate_limiter.as_ref().map(|_| "configured"), + ) + .finish() + } } #[derive(Clone)] @@ -66,6 +86,20 @@ struct PDRequestContext<'a> { } impl PDRouter { + fn enforce_token_rate_limit( + &self, + tenant_meta: &TenantRequestMeta, + body: &T, + model_id: &str, + ) -> Option { + let limiter = self.app_token_rate_limiter.as_ref()?; + let estimated_tokens = count_tokens(&self.tokenizer_registry, body, model_id); + limiter + .check_and_consume(tenant_meta.tenant_key().as_str(), estimated_tokens) + .err() + .map(rate_limit_exceeded_response) + } + async fn proxy_to_first_prefill_worker( &self, endpoint: &str, @@ -168,6 +202,8 @@ impl PDRouter { client: ctx.client.clone(), retry_config: ctx.router_config.effective_retry_config(), api_key: ctx.router_config.api_key.clone(), + tokenizer_registry: ctx.tokenizer_registry.clone(), + app_token_rate_limiter: ctx.token_rate_limiter.clone(), }) } @@ -1322,10 +1358,13 @@ impl RouterTrait for PDRouter { async fn route_generate( &self, headers: Option<&HeaderMap>, - _tenant_meta: &TenantRequestMeta, + tenant_meta: &TenantRequestMeta, body: &GenerateRequest, model_id: &str, ) -> Response { + if let Some(response) = self.enforce_token_rate_limit(tenant_meta, body, model_id) { + return response; + } let is_stream = body.stream; let return_logprob = body.return_logprob.unwrap_or(false); @@ -1353,10 +1392,13 @@ impl RouterTrait for PDRouter { async fn route_chat( &self, headers: Option<&HeaderMap>, - _tenant_meta: &TenantRequestMeta, + tenant_meta: &TenantRequestMeta, body: &ChatCompletionRequest, model_id: &str, ) -> Response { + if let Some(response) = self.enforce_token_rate_limit(tenant_meta, body, model_id) { + return response; + } let is_stream = body.stream; let return_logprob = body.logprobs; @@ -1396,10 +1438,13 @@ impl RouterTrait for PDRouter { async fn route_completion( &self, headers: Option<&HeaderMap>, - _tenant_meta: &TenantRequestMeta, + tenant_meta: &TenantRequestMeta, body: &CompletionRequest, model_id: &str, ) -> Response { + if let Some(response) = self.enforce_token_rate_limit(tenant_meta, body, model_id) { + return response; + } let is_stream = body.stream; let return_logprob = body.logprobs.is_some(); @@ -1431,10 +1476,13 @@ impl RouterTrait for PDRouter { async fn route_rerank( &self, headers: Option<&HeaderMap>, - _tenant_meta: &TenantRequestMeta, + tenant_meta: &TenantRequestMeta, body: &RerankRequest, model_id: &str, ) -> Response { + if let Some(response) = self.enforce_token_rate_limit(tenant_meta, body, model_id) { + return response; + } // Extract text for cache-aware routing let req_text = if self.policies_need_request_text() { Some(body.query.clone()) @@ -1478,6 +1526,8 @@ mod tests { client: Client::new(), retry_config: RetryConfig::default(), api_key: Some("test_api_key".to_string()), + tokenizer_registry: Arc::new(llm_tokenizer::registry::TokenizerRegistry::new()), + app_token_rate_limiter: None, } } diff --git a/model_gateway/src/routers/http/router.rs b/model_gateway/src/routers/http/router.rs index 0799467f1e..d1ae1e41d3 100644 --- a/model_gateway/src/routers/http/router.rs +++ b/model_gateway/src/routers/http/router.rs @@ -38,10 +38,12 @@ use crate::{ otel_trace::inject_trace_context_http, }, policies::{PolicyRegistry, SelectWorkerInfo}, + rate_limit::{rate_limit_exceeded_response, LocalTokenRateLimiter}, routers::{ common::{ header_utils, retry::{is_retryable_status, RetryExecutor}, + token_count::count_tokens, }, error::{self, extract_error_code_from_response}, grpc::utils::{error_type_from_status, route_to_endpoint}, @@ -56,6 +58,8 @@ pub struct Router { policy_registry: Arc, client: Client, retry_config: RetryConfig, + tokenizer_registry: Arc, + app_token_rate_limiter: Option>, } impl std::fmt::Debug for Router { @@ -65,11 +69,29 @@ impl std::fmt::Debug for Router { .field("policy_registry", &self.policy_registry) .field("client", &self.client) .field("retry_config", &self.retry_config) + .field("tokenizer_registry", &"configured") + .field( + "app_token_rate_limiter", + &self.app_token_rate_limiter.as_ref().map(|_| "configured"), + ) .finish() } } impl Router { + fn enforce_token_rate_limit( + &self, + tenant_meta: &TenantRequestMeta, + body: &T, + model_id: &str, + ) -> Option { + let limiter = self.app_token_rate_limiter.as_ref()?; + let estimated_tokens = count_tokens(&self.tokenizer_registry, body, model_id); + limiter + .check_and_consume(tenant_meta.tenant_key().as_str(), estimated_tokens) + .err() + .map(rate_limit_exceeded_response) + } /// Create a new router with injected policy and client #[expect( clippy::unused_async, @@ -81,6 +103,8 @@ impl Router { policy_registry: ctx.policy_registry.clone(), client: ctx.client.clone(), retry_config: ctx.router_config.effective_retry_config(), + tokenizer_registry: ctx.tokenizer_registry.clone(), + app_token_rate_limiter: ctx.token_rate_limiter.clone(), }) } @@ -1107,10 +1131,13 @@ impl RouterTrait for Router { async fn route_generate( &self, headers: Option<&HeaderMap>, - _tenant_meta: &TenantRequestMeta, + tenant_meta: &TenantRequestMeta, body: &GenerateRequest, model_id: &str, ) -> Response { + if let Some(response) = self.enforce_token_rate_limit(tenant_meta, body, model_id) { + return response; + } self.route_typed_request(headers, body, "/generate", model_id) .await } @@ -1118,10 +1145,13 @@ impl RouterTrait for Router { async fn route_chat( &self, headers: Option<&HeaderMap>, - _tenant_meta: &TenantRequestMeta, + tenant_meta: &TenantRequestMeta, body: &ChatCompletionRequest, model_id: &str, ) -> Response { + if let Some(response) = self.enforce_token_rate_limit(tenant_meta, body, model_id) { + return response; + } self.route_typed_request(headers, body, "/v1/chat/completions", model_id) .await } @@ -1140,10 +1170,13 @@ impl RouterTrait for Router { async fn route_completion( &self, headers: Option<&HeaderMap>, - _tenant_meta: &TenantRequestMeta, + tenant_meta: &TenantRequestMeta, body: &CompletionRequest, model_id: &str, ) -> Response { + if let Some(response) = self.enforce_token_rate_limit(tenant_meta, body, model_id) { + return response; + } self.route_typed_request(headers, body, "/v1/completions", model_id) .await } @@ -1151,10 +1184,13 @@ impl RouterTrait for Router { async fn route_responses( &self, headers: Option<&HeaderMap>, - _tenant_meta: &TenantRequestMeta, + tenant_meta: &TenantRequestMeta, body: &ResponsesRequest, model_id: &str, ) -> Response { + if let Some(response) = self.enforce_token_rate_limit(tenant_meta, body, model_id) { + return response; + } self.route_typed_request(headers, body, "/v1/responses", model_id) .await } @@ -1167,10 +1203,13 @@ impl RouterTrait for Router { async fn route_embeddings( &self, headers: Option<&HeaderMap>, - _tenant_meta: &TenantRequestMeta, + tenant_meta: &TenantRequestMeta, body: &EmbeddingRequest, model_id: &str, ) -> Response { + if let Some(response) = self.enforce_token_rate_limit(tenant_meta, body, model_id) { + return response; + } self.route_typed_request(headers, body, "/v1/embeddings", model_id) .await } @@ -1178,10 +1217,13 @@ impl RouterTrait for Router { async fn route_classify( &self, headers: Option<&HeaderMap>, - _tenant_meta: &TenantRequestMeta, + tenant_meta: &TenantRequestMeta, body: &ClassifyRequest, model_id: &str, ) -> Response { + if let Some(response) = self.enforce_token_rate_limit(tenant_meta, body, model_id) { + return response; + } self.route_typed_request(headers, body, "/v1/classify", model_id) .await } @@ -1189,11 +1231,14 @@ impl RouterTrait for Router { async fn route_audio_transcriptions( &self, headers: Option<&HeaderMap>, - _tenant_meta: &TenantRequestMeta, + tenant_meta: &TenantRequestMeta, body: &TranscriptionRequest, audio: AudioFile, model_id: &str, ) -> Response { + if let Some(response) = self.enforce_token_rate_limit(tenant_meta, body, model_id) { + return response; + } self.route_multipart_transcription( headers, body, @@ -1207,10 +1252,13 @@ impl RouterTrait for Router { async fn route_rerank( &self, headers: Option<&HeaderMap>, - _tenant_meta: &TenantRequestMeta, + tenant_meta: &TenantRequestMeta, body: &RerankRequest, model_id: &str, ) -> Response { + if let Some(response) = self.enforce_token_rate_limit(tenant_meta, body, model_id) { + return response; + } let response = self .route_typed_request(headers, body, "/v1/rerank", model_id) .await; @@ -1271,6 +1319,8 @@ mod tests { policy_registry, client: Client::new(), retry_config: RetryConfig::default(), + tokenizer_registry: Arc::new(llm_tokenizer::registry::TokenizerRegistry::new()), + app_token_rate_limiter: None, } } diff --git a/model_gateway/src/service_discovery.rs b/model_gateway/src/service_discovery.rs index f92274d8a7..b3fd5d03ed 100644 --- a/model_gateway/src/service_discovery.rs +++ b/model_gateway/src/service_discovery.rs @@ -1285,6 +1285,7 @@ mod tests { client: reqwest::Client::new(), router_config: router_config.clone(), rate_limiter: Some(Arc::new(TokenBucket::new(1000, 1000))), + token_rate_limiter: None, worker_registry: worker_registry.clone(), policy_registry: Arc::new(crate::policies::PolicyRegistry::new( router_config.policy.clone(), diff --git a/model_gateway/src/workflow/steps/local/drain_workers.rs b/model_gateway/src/workflow/steps/local/drain_workers.rs index 78752fbdd4..5dd4cabea3 100644 --- a/model_gateway/src/workflow/steps/local/drain_workers.rs +++ b/model_gateway/src/workflow/steps/local/drain_workers.rs @@ -149,6 +149,7 @@ mod tests { client: reqwest::Client::new(), router_config: router_config.clone(), rate_limiter: Some(Arc::new(TokenBucket::new(1000, 1000))), + token_rate_limiter: None, worker_registry: Arc::clone(®istry), policy_registry: Arc::new(crate::policies::PolicyRegistry::new( router_config.policy.clone(),