Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
111 changes: 57 additions & 54 deletions bindings/python/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -549,6 +549,23 @@ struct Router {
prefill_max_inflight_requests_per_worker: i32,
prefill_queue_size: Option<usize>,
prefill_queue_timeout_secs: Option<u64>,
/// The keyword-only `discovery` mapping, read by the same rules as
/// `RouterConfig.discovery`.
discovery: Option<config::DiscoveryConfig>,
}

/// 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<Option<config::DiscoveryConfig>> {
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 {
Expand Down Expand Up @@ -612,7 +629,8 @@ impl Router {

pub fn to_router_config(&self) -> config::ConfigResult<config::RouterConfig> {
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
Expand Down Expand Up @@ -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,
Expand All @@ -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 {
Expand Down Expand Up @@ -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,
})
Expand Down Expand Up @@ -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<String>,
policy: PolicyType,
Expand Down Expand Up @@ -1335,7 +1352,20 @@ impl Router {
prefill_max_inflight_requests_per_worker: i32,
prefill_queue_size: Option<usize>,
prefill_queue_timeout_secs: Option<u64>,
discovery: Option<Bound<'_, PyAny>>,
) -> PyResult<Self> {
// 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 {
Expand Down Expand Up @@ -1516,6 +1546,7 @@ impl Router {
prefill_max_inflight_requests_per_worker,
prefill_queue_size,
prefill_queue_timeout_secs,
discovery,
})
}

Expand All @@ -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),
Expand Down
2 changes: 1 addition & 1 deletion bindings/python/src/smg/launch_router.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
14 changes: 14 additions & 0 deletions bindings/python/src/smg/router.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -313,6 +317,16 @@ 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.
# 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"] = (
[]
Expand Down
37 changes: 35 additions & 2 deletions bindings/python/src/smg/router_args.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]:
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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",
Expand Down Expand Up @@ -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:
Expand Down
37 changes: 37 additions & 0 deletions bindings/python/tests/test_arg_parser.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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 = [
Expand Down Expand Up @@ -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):
Expand Down
34 changes: 34 additions & 0 deletions bindings/python/tests/test_router_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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"})
Loading
Loading