use crate::route::{RoutePolicy, RouteTier};
use crate::secret::SecretRef;
use serde::{Deserialize, Serialize};
use std::collections::BTreeMap;
use std::path::{Path, PathBuf};
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Default)]
pub struct ModelEntry {
#[serde(default)]
pub provider: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub vendor: Option<String>,
#[serde(default)]
pub model: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub base_url: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub secret: Option<SecretRef>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub capabilities: Vec<String>,
#[serde(default, skip_serializing_if = "serde_json::Value::is_null")]
pub params: serde_json::Value,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub tier: Option<RouteTier>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub cost_per_1k_tokens: Option<f64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub input_cost_per_1k: Option<f64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub output_cost_per_1k: Option<f64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub context_window: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub priced_at: Option<chrono::DateTime<chrono::Utc>>,
}
fn vendor_label_of_url(base_url: Option<&str>) -> Option<String> {
let host = base_url?
.split("//")
.nth(1)
.unwrap_or(base_url?)
.split(['/', ':'])
.next()
.unwrap_or("");
let label = host.strip_prefix("api.").unwrap_or(host);
let first = label.split('.').next().unwrap_or("");
(!first.is_empty()).then(|| first.to_string())
}
impl ModelEntry {
pub fn effective_costs(&self) -> (Option<f64>, Option<f64>) {
let output = self.output_cost_per_1k.or(self.cost_per_1k_tokens);
let input = self.input_cost_per_1k.or(self.cost_per_1k_tokens);
(input, output)
}
pub fn vendor_candidates(&self) -> Vec<String> {
let mut out: Vec<String> = Vec::with_capacity(3);
let mut push = |v: &str| {
if !v.is_empty() && !out.iter().any(|e| e == v) {
out.push(v.to_string());
}
};
if let Some(v) = self.vendor.as_deref() {
push(v);
}
if let Some(label) = vendor_label_of_url(self.base_url.as_deref()) {
push(&label);
}
push(&self.provider);
out
}
pub fn is_priced(&self) -> bool {
let (input, output) = self.effective_costs();
input.is_some() || output.is_some()
}
pub fn stamp_priced_at(&mut self, now: chrono::DateTime<chrono::Utc>) {
if self.is_priced() && self.priced_at.is_none_or(|prev| prev < now) {
self.priced_at = Some(now);
}
}
pub fn price_age(&self, now: chrono::DateTime<chrono::Utc>) -> Option<chrono::TimeDelta> {
self.priced_at.map(|at| now - at)
}
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Default)]
pub struct RoleEntry {
pub primary: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub fallback: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub cost_budget_per_day_usd: Option<f64>,
#[serde(default)]
pub privacy_local_only: bool,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub route_policy: Option<RoutePolicy>,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct ModelRegistry {
pub schema_version: u32,
#[serde(default)]
pub models: BTreeMap<String, ModelEntry>,
#[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
pub roles: BTreeMap<String, RoleEntry>,
}
impl Default for ModelRegistry {
fn default() -> Self {
Self {
schema_version: 1,
models: BTreeMap::new(),
roles: BTreeMap::new(),
}
}
}
impl ModelRegistry {
pub fn load_from(path: &Path) -> anyhow::Result<Self> {
if !path.exists() {
return Ok(Self::default());
}
let body = std::fs::read_to_string(path)?;
if body.trim().is_empty() {
return Ok(Self::default());
}
Ok(serde_yaml_ng::from_str(&body)?)
}
pub fn save_to(&self, path: &Path) -> anyhow::Result<()> {
if let Some(parent) = path.parent() {
std::fs::create_dir_all(parent)?;
}
let body = serde_yaml_ng::to_string(self)?;
let tmp = path.with_extension("yaml.tmp");
std::fs::write(&tmp, body)?;
std::fs::rename(&tmp, path)?;
Ok(())
}
pub fn default_path() -> anyhow::Result<PathBuf> {
if let Ok(p) = std::env::var("MUR_HOME")
&& !p.is_empty()
{
return Ok(PathBuf::from(p).join("models.yaml"));
}
let home = dirs::home_dir().ok_or_else(|| anyhow::anyhow!("no home dir"))?;
Ok(home.join(".mur/models.yaml"))
}
pub fn resolve_role(&self, role: &str) -> Option<&str> {
let entry = self.roles.get(role)?;
if self.models.contains_key(&entry.primary) {
return Some(&entry.primary);
}
if let Some(fb) = &entry.fallback
&& self.models.contains_key(fb)
{
return Some(fb);
}
None
}
}
use crate::agent::AgentProfile;
use crate::config::{DEFAULT_ROUTING_THRESHOLD, ModelSwitchConfig, RoutingConfig};
pub fn resolve_model_refs(
profile: &AgentProfile,
cfg: &ModelSwitchConfig,
routed_primary: Option<String>,
) -> Vec<String> {
let primary = routed_primary
.or_else(|| profile.model_ref.clone())
.or_else(|| cfg.default.clone());
let chain = if !profile.fallback_chain.is_empty() {
profile.fallback_chain.clone()
} else {
cfg.fallback_chain.clone()
};
let mut out: Vec<String> = Vec::new();
if let Some(p) = primary {
out.push(p);
}
for r in chain {
if !out.contains(&r) {
out.push(r);
}
}
out
}
pub fn choose_by_difficulty(est_input_tokens: u32, r: &RoutingConfig) -> Option<String> {
let threshold = r
.threshold_input_tokens
.unwrap_or(DEFAULT_ROUTING_THRESHOLD);
match (r.cheap.as_ref(), r.frontier.as_ref()) {
(Some(cheap), Some(frontier)) => Some(if est_input_tokens > threshold {
frontier.clone()
} else {
cheap.clone()
}),
_ => None,
}
}
pub const CAP_CHAT: &str = "chat";
pub const CAP_TOOLS: &str = "tools";
pub const CAP_VISION: &str = "vision";
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Requirement {
Vision,
Tools,
}
impl Requirement {
pub fn capability(self) -> &'static str {
match self {
Requirement::Vision => CAP_VISION,
Requirement::Tools => CAP_TOOLS,
}
}
fn permitted_when_undeclared(self) -> bool {
match self {
Requirement::Vision => false,
Requirement::Tools => true,
}
}
}
pub fn satisfies(e: &ModelEntry, reqs: &[Requirement]) -> bool {
let chat_capable = e.capabilities.is_empty() || e.capabilities.iter().any(|c| c == CAP_CHAT);
if !chat_capable {
return false;
}
reqs.iter().all(|r| {
if e.capabilities.is_empty() {
r.permitted_when_undeclared()
} else {
e.capabilities.iter().any(|c| c == r.capability())
}
})
}
pub fn pick_cheap_model(
reg: &ModelRegistry,
exclude: Option<&str>,
reqs: &[Requirement],
) -> Option<String> {
reg.models
.iter()
.filter(|(k, _)| exclude != Some(k.as_str()))
.filter(|(_, e)| satisfies(e, reqs))
.filter_map(|(k, e)| {
let (input, output) = e.effective_costs();
output.or(input).map(|c| (c, k.clone()))
})
.min_by(|a, b| a.0.partial_cmp(&b.0).unwrap_or(std::cmp::Ordering::Equal))
.map(|(_, k)| k)
}
#[cfg(test)]
mod tests {
#[test]
fn vendor_candidates_prefer_the_recorded_vendor_then_the_host_then_provider() {
let e = ModelEntry {
provider: "openai".into(),
vendor: Some("deepseek".into()),
base_url: Some("https://api.deepseek.com/v1".into()),
..Default::default()
};
assert_eq!(e.vendor_candidates(), vec!["deepseek", "openai"]);
let legacy = ModelEntry {
provider: "openai".into(),
base_url: Some("https://api.deepseek.com/v1".into()),
..Default::default()
};
assert_eq!(legacy.vendor_candidates(), vec!["deepseek", "openai"]);
let bare = ModelEntry {
provider: "anthropic".into(),
..Default::default()
};
assert_eq!(bare.vendor_candidates(), vec!["anthropic"]);
let same = ModelEntry {
provider: "openai".into(),
base_url: Some("https://api.openai.com/v1".into()),
..Default::default()
};
assert_eq!(same.vendor_candidates(), vec!["openai"]);
}
#[test]
fn vendor_is_omitted_from_yaml_when_absent_and_round_trips_when_set() {
let bare = ModelEntry {
provider: "anthropic".into(),
model: "claude-opus-5".into(),
..Default::default()
};
let y = serde_yaml_ng::to_string(&bare).unwrap();
assert!(!y.contains("vendor"), "{y}");
let tagged = ModelEntry {
provider: "openai".into(),
vendor: Some("groq".into()),
model: "llama-3.3".into(),
..Default::default()
};
let y = serde_yaml_ng::to_string(&tagged).unwrap();
let back: ModelEntry = serde_yaml_ng::from_str(&y).unwrap();
assert_eq!(back.vendor.as_deref(), Some("groq"));
}
use super::*;
#[test]
fn parses_full_registry() {
let yaml = r#"
schema_version: 1
models:
anthropic_opus_4_7:
provider: anthropic
model: claude-opus-4-7
secret: env:ANTHROPIC_API_KEY
capabilities: [chat, tools]
ollama_llama3:
provider: ollama
model: llama3.2:3b
base_url: http://127.0.0.1:11434
"#;
let r: ModelRegistry = serde_yaml_ng::from_str(yaml).unwrap();
assert_eq!(r.schema_version, 1);
assert_eq!(r.models.len(), 2);
let opus = r.models.get("anthropic_opus_4_7").unwrap();
assert_eq!(opus.provider, "anthropic");
assert_eq!(
opus.secret,
Some(SecretRef::Env("ANTHROPIC_API_KEY".into()))
);
assert!(r.models["ollama_llama3"].secret.is_none());
}
#[test]
fn round_trip_preserves_shape() {
let mut r = ModelRegistry::default();
r.models.insert(
"foo".into(),
ModelEntry {
provider: "anthropic".into(),
model: "claude-opus-4-7".into(),
base_url: None,
secret: Some(SecretRef::Keychain {
service: "mur".into(),
account: "anthropic".into(),
}),
capabilities: vec!["chat".into()],
params: serde_json::Value::Null,
tier: None,
cost_per_1k_tokens: None,
input_cost_per_1k: None,
output_cost_per_1k: None,
context_window: None,
priced_at: None,
..Default::default()
},
);
let s = serde_yaml_ng::to_string(&r).unwrap();
let parsed: ModelRegistry = serde_yaml_ng::from_str(&s).unwrap();
assert_eq!(r, parsed);
}
#[test]
fn rejects_unknown_secret_scheme() {
let yaml = r#"
schema_version: 1
models:
bad:
provider: x
model: y
secret: bogus:value
"#;
let r: Result<ModelRegistry, _> = serde_yaml_ng::from_str(yaml);
assert!(r.is_err(), "should reject unknown scheme");
}
#[test]
fn test_registry_roundtrip_with_roles() {
let yaml = r#"
schema_version: 1
models:
haiku:
provider: anthropic
model: claude-haiku-4-5
roles:
reflector:
primary: haiku
fallback: null
cost_budget_per_day_usd: 0.5
"#;
let reg: ModelRegistry = serde_yaml_ng::from_str(yaml).unwrap();
assert_eq!(reg.roles["reflector"].primary, "haiku");
let back = serde_yaml_ng::to_string(®).unwrap();
let reg2: ModelRegistry = serde_yaml_ng::from_str(&back).unwrap();
assert_eq!(reg, reg2);
}
#[test]
fn test_resolve_role_primary() {
let mut reg = ModelRegistry::default();
reg.models.insert(
"haiku".into(),
ModelEntry {
provider: "anthropic".into(),
model: "claude-haiku-4-5".into(),
base_url: None,
secret: None,
capabilities: vec![],
params: serde_json::Value::Null,
tier: None,
cost_per_1k_tokens: None,
input_cost_per_1k: None,
output_cost_per_1k: None,
context_window: None,
priced_at: None,
..Default::default()
},
);
reg.roles.insert(
"reflector".into(),
RoleEntry {
primary: "haiku".into(),
fallback: None,
..Default::default()
},
);
assert_eq!(reg.resolve_role("reflector"), Some("haiku"));
}
#[test]
fn test_resolve_role_fallback() {
let mut reg = ModelRegistry::default();
reg.models.insert(
"haiku".into(),
ModelEntry {
provider: "anthropic".into(),
model: "claude-haiku-4-5".into(),
base_url: None,
secret: None,
capabilities: vec![],
params: serde_json::Value::Null,
tier: None,
cost_per_1k_tokens: None,
input_cost_per_1k: None,
output_cost_per_1k: None,
context_window: None,
priced_at: None,
..Default::default()
},
);
reg.roles.insert(
"reflector".into(),
RoleEntry {
primary: "nonexistent".into(),
fallback: Some("haiku".into()),
..Default::default()
},
);
assert_eq!(reg.resolve_role("reflector"), Some("haiku"));
}
#[test]
fn test_resolve_role_none() {
let reg = ModelRegistry::default();
assert_eq!(reg.resolve_role("reflector"), None);
}
#[test]
fn model_entry_parses_tier_field() {
let yaml = r#"
schema_version: 1
models:
haiku:
provider: anthropic
model: claude-haiku-4-5
tier: local
opus:
provider: anthropic
model: claude-opus-4-7
tier: frontier
cost_per_1k_tokens: 0.015
"#;
let r: ModelRegistry = serde_yaml_ng::from_str(yaml).unwrap();
assert_eq!(r.models["haiku"].tier, Some(RouteTier::Local));
assert_eq!(r.models["opus"].tier, Some(RouteTier::Frontier));
assert_eq!(r.models["opus"].cost_per_1k_tokens, Some(0.015));
let mut r2 = ModelRegistry::default();
r2.models.insert(
"x".into(),
ModelEntry {
provider: "ollama".into(),
model: "llama3".into(),
base_url: None,
secret: None,
capabilities: vec![],
params: serde_json::Value::Null,
tier: None,
cost_per_1k_tokens: None,
input_cost_per_1k: None,
output_cost_per_1k: None,
context_window: None,
priced_at: None,
..Default::default()
},
);
let yaml = serde_yaml_ng::to_string(&r2).unwrap();
assert!(
!yaml.contains("tier:"),
"absent tier should not be serialized: {yaml}"
);
}
#[test]
fn role_entry_parses_route_policy() {
let yaml = r#"
schema_version: 1
models:
haiku:
provider: anthropic
model: claude-haiku-4-5
opus:
provider: anthropic
model: claude-opus-4-7
roles:
dev:
primary: opus
route_policy: !force_frontier
model_id: opus
reflector:
primary: haiku
route_policy: prefer_local
curator:
primary: haiku
route_policy: force_local
chat:
primary: haiku
"#;
let r: ModelRegistry = serde_yaml_ng::from_str(yaml).unwrap();
assert_eq!(
r.roles["dev"].route_policy,
Some(RoutePolicy::ForceFrontier {
model_id: "opus".into()
})
);
assert_eq!(
r.roles["reflector"].route_policy,
Some(RoutePolicy::PreferLocal)
);
assert_eq!(
r.roles["curator"].route_policy,
Some(RoutePolicy::ForceLocal)
);
assert_eq!(r.roles["chat"].route_policy, None);
}
#[test]
fn parses_split_cost_fields() {
let yaml = r#"
schema_version: 1
models:
opus:
provider: anthropic
model: claude-opus-4-8
input_cost_per_1k: 0.005
output_cost_per_1k: 0.025
context_window: 200000
"#;
let r: ModelRegistry = serde_yaml_ng::from_str(yaml).unwrap();
let e = r.models.get("opus").unwrap();
assert_eq!(e.input_cost_per_1k, Some(0.005));
assert_eq!(e.output_cost_per_1k, Some(0.025));
assert_eq!(e.context_window, Some(200_000));
}
#[test]
fn default_model_entry_is_empty() {
let e = ModelEntry::default();
assert!(e.provider.is_empty());
assert_eq!(e.input_cost_per_1k, None);
assert_eq!(e.output_cost_per_1k, None);
assert_eq!(e.context_window, None);
}
#[test]
fn effective_costs_fallback_matrix() {
let mut e = ModelEntry {
cost_per_1k_tokens: Some(0.01),
..Default::default()
};
assert_eq!(e.effective_costs(), (Some(0.01), Some(0.01)));
e = ModelEntry {
input_cost_per_1k: Some(0.005),
output_cost_per_1k: Some(0.025),
..Default::default()
};
assert_eq!(e.effective_costs(), (Some(0.005), Some(0.025)));
e = ModelEntry {
cost_per_1k_tokens: Some(0.01),
input_cost_per_1k: Some(0.005),
output_cost_per_1k: Some(0.025),
..Default::default()
};
assert_eq!(e.effective_costs(), (Some(0.005), Some(0.025)));
e = ModelEntry::default();
assert_eq!(e.effective_costs(), (None, None));
}
}
#[cfg(test)]
mod io_tests {
use super::*;
use tempfile::tempdir;
#[test]
fn load_returns_empty_when_file_missing() {
let dir = tempdir().unwrap();
let r = ModelRegistry::load_from(&dir.path().join("nope.yaml")).unwrap();
assert_eq!(r.models.len(), 0);
assert_eq!(r.schema_version, 1);
}
#[test]
fn save_then_load_round_trips() {
let dir = tempdir().unwrap();
let p = dir.path().join("models.yaml");
let mut r = ModelRegistry::default();
r.models.insert(
"x".into(),
ModelEntry {
provider: "ollama".into(),
model: "llama3.2:3b".into(),
base_url: None,
secret: None,
capabilities: vec![],
params: serde_json::Value::Null,
tier: None,
cost_per_1k_tokens: None,
input_cost_per_1k: None,
output_cost_per_1k: None,
context_window: None,
priced_at: None,
..Default::default()
},
);
r.save_to(&p).unwrap();
let r2 = ModelRegistry::load_from(&p).unwrap();
assert_eq!(r, r2);
}
#[test]
fn save_uses_atomic_rename() {
let dir = tempdir().unwrap();
let p = dir.path().join("models.yaml");
ModelRegistry::default().save_to(&p).unwrap();
let temp = dir.path().join("models.yaml.tmp");
assert!(!temp.exists(), "atomic temp left behind");
}
}
#[cfg(test)]
mod switch_tests {
use super::*;
use crate::agent::AgentProfile;
use crate::config::{ModelSwitchConfig, RoutingConfig};
fn profile(model_ref: Option<&str>, chain: &[&str]) -> AgentProfile {
let mut p = AgentProfile::default_for_tests();
p.model_ref = model_ref.map(|s| s.to_string());
p.fallback_chain = chain.iter().map(|s| s.to_string()).collect();
p
}
#[test]
fn per_agent_primary_and_chain_win_over_global() {
let cfg = ModelSwitchConfig {
default: Some("global_default".into()),
fallback_chain: vec!["g1".into(), "g2".into()],
..Default::default()
};
let p = profile(Some("agent_primary"), &["agent_primary", "agent_fb"]);
assert_eq!(
resolve_model_refs(&p, &cfg, None),
vec!["agent_primary", "agent_fb"]
);
}
#[test]
fn falls_back_to_global_default_and_chain() {
let cfg = ModelSwitchConfig {
default: Some("global_default".into()),
fallback_chain: vec!["g1".into(), "global_default".into()],
..Default::default()
};
let p = profile(None, &[]); assert_eq!(
resolve_model_refs(&p, &cfg, None),
vec!["global_default", "g1"]
);
}
#[test]
fn routed_primary_overrides_model_ref() {
let cfg = ModelSwitchConfig {
fallback_chain: vec!["g1".into()],
..Default::default()
};
let p = profile(Some("agent_primary"), &[]);
assert_eq!(
resolve_model_refs(&p, &cfg, Some("frontier".into())),
vec!["frontier", "g1"]
);
}
#[test]
fn no_config_no_agent_yields_empty() {
let cfg = ModelSwitchConfig::default();
assert!(resolve_model_refs(&profile(None, &[]), &cfg, None).is_empty());
}
#[test]
fn difficulty_picks_frontier_over_threshold() {
let r = RoutingConfig {
enabled: true,
cheap: Some("cheap".into()),
frontier: Some("frontier".into()),
threshold_input_tokens: Some(1000),
};
assert_eq!(choose_by_difficulty(1500, &r), Some("frontier".into()));
assert_eq!(choose_by_difficulty(500, &r), Some("cheap".into()));
let bad = RoutingConfig {
enabled: true,
cheap: Some("c".into()),
frontier: None,
threshold_input_tokens: None,
};
assert_eq!(choose_by_difficulty(9999, &bad), None);
}
#[test]
fn pick_cheap_model_lowest_cost_chat_excluding_primary() {
let mut reg = ModelRegistry::default();
let mk = |cost: f64, caps: &[&str]| ModelEntry {
provider: "x".into(),
model: "m".into(),
capabilities: caps.iter().map(|s| s.to_string()).collect(),
cost_per_1k_tokens: Some(cost),
..Default::default()
};
reg.models.insert("frontier".into(), mk(0.01, &["chat"]));
reg.models.insert("cheap".into(), mk(0.0001, &["chat"]));
reg.models
.insert("embed".into(), mk(0.00001, &["embedding"])); assert_eq!(
pick_cheap_model(®, Some("cheap"), &[]),
Some("frontier".into())
); assert_eq!(pick_cheap_model(®, None, &[]), Some("cheap".into()));
let mut empty = ModelRegistry::default();
empty.models.insert("e".into(), mk(0.0, &["embedding"]));
assert_eq!(pick_cheap_model(&empty, None, &[]), None);
}
#[test]
fn satisfies_is_permissive_at_baseline_and_fail_closed_above_it() {
let mk = |caps: &[&str]| ModelEntry {
provider: "x".into(),
model: "m".into(),
capabilities: caps.iter().map(|s| s.to_string()).collect(),
..Default::default()
};
assert!(satisfies(&mk(&[]), &[]));
assert!(satisfies(&mk(&["chat"]), &[]));
assert!(!satisfies(&mk(&["embedding"]), &[]));
assert!(!satisfies(&mk(&[]), &[Requirement::Vision]));
assert!(!satisfies(&mk(&["chat"]), &[Requirement::Vision]));
assert!(satisfies(&mk(&["chat", "vision"]), &[Requirement::Vision]));
assert!(!satisfies(&mk(&["chat", "vision"]), &[Requirement::Tools]));
assert!(satisfies(
&mk(&["chat", "vision", "tools"]),
&[Requirement::Vision, Requirement::Tools]
));
}
#[test]
fn undeclared_capabilities_pass_tools_but_never_vision() {
let mk = |caps: &[&str]| ModelEntry {
provider: "x".into(),
model: "m".into(),
capabilities: caps.iter().map(|s| s.to_string()).collect(),
..Default::default()
};
assert!(satisfies(&mk(&[]), &[Requirement::Tools]));
assert!(!satisfies(&mk(&[]), &[Requirement::Vision]));
assert!(!satisfies(
&mk(&[]),
&[Requirement::Vision, Requirement::Tools]
));
assert!(!satisfies(&mk(&["chat"]), &[Requirement::Tools]));
assert!(satisfies(&mk(&["chat", "tools"]), &[Requirement::Tools]));
}
#[test]
fn pick_cheap_model_declines_when_no_entry_declares_the_requirement() {
let mk = |cost: f64, caps: &[&str]| ModelEntry {
provider: "x".into(),
model: "m".into(),
capabilities: caps.iter().map(|s| s.to_string()).collect(),
cost_per_1k_tokens: Some(cost),
..Default::default()
};
let mut reg = ModelRegistry::default();
reg.models
.insert("cheap_text".into(), mk(0.0001, &["chat"]));
reg.models.insert("legacy".into(), mk(0.0002, &[]));
reg.models
.insert("frontier".into(), mk(0.01, &["chat", "vision"]));
assert_eq!(pick_cheap_model(®, None, &[]), Some("cheap_text".into()));
assert_eq!(
pick_cheap_model(®, None, &[Requirement::Vision]),
Some("frontier".into())
);
let mut blind = ModelRegistry::default();
blind
.models
.insert("cheap_text".into(), mk(0.0001, &["chat"]));
blind.models.insert("legacy".into(), mk(0.0002, &[]));
assert_eq!(pick_cheap_model(&blind, None, &[Requirement::Vision]), None);
}
#[test]
fn priced_at_stamps_only_priced_entries() {
let now = chrono::Utc::now();
let mut unpriced = ModelEntry {
provider: "openai".into(),
model: "local-thing".into(),
..Default::default()
};
unpriced.stamp_priced_at(now);
assert_eq!(unpriced.priced_at, None);
assert_eq!(unpriced.price_age(now), None);
let mut priced = ModelEntry {
output_cost_per_1k: Some(0.025),
..unpriced.clone()
};
priced.stamp_priced_at(now);
assert_eq!(priced.priced_at, Some(now));
let mut legacy = ModelEntry {
cost_per_1k_tokens: Some(0.01),
..unpriced.clone()
};
legacy.stamp_priced_at(now);
assert!(legacy.priced_at.is_some());
let earlier = now - chrono::TimeDelta::days(30);
priced.stamp_priced_at(earlier);
assert_eq!(priced.priced_at, Some(now));
}
#[test]
fn registry_without_priced_at_still_loads_and_reports_unknown_age() {
let yaml = r#"
schema_version: 1
models:
opus:
provider: anthropic
model: claude-opus-5
input_cost_per_1k: 0.005
output_cost_per_1k: 0.025
"#;
let reg: ModelRegistry = serde_yaml_ng::from_str(yaml).unwrap();
let e = ®.models["opus"];
assert_eq!(e.priced_at, None);
assert_eq!(e.price_age(chrono::Utc::now()), None);
let out = serde_yaml_ng::to_string(®).unwrap();
assert!(!out.contains("priced_at"), "{out}");
}
#[test]
fn pick_cheap_model_sees_split_cost_entries() {
let mut reg = ModelRegistry::default();
let split = |input: f64, output: f64| ModelEntry {
provider: "x".into(),
model: "m".into(),
capabilities: vec!["chat".into()],
input_cost_per_1k: Some(input),
output_cost_per_1k: Some(output),
..Default::default()
};
reg.models.insert("dear".into(), split(0.005, 0.025));
reg.models.insert("cheap".into(), split(0.0001, 0.0004));
assert_eq!(pick_cheap_model(®, None, &[]), Some("cheap".into()));
let mut input_only = ModelRegistry::default();
input_only.models.insert(
"in".into(),
ModelEntry {
provider: "x".into(),
model: "m".into(),
capabilities: vec!["chat".into()],
input_cost_per_1k: Some(0.002),
..Default::default()
},
);
assert_eq!(pick_cheap_model(&input_only, None, &[]), Some("in".into()));
}
}