use crate::driver_registry::{DiscoveredModel, DriverId, DriverRegistry, ProviderConfig};
use crate::error::{AgentLoopError, Result};
use crate::model_profiles::get_model_profile;
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct DiscoveredProviderModel {
pub model_id: String,
pub display_name: Option<String>,
pub description: Option<String>,
}
pub async fn discover_provider_models(
registry: &DriverRegistry,
config: &ProviderConfig,
) -> Result<Option<Vec<DiscoveredProviderModel>>> {
if matches!(config.provider_type, DriverId::LlmSim | DriverId::Bedrock) {
return Ok(None);
}
let driver = registry.create_chat_driver(config)?;
let models = match driver.list_models().await? {
Some(models) => Some(models),
None => match &config.base_url {
Some(base_url) => {
list_openai_compatible_models(base_url, config.api_key.as_deref()).await?
}
None => None,
},
};
let Some(models) = models else {
return Ok(None);
};
Ok(Some(normalize_and_enrich(&config.provider_type, models)))
}
pub fn normalize_and_enrich(
provider_type: &DriverId,
mut models: Vec<DiscoveredModel>,
) -> Vec<DiscoveredProviderModel> {
for model in models.iter_mut() {
if let Some(bare) = model.model_id.strip_prefix("models/") {
model.model_id = bare.to_string();
}
}
models.sort_by(|a, b| {
b.created_at
.cmp(&a.created_at)
.then_with(|| a.model_id.cmp(&b.model_id))
});
enrich_with_profiles(provider_type, models)
}
pub fn enrich_with_profiles(
provider_type: &DriverId,
models: Vec<DiscoveredModel>,
) -> Vec<DiscoveredProviderModel> {
models
.into_iter()
.map(|model| {
let core_profile = get_model_profile(provider_type, &model.model_id);
let api_profile = model.discovered_profile;
let display_name = model
.display_name
.filter(|name| !name.is_empty() && *name != model.model_id)
.or_else(|| core_profile.as_ref().map(|profile| profile.name.clone()));
let description = core_profile
.as_ref()
.and_then(|profile| profile.description.clone())
.or_else(|| {
api_profile
.as_ref()
.and_then(|profile| profile.description.clone())
});
DiscoveredProviderModel {
model_id: model.model_id,
display_name,
description,
}
})
.collect()
}
#[derive(serde::Deserialize)]
struct OpenAiCompatibleModelsResponse {
data: Vec<OpenAiCompatibleModel>,
}
#[derive(serde::Deserialize)]
struct OpenAiCompatibleModel {
id: String,
#[serde(default)]
created: Option<i64>,
#[serde(default)]
owned_by: Option<String>,
}
pub async fn list_openai_compatible_models(
base_url: &str,
api_key: Option<&str>,
) -> Result<Option<Vec<DiscoveredModel>>> {
let url = format!("{}/models", base_url.trim_end_matches('/'));
let mut request = reqwest::Client::new().get(&url);
if let Some(key) = api_key {
request = request.bearer_auth(key);
}
let response = request
.send()
.await
.map_err(|error| AgentLoopError::llm(format!("fetch models from {url}: {error}")))?;
if !response.status().is_success() {
return Err(AgentLoopError::llm(format!(
"models API at {url} returned {}",
response.status()
)));
}
let parsed: OpenAiCompatibleModelsResponse = response.json().await.map_err(|error| {
AgentLoopError::llm(format!("parse models response from {url}: {error}"))
})?;
let models = parsed
.data
.into_iter()
.map(|model| DiscoveredModel {
created_at: model
.created
.and_then(|ts| chrono::DateTime::from_timestamp(ts, 0)),
display_name: None,
owned_by: model.owned_by,
model_id: model.id,
discovered_profile: None,
})
.collect();
Ok(Some(models))
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct RankedDiscoveredModels {
pub models: Vec<DiscoveredProviderModel>,
pub recommended_count: usize,
}
const RECOMMENDED_CAP: usize = 20;
pub fn rank_discovered_models(
provider_type: &DriverId,
models: Vec<DiscoveredProviderModel>,
current_model: Option<&str>,
curated: &[&str],
) -> RankedDiscoveredModels {
if matches!(provider_type, DriverId::OpenRouter) {
rank_aggregator_models(provider_type, models, current_model, curated)
} else {
RankedDiscoveredModels {
recommended_count: 0,
models,
}
}
}
fn rank_aggregator_models(
provider_type: &DriverId,
models: Vec<DiscoveredProviderModel>,
current_model: Option<&str>,
curated: &[&str],
) -> RankedDiscoveredModels {
let mut recommended_ids: Vec<String> = Vec::new();
for suggestion in curated {
let bare = bare_model_id(suggestion);
if models.iter().any(|model| model.model_id == bare) {
push_unique(&mut recommended_ids, bare.to_string());
}
}
if let Some(current) = current_model.map(bare_model_id)
&& models.iter().any(|model| model.model_id == current)
{
push_unique(&mut recommended_ids, current.to_string());
}
let mut profile_candidates: Vec<String> = models
.iter()
.filter(|model| {
!recommended_ids.contains(&model.model_id)
&& is_major_vendor_model(&model.model_id)
&& get_model_profile(provider_type, &model.model_id).is_some()
})
.map(|model| model.model_id.clone())
.collect();
profile_candidates.sort();
for model_id in profile_candidates {
if recommended_ids.len() >= RECOMMENDED_CAP {
break;
}
push_unique(&mut recommended_ids, model_id);
}
let recommended_count = recommended_ids.len();
let mut ranked = Vec::with_capacity(models.len());
for model_id in &recommended_ids {
if let Some(index) = models.iter().position(|model| &model.model_id == model_id) {
ranked.push(models[index].clone());
}
}
let mut rest: Vec<DiscoveredProviderModel> = models
.into_iter()
.filter(|model| !recommended_ids.contains(&model.model_id))
.collect();
rest.sort_by(|a, b| a.model_id.cmp(&b.model_id));
ranked.extend(rest);
RankedDiscoveredModels {
models: ranked,
recommended_count,
}
}
pub fn bare_model_id(spec: &str) -> &str {
spec.split_whitespace().next().unwrap_or(spec)
}
fn push_unique(ids: &mut Vec<String>, id: String) {
if !ids.contains(&id) {
ids.push(id);
}
}
fn is_major_vendor_model(model_id: &str) -> bool {
model_id.starts_with("openai/")
|| model_id.starts_with("anthropic/")
|| model_id.starts_with("google/")
|| model_id.starts_with("nvidia/")
}
#[cfg(test)]
mod tests {
use super::*;
fn bare_discovered(model_id: &str) -> DiscoveredModel {
DiscoveredModel {
model_id: model_id.to_string(),
display_name: None,
created_at: None,
owned_by: None,
discovered_profile: None,
}
}
fn model(id: &str) -> DiscoveredProviderModel {
DiscoveredProviderModel {
model_id: id.to_string(),
display_name: None,
description: None,
}
}
#[tokio::test]
async fn discovery_is_unsupported_for_llmsim() {
let registry = DriverRegistry::new();
let result = discover_provider_models(®istry, &ProviderConfig::new(DriverId::LlmSim))
.await
.expect("llmsim discovery should not error");
assert!(result.is_none());
}
#[test]
fn enrichment_fills_names_and_descriptions_from_profiles() {
let enriched = enrich_with_profiles(&DriverId::OpenAI, vec![bare_discovered("gpt-5.5")]);
assert_eq!(enriched.len(), 1);
assert_eq!(enriched[0].model_id, "gpt-5.5");
assert_eq!(enriched[0].display_name.as_deref(), Some("GPT-5.5"));
assert!(
enriched[0].description.is_some(),
"profile description should be carried over"
);
}
#[test]
fn enrichment_prefers_api_display_name_over_profile() {
let mut discovered = bare_discovered("gpt-5.5");
discovered.display_name = Some("GPT-5.5 (via gateway)".to_string());
let enriched = enrich_with_profiles(&DriverId::OpenAI, vec![discovered]);
assert_eq!(
enriched[0].display_name.as_deref(),
Some("GPT-5.5 (via gateway)")
);
}
#[test]
fn enrichment_keeps_unknown_models_with_bare_ids() {
let enriched = enrich_with_profiles(
&DriverId::OpenAI,
vec![bare_discovered("totally-new-model")],
);
assert_eq!(enriched[0].model_id, "totally-new-model");
assert!(enriched[0].display_name.is_none());
assert!(enriched[0].description.is_none());
}
#[test]
fn normalization_strips_gemini_style_prefixes_and_sorts_newest_first() {
let mut older = bare_discovered("models/qwen3");
older.created_at = chrono::DateTime::from_timestamp(1_600_000_000, 0);
let mut newer = bare_discovered("llama3.2:latest");
newer.created_at = chrono::DateTime::from_timestamp(1_700_000_000, 0);
let normalized = normalize_and_enrich(&DriverId::OpenAI, vec![older, newer]);
let ids: Vec<&str> = normalized.iter().map(|m| m.model_id.as_str()).collect();
assert_eq!(ids, &["llama3.2:latest", "qwen3"]);
}
#[test]
fn aggregator_ranking_puts_curated_and_current_first_then_sorts_rest() {
let ranked = rank_discovered_models(
&DriverId::OpenRouter,
vec![
model("zai/glm-5"),
model("openai/gpt-5.5"),
model("anthropic/claude-opus-4-8"),
model("moon/kimi-k3"),
],
Some("moon/kimi-k3"),
&["openai/gpt-5.5", "anthropic/claude-opus-4-8"],
);
assert_eq!(ranked.recommended_count, 3);
let ids: Vec<&str> = ranked.models.iter().map(|m| m.model_id.as_str()).collect();
assert_eq!(
ids,
&[
"openai/gpt-5.5",
"anthropic/claude-opus-4-8",
"moon/kimi-k3",
"zai/glm-5",
]
);
}
#[test]
fn curated_ids_absent_from_the_catalog_are_not_recommended() {
let ranked = rank_discovered_models(
&DriverId::OpenRouter,
vec![model("zai/glm-5")],
None,
&["openai/gpt-5.5"],
);
assert_eq!(ranked.recommended_count, 0);
assert_eq!(ranked.models.len(), 1);
}
#[test]
fn single_vendor_providers_keep_discovery_order() {
let ranked = rank_discovered_models(
&DriverId::OpenAI,
vec![model("gpt-5.5"), model("gpt-5.2")],
None,
&["gpt-5.2"],
);
assert_eq!(ranked.recommended_count, 0);
let ids: Vec<&str> = ranked.models.iter().map(|m| m.model_id.as_str()).collect();
assert_eq!(ids, &["gpt-5.5", "gpt-5.2"]);
}
#[test]
fn bare_model_id_strips_reasoning_effort_suffix() {
assert_eq!(
bare_model_id("nvidia/nemotron-3-super-120b-a12b high"),
"nvidia/nemotron-3-super-120b-a12b"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn openai_compatible_fallback_lists_models() {
use std::io::{Read, Write};
let listener = std::net::TcpListener::bind("127.0.0.1:0").expect("bind mock server");
let addr = listener.local_addr().expect("mock server addr");
let server = std::thread::spawn(move || {
let (mut stream, _) = listener.accept().expect("accept");
let mut buf = [0u8; 4096];
let _ = stream.read(&mut buf);
let body = r#"{"object":"list","data":[
{"id":"llama3.2:latest","object":"model","created":1700000000,"owned_by":"library"},
{"id":"models/qwen3","object":"model"}
]}"#;
let response = format!(
"HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{body}",
body.len(),
);
stream
.write_all(response.as_bytes())
.expect("write response");
});
let discovered = list_openai_compatible_models(&format!("http://{addr}/v1"), None)
.await
.expect("fallback discovery should succeed")
.expect("endpoint lists models");
server.join().expect("mock server thread");
let presented = normalize_and_enrich(&DriverId::OpenAI, discovered);
let ids: Vec<&str> = presented.iter().map(|m| m.model_id.as_str()).collect();
assert!(ids.contains(&"llama3.2:latest"), "ids: {ids:?}");
assert!(ids.contains(&"qwen3"), "ids: {ids:?}");
}
}