use crate::capabilities::ModelCapabilities;
use serde::{Deserialize, Serialize};
use std::collections::HashSet;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub enum Capability {
Text,
Vision,
Audio,
Tools,
JsonMode,
Reasoning,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum TaskModality {
Text,
Vision,
Audio,
}
#[derive(Debug, Clone, Default)]
pub struct RouteConstraints {
pub min_context_window: Option<u32>,
pub require_tools: bool,
pub require_json_mode: bool,
pub require_reasoning: bool,
pub max_cost_tier: Option<String>,
pub min_speed_tier: Option<String>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize)]
pub enum SpeedTier {
Slow,
Medium,
Fast,
Realtime,
}
impl SpeedTier {
fn from_str(s: &str) -> Option<Self> {
match s.to_ascii_lowercase().as_str() {
"slow" => Some(SpeedTier::Slow),
"medium" => Some(SpeedTier::Medium),
"fast" => Some(SpeedTier::Fast),
"realtime" => Some(SpeedTier::Realtime),
_ => None,
}
}
fn label(&self) -> &'static str {
match self {
SpeedTier::Slow => "slow",
SpeedTier::Medium => "medium",
SpeedTier::Fast => "fast",
SpeedTier::Realtime => "realtime",
}
}
}
impl Capability {
fn label(&self) -> &'static str {
match self {
Capability::Text => "text",
Capability::Vision => "vision",
Capability::Audio => "audio",
Capability::Tools => "tools",
Capability::JsonMode => "json",
Capability::Reasoning => "reasoning",
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ModelProfile {
pub model_id: String,
pub provider: String,
pub capabilities: HashSet<Capability>,
pub context_window: u32,
pub speed_tier: SpeedTier,
pub cost_tier: String,
#[serde(default)]
pub notes: String,
#[serde(default)]
pub quality_score: Option<u8>,
}
#[derive(Default)]
pub struct ModelCatalog {
static_profiles: Vec<ModelProfile>,
runtime_profiles: Vec<ModelProfile>,
}
impl ModelCatalog {
pub fn new() -> Self {
Self::default()
}
pub fn with_static_table(models: Vec<(String, String, ModelCapabilities)>) -> Self {
let mut catalog = Self::new();
for (model_id, provider, caps) in models {
catalog
.static_profiles
.push(profile_from_capabilities(&model_id, &provider, &caps, None));
}
catalog
}
pub fn register_runtime(&mut self, profile: ModelProfile) {
self.runtime_profiles
.retain(|p| !(p.provider == profile.provider && p.model_id == profile.model_id));
self.runtime_profiles.push(profile);
}
pub fn extend_runtime<I: IntoIterator<Item = ModelProfile>>(&mut self, profiles: I) {
for p in profiles {
self.register_runtime(p);
}
}
pub fn profiles(&self) -> Vec<ModelProfile> {
let mut out = Vec::with_capacity(self.static_profiles.len() + self.runtime_profiles.len());
out.extend(self.runtime_profiles.iter().cloned());
for s in &self.static_profiles {
if !out
.iter()
.any(|p| p.provider == s.provider && p.model_id == s.model_id)
{
out.push(s.clone());
}
}
out
}
pub fn get(&self, provider: &str, model_id: &str) -> Option<ModelProfile> {
self.profiles()
.into_iter()
.find(|p| p.provider == provider && p.model_id == model_id)
}
pub fn lean_hint(&self) -> String {
let mut line = String::from("Models:");
for p in self.profiles() {
line.push_str(&format!(" {}({});", p.model_id, capability_labels(&p)));
}
line
}
pub fn describe_full(&self, provider: &str, model_id: &str) -> Option<String> {
let p = self.get(provider, model_id)?;
Some(format!(
"{} [{}] context={} speed={} cost={}{}",
p.model_id,
p.provider,
p.context_window,
p.speed_tier.label(),
p.cost_tier,
if p.notes.is_empty() {
String::new()
} else {
format!(" notes={}", p.notes)
},
))
}
pub fn route(&self, modality: TaskModality, constraints: &RouteConstraints) -> Option<String> {
let required = required_capability(modality)?;
self.profiles()
.into_iter()
.filter(|p| p.capabilities.contains(&required))
.filter(|p| {
constraints
.min_context_window
.is_none_or(|w| p.context_window >= w)
})
.filter(|p| !constraints.require_tools || p.capabilities.contains(&Capability::Tools))
.filter(|p| {
!constraints.require_json_mode || p.capabilities.contains(&Capability::JsonMode)
})
.filter(|p| {
!constraints.require_reasoning || p.capabilities.contains(&Capability::Reasoning)
})
.filter(|p| {
constraints
.max_cost_tier
.as_deref()
.is_none_or(|max| cost_at_most(&p.cost_tier, max))
})
.filter(|p| {
constraints
.min_speed_tier
.as_deref()
.and_then(SpeedTier::from_str)
.is_none_or(|min| p.speed_tier >= min)
})
.min_by_key(cost_rank_and_speed)
.map(|p| p.model_id)
}
}
impl From<&NvidiaLikeEntry> for ModelProfile {
fn from(e: &NvidiaLikeEntry) -> Self {
ModelProfile {
model_id: e.id.clone(),
provider: "nvidia".to_string(),
capabilities: HashSet::from([Capability::Text]),
context_window: 128_000,
speed_tier: SpeedTier::Medium,
cost_tier: "low".to_string(),
notes: format!("owned_by {}", e.owned_by),
quality_score: Some(e.quality_score),
}
}
}
#[derive(Debug, Clone)]
pub struct NvidiaLikeEntry {
pub id: String,
pub owned_by: String,
pub quality_score: u8,
}
impl From<&crate::nvidia_catalog::CatalogEntry> for NvidiaLikeEntry {
fn from(e: &crate::nvidia_catalog::CatalogEntry) -> Self {
NvidiaLikeEntry {
id: e.id.clone(),
owned_by: e.owned_by.clone(),
quality_score: e.quality_score,
}
}
}
#[allow(clippy::too_many_lines)]
fn profile_from_capabilities(
model_id: &str,
provider: &str,
caps: &ModelCapabilities,
quality_score: Option<u8>,
) -> ModelProfile {
let mut capabilities = HashSet::new();
capabilities.insert(Capability::Text);
if caps.supports_vision {
capabilities.insert(Capability::Vision);
}
if caps.supports_audio {
capabilities.insert(Capability::Audio);
}
if caps.supports_tools {
capabilities.insert(Capability::Tools);
}
if caps.supports_json_mode {
capabilities.insert(Capability::JsonMode);
}
if caps.supports_reasoning {
capabilities.insert(Capability::Reasoning);
}
ModelProfile {
model_id: model_id.to_string(),
provider: provider.to_string(),
capabilities,
context_window: caps.context_window,
speed_tier: SpeedTier::from_str(&caps.speed_tier).unwrap_or(SpeedTier::Medium),
cost_tier: caps.cost_tier.clone(),
notes: caps
.family
.as_ref()
.map(|f| format!("family {f}"))
.unwrap_or_default(),
quality_score,
}
}
pub fn profile_for(model_id: &str, provider: &str, caps: &ModelCapabilities) -> ModelProfile {
profile_from_capabilities(model_id, provider, caps, None)
}
fn capability_labels(p: &ModelProfile) -> String {
let mut labels: Vec<&str> = p.capabilities.iter().map(Capability::label).collect();
labels.sort_unstable();
labels.join(",")
}
const fn required_capability(modality: TaskModality) -> Option<Capability> {
match modality {
TaskModality::Text => Some(Capability::Text),
TaskModality::Vision => Some(Capability::Vision),
TaskModality::Audio => Some(Capability::Audio),
}
}
fn cost_rank(tier: &str) -> u8 {
const COST_ORDER: [&str; 5] = ["free", "low", "medium", "high", "premium"];
COST_ORDER
.iter()
.position(|t| *t == tier.to_ascii_lowercase())
.unwrap_or(COST_ORDER.len()) as u8
}
fn cost_at_most(actual: &str, max: &str) -> bool {
cost_rank(actual) <= cost_rank(max)
}
fn cost_rank_and_speed(p: &ModelProfile) -> (u8, u8) {
let speed_rank = u8::MAX - (p.speed_tier as u8);
(cost_rank(&p.cost_tier), speed_rank)
}
#[cfg(test)]
pub(crate) mod test_support {
use super::{Capability, ModelCatalog, ModelProfile, SpeedTier};
pub(crate) fn profile(
model_id: &str,
provider: &str,
caps: &[Capability],
cost: &str,
speed: SpeedTier,
) -> ModelProfile {
ModelProfile {
model_id: model_id.to_string(),
provider: provider.to_string(),
capabilities: caps.iter().copied().collect(),
context_window: 32_768,
speed_tier: speed,
cost_tier: cost.to_string(),
notes: String::new(),
quality_score: None,
}
}
pub(crate) fn catalog() -> ModelCatalog {
let mut c = ModelCatalog::new();
c.extend_runtime([
profile(
"small-text",
"static",
&[Capability::Text],
"free",
SpeedTier::Fast,
),
profile(
"big-tools",
"static",
&[Capability::Text, Capability::Tools],
"high",
SpeedTier::Slow,
),
profile(
"vision-pro",
"cloud",
&[Capability::Text, Capability::Vision],
"premium",
SpeedTier::Medium,
),
profile(
"audio-max",
"cloud",
&[Capability::Text, Capability::Audio],
"premium",
SpeedTier::Slow,
),
]);
c
}
}
#[cfg(test)]
mod tests {
use super::test_support::{catalog, profile};
use super::{
Capability, ModelCatalog, ModelProfile, RouteConstraints, SpeedTier, TaskModality,
};
use crate::nvidia_catalog::{CatalogEntry, NvidiaConfig};
#[test]
fn lean_hint_under_budget() {
let hint = catalog().lean_hint();
assert!(hint.starts_with("Models:"));
for needle in [
"small-text(text)",
"big-tools(text,tools)",
"vision-pro(text,vision)",
"audio-max(audio,text)",
] {
assert!(hint.contains(needle), "missing `{needle}` in `{hint}`");
}
let words = hint.split_whitespace().count();
assert!(
words <= 50,
"lean_hint exceeded token budget: {words} tokens: {hint}"
);
}
#[test]
fn route_prefers_capable_cheapest() {
let c = catalog();
let picked = c
.route(TaskModality::Text, &RouteConstraints::default())
.unwrap();
assert_eq!(picked, "small-text");
let reqs = RouteConstraints {
require_tools: true,
..RouteConstraints::default()
};
assert_eq!(c.route(TaskModality::Text, &reqs).unwrap(), "big-tools");
}
#[test]
fn unknown_modality_returns_none() {
let c = ModelCatalog::new();
assert_eq!(
c.route(TaskModality::Text, &RouteConstraints::default()),
None
);
assert_eq!(
c.route(TaskModality::Vision, &RouteConstraints::default()),
None
);
assert_eq!(
c.route(TaskModality::Audio, &RouteConstraints::default()),
None
);
let mut audio_less = ModelCatalog::new();
audio_less.extend_runtime([profile(
"text-only",
"static",
&[Capability::Text],
"free",
SpeedTier::Fast,
)]);
assert_eq!(
audio_less.route(TaskModality::Audio, &RouteConstraints::default()),
None
);
}
#[test]
fn nvidia_catalog_still_works() {
let entry = CatalogEntry {
id: "meta/llama-3.3-70b-instruct".to_string(),
owned_by: "meta".to_string(),
created: 1_700_000_000,
quality_score: 93,
};
let mut c = ModelCatalog::new();
let converted: ModelProfile = (&super::NvidiaLikeEntry::from(&entry)).into();
assert_eq!(converted.provider, "nvidia");
assert!(converted.quality_score == Some(93));
c.extend_runtime([converted]);
let cfg = NvidiaConfig::default();
assert_eq!(cfg.default_model, "meta/llama-3.3-70b-instruct");
let cache = crate::nvidia_catalog::NvidiaCatalogCache::new(cfg);
assert!(cache.snapshot().is_empty());
assert_eq!(
c.get("nvidia", "meta/llama-3.3-70b-instruct")
.unwrap()
.model_id,
"meta/llama-3.3-70b-instruct"
);
}
#[test]
fn runtime_shadows_static_per_key() {
let mut c = catalog();
c.register_runtime(profile(
"small-text",
"static",
&[Capability::Text],
"low",
SpeedTier::Realtime,
));
let merged = c.profiles();
assert_eq!(
merged.iter().filter(|p| p.model_id == "small-text").count(),
1
);
assert_eq!(
merged
.iter()
.find(|p| p.model_id == "small-text")
.unwrap()
.speed_tier,
SpeedTier::Realtime
);
}
#[test]
fn describe_full_includes_all_fields() {
let c = catalog();
let d = c.describe_full("cloud", "vision-pro").unwrap();
assert!(d.contains("vision-pro"));
assert!(d.contains("[cloud]"));
assert!(d.contains("context=32768"));
assert!(d.contains("speed=medium"));
assert!(d.contains("cost=premium"));
assert!(c.describe_full("cloud", "nope").is_none());
}
#[test]
fn route_respects_context_window_and_cost_ceiling() {
let c = catalog();
let reqs = RouteConstraints {
min_context_window: Some(64_000),
..Default::default()
};
assert!(c.route(TaskModality::Text, &reqs).is_none());
let reqs = RouteConstraints {
max_cost_tier: Some("low".to_string()),
..Default::default()
};
assert_eq!(c.route(TaskModality::Text, &reqs).unwrap(), "small-text");
}
#[test]
fn every_profile_advertises_text_baseline() {
for p in catalog().profiles() {
assert!(
p.capabilities.contains(&Capability::Text),
"profile {} lost the text baseline",
p.model_id
);
}
}
}