From 2cddd26f467eee4fd39d5b22f1259ee5387fc02b Mon Sep 17 00:00:00 2001 From: XinyueZhang369 Date: Mon, 5 Oct 2026 14:24:28 -0700 Subject: [PATCH 1/2] feat(discovery)!: tagged discovery config with one runtime conversion `RouterConfig.discovery` becomes a tagged provider configuration, and the CLI and the Python binding now reach the runtime through one conversion. - `DiscoveryConfig` is an enum tagged by `provider`, with one variant, `Kubernetes(KubernetesDiscoveryConfig)`: the old fields minus `enabled`. Presence means enabled. The legacy flat object still reads: no `provider` means Kubernetes, and `enabled: false` becomes `None` at deserialization, so writing the config back cannot turn discovery on. Canonical output writes `provider` and no `enabled`. - `RuntimeDiscoveryConfig::from_config` is the one config-to-runtime conversion. `main.rs` and the Python binding used to hand-build the runtime config beside the serialized one, the duplication that once let the KV annotation flags reach only one of the two. Mesh-router config is likewise derived once, by `MeshDiscoveryConfig::from_discovery`. - `--discovery-provider kubernetes` selects the provider, with `--service-discovery` kept as its legacy spelling; giving both is a usage error. IGW auto-enable and the worker-auto-recovery default follow the selected provider, whichever spelling chose it. The Python CLI mirrors this, and the binding gains a keyword-only `discovery=` mapping read by the same rules as `RouterConfig.discovery`. - Validation dispatches per provider; the Kubernetes checks are unchanged. - The in-cluster kind gateway now starts through `--discovery-provider kubernetes`; the host-process journey keeps `--service-discovery`. Breaking for Rust callers: construct `DiscoveryConfig::Kubernetes(..)` (or `.into()` a `KubernetesDiscoveryConfig`) instead of a struct literal; `ServiceDiscoveryConfig` loses `enabled`; `ServerConfig::service_discovery_config` is now `Option`. YAML read into `RouterConfig` must quote label values that look like numbers or booleans (`version: "1"`), as a Kubernetes manifest already must. Signed-off-by: XinyueZhang369 --- bindings/python/src/lib.rs | 111 ++++---- bindings/python/src/smg/launch_router.py | 2 +- bindings/python/src/smg/router.py | 7 + bindings/python/src/smg/router_args.py | 37 ++- bindings/python/tests/test_arg_parser.py | 37 +++ bindings/python/tests/test_router_config.py | 34 +++ .../python/tests/test_startup_sequence.py | 36 ++- e2e_test/kind_discovery/conftest.py | 2 + e2e_test/kind_discovery/in_cluster.yaml | 5 +- model_gateway/src/config/builder.rs | 17 +- model_gateway/src/config/types.rs | 248 ++++++++++++++--- model_gateway/src/config/validation.rs | 50 ++-- model_gateway/src/main.rs | 263 +++++++++++------- .../src/mesh_discovery/kubernetes.rs | 15 + model_gateway/src/server.rs | 28 +- .../src/service_discovery/kubernetes.rs | 15 +- model_gateway/src/service_discovery/mod.rs | 8 +- .../src/service_discovery/runtime.rs | 177 ++++++++++++ model_gateway/tests/k8s_discovery_test.rs | 1 - 19 files changed, 837 insertions(+), 256 deletions(-) create mode 100644 model_gateway/src/service_discovery/runtime.rs diff --git a/bindings/python/src/lib.rs b/bindings/python/src/lib.rs index 525b6521cd..2d2cafae6b 100755 --- a/bindings/python/src/lib.rs +++ b/bindings/python/src/lib.rs @@ -549,6 +549,23 @@ struct Router { prefill_max_inflight_requests_per_worker: i32, prefill_queue_size: Option, prefill_queue_timeout_secs: Option, + /// The keyword-only `discovery` mapping, read by the same rules as + /// `RouterConfig.discovery`. + discovery: Option, +} + +/// Read the keyword-only `discovery` mapping by the same rules as +/// `RouterConfig.discovery`. It goes through JSON, so nested mappings and +/// lists convert without a hand-written walker. +fn parse_discovery(mapping: &Bound<'_, PyAny>) -> PyResult> { + let invalid = |e: serde_json::Error| { + pyo3::exceptions::PyValueError::new_err(format!("Invalid discovery mapping: {e}")) + }; + let json: String = PyModule::import(mapping.py(), "json")? + .call_method1("dumps", (mapping,))? + .extract()?; + let value: serde_json::Value = serde_json::from_str(&json).map_err(invalid)?; + config::deserialize_discovery(value).map_err(invalid) } impl Router { @@ -612,7 +629,8 @@ impl Router { pub fn to_router_config(&self) -> config::ConfigResult { use config::{ - DiscoveryConfig, MetricsConfig, PolicyConfig as ConfigPolicyConfig, RoutingMode, + DiscoveryConfig, KubernetesDiscoveryConfig, MetricsConfig, + PolicyConfig as ConfigPolicyConfig, RoutingMode, }; // Validate the transport mode up front. The CLI (value_parser) and the @@ -755,8 +773,7 @@ impl Router { let policy = convert_policy(&self.policy)?; let discovery = if self.service_discovery { - Some(DiscoveryConfig { - enabled: true, + Some(DiscoveryConfig::Kubernetes(KubernetesDiscoveryConfig { namespace: self.service_discovery_namespace.clone(), port: self.service_discovery_port, check_interval_secs: 60, @@ -771,10 +788,11 @@ impl Router { router_selector: self.router_selector.clone(), router_mesh_port_annotation: "sglang.ai/mesh-port".to_string(), model_id_source: self.model_id_from.clone(), - }) + })) } else { - None + self.discovery.clone() }; + let has_discovery = discovery.is_some(); let metrics = match (self.prometheus_port, self.prometheus_host.as_ref()) { (Some(port), Some(host)) => Some(MetricsConfig { @@ -915,7 +933,7 @@ impl Router { // service discovery, which is what re-adds a removed worker. remove_unhealthy_workers: config::resolve_worker_auto_recovery( self.remove_unhealthy_workers, - self.service_discovery, + has_discovery, ), drain_settle_secs: self.drain_settle_secs, }) @@ -1166,12 +1184,11 @@ impl Router { prefill_max_inflight_requests_per_worker = -1, prefill_queue_size = None, prefill_queue_timeout_secs = None, + // Keyword-only, so it never takes a positional slot. + *, + discovery = None, ))] #[expect(clippy::too_many_arguments)] - #[expect( - clippy::unnecessary_wraps, - reason = "PyO3 #[new] method signature requires PyResult" - )] fn new( worker_urls: Vec, policy: PolicyType, @@ -1335,7 +1352,20 @@ impl Router { prefill_max_inflight_requests_per_worker: i32, prefill_queue_size: Option, prefill_queue_timeout_secs: Option, + discovery: Option>, ) -> PyResult { + // Two spellings of one choice: refuse both rather than pick one. + if service_discovery && discovery.is_some() { + return Err(pyo3::exceptions::PyValueError::new_err( + "Pass service_discovery=True or a discovery mapping, not both", + )); + } + let discovery = discovery + .as_ref() + .map(parse_discovery) + .transpose()? + .flatten(); + let mut all_urls = worker_urls.clone(); if let Some(ref encode_urls) = encode_urls { @@ -1516,6 +1546,7 @@ impl Router { prefill_max_inflight_requests_per_worker, prefill_queue_size, prefill_queue_timeout_secs, + discovery, }) } @@ -1530,51 +1561,23 @@ impl Router { pyo3::exceptions::PyValueError::new_err(format!("Configuration validation failed: {e}")) })?; - let model_id_source = self - .model_id_from - .as_deref() - .map(|s| { - service_discovery::ModelIdSource::parse(s).map_err(|e| { - pyo3::exceptions::PyValueError::new_err(format!( - "Invalid --model-id-from value '{s}': {e}" - )) - }) - }) - .transpose()?; - - let service_discovery_config = if self.service_discovery { - Some(service_discovery::ServiceDiscoveryConfig { - enabled: true, - selector: self.selector.clone(), - check_interval: std::time::Duration::from_secs(60), - port: self.service_discovery_port, - namespace: self.service_discovery_namespace.clone(), - disaggregated_mode: self.pd_disaggregation || self.epd_disaggregation, - encode_selector: self.encode_selector.clone(), - prefill_selector: self.prefill_selector.clone(), - decode_selector: self.decode_selector.clone(), - bootstrap_port_annotation: self.bootstrap_port_annotation.clone(), - worker_ports_annotation: self.worker_ports_annotation.clone(), - kv_connector_annotation: self.kv_connector_annotation.clone(), - kv_engine_id_annotation: self.kv_engine_id_annotation.clone(), - model_id_source, - }) - } else { - None - }; - - // Mesh-router discovery now has its own task and lifetime, but stays - // gated on the legacy service-discovery flag until it gains its own - // config surface with the tagged provider configuration. - let mesh_discovery_config = if self.service_discovery && !self.router_selector.is_empty() { - Some(mesh_discovery::MeshDiscoveryConfig { - namespace: self.service_discovery_namespace.clone(), - router_selector: self.router_selector.clone(), - router_mesh_port_annotation: "sglang.ai/mesh-port".to_string(), + let service_discovery_config = router_config + .discovery + .as_ref() + .map(|discovery| { + service_discovery::RuntimeDiscoveryConfig::from_config( + discovery, + &router_config.mode, + ) }) - } else { - None - }; + .transpose() + .map_err(|e| { + pyo3::exceptions::PyValueError::new_err(format!("Configuration error: {e}")) + })?; + let mesh_discovery_config = router_config + .discovery + .as_ref() + .and_then(mesh_discovery::MeshDiscoveryConfig::from_discovery); let prometheus_config = Some(PrometheusConfig { port: self.prometheus_port.unwrap_or(29000), diff --git a/bindings/python/src/smg/launch_router.py b/bindings/python/src/smg/launch_router.py index 3e9af08814..38bafa4966 100644 --- a/bindings/python/src/smg/launch_router.py +++ b/bindings/python/src/smg/launch_router.py @@ -40,7 +40,7 @@ def launch_router(args: argparse.Namespace | RouterArgs) -> None: if Router is None: raise RuntimeError("Rust Router is not installed") - if router_args.service_discovery and not router_args.enable_igw: + if router_args.selected_discovery_provider() and not router_args.enable_igw: logger.info("IGW mode automatically enabled because service discovery is turned on") router_args.enable_igw = True diff --git a/bindings/python/src/smg/router.py b/bindings/python/src/smg/router.py index bc72caa2e0..cb799d8bc7 100644 --- a/bindings/python/src/smg/router.py +++ b/bindings/python/src/smg/router.py @@ -228,6 +228,10 @@ class Router: service_discovery: Enable Kubernetes service discovery. When enabled, the router will automatically discover worker pods based on the selector. Default: False + discovery: The worker discovery provider's configuration, keyword-only: + a mapping with a ``provider`` key and that provider's fields, read + like ``RouterConfig.discovery`` (no ``provider`` means Kubernetes). + Mutually exclusive with ``service_discovery``. Default: None selector: Dictionary mapping of label keys to values for Kubernetes pod selection. Example: {"app": "sglang-worker"}. Default: {} service_discovery_port: Port to use for service discovery. The router will @@ -313,6 +317,9 @@ def from_args(args: RouterArgs) -> Router: """Create a router from a RouterArgs instance.""" args_dict = vars(args).copy() + # Kubernetes, by either spelling, reaches Rust as service_discovery. + args_dict["service_discovery"] = args.selected_discovery_provider() == "kubernetes" + args_dict.pop("discovery_provider") # Convert RouterArgs to _Router parameters args_dict["worker_urls"] = ( [] diff --git a/bindings/python/src/smg/router_args.py b/bindings/python/src/smg/router_args.py index edf0a6537d..76c0cde0dd 100644 --- a/bindings/python/src/smg/router_args.py +++ b/bindings/python/src/smg/router_args.py @@ -24,6 +24,8 @@ PREFILL_POLICY_CHOICES = [*COMMON_POLICY_CHOICES, "bucket"] ENCODE_POLICY_CHOICES = ["random", "round_robin", "consistent_hashing"] +# Worker discovery providers --discovery-provider accepts. +DISCOVERY_PROVIDER_CHOICES = ["kubernetes"] def _parse_int_csv(value: str) -> list[int]: @@ -286,6 +288,9 @@ class RouterArgs: prefill_max_inflight_requests_per_worker: int = -1 prefill_queue_size: int | None = None prefill_queue_timeout_secs: int | None = None + # The worker discovery provider. service_discovery=True is the legacy + # spelling of discovery_provider="kubernetes"; set one or the other. + discovery_provider: str | None = None @staticmethod def add_cli_args( @@ -1079,10 +1084,21 @@ def add_cli_args( ) # Service discovery configuration - k8s_group.add_argument( + discovery_selection = k8s_group.add_mutually_exclusive_group() + discovery_selection.add_argument( f"--{prefix}service-discovery", action="store_true", - help="Enable Kubernetes service discovery", + help=( + "Enable Kubernetes service discovery (the legacy spelling of" + " --discovery-provider kubernetes)" + ), + ) + discovery_selection.add_argument( + f"--{prefix}discovery-provider", + type=str, + choices=DISCOVERY_PROVIDER_CHOICES, + default=None, + help="Worker discovery provider. Give this or --service-discovery, not both", ) k8s_group.add_argument( f"--{prefix}selector", @@ -1879,6 +1895,23 @@ def from_cli_args(cls, args: argparse.Namespace, use_router_prefix: bool = False return cls(**args_dict) + def selected_discovery_provider(self) -> str | None: + """The worker discovery provider selected by either spelling. + + ``service_discovery=True`` is ``discovery_provider="kubernetes"``. + Setting both is an error rather than a precedence rule, as on the CLI. + """ + if self.discovery_provider is not None: + if self.service_discovery: + raise ValueError("Set service_discovery=True or discovery_provider, not both") + if self.discovery_provider not in DISCOVERY_PROVIDER_CHOICES: + raise ValueError( + f"Unknown discovery provider {self.discovery_provider!r};" + f" expected one of {DISCOVERY_PROVIDER_CHOICES}" + ) + return self.discovery_provider + return "kubernetes" if self.service_discovery else None + def _validate_router_args(self): # Validate configuration based on mode if self.epd_disaggregation: diff --git a/bindings/python/tests/test_arg_parser.py b/bindings/python/tests/test_arg_parser.py index e7682a16e6..c7ca855721 100644 --- a/bindings/python/tests/test_arg_parser.py +++ b/bindings/python/tests/test_arg_parser.py @@ -35,6 +35,8 @@ def test_default_values(self): # Test service discovery defaults assert args.service_discovery is False + assert args.discovery_provider is None + assert args.selected_discovery_provider() is None assert args.selector == {} assert args.service_discovery_port == 80 assert args.service_discovery_namespace is None @@ -861,6 +863,40 @@ def test_parse_pd_prefill_admission_args(self): assert router_args.prefill_queue_size == 13 assert router_args.prefill_queue_timeout_secs == 17 + def test_parse_discovery_provider_kubernetes(self): + """--discovery-provider kubernetes selects Kubernetes without the legacy flag.""" + router_args = parse_router_args( + ["--discovery-provider", "kubernetes", "--selector", "app=worker"] + ) + + assert router_args.discovery_provider == "kubernetes" + assert router_args.service_discovery is False + assert router_args.selected_discovery_provider() == "kubernetes" + assert router_args.selector == {"app": "worker"} + + def test_service_discovery_and_discovery_provider_are_exclusive(self): + """Both spellings of one choice is a usage error, not a precedence rule.""" + with pytest.raises(SystemExit): + parse_router_args(["--service-discovery", "--discovery-provider", "kubernetes"]) + + def test_unknown_discovery_provider_is_rejected(self): + with pytest.raises(SystemExit): + parse_router_args(["--discovery-provider", "zookeeper"]) + + def test_selected_discovery_provider_programmatic(self): + """RouterArgs built in code follows the CLI's rules.""" + assert RouterArgs(service_discovery=True).selected_discovery_provider() == "kubernetes" + assert ( + RouterArgs(discovery_provider="kubernetes").selected_discovery_provider() + == "kubernetes" + ) + with pytest.raises(ValueError, match="not both"): + RouterArgs( + service_discovery=True, discovery_provider="kubernetes" + ).selected_discovery_provider() + with pytest.raises(ValueError, match="Unknown discovery provider"): + RouterArgs(discovery_provider="zookeeper").selected_discovery_provider() + def test_parse_service_discovery_args(self): """Test parsing service discovery arguments.""" args_a = [ @@ -1541,6 +1577,7 @@ class TestRouterArgsFieldOrder: "prefill_max_inflight_requests_per_worker", "prefill_queue_size", "prefill_queue_timeout_secs", + "discovery_provider", ] def test_complete_field_sequence_is_frozen(self): diff --git a/bindings/python/tests/test_router_config.py b/bindings/python/tests/test_router_config.py index 3d4d658c7b..5564292bc3 100644 --- a/bindings/python/tests/test_router_config.py +++ b/bindings/python/tests/test_router_config.py @@ -11,6 +11,7 @@ from smg.launch_router import RouterArgs, launch_router from smg.router import policy_from_str from smg.smg_rs import PolicyType +from smg.smg_rs import Router as _Router class TestRouterConfigValidation: @@ -435,3 +436,36 @@ def test_config_with_empty_dicts(self): assert args.prefill_selector == {} assert args.decode_selector == {} assert args.storage_context_headers == {} + + +class TestDiscoveryMapping: + """The binding's keyword-only `discovery` mapping.""" + + @staticmethod + def kubernetes_fields(): + return { + "port": 8000, + "check_interval_secs": 60, + "selector": {"app": "worker"}, + "prefill_selector": {}, + "decode_selector": {}, + "bootstrap_port_annotation": "sglang.ai/bootstrap-port", + } + + def test_tagged_kubernetes_mapping_is_accepted(self): + _Router(worker_urls=[], discovery={"provider": "kubernetes", **self.kubernetes_fields()}) + + def test_untagged_mapping_reads_as_kubernetes(self): + _Router(worker_urls=[], discovery=self.kubernetes_fields()) + + def test_mapping_and_service_discovery_conflict(self): + with pytest.raises(ValueError, match="not both"): + _Router( + worker_urls=[], + service_discovery=True, + discovery=self.kubernetes_fields(), + ) + + def test_unknown_provider_is_rejected(self): + with pytest.raises(ValueError, match="Invalid discovery mapping"): + _Router(worker_urls=[], discovery={"provider": "zookeeper"}) diff --git a/bindings/python/tests/test_startup_sequence.py b/bindings/python/tests/test_startup_sequence.py index 3b908dd91b..435cb336d6 100644 --- a/bindings/python/tests/test_startup_sequence.py +++ b/bindings/python/tests/test_startup_sequence.py @@ -10,7 +10,7 @@ import pytest from smg.launch_router import RouterArgs, launch_router -from smg.router import policy_from_str +from smg.router import Router, policy_from_str # Local helper mirroring the router logger setup used in production @@ -261,6 +261,40 @@ def fake_from_args(router_args): assert captured_args["enable_igw"] is True mock_router_instance.start.assert_called_once() + def test_discovery_provider_kubernetes_enables_igw(self): + """IGW follows the selected provider, whichever spelling selected it.""" + args = RouterArgs(discovery_provider="kubernetes", selector={"app": "worker"}) + + with patch("smg.launch_router.Router") as router_mod: + captured_args = {} + + def fake_from_args(router_args): + captured_args["enable_igw"] = router_args.enable_igw + return MagicMock() + + router_mod.from_args = MagicMock(side_effect=fake_from_args) + + launch_router(args) + + assert captured_args["enable_igw"] is True + + def test_from_args_passes_discovery_provider_as_service_discovery(self): + """The Rust binding receives Kubernetes discovery by its legacy name.""" + args = RouterArgs( + discovery_provider="kubernetes", + selector={"app": "worker"}, + worker_urls=["http://worker:8000"], + ) + + with patch("smg.router._Router") as rust_router: + Router.from_args(args) + + kwargs = rust_router.call_args.kwargs + assert kwargs["service_discovery"] is True + assert "discovery_provider" not in kwargs + # Discovery supplies the workers, exactly as under --service-discovery. + assert kwargs["worker_urls"] == [] + def test_router_initialization_with_retry_config(self): """Test router initialization with retry configuration.""" args = RouterArgs( diff --git a/e2e_test/kind_discovery/conftest.py b/e2e_test/kind_discovery/conftest.py index 1297565614..6d4b679d28 100644 --- a/e2e_test/kind_discovery/conftest.py +++ b/e2e_test/kind_discovery/conftest.py @@ -51,6 +51,8 @@ def start(self, extra_args: tuple[str, ...] = ()) -> None: "127.0.0.1", "--port", str(SMG_PORT), + # The legacy spelling; in_cluster.yaml uses the tagged + # --discovery-provider kubernetes, so a run covers both. "--service-discovery", "--selector", "app=smg-kind-e2e", diff --git a/e2e_test/kind_discovery/in_cluster.yaml b/e2e_test/kind_discovery/in_cluster.yaml index 2616f2df06..8be20db588 100644 --- a/e2e_test/kind_discovery/in_cluster.yaml +++ b/e2e_test/kind_discovery/in_cluster.yaml @@ -56,7 +56,10 @@ spec: - 0.0.0.0 - --port - "3009" - - --service-discovery + # The tagged spelling. The host-process journey (conftest.py) + # keeps the legacy --service-discovery, so a run covers both. + - --discovery-provider + - kubernetes - --selector - app=engines-incluster - --service-discovery-port diff --git a/model_gateway/src/config/builder.rs b/model_gateway/src/config/builder.rs index 15e211a485..81e7acd590 100644 --- a/model_gateway/src/config/builder.rs +++ b/model_gateway/src/config/builder.rs @@ -5,9 +5,9 @@ use smg_mcp::McpConfig; use super::{ CacheIndexKind, CircuitBreakerConfig, ConfigError, ConfigResult, DiscoveryConfig, - HealthCheckConfig, HistoryBackend, MetricsConfig, OracleConfig, PdPairingMode, PolicyConfig, - PostgresConfig, RedisConfig, RetryConfig, RouterConfig, RoutingKeyOverrideConfig, RoutingMode, - TenantApiKeyEntry, TokenizerCacheConfig, TraceConfig, + HealthCheckConfig, HistoryBackend, KubernetesDiscoveryConfig, MetricsConfig, OracleConfig, + PdPairingMode, PolicyConfig, PostgresConfig, RedisConfig, RetryConfig, RouterConfig, + RoutingKeyOverrideConfig, RoutingMode, TenantApiKeyEntry, TokenizerCacheConfig, TraceConfig, }; use crate::worker::{ConnectionMode, RuntimeType}; @@ -504,17 +504,14 @@ impl RouterConfigBuilder { // ==================== Discovery ==================== - pub fn discovery_config(mut self, discovery: DiscoveryConfig) -> Self { - self.config.discovery = Some(discovery); + pub fn discovery_config(mut self, discovery: impl Into) -> Self { + self.config.discovery = Some(discovery.into()); self } - /// With default settings + /// Kubernetes discovery with default settings pub fn enable_discovery(mut self) -> Self { - self.config.discovery = Some(DiscoveryConfig { - enabled: true, - ..Default::default() - }); + self.config.discovery = Some(KubernetesDiscoveryConfig::default().into()); self } diff --git a/model_gateway/src/config/types.rs b/model_gateway/src/config/types.rs index 238a9c4980..3883d1b71d 100755 --- a/model_gateway/src/config/types.rs +++ b/model_gateway/src/config/types.rs @@ -2,7 +2,7 @@ use std::collections::HashMap; use openai_protocol::worker::HealthCheckConfig as ProtocolHealthCheckConfig; pub use openai_protocol::worker::{MmProcessingMode, TransportMode}; -use serde::{Deserialize, Serialize}; +use serde::{de, Deserialize, Deserializer, Serialize}; // Re-export storage config types from data_connector pub use smg_data_connector::{ HistoryBackend, OracleConfig, PostgresConfig, RedisConfig, SchemaConfig, @@ -205,6 +205,7 @@ pub struct RouterConfig { /// `api_key` rather than replacing it. #[serde(default)] pub tenant_api_keys: Vec, + #[serde(default, deserialize_with = "deserialize_discovery")] pub discovery: Option, pub metrics: Option, pub trace_config: Option, @@ -966,10 +967,78 @@ impl PolicyConfig { } } -/// Service discovery configuration -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct DiscoveryConfig { - pub enabled: bool, +/// Worker discovery: the one provider that finds this router's workers, and +/// that provider's settings. Presence means discovery is on; there is no +/// `enabled` flag. +/// +/// Serialized as a `provider` tag beside the provider's own flat fields: +/// +/// ```yaml +/// discovery: +/// provider: kubernetes +/// selector: { app: sglang } +/// port: 8000 +/// ``` +/// +/// [`RouterConfig::discovery`] also reads the flat Kubernetes object that +/// predates the tag; see [`deserialize_discovery`]. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[serde(tag = "provider", rename_all = "snake_case")] +pub enum DiscoveryConfig { + Kubernetes(KubernetesDiscoveryConfig), +} + +impl From for DiscoveryConfig { + fn from(config: KubernetesDiscoveryConfig) -> Self { + Self::Kubernetes(config) + } +} + +/// Deserialize [`RouterConfig::discovery`] from either shape it has had. +/// +/// The flat Kubernetes object that predates the `provider` tag has no tag and +/// carries an `enabled` flag. A missing tag means Kubernetes, the only +/// provider that object could describe. `enabled: false` becomes `None` here, +/// at deserialization, because canonical output has no `enabled` to carry the +/// `false`: a disabled config that is read and written back must stay +/// disabled. +/// +/// Fields are read only after the object is buffered to find its tag, so a +/// YAML scalar keeps the type YAML gave it: a label value that looks like a +/// number or boolean must be quoted (`version: "1"`), as in a Kubernetes +/// manifest. The flat struct once read such values as strings. +/// +/// Public so every input surface — the Python binding's `discovery` mapping +/// included — reads discovery by the same rules. +pub fn deserialize_discovery<'de, D>(deserializer: D) -> Result, D::Error> +where + D: Deserializer<'de>, +{ + let Some(mut value) = Option::::deserialize(deserializer)? else { + return Ok(None); + }; + if let Some(fields) = value.as_object_mut() { + match fields.remove("enabled") { + None | Some(serde_json::Value::Bool(true)) => {} + Some(serde_json::Value::Bool(false)) => return Ok(None), + Some(other) => { + return Err(de::Error::custom(format!( + "discovery.enabled must be a boolean, got {other}" + ))); + } + } + fields + .entry("provider") + .or_insert_with(|| serde_json::Value::from("kubernetes")); + } + DiscoveryConfig::deserialize(value) + .map(Some) + .map_err(de::Error::custom) +} + +/// Kubernetes worker discovery: Pods matching a selector become workers. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub struct KubernetesDiscoveryConfig { /// None = all namespaces pub namespace: Option, pub port: u16, @@ -1019,10 +1088,9 @@ fn default_kv_engine_id_annotation() -> String { "smg.ai/kv-engine-id".to_string() } -impl Default for DiscoveryConfig { +impl Default for KubernetesDiscoveryConfig { fn default() -> Self { Self { - enabled: false, namespace: None, port: 8000, check_interval_secs: 120, @@ -1295,7 +1363,7 @@ impl RouterConfig { /// Check if service discovery is enabled pub fn has_service_discovery(&self) -> bool { - self.discovery.as_ref().is_some_and(|d| d.enabled) + self.discovery.is_some() } /// Check if metrics are enabled @@ -2040,9 +2108,8 @@ mod tests { #[test] fn test_discovery_config_default() { - let config = DiscoveryConfig::default(); + let config = KubernetesDiscoveryConfig::default(); - assert!(!config.enabled); assert!(config.namespace.is_none()); assert_eq!(config.port, 8000); assert_eq!(config.check_interval_secs, 120); @@ -2061,8 +2128,7 @@ mod tests { selector.insert("app".to_string(), "sglang".to_string()); selector.insert("role".to_string(), "worker".to_string()); - let config = DiscoveryConfig { - enabled: true, + let config = KubernetesDiscoveryConfig { namespace: Some("default".to_string()), port: 9000, check_interval_secs: 30, @@ -2079,7 +2145,6 @@ mod tests { model_id_source: None, }; - assert!(config.enabled); assert_eq!(config.namespace, Some("default".to_string())); assert_eq!(config.port, 9000); assert_eq!(config.selector.len(), 2); @@ -2088,19 +2153,147 @@ mod tests { #[test] fn test_discovery_config_namespace() { - let config = DiscoveryConfig { + let config = KubernetesDiscoveryConfig { namespace: None, ..Default::default() }; assert!(config.namespace.is_none()); - let config = DiscoveryConfig { + let config = KubernetesDiscoveryConfig { namespace: Some("production".to_string()), ..Default::default() }; assert_eq!(config.namespace, Some("production".to_string())); } + /// `discovery` as `RouterConfig` reads it, from a JSON object. + fn read_discovery(discovery: serde_json::Value) -> Option { + let mut config = serde_json::to_value(RouterConfig::default()).unwrap(); + config["discovery"] = discovery; + serde_json::from_value::(config) + .unwrap() + .discovery + } + + /// The fields the discovery object has always required. + fn kubernetes_fields() -> serde_json::Value { + serde_json::json!({ + "namespace": "prod", + "port": 8000, + "check_interval_secs": 60, + "selector": { "app": "sglang" }, + "prefill_selector": {}, + "decode_selector": {}, + "bootstrap_port_annotation": "sglang.ai/bootstrap-port", + }) + } + + fn with(mut base: serde_json::Value, key: &str, value: serde_json::Value) -> serde_json::Value { + base[key] = value; + base + } + + #[test] + fn legacy_flat_discovery_reads_as_kubernetes() { + let legacy = with(kubernetes_fields(), "enabled", true.into()); + let Some(DiscoveryConfig::Kubernetes(kubernetes)) = read_discovery(legacy) else { + panic!("expected Kubernetes discovery"); + }; + assert_eq!(kubernetes.namespace.as_deref(), Some("prod")); + assert_eq!( + kubernetes.selector.get("app").map(String::as_str), + Some("sglang") + ); + } + + /// A missing tag means Kubernetes even without the legacy `enabled`: + /// presence alone turns discovery on. + #[test] + fn untagged_discovery_without_enabled_reads_as_kubernetes() { + assert!(matches!( + read_discovery(kubernetes_fields()), + Some(DiscoveryConfig::Kubernetes(_)) + )); + } + + #[test] + fn tagged_kubernetes_discovery_reads_as_kubernetes() { + let tagged = with(kubernetes_fields(), "provider", "kubernetes".into()); + assert!(matches!( + read_discovery(tagged), + Some(DiscoveryConfig::Kubernetes(_)) + )); + } + + /// `enabled: false` is dropped at deserialization, so writing the config + /// back out cannot turn discovery on. + #[test] + fn disabled_legacy_discovery_reads_as_none_and_stays_none() { + let disabled = with(kubernetes_fields(), "enabled", false.into()); + assert!(read_discovery(disabled.clone()).is_none()); + + let mut config = serde_json::to_value(RouterConfig::default()).unwrap(); + config["discovery"] = disabled; + let config: RouterConfig = serde_json::from_value(config).unwrap(); + let reread: RouterConfig = + serde_json::from_str(&serde_json::to_string(&config).unwrap()).unwrap(); + assert!(reread.discovery.is_none()); + } + + /// Canonical output carries the tag and no `enabled`, and reads back to + /// the same config. + #[test] + fn discovery_serializes_tagged_without_enabled() { + let discovery = DiscoveryConfig::Kubernetes(KubernetesDiscoveryConfig { + namespace: Some("prod".to_string()), + ..Default::default() + }); + let written = serde_json::to_value(&discovery).unwrap(); + assert_eq!(written["provider"], "kubernetes"); + assert!(written.get("enabled").is_none()); + assert_eq!(read_discovery(written), Some(discovery)); + } + + #[test] + fn legacy_flat_discovery_reads_from_yaml() { + let yaml = " +discovery: + enabled: true + port: 8000 + check_interval_secs: 60 + selector: { app: sglang } + prefill_selector: {} + decode_selector: {} + bootstrap_port_annotation: sglang.ai/bootstrap-port +"; + let mut config = serde_json::to_value(RouterConfig::default()).unwrap(); + let overlay: serde_json::Value = serde_yaml::from_str(yaml).unwrap(); + for (key, value) in overlay.as_object().unwrap() { + config[key] = value.clone(); + } + let config: RouterConfig = serde_json::from_value(config).unwrap(); + assert!(matches!( + config.discovery, + Some(DiscoveryConfig::Kubernetes(_)) + )); + } + + #[test] + fn unknown_discovery_provider_is_rejected() { + let mut config = serde_json::to_value(RouterConfig::default()).unwrap(); + config["discovery"] = with(kubernetes_fields(), "provider", "zookeeper".into()); + let err = serde_json::from_value::(config).unwrap_err(); + assert!(err.to_string().contains("zookeeper"), "{err}"); + } + + #[test] + fn non_boolean_discovery_enabled_is_rejected() { + let mut config = serde_json::to_value(RouterConfig::default()).unwrap(); + config["discovery"] = with(kubernetes_fields(), "enabled", "yes".into()); + let err = serde_json::from_value::(config).unwrap_err(); + assert!(err.to_string().contains("discovery.enabled"), "{err}"); + } + #[test] fn test_metrics_config_default() { let config = MetricsConfig::default(); @@ -2157,14 +2350,6 @@ mod tests { let config = RouterConfig::default(); assert!(!config.has_service_discovery()); - let config = RouterConfig::builder() - .discovery_config(DiscoveryConfig { - enabled: false, - ..Default::default() - }) - .build_unchecked(); - assert!(!config.has_service_discovery()); - let config = RouterConfig::builder().enable_discovery().build_unchecked(); assert!(config.has_service_discovery()); } @@ -2269,8 +2454,7 @@ mod tests { .request_timeout_secs(120) .worker_startup_timeout_secs(60) .worker_startup_check_interval_secs(5) - .discovery_config(DiscoveryConfig { - enabled: true, + .discovery_config(KubernetesDiscoveryConfig { namespace: Some("sglang".to_string()), ..Default::default() }) @@ -2307,8 +2491,7 @@ mod tests { .request_timeout_secs(300) .worker_startup_timeout_secs(180) .worker_startup_check_interval_secs(15) - .discovery_config(DiscoveryConfig { - enabled: true, + .discovery_config(KubernetesDiscoveryConfig { namespace: None, port: 8080, check_interval_secs: 45, @@ -2344,8 +2527,7 @@ mod tests { .request_timeout_secs(900) .worker_startup_timeout_secs(600) .worker_startup_check_interval_secs(20) - .discovery_config(DiscoveryConfig { - enabled: true, + .discovery_config(KubernetesDiscoveryConfig { namespace: Some("production".to_string()), port: 8443, check_interval_secs: 120, @@ -2378,10 +2560,10 @@ mod tests { assert_eq!(deserialized.host, "::1"); assert_eq!(deserialized.port, 8888); - assert_eq!( - deserialized.discovery.unwrap().namespace, - Some("production".to_string()) - ); + let Some(DiscoveryConfig::Kubernetes(discovery)) = deserialized.discovery else { + panic!("expected Kubernetes discovery"); + }; + assert_eq!(discovery.namespace, Some("production".to_string())); } #[test] diff --git a/model_gateway/src/config/validation.rs b/model_gateway/src/config/validation.rs index bed4214a11..554c28fc80 100644 --- a/model_gateway/src/config/validation.rs +++ b/model_gateway/src/config/validation.rs @@ -892,10 +892,17 @@ impl ConfigValidator { } fn validate_discovery(discovery: &DiscoveryConfig, mode: &RoutingMode) -> ConfigResult<()> { - if !discovery.enabled { - return Ok(()); + match discovery { + DiscoveryConfig::Kubernetes(kubernetes) => { + Self::validate_kubernetes_discovery(kubernetes, mode) + } } + } + fn validate_kubernetes_discovery( + discovery: &KubernetesDiscoveryConfig, + mode: &RoutingMode, + ) -> ConfigResult<()> { if discovery.port == 0 { return Err(ConfigError::InvalidValue { field: "discovery.port".to_string(), @@ -1151,7 +1158,7 @@ impl ConfigValidator { } fn validate_compatibility(config: &RouterConfig) -> ConfigResult<()> { - let has_service_discovery = config.discovery.as_ref().is_some_and(|d| d.enabled); + let has_service_discovery = config.discovery.is_some(); let invalid_decode_policy = match &config.mode { RoutingMode::PrefillDecode { decode_policy, .. } if !config.enable_igw => { !has_service_discovery && matches!(decode_policy, Some(PolicyConfig::Bucket { .. })) @@ -1270,10 +1277,7 @@ mod tests { bucket_adjust_interval_secs: 5, }; let mut config = RouterConfig { - discovery: Some(DiscoveryConfig { - enabled: true, - ..Default::default() - }), + discovery: Some(KubernetesDiscoveryConfig::default().into()), mode: RoutingMode::PrefillDecode { prefill_urls: vec![], decode_urls: vec![], @@ -1649,13 +1653,15 @@ mod tests { ); // Enable service discovery - config.discovery = Some(DiscoveryConfig { - enabled: true, - selector: vec![("app".to_string(), "test".to_string())] - .into_iter() - .collect(), - ..Default::default() - }); + config.discovery = Some( + KubernetesDiscoveryConfig { + selector: vec![("app".to_string(), "test".to_string())] + .into_iter() + .collect(), + ..Default::default() + } + .into(), + ); // Should pass validation since service discovery is enabled assert!(ConfigValidator::validate(&config).is_ok()); @@ -1676,13 +1682,15 @@ mod tests { ), ] { let mut config = regular_mode_config(); - config.discovery = Some(DiscoveryConfig { - enabled: true, - selector: [("app".to_string(), "worker".to_string())].into(), - kv_connector_annotation: kv_connector_annotation.to_string(), - kv_engine_id_annotation: kv_engine_id_annotation.to_string(), - ..Default::default() - }); + config.discovery = Some( + KubernetesDiscoveryConfig { + selector: [("app".to_string(), "worker".to_string())].into(), + kv_connector_annotation: kv_connector_annotation.to_string(), + kv_engine_id_annotation: kv_engine_id_annotation.to_string(), + ..Default::default() + } + .into(), + ); assert!(matches!( ConfigValidator::validate(&config), diff --git a/model_gateway/src/main.rs b/model_gateway/src/main.rs index e4b9d8afd8..1c398ba9c8 100644 --- a/model_gateway/src/main.rs +++ b/model_gateway/src/main.rs @@ -14,9 +14,9 @@ use smg::{ config::{ resolve_worker_auto_recovery, validate_mesh_server_name, CacheIndexKind, CircuitBreakerConfig, ConfigError, ConfigResult, DiscoveryConfig, HealthCheckConfig, - HistoryBackend, ManualAssignmentMode, MetricsConfig, OracleConfig, PdPairingMode, - PolicyConfig, PostgresConfig, RedisConfig, RetryConfig, RouterConfig, - RoutingKeyOverrideConfig, RoutingMode, SchemaConfig, TenantApiKeyEntry, + HistoryBackend, KubernetesDiscoveryConfig, ManualAssignmentMode, MetricsConfig, + OracleConfig, PdPairingMode, PolicyConfig, PostgresConfig, RedisConfig, RetryConfig, + RouterConfig, RoutingKeyOverrideConfig, RoutingMode, SchemaConfig, TenantApiKeyEntry, TokenizerCacheConfig, TraceConfig, }, mesh_discovery::MeshDiscoveryConfig, @@ -25,7 +25,7 @@ use smg::{ otel_trace::{is_otel_enabled, shutdown_otel}, }, server::{self, ServerConfig}, - service_discovery::{ModelIdSource, ServiceDiscoveryConfig}, + service_discovery::{ModelIdSource, RuntimeDiscoveryConfig}, version, worker::{ConnectionMode, RuntimeType}, }; @@ -109,6 +109,13 @@ impl std::fmt::Display for Backend { } } +/// A worker discovery provider, as `--discovery-provider` names it. +#[derive(Copy, Clone, Debug, Eq, PartialEq, ValueEnum)] +pub enum DiscoveryProvider { + #[value(name = "kubernetes")] + Kubernetes, +} + #[derive(Parser, Debug)] #[command(name = "shepherd-model-gateway", alias = "smg", alias = "amg")] #[command(about = "Shepherd Model Gateway - High-performance inference gateway")] @@ -666,7 +673,8 @@ struct CliArgs { log_mm_timing: bool, // ==================== Service Discovery (Kubernetes) ==================== - /// Enable Kubernetes service discovery + /// Enable Kubernetes service discovery (the legacy spelling of + /// `--discovery-provider kubernetes`) #[arg( long, default_value_t = false, @@ -674,6 +682,15 @@ struct CliArgs { )] service_discovery: bool, + /// Worker discovery provider. Give this or `--service-discovery`, not both + #[arg( + long, + value_enum, + conflicts_with = "service_discovery", + help_heading = "Service Discovery (Kubernetes)" + )] + discovery_provider: Option, + /// Label selector for Kubernetes service discovery (format: key=value) #[arg(long, num_args = 0.., help_heading = "Service Discovery (Kubernetes)")] selector: Vec, @@ -1410,6 +1427,24 @@ impl CliArgs { .unwrap_or(ConnectionMode::Http) } + /// The worker discovery provider selected by either spelling: + /// `--service-discovery` is `--discovery-provider kubernetes`. + fn selected_discovery_provider(&self) -> Option { + if self.service_discovery { + Some(DiscoveryProvider::Kubernetes) + } else { + self.discovery_provider + } + } + + /// Selecting a discovery provider, by either spelling, turns IGW mode on. + /// Returns whether this call turned it on. + fn enable_igw_for_discovery(&mut self) -> bool { + let enable = self.selected_discovery_provider().is_some() && !self.enable_igw; + self.enable_igw |= enable; + enable + } + fn parse_selector(selector_list: &[String]) -> HashMap { let mut map = HashMap::new(); for item in selector_list { @@ -1778,27 +1813,28 @@ impl CliArgs { let policy = self.parse_policy(&self.policy); - let discovery = if self.service_discovery { - Some(DiscoveryConfig { - enabled: true, - namespace: self.service_discovery_namespace.clone(), - port: self.service_discovery_port, - check_interval_secs: 60, - selector: Self::parse_selector(&self.selector), - encode_selector: Self::parse_selector(&self.encode_selector), - prefill_selector: Self::parse_selector(&self.prefill_selector), - decode_selector: Self::parse_selector(&self.decode_selector), - bootstrap_port_annotation: "sglang.ai/bootstrap-port".to_string(), - worker_ports_annotation: "smg.ai/worker-ports".to_string(), - kv_connector_annotation: self.kv_connector_annotation.clone(), - kv_engine_id_annotation: self.kv_engine_id_annotation.clone(), - router_selector: Self::parse_selector(&self.router_selector), - router_mesh_port_annotation: "sglang.ai/mesh-port".to_string(), - model_id_source: self.model_id_from.clone(), - }) - } else { - None - }; + let discovery = self + .selected_discovery_provider() + .map(|provider| match provider { + DiscoveryProvider::Kubernetes => { + DiscoveryConfig::Kubernetes(KubernetesDiscoveryConfig { + namespace: self.service_discovery_namespace.clone(), + port: self.service_discovery_port, + check_interval_secs: 60, + selector: Self::parse_selector(&self.selector), + encode_selector: Self::parse_selector(&self.encode_selector), + prefill_selector: Self::parse_selector(&self.prefill_selector), + decode_selector: Self::parse_selector(&self.decode_selector), + bootstrap_port_annotation: "sglang.ai/bootstrap-port".to_string(), + worker_ports_annotation: "smg.ai/worker-ports".to_string(), + kv_connector_annotation: self.kv_connector_annotation.clone(), + kv_engine_id_annotation: self.kv_engine_id_annotation.clone(), + router_selector: Self::parse_selector(&self.router_selector), + router_mesh_port_annotation: "sglang.ai/mesh-port".to_string(), + model_id_source: self.model_id_from.clone(), + }) + } + }); let metrics = Some(MetricsConfig { port: self.prometheus_port, @@ -1984,7 +2020,7 @@ impl CliArgs { disable_health_check: self.disable_health_check, remove_unhealthy_workers: resolve_worker_auto_recovery( self.remove_unhealthy_workers, - self.service_discovery, + self.selected_discovery_provider().is_some(), ), drain_settle_secs: self.drain_settle_secs, }) @@ -2053,74 +2089,15 @@ impl CliArgs { } fn to_server_config(&self, router_config: RouterConfig) -> ConfigResult { - let service_discovery_config = if self.service_discovery { - let (kv_connector_annotation, kv_engine_id_annotation) = router_config - .discovery - .as_ref() - .map(|d| { - ( - d.kv_connector_annotation.clone(), - d.kv_engine_id_annotation.clone(), - ) - }) - .unwrap_or_else(|| { - ( - self.kv_connector_annotation.clone(), - self.kv_engine_id_annotation.clone(), - ) - }); - - let model_id_source = self - .model_id_from - .as_deref() - .or_else(|| { - router_config - .discovery - .as_ref() - .and_then(|d| d.model_id_source.as_deref()) - }) - .map(|s| { - ModelIdSource::parse(s).map_err(|e| ConfigError::InvalidValue { - field: "model_id_source".to_string(), - value: s.to_string(), - reason: e, - }) - }) - .transpose()?; - - Some(ServiceDiscoveryConfig { - enabled: true, - selector: Self::parse_selector(&self.selector), - check_interval: std::time::Duration::from_secs(60), - port: self.service_discovery_port, - namespace: self.service_discovery_namespace.clone(), - disaggregated_mode: self.pd_disaggregation || self.epd_disaggregation, - encode_selector: Self::parse_selector(&self.encode_selector), - prefill_selector: Self::parse_selector(&self.prefill_selector), - decode_selector: Self::parse_selector(&self.decode_selector), - bootstrap_port_annotation: "sglang.ai/bootstrap-port".to_string(), - worker_ports_annotation: "smg.ai/worker-ports".to_string(), - kv_connector_annotation, - kv_engine_id_annotation, - model_id_source, - }) - } else { - None - }; - - // Mesh-router discovery now has its own task and lifetime, but its - // configuration still arrives inside `discovery`, so it stays reachable - // only under `--service-discovery`. Moving it onto its own config - // surface lands with the tagged provider configuration. + let service_discovery_config = router_config + .discovery + .as_ref() + .map(|discovery| RuntimeDiscoveryConfig::from_config(discovery, &router_config.mode)) + .transpose()?; let mesh_discovery_config = router_config .discovery .as_ref() - .map(|d| MeshDiscoveryConfig { - namespace: d.namespace.clone(), - router_selector: d.router_selector.clone(), - router_mesh_port_annotation: d.router_mesh_port_annotation.clone(), - }) - .filter(MeshDiscoveryConfig::is_enabled); + .and_then(MeshDiscoveryConfig::from_discovery); let prometheus_config = Some(PrometheusConfig { port: self.prometheus_port, @@ -2223,10 +2200,8 @@ fn main() -> Result<(), Box> { None => cli.router_args, }; - // Automatically enable IGW mode when service discovery is turned on - if cli_args.service_discovery && !cli_args.enable_igw { + if cli_args.enable_igw_for_discovery() { println!("INFO: IGW mode automatically enabled because service discovery is turned on"); - cli_args.enable_igw = true; } let mode_str = if cli_args.enable_igw { @@ -2538,9 +2513,9 @@ mod tests { } /// Mesh-router discovery has its own task and lifetime, but is still - /// configured through `discovery`, which only `--service-discovery` - /// populates. `--router-selector` alone therefore still does nothing; it - /// gains its own config surface with the tagged provider configuration. + /// configured through Kubernetes `discovery`, which only a selected + /// Kubernetes provider populates. `--router-selector` alone therefore + /// still does nothing until router discovery gets its own config surface. #[test] fn router_selector_alone_does_not_yet_configure_mesh_discovery() { let cli = cli_args_from(&["--router-selector", "role=router"]); @@ -2632,23 +2607,103 @@ mod tests { "example.com/engine-id", ]); let router = cli.to_router_config(vec![], vec![]).unwrap(); - let discovery = router.discovery.as_ref().unwrap(); + let Some(DiscoveryConfig::Kubernetes(discovery)) = &router.discovery else { + panic!("expected Kubernetes discovery"); + }; assert_eq!(discovery.kv_connector_annotation, "example.com/connector"); assert_eq!(discovery.kv_engine_id_annotation, "example.com/engine-id"); let server = cli.to_server_config(router).unwrap(); - let discovery = server.service_discovery_config.as_ref().unwrap(); + let Some(RuntimeDiscoveryConfig::Kubernetes(discovery)) = &server.service_discovery_config + else { + panic!("expected Kubernetes runtime discovery"); + }; assert_eq!(discovery.kv_connector_annotation, "example.com/connector"); assert_eq!(discovery.kv_engine_id_annotation, "example.com/engine-id"); let defaults = cli_args_from(&["--service-discovery", "--selector", "app=worker"]) .to_router_config(vec![], vec![]) .unwrap(); - let defaults = defaults.discovery.as_ref().unwrap(); + let Some(DiscoveryConfig::Kubernetes(defaults)) = &defaults.discovery else { + panic!("expected Kubernetes discovery"); + }; assert_eq!(defaults.kv_connector_annotation, "smg.ai/kv-connector"); assert_eq!(defaults.kv_engine_id_annotation, "smg.ai/kv-engine-id"); } + /// `--discovery-provider kubernetes` is `--service-discovery` under its + /// new name: the same detail flags build the same configuration, worker + /// and router discovery alike. + #[test] + fn discovery_provider_kubernetes_matches_service_discovery() { + let details = [ + "--selector", + "app=worker", + "--service-discovery-port", + "9000", + "--service-discovery-namespace", + "prod", + "--kv-connector-annotation", + "example.com/connector", + "--model-id-from", + "namespace", + "--router-selector", + "role=router", + ]; + let build = |selection: &[&str]| { + let args: Vec<&str> = selection.iter().chain(details.iter()).copied().collect(); + let cli = cli_args_from(&args); + let router = cli.to_router_config(vec![], vec![]).unwrap(); + let discovery = router.discovery.clone(); + let server = cli.to_server_config(router).unwrap(); + ( + discovery, + format!("{:?}", server.service_discovery_config), + format!("{:?}", server.mesh_discovery_config), + ) + }; + + let legacy = build(&["--service-discovery"]); + let tagged = build(&["--discovery-provider", "kubernetes"]); + assert!(matches!(legacy.0, Some(DiscoveryConfig::Kubernetes(_)))); + assert_eq!(legacy, tagged); + } + + /// Two spellings of one choice: giving both is a usage error, not a + /// precedence rule to remember. + #[test] + fn service_discovery_and_discovery_provider_conflict() { + let err = Cli::try_parse_from([ + "smg", + "--service-discovery", + "--discovery-provider", + "kubernetes", + ]) + .expect_err("both spellings must be rejected"); + assert_eq!(err.kind(), clap::error::ErrorKind::ArgumentConflict); + } + + /// IGW mode follows the selected provider, whichever spelling selected it. + #[test] + fn selecting_a_discovery_provider_enables_igw() { + for selection in [ + &["--service-discovery"][..], + &["--discovery-provider", "kubernetes"][..], + ] { + let mut cli = cli_args_from(selection); + assert!(cli.enable_igw_for_discovery(), "{selection:?}"); + assert!(cli.enable_igw, "{selection:?}"); + } + + let mut already_on = cli_args_from(&["--service-discovery", "--enable-igw"]); + assert!(!already_on.enable_igw_for_discovery()); + assert!(already_on.enable_igw); + + let mut no_discovery = cli_args_from(&[]); + assert!(!no_discovery.enable_igw_for_discovery()); + assert!(!no_discovery.enable_igw); + } + /// `--worker-auto-recovery` defaults to the `--service-discovery` /// setting: recovery works by removal plus discovery re-registration, so /// it is on exactly when discovery can complete that loop, and off when @@ -2661,6 +2716,12 @@ mod tests { .unwrap(); assert!(derived_on.health_check.remove_unhealthy_workers); + let derived_on_tagged = + cli_args_from(&["--discovery-provider", "kubernetes", "--selector", "app=w"]) + .to_router_config(vec![], vec![]) + .unwrap(); + assert!(derived_on_tagged.health_check.remove_unhealthy_workers); + let derived_off = cli_args_from(&[]).to_router_config(vec![], vec![]).unwrap(); assert!(!derived_off.health_check.remove_unhealthy_workers); diff --git a/model_gateway/src/mesh_discovery/kubernetes.rs b/model_gateway/src/mesh_discovery/kubernetes.rs index 39322066ec..d777b5c95d 100644 --- a/model_gateway/src/mesh_discovery/kubernetes.rs +++ b/model_gateway/src/mesh_discovery/kubernetes.rs @@ -27,6 +27,8 @@ use smg_mesh::{ use tokio::task; use tracing::{debug, error, info, warn}; +use crate::config::DiscoveryConfig; + /// Configuration for Kubernetes router-peer discovery. #[derive(Debug, Clone)] pub struct MeshDiscoveryConfig { @@ -49,6 +51,19 @@ impl Default for MeshDiscoveryConfig { } impl MeshDiscoveryConfig { + /// Router discovery as configured today: inside Kubernetes worker + /// discovery, so it runs only alongside it. `None` without a router + /// selector. + pub fn from_discovery(discovery: &DiscoveryConfig) -> Option { + let DiscoveryConfig::Kubernetes(kubernetes) = discovery; + Some(Self { + namespace: kubernetes.namespace.clone(), + router_selector: kubernetes.router_selector.clone(), + router_mesh_port_annotation: kubernetes.router_mesh_port_annotation.clone(), + }) + .filter(Self::is_enabled) + } + /// Router discovery only runs with a selector; without one there is no way /// to tell a router Pod from any other Pod in the namespace. pub fn is_enabled(&self) -> bool { diff --git a/model_gateway/src/server.rs b/model_gateway/src/server.rs index df1d8635c9..bb1668cd24 100644 --- a/model_gateway/src/server.rs +++ b/model_gateway/src/server.rs @@ -63,7 +63,7 @@ use crate::{ http::router::{stream_eligible_request_bodies, StreamBodyState}, RouterTrait, }, - service_discovery::{start_service_discovery, ServiceDiscoveryConfig}, + service_discovery::{start_service_discovery, RuntimeDiscoveryConfig}, wasm::route::{add_wasm_module, list_wasm_modules, remove_wasm_module}, worker::{ manager::{WorkerManager, WorkerManagerConfig}, @@ -755,7 +755,7 @@ pub struct ServerConfig { pub log_dir: Option, pub log_level: Option, pub log_json: bool, - pub service_discovery_config: Option, + pub service_discovery_config: Option, /// Kubernetes discovery of SMG mesh router peers. Independent of the /// worker discovery provider: either may run without the other. pub mesh_discovery_config: Option, @@ -1484,19 +1484,17 @@ pub async fn startup(config: ServerConfig) -> Result<(), Box { - info!("Service discovery started"); - discovery_tasks - .0 - .push(supervise_discovery("Worker discovery", handle)); - } - Err(e) => { - error!("Failed to start service discovery: {e}"); - warn!("Continuing without service discovery"); - } + let app_context_arc = Arc::clone(&app_state.context); + match start_service_discovery(service_discovery_config, app_context_arc).await { + Ok(handle) => { + info!("Service discovery started"); + discovery_tasks + .0 + .push(supervise_discovery("Worker discovery", handle)); + } + Err(e) => { + error!("Failed to start service discovery: {e}"); + warn!("Continuing without service discovery"); } } } diff --git a/model_gateway/src/service_discovery/kubernetes.rs b/model_gateway/src/service_discovery/kubernetes.rs index d8432f2cf4..e46b2ae2e2 100644 --- a/model_gateway/src/service_discovery/kubernetes.rs +++ b/model_gateway/src/service_discovery/kubernetes.rs @@ -91,7 +91,6 @@ impl ModelIdSource { #[derive(Debug, Clone)] pub struct ServiceDiscoveryConfig { - pub enabled: bool, pub selector: HashMap, pub check_interval: Duration, pub port: u16, @@ -177,7 +176,6 @@ fn build_watcher_config(label_selector: &str) -> Config { impl Default for ServiceDiscoveryConfig { fn default() -> Self { ServiceDiscoveryConfig { - enabled: false, selector: HashMap::new(), check_interval: Duration::from_secs(60), port: 8000, @@ -461,18 +459,10 @@ fn resolve_bootstrap_ports( } } -pub async fn start_service_discovery( +pub(super) async fn start_kubernetes_discovery( config: ServiceDiscoveryConfig, app_context: Arc, ) -> Result, kube::Error> { - if !config.enabled { - return Err(kube::Error::Api( - kube::core::Status::failure("Service discovery is disabled", "ConfigurationError") - .with_code(400) - .boxed(), - )); - } - let _ = ring::default_provider().install_default(); let client = Client::try_default().await?; @@ -793,7 +783,6 @@ mod tests { decode_selector.insert("component".to_string(), "decode".to_string()); ServiceDiscoveryConfig { - enabled: true, selector: HashMap::new(), check_interval: Duration::from_secs(60), port: 8080, @@ -861,7 +850,6 @@ mod tests { #[test] fn test_service_discovery_config_default() { let config = ServiceDiscoveryConfig::default(); - assert!(!config.enabled); assert!(config.selector.is_empty()); assert_eq!(config.check_interval, Duration::from_secs(60)); assert_eq!(config.port, 8000); @@ -1635,7 +1623,6 @@ mod tests { let mut selector = HashMap::new(); selector.insert("app".to_string(), "sglang".to_string()); ServiceDiscoveryConfig { - enabled: true, selector, disaggregated_mode: false, ..Default::default() diff --git a/model_gateway/src/service_discovery/mod.rs b/model_gateway/src/service_discovery/mod.rs index 35b3b93294..81d13d8526 100644 --- a/model_gateway/src/service_discovery/mod.rs +++ b/model_gateway/src/service_discovery/mod.rs @@ -10,6 +10,9 @@ //! - [`reconciler`] owns ownership, the registry diff and the `JobQueue` //! submissions. It sees only that contract — never a `Pod`, a Pod UID, or //! which provider is running beyond its kind. +//! - [`runtime`] resolves configuration into the selected provider's runtime +//! settings, through the one conversion every input surface shares, and +//! starts that provider. //! //! SMG mesh-router peer discovery is a different concern — it discovers router //! peers, not inference workers — and lives in [`crate::mesh_discovery`]. @@ -17,13 +20,14 @@ mod kubernetes; mod provider; mod reconciler; +mod runtime; #[cfg(test)] mod testing; #[cfg(feature = "test-util")] pub use kubernetes::start_service_discovery_with_client; pub use kubernetes::{ - start_service_discovery, ModelIdSource, PodInfo, PodType, ServiceDiscoveryConfig, - POD_NAME_LABEL, POD_UID_LABEL, + ModelIdSource, PodInfo, PodType, ServiceDiscoveryConfig, POD_NAME_LABEL, POD_UID_LABEL, }; pub use reconciler::{DISCOVERY_ID_LABEL, DISCOVERY_PROVIDER_LABEL, DISCOVERY_SPEC_HASH_LABEL}; +pub use runtime::{start_service_discovery, RuntimeDiscoveryConfig}; diff --git a/model_gateway/src/service_discovery/runtime.rs b/model_gateway/src/service_discovery/runtime.rs new file mode 100644 index 0000000000..ed921852ae --- /dev/null +++ b/model_gateway/src/service_discovery/runtime.rs @@ -0,0 +1,177 @@ +//! Worker discovery as the running gateway uses it, and the one conversion +//! that produces it from configuration. + +use std::{error::Error, sync::Arc, time::Duration}; + +use tokio::task; + +use super::kubernetes::{self, ModelIdSource, ServiceDiscoveryConfig}; +use crate::{ + app_context::AppContext, + config::{ConfigError, ConfigResult, DiscoveryConfig, KubernetesDiscoveryConfig, RoutingMode}, +}; + +/// The selected worker-discovery provider, resolved for the running gateway. +/// +/// `Option` is on versus off; a variant carries no +/// `enabled` flag of its own. +#[derive(Debug, Clone)] +pub enum RuntimeDiscoveryConfig { + Kubernetes(ServiceDiscoveryConfig), +} + +impl RuntimeDiscoveryConfig { + /// Resolve configuration into its runtime form. + /// + /// Every input surface — the CLI and the Python binding — converts + /// through here, so a setting cannot reach the runtime through one surface + /// and be dropped by another. Values arrive as the surface built them, + /// defaults included: the CLI's port 80 and `KubernetesDiscoveryConfig`'s + /// port 8000 are each kept, not reconciled. + pub fn from_config(discovery: &DiscoveryConfig, mode: &RoutingMode) -> ConfigResult { + match discovery { + DiscoveryConfig::Kubernetes(kubernetes) => { + Ok(Self::Kubernetes(kubernetes_runtime(kubernetes, mode)?)) + } + } + } +} + +fn kubernetes_runtime( + config: &KubernetesDiscoveryConfig, + mode: &RoutingMode, +) -> ConfigResult { + let model_id_source = config + .model_id_source + .as_deref() + .map(|source| { + ModelIdSource::parse(source).map_err(|reason| ConfigError::InvalidValue { + field: "discovery.model_id_source".to_string(), + value: source.to_string(), + reason, + }) + }) + .transpose()?; + + Ok(ServiceDiscoveryConfig { + selector: config.selector.clone(), + check_interval: Duration::from_secs(config.check_interval_secs), + port: config.port, + namespace: config.namespace.clone(), + disaggregated_mode: matches!( + mode, + RoutingMode::PrefillDecode { .. } | RoutingMode::EncodePrefillDecode { .. } + ), + encode_selector: config.encode_selector.clone(), + prefill_selector: config.prefill_selector.clone(), + decode_selector: config.decode_selector.clone(), + bootstrap_port_annotation: config.bootstrap_port_annotation.clone(), + worker_ports_annotation: config.worker_ports_annotation.clone(), + kv_connector_annotation: config.kv_connector_annotation.clone(), + kv_engine_id_annotation: config.kv_engine_id_annotation.clone(), + model_id_source, + }) +} + +/// Start the selected provider's discovery task. +pub async fn start_service_discovery( + config: RuntimeDiscoveryConfig, + app_context: Arc, +) -> Result, Box> { + match config { + RuntimeDiscoveryConfig::Kubernetes(config) => { + Ok(kubernetes::start_kubernetes_discovery(config, app_context).await?) + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn regular() -> RoutingMode { + RoutingMode::Regular { + worker_urls: vec![], + } + } + + fn kubernetes(config: &DiscoveryConfig, mode: &RoutingMode) -> ServiceDiscoveryConfig { + let RuntimeDiscoveryConfig::Kubernetes(runtime) = + RuntimeDiscoveryConfig::from_config(config, mode).unwrap(); + runtime + } + + /// Every Kubernetes setting reaches the runtime, the KV annotation names + /// included: those once flowed through only one of two hand-written + /// conversions. + #[test] + fn kubernetes_settings_reach_the_runtime() { + let config = DiscoveryConfig::Kubernetes(KubernetesDiscoveryConfig { + namespace: Some("prod".to_string()), + port: 9000, + check_interval_secs: 45, + selector: [("app".to_string(), "worker".to_string())].into(), + bootstrap_port_annotation: "example.com/bootstrap".to_string(), + worker_ports_annotation: "example.com/ports".to_string(), + kv_connector_annotation: "example.com/connector".to_string(), + kv_engine_id_annotation: "example.com/engine-id".to_string(), + model_id_source: Some("namespace".to_string()), + ..Default::default() + }); + + let runtime = kubernetes(&config, ®ular()); + + assert_eq!(runtime.namespace.as_deref(), Some("prod")); + assert_eq!(runtime.port, 9000); + assert_eq!(runtime.check_interval, Duration::from_secs(45)); + assert_eq!( + runtime.selector.get("app").map(String::as_str), + Some("worker") + ); + assert_eq!(runtime.bootstrap_port_annotation, "example.com/bootstrap"); + assert_eq!(runtime.worker_ports_annotation, "example.com/ports"); + assert_eq!(runtime.kv_connector_annotation, "example.com/connector"); + assert_eq!(runtime.kv_engine_id_annotation, "example.com/engine-id"); + assert!(matches!( + runtime.model_id_source, + Some(ModelIdSource::Namespace) + )); + assert!(!runtime.disaggregated_mode); + } + + #[test] + fn disaggregated_modes_select_role_selectors() { + let config = DiscoveryConfig::Kubernetes(KubernetesDiscoveryConfig::default()); + for mode in [ + RoutingMode::PrefillDecode { + prefill_urls: vec![], + decode_urls: vec![], + prefill_policy: None, + decode_policy: None, + }, + RoutingMode::EncodePrefillDecode { + encode_urls: vec![], + prefill_urls: vec![], + decode_urls: vec![], + encode_policy: None, + prefill_policy: None, + decode_policy: None, + }, + ] { + assert!(kubernetes(&config, &mode).disaggregated_mode, "{mode:?}"); + } + } + + #[test] + fn invalid_model_id_source_is_a_config_error() { + let config = DiscoveryConfig::Kubernetes(KubernetesDiscoveryConfig { + model_id_source: Some("pod-name".to_string()), + ..Default::default() + }); + let err = RuntimeDiscoveryConfig::from_config(&config, ®ular()).unwrap_err(); + assert!( + matches!(err, ConfigError::InvalidValue { ref field, .. } if field == "discovery.model_id_source"), + "{err}" + ); + } +} diff --git a/model_gateway/tests/k8s_discovery_test.rs b/model_gateway/tests/k8s_discovery_test.rs index b31f1b72bc..016f35975c 100644 --- a/model_gateway/tests/k8s_discovery_test.rs +++ b/model_gateway/tests/k8s_discovery_test.rs @@ -236,7 +236,6 @@ fn worker_pod(name: &str, uid: &str, ports: &str, ready: bool) -> Pod { fn discovery_config() -> ServiceDiscoveryConfig { ServiceDiscoveryConfig { - enabled: true, selector: [("app".to_string(), "smg-it".to_string())].into(), check_interval: Duration::from_millis(300), ..Default::default() From 12051b02d8a5052e5e631b7f19da49cf5a188cde Mon Sep 17 00:00:00 2001 From: XinyueZhang369 Date: Tue, 6 Oct 2026 17:39:01 -0700 Subject: [PATCH 2/2] fix(discovery): reject a provider Router.from_args cannot pass on `Router.from_args` mapped the selected provider to `service_discovery` with `== "kubernetes"`, so any other provider became `service_discovery=False`: discovery would not start and nothing would say so. It can't happen today, since only `kubernetes` is accepted, but adding a provider to the CLI choices without wiring it here would hit it. Raise instead. The new test extends the choices without wiring `from_args`, the case it guards; it fails with the old mapping. Signed-off-by: XinyueZhang369 --- bindings/python/src/smg/router.py | 9 ++++++++- bindings/python/tests/test_startup_sequence.py | 18 ++++++++++++++++++ 2 files changed, 26 insertions(+), 1 deletion(-) diff --git a/bindings/python/src/smg/router.py b/bindings/python/src/smg/router.py index cb799d8bc7..718564f119 100644 --- a/bindings/python/src/smg/router.py +++ b/bindings/python/src/smg/router.py @@ -318,7 +318,14 @@ def from_args(args: RouterArgs) -> Router: args_dict = vars(args).copy() # Kubernetes, by either spelling, reaches Rust as service_discovery. - args_dict["service_discovery"] = args.selected_discovery_provider() == "kubernetes" + # A provider this cannot pass on fails here instead of quietly + # starting the router without discovery. + provider = args.selected_discovery_provider() + if provider not in (None, "kubernetes"): + raise ValueError( + f"Router.from_args cannot pass discovery provider {provider!r} to Rust" + ) + args_dict["service_discovery"] = provider == "kubernetes" args_dict.pop("discovery_provider") # Convert RouterArgs to _Router parameters args_dict["worker_urls"] = ( diff --git a/bindings/python/tests/test_startup_sequence.py b/bindings/python/tests/test_startup_sequence.py index 435cb336d6..a2afc26374 100644 --- a/bindings/python/tests/test_startup_sequence.py +++ b/bindings/python/tests/test_startup_sequence.py @@ -295,6 +295,24 @@ def test_from_args_passes_discovery_provider_as_service_discovery(self): # Discovery supplies the workers, exactly as under --service-discovery. assert kwargs["worker_urls"] == [] + def test_from_args_rejects_a_provider_it_cannot_pass(self): + """A provider accepted on the CLI but not wired through fails loudly. + + Simulates adding a provider to the CLI choices without teaching + from_args to pass it on: that must not start the router without + discovery. + """ + args = RouterArgs(discovery_provider="file") + + with ( + patch("smg.router_args.DISCOVERY_PROVIDER_CHOICES", ["kubernetes", "file"]), + patch("smg.router._Router") as rust_router, + pytest.raises(ValueError, match="cannot pass discovery provider 'file'"), + ): + Router.from_args(args) + + rust_router.assert_not_called() + def test_router_initialization_with_retry_config(self): """Test router initialization with retry configuration.""" args = RouterArgs(