From 0156cdc2ccb48a161b4e449159d97a3651babf51 Mon Sep 17 00:00:00 2001 From: Chang Su <8605658+CatherineSue@users.noreply.github.com> Date: Tue, 1 Sep 2026 11:33:25 -0700 Subject: [PATCH] feat(router): add --mm-per-request-image-limit override for spec image limits The multimodal model specs hardcode per-request modality limits (e.g. qwen3_vl caps images at 10) with no router-level override. A vLLM engine configured with --limit-mm-per-prompt.image 128 still gets every >10-image request rejected at the gateway with HTTP 400 invalid_multimodal_request, so the gateway silently contradicts the engine's configured capability. Add a --mm-per-request-image-limit flag (Rust binary and Python launcher) that replaces each spec's built-in image limit for all models, taking precedence over the SMG_IMAGE_MAX_COUNT env override. The registry trait gains validate_media_request_with_limits, an object-safe variant taking per-modality caller overrides, and the router passes its configured map through MultimodalComponents into the single validation call site. Unset keeps today's spec-default behavior; zero is rejected at clap parse time and in RouterConfig::validate. Signed-off-by: Chang Su <8605658+CatherineSue@users.noreply.github.com> --- bindings/python/src/lib.rs | 5 ++ bindings/python/src/smg/router_args.py | 12 ++++ bindings/python/tests/test_arg_parser.py | 9 +++ crates/multimodal/src/registry/traits.rs | 67 ++++++++++++++++++- model_gateway/src/config/builder.rs | 6 ++ model_gateway/src/config/types.rs | 5 ++ model_gateway/src/config/validation.rs | 26 +++++++ model_gateway/src/main.rs | 16 +++++ .../src/routers/grpc/multimodal/config.rs | 14 +++- .../src/routers/grpc/multimodal/plan.rs | 10 +-- model_gateway/src/routers/grpc/router.rs | 5 +- 11 files changed, 166 insertions(+), 9 deletions(-) diff --git a/bindings/python/src/lib.rs b/bindings/python/src/lib.rs index d72f326b38..2b1adb3e5a 100755 --- a/bindings/python/src/lib.rs +++ b/bindings/python/src/lib.rs @@ -527,6 +527,7 @@ struct Router { max_buffered_request_bytes: u64, kv_connector_annotation: String, kv_engine_id_annotation: String, + mm_per_request_image_limit: Option, } impl Router { @@ -918,6 +919,7 @@ impl Router { .stream_body_stall_timeout_secs(self.stream_body_stall_timeout_secs) .multimodal_tensor_transport(multimodal_tensor_transport) .multimodal_shm_min_bytes(self.multimodal_shm_min_bytes) + .mm_per_request_image_limit(self.mm_per_request_image_limit) .routing_key_override(config::RoutingKeyOverrideConfig { enabled: self.routing_key_override, eviction_interval_secs: self.eviction_interval_secs, @@ -1093,6 +1095,7 @@ impl Router { max_buffered_request_bytes = 1_048_576, kv_connector_annotation = String::from("smg.ai/kv-connector"), kv_engine_id_annotation = String::from("smg.ai/kv-engine-id"), + mm_per_request_image_limit = None, ))] #[expect(clippy::too_many_arguments)] #[expect( @@ -1245,6 +1248,7 @@ impl Router { max_buffered_request_bytes: u64, kv_connector_annotation: String, kv_engine_id_annotation: String, + mm_per_request_image_limit: Option, ) -> PyResult { let mut all_urls = worker_urls.clone(); @@ -1411,6 +1415,7 @@ impl Router { max_buffered_request_bytes, kv_connector_annotation, kv_engine_id_annotation, + mm_per_request_image_limit, }) } diff --git a/bindings/python/src/smg/router_args.py b/bindings/python/src/smg/router_args.py index dd52e7a2e4..6b01f19847 100644 --- a/bindings/python/src/smg/router_args.py +++ b/bindings/python/src/smg/router_args.py @@ -249,6 +249,8 @@ class RouterArgs: max_buffered_request_bytes: int = 1048576 kv_connector_annotation: str = "smg.ai/kv-connector" kv_engine_id_annotation: str = "smg.ai/kv-engine-id" + # Per-request image-count limit replacing model spec limits; None keeps spec limits + mm_per_request_image_limit: int | None = None @staticmethod def add_cli_args( @@ -879,6 +881,16 @@ def add_cli_args( default=RouterArgs.multimodal_shm_min_bytes, help="Minimum multimodal tensor size (bytes) before the SHM transport is used", ) + parser.add_argument( + f"--{prefix}mm-per-request-image-limit", + type=int, + default=RouterArgs.mm_per_request_image_limit, + help=( + "Per-request image-count limit applied to every model, replacing the" + " model spec's built-in limit (e.g. to match the engine's" + " --limit-mm-per-prompt). Must be >= 1; unset keeps spec limits." + ), + ) # Logging configuration logging_group.add_argument( diff --git a/bindings/python/tests/test_arg_parser.py b/bindings/python/tests/test_arg_parser.py index 2886fac8f1..3ba1491df9 100644 --- a/bindings/python/tests/test_arg_parser.py +++ b/bindings/python/tests/test_arg_parser.py @@ -539,6 +539,14 @@ def test_parse_basic_args(self): assert router_args.worker_urls == ["http://worker1:8000", "http://worker2:8000"] assert router_args.policy == "round_robin" + def test_parse_mm_per_request_image_limit(self): + """Deployment-wide image-count limit; unset keeps model spec limits.""" + router_args = parse_router_args(["--mm-per-request-image-limit", "128"]) + assert router_args.mm_per_request_image_limit == 128 + + defaults = parse_router_args([]) + assert defaults.mm_per_request_image_limit is None + def test_parse_routing_key_headers(self): """Ordered list flag; unset keeps the x-smg-routing-key default.""" router_args = parse_router_args( @@ -1431,6 +1439,7 @@ class TestRouterArgsFieldOrder: "max_buffered_request_bytes", "kv_connector_annotation", "kv_engine_id_annotation", + "mm_per_request_image_limit", ] def test_complete_field_sequence_is_frozen(self): diff --git a/crates/multimodal/src/registry/traits.rs b/crates/multimodal/src/registry/traits.rs index b9c9f892a6..b91abf131d 100644 --- a/crates/multimodal/src/registry/traits.rs +++ b/crates/multimodal/src/registry/traits.rs @@ -218,9 +218,25 @@ pub trait ModelProcessorSpec: Send + Sync { &self, metadata: &ModelMetadata, requested: &[(Modality, usize)], + ) -> RegistryResult<()> { + self.validate_media_request_with_limits(metadata, requested, &HashMap::new()) + } + + /// [`Self::validate_media_request`] with caller-supplied per-modality limit + /// overrides; each replaces the spec limit and beats the env override. + fn validate_media_request_with_limits( + &self, + metadata: &ModelMetadata, + requested: &[(Modality, usize)], + limit_overrides: &HashMap, ) -> RegistryResult<()> { let limits = self.modality_limits(metadata)?; - check_media_counts(self.name(), &limits, requested, modality_limit_override) + check_media_counts(self.name(), &limits, requested, |modality| { + limit_overrides + .get(&modality) + .copied() + .or_else(|| modality_limit_override(modality)) + }) } fn processor_kwargs(&self, metadata: &ModelMetadata) -> RegistryResult; @@ -408,6 +424,55 @@ mod tests { ); } + #[test] + fn caller_limit_override_replaces_spec_limit() { + let tokenizer = TestTokenizer::new(&[]); + let config = json!({}); + let metadata = ModelMetadata { + model_id: "test-model", + tokenizer: &tokenizer, + config: &config, + }; + let overrides = HashMap::from([(Modality::Image, 5)]); + + // TestSpec declares Image=2; the caller override wins in both directions. + assert_eq!( + TestSpec.validate_media_request_with_limits( + &metadata, + &[(Modality::Image, 5)], + &overrides + ), + Ok(()) + ); + assert_eq!( + TestSpec.validate_media_request_with_limits( + &metadata, + &[(Modality::Image, 6)], + &overrides + ), + Err(ModelRegistryError::ModalityLimitExceeded { + spec: "test", + modality: Modality::Image, + limit: 5, + requested: 6, + }) + ); + // Unoverridden modalities keep the spec limit. + assert_eq!( + TestSpec.validate_media_request_with_limits( + &metadata, + &[(Modality::Audio, 2)], + &overrides + ), + Err(ModelRegistryError::ModalityLimitExceeded { + spec: "test", + modality: Modality::Audio, + limit: 1, + requested: 2, + }) + ); + } + #[test] fn limit_override_raises_declared_limit() { let limits = HashMap::from([(Modality::Image, 2)]); diff --git a/model_gateway/src/config/builder.rs b/model_gateway/src/config/builder.rs index 352224f238..6dd32c9cc7 100644 --- a/model_gateway/src/config/builder.rs +++ b/model_gateway/src/config/builder.rs @@ -310,6 +310,12 @@ impl RouterConfigBuilder { self } + /// Per-request image-count limit replacing each model spec's built-in limit. + pub fn mm_per_request_image_limit(mut self, limit: Option) -> Self { + self.config.mm_per_request_image_limit = limit; + self + } + // ==================== Rate Limiting ==================== pub fn max_concurrent_requests(mut self, max: i32) -> Self { diff --git a/model_gateway/src/config/types.rs b/model_gateway/src/config/types.rs index e849414189..455489399c 100755 --- a/model_gateway/src/config/types.rs +++ b/model_gateway/src/config/types.rs @@ -147,6 +147,10 @@ pub struct RouterConfig { /// to `SMG_MM_SHM_MIN_BYTES`, then 64 KiB. #[serde(default, skip_serializing_if = "Option::is_none")] pub multimodal_shm_min_bytes: Option, + /// Per-request image-count limit applied to every model, replacing each + /// spec's built-in limit; beats `SMG_IMAGE_MAX_COUNT`. Unset keeps spec limits. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub mm_per_request_image_limit: Option, pub dp_aware: bool, #[serde(default)] pub dp_minimum_tokens_scheduler: bool, @@ -1094,6 +1098,7 @@ impl Default for RouterConfig { engine_metrics: false, multimodal_tensor_transport: None, multimodal_shm_min_bytes: None, + mm_per_request_image_limit: None, dp_aware: false, dp_minimum_tokens_scheduler: false, api_key: None, diff --git a/model_gateway/src/config/validation.rs b/model_gateway/src/config/validation.rs index b5281e01e3..028e08ae81 100644 --- a/model_gateway/src/config/validation.rs +++ b/model_gateway/src/config/validation.rs @@ -732,6 +732,14 @@ impl ConfigValidator { }); } + if config.mm_per_request_image_limit == Some(0) { + return Err(ConfigError::InvalidValue { + field: "mm_per_request_image_limit".to_string(), + value: "0".to_string(), + reason: "Must be at least 1".to_string(), + }); + } + // A zero-capacity job channel panics at construction and a // zero-permit dispatcher never dequeues; reject both here so a // config-file value fails as early as the CLI parsers do. @@ -1213,6 +1221,24 @@ mod tests { )); } + #[test] + fn zero_mm_per_request_image_limit_is_rejected() { + let config = RouterConfig { + mm_per_request_image_limit: Some(0), + ..Default::default() + }; + assert!(matches!( + ConfigValidator::validate(&config), + Err(ConfigError::InvalidValue { ref field, .. }) if field == "mm_per_request_image_limit" + )); + + let config = RouterConfig { + mm_per_request_image_limit: Some(128), + ..Default::default() + }; + assert!(ConfigValidator::validate(&config).is_ok()); + } + #[test] fn mesh_server_name_with_colon_is_rejected() { assert!(matches!( diff --git a/model_gateway/src/main.rs b/model_gateway/src/main.rs index 4d913fe8c3..5680094e76 100644 --- a/model_gateway/src/main.rs +++ b/model_gateway/src/main.rs @@ -557,6 +557,11 @@ struct CliArgs { #[arg(long, help_heading = "Multimodal")] multimodal_shm_min_bytes: Option, + /// Per-request image-count limit applied to every model, replacing each + /// spec's built-in limit (e.g. to match the engine's `--limit-mm-per-prompt`). + #[arg(long, value_parser = clap::value_parser!(u64).range(1..), help_heading = "Multimodal")] + mm_per_request_image_limit: Option, + // ==================== Service Discovery (Kubernetes) ==================== /// Enable Kubernetes service discovery #[arg( @@ -1816,6 +1821,7 @@ impl CliArgs { .engine_metrics(self.engine_metrics) .multimodal_tensor_transport(self.multimodal_tensor_transport) .multimodal_shm_min_bytes(self.multimodal_shm_min_bytes) + .mm_per_request_image_limit(self.mm_per_request_image_limit.map(|v| v as usize)) .max_concurrent_requests(self.max_concurrent_requests) .queue_size(self.queue_size) .queue_timeout_secs(self.queue_timeout_secs) @@ -2689,6 +2695,8 @@ mod tests { "shm", "--multimodal-shm-min-bytes", "1024", + "--mm-per-request-image-limit", + "128", ]); let router_config = cli.to_router_config(vec![], vec![]).unwrap(); @@ -2698,6 +2706,7 @@ mod tests { "transport mode must reach RouterConfig via to_router_config" ); assert_eq!(router_config.multimodal_shm_min_bytes, Some(1024)); + assert_eq!(router_config.mm_per_request_image_limit, Some(128)); let server_config = cli.to_server_config(router_config).unwrap(); assert_eq!( @@ -2709,6 +2718,13 @@ mod tests { server_config.router_config.multimodal_shm_min_bytes, Some(1024) ); + assert_eq!( + server_config.router_config.mm_per_request_image_limit, + Some(128) + ); + + // clap rejects a zero limit outright. + assert!(Cli::try_parse_from(["smg", "--mm-per-request-image-limit", "0"]).is_err()); } /// Default is off: the flag stays false through both conversions so diff --git a/model_gateway/src/routers/grpc/multimodal/config.rs b/model_gateway/src/routers/grpc/multimodal/config.rs index 3244de5eca..32ddeb46fa 100644 --- a/model_gateway/src/routers/grpc/multimodal/config.rs +++ b/model_gateway/src/routers/grpc/multimodal/config.rs @@ -1,12 +1,12 @@ //! Multimodal model configuration: the shared config-file registry and the //! per-router component bundle (media connector + processor/model registries). -use std::{path::Path, sync::Arc}; +use std::{collections::HashMap, path::Path, sync::Arc}; use anyhow::{Context, Result}; use dashmap::DashMap; use llm_multimodal::{ - MediaConnector, MediaConnectorConfig, ModelRegistry, PreProcessorConfig, + MediaConnector, MediaConnectorConfig, Modality, ModelRegistry, PreProcessorConfig, VisionProcessorRegistry, }; use tracing::{debug, warn}; @@ -241,12 +241,17 @@ pub(crate) struct MultimodalComponents { pub config_registry: Arc, /// Optional host-DRAM cache of preprocessed per-image encoder inputs. pub pixel_cache: Option>, + /// Router-configured per-modality media-count limits replacing spec limits. + pub modality_limit_overrides: HashMap, } impl MultimodalComponents { /// Create multimodal components with default registries and a reference /// to the shared `MultimodalConfigRegistry` owned by `AppContext`. - pub fn new(config_registry: Arc) -> Result { + pub fn new( + config_registry: Arc, + image_limit_override: Option, + ) -> Result { let client = reqwest::Client::builder() .timeout(std::time::Duration::from_secs(30)) .build() @@ -260,6 +265,9 @@ impl MultimodalComponents { model_registry: Arc::new(ModelRegistry::default()), config_registry, pixel_cache: pixel_cache_from_env(), + modality_limit_overrides: image_limit_override + .map(|limit| HashMap::from([(Modality::Image, limit)])) + .unwrap_or_default(), }) } } diff --git a/model_gateway/src/routers/grpc/multimodal/plan.rs b/model_gateway/src/routers/grpc/multimodal/plan.rs index ffc85d9da7..541b469e81 100644 --- a/model_gateway/src/routers/grpc/multimodal/plan.rs +++ b/model_gateway/src/routers/grpc/multimodal/plan.rs @@ -120,10 +120,12 @@ pub(crate) async fn prepare_placeholder_tokens( .iter() .map(|&modality| (modality, plan.count(modality))) .collect::>(); - spec.validate_media_request(&metadata, &requested) - .map_err(|error| { - anyhow::anyhow!("invalid media request for model {}: {error}", spec.name()) - })?; + spec.validate_media_request_with_limits( + &metadata, + &requested, + &components.modality_limit_overrides, + ) + .map_err(|error| anyhow::anyhow!("invalid media request for model {}: {error}", spec.name()))?; let mut placeholders = PlaceholderTokens::default(); for &modality in plan.modalities() { let token = spec diff --git a/model_gateway/src/routers/grpc/router.rs b/model_gateway/src/routers/grpc/router.rs index 6ad0df17be..7b4fe4baf1 100644 --- a/model_gateway/src/routers/grpc/router.rs +++ b/model_gateway/src/routers/grpc/router.rs @@ -337,7 +337,10 @@ impl GrpcRouter { let policy_registry = ctx.policy_registry.clone(); // Create multimodal components (best-effort; non-fatal if initialization fails) - let multimodal = match MultimodalComponents::new(ctx.multimodal_config_registry.clone()) { + let multimodal = match MultimodalComponents::new( + ctx.multimodal_config_registry.clone(), + ctx.router_config.mm_per_request_image_limit, + ) { Ok(mc) => Some(Arc::new(mc)), Err(e) => { tracing::warn!("Multimodal components initialization failed (non-fatal): {e}");