[model-gateway] add audio and moderation in model card (#14263)
This commit is contained in:
@@ -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");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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");
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user