use crate::runtime::ProviderChoice;
use crate::settings::Settings;
use anyhow::{Context, Result, anyhow};
use everruns_core::DriverId;
use everruns_core::driver_registry::{DiscoveredModel, DriverRegistry, ProviderConfig};
use everruns_core::get_model_profile;
use std::collections::HashSet;
#[derive(Clone, Debug, PartialEq, Eq)]
pub(crate) struct DiscoveredProviderModel {
pub model_id: String,
pub display_name: Option<String>,
pub description: Option<String>,
}
pub(crate) async fn discover_provider_models(
choice: &ProviderChoice,
settings: &Settings,
) -> Result<Option<Vec<DiscoveredProviderModel>>> {
if matches!(choice, ProviderChoice::Sim | ProviderChoice::Codex { .. }) {
return Ok(None);
}
let target = choice.model_with_provider(settings)?;
if matches!(target.provider_type, DriverId::Bedrock | DriverId::LlmSim) {
return Ok(None);
}
let mut config = ProviderConfig::new(target.provider_type.clone());
if let Some(key) = &target.api_key {
config = config.with_api_key(key);
}
if let Some(base_url) = &target.base_url {
config = config.with_base_url(base_url);
}
let mut registry = DriverRegistry::new();
everruns_anthropic::register_driver(&mut registry);
everruns_openai::register_driver(&mut registry);
everruns_openrouter::register_driver(&mut registry);
let driver = registry.create_chat_driver(&config)?;
let models = match driver.list_models().await? {
Some(models) => Some(models),
None => match &target.base_url {
Some(base_url) => {
list_openai_compatible_models(base_url, target.api_key.as_deref()).await?
}
None => None,
},
};
let Some(mut models) = models else {
return Ok(None);
};
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))
});
Ok(Some(enrich_with_profiles(&target.provider_type, models)))
}
const DISCOVERY_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(5);
fn bare_model_id_from_spec(spec: &str) -> &str {
spec.split_whitespace().next().unwrap_or(spec)
}
pub(crate) async fn reconcile_provider_with_catalog(
choice: ProviderChoice,
settings: &Settings,
) -> (ProviderChoice, Vec<String>) {
let discovered = tokio::time::timeout(
DISCOVERY_TIMEOUT,
discover_provider_models(&choice, settings),
)
.await;
let Ok(Ok(Some(models))) = discovered else {
return (choice, vec![]);
};
if models.is_empty() {
return (choice, vec![]);
}
let ids: HashSet<String> = models.iter().map(|m| m.model_id.clone()).collect();
let model_id = choice.model_id().to_string();
if ids.contains(&model_id) {
return (choice, vec![]);
}
let fallback = pick_catalog_fallback(&choice, &ids, &models);
let note = format!(
"model \"{model_id}\" not available on {}; using {} instead",
choice.provider_name(),
fallback.label()
);
(fallback, vec![note])
}
fn pick_catalog_fallback(
choice: &ProviderChoice,
ids: &HashSet<String>,
models: &[DiscoveredProviderModel],
) -> ProviderChoice {
let provider = choice.provider_name();
let default =
ProviderChoice::default_for_provider_name(provider).unwrap_or_else(|_| choice.clone());
if ids.contains(default.model_id()) {
return default;
}
for suggestion in ProviderChoice::model_suggestions_for_provider(provider) {
let bare = bare_model_id_from_spec(suggestion);
if ids.contains(bare)
&& let Ok(resolved) = default.resolve_model_spec(suggestion)
{
return resolved;
}
}
if let Some(first) = models.first()
&& let Ok(resolved) = default.resolve_model_spec(&first.model_id)
{
return resolved;
}
default
}
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>,
}
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
.with_context(|| format!("fetch models from {url}"))?;
if !response.status().is_success() {
return Err(anyhow!(
"models API at {url} returned {}",
response.status()
));
}
let parsed: OpenAiCompatibleModelsResponse = response
.json()
.await
.with_context(|| format!("parse models response from {url}"))?;
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))
}
#[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,
}
}
#[tokio::test]
async fn discovery_is_unsupported_for_llmsim() {
let result = discover_provider_models(&ProviderChoice::Sim, &Settings::default())
.await
.expect("llmsim discovery should not error");
assert!(result.is_none());
}
#[test]
fn enrichment_fills_names_and_descriptions_from_core_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
.as_deref()
.is_some_and(|description| description.contains("flagship")),
"core profile description should be carried over: {:?}",
enriched[0].description
);
}
#[test]
fn enrichment_prefers_api_display_name_over_core_profile() {
let mut model = bare_discovered("gpt-5.5");
model.display_name = Some("GPT-5.5 (via gateway)".to_string());
let enriched = enrich_with_profiles(&DriverId::OpenAI, vec![model]);
assert_eq!(
enriched[0].display_name.as_deref(),
Some("GPT-5.5 (via gateway)")
);
}
#[test]
fn pick_catalog_fallback_prefers_provider_default_when_available() {
use super::pick_catalog_fallback;
use std::collections::HashSet;
let choice = ProviderChoice::default_for_provider_name("openai").unwrap();
let ids: HashSet<String> = ["gpt-5.5", "gpt-4o"]
.into_iter()
.map(str::to_string)
.collect();
let models = vec![
DiscoveredProviderModel {
model_id: "gpt-4o".to_string(),
display_name: None,
description: None,
},
DiscoveredProviderModel {
model_id: "gpt-5.5".to_string(),
display_name: None,
description: None,
},
];
let unavailable = ProviderChoice::OpenAi {
model: "gpt-4o-mini".to_string(),
reasoning_effort: Some("medium".to_string()),
};
let fallback = pick_catalog_fallback(&unavailable, &ids, &models);
assert_eq!(fallback.label(), choice.label());
}
#[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());
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn discovery_falls_back_to_openai_compatible_endpoint() {
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 provider = ProviderChoice::Ollama {
model: "llama3.2".to_string(),
base_url: format!("http://{addr}/v1"),
reasoning_effort: None,
};
let models = discover_provider_models(&provider, &Settings::default())
.await
.expect("fallback discovery should succeed")
.expect("openai-compatible endpoint lists models");
server.join().expect("mock server thread");
let ids: Vec<&str> = models.iter().map(|m| m.model_id.as_str()).collect();
assert!(ids.contains(&"llama3.2:latest"), "ids: {ids:?}");
assert!(ids.contains(&"qwen3"), "ids: {ids:?}");
}
#[tokio::test]
#[ignore = "requires OPENROUTER_API_KEY; performs a live models API call"]
async fn discovery_openrouter_live() {
if std::env::var("OPENROUTER_API_KEY")
.map(|v| v.is_empty())
.unwrap_or(true)
{
eprintln!("skipping: OPENROUTER_API_KEY not set");
return;
}
let provider = ProviderChoice::default_for_provider_name("openrouter").unwrap();
let models = discover_provider_models(&provider, &Settings::default())
.await
.expect("openrouter discovery should succeed")
.expect("openrouter supports model listing");
assert!(!models.is_empty(), "openrouter should report models");
assert!(
models.iter().all(|m| !m.model_id.is_empty()),
"every discovered model needs an id"
);
}
}