diff --git a/model_gateway/src/app_context.rs b/model_gateway/src/app_context.rs index d3877c4a2..963a26df6 100644 --- a/model_gateway/src/app_context.rs +++ b/model_gateway/src/app_context.rs @@ -22,6 +22,7 @@ use crate::{ rate_limit::RateLimitManager, routers::{ common::{openai_bridge::FormatRegistry, realtime::RealtimeRegistry}, + grpc::adaptive_admission::AdaptiveAdmissionController, grpc::multimodal::MultimodalConfigRegistry, router_manager::RouterManager, }, @@ -65,6 +66,7 @@ pub struct AppContext { pub conversation_storage: Arc, pub conversation_item_storage: Arc, pub worker_monitor: Option>, + pub(crate) adaptive_admission: Option>, pub configured_reasoning_parser: Option, pub configured_tool_parser: Option, pub worker_job_queue: Arc>>, @@ -338,6 +340,20 @@ impl AppContextBuilder { let worker_job_queue = self .worker_job_queue .ok_or(AppContextBuildError::MissingField("worker_job_queue"))?; + let worker_monitor = self.worker_monitor; + let adaptive_admission = + if router_config.adaptive_admission.mode == crate::config::AdaptiveAdmissionMode::Off { + None + } else { + let controller = AdaptiveAdmissionController::new( + router_config.adaptive_admission.clone(), + worker_registry.clone(), + ); + if let Some(monitor) = &worker_monitor { + controller.start_load_updates(monitor.subscribe()); + } + Some(controller) + }; // Create WorkerService from the already-built components let worker_service = Arc::new(WorkerService::new( @@ -373,7 +389,8 @@ impl AppContextBuilder { conversation_item_storage: self.conversation_item_storage.ok_or( AppContextBuildError::MissingField("conversation_item_storage"), )?, - worker_monitor: self.worker_monitor, + worker_monitor, + adaptive_admission, configured_reasoning_parser, configured_tool_parser, worker_job_queue, @@ -605,7 +622,8 @@ impl AppContextBuilder { .clone(), client.clone(), config.load_monitor_interval_secs, - config.engine_metrics, + config.engine_metrics + || config.adaptive_admission.mode != crate::config::AdaptiveAdmissionMode::Off, ))); Ok(self) } diff --git a/model_gateway/src/config/builder.rs b/model_gateway/src/config/builder.rs index d107abf98..4641c30f7 100644 --- a/model_gateway/src/config/builder.rs +++ b/model_gateway/src/config/builder.rs @@ -289,6 +289,11 @@ impl RouterConfigBuilder { self } + pub fn adaptive_admission(mut self, config: super::types::AdaptiveAdmissionConfig) -> Self { + self.config.adaptive_admission = config; + self + } + // ==================== Tenant Rate Limit ==================== pub fn tenant_rate_limit_enabled(mut self, enabled: bool) -> Self { diff --git a/model_gateway/src/config/types.rs b/model_gateway/src/config/types.rs index 7d5c5677c..c79aab00c 100644 --- a/model_gateway/src/config/types.rs +++ b/model_gateway/src/config/types.rs @@ -11,6 +11,100 @@ pub use smg_data_connector::{ use super::{validation::ConfigValidator, ConfigResult}; use crate::{tenant::DEFAULT_TENANT_HEADER_NAME, worker::ConnectionMode}; +/// Runtime mode for predictive, token-work admission. +/// +/// `shadow` learns from completed requests and records the decision that would +/// have been made, but never delays or rejects a request. `enforce` is kept as +/// an explicit operator action so a newly deployed estimator cannot change +/// serving behavior before its calibration is measured. +#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum AdaptiveAdmissionMode { + #[default] + Off, + Shadow, + Enforce, +} + +impl std::str::FromStr for AdaptiveAdmissionMode { + type Err = String; + + fn from_str(value: &str) -> Result { + match value.trim().to_ascii_lowercase().as_str() { + "off" => Ok(Self::Off), + "shadow" => Ok(Self::Shadow), + "enforce" => Ok(Self::Enforce), + _ => Err(format!( + "adaptive admission mode must be one of off, shadow, enforce; got {value:?}" + )), + } + } +} + +fn default_adaptive_work_horizon_secs() -> f64 { + 30.0 +} + +fn default_adaptive_estimator_half_life_secs() -> f64 { + 900.0 +} + +fn default_adaptive_prior_observations() -> f64 { + 20.0 +} + +fn default_adaptive_max_segments() -> usize { + 50_000 +} + +fn default_adaptive_min_load_coverage() -> f64 { + 0.8 +} + +fn default_adaptive_cold_start_output_tokens() -> u32 { + 4096 +} + +/// Predictive token-work admission settings. +/// +/// The work horizon is an operator-facing latency objective rather than a +/// request-concurrency guess: recent per-replica generation capacity multiplied +/// by healthy replica count and this horizon gives the maximum predicted +/// outstanding decode work. Router reservations are reconciled with engine +/// running and waiting counts. Missing or insufficient telemetry fails open to +/// the existing priority scheduler. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub struct AdaptiveAdmissionConfig { + #[serde(default)] + pub mode: AdaptiveAdmissionMode, + #[serde(default = "default_adaptive_work_horizon_secs")] + pub work_horizon_secs: f64, + #[serde(default = "default_adaptive_estimator_half_life_secs")] + pub estimator_half_life_secs: f64, + #[serde(default = "default_adaptive_prior_observations")] + pub prior_observations: f64, + #[serde(default = "default_adaptive_max_segments")] + pub max_segments: usize, + #[serde(default = "default_adaptive_min_load_coverage")] + pub min_load_coverage: f64, + #[serde(default = "default_adaptive_cold_start_output_tokens")] + pub cold_start_output_tokens: u32, +} + +impl Default for AdaptiveAdmissionConfig { + fn default() -> Self { + Self { + mode: AdaptiveAdmissionMode::Off, + work_horizon_secs: default_adaptive_work_horizon_secs(), + estimator_half_life_secs: default_adaptive_estimator_half_life_secs(), + prior_observations: default_adaptive_prior_observations(), + max_segments: default_adaptive_max_segments(), + min_load_coverage: default_adaptive_min_load_coverage(), + cold_start_output_tokens: default_adaptive_cold_start_output_tokens(), + } + } +} + /// Main router configuration #[derive(Debug, Clone, Serialize, Deserialize)] pub struct RouterConfig { @@ -104,6 +198,10 @@ pub struct RouterConfig { /// by inflight; the remainder bucket under `tenant="other"`). #[serde(default = "default_priority_scheduler_tenant_metric_top_n")] pub priority_scheduler_tenant_metric_top_n: u32, + /// Optional predictive token-work admission. Off by default. Shadow mode + /// is behavior-preserving and is the required first deployment state. + #[serde(default)] + pub adaptive_admission: AdaptiveAdmissionConfig, /// Enable per-tenant LLM token/request rate limiting. When false /// (default), no rate limiter is constructed — zero behavior change /// for existing deployments. @@ -842,6 +940,7 @@ impl Default for RouterConfig { priority_scheduler_config: None, priority_scheduler_tenant_metric_top_n: default_priority_scheduler_tenant_metric_top_n( ), + adaptive_admission: AdaptiveAdmissionConfig::default(), tenant_rate_limit_enabled: false, tenant_rate_limit_config: None, cors_allowed_origins: vec![], diff --git a/model_gateway/src/config/validation.rs b/model_gateway/src/config/validation.rs index 9b0b1dac3..b5a3071e3 100755 --- a/model_gateway/src/config/validation.rs +++ b/model_gateway/src/config/validation.rs @@ -724,6 +724,55 @@ impl ConfigValidator { }); } + let adaptive = &config.adaptive_admission; + if !adaptive.work_horizon_secs.is_finite() || adaptive.work_horizon_secs <= 0.0 { + return Err(ConfigError::InvalidValue { + field: "adaptive_admission.work_horizon_secs".to_string(), + value: adaptive.work_horizon_secs.to_string(), + reason: "Must be finite and > 0".to_string(), + }); + } + if !adaptive.estimator_half_life_secs.is_finite() + || adaptive.estimator_half_life_secs <= 0.0 + { + return Err(ConfigError::InvalidValue { + field: "adaptive_admission.estimator_half_life_secs".to_string(), + value: adaptive.estimator_half_life_secs.to_string(), + reason: "Must be finite and > 0".to_string(), + }); + } + if !adaptive.prior_observations.is_finite() || adaptive.prior_observations < 0.0 { + return Err(ConfigError::InvalidValue { + field: "adaptive_admission.prior_observations".to_string(), + value: adaptive.prior_observations.to_string(), + reason: "Must be finite and >= 0".to_string(), + }); + } + if adaptive.max_segments < 4 { + return Err(ConfigError::InvalidValue { + field: "adaptive_admission.max_segments".to_string(), + value: adaptive.max_segments.to_string(), + reason: "Must be >= 4".to_string(), + }); + } + if !adaptive.min_load_coverage.is_finite() + || adaptive.min_load_coverage <= 0.0 + || adaptive.min_load_coverage > 1.0 + { + return Err(ConfigError::InvalidValue { + field: "adaptive_admission.min_load_coverage".to_string(), + value: adaptive.min_load_coverage.to_string(), + reason: "Must be finite and in (0, 1]".to_string(), + }); + } + if adaptive.cold_start_output_tokens == 0 { + return Err(ConfigError::InvalidValue { + field: "adaptive_admission.cold_start_output_tokens".to_string(), + value: adaptive.cold_start_output_tokens.to_string(), + reason: "Must be > 0".to_string(), + }); + } + Ok(()) } diff --git a/model_gateway/src/main.rs b/model_gateway/src/main.rs index 926ba171a..01b9507e6 100755 --- a/model_gateway/src/main.rs +++ b/model_gateway/src/main.rs @@ -5,11 +5,11 @@ use openai_protocol::worker::TransportMode; use rand::{distr::Alphanumeric, RngExt}; use smg::{ config::{ - validate_mesh_server_name, CircuitBreakerConfig, ConfigError, ConfigResult, - DiscoveryConfig, HealthCheckConfig, HistoryBackend, ManualAssignmentMode, MetricsConfig, - OracleConfig, PolicyConfig, PostgresConfig, RedisConfig, RetryConfig, RouterConfig, - RoutingKeyOverrideConfig, RoutingMode, SchemaConfig, TenantApiKeyEntry, - TokenizerCacheConfig, TraceConfig, + validate_mesh_server_name, AdaptiveAdmissionConfig, AdaptiveAdmissionMode, + CircuitBreakerConfig, ConfigError, ConfigResult, DiscoveryConfig, HealthCheckConfig, + HistoryBackend, ManualAssignmentMode, MetricsConfig, OracleConfig, PolicyConfig, + PostgresConfig, RedisConfig, RetryConfig, RouterConfig, RoutingKeyOverrideConfig, + RoutingMode, SchemaConfig, TenantApiKeyEntry, TokenizerCacheConfig, TraceConfig, }, observability::{ metrics::PrometheusConfig, @@ -97,6 +97,34 @@ impl std::fmt::Display for Backend { } } +#[derive(Copy, Clone, Debug, Default, Eq, PartialEq, ValueEnum)] +enum AdaptiveAdmissionCliMode { + #[default] + Off, + Shadow, + Enforce, +} + +impl std::fmt::Display for AdaptiveAdmissionCliMode { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.write_str(match self { + Self::Off => "off", + Self::Shadow => "shadow", + Self::Enforce => "enforce", + }) + } +} + +impl From for AdaptiveAdmissionMode { + fn from(value: AdaptiveAdmissionCliMode) -> Self { + match value { + AdaptiveAdmissionCliMode::Off => Self::Off, + AdaptiveAdmissionCliMode::Shadow => Self::Shadow, + AdaptiveAdmissionCliMode::Enforce => Self::Enforce, + } + } +} + #[derive(Parser, Debug)] #[command(name = "shepherd-model-gateway", alias = "smg", alias = "amg")] #[command(about = "Shepherd Model Gateway - High-performance inference gateway")] @@ -488,6 +516,42 @@ struct CliArgs { #[arg(long, default_value_t = 32, help_heading = "Priority Scheduler")] priority_scheduler_tenant_metric_top_n: u32, + // ==================== Adaptive Admission ==================== + /// Predictive token-work admission mode. Shadow mode learns and records + /// hypothetical decisions without delaying or rejecting requests. + #[arg( + long, + value_enum, + default_value_t = AdaptiveAdmissionCliMode::Off, + help_heading = "Adaptive Admission" + )] + adaptive_admission_mode: AdaptiveAdmissionCliMode, + + /// Maximum predicted outstanding decode-work horizon in seconds. + #[arg(long, default_value_t = 30.0, help_heading = "Adaptive Admission")] + adaptive_admission_work_horizon_secs: f64, + + /// Half-life for recency weighting of observed output lengths. + #[arg(long, default_value_t = 900.0, help_heading = "Adaptive Admission")] + adaptive_admission_estimator_half_life_secs: f64, + + /// Hierarchical shrinkage strength in effective observations. + #[arg(long, default_value_t = 20.0, help_heading = "Adaptive Admission")] + adaptive_admission_prior_observations: f64, + + /// Maximum in-memory predictor segments before stale-segment eviction. + #[arg(long, default_value_t = 50000, help_heading = "Adaptive Admission")] + adaptive_admission_max_segments: usize, + + /// Minimum fraction of healthy replicas with fresh load telemetry before + /// a token-work decision is considered usable. + #[arg(long, default_value_t = 0.8, help_heading = "Adaptive Admission")] + adaptive_admission_min_load_coverage: f64, + + /// Cold-start output-token prediction before a model has observations. + #[arg(long, default_value_t = 4096, help_heading = "Adaptive Admission")] + adaptive_admission_cold_start_output_tokens: u32, + // ==================== Tenant Rate Limit ==================== /// Enable per-tenant LLM token/request rate limiting. When unset /// (default), no rate limiter is constructed. @@ -1490,6 +1554,15 @@ impl CliArgs { .priority_scheduler_default_max_class(self.priority_scheduler_default_max_class.clone()) .priority_scheduler_config(self.priority_scheduler_config.clone()) .priority_scheduler_tenant_metric_top_n(self.priority_scheduler_tenant_metric_top_n) + .adaptive_admission(AdaptiveAdmissionConfig { + mode: self.adaptive_admission_mode.into(), + work_horizon_secs: self.adaptive_admission_work_horizon_secs, + estimator_half_life_secs: self.adaptive_admission_estimator_half_life_secs, + prior_observations: self.adaptive_admission_prior_observations, + max_segments: self.adaptive_admission_max_segments, + min_load_coverage: self.adaptive_admission_min_load_coverage, + cold_start_output_tokens: self.adaptive_admission_cold_start_output_tokens, + }) .tenant_rate_limit_enabled(self.tenant_rate_limit_enabled) .tenant_rate_limit_config(self.tenant_rate_limit_config.clone()) .cors_allowed_origins(self.cors_allowed_origins.clone()) diff --git a/model_gateway/src/observability/metrics.rs b/model_gateway/src/observability/metrics.rs index 3857fb22b..e78ae1716 100644 --- a/model_gateway/src/observability/metrics.rs +++ b/model_gateway/src/observability/metrics.rs @@ -480,7 +480,9 @@ pub(crate) fn init_metrics() { // Priority scheduler metrics (no-op at scrape time unless the scheduler // is enabled and recording). use crate::middleware::scheduler::metrics as scheduler_metrics; + use crate::routers::grpc::adaptive_admission as adaptive_admission_metrics; scheduler_metrics::describe(); + adaptive_admission_metrics::describe_metrics(); } #[expect( diff --git a/model_gateway/src/routers/grpc/adaptive_admission.rs b/model_gateway/src/routers/grpc/adaptive_admission.rs new file mode 100644 index 000000000..b2c61dee0 --- /dev/null +++ b/model_gateway/src/routers/grpc/adaptive_admission.rs @@ -0,0 +1,960 @@ +//! Adaptive, token-work admission for the gRPC serving path. +//! +//! The existing priority scheduler remains the infrastructure safety layer. +//! This controller estimates request work after tokenization, learns output +//! length online, and compares predicted outstanding decode work with fresh +//! aggregate engine throughput. Shadow mode exercises the complete state +//! machine without delaying or rejecting traffic. + +use std::{ + collections::HashMap, + sync::{ + atomic::{AtomicBool, Ordering}, + Arc, Weak, + }, + time::Instant, +}; + +use metrics::{counter, describe_counter, describe_gauge, describe_histogram, gauge, histogram}; +use openai_protocol::worker::WorkerLoadResponse; +use parking_lot::Mutex; +use tokio::sync::watch; + +use crate::{ + config::{AdaptiveAdmissionConfig, AdaptiveAdmissionMode}, + observability::metrics::intern_string, + worker::WorkerRegistry, +}; + +const ADMISSION_PARTITION_LABEL: &str = "admission_partition"; + +const PREDICTIONS_TOTAL: &str = "smg_adaptive_admission_predictions_total"; +const PREDICTED_OUTPUT_TOKENS: &str = "smg_adaptive_admission_predicted_output_tokens"; +const OBSERVED_OUTPUT_TOKENS: &str = "smg_adaptive_admission_observed_output_tokens"; +const ABSOLUTE_ERROR_TOKENS: &str = "smg_adaptive_admission_absolute_error_tokens"; +const DECISIONS_TOTAL: &str = "smg_adaptive_admission_decisions_total"; +const OUTSTANDING_TOKENS: &str = "smg_adaptive_admission_outstanding_tokens"; +const ROUTER_OUTSTANDING_TOKENS: &str = "smg_adaptive_admission_router_outstanding_tokens"; +const ENGINE_ESTIMATED_TOKENS: &str = "smg_adaptive_admission_engine_estimated_tokens"; +const WORK_BUDGET_TOKENS: &str = "smg_adaptive_admission_work_budget_tokens"; +const DRAIN_SECONDS: &str = "smg_adaptive_admission_predicted_drain_seconds"; +const LOAD_COVERAGE: &str = "smg_adaptive_admission_load_coverage"; +const GEN_THROUGHPUT: &str = "smg_adaptive_admission_generation_tokens_per_second"; +const LEARNED_CAPACITY: &str = "smg_adaptive_admission_learned_capacity_tokens_per_second"; +const ENGINE_RUNNING: &str = "smg_adaptive_admission_engine_running_requests"; +const ENGINE_WAITING: &str = "smg_adaptive_admission_engine_waiting_requests"; +const ENGINE_WAITING_TOKENS: &str = "smg_adaptive_admission_engine_waiting_uncached_tokens"; +const ENGINE_TOKEN_USAGE: &str = "smg_adaptive_admission_engine_max_token_usage"; +const SEGMENTS: &str = "smg_adaptive_admission_estimator_segments"; + +pub(crate) fn describe_metrics() { + describe_counter!( + PREDICTIONS_TOTAL, + "Adaptive output-token predictions by model and fallback level" + ); + describe_histogram!( + PREDICTED_OUTPUT_TOKENS, + "Predicted completion tokens per adaptive-admission request" + ); + describe_histogram!( + OBSERVED_OUTPUT_TOKENS, + "Observed completion tokens for requests learned by adaptive admission" + ); + describe_histogram!( + ABSOLUTE_ERROR_TOKENS, + "Absolute adaptive output-token prediction error" + ); + describe_counter!( + DECISIONS_TOTAL, + "Adaptive token-work admission decisions, including shadow decisions" + ); + describe_gauge!( + OUTSTANDING_TOKENS, + "Effective outstanding output-token estimate used for adaptive admission" + ); + describe_gauge!( + ROUTER_OUTSTANDING_TOKENS, + "Predicted output tokens reserved by this router process" + ); + describe_gauge!( + ENGINE_ESTIMATED_TOKENS, + "Estimated output tokens already running or waiting on engines" + ); + describe_gauge!( + WORK_BUDGET_TOKENS, + "Live output-token work budget derived from engine throughput and the configured horizon" + ); + describe_gauge!( + DRAIN_SECONDS, + "Predicted seconds to drain outstanding output work at current engine throughput" + ); + describe_gauge!( + LOAD_COVERAGE, + "Fraction of healthy replicas with fresh engine-load telemetry" + ); + describe_gauge!( + GEN_THROUGHPUT, + "Aggregate generation throughput reported by engines in an admission partition" + ); + describe_gauge!( + LEARNED_CAPACITY, + "Recent decayed-peak generation capacity learned from engine telemetry" + ); + describe_gauge!( + ENGINE_RUNNING, + "Engine-reported running requests in an admission partition" + ); + describe_gauge!( + ENGINE_WAITING, + "Engine-reported waiting requests in an admission partition" + ); + describe_gauge!( + ENGINE_WAITING_TOKENS, + "Engine-reported waiting uncached tokens in an admission partition" + ); + describe_gauge!( + ENGINE_TOKEN_USAGE, + "Maximum engine-reported token usage in an admission partition" + ); + describe_gauge!( + SEGMENTS, + "Current bounded in-memory output estimator segment count" + ); +} + +#[derive(Debug, Clone)] +pub(crate) struct PredictionFeatures { + pub model: String, + pub user: String, + pub workload_type: String, + pub endpoint: &'static str, + pub prompt_tokens: u32, + /// Total upper bound after multiplying per-completion limits by request + /// multiplicity. `None` means the client did not supply an upper bound. + pub max_output_tokens: Option, + pub generation_flags: u16, +} + +impl PredictionFeatures { + fn prompt_bucket(&self) -> u8 { + if self.prompt_tokens == 0 { + 0 + } else { + (u32::BITS - self.prompt_tokens.leading_zeros()) as u8 + } + } + + fn output_limit_bucket(&self) -> u8 { + self.max_output_tokens.map_or(0, |tokens| { + if tokens == 0 { + 0 + } else { + (u32::BITS - tokens.leading_zeros()) as u8 + } + }) + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum PredictionSource { + ColdStart, + Model, + UserOrWorkload, + UserWorkload, + Full, +} + +impl PredictionSource { + fn as_str(self) -> &'static str { + match self { + Self::ColdStart => "cold_start", + Self::Model => "model", + Self::UserOrWorkload => "user_or_workload", + Self::UserWorkload => "user_workload", + Self::Full => "full", + } + } +} + +#[derive(Debug, Clone)] +struct Prediction { + output_tokens: u32, + model_output_tokens: u32, + source: PredictionSource, +} + +#[derive(Debug, Clone, PartialEq, Eq, Hash)] +enum SegmentKey { + Model(String), + Workload(String, String), + User(String, String), + UserWorkload(String, String, String), + Full { + model: String, + user: String, + workload: String, + endpoint: &'static str, + prompt_bucket: u8, + output_limit_bucket: u8, + generation_flags: u16, + }, +} + +#[derive(Debug, Clone)] +struct DecayedMean { + weight: f64, + weighted_sum: f64, + last_update: Instant, +} + +impl DecayedMean { + fn new(value: f64, now: Instant) -> Self { + Self { + weight: 1.0, + weighted_sum: value, + last_update: now, + } + } + + fn decay_factor(&self, now: Instant, half_life_secs: f64) -> f64 { + let elapsed = now + .saturating_duration_since(self.last_update) + .as_secs_f64(); + 2.0_f64.powf(-elapsed / half_life_secs) + } + + fn effective(&self, now: Instant, half_life_secs: f64) -> (f64, f64) { + let decay = self.decay_factor(now, half_life_secs); + (self.weight * decay, self.weighted_sum * decay) + } + + fn update(&mut self, value: f64, now: Instant, half_life_secs: f64) { + let decay = self.decay_factor(now, half_life_secs); + self.weight = self.weight * decay + 1.0; + self.weighted_sum = self.weighted_sum * decay + value; + self.last_update = now; + } +} + +#[derive(Debug)] +struct HierarchicalPredictor { + half_life_secs: f64, + prior_observations: f64, + cold_start_output_tokens: u32, + max_segments: usize, + segments: HashMap, +} + +impl HierarchicalPredictor { + fn new(config: &AdaptiveAdmissionConfig) -> Self { + Self { + half_life_secs: config.estimator_half_life_secs, + prior_observations: config.prior_observations, + cold_start_output_tokens: config.cold_start_output_tokens, + max_segments: config.max_segments, + segments: HashMap::new(), + } + } + + fn keys(features: &PredictionFeatures) -> [SegmentKey; 5] { + [ + SegmentKey::Model(features.model.clone()), + SegmentKey::Workload(features.model.clone(), features.workload_type.clone()), + SegmentKey::User(features.model.clone(), features.user.clone()), + SegmentKey::UserWorkload( + features.model.clone(), + features.user.clone(), + features.workload_type.clone(), + ), + SegmentKey::Full { + model: features.model.clone(), + user: features.user.clone(), + workload: features.workload_type.clone(), + endpoint: features.endpoint, + prompt_bucket: features.prompt_bucket(), + output_limit_bucket: features.output_limit_bucket(), + generation_flags: features.generation_flags, + }, + ] + } + + fn estimate(&self, key: &SegmentKey, now: Instant) -> Option<(f64, f64)> { + let (weight, sum) = self.segments.get(key)?.effective(now, self.half_life_secs); + (weight > f64::EPSILON).then_some((sum / weight, weight)) + } + + fn blend(&self, prior: f64, estimate: Option<(f64, f64)>) -> (f64, bool) { + let Some((mean, weight)) = estimate else { + return (prior, false); + }; + let denominator = weight + self.prior_observations; + if denominator <= f64::EPSILON { + return (mean, true); + } + ( + (mean * weight + prior * self.prior_observations) / denominator, + true, + ) + } + + fn predict_at(&self, features: &PredictionFeatures, now: Instant) -> Prediction { + let [model, workload, user, user_workload, full] = Self::keys(features); + let cold = f64::from(self.cold_start_output_tokens); + let (model_prediction, has_model) = self.blend(cold, self.estimate(&model, now)); + + let workload_estimate = self.estimate(&workload, now); + let user_estimate = self.estimate(&user, now); + let parent = match (workload_estimate, user_estimate) { + (Some((workload_mean, workload_weight)), Some((user_mean, user_weight))) => { + let total = workload_weight + user_weight; + if total <= f64::EPSILON { + model_prediction + } else { + (workload_mean * workload_weight + user_mean * user_weight) / total + } + } + (Some((mean, _)), None) | (None, Some((mean, _))) => mean, + (None, None) => model_prediction, + }; + let has_user_or_workload = workload_estimate.is_some() || user_estimate.is_some(); + let (user_workload_prediction, has_user_workload) = + self.blend(parent, self.estimate(&user_workload, now)); + let (full_prediction, has_full) = + self.blend(user_workload_prediction, self.estimate(&full, now)); + + let source = if has_full { + PredictionSource::Full + } else if has_user_workload { + PredictionSource::UserWorkload + } else if has_user_or_workload { + PredictionSource::UserOrWorkload + } else if has_model { + PredictionSource::Model + } else { + PredictionSource::ColdStart + }; + let mut output_tokens = full_prediction.round().clamp(1.0, f64::from(u32::MAX)) as u32; + if let Some(maximum) = features.max_output_tokens { + output_tokens = output_tokens.min(maximum.max(1)); + } + Prediction { + output_tokens, + model_output_tokens: model_prediction.round().clamp(1.0, f64::from(u32::MAX)) as u32, + source, + } + } + + fn observe_at(&mut self, features: &PredictionFeatures, output_tokens: u32, now: Instant) { + for key in Self::keys(features) { + self.segments + .entry(key) + .and_modify(|mean| { + mean.update(f64::from(output_tokens), now, self.half_life_secs); + }) + .or_insert_with(|| DecayedMean::new(f64::from(output_tokens), now)); + } + self.evict_oldest(); + } + + fn evict_oldest(&mut self) { + if self.segments.len() <= self.max_segments { + return; + } + // Evict a batch so a stream of one-off users does not sort the entire + // table on every completion after reaching the limit. + let batch = (self.max_segments / 10).max(1); + let target = self.max_segments.saturating_sub(batch); + let excess = self.segments.len().saturating_sub(target); + let mut oldest: Vec<_> = self + .segments + .iter() + .map(|(key, value)| (key.clone(), value.last_update)) + .collect(); + oldest.sort_unstable_by_key(|(_, updated)| *updated); + for (key, _) in oldest.into_iter().take(excess) { + self.segments.remove(&key); + } + } +} + +#[derive(Debug, Clone, Default)] +struct PartitionLoad { + healthy_replicas: u32, + observed_replicas: u32, + generation_tokens_per_second: f64, + learned_capacity_tokens_per_second: f64, + running_requests: i64, + waiting_requests: i64, + waiting_uncached_tokens: i64, + max_token_usage: f64, +} + +#[derive(Debug, Clone)] +struct CapacityEstimate { + per_replica_tokens_per_second: f64, + last_update: Instant, +} + +impl CapacityEstimate { + fn effective(&self, now: Instant, half_life_secs: f64) -> f64 { + let elapsed = now + .saturating_duration_since(self.last_update) + .as_secs_f64(); + self.per_replica_tokens_per_second * 2.0_f64.powf(-elapsed / half_life_secs) + } + + fn observe(&mut self, value: f64, now: Instant, half_life_secs: f64) { + self.per_replica_tokens_per_second = self.effective(now, half_life_secs).max(value); + self.last_update = now; + } +} + +impl PartitionLoad { + fn coverage(&self) -> f64 { + if self.healthy_replicas == 0 { + 0.0 + } else { + f64::from(self.observed_replicas) / f64::from(self.healthy_replicas) + } + } +} + +#[derive(Debug, Default)] +struct WorkState { + outstanding_tokens: HashMap, + loads: HashMap, + capacities: HashMap, +} + +#[derive(Debug)] +pub(crate) struct AdaptiveAdmissionController { + config: AdaptiveAdmissionConfig, + predictor: Mutex, + work: Mutex, + registry: Arc, +} + +impl AdaptiveAdmissionController { + pub(crate) fn new(config: AdaptiveAdmissionConfig, registry: Arc) -> Arc { + Arc::new(Self { + predictor: Mutex::new(HierarchicalPredictor::new(&config)), + config, + work: Mutex::new(WorkState::default()), + registry, + }) + } + + pub(crate) fn mode(&self) -> AdaptiveAdmissionMode { + self.config.mode + } + + pub(crate) fn start_load_updates( + self: &Arc, + mut loads: watch::Receiver>, + ) { + self.update_loads(&loads.borrow()); + let controller = Arc::downgrade(self); + #[expect( + clippy::disallowed_methods, + reason = "controller task holds only a weak reference and exits with the gateway" + )] + tokio::spawn(async move { + loop { + if loads.changed().await.is_err() { + break; + } + let Some(controller) = controller.upgrade() else { + break; + }; + controller.update_loads(&loads.borrow()); + } + }); + } + + fn update_loads(&self, loads: &HashMap) { + let mut partitions: HashMap = HashMap::new(); + for worker in self + .registry + .get_all() + .into_iter() + .filter(|w| w.is_healthy()) + { + let partition = worker + .metadata() + .spec + .labels + .get(ADMISSION_PARTITION_LABEL) + .map(String::as_str) + .filter(|value| !value.trim().is_empty()) + .unwrap_or_else(|| worker.model_id()) + .to_string(); + let aggregate = partitions.entry(partition).or_default(); + aggregate.healthy_replicas = aggregate.healthy_replicas.saturating_add(1); + let Some(load) = loads.get(worker.url()) else { + continue; + }; + aggregate.observed_replicas = aggregate.observed_replicas.saturating_add(1); + aggregate.generation_tokens_per_second += load.total_gen_throughput().max(0.0); + aggregate.running_requests += load + .loads + .iter() + .map(|rank| i64::from(rank.num_running_reqs.max(0))) + .sum::(); + aggregate.waiting_requests += load + .loads + .iter() + .map(|rank| i64::from(rank.num_waiting_reqs.max(0))) + .sum::(); + aggregate.waiting_uncached_tokens += load.total_waiting_uncached_tokens().max(0); + aggregate.max_token_usage = aggregate + .max_token_usage + .max(load.effective_token_usage().clamp(0.0, 1.0)); + } + + let now = Instant::now(); + let mut work = self.work.lock(); + work.capacities + .retain(|partition, _| partitions.contains_key(partition)); + for (partition, load) in &mut partitions { + if load.observed_replicas > 0 && load.generation_tokens_per_second > 0.0 { + let per_replica = + load.generation_tokens_per_second / f64::from(load.observed_replicas); + work.capacities + .entry(partition.clone()) + .and_modify(|capacity| { + capacity.observe(per_replica, now, self.config.estimator_half_life_secs); + }) + .or_insert(CapacityEstimate { + per_replica_tokens_per_second: per_replica, + last_update: now, + }); + } + if let Some(capacity) = work.capacities.get(partition) { + load.learned_capacity_tokens_per_second = capacity + .effective(now, self.config.estimator_half_life_secs) + * f64::from(load.healthy_replicas); + } + } + work.loads = partitions; + for (partition, load) in &work.loads { + let partition_label = intern_string(partition); + gauge!(LOAD_COVERAGE, "partition" => Arc::clone(&partition_label)).set(load.coverage()); + gauge!(GEN_THROUGHPUT, "partition" => Arc::clone(&partition_label)) + .set(load.generation_tokens_per_second); + gauge!(LEARNED_CAPACITY, "partition" => Arc::clone(&partition_label)) + .set(load.learned_capacity_tokens_per_second); + gauge!(ENGINE_RUNNING, "partition" => Arc::clone(&partition_label)) + .set(load.running_requests as f64); + gauge!(ENGINE_WAITING, "partition" => Arc::clone(&partition_label)) + .set(load.waiting_requests as f64); + gauge!(ENGINE_WAITING_TOKENS, "partition" => Arc::clone(&partition_label)) + .set(load.waiting_uncached_tokens as f64); + gauge!(ENGINE_TOKEN_USAGE, "partition" => partition_label).set(load.max_token_usage); + } + } + + pub(crate) fn begin( + self: &Arc, + partition: String, + features: PredictionFeatures, + ) -> AdaptiveRequestTracker { + // Comet injects a trusted partition header, but standalone SMG users + // can send arbitrary headers. Only retain a selector already known + // from fresh worker telemetry; otherwise fall back to the request's + // model. This keeps state and metric cardinality bounded. + let partition = { + let work = self.work.lock(); + if partition == features.model || work.loads.contains_key(&partition) { + partition + } else { + features.model.clone() + } + }; + let now = Instant::now(); + let prediction = self.predictor.lock().predict_at(&features, now); + let model_label = intern_string(&features.model); + counter!( + PREDICTIONS_TOTAL, + "model" => Arc::clone(&model_label), + "source" => prediction.source.as_str() + ) + .increment(1); + histogram!(PREDICTED_OUTPUT_TOKENS, "model" => model_label) + .record(f64::from(prediction.output_tokens)); + + let decision = { + let mut work = self.work.lock(); + let load = work.loads.get(&partition).cloned().unwrap_or_default(); + let outstanding = work + .outstanding_tokens + .entry(partition.clone()) + .or_default(); + let prior_outstanding = *outstanding; + *outstanding = outstanding.saturating_add(u64::from(prediction.output_tokens)); + let coverage = load.coverage(); + let telemetry_usable = coverage >= self.config.min_load_coverage + && load.learned_capacity_tokens_per_second.is_finite() + && load.learned_capacity_tokens_per_second > 0.0; + let budget = if telemetry_usable { + load.learned_capacity_tokens_per_second * self.config.work_horizon_secs + } else { + f64::INFINITY + }; + let router_outstanding = *outstanding as f64; + let engine_request_count = load + .running_requests + .saturating_add(load.waiting_requests) + .max(0) as f64; + let engine_estimated = engine_request_count * f64::from(prediction.model_output_tokens); + // `max` avoids counting work both in the router reservation table + // and in a later engine poll. The incoming request is not yet in + // the engine snapshot, so include it in the engine-side bound. + let projected = + router_outstanding.max(engine_estimated + f64::from(prediction.output_tokens)); + let drain_seconds = if telemetry_usable { + projected / load.learned_capacity_tokens_per_second + } else { + 0.0 + }; + let would_admit = !telemetry_usable || projected <= budget || prior_outstanding == 0; + + let partition_label = intern_string(&partition); + gauge!(OUTSTANDING_TOKENS, "partition" => Arc::clone(&partition_label)).set(projected); + gauge!(ROUTER_OUTSTANDING_TOKENS, "partition" => Arc::clone(&partition_label)) + .set(router_outstanding); + gauge!(ENGINE_ESTIMATED_TOKENS, "partition" => Arc::clone(&partition_label)) + .set(engine_estimated); + gauge!(WORK_BUDGET_TOKENS, "partition" => Arc::clone(&partition_label)) + .set(if budget.is_finite() { budget } else { 0.0 }); + gauge!(DRAIN_SECONDS, "partition" => partition_label).set(drain_seconds); + AdmissionDecision { + would_admit, + telemetry_usable, + retry_after_secs: if would_admit || !telemetry_usable { + 0 + } else { + ((projected - budget) / load.learned_capacity_tokens_per_second) + .ceil() + .clamp(1.0, f64::from(u32::MAX)) as u32 + }, + } + }; + + let outcome = if !decision.telemetry_usable { + "telemetry_fallback" + } else if decision.would_admit { + "would_admit" + } else { + "would_reject" + }; + counter!( + DECISIONS_TOTAL, + "partition" => intern_string(&partition), + "mode" => match self.config.mode { + AdaptiveAdmissionMode::Off => "off", + AdaptiveAdmissionMode::Shadow => "shadow", + AdaptiveAdmissionMode::Enforce => "enforce", + }, + "outcome" => outcome + ) + .increment(1); + + AdaptiveRequestTracker { + inner: Some(TrackerInner { + controller: Arc::downgrade(self), + partition, + features, + prediction, + decision, + resolved: AtomicBool::new(false), + }), + } + } + + fn finish(&self, inner: &TrackerInner, observed_output_tokens: Option) { + { + let mut work = self.work.lock(); + let outstanding = work + .outstanding_tokens + .entry(inner.partition.clone()) + .or_default(); + *outstanding = outstanding.saturating_sub(u64::from(inner.prediction.output_tokens)); + gauge!(OUTSTANDING_TOKENS, "partition" => intern_string(&inner.partition)) + .set(*outstanding as f64); + } + let Some(observed) = observed_output_tokens else { + return; + }; + self.predictor + .lock() + .observe_at(&inner.features, observed, Instant::now()); + let model_label = intern_string(&inner.features.model); + histogram!(OBSERVED_OUTPUT_TOKENS, "model" => Arc::clone(&model_label)) + .record(f64::from(observed)); + histogram!(ABSOLUTE_ERROR_TOKENS, "model" => model_label) + .record((f64::from(observed) - f64::from(inner.prediction.output_tokens)).abs()); + gauge!(SEGMENTS).set(self.predictor.lock().segments.len() as f64); + } +} + +#[derive(Debug, Clone, Copy)] +struct AdmissionDecision { + would_admit: bool, + telemetry_usable: bool, + retry_after_secs: u32, +} + +struct TrackerInner { + controller: Weak, + partition: String, + features: PredictionFeatures, + prediction: Prediction, + decision: AdmissionDecision, + resolved: AtomicBool, +} + +/// Per-request adaptive-admission state. Dropping an unfinished tracker +/// releases its predicted work without teaching the estimator from a partial +/// or failed response. +pub(crate) struct AdaptiveRequestTracker { + inner: Option, +} + +impl AdaptiveRequestTracker { + pub(crate) fn should_reject(&self) -> bool { + let Some(inner) = &self.inner else { + return false; + }; + inner.controller.upgrade().is_some_and(|controller| { + controller.mode() == AdaptiveAdmissionMode::Enforce + && inner.decision.telemetry_usable + && !inner.decision.would_admit + }) + } + + pub(crate) fn retry_after_secs(&self) -> u32 { + self.inner + .as_ref() + .map_or(0, |inner| inner.decision.retry_after_secs) + } + + pub(crate) fn complete(mut self, observed_output_tokens: u32) { + self.resolve(Some(observed_output_tokens)); + } + + fn resolve(&mut self, observed_output_tokens: Option) { + let Some(inner) = self.inner.take() else { + return; + }; + if inner.resolved.swap(true, Ordering::AcqRel) { + return; + } + if let Some(controller) = inner.controller.upgrade() { + controller.finish(&inner, observed_output_tokens); + } + } +} + +impl Drop for AdaptiveRequestTracker { + fn drop(&mut self) { + self.resolve(None); + } +} + +#[cfg(test)] +mod tests { + use std::time::Duration; + + use super::*; + + fn config() -> AdaptiveAdmissionConfig { + AdaptiveAdmissionConfig { + mode: AdaptiveAdmissionMode::Shadow, + work_horizon_secs: 10.0, + estimator_half_life_secs: 60.0, + prior_observations: 2.0, + max_segments: 20, + min_load_coverage: 0.8, + cold_start_output_tokens: 100, + } + } + + fn features(user: &str, prompt_tokens: u32, maximum: Option) -> PredictionFeatures { + PredictionFeatures { + model: "model".to_string(), + user: user.to_string(), + workload_type: "rollout".to_string(), + endpoint: "chat", + prompt_tokens, + max_output_tokens: maximum, + generation_flags: 0, + } + } + + #[test] + fn cold_start_is_clamped_by_request_limit() { + let predictor = HierarchicalPredictor::new(&config()); + let prediction = predictor.predict_at(&features("u", 10, Some(32)), Instant::now()); + assert_eq!(prediction.output_tokens, 32); + assert_eq!(prediction.source, PredictionSource::ColdStart); + } + + #[test] + fn user_history_beats_model_mean_and_decays() { + let mut predictor = HierarchicalPredictor::new(&config()); + let start = Instant::now(); + let user_a = features("a", 1000, None); + let user_b = features("b", 1000, None); + for i in 0..20 { + let now = start + Duration::from_secs(i); + predictor.observe_at(&user_a, 20, now); + predictor.observe_at(&user_b, 200, now); + } + let prediction = predictor.predict_at(&user_a, start + Duration::from_secs(21)); + assert!(prediction.output_tokens < 80, "{prediction:?}"); + assert!(matches!( + prediction.source, + PredictionSource::UserWorkload | PredictionSource::Full + )); + + predictor.observe_at(&user_a, 400, start + Duration::from_secs(600)); + let shifted = predictor.predict_at(&user_a, start + Duration::from_secs(601)); + assert!(shifted.output_tokens > prediction.output_tokens); + } + + #[test] + fn estimator_state_is_bounded() { + let mut predictor = HierarchicalPredictor::new(&config()); + let start = Instant::now(); + for i in 0..100 { + predictor.observe_at( + &features(&format!("user-{i}"), i + 1, None), + i + 1, + start + Duration::from_secs(u64::from(i)), + ); + } + assert!(predictor.segments.len() <= predictor.max_segments); + assert!(predictor.segments.len() >= predictor.max_segments / 2); + } + + #[test] + fn learned_capacity_tracks_recent_peak_and_decays() { + let start = Instant::now(); + let mut capacity = CapacityEstimate { + per_replica_tokens_per_second: 100.0, + last_update: start, + }; + capacity.observe(40.0, start + Duration::from_secs(60), 60.0); + assert!((capacity.per_replica_tokens_per_second - 50.0).abs() < f64::EPSILON); + capacity.observe(80.0, start + Duration::from_secs(61), 60.0); + assert!((capacity.per_replica_tokens_per_second - 80.0).abs() < f64::EPSILON); + } + + #[test] + fn unknown_partition_header_falls_back_to_model() { + let controller = + AdaptiveAdmissionController::new(config(), Arc::new(WorkerRegistry::new())); + let tracker = controller.begin("attacker-controlled".to_string(), features("a", 10, None)); + assert_eq!(tracker.inner.as_ref().unwrap().partition, "model"); + } + + #[test] + fn work_horizon_decision_uses_throughput_not_request_count() { + let registry = Arc::new(WorkerRegistry::new()); + let controller = AdaptiveAdmissionController::new(config(), registry); + controller.work.lock().loads.insert( + "model".to_string(), + PartitionLoad { + healthy_replicas: 1, + observed_replicas: 1, + generation_tokens_per_second: 10.0, + learned_capacity_tokens_per_second: 10.0, + ..PartitionLoad::default() + }, + ); + + let first = controller.begin("model".to_string(), features("a", 10, None)); + assert!(!first.should_reject(), "shadow mode never rejects"); + let second = controller.begin("model".to_string(), features("b", 10, None)); + assert!(!second.inner.as_ref().unwrap().decision.would_admit); + drop(second); + drop(first); + assert_eq!( + controller.work.lock().outstanding_tokens.get("model"), + Some(&0) + ); + } + + #[test] + fn enforce_rejects_only_after_live_work_budget_is_exhausted() { + let mut settings = config(); + settings.mode = AdaptiveAdmissionMode::Enforce; + let controller = + AdaptiveAdmissionController::new(settings, Arc::new(WorkerRegistry::new())); + controller.work.lock().loads.insert( + "model".to_string(), + PartitionLoad { + healthy_replicas: 1, + observed_replicas: 1, + generation_tokens_per_second: 10.0, + learned_capacity_tokens_per_second: 10.0, + ..PartitionLoad::default() + }, + ); + + let first = controller.begin("model".to_string(), features("a", 10, None)); + assert!(!first.should_reject()); + let second = controller.begin("model".to_string(), features("b", 10, None)); + assert!(second.should_reject()); + assert_eq!(second.retry_after_secs(), 10); + } + + #[test] + fn enforce_fails_open_when_engine_load_coverage_is_incomplete() { + let mut settings = config(); + settings.mode = AdaptiveAdmissionMode::Enforce; + let controller = + AdaptiveAdmissionController::new(settings, Arc::new(WorkerRegistry::new())); + controller.work.lock().loads.insert( + "model".to_string(), + PartitionLoad { + healthy_replicas: 2, + observed_replicas: 1, + generation_tokens_per_second: 10.0, + learned_capacity_tokens_per_second: 20.0, + ..PartitionLoad::default() + }, + ); + + let first = controller.begin("model".to_string(), features("a", 10, None)); + let second = controller.begin("model".to_string(), features("b", 10, None)); + assert!(!first.should_reject()); + assert!(!second.should_reject()); + assert!(!second.inner.as_ref().unwrap().decision.telemetry_usable); + } + + #[test] + fn engine_backlog_survives_router_reservation_loss() { + let mut settings = config(); + settings.mode = AdaptiveAdmissionMode::Enforce; + settings.work_horizon_secs = 30.0; + let controller = + AdaptiveAdmissionController::new(settings, Arc::new(WorkerRegistry::new())); + controller.work.lock().loads.insert( + "model".to_string(), + PartitionLoad { + healthy_replicas: 1, + observed_replicas: 1, + generation_tokens_per_second: 10.0, + learned_capacity_tokens_per_second: 10.0, + running_requests: 3, + ..PartitionLoad::default() + }, + ); + + let starvation_probe = controller.begin("model".to_string(), features("a", 10, None)); + assert!(!starvation_probe.should_reject()); + let next = controller.begin("model".to_string(), features("b", 10, None)); + assert!(next.should_reject()); + } +} diff --git a/model_gateway/src/routers/grpc/common/stages/adaptive_admission.rs b/model_gateway/src/routers/grpc/common/stages/adaptive_admission.rs new file mode 100644 index 000000000..fcd34b42f --- /dev/null +++ b/model_gateway/src/routers/grpc/common/stages/adaptive_admission.rs @@ -0,0 +1,243 @@ +//! Adaptive predicted-work admission after request tokenization. + +use async_trait::async_trait; +use axum::{ + http::{header::RETRY_AFTER, HeaderValue, StatusCode}, + response::Response, +}; + +use super::PipelineStage; +use crate::{ + middleware::scheduler::ADMISSION_PARTITION_HEADER, + routers::{ + error, + grpc::{ + adaptive_admission::PredictionFeatures, + context::{RequestContext, RequestType}, + }, + }, +}; + +const COMET_USER_HEADER: &str = "x-comet-user"; +const COMET_WORKLOAD_TYPE_HEADER: &str = "x-comet-workload-type"; + +const FLAG_MULTIPLE_COMPLETIONS: u16 = 1 << 0; +const FLAG_TOOLS: u16 = 1 << 1; +const FLAG_STRUCTURED_OUTPUT: u16 = 1 << 2; +const FLAG_REASONING: u16 = 1 << 3; +const FLAG_STREAMING: u16 = 1 << 4; + +pub(crate) struct AdaptiveAdmissionStage; + +#[derive(Clone, Copy)] +struct GenerationShape { + endpoint: &'static str, + max_output_tokens: Option, + flags: u16, +} + +impl GenerationShape { + fn for_request(request: &RequestType) -> Option { + match request { + RequestType::Chat(request) => { + #[expect( + deprecated, + reason = "max_tokens remains an OpenAI compatibility fallback" + )] + let per_completion_limit = request.max_completion_tokens.or(request.max_tokens); + let multiplicity = request.n.unwrap_or(1).max(1); + let mut flags = 0; + if multiplicity > 1 { + flags |= FLAG_MULTIPLE_COMPLETIONS; + } + if request + .tools + .as_ref() + .is_some_and(|tools| !tools.is_empty()) + { + flags |= FLAG_TOOLS; + } + if request.response_format.is_some() + || request.regex.is_some() + || request.ebnf.is_some() + { + flags |= FLAG_STRUCTURED_OUTPUT; + } + if request.reasoning_effort.is_some() { + flags |= FLAG_REASONING; + } + if request.stream { + flags |= FLAG_STREAMING; + } + Some(Self { + endpoint: "chat", + max_output_tokens: multiplied_limit(per_completion_limit, multiplicity), + flags, + }) + } + RequestType::Generate(request) => { + let params = request.sampling_params.as_ref(); + let multiplicity = params.and_then(|p| p.n).unwrap_or(1).max(1); + let mut flags = 0; + if multiplicity > 1 { + flags |= FLAG_MULTIPLE_COMPLETIONS; + } + if params.is_some_and(|p| { + p.json_schema.is_some() || p.regex.is_some() || p.ebnf.is_some() + }) { + flags |= FLAG_STRUCTURED_OUTPUT; + } + if request.stream { + flags |= FLAG_STREAMING; + } + Some(Self { + endpoint: "generate", + max_output_tokens: multiplied_limit( + params.and_then(|p| p.max_new_tokens), + multiplicity, + ), + flags, + }) + } + RequestType::Completion(request) => { + let returned = request.n.unwrap_or(1).max(1); + let generated = request.best_of.unwrap_or(returned).max(returned); + let mut flags = 0; + if generated > 1 { + flags |= FLAG_MULTIPLE_COMPLETIONS; + } + if request.stream { + flags |= FLAG_STREAMING; + } + Some(Self { + endpoint: "completion", + max_output_tokens: multiplied_limit(request.max_tokens, generated), + flags, + }) + } + RequestType::Messages(request) => { + let mut flags = 0; + if request + .tools + .as_ref() + .is_some_and(|tools| !tools.is_empty()) + { + flags |= FLAG_TOOLS; + } + if request.thinking.is_some() { + flags |= FLAG_REASONING; + } + if request.is_stream() { + flags |= FLAG_STREAMING; + } + Some(Self { + endpoint: "messages", + max_output_tokens: Some(request.max_tokens), + flags, + }) + } + RequestType::Responses(request) => { + let mut flags = 0; + if request + .tools + .as_ref() + .is_some_and(|tools| !tools.is_empty()) + { + flags |= FLAG_TOOLS; + } + if request.text.is_some() { + flags |= FLAG_STRUCTURED_OUTPUT; + } + if request.reasoning.is_some() { + flags |= FLAG_REASONING; + } + if request.stream.unwrap_or(false) { + flags |= FLAG_STREAMING; + } + Some(Self { + endpoint: "responses", + max_output_tokens: request.max_output_tokens, + flags, + }) + } + RequestType::Embedding(_) | RequestType::Classify(_) => None, + } + } +} + +fn multiplied_limit(per_completion: Option, multiplicity: u32) -> Option { + per_completion.map(|limit| limit.saturating_mul(multiplicity)) +} + +fn trusted_header(ctx: &RequestContext, name: &str) -> Option { + ctx.input + .headers + .as_ref() + .and_then(|headers| headers.get(name)) + .and_then(|value| value.to_str().ok()) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(str::to_string) +} + +fn rejection_response(retry_after_secs: u32) -> Response { + let mut response = error::create_error( + StatusCode::TOO_MANY_REQUESTS, + "adaptive_admission_saturated", + "model fleet is temporarily saturated", + ); + if let Ok(value) = HeaderValue::from_str(&retry_after_secs.max(1).to_string()) { + response.headers_mut().insert(RETRY_AFTER, value); + } + response +} + +#[async_trait] +impl PipelineStage for AdaptiveAdmissionStage { + async fn execute(&self, ctx: &mut RequestContext) -> Result, Response> { + let Some(controller) = ctx.components.adaptive_admission.clone() else { + return Ok(None); + }; + let Some(shape) = GenerationShape::for_request(&ctx.input.request_type) else { + return Ok(None); + }; + let prompt_tokens = ctx + .state + .preparation + .as_ref() + .map_or(0, |preparation| preparation.total_token_count()); + let partition = trusted_header(ctx, ADMISSION_PARTITION_HEADER) + .unwrap_or_else(|| ctx.input.model_id.clone()); + let user = trusted_header(ctx, COMET_USER_HEADER) + .or_else(|| { + ctx.input + .tenant_request_meta + .as_ref() + .map(|meta| meta.tenant_key().as_str().to_string()) + }) + .unwrap_or_else(|| "anonymous".to_string()); + let workload_type = trusted_header(ctx, COMET_WORKLOAD_TYPE_HEADER) + .unwrap_or_else(|| "unclassified".to_string()); + let tracker = controller.begin( + partition, + PredictionFeatures { + model: ctx.input.model_id.clone(), + user, + workload_type, + endpoint: shape.endpoint, + prompt_tokens, + max_output_tokens: shape.max_output_tokens, + generation_flags: shape.flags, + }, + ); + if tracker.should_reject() { + return Err(rejection_response(tracker.retry_after_secs())); + } + ctx.state.adaptive_request = Some(tracker); + Ok(None) + } + + fn name(&self) -> &'static str { + "AdaptiveAdmission" + } +} diff --git a/model_gateway/src/routers/grpc/common/stages/mod.rs b/model_gateway/src/routers/grpc/common/stages/mod.rs index aa35d8edb..d31ba8051 100644 --- a/model_gateway/src/routers/grpc/common/stages/mod.rs +++ b/model_gateway/src/routers/grpc/common/stages/mod.rs @@ -39,6 +39,7 @@ pub trait PipelineStage: Send + Sync { } } +mod adaptive_admission; mod client_acquisition; mod dispatch_metadata; pub(crate) mod encode; @@ -47,6 +48,7 @@ mod request_execution; mod worker_selection; // Export stage implementations +pub(crate) use adaptive_admission::AdaptiveAdmissionStage; pub(crate) use client_acquisition::ClientAcquisitionStage; pub(crate) use dispatch_metadata::DispatchMetadataStage; pub(crate) use encode::EncodeStage; diff --git a/model_gateway/src/routers/grpc/context.rs b/model_gateway/src/routers/grpc/context.rs index 7a00d2c3e..ab19bb7e5 100644 --- a/model_gateway/src/routers/grpc/context.rs +++ b/model_gateway/src/routers/grpc/context.rs @@ -22,6 +22,7 @@ use tool_parser::ParserFactory as ToolParserFactory; use tracing::debug; use super::{ + adaptive_admission::{AdaptiveAdmissionController, AdaptiveRequestTracker}, client::GrpcClient, common::stages::encode::EncodeDispatchPlan, multimodal::{MultimodalComponents, MultimodalIntermediate}, @@ -146,6 +147,7 @@ pub(crate) struct SharedComponents { pub configured_reasoning_parser: Option, /// Multimodal processing components (initialized at router creation) pub multimodal: Option>, + pub adaptive_admission: Option>, } /// Mutable processing state (evolves through pipeline stages) @@ -168,6 +170,12 @@ pub(crate) struct ProcessingState { /// This avoids redundant registry lookups across pipeline stages. pub tokenizer: Option>, + /// Predicted token-work reservation. Non-streaming response processing + /// completes it from the final usage counters; streaming processing moves + /// it into the background stream task. Any earlier error drops it and + /// releases predicted outstanding work without training on a partial run. + pub adaptive_request: Option, + // Stage 2: Worker selection outputs pub workers: Option, @@ -329,6 +337,18 @@ pub(crate) struct CompletionItem { } impl PreparationOutput { + /// Total prompt tokens represented by this request. Completion batches + /// sum every prompt because they fan out into independent backend work. + pub fn total_token_count(&self) -> u32 { + let count = match self { + Self::Completion { items, .. } => { + items.iter().map(|item| item.token_ids.len()).sum::() + } + _ => self.token_ids().len(), + }; + u32::try_from(count).unwrap_or(u32::MAX) + } + /// Token IDs (common to all variants). Batched completions expose the /// first prompt's tokens as the routing-affinity proxy. pub fn token_ids(&self) -> &[u32] { diff --git a/model_gateway/src/routers/grpc/harmony/responses/streaming.rs b/model_gateway/src/routers/grpc/harmony/responses/streaming.rs index d81205444..d0ea04267 100644 --- a/model_gateway/src/routers/grpc/harmony/responses/streaming.rs +++ b/model_gateway/src/routers/grpc/harmony/responses/streaming.rs @@ -220,7 +220,7 @@ async fn execute_mcp_tool_loop_streaming( ); // Execute pipeline and get stream + load guards - let (execution_result, _load_guards) = match ctx + let (execution_result, _load_guards, adaptive_request) = match ctx .pipeline .execute_harmony_responses_streaming( ¤t_request, @@ -266,6 +266,10 @@ async fn execute_mcp_tool_loop_streaming( usage, request_id: _, } => { + if let Some(tracker) = adaptive_request { + tracker.complete(usage.completion_tokens); + } + debug!( tool_call_count = tool_calls.len(), has_analysis = analysis.is_some(), @@ -384,6 +388,10 @@ async fn execute_mcp_tool_loop_streaming( // Continue loop } ResponsesIterationResult::Completed { response, usage } => { + if let Some(tracker) = adaptive_request { + tracker.complete(usage.completion_tokens); + } + debug!( output_items = response.output.len(), input_tokens = usage.prompt_tokens, @@ -440,7 +448,7 @@ async fn execute_without_mcp_streaming( debug!("No MCP tools - executing single iteration"); // Execute pipeline and get stream + load guards - let (execution_result, _load_guards) = match ctx + let (execution_result, _load_guards, adaptive_request) = match ctx .pipeline .execute_harmony_responses_streaming(current_request, ctx, Some(tenant_request_meta)) .await @@ -479,6 +487,9 @@ async fn execute_without_mcp_streaming( ResponsesIterationResult::ToolCallsFound { usage, .. } => usage, ResponsesIterationResult::Completed { usage, .. } => usage, }; + if let Some(tracker) = adaptive_request { + tracker.complete(usage.completion_tokens); + } // Finalize response from emitter's accumulated data let final_response = emitter.finalize(Some(usage.clone())); diff --git a/model_gateway/src/routers/grpc/harmony/stages/response_processing.rs b/model_gateway/src/routers/grpc/harmony/stages/response_processing.rs index f745ca589..c6cf7ebd5 100644 --- a/model_gateway/src/routers/grpc/harmony/stages/response_processing.rs +++ b/model_gateway/src/routers/grpc/harmony/stages/response_processing.rs @@ -80,6 +80,7 @@ impl PipelineStage for HarmonyResponseProcessingStage { execution_result, ctx.chat_request_arc(), dispatch, + ctx.state.adaptive_request.take(), ); // Attach load guards to response body for proper RAII lifecycle @@ -98,6 +99,12 @@ impl PipelineStage for HarmonyResponseProcessingStage { .process_non_streaming_chat_response(execution_result, chat_request, dispatch) .await?; + if let (Some(tracker), Some(usage)) = + (ctx.state.adaptive_request.take(), response.usage.as_ref()) + { + tracker.complete(usage.completion_tokens); + } + ctx.state.response.final_response = Some(FinalResponse::Chat(response)); Ok(None) } diff --git a/model_gateway/src/routers/grpc/harmony/streaming.rs b/model_gateway/src/routers/grpc/harmony/streaming.rs index 76378c6ac..0f9c32f9d 100644 --- a/model_gateway/src/routers/grpc/harmony/streaming.rs +++ b/model_gateway/src/routers/grpc/harmony/streaming.rs @@ -38,6 +38,7 @@ use crate::{ sse::SseEncoder, }, grpc::{ + adaptive_admission::AdaptiveRequestTracker, common::{ response_formatting::CompletionTokenTracker, responses::{ @@ -105,6 +106,7 @@ impl HarmonyStreamingProcessor { execution_result: context::ExecutionResult, chat_request: Arc, dispatch: context::DispatchMetadata, + adaptive_request: Option, ) -> Response { // Create SSE channel let (tx, rx) = mpsc::unbounded_channel::>(); @@ -116,9 +118,16 @@ impl HarmonyStreamingProcessor { let result = Self::process_single_stream(stream, dispatch, chat_request, &tx).await; - if let Err(e) = result { - error!("Harmony streaming error: {}", e); - utils::send_error_sse(&tx, &e, "internal_error"); + match result { + Ok(tokens) => { + if let Some(tracker) = adaptive_request { + tracker.complete(tokens); + } + } + Err(e) => { + error!("Harmony streaming error: {}", e); + utils::send_error_sse(&tx, &e, "internal_error"); + } } let _ = tx.send(Ok(SseEncoder::done())); @@ -140,9 +149,16 @@ impl HarmonyStreamingProcessor { ) .await; - if let Err(e) = result { - error!("Harmony prefill/decode streaming error: {}", e); - utils::send_error_sse(&tx, &e, "internal_error"); + match result { + Ok(tokens) => { + if let Some(tracker) = adaptive_request { + tracker.complete(tokens); + } + } + Err(e) => { + error!("Harmony prefill/decode streaming error: {}", e); + utils::send_error_sse(&tx, &e, "internal_error"); + } } let _ = tx.send(Ok(SseEncoder::done())); @@ -179,7 +195,7 @@ impl HarmonyStreamingProcessor { dispatch: context::DispatchMetadata, original_request: Arc, tx: &mpsc::UnboundedSender>, - ) -> Result<(), String> { + ) -> Result { let mut prompt_tokens = HashMap::new(); let mut cached_tokens = HashMap::new(); Self::process_chat_decode_stream( @@ -200,7 +216,7 @@ impl HarmonyStreamingProcessor { dispatch: context::DispatchMetadata, original_request: Arc, tx: &mpsc::UnboundedSender>, - ) -> Result<(), String> { + ) -> Result { // Phase 1: Process prefill stream (collect metadata) let mut prompt_tokens: HashMap = HashMap::new(); let mut cached_tokens: HashMap = HashMap::new(); @@ -215,7 +231,7 @@ impl HarmonyStreamingProcessor { } // Phase 2: Decode (shared helper) - Self::process_chat_decode_stream( + let completion_tokens = Self::process_chat_decode_stream( decode_stream, &dispatch, &original_request, @@ -228,7 +244,7 @@ impl HarmonyStreamingProcessor { // Mark prefill stream completed AFTER decode succeeds // This ensures that if client disconnects during decode, BOTH streams send abort prefill_stream.mark_completed(); - Ok(()) + Ok(completion_tokens) } /// Process the decode phase of a Chat Completion stream. @@ -244,7 +260,7 @@ impl HarmonyStreamingProcessor { tx: &mpsc::UnboundedSender>, prompt_tokens: &mut HashMap, cached_tokens: &mut HashMap, - ) -> Result<(), String> { + ) -> Result { // Timing for metrics let start_time = Instant::now(); let mut first_token_time: Option = None; @@ -389,7 +405,7 @@ impl HarmonyStreamingProcessor { output_tokens: total_completion as u64, }); - Ok(()) + Ok(total_completion) } /// Emit a chunk delta from Harmony channels diff --git a/model_gateway/src/routers/grpc/mod.rs b/model_gateway/src/routers/grpc/mod.rs index e2cfdabcc..53c823cdc 100644 --- a/model_gateway/src/routers/grpc/mod.rs +++ b/model_gateway/src/routers/grpc/mod.rs @@ -5,6 +5,7 @@ use openai_protocol::{chat::ChatCompletionRequest, common::StringOrArray}; use crate::routers::error; +pub(crate) mod adaptive_admission; pub mod client; // Used by core/ pub(crate) mod common; pub(crate) mod context; diff --git a/model_gateway/src/routers/grpc/pipeline.rs b/model_gateway/src/routers/grpc/pipeline.rs index e71715904..14655986a 100644 --- a/model_gateway/src/routers/grpc/pipeline.rs +++ b/model_gateway/src/routers/grpc/pipeline.rs @@ -259,6 +259,7 @@ impl RequestPipeline { let (processor, streaming_processor) = deps.configured_processors(backend); let mut stages: Vec> = vec![ Box::new(ChatGeneratePreparationStage::new()), + Box::new(AdaptiveAdmissionStage), Box::new(WorkerSelectionStage::new( deps.worker_registry.clone(), deps.policy_registry.clone(), @@ -287,6 +288,7 @@ impl RequestPipeline { let (processor, streaming_processor) = deps.configured_processors(backend); let mut stages: Vec> = vec![ Box::new(MessagePreparationStage), + Box::new(AdaptiveAdmissionStage), Box::new(WorkerSelectionStage::new( deps.worker_registry.clone(), deps.policy_registry.clone(), @@ -316,6 +318,7 @@ impl RequestPipeline { let (processor, streaming_processor) = PipelineDeps::default_processors(backend); let mut stages: Vec> = vec![ Box::new(CompletionPreparationStage), + Box::new(AdaptiveAdmissionStage), Box::new(WorkerSelectionStage::new( deps.worker_registry.clone(), deps.policy_registry.clone(), @@ -347,6 +350,7 @@ impl RequestPipeline { } vec![ Box::new(harmony::stages::HarmonyPreparationStage::new()), + Box::new(AdaptiveAdmissionStage), Box::new(WorkerSelectionStage::new( deps.worker_registry.clone(), deps.policy_registry.clone(), @@ -1109,7 +1113,8 @@ impl RequestPipeline { // Extract ResponsesIterationResult from context // This should have been set by HarmonyResponseProcessingStage - ctx.state + let result = ctx + .state .response .responses_iteration_result .take() @@ -1122,7 +1127,17 @@ impl RequestPipeline { "no_responses_iteration_result", "No ResponsesIterationResult produced by pipeline", ) - }) + })?; + if let Some(tracker) = ctx.state.adaptive_request.take() { + let completion_tokens = match &result { + harmony::ResponsesIterationResult::ToolCallsFound { usage, .. } + | harmony::ResponsesIterationResult::Completed { usage, .. } => { + usage.completion_tokens + } + }; + tracker.complete(completion_tokens); + } + Ok(result) } /// Execute Harmony Responses pipeline iteration with streaming support @@ -1135,7 +1150,14 @@ impl RequestPipeline { request: &openai_protocol::responses::ResponsesRequest, harmony_ctx: &ResponsesContext, tenant_request_meta: Option, - ) -> Result<(ExecutionResult, Option), Response> { + ) -> Result< + ( + ExecutionResult, + Option, + Option, + ), + Response, + > { // Create RequestContext for this Responses request let mut ctx = RequestContext::for_responses( Arc::new(request.clone()), @@ -1181,8 +1203,9 @@ impl RequestPipeline { })?; let load_guards = ctx.state.load_guards.take(); + let adaptive_request = ctx.state.adaptive_request.take(); - Ok((execution_result, load_guards)) + Ok((execution_result, load_guards, adaptive_request)) } } @@ -1216,6 +1239,7 @@ mod build_parity_tests { (Endpoint::Chat, Mode::Regular) => ( v(&[ "ChatGeneratePreparationStage", + "AdaptiveAdmissionStage", "WorkerSelectionStage(Regular)", "ClientAcquisitionStage", "ChatGenerateRequestBuildingStage(ChatRequestBuildingStage(inject_pd_metadata=false, Single), GenerateRequestBuildingStage(inject_pd_metadata=false, Single))", @@ -1228,6 +1252,7 @@ mod build_parity_tests { (Endpoint::Chat, Mode::PrefillDecode) => ( v(&[ "ChatGeneratePreparationStage", + "AdaptiveAdmissionStage", "WorkerSelectionStage(PrefillDecode)", "ClientAcquisitionStage", "ChatGenerateRequestBuildingStage(ChatRequestBuildingStage(inject_pd_metadata=true, PrefillDecode), GenerateRequestBuildingStage(inject_pd_metadata=true, PrefillDecode))", @@ -1240,6 +1265,7 @@ mod build_parity_tests { (Endpoint::Chat, Mode::EncodePrefillDecode) => ( v(&[ "ChatGeneratePreparationStage", + "AdaptiveAdmissionStage", "WorkerSelectionStage(EncodePrefillDecode)", "ClientAcquisitionStage", "EncodeStage", @@ -1253,6 +1279,7 @@ mod build_parity_tests { (Endpoint::Messages, Mode::Regular) => ( v(&[ "MessagePreparationStage", + "AdaptiveAdmissionStage", "WorkerSelectionStage(Regular)", "ClientAcquisitionStage", "MessageRequestBuildingStage(inject_pd_metadata=false, Single)", @@ -1265,6 +1292,7 @@ mod build_parity_tests { (Endpoint::Messages, Mode::PrefillDecode) => ( v(&[ "MessagePreparationStage", + "AdaptiveAdmissionStage", "WorkerSelectionStage(PrefillDecode)", "ClientAcquisitionStage", "MessageRequestBuildingStage(inject_pd_metadata=true, PrefillDecode)", @@ -1277,6 +1305,7 @@ mod build_parity_tests { (Endpoint::Messages, Mode::EncodePrefillDecode) => ( v(&[ "MessagePreparationStage", + "AdaptiveAdmissionStage", "WorkerSelectionStage(EncodePrefillDecode)", "ClientAcquisitionStage", "EncodeStage", @@ -1290,6 +1319,7 @@ mod build_parity_tests { (Endpoint::Completion, Mode::Regular) => ( v(&[ "CompletionPreparationStage", + "AdaptiveAdmissionStage", "WorkerSelectionStage(Regular)", "ClientAcquisitionStage", "CompletionRequestBuildingStage(inject_pd_metadata=false, Single)", @@ -1302,6 +1332,7 @@ mod build_parity_tests { (Endpoint::Completion, Mode::PrefillDecode) => ( v(&[ "CompletionPreparationStage", + "AdaptiveAdmissionStage", "WorkerSelectionStage(PrefillDecode)", "ClientAcquisitionStage", "CompletionRequestBuildingStage(inject_pd_metadata=true, PrefillDecode)", @@ -1314,6 +1345,7 @@ mod build_parity_tests { (Endpoint::Completion, Mode::EncodePrefillDecode) => ( v(&[ "CompletionPreparationStage", + "AdaptiveAdmissionStage", "WorkerSelectionStage(EncodePrefillDecode)", "ClientAcquisitionStage", "EncodeStage", @@ -1327,6 +1359,7 @@ mod build_parity_tests { (Endpoint::Harmony, Mode::Regular) => ( v(&[ "HarmonyPreparationStage", + "AdaptiveAdmissionStage", "WorkerSelectionStage(Regular)", "ClientAcquisitionStage", "HarmonyRequestBuildingStage(inject_pd_metadata=false, Single)", @@ -1339,6 +1372,7 @@ mod build_parity_tests { (Endpoint::Harmony, Mode::PrefillDecode) => ( v(&[ "HarmonyPreparationStage", + "AdaptiveAdmissionStage", "WorkerSelectionStage(PrefillDecode)", "ClientAcquisitionStage", "HarmonyRequestBuildingStage(inject_pd_metadata=true, PrefillDecode)", @@ -1491,6 +1525,7 @@ mod alias_pipeline_tests { configured_tool_parser: None, configured_reasoning_parser: None, multimodal: None, + adaptive_admission: None, }); let request: GenerateRequest = serde_json::from_value(json!({ "model": MODEL_ALIAS, diff --git a/model_gateway/src/routers/grpc/regular/stages/chat/response_processing.rs b/model_gateway/src/routers/grpc/regular/stages/chat/response_processing.rs index 09a678ebc..efb2fc36d 100644 --- a/model_gateway/src/routers/grpc/regular/stages/chat/response_processing.rs +++ b/model_gateway/src/routers/grpc/regular/stages/chat/response_processing.rs @@ -107,6 +107,7 @@ impl ChatResponseProcessingStage { dispatch, tokenizer, skip_special_tokens, + ctx.state.adaptive_request.take(), ); // Attach load guards to response body for proper RAII lifecycle @@ -146,6 +147,12 @@ impl ChatResponseProcessingStage { ) .await?; + if let (Some(tracker), Some(usage)) = + (ctx.state.adaptive_request.take(), response.usage.as_ref()) + { + tracker.complete(usage.completion_tokens); + } + // Store the final response ctx.state.response.final_response = Some(FinalResponse::Chat(response)); diff --git a/model_gateway/src/routers/grpc/regular/stages/completion/response_processing.rs b/model_gateway/src/routers/grpc/regular/stages/completion/response_processing.rs index e373c3066..8514fe45d 100644 --- a/model_gateway/src/routers/grpc/regular/stages/completion/response_processing.rs +++ b/model_gateway/src/routers/grpc/regular/stages/completion/response_processing.rs @@ -88,6 +88,7 @@ impl PipelineStage for CompletionResponseProcessingStage { ctx.completion_request_arc(), dispatch, tokenizer, + ctx.state.adaptive_request.take(), ); let response = match ctx.state.load_guards.take() { @@ -123,6 +124,12 @@ impl PipelineStage for CompletionResponseProcessingStage { ) .await?; + if let (Some(tracker), Some(usage)) = + (ctx.state.adaptive_request.take(), response.usage.as_ref()) + { + tracker.complete(usage.completion_tokens); + } + ctx.state.response.final_response = Some(FinalResponse::Completion(response)); Ok(None) diff --git a/model_gateway/src/routers/grpc/regular/stages/generate/response_processing.rs b/model_gateway/src/routers/grpc/regular/stages/generate/response_processing.rs index b3849f737..154741b75 100644 --- a/model_gateway/src/routers/grpc/regular/stages/generate/response_processing.rs +++ b/model_gateway/src/routers/grpc/regular/stages/generate/response_processing.rs @@ -99,6 +99,7 @@ impl GenerateResponseProcessingStage { ctx.generate_request_arc(), // Cheap Arc clone (8 bytes) dispatch, tokenizer, + ctx.state.adaptive_request.take(), ); // Attach load guards to response body for proper RAII lifecycle @@ -137,6 +138,13 @@ impl GenerateResponseProcessingStage { ) .await?; + if let Some(tracker) = ctx.state.adaptive_request.take() { + let completion_tokens = result_array.iter().fold(0u32, |total, response| { + total.saturating_add(response.meta_info.completion_tokens) + }); + tracker.complete(completion_tokens); + } + // Store the final response ctx.state.response.final_response = Some(FinalResponse::Generate(result_array)); diff --git a/model_gateway/src/routers/grpc/regular/stages/messages/response_processing.rs b/model_gateway/src/routers/grpc/regular/stages/messages/response_processing.rs index fa5b4ce6b..98f1a0409 100644 --- a/model_gateway/src/routers/grpc/regular/stages/messages/response_processing.rs +++ b/model_gateway/src/routers/grpc/regular/stages/messages/response_processing.rs @@ -94,6 +94,7 @@ impl PipelineStage for MessageResponseProcessingStage { dispatch, tokenizer, skip_special_tokens, + ctx.state.adaptive_request.take(), ); // Attach load guards for RAII lifecycle @@ -130,6 +131,10 @@ impl PipelineStage for MessageResponseProcessingStage { ) .await?; + if let Some(tracker) = ctx.state.adaptive_request.take() { + tracker.complete(response.usage.output_tokens); + } + // Store the final response ctx.state.response.final_response = Some(FinalResponse::Messages(response)); diff --git a/model_gateway/src/routers/grpc/regular/streaming.rs b/model_gateway/src/routers/grpc/regular/streaming.rs index 505a8dc3b..5c7a3b396 100644 --- a/model_gateway/src/routers/grpc/regular/streaming.rs +++ b/model_gateway/src/routers/grpc/regular/streaming.rs @@ -39,6 +39,7 @@ use crate::{ routers::{ common::sse::SseEncoder, grpc::{ + adaptive_admission::AdaptiveRequestTracker, common::{response_formatting::CompletionTokenTracker, responses::build_sse_response}, context, proto_wrapper::{ProtoResponseVariant, ProtoStream}, @@ -120,6 +121,7 @@ impl StreamingProcessor { dispatch: context::DispatchMetadata, tokenizer: Arc, skip_special_tokens: bool, + adaptive_request: Option, ) -> Response { use bytes::Bytes; use tokio::sync::mpsc; @@ -157,8 +159,13 @@ impl StreamingProcessor { ) .await; - if let Err(e) = result { - utils::send_error_sse(&tx, &e, "internal_error"); + match result { + Ok(tokens) => { + if let Some(tracker) = adaptive_request { + tracker.complete(tokens); + } + } + Err(e) => utils::send_error_sse(&tx, &e, "internal_error"), } let _ = tx.send(Ok(Bytes::from("data: [DONE]\n\n"))); @@ -189,8 +196,13 @@ impl StreamingProcessor { ) .await; - if let Err(e) = result { - utils::send_error_sse(&tx, &e, "internal_error"); + match result { + Ok(tokens) => { + if let Some(tracker) = adaptive_request { + tracker.complete(tokens); + } + } + Err(e) => utils::send_error_sse(&tx, &e, "internal_error"), } let _ = tx.send(Ok(Bytes::from("data: [DONE]\n\n"))); @@ -228,7 +240,7 @@ impl StreamingProcessor { stop_params: (Option, Option>, bool, bool, bool), original_request: Arc, tx: &UnboundedSender>, - ) -> Result<(), String> { + ) -> Result { self.process_streaming_chunks_inner( grpc_stream, dispatch, @@ -254,7 +266,7 @@ impl StreamingProcessor { original_request: Arc, tx: &UnboundedSender>, pd_timing: Option, - ) -> Result<(), String> { + ) -> Result { // Metrics timing let start_time = Instant::now(); let mut first_token_time: Option = None; @@ -679,7 +691,7 @@ impl StreamingProcessor { output_tokens: total_completion as u64, }); - Ok(()) + Ok(total_completion) } /// Process prefill/decode streaming chunks (prefill + decode) - PD mode @@ -694,7 +706,7 @@ impl StreamingProcessor { original_request: Arc, tx: &UnboundedSender>, pd_timing: context::PdTiming, - ) -> Result<(), String> { + ) -> Result { // Phase 1.5: Collect input_logprobs from prefill stream if requested if original_request.logprobs { while let Some(response) = prefill_stream.next().await { @@ -747,6 +759,7 @@ impl StreamingProcessor { generate_request: Arc, dispatch: context::DispatchMetadata, tokenizer: Arc, + adaptive_request: Option, ) -> Response { // Create SSE channel let (tx, rx) = mpsc::unbounded_channel::>(); @@ -775,8 +788,13 @@ impl StreamingProcessor { let result = Self::process_generate_streaming(tokenizer, stream, ctx, &tx).await; - if let Err(e) = result { - utils::send_error_sse(&tx, &e, "internal_error"); + match result { + Ok(tokens) => { + if let Some(tracker) = adaptive_request { + tracker.complete(tokens); + } + } + Err(e) => utils::send_error_sse(&tx, &e, "internal_error"), } let _ = tx.send(Ok(Bytes::from("data: [DONE]\n\n"))); @@ -799,8 +817,13 @@ impl StreamingProcessor { ) .await; - if let Err(e) = result { - utils::send_error_sse(&tx, &e, "internal_error"); + match result { + Ok(tokens) => { + if let Some(tracker) = adaptive_request { + tracker.complete(tokens); + } + } + Err(e) => utils::send_error_sse(&tx, &e, "internal_error"), } let _ = tx.send(Ok(Bytes::from("data: [DONE]\n\n"))); @@ -836,7 +859,7 @@ impl StreamingProcessor { mut stream: ProtoStream, ctx: GenerateStreamContext, tx: &UnboundedSender>, - ) -> Result<(), String> { + ) -> Result { let start_time = Instant::now(); let mut first_token_time: Option = None; @@ -941,7 +964,7 @@ impl StreamingProcessor { let total_completion: u32 = completion_tokens_map.values().sum(); Self::record_generate_metrics(start_time, first_token_time, total_completion, &ctx); - Ok(()) + Ok(total_completion) } /// Process prefill/decode streaming for generate endpoint (PD mode with logprobs support) @@ -952,7 +975,7 @@ impl StreamingProcessor { ctx: GenerateStreamContext, tx: &UnboundedSender>, pd_timing: context::PdTiming, - ) -> Result<(), String> { + ) -> Result { // Collect input_logprobs from prefill stream if requested let input_token_logprobs = if ctx.return_logprob { let mut input_logprobs = None; @@ -1006,7 +1029,7 @@ impl StreamingProcessor { input_token_logprobs: Option>>>, tx: &UnboundedSender>, pd_timing: Option, - ) -> Result<(), String> { + ) -> Result { let start_time = Instant::now(); let mut first_token_time: Option = None; @@ -1161,7 +1184,7 @@ impl StreamingProcessor { let total_completion: u32 = completion_tokens_map.values().sum(); Self::record_generate_metrics(start_time, first_token_time, total_completion, &ctx); - Ok(()) + Ok(total_completion) } // ======================================================================== @@ -1576,6 +1599,7 @@ impl StreamingProcessor { dispatch: context::DispatchMetadata, tokenizer: Arc, skip_special_tokens: bool, + adaptive_request: Option, ) -> Response { let stop_params = ( messages_request @@ -1611,15 +1635,22 @@ impl StreamingProcessor { ) .await; - if let Err(e) = result { - let error_event = MessageStreamEvent::Error { - error: messages::ErrorResponse { - error_type: "api_error".to_string(), - message: e, - }, - }; - let mut buf = Vec::with_capacity(256); - let _ = Self::send_messages_event(&tx, &mut buf, &error_event); + match result { + Ok(tokens) => { + if let Some(tracker) = adaptive_request { + tracker.complete(tokens); + } + } + Err(e) => { + let error_event = MessageStreamEvent::Error { + error: messages::ErrorResponse { + error_type: "api_error".to_string(), + message: e, + }, + }; + let mut buf = Vec::with_capacity(256); + let _ = Self::send_messages_event(&tx, &mut buf, &error_event); + } } // No data: [DONE] — Anthropic uses message_stop instead }); @@ -1649,15 +1680,22 @@ impl StreamingProcessor { ) .await; - if let Err(e) = result { - let error_event = MessageStreamEvent::Error { - error: messages::ErrorResponse { - error_type: "api_error".to_string(), - message: e, - }, - }; - let mut buf = Vec::with_capacity(256); - let _ = Self::send_messages_event(&tx, &mut buf, &error_event); + match result { + Ok(tokens) => { + if let Some(tracker) = adaptive_request { + tracker.complete(tokens); + } + } + Err(e) => { + let error_event = MessageStreamEvent::Error { + error: messages::ErrorResponse { + error_type: "api_error".to_string(), + message: e, + }, + }; + let mut buf = Vec::with_capacity(256); + let _ = Self::send_messages_event(&tx, &mut buf, &error_event); + } } }); } @@ -1699,7 +1737,7 @@ impl StreamingProcessor { stop_params: (Option, Option>, bool, bool, bool), original_request: Arc, tx: &UnboundedSender>, - ) -> Result<(), String> { + ) -> Result { let start_time = Instant::now(); let mut first_token_time: Option = None; @@ -2310,7 +2348,7 @@ impl StreamingProcessor { output_tokens: u64::from(completion_tokens.total()), }); - Ok(()) + Ok(completion_tokens.total()) } /// Process prefill/decode streaming chunks for Messages API (PD mode). @@ -2327,7 +2365,7 @@ impl StreamingProcessor { stop_params: (Option, Option>, bool, bool, bool), original_request: Arc, tx: &UnboundedSender>, - ) -> Result<(), String> { + ) -> Result { // Consume prefill stream (Messages API does not expose prompt logprobs) while let Some(response) = prefill_stream.next().await { let gen_response = @@ -2372,6 +2410,7 @@ impl StreamingProcessor { completion_request: Arc, dispatch: context::DispatchMetadata, tokenizer: Arc, + adaptive_request: Option, ) -> Response { let (tx, rx) = mpsc::unbounded_channel::>(); @@ -2511,6 +2550,9 @@ impl StreamingProcessor { input_tokens: Some(total_prompt as u64), output_tokens: total_completion as u64, }); + if let Some(tracker) = adaptive_request { + tracker.complete(total_completion); + } } } diff --git a/model_gateway/src/routers/grpc/router.rs b/model_gateway/src/routers/grpc/router.rs index 9b84c4d4b..c68c3a5be 100644 --- a/model_gateway/src/routers/grpc/router.rs +++ b/model_gateway/src/routers/grpc/router.rs @@ -354,6 +354,7 @@ impl GrpcRouter { configured_tool_parser: ctx.configured_tool_parser.clone(), configured_reasoning_parser: ctx.configured_reasoning_parser.clone(), multimodal, + adaptive_admission: ctx.adaptive_admission.clone(), }); // Deps for the parser-consuming endpoints (chat/messages/harmony). diff --git a/model_gateway/src/service_discovery.rs b/model_gateway/src/service_discovery.rs index bacc6c3f1..3c806a514 100644 --- a/model_gateway/src/service_discovery.rs +++ b/model_gateway/src/service_discovery.rs @@ -1336,6 +1336,7 @@ mod tests { smg_data_connector::MemoryConversationItemStorage::new(), ), worker_monitor: None, + adaptive_admission: None, configured_reasoning_parser: None, configured_tool_parser: None, worker_job_queue: worker_job_queue.clone(), diff --git a/model_gateway/src/workflow/steps/local/drain_workers.rs b/model_gateway/src/workflow/steps/local/drain_workers.rs index 0ba0d69e0..ef6df9725 100644 --- a/model_gateway/src/workflow/steps/local/drain_workers.rs +++ b/model_gateway/src/workflow/steps/local/drain_workers.rs @@ -163,6 +163,7 @@ mod tests { smg_data_connector::MemoryConversationItemStorage::new(), ), worker_monitor: None, + adaptive_admission: None, configured_reasoning_parser: None, configured_tool_parser: None, worker_job_queue: Arc::clone(&job_queue),