use crate::config::error::{LoadError, Warning};
use schemars::JsonSchema;
use serde::{Deserialize, Serialize};
use std::collections::{BTreeMap, BTreeSet};
use std::fs;
use std::path::{Path, PathBuf};
#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema, PartialEq, Eq, Default)]
pub struct Models {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub adapter: Option<PathBuf>,
#[serde(default)]
pub models: BTreeMap<String, Model>,
}
#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema, PartialEq, Eq)]
pub struct Model {
pub provider: String,
pub model_id: String,
pub capabilities: Capabilities,
pub context_window: u32,
}
#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema, PartialEq, Eq)]
#[serde(transparent)]
pub struct Capabilities(pub Vec<String>);
pub fn known_capabilities() -> BTreeSet<&'static str> {
[
"tool_use_native",
"prompt_caching",
"streaming",
"stop_sequences",
]
.into_iter()
.collect()
}
impl Models {
pub fn load(path: &Path) -> Result<(Self, Vec<Warning>), LoadError> {
let raw = fs::read_to_string(path).map_err(|source| LoadError::Io {
path: path.to_path_buf(),
source,
})?;
let parsed: Self = serde_yaml_ng::from_str(&raw).map_err(|source| LoadError::Yaml {
path: path.to_path_buf(),
source,
})?;
let warnings = parsed.collect_warnings(path);
Ok((parsed, warnings))
}
fn collect_warnings(&self, path: &Path) -> Vec<Warning> {
let known = known_capabilities();
let mut out = Vec::new();
for (name, model) in &self.models {
for cap in &model.capabilities.0 {
if !known.contains(cap.as_str()) {
out.push(Warning::new(
path,
format!("models.{name}.capabilities"),
format!(
"capability {cap:?} is not in the seeded registry; \
add it if intended (extend-only)"
),
));
}
}
}
out
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::io::Write;
use tempfile::NamedTempFile;
fn write_yaml(s: &str) -> NamedTempFile {
let mut f = NamedTempFile::new().unwrap();
f.write_all(s.as_bytes()).unwrap();
f
}
const ARCH_EXAMPLE: &str = r#"
models:
claude-sonnet-5:
provider: anthropic
model_id: claude-sonnet-5
capabilities: [tool_use_native, prompt_caching, streaming, stop_sequences]
context_window: 1000000
"#;
#[test]
fn parses_arch_example() {
let f = write_yaml(ARCH_EXAMPLE);
let (m, warnings) = Models::load(f.path()).unwrap();
assert!(warnings.is_empty());
assert!(m.adapter.is_none());
assert_eq!(m.models.len(), 1);
assert_eq!(m.models["claude-sonnet-5"].provider, "anthropic");
assert_eq!(m.models["claude-sonnet-5"].context_window, 1_000_000);
}
#[test]
fn parses_adapter_override() {
let yaml = r#"
adapter: /usr/local/bin/bz
models: {}
"#;
let f = write_yaml(yaml);
let (m, _) = Models::load(f.path()).unwrap();
assert_eq!(m.adapter.as_deref(), Some(Path::new("/usr/local/bin/bz")));
}
#[test]
fn warns_on_unknown_capability() {
let yaml = r#"
models:
acme-llm:
provider: acme
model_id: llm-1
capabilities: [time_travel]
context_window: 8000
"#;
let f = write_yaml(yaml);
let (_, warnings) = Models::load(f.path()).unwrap();
assert_eq!(warnings.len(), 1);
assert!(warnings[0].message.contains("time_travel"));
}
#[test]
fn surfaces_yaml_parse_errors() {
let f = write_yaml("not: [valid: yaml");
let err = Models::load(f.path()).unwrap_err();
assert!(matches!(err, LoadError::Yaml { .. }));
}
#[test]
fn surfaces_io_errors() {
let err = Models::load(Path::new("/no/such/models.yaml")).unwrap_err();
assert!(matches!(err, LoadError::Io { .. }));
}
}