[model-gateway] add audio and moderation in model card (#14263)

This commit is contained in:
Simo Lin
2025-12-01 18:36:38 -08:00
committed by GitHub
parent 8fe8b63576
commit cd4151abc7
3 changed files with 60 additions and 703 deletions

View File

@@ -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");
}
}

View File

@@ -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);
}
}

View File

@@ -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");
}
}