diff --git a/sgl-router/src/core/worker.rs b/sgl-router/src/core/worker.rs index 1c42092f2..621127e54 100644 --- a/sgl-router/src/core/worker.rs +++ b/sgl-router/src/core/worker.rs @@ -171,11 +171,17 @@ pub trait Worker: Send + Sync + fmt::Debug { } /// Get the model ID this worker serves + /// Checks ModelCards first, then falls back to labels fn model_id(&self) -> &str { + // Check ModelCards first self.metadata() - .labels - .get("model_id") - .map(|s| s.as_str()) + .models + .first() + .map(|m| m.id.as_str()) + .or_else(|| { + // Fall back to labels + self.metadata().labels.get("model_id").map(|s| s.as_str()) + }) .unwrap_or("unknown") } @@ -197,36 +203,61 @@ pub trait Worker: Send + Sync + fmt::Debug { .unwrap_or(1.0) } - /// Get the tokenizer path for this worker (gRPC mode only) - fn tokenizer_path(&self) -> Option<&str> { + /// Get tokenizer path for a specific model. + fn tokenizer_path(&self, model_id: &str) -> Option<&str> { self.metadata() - .labels - .get("tokenizer_path") - .map(|s| s.as_str()) + .find_model(model_id) + .and_then(|m| m.tokenizer_path.as_deref()) } - /// Get the reasoning parser type for this worker (gRPC mode only) - fn reasoning_parser(&self) -> Option<&str> { + /// Get reasoning parser for a specific model. + fn reasoning_parser(&self, model_id: &str) -> Option<&str> { self.metadata() - .labels - .get("reasoning_parser") - .map(|s| s.as_str()) + .find_model(model_id) + .and_then(|m| m.reasoning_parser.as_deref()) } - /// Get the tool parser type for this worker (gRPC mode only) - fn tool_parser(&self) -> Option<&str> { + /// Get tool parser for a specific model. + fn tool_parser(&self, model_id: &str) -> Option<&str> { self.metadata() - .labels - .get("tool_parser") - .map(|s| s.as_str()) + .find_model(model_id) + .and_then(|m| m.tool_parser.as_deref()) } - /// Get the chat template for this worker (gRPC mode only) - fn chat_template(&self) -> Option<&str> { + /// Get chat template for a specific model. + fn chat_template(&self, model_id: &str) -> Option<&str> { self.metadata() - .labels - .get("chat_template") - .map(|s| s.as_str()) + .find_model(model_id) + .and_then(|m| m.chat_template.as_deref()) + } + + /// Get the default provider type for this worker. + /// `None` means native/passthrough. + fn default_provider(&self) -> Option<&ProviderType> { + self.metadata().default_provider.as_ref() + } + + /// Get provider for a specific model. + /// Priority: ModelCard.provider > worker.default_provider + fn provider_for_model(&self, model_id: &str) -> Option<&ProviderType> { + self.metadata().provider_for_model(model_id) + } + + /// Check if this worker supports a specific model. + /// If models list is empty, worker accepts any model. + fn supports_model(&self, model_id: &str) -> bool { + self.metadata().supports_model(model_id) + } + + /// Check if this worker supports an endpoint for a given model. + /// Falls back to default_model_type if model not found. + fn supports_endpoint(&self, model_id: &str, endpoint: Endpoint) -> bool { + self.metadata().supports_endpoint(model_id, endpoint) + } + + /// Get all models this worker can serve. + fn models(&self) -> &[ModelCard] { + &self.metadata().models } /// Get or create a gRPC client for this worker @@ -1015,10 +1046,11 @@ pub fn worker_to_info(worker: &Arc) -> WorkerInfo { ConnectionMode::Http => None, }; + let model_id = worker.model_id(); WorkerInfo { id: worker.url().to_string(), url: worker.url().to_string(), - model_id: worker.model_id().to_string(), + model_id: model_id.to_string(), priority: worker.priority(), cost: worker.cost(), worker_type: worker_type_str.to_string(), @@ -1026,10 +1058,10 @@ pub fn worker_to_info(worker: &Arc) -> WorkerInfo { load: worker.load(), connection_mode: format!("{:?}", worker.connection_mode()), runtime_type, - tokenizer_path: worker.tokenizer_path().map(String::from), - reasoning_parser: worker.reasoning_parser().map(String::from), - tool_parser: worker.tool_parser().map(String::from), - chat_template: worker.chat_template().map(String::from), + tokenizer_path: worker.tokenizer_path(model_id).map(String::from), + reasoning_parser: worker.reasoning_parser(model_id).map(String::from), + tool_parser: worker.tool_parser(model_id).map(String::from), + chat_template: worker.chat_template(model_id).map(String::from), bootstrap_port, metadata: worker.metadata().labels.clone(), job_status: None, @@ -2041,4 +2073,156 @@ mod tests { let ids: Vec<&str> = metadata.model_ids().collect(); assert_eq!(ids, vec!["model-a", "model-b", "model-c"]); } + + // === Phase 1.4: Worker trait model-aware methods tests === + + #[test] + fn test_worker_tokenizer_path() { + use super::ModelCard; + use crate::core::BasicWorkerBuilder; + + // Create a worker with a ModelCard that has tokenizer_path + let model_card = + ModelCard::new("my-model").with_tokenizer_path("my-model/tokenizer".to_string()); + + let worker = BasicWorkerBuilder::new("http://test:8080") + .model(model_card) + .build(); + + // Should find the tokenizer_path from the ModelCard + assert_eq!( + worker.tokenizer_path("my-model"), + Some("my-model/tokenizer") + ); + + // Unknown model should return None + assert_eq!(worker.tokenizer_path("unknown-model"), None); + } + + #[test] + fn test_worker_model_aware_methods_with_model_cards() { + use super::{ModelCard, ProviderType}; + use crate::core::BasicWorkerBuilder; + + // Build worker (labels are not used for model config anymore) + let mut worker = BasicWorkerBuilder::new("http://test:8080").build(); + + // Add model cards to the worker's metadata + let model_with_config = ModelCard::new("gpt-4o") + .with_tokenizer_path("gpt4o/tokenizer") + .with_chat_template("gpt4o_template") + .with_reasoning_parser("gpt4o_reasoning") + .with_tool_parser("gpt4o_tools") + .with_provider(ProviderType::OpenAI); + + let model_without_config = ModelCard::new("llama-3.1"); + + worker.metadata.models = vec![model_with_config, model_without_config]; + + // Model with explicit config should use ModelCard values + assert_eq!(worker.tokenizer_path("gpt-4o"), Some("gpt4o/tokenizer")); + assert_eq!(worker.chat_template("gpt-4o"), Some("gpt4o_template")); + assert_eq!(worker.reasoning_parser("gpt-4o"), Some("gpt4o_reasoning")); + assert_eq!(worker.tool_parser("gpt-4o"), Some("gpt4o_tools")); + assert_eq!( + worker.provider_for_model("gpt-4o"), + Some(&ProviderType::OpenAI) + ); + + // Model without explicit config should return None (no fallback to labels) + assert_eq!(worker.tokenizer_path("llama-3.1"), None); + assert_eq!(worker.chat_template("llama-3.1"), None); + assert_eq!(worker.reasoning_parser("llama-3.1"), None); + assert_eq!(worker.tool_parser("llama-3.1"), None); + + // Unknown model should return None + assert_eq!(worker.tokenizer_path("unknown"), None); + } + + #[test] + fn test_worker_supports_model_and_endpoint() { + use super::{Endpoint, ModelCard}; + use crate::core::BasicWorkerBuilder; + + let mut worker = BasicWorkerBuilder::new("http://test:8080").build(); + + // Empty models list - accepts any model + assert!(worker.supports_model("any-model")); + + // Add specific models + let llm_model = ModelCard::new("gpt-4o").with_model_type(ModelType::LLM); + let embed_model = ModelCard::new("text-embedding").with_model_type(ModelType::EMBEDDINGS); + + worker.metadata.models = vec![llm_model, embed_model]; + + // Now only listed models are supported + assert!(worker.supports_model("gpt-4o")); + assert!(worker.supports_model("text-embedding")); + assert!(!worker.supports_model("unknown-model")); + + // Check endpoint support + assert!(worker.supports_endpoint("gpt-4o", Endpoint::Chat)); + assert!(!worker.supports_endpoint("gpt-4o", Endpoint::Embeddings)); + assert!(worker.supports_endpoint("text-embedding", Endpoint::Embeddings)); + assert!(!worker.supports_endpoint("text-embedding", Endpoint::Chat)); + } + + #[test] + fn test_worker_models_accessor() { + use super::ModelCard; + use crate::core::BasicWorkerBuilder; + + let mut worker = BasicWorkerBuilder::new("http://test:8080").build(); + + // Initially empty + assert!(worker.models().is_empty()); + + // Add models + worker.metadata.models = vec![ModelCard::new("model-a"), ModelCard::new("model-b")]; + + assert_eq!(worker.models().len(), 2); + assert_eq!(worker.models()[0].id, "model-a"); + assert_eq!(worker.models()[1].id, "model-b"); + } + + #[test] + fn test_worker_default_provider() { + use super::ProviderType; + use crate::core::BasicWorkerBuilder; + + let mut worker = BasicWorkerBuilder::new("http://test:8080").build(); + + // Default is None (native/passthrough) + assert!(worker.default_provider().is_none()); + + // Set a default provider + worker.metadata.default_provider = Some(ProviderType::OpenAI); + assert_eq!(worker.default_provider(), Some(&ProviderType::OpenAI)); + } + + #[test] + fn test_worker_model_id_with_model_cards() { + use super::ModelCard; + use crate::core::BasicWorkerBuilder; + + // Test 1: No models, no labels → "unknown" + let worker = BasicWorkerBuilder::new("http://test:8080").build(); + assert_eq!(worker.model_id(), "unknown"); + + // Test 2: No models but has label → uses label + let worker = BasicWorkerBuilder::new("http://test:8080") + .label("model_id", "label-model") + .build(); + assert_eq!(worker.model_id(), "label-model"); + + // Test 3: Has ModelCards → uses first ModelCard + let mut worker = BasicWorkerBuilder::new("http://test:8080") + .label("model_id", "label-model") + .build(); + worker.metadata.models = vec![ + ModelCard::new("card-model-1"), + ModelCard::new("card-model-2"), + ]; + assert_eq!(worker.model_id(), "card-model-1"); + } } diff --git a/sgl-router/src/core/worker_builder.rs b/sgl-router/src/core/worker_builder.rs index dfdfb071e..de0354950 100644 --- a/sgl-router/src/core/worker_builder.rs +++ b/sgl-router/src/core/worker_builder.rs @@ -2,6 +2,7 @@ use std::collections::HashMap; use super::{ circuit_breaker::{CircuitBreaker, CircuitBreakerConfig}, + model_card::ModelCard, model_type::ModelType, worker::{ BasicWorker, ConnectionMode, DPAwareWorker, HealthConfig, RuntimeType, WorkerMetadata, @@ -18,6 +19,7 @@ pub struct BasicWorkerBuilder { connection_mode: ConnectionMode, runtime_type: RuntimeType, labels: HashMap, + models: Vec, health_config: HealthConfig, circuit_breaker_config: CircuitBreakerConfig, grpc_client: Option, @@ -33,6 +35,7 @@ impl BasicWorkerBuilder { connection_mode: ConnectionMode::Http, runtime_type: RuntimeType::default(), labels: HashMap::new(), + models: Vec::new(), health_config: HealthConfig::default(), circuit_breaker_config: CircuitBreakerConfig::default(), grpc_client: None, @@ -48,6 +51,7 @@ impl BasicWorkerBuilder { connection_mode: ConnectionMode::Http, runtime_type: RuntimeType::default(), labels: HashMap::new(), + models: Vec::new(), health_config: HealthConfig::default(), circuit_breaker_config: CircuitBreakerConfig::default(), grpc_client: None, @@ -108,6 +112,18 @@ impl BasicWorkerBuilder { self } + /// Set models this worker can serve + pub fn models(mut self, models: Vec) -> Self { + self.models = models; + self + } + + /// Add a single model this worker can serve + pub fn model(mut self, model: ModelCard) -> Self { + self.models.push(model); + self + } + /// Build the BasicWorker instance pub fn build(self) -> BasicWorker { use std::sync::{ @@ -155,7 +171,7 @@ impl BasicWorkerBuilder { health_config: self.health_config, bootstrap_host, bootstrap_port, - models: Vec::new(), // Empty = accepts any model + models: self.models, // Empty = accepts any model default_provider: None, // Native/passthrough default_model_type: ModelType::LLM, // Standard LLM capabilities }; @@ -185,6 +201,7 @@ pub struct DPAwareWorkerBuilder { connection_mode: ConnectionMode, runtime_type: RuntimeType, labels: HashMap, + models: Vec, health_config: HealthConfig, circuit_breaker_config: CircuitBreakerConfig, grpc_client: Option, @@ -202,6 +219,7 @@ impl DPAwareWorkerBuilder { connection_mode: ConnectionMode::Http, runtime_type: RuntimeType::default(), labels: HashMap::new(), + models: Vec::new(), health_config: HealthConfig::default(), circuit_breaker_config: CircuitBreakerConfig::default(), grpc_client: None, @@ -224,6 +242,7 @@ impl DPAwareWorkerBuilder { connection_mode: ConnectionMode::Http, runtime_type: RuntimeType::default(), labels: HashMap::new(), + models: Vec::new(), health_config: HealthConfig::default(), circuit_breaker_config: CircuitBreakerConfig::default(), grpc_client: None, @@ -284,10 +303,23 @@ impl DPAwareWorkerBuilder { self } + /// Set models this worker can serve + pub fn models(mut self, models: Vec) -> Self { + self.models = models; + self + } + + /// Add a single model this worker can serve + pub fn model(mut self, model: ModelCard) -> Self { + self.models.push(model); + self + } + /// Build the DPAwareWorker instance pub fn build(self) -> DPAwareWorker { let worker_url = format!("{}@{}", self.base_url, self.dp_rank); let mut builder = BasicWorkerBuilder::new(worker_url) + .models(self.models) .worker_type(self.worker_type) .connection_mode(self.connection_mode) .runtime_type(self.runtime_type) diff --git a/sgl-router/src/core/workflow/steps/worker_registration.rs b/sgl-router/src/core/workflow/steps/worker_registration.rs index af21c30e2..7b380f072 100644 --- a/sgl-router/src/core/workflow/steps/worker_registration.rs +++ b/sgl-router/src/core/workflow/steps/worker_registration.rs @@ -24,7 +24,7 @@ use crate::{ app_context::AppContext, core::{ workflow::*, BasicWorkerBuilder, CircuitBreakerConfig, ConnectionMode, - DPAwareWorkerBuilder, HealthConfig, RuntimeType, Worker, WorkerType, + DPAwareWorkerBuilder, HealthConfig, ModelCard, RuntimeType, Worker, WorkerType, }, protocols::worker_spec::WorkerConfigRequest, routers::grpc::client::GrpcClient, @@ -474,29 +474,14 @@ impl StepExecutor for CreateWorkerStep { }); } - // Build labels from config + // Build labels from config (non-model-specific labels only) let mut config_labels = config.labels.clone(); - if let Some(model_id) = &config.model_id { - config_labels.insert("model_id".to_string(), model_id.clone()); - } if let Some(priority) = config.priority { config_labels.insert("priority".to_string(), priority.to_string()); } if let Some(cost) = config.cost { config_labels.insert("cost".to_string(), cost.to_string()); } - if let Some(ref tokenizer_path) = config.tokenizer_path { - config_labels.insert("tokenizer_path".to_string(), tokenizer_path.clone()); - } - if let Some(ref reasoning_parser) = config.reasoning_parser { - config_labels.insert("reasoning_parser".to_string(), reasoning_parser.clone()); - } - if let Some(ref tool_parser) = config.tool_parser { - config_labels.insert("tool_parser".to_string(), tool_parser.clone()); - } - if let Some(ref chat_template) = config.chat_template { - config_labels.insert("chat_template".to_string(), chat_template.clone()); - } // Merge: discovered labels first, then config labels (config takes precedence) let mut final_labels = discovered_labels.as_ref().clone(); @@ -504,19 +489,36 @@ impl StepExecutor for CreateWorkerStep { final_labels.insert(key.clone(), value.clone()); } - // Derive model_id if not already set - if !final_labels.contains_key("model_id") { - let derived_model_id = final_labels - .get("served_model_name") - .or_else(|| final_labels.get("model_path")) - .cloned(); + // Determine model_id: config > served_model_name > model_path > "unknown" + let model_id = config + .model_id + .clone() + .or_else(|| final_labels.get("served_model_name").cloned()) + .or_else(|| final_labels.get("model_path").cloned()) + .unwrap_or_else(|| "unknown".to_string()); - if let Some(model_id) = derived_model_id { - debug!("Derived model_id from metadata: {}", model_id); - final_labels.insert("model_id".to_string(), model_id); - } + if model_id != "unknown" { + debug!("Using model_id: {}", model_id); } + // Create ModelCard with model-specific configuration + let model_card = { + let mut card = ModelCard::new(&model_id); + if let Some(ref tokenizer_path) = config.tokenizer_path { + card = card.with_tokenizer_path(tokenizer_path.clone()); + } + if let Some(ref reasoning_parser) = config.reasoning_parser { + card = card.with_reasoning_parser(reasoning_parser.clone()); + } + if let Some(ref tool_parser) = config.tool_parser { + card = card.with_tool_parser(tool_parser.clone()); + } + if let Some(ref chat_template) = config.chat_template { + card = card.with_chat_template(chat_template.clone()); + } + card + }; + debug!( "Creating worker {} with {} discovered + {} config = {} final labels", config.url, @@ -621,6 +623,7 @@ impl StepExecutor for CreateWorkerStep { for rank in 0..dp_info.dp_size { let mut builder = DPAwareWorkerBuilder::new(normalized_url.clone(), rank, dp_info.dp_size) + .model(model_card.clone()) .worker_type(worker_type.clone()) .connection_mode(connection_mode.as_ref().clone()) .runtime_type(runtime_type.clone()) @@ -656,6 +659,7 @@ impl StepExecutor for CreateWorkerStep { } else { // Non-DP-aware path: Create single worker let mut builder = BasicWorkerBuilder::new(normalized_url.clone()) + .model(model_card) .worker_type(worker_type) .connection_mode(connection_mode.as_ref().clone()) .runtime_type(runtime_type)