use schemars::JsonSchema;
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Default, Serialize, Deserialize, JsonSchema, PartialEq, Eq)]
#[serde(rename_all = "snake_case")]
pub enum ModelSource {
#[default]
Builtin,
Manual,
ModelsDev { refreshed_at: String },
}
#[derive(Debug, Clone, Default, Serialize, Deserialize, JsonSchema, PartialEq)]
pub struct ModelPricing {
pub input_per_1m: Option<f64>,
pub output_per_1m: Option<f64>,
pub cache_read_per_1m: Option<f64>,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize, JsonSchema, PartialEq, Eq)]
pub struct ModelCapabilities {
#[serde(default)]
pub tools: bool,
#[serde(default)]
pub structured_output: bool,
#[serde(default)]
pub reasoning: bool,
#[serde(default)]
pub image_input: bool,
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum ImageInputCapability {
Supported,
Unsupported,
#[default]
Unknown,
}
impl ImageInputCapability {
pub fn from_metadata(metadata: Option<&ModelMetadata>) -> Self {
match metadata {
Some(m) if m.capabilities.image_input => Self::Supported,
Some(_) => Self::Unsupported,
None => Self::Unknown,
}
}
pub fn allows_attachment(self) -> bool {
matches!(self, Self::Supported)
}
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq, JsonSchema)]
#[serde(rename_all = "lowercase")]
pub enum ReasoningEffort {
Low,
Medium,
High,
}
#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema, PartialEq, Eq)]
pub enum CatalogProviderProtocol {
#[serde(rename = "anthropic-messages")]
AnthropicMessages,
#[serde(rename = "openai-chat")]
OpenAIChat,
}
#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema)]
pub struct ModelMetadata {
pub id: String,
pub provider: String,
pub context_limit: Option<u32>,
pub output_limit: Option<u32>,
#[serde(default)]
pub pricing: Option<ModelPricing>,
#[serde(default)]
pub capabilities: ModelCapabilities,
pub release_date: Option<String>,
#[serde(default)]
pub source: ModelSource,
#[serde(default)]
pub variants: Vec<VariantDef>,
}
pub fn find_model<'a>(models: &'a [ModelMetadata], id: &str) -> Option<&'a ModelMetadata> {
models.iter().find(|m| m.id == id)
}
pub fn find_model_by_provider<'a>(
models: &'a [ModelMetadata],
provider: &str,
id: &str,
) -> Option<&'a ModelMetadata> {
models.iter().find(|m| m.provider == provider && m.id == id)
}
pub fn models_with_id<'a>(models: &'a [ModelMetadata], id: &str) -> Vec<&'a ModelMetadata> {
models.iter().filter(|m| m.id == id).collect()
}
#[derive(Debug, Clone, Default, Serialize, Deserialize, JsonSchema, PartialEq, Eq)]
pub struct ProviderInfo {
pub id: String,
pub name: String,
pub api_base_url: Option<String>,
#[serde(default)]
pub protocol: Option<CatalogProviderProtocol>,
pub env_var: Option<String>,
pub doc_url: Option<String>,
#[serde(default)]
pub source: ProviderSource,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize, JsonSchema, PartialEq, Eq)]
#[serde(rename_all = "snake_case")]
pub enum ProviderSource {
#[default]
Builtin,
ModelsDev { refreshed_at: String },
}
#[cfg(test)]
#[allow(clippy::unwrap_used)]
mod tests {
use super::*;
#[test]
fn test_model_source_default_is_builtin() {
assert_eq!(ModelSource::default(), ModelSource::Builtin);
}
#[test]
fn test_model_capabilities_default_all_false() {
let caps = ModelCapabilities::default();
assert!(!caps.tools);
assert!(!caps.structured_output);
assert!(!caps.reasoning);
assert!(!caps.image_input);
}
#[test]
fn test_model_pricing_default_all_none() {
let pricing = ModelPricing::default();
assert!(pricing.input_per_1m.is_none());
assert!(pricing.output_per_1m.is_none());
assert!(pricing.cache_read_per_1m.is_none());
}
#[test]
fn test_find_model_by_provider_resolves() {
let models = vec![
ModelMetadata {
id: "glm-5.2".to_string(),
provider: "zhipu".to_string(),
context_limit: Some(128_000),
output_limit: None,
pricing: None,
capabilities: ModelCapabilities::default(),
release_date: None,
variants: vec![],
source: ModelSource::Builtin,
},
ModelMetadata {
id: "glm-5.2".to_string(),
provider: "zai".to_string(),
context_limit: Some(128_000),
output_limit: None,
pricing: None,
capabilities: ModelCapabilities::default(),
release_date: None,
variants: vec![],
source: ModelSource::Builtin,
},
];
let zhipu = find_model_by_provider(&models, "zhipu", "glm-5.2");
assert!(zhipu.is_some());
assert_eq!(zhipu.expect("operation should succeed").provider, "zhipu");
let zai = find_model_by_provider(&models, "zai", "glm-5.2");
assert!(zai.is_some());
assert_eq!(zai.expect("operation should succeed").provider, "zai");
assert!(find_model_by_provider(&models, "openai", "glm-5.2").is_none());
}
#[test]
fn test_models_with_id_detects_ambiguity() {
let models = vec![
ModelMetadata {
id: "shared".to_string(),
provider: "a".to_string(),
context_limit: None,
output_limit: None,
pricing: None,
capabilities: ModelCapabilities::default(),
release_date: None,
variants: vec![],
source: ModelSource::Builtin,
},
ModelMetadata {
id: "shared".to_string(),
provider: "b".to_string(),
context_limit: None,
output_limit: None,
pricing: None,
capabilities: ModelCapabilities::default(),
release_date: None,
variants: vec![],
source: ModelSource::Builtin,
},
ModelMetadata {
id: "unique".to_string(),
provider: "a".to_string(),
context_limit: None,
output_limit: None,
pricing: None,
capabilities: ModelCapabilities::default(),
release_date: None,
variants: vec![],
source: ModelSource::Builtin,
},
];
assert_eq!(models_with_id(&models, "shared").len(), 2);
assert_eq!(models_with_id(&models, "unique").len(), 1);
assert!(models_with_id(&models, "missing").is_empty());
}
#[test]
fn test_provider_info_default() {
let info = ProviderInfo::default();
assert!(info.id.is_empty());
assert!(info.name.is_empty());
assert!(info.api_base_url.is_none());
assert!(info.protocol.is_none());
assert!(info.env_var.is_none());
assert!(info.doc_url.is_none());
assert_eq!(info.source, ProviderSource::Builtin);
}
#[test]
fn test_model_metadata_serde_roundtrip() {
let meta = ModelMetadata {
id: "test-model".to_string(),
provider: "test".to_string(),
context_limit: Some(200_000),
output_limit: Some(8_192),
pricing: Some(ModelPricing {
input_per_1m: Some(3.0),
output_per_1m: Some(15.0),
cache_read_per_1m: Some(0.3),
}),
capabilities: ModelCapabilities {
tools: true,
structured_output: false,
reasoning: true,
image_input: true,
},
release_date: Some("2025-01-01".to_string()),
variants: vec![],
source: ModelSource::ModelsDev {
refreshed_at: "2025-07-03T00:00:00Z".to_string(),
},
};
let json = serde_json::to_string(&meta).expect("serialize");
let roundtrip: ModelMetadata = serde_json::from_str(&json).expect("deserialize");
assert_eq!(meta.id, roundtrip.id);
assert_eq!(meta.provider, roundtrip.provider);
assert_eq!(meta.context_limit, roundtrip.context_limit);
assert_eq!(meta.output_limit, roundtrip.output_limit);
assert_eq!(meta.capabilities, roundtrip.capabilities);
assert_eq!(meta.source, roundtrip.source);
}
#[test]
fn image_input_capability_supported_when_metadata_image_input_true() {
let metadata = ModelMetadata {
id: "test-model".into(),
provider: "test".into(),
context_limit: None,
output_limit: None,
pricing: None,
capabilities: ModelCapabilities {
image_input: true,
..Default::default()
},
release_date: None,
source: ModelSource::default(),
variants: vec![],
};
let cap = ImageInputCapability::from_metadata(Some(&metadata));
assert_eq!(cap, ImageInputCapability::Supported);
assert!(cap.allows_attachment());
}
#[test]
fn image_input_capability_unsupported_when_metadata_image_input_false() {
let metadata = ModelMetadata {
id: "test-model".into(),
provider: "test".into(),
context_limit: None,
output_limit: None,
pricing: None,
capabilities: ModelCapabilities {
image_input: false,
..Default::default()
},
release_date: None,
source: ModelSource::default(),
variants: vec![],
};
let cap = ImageInputCapability::from_metadata(Some(&metadata));
assert_eq!(cap, ImageInputCapability::Unsupported);
assert!(!cap.allows_attachment());
}
#[test]
fn image_input_capability_unknown_when_no_metadata() {
let cap = ImageInputCapability::from_metadata(None);
assert_eq!(cap, ImageInputCapability::Unknown);
assert!(!cap.allows_attachment());
}
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq, JsonSchema)]
pub struct VariantDef {
pub id: String,
pub label: String,
#[serde(default)]
pub reasoning_effort: Option<ReasoningEffort>,
}