diff --git a/sgl-model-gateway/bindings/python/src/lib.rs b/sgl-model-gateway/bindings/python/src/lib.rs index de75ed202..f0378fbcd 100644 --- a/sgl-model-gateway/bindings/python/src/lib.rs +++ b/sgl-model-gateway/bindings/python/src/lib.rs @@ -973,7 +973,7 @@ impl Router { })?; let service_discovery_config = if self.service_discovery { - Some(service_discovery::ServiceDiscoveryConfig { + let config = service_discovery::ServiceDiscoveryConfig { enabled: true, selector: self.selector.clone(), check_interval: std::time::Duration::from_secs(60), @@ -985,7 +985,10 @@ impl Router { bootstrap_port_annotation: self.bootstrap_port_annotation.clone(), router_selector: HashMap::new(), router_mesh_port_annotation: "sglang.ai/mesh-port".to_string(), - }) + igw_mode: self.enable_igw, + }; + config.warn_if_misconfigured(); + Some(config) } else { None }; diff --git a/sgl-model-gateway/src/main.rs b/sgl-model-gateway/src/main.rs index fe1bcf4d3..b389bdad1 100644 --- a/sgl-model-gateway/src/main.rs +++ b/sgl-model-gateway/src/main.rs @@ -1089,9 +1089,11 @@ impl CliArgs { }) .unwrap_or_else(|| (HashMap::new(), "sglang.ai/mesh-port".to_string())); - Some(ServiceDiscoveryConfig { + let selector = Self::parse_selector(&self.selector); + + let service_discovery_config = ServiceDiscoveryConfig { enabled: true, - selector: Self::parse_selector(&self.selector), + selector, check_interval: std::time::Duration::from_secs(60), port: self.service_discovery_port, namespace: self.service_discovery_namespace.clone(), @@ -1101,7 +1103,10 @@ impl CliArgs { bootstrap_port_annotation: "sglang.ai/bootstrap-port".to_string(), router_selector, router_mesh_port_annotation, - }) + igw_mode: self.enable_igw, + }; + service_discovery_config.warn_if_misconfigured(); + Some(service_discovery_config) } else { None }; diff --git a/sgl-model-gateway/src/service_discovery.rs b/sgl-model-gateway/src/service_discovery.rs index 25dc63ddc..c7c49e95f 100644 --- a/sgl-model-gateway/src/service_discovery.rs +++ b/sgl-model-gateway/src/service_discovery.rs @@ -45,6 +45,8 @@ pub struct ServiceDiscoveryConfig { // Router node discovery for mesh pub router_selector: HashMap, pub router_mesh_port_annotation: String, + // When true (IGW mode), also discover selector pods as Regular workers alongside PD workers + pub igw_mode: bool, } impl Default for ServiceDiscoveryConfig { @@ -61,6 +63,19 @@ impl Default for ServiceDiscoveryConfig { bootstrap_port_annotation: "sglang.ai/bootstrap-port".to_string(), router_selector: HashMap::new(), router_mesh_port_annotation: "sglang.ai/ha-port".to_string(), + igw_mode: false, + } + } +} + +impl ServiceDiscoveryConfig { + pub fn warn_if_misconfigured(&self) { + if self.pd_mode && !self.igw_mode && !self.selector.is_empty() { + warn!( + "--selector is set in PD mode without IGW mode enabled; \ + regular worker discovery alongside PD workers requires IGW mode, \ + selector will be ignored" + ); } } } @@ -99,11 +114,18 @@ impl PodInfo { pub fn should_include(pod: &Pod, config: &ServiceDiscoveryConfig) -> bool { if config.pd_mode { if config.prefill_selector.is_empty() && config.decode_selector.is_empty() { - warn!("PD mode enabled but both prefill_selector and decode_selector are empty"); - return false; + if !(config.igw_mode && !config.selector.is_empty()) { + warn!("PD mode enabled but both prefill_selector and decode_selector are empty"); + return false; + } } - Self::matches_selector(pod, &config.prefill_selector) - || Self::matches_selector(pod, &config.decode_selector) + let matches_pd = Self::matches_selector(pod, &config.prefill_selector) + || Self::matches_selector(pod, &config.decode_selector); + // In IGW mode, also discover regular workers via the selector field + let matches_regular = config.igw_mode + && !config.selector.is_empty() + && Self::matches_selector(pod, &config.selector); + matches_pd || matches_regular } else { if config.selector.is_empty() { warn!("Regular mode enabled but selector is empty"); @@ -884,6 +906,35 @@ mod tests { bootstrap_port_annotation: "sglang.ai/bootstrap-port".to_string(), router_selector: HashMap::new(), router_mesh_port_annotation: "sglang.ai/ha-port".to_string(), + igw_mode: false, + } + } + + fn create_regular_k8s_pod(name: &str, ip: &str) -> Pod { + let mut labels = std::collections::BTreeMap::new(); + labels.insert("app".to_string(), "regular-worker".to_string()); + + Pod { + metadata: ObjectMeta { + name: Some(name.to_string()), + labels: Some(labels), + ..Default::default() + }, + spec: Some(PodSpec::default()), + status: Some(PodStatus { + pod_ip: Some(ip.to_string()), + phase: Some("Running".to_string()), + conditions: Some(vec![PodCondition { + type_: "Ready".to_string(), + status: "True".to_string(), + last_probe_time: None, + last_transition_time: None, + message: None, + reason: None, + observed_generation: None, + }]), + ..Default::default() + }), } } @@ -910,6 +961,33 @@ mod tests { assert!(PodInfo::should_include(®ular_pod, ®ular_config)); } + #[test] + fn test_should_include_regular_pod_in_pd_igw_mode() { + let mut config = create_pd_config(); + config.igw_mode = true; + config + .selector + .insert("app".to_string(), "regular-worker".to_string()); + + let regular_pod = create_regular_k8s_pod("regular-pod", "10.0.0.5"); + assert!(PodInfo::should_include(®ular_pod, &config)); + + let pod_info = PodInfo::from_pod(®ular_pod, Some(&config)).unwrap(); + assert_eq!(pod_info.pod_type, Some(PodType::Regular)); + } + + #[test] + fn test_should_exclude_regular_pod_in_pd_mode_without_igw() { + let mut config = create_pd_config(); + config.igw_mode = false; + config + .selector + .insert("app".to_string(), "regular-worker".to_string()); + + let regular_pod = create_regular_k8s_pod("regular-pod", "10.0.0.5"); + assert!(!PodInfo::should_include(®ular_pod, &config)); + } + #[test] fn test_service_discovery_config_default() { let config = ServiceDiscoveryConfig::default(); @@ -1441,4 +1519,71 @@ mod tests { // Pod should be removed from tracking assert!(!tracked_pods.lock().unwrap().contains(&pod_info)); } + + #[test] + fn test_should_include_mixed_pd_igw_regular_pod_included() { + let mut regular_selector = HashMap::new(); + regular_selector.insert("app".to_string(), "regular-worker".to_string()); + + let mut prefill_selector = HashMap::new(); + prefill_selector.insert("app".to_string(), "sglang".to_string()); + prefill_selector.insert("component".to_string(), "prefill".to_string()); + + let mut decode_selector = HashMap::new(); + decode_selector.insert("app".to_string(), "sglang".to_string()); + decode_selector.insert("component".to_string(), "decode".to_string()); + + let config = ServiceDiscoveryConfig { + enabled: true, + selector: regular_selector, + check_interval: Duration::from_secs(60), + port: 8080, + namespace: None, + pd_mode: true, + prefill_selector, + decode_selector, + bootstrap_port_annotation: "sglang.ai/bootstrap-port".to_string(), + router_selector: HashMap::new(), + router_mesh_port_annotation: "sglang.ai/ha-port".to_string(), + igw_mode: true, + }; + + let regular_pod = create_regular_k8s_pod("regular-pod", "10.0.1.1"); + assert!(PodInfo::should_include(®ular_pod, &config)); + + let pod_info = PodInfo::from_pod(®ular_pod, Some(&config)).unwrap(); + assert_eq!(pod_info.pod_type, Some(PodType::Regular)); + } + + #[test] + fn test_should_include_mixed_pd_no_igw_regular_pod_excluded() { + let mut regular_selector = HashMap::new(); + regular_selector.insert("app".to_string(), "regular-worker".to_string()); + + let mut prefill_selector = HashMap::new(); + prefill_selector.insert("app".to_string(), "sglang".to_string()); + prefill_selector.insert("component".to_string(), "prefill".to_string()); + + let mut decode_selector = HashMap::new(); + decode_selector.insert("app".to_string(), "sglang".to_string()); + decode_selector.insert("component".to_string(), "decode".to_string()); + + let config = ServiceDiscoveryConfig { + enabled: true, + selector: regular_selector, + check_interval: Duration::from_secs(60), + port: 8080, + namespace: None, + pd_mode: true, + prefill_selector, + decode_selector, + bootstrap_port_annotation: "sglang.ai/bootstrap-port".to_string(), + router_selector: HashMap::new(), + router_mesh_port_annotation: "sglang.ai/ha-port".to_string(), + igw_mode: false, + }; + + let regular_pod = create_regular_k8s_pod("regular-pod", "10.0.1.1"); + assert!(!PodInfo::should_include(®ular_pod, &config)); + } }