From cd4151abc7943f81ef4ec5f1f869a92e9b8e55af Mon Sep 17 00:00:00 2001 From: Simo Lin Date: Mon, 1 Dec 2025 18:36:38 -0800 Subject: [PATCH] [model-gateway] add audio and moderation in model card (#14263) --- sgl-router/src/core/model_card.rs | 226 ----------------------- sgl-router/src/core/model_type.rs | 242 ++++++------------------ sgl-router/src/core/worker.rs | 295 +----------------------------- 3 files changed, 60 insertions(+), 703 deletions(-) diff --git a/sgl-router/src/core/model_card.rs b/sgl-router/src/core/model_card.rs index bd9809bdd..a2c27360d 100644 --- a/sgl-router/src/core/model_card.rs +++ b/sgl-router/src/core/model_card.rs @@ -321,229 +321,3 @@ impl std::fmt::Display for ModelCard { write!(f, "{}", self.name()) } } - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_model_card_new() { - let card = ModelCard::new("llama-3.1-8b"); - assert_eq!(card.id, "llama-3.1-8b"); - assert_eq!(card.model_type, ModelType::LLM); - assert!(card.provider.is_none()); // Native by default - assert!(card.aliases.is_empty()); - } - - #[test] - fn test_model_card_builder() { - let card = ModelCard::new("meta-llama/Llama-3.1-8B-Instruct") - .with_display_name("Llama 3.1 8B") - .with_alias("llama-3.1-8b") - .with_alias("llama3.1") - .with_model_type(ModelType::VISION_LLM) - .with_context_length(128_000) - .with_tokenizer_path("meta-llama/Llama-3.1-8B-Instruct") - .with_reasoning_parser("deepseek") - .with_tool_parser("llama"); - - assert_eq!(card.name(), "Llama 3.1 8B"); - assert_eq!(card.aliases.len(), 2); - assert!(card.supports_vision()); - assert_eq!(card.context_length, Some(128_000)); - assert_eq!( - card.tokenizer_path.as_deref(), - Some("meta-llama/Llama-3.1-8B-Instruct") - ); - assert_eq!(card.reasoning_parser.as_deref(), Some("deepseek")); - assert_eq!(card.tool_parser.as_deref(), Some("llama")); - assert!(card.is_native()); // No provider set - } - - #[test] - fn test_model_card_with_provider() { - let card = ModelCard::new("gpt-4o").with_provider(ProviderType::OpenAI); - - assert!(card.has_external_provider()); - assert!(!card.is_native()); - assert_eq!(card.provider, Some(ProviderType::OpenAI)); - } - - #[test] - fn test_model_card_matches() { - let card = ModelCard::new("gpt-4o") - .with_alias("gpt-4o-2024-08-06") - .with_alias("gpt-4-turbo"); - - assert!(card.matches("gpt-4o")); - assert!(card.matches("gpt-4o-2024-08-06")); - assert!(card.matches("gpt-4-turbo")); - assert!(!card.matches("gpt-3.5")); - } - - #[test] - fn test_model_card_supports_endpoint() { - let llm = ModelCard::new("llama").with_model_type(ModelType::LLM); - assert!(llm.supports_endpoint(Endpoint::Chat)); - assert!(llm.supports_endpoint(Endpoint::Completions)); - assert!(!llm.supports_endpoint(Endpoint::Embeddings)); - - let embed = ModelCard::new("bge").with_model_type(ModelType::EMBED_MODEL); - assert!(embed.supports_endpoint(Endpoint::Embeddings)); - assert!(!embed.supports_endpoint(Endpoint::Chat)); - } - - #[test] - fn test_model_card_name_fallback() { - let card_with_name = ModelCard::new("model-id").with_display_name("Display Name"); - assert_eq!(card_with_name.name(), "Display Name"); - - let card_without_name = ModelCard::new("model-id"); - assert_eq!(card_without_name.name(), "model-id"); - } - - #[test] - fn test_model_card_is_llm() { - let llm = ModelCard::new("llama").with_model_type(ModelType::LLM); - assert!(llm.is_llm()); - - let embed = ModelCard::new("bge").with_model_type(ModelType::EMBED_MODEL); - assert!(!embed.is_llm()); - } - - #[test] - fn test_model_card_default() { - let card = ModelCard::default(); - assert_eq!(card.id, "default"); - assert_eq!(card.model_type, ModelType::LLM); - assert!(card.provider.is_none()); - } - - #[test] - fn test_model_card_serialization() { - let card = ModelCard::new("gpt-4o") - .with_display_name("GPT-4o") - .with_alias("gpt4o") - .with_provider(ProviderType::OpenAI) - .with_context_length(128_000); - - let json = serde_json::to_string(&card).unwrap(); - let deserialized: ModelCard = serde_json::from_str(&json).unwrap(); - - assert_eq!(deserialized.id, "gpt-4o"); - assert_eq!(deserialized.display_name, Some("GPT-4o".to_string())); - assert_eq!(deserialized.aliases, vec!["gpt4o"]); - assert_eq!(deserialized.provider, Some(ProviderType::OpenAI)); - assert_eq!(deserialized.context_length, Some(128_000)); - } - - #[test] - fn test_model_card_serialization_native() { - // Native model (no provider) should not include provider in JSON - let card = ModelCard::new("llama-3.1"); - let json = serde_json::to_string(&card).unwrap(); - - // Provider should be omitted from JSON when None - assert!(!json.contains("provider")); - - let deserialized: ModelCard = serde_json::from_str(&json).unwrap(); - assert!(deserialized.provider.is_none()); - } - - #[test] - fn test_model_card_display() { - let card = ModelCard::new("model-id").with_display_name("Display Name"); - assert_eq!(format!("{}", card), "Display Name"); - } - - // === ProviderType tests === - - #[test] - fn test_provider_type_as_str() { - assert_eq!(ProviderType::OpenAI.as_str(), "openai"); - assert_eq!(ProviderType::XAI.as_str(), "xai"); - assert_eq!(ProviderType::Anthropic.as_str(), "anthropic"); - assert_eq!(ProviderType::Gemini.as_str(), "gemini"); - assert_eq!( - ProviderType::Custom("my-provider".to_string()).as_str(), - "my-provider" - ); - } - - #[test] - fn test_provider_type_from_model_name() { - // External providers - assert_eq!( - ProviderType::from_model_name("grok-beta"), - Some(ProviderType::XAI) - ); - assert_eq!( - ProviderType::from_model_name("Grok-2"), - Some(ProviderType::XAI) - ); - assert_eq!( - ProviderType::from_model_name("gemini-pro"), - Some(ProviderType::Gemini) - ); - assert_eq!( - ProviderType::from_model_name("claude-3-opus"), - Some(ProviderType::Anthropic) - ); - assert_eq!( - ProviderType::from_model_name("gpt-4o"), - Some(ProviderType::OpenAI) - ); - assert_eq!( - ProviderType::from_model_name("o1-preview"), - Some(ProviderType::OpenAI) - ); - assert_eq!( - ProviderType::from_model_name("o3-mini"), - Some(ProviderType::OpenAI) - ); - - // Native/local models - no provider - assert_eq!(ProviderType::from_model_name("llama-3.1"), None); - assert_eq!(ProviderType::from_model_name("mistral-7b"), None); - assert_eq!(ProviderType::from_model_name("deepseek-r1"), None); - assert_eq!(ProviderType::from_model_name("qwen-2.5"), None); - } - - #[test] - fn test_provider_type_serialization() { - // Test standard variants - let openai = ProviderType::OpenAI; - let json = serde_json::to_string(&openai).unwrap(); - assert_eq!(json, "\"openai\""); - let deserialized: ProviderType = serde_json::from_str(&json).unwrap(); - assert_eq!(deserialized, ProviderType::OpenAI); - - // Test custom variant - let custom = ProviderType::Custom("my-provider".to_string()); - let json = serde_json::to_string(&custom).unwrap(); - let deserialized: ProviderType = serde_json::from_str(&json).unwrap(); - assert_eq!( - deserialized, - ProviderType::Custom("my-provider".to_string()) - ); - } - - #[test] - fn test_provider_type_aliases() { - // Test serde aliases - let xai: ProviderType = serde_json::from_str("\"grok\"").unwrap(); - assert_eq!(xai, ProviderType::XAI); - - let anthropic: ProviderType = serde_json::from_str("\"claude\"").unwrap(); - assert_eq!(anthropic, ProviderType::Anthropic); - - let gemini: ProviderType = serde_json::from_str("\"google\"").unwrap(); - assert_eq!(gemini, ProviderType::Gemini); - } - - #[test] - fn test_provider_type_display() { - assert_eq!(format!("{}", ProviderType::OpenAI), "openai"); - assert_eq!(format!("{}", ProviderType::XAI), "xai"); - } -} diff --git a/sgl-router/src/core/model_type.rs b/sgl-router/src/core/model_type.rs index 912ad4be8..9e9381971 100644 --- a/sgl-router/src/core/model_type.rs +++ b/sgl-router/src/core/model_type.rs @@ -30,9 +30,12 @@ bitflags! { const TOOLS = 1 << 7; /// Reasoning/thinking support (e.g., o1, DeepSeek-R1) const REASONING = 1 << 8; - - // === Convenience combinations === - // Note: Within bitflags! macro, we must use .bits() for combining flags + /// Image generation (DALL-E, Sora, gpt-image) + const IMAGE_GEN = 1 << 9; + /// Audio models (TTS, Whisper, realtime, transcribe) + const AUDIO = 1 << 10; + /// Content moderation models + const MODERATION = 1 << 11; /// Standard LLM: chat + completions + responses + tools const LLM = Self::CHAT.bits() | Self::COMPLETIONS.bits() @@ -52,6 +55,15 @@ bitflags! { /// Reranker model only const RERANK_MODEL = Self::RERANK.bits(); + + /// Image generation model only (DALL-E, Sora, gpt-image) + const IMAGE_MODEL = Self::IMAGE_GEN.bits(); + + /// Audio model only (TTS, Whisper, realtime) + const AUDIO_MODEL = Self::AUDIO.bits(); + + /// Content moderation model only + const MODERATION_MODEL = Self::MODERATION.bits(); } } @@ -67,6 +79,9 @@ const CAPABILITY_NAMES: &[(ModelType, &str)] = &[ (ModelType::VISION, "vision"), (ModelType::TOOLS, "tools"), (ModelType::REASONING, "reasoning"), + (ModelType::IMAGE_GEN, "image_gen"), + (ModelType::AUDIO, "audio"), + (ModelType::MODERATION, "moderation"), ]; impl ModelType { @@ -124,6 +139,24 @@ impl ModelType { self.contains(Self::REASONING) } + /// Check if this model type supports image generation + #[inline] + pub fn supports_image_gen(&self) -> bool { + self.contains(Self::IMAGE_GEN) + } + + /// Check if this model type supports audio (TTS, Whisper, etc.) + #[inline] + pub fn supports_audio(&self) -> bool { + self.contains(Self::AUDIO) + } + + /// Check if this model type supports content moderation + #[inline] + pub fn supports_moderation(&self) -> bool { + self.contains(Self::MODERATION) + } + /// Check if this model type supports a given endpoint pub fn supports_endpoint(&self, endpoint: Endpoint) -> bool { match endpoint { @@ -165,6 +198,24 @@ impl ModelType { pub fn is_reranker(&self) -> bool { self.supports_rerank() && !self.supports_chat() } + + /// Check if this is an image generation model + #[inline] + pub fn is_image_model(&self) -> bool { + self.supports_image_gen() && !self.supports_chat() + } + + /// Check if this is an audio model + #[inline] + pub fn is_audio_model(&self) -> bool { + self.supports_audio() && !self.supports_chat() + } + + /// Check if this is a moderation model + #[inline] + pub fn is_moderation_model(&self) -> bool { + self.supports_moderation() && !self.supports_chat() + } } impl std::fmt::Display for ModelType { @@ -279,188 +330,3 @@ impl std::fmt::Display for Endpoint { } } } - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_model_type_individual_flags() { - let chat_only = ModelType::CHAT; - assert!(chat_only.supports_chat()); - assert!(!chat_only.supports_completions()); - assert!(!chat_only.supports_embeddings()); - } - - #[test] - fn test_model_type_combinations() { - let llm = ModelType::LLM; - assert!(llm.supports_chat()); - assert!(llm.supports_completions()); - assert!(llm.supports_responses()); - assert!(llm.supports_tools()); - assert!(!llm.supports_embeddings()); - assert!(!llm.supports_vision()); - } - - #[test] - fn test_model_type_custom_combination() { - let custom = ModelType::CHAT | ModelType::EMBEDDINGS; - assert!(custom.supports_chat()); - assert!(custom.supports_embeddings()); - assert!(!custom.supports_completions()); - assert!(!custom.supports_tools()); - } - - #[test] - fn test_model_type_vision_llm() { - let vision = ModelType::VISION_LLM; - assert!(vision.supports_chat()); - assert!(vision.supports_completions()); - assert!(vision.supports_responses()); - assert!(vision.supports_tools()); - assert!(vision.supports_vision()); - assert!(!vision.supports_embeddings()); - assert!(!vision.supports_reasoning()); - } - - #[test] - fn test_model_type_reasoning_llm() { - let reasoning = ModelType::REASONING_LLM; - assert!(reasoning.supports_chat()); - assert!(reasoning.supports_reasoning()); - assert!(!reasoning.supports_vision()); - } - - #[test] - fn test_model_type_full_llm() { - let full = ModelType::FULL_LLM; - assert!(full.supports_chat()); - assert!(full.supports_completions()); - assert!(full.supports_responses()); - assert!(full.supports_tools()); - assert!(full.supports_vision()); - assert!(full.supports_reasoning()); - assert!(!full.supports_embeddings()); - } - - #[test] - fn test_model_type_supports_endpoint() { - let llm = ModelType::LLM; - assert!(llm.supports_endpoint(Endpoint::Chat)); - assert!(llm.supports_endpoint(Endpoint::Completions)); - assert!(llm.supports_endpoint(Endpoint::Responses)); - assert!(llm.supports_endpoint(Endpoint::Models)); // Always true - assert!(!llm.supports_endpoint(Endpoint::Embeddings)); - assert!(!llm.supports_endpoint(Endpoint::Rerank)); - } - - #[test] - fn test_model_type_as_capability_names() { - let llm = ModelType::LLM; - let names = llm.as_capability_names(); - assert!(names.contains(&"chat")); - assert!(names.contains(&"completions")); - assert!(names.contains(&"responses")); - assert!(names.contains(&"tools")); - assert!(!names.contains(&"embeddings")); - } - - #[test] - fn test_model_type_display() { - let llm = ModelType::LLM; - let display = llm.to_string(); - assert!(display.contains("chat")); - assert!(display.contains("completions")); - } - - #[test] - fn test_model_type_is_llm() { - assert!(ModelType::LLM.is_llm()); - assert!(ModelType::CHAT.is_llm()); - assert!(!ModelType::EMBEDDINGS.is_llm()); - assert!(!ModelType::RERANK.is_llm()); - } - - #[test] - fn test_model_type_is_embedding_model() { - assert!(ModelType::EMBED_MODEL.is_embedding_model()); - assert!(ModelType::EMBEDDINGS.is_embedding_model()); - // LLM with embeddings is not an "embedding model" - let llm_with_embed = ModelType::LLM | ModelType::EMBEDDINGS; - assert!(!llm_with_embed.is_embedding_model()); - } - - #[test] - fn test_model_type_is_reranker() { - assert!(ModelType::RERANK_MODEL.is_reranker()); - assert!(ModelType::RERANK.is_reranker()); - // LLM with rerank is not a "reranker" - let llm_with_rerank = ModelType::LLM | ModelType::RERANK; - assert!(!llm_with_rerank.is_reranker()); - } - - #[test] - fn test_model_type_default() { - let default = ModelType::default(); - assert!(default.is_empty()); - assert!(!default.supports_chat()); - } - - #[test] - fn test_model_type_serialization() { - let llm = ModelType::LLM; - let json = serde_json::to_string(&llm).unwrap(); - let deserialized: ModelType = serde_json::from_str(&json).unwrap(); - assert_eq!(llm, deserialized); - } - - #[test] - fn test_endpoint_path() { - assert_eq!(Endpoint::Chat.path(), "/v1/chat/completions"); - assert_eq!(Endpoint::Embeddings.path(), "/v1/embeddings"); - assert_eq!(Endpoint::Generate.path(), "/generate"); - } - - #[test] - fn test_endpoint_from_path() { - assert_eq!( - Endpoint::from_path("/v1/chat/completions"), - Some(Endpoint::Chat) - ); - assert_eq!( - Endpoint::from_path("/v1/chat/completions/"), - Some(Endpoint::Chat) - ); - assert_eq!( - Endpoint::from_path("/v1/embeddings"), - Some(Endpoint::Embeddings) - ); - assert_eq!(Endpoint::from_path("/unknown"), None); - } - - #[test] - fn test_endpoint_required_capability() { - assert_eq!(Endpoint::Chat.required_capability(), Some(ModelType::CHAT)); - assert_eq!( - Endpoint::Embeddings.required_capability(), - Some(ModelType::EMBEDDINGS) - ); - assert_eq!(Endpoint::Models.required_capability(), None); - } - - #[test] - fn test_endpoint_display() { - assert_eq!(Endpoint::Chat.to_string(), "chat"); - assert_eq!(Endpoint::Embeddings.to_string(), "embeddings"); - } - - #[test] - fn test_endpoint_serialization() { - let endpoint = Endpoint::Chat; - let json = serde_json::to_string(&endpoint).unwrap(); - assert_eq!(json, "\"chat\""); - let deserialized: Endpoint = serde_json::from_str(&json).unwrap(); - assert_eq!(endpoint, deserialized); - } -} diff --git a/sgl-router/src/core/worker.rs b/sgl-router/src/core/worker.rs index 621127e54..c1457df57 100644 --- a/sgl-router/src/core/worker.rs +++ b/sgl-router/src/core/worker.rs @@ -315,7 +315,7 @@ impl fmt::Display for ConnectionMode { } } -/// Runtime implementation type for gRPC workers +/// Runtime implementation type for workers #[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize, Default)] #[serde(rename_all = "lowercase")] pub enum RuntimeType { @@ -324,6 +324,9 @@ pub enum RuntimeType { Sglang, /// vLLM runtime Vllm, + /// External OpenAI-compatible API (not local inference) + /// Used for routing to external providers like OpenAI, Azure OpenAI, xAI, etc. + External, } impl fmt::Display for RuntimeType { @@ -331,6 +334,7 @@ impl fmt::Display for RuntimeType { match self { RuntimeType::Sglang => write!(f, "sglang"), RuntimeType::Vllm => write!(f, "vllm"), + RuntimeType::External => write!(f, "external"), } } } @@ -342,6 +346,7 @@ impl std::str::FromStr for RuntimeType { match s.to_lowercase().as_str() { "sglang" => Ok(RuntimeType::Sglang), "vllm" => Ok(RuntimeType::Vllm), + "external" => Ok(RuntimeType::External), _ => Err(format!("Unknown runtime type: {}", s)), } } @@ -1937,292 +1942,4 @@ mod tests { // Not found assert!(metadata.find_model("unknown-model").is_none()); } - - #[test] - fn test_worker_metadata_supports_model_with_list() { - use super::ModelCard; - - let model1 = ModelCard::new("model-a").with_alias("alias-a"); - let model2 = ModelCard::new("model-b"); - - let metadata = WorkerMetadata { - url: "http://test:8080".to_string(), - worker_type: WorkerType::Regular, - connection_mode: ConnectionMode::Http, - runtime_type: RuntimeType::default(), - labels: std::collections::HashMap::new(), - health_config: HealthConfig::default(), - api_key: None, - bootstrap_host: "test".to_string(), - bootstrap_port: None, - models: vec![model1, model2], - default_provider: None, - default_model_type: ModelType::LLM, - }; - - // Should support listed models - assert!(metadata.supports_model("model-a")); - assert!(metadata.supports_model("alias-a")); - assert!(metadata.supports_model("model-b")); - - // Should not support unlisted models - assert!(!metadata.supports_model("model-c")); - } - - #[test] - fn test_worker_metadata_supports_endpoint() { - use super::{Endpoint, ModelCard}; - - let embed_model = - ModelCard::new("text-embedding-3-small").with_model_type(ModelType::EMBEDDINGS); - let llm_model = ModelCard::new("gpt-4o").with_model_type(ModelType::LLM); - - let metadata = WorkerMetadata { - url: "http://test:8080".to_string(), - worker_type: WorkerType::Regular, - connection_mode: ConnectionMode::Http, - runtime_type: RuntimeType::default(), - labels: std::collections::HashMap::new(), - health_config: HealthConfig::default(), - api_key: None, - bootstrap_host: "test".to_string(), - bootstrap_port: None, - models: vec![embed_model, llm_model], - default_provider: None, - default_model_type: ModelType::LLM, - }; - - // Embedding model supports embeddings but not chat - assert!(metadata.supports_endpoint("text-embedding-3-small", Endpoint::Embeddings)); - assert!(!metadata.supports_endpoint("text-embedding-3-small", Endpoint::Chat)); - - // LLM model supports chat but not embeddings - assert!(metadata.supports_endpoint("gpt-4o", Endpoint::Chat)); - assert!(!metadata.supports_endpoint("gpt-4o", Endpoint::Embeddings)); - - // Unknown model falls back to default_model_type (LLM) - assert!(metadata.supports_endpoint("unknown", Endpoint::Chat)); - assert!(!metadata.supports_endpoint("unknown", Endpoint::Embeddings)); - } - - #[test] - fn test_worker_metadata_provider_for_model() { - use super::{ModelCard, ProviderType}; - - let openai_model = ModelCard::new("gpt-4o").with_provider(ProviderType::OpenAI); - let native_model = ModelCard::new("llama-3.1"); // No provider = native - - let metadata = WorkerMetadata { - url: "http://test:8080".to_string(), - worker_type: WorkerType::Regular, - connection_mode: ConnectionMode::Http, - runtime_type: RuntimeType::default(), - labels: std::collections::HashMap::new(), - health_config: HealthConfig::default(), - api_key: None, - bootstrap_host: "test".to_string(), - bootstrap_port: None, - models: vec![openai_model, native_model], - default_provider: Some(ProviderType::XAI), // Default for unknown models - default_model_type: ModelType::LLM, - }; - - // OpenAI model returns OpenAI provider - assert_eq!( - metadata.provider_for_model("gpt-4o"), - Some(&ProviderType::OpenAI) - ); - - // Native model returns None (model has no provider) - // But falls back to worker's default_provider - assert_eq!( - metadata.provider_for_model("llama-3.1"), - Some(&ProviderType::XAI) - ); - - // Unknown model returns worker's default_provider - assert_eq!( - metadata.provider_for_model("unknown"), - Some(&ProviderType::XAI) - ); - } - - #[test] - fn test_worker_metadata_model_ids() { - use super::ModelCard; - - let model1 = ModelCard::new("model-a"); - let model2 = ModelCard::new("model-b"); - let model3 = ModelCard::new("model-c"); - - let metadata = WorkerMetadata { - url: "http://test:8080".to_string(), - worker_type: WorkerType::Regular, - connection_mode: ConnectionMode::Http, - runtime_type: RuntimeType::default(), - labels: std::collections::HashMap::new(), - health_config: HealthConfig::default(), - api_key: None, - bootstrap_host: "test".to_string(), - bootstrap_port: None, - models: vec![model1, model2, model3], - default_provider: None, - default_model_type: ModelType::LLM, - }; - - 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"); - } }