Skip to main content

everruns_provider/
model_discovery.rs

1//! Provider model discovery and display ranking.
2//!
3//! Ported from yolop, where every host that offers a model picker had to
4//! reimplement the same three steps: ask the driver for a catalog, fall back to
5//! the OpenAI-compatible `GET <base>/models` for endpoints the drivers decline,
6//! and merge the answer with [`crate::model_profiles`] so bare ids still render
7//! human-readable names. None of that is host-specific, so it lives beside the
8//! driver registry and the profile registry it depends on.
9//!
10//! Discovery is deliberately three-valued: `Ok(None)` means "this provider has
11//! no catalog to offer" and callers should keep their curated suggestions,
12//! while `Err` means the catalog request itself failed.
13
14use crate::driver_registry::{DiscoveredModel, DriverId, DriverRegistry, ProviderConfig};
15use crate::error::{AgentLoopError, Result};
16use crate::model_profiles::get_model_profile;
17
18/// One model offered by a provider, ready for display: the bare id plus
19/// human-readable metadata merged from the provider's API response and the
20/// model profile registry.
21#[derive(Clone, Debug, PartialEq, Eq)]
22pub struct DiscoveredProviderModel {
23    /// Bare model id, as chat calls and profile lookups expect it.
24    pub model_id: String,
25    /// Human-readable name, when the provider or a profile supplies one.
26    pub display_name: Option<String>,
27    /// Short description, when a profile or the provider supplies one.
28    pub description: Option<String>,
29}
30
31/// Query a provider's models API through its driver.
32///
33/// Returns `Ok(None)` when the provider (or its custom endpoint) does not
34/// support model listing; callers should fall back to curated suggestions in
35/// that case rather than treating it as an error.
36///
37/// Drivers that decline listing for an unrecognized OpenAI-compatible endpoint
38/// (Ollama, Gemini's OpenAI surface, proxies) are retried against
39/// [`list_openai_compatible_models`] when the config carries a base URL.
40pub async fn discover_provider_models(
41    registry: &DriverRegistry,
42    config: &ProviderConfig,
43) -> Result<Option<Vec<DiscoveredProviderModel>>> {
44    // Neither of these has a catalog: the simulator has no API, and Bedrock
45    // model access is an account-level IAM concern rather than a listable one.
46    if matches!(config.provider_type, DriverId::LlmSim | DriverId::Bedrock) {
47        return Ok(None);
48    }
49
50    let driver = registry.create_chat_driver(config)?;
51    let models = match driver.list_models().await? {
52        Some(models) => Some(models),
53        None => match &config.base_url {
54            Some(base_url) => {
55                list_openai_compatible_models(base_url, config.api_key.as_deref()).await?
56            }
57            None => None,
58        },
59    };
60    let Some(models) = models else {
61        return Ok(None);
62    };
63
64    Ok(Some(normalize_and_enrich(&config.provider_type, models)))
65}
66
67/// Normalize discovered ids, sort newest-first, and merge in profile metadata.
68///
69/// Split out from [`discover_provider_models`] so a host that obtained a
70/// catalog some other way (a cached response, a proxy's own endpoint) gets the
71/// same presentation.
72pub fn normalize_and_enrich(
73    provider_type: &DriverId,
74    mut models: Vec<DiscoveredModel>,
75) -> Vec<DiscoveredProviderModel> {
76    for model in models.iter_mut() {
77        // Gemini's OpenAI-compatible surface reports ids as `models/<id>`; the
78        // bare id is what chat calls and profile lookups expect.
79        if let Some(bare) = model.model_id.strip_prefix("models/") {
80            model.model_id = bare.to_string();
81        }
82    }
83    models.sort_by(|a, b| {
84        b.created_at
85            .cmp(&a.created_at)
86            .then_with(|| a.model_id.cmp(&b.model_id))
87    });
88    enrich_with_profiles(provider_type, models)
89}
90
91/// Merge each discovered model with metadata from the model profile registry.
92///
93/// The curated profile wins for descriptions (short, written for display); the
94/// provider's API response wins for display names, since it knows its own
95/// catalog best (e.g. OpenRouter's `name` field), with the profile filling the
96/// gap for APIs that return bare ids (e.g. OpenAI).
97pub fn enrich_with_profiles(
98    provider_type: &DriverId,
99    models: Vec<DiscoveredModel>,
100) -> Vec<DiscoveredProviderModel> {
101    models
102        .into_iter()
103        .map(|model| {
104            let core_profile = get_model_profile(provider_type, &model.model_id);
105            let api_profile = model.discovered_profile;
106            let display_name = model
107                .display_name
108                .filter(|name| !name.is_empty() && *name != model.model_id)
109                .or_else(|| core_profile.as_ref().map(|profile| profile.name.clone()));
110            let description = core_profile
111                .as_ref()
112                .and_then(|profile| profile.description.clone())
113                .or_else(|| {
114                    api_profile
115                        .as_ref()
116                        .and_then(|profile| profile.description.clone())
117                });
118            DiscoveredProviderModel {
119                model_id: model.model_id,
120                display_name,
121                description,
122            }
123        })
124        .collect()
125}
126
127#[derive(serde::Deserialize)]
128struct OpenAiCompatibleModelsResponse {
129    data: Vec<OpenAiCompatibleModel>,
130}
131
132#[derive(serde::Deserialize)]
133struct OpenAiCompatibleModel {
134    id: String,
135    #[serde(default)]
136    created: Option<i64>,
137    #[serde(default)]
138    owned_by: Option<String>,
139}
140
141/// Discovery fallback for OpenAI-compatible endpoints no driver recognizes:
142/// `GET <base>/models` with bearer auth.
143pub async fn list_openai_compatible_models(
144    base_url: &str,
145    api_key: Option<&str>,
146) -> Result<Option<Vec<DiscoveredModel>>> {
147    let url = format!("{}/models", base_url.trim_end_matches('/'));
148    let mut request = reqwest::Client::new().get(&url);
149    if let Some(key) = api_key {
150        request = request.bearer_auth(key);
151    }
152    let response = request
153        .send()
154        .await
155        .map_err(|error| AgentLoopError::llm(format!("fetch models from {url}: {error}")))?;
156    if !response.status().is_success() {
157        return Err(AgentLoopError::llm(format!(
158            "models API at {url} returned {}",
159            response.status()
160        )));
161    }
162    let parsed: OpenAiCompatibleModelsResponse = response.json().await.map_err(|error| {
163        AgentLoopError::llm(format!("parse models response from {url}: {error}"))
164    })?;
165    let models = parsed
166        .data
167        .into_iter()
168        .map(|model| DiscoveredModel {
169            created_at: model
170                .created
171                .and_then(|ts| chrono::DateTime::from_timestamp(ts, 0)),
172            display_name: None,
173            owned_by: model.owned_by,
174            model_id: model.id,
175            discovered_profile: None,
176        })
177        .collect();
178    Ok(Some(models))
179}
180
181/// Models reordered for display, plus how many leading entries belong in the
182/// recommended section.
183#[derive(Clone, Debug, PartialEq, Eq)]
184pub struct RankedDiscoveredModels {
185    /// Recommended models first, then the rest of the catalog.
186    pub models: Vec<DiscoveredProviderModel>,
187    /// How many leading entries of `models` are recommendations.
188    pub recommended_count: usize,
189}
190
191/// Cap on the recommended block, so it stays a shortlist rather than a second
192/// full catalog.
193const RECOMMENDED_CAP: usize = 20;
194
195/// Reorder discovered models for a picker.
196///
197/// Aggregator catalogs (OpenRouter lists several hundred models) get a short
198/// recommended block — `curated` ids that are actually offered, then the active
199/// model, then profile-known flagships from major vendors — followed by the rest
200/// sorted by id. Single-vendor providers already return a useful order
201/// (newest-first from discovery) and are left alone.
202///
203/// `curated` entries may carry a trailing reasoning-effort suffix
204/// (`"vendor/model high"`); only the leading token is matched.
205pub fn rank_discovered_models(
206    provider_type: &DriverId,
207    models: Vec<DiscoveredProviderModel>,
208    current_model: Option<&str>,
209    curated: &[&str],
210) -> RankedDiscoveredModels {
211    if matches!(provider_type, DriverId::OpenRouter) {
212        rank_aggregator_models(provider_type, models, current_model, curated)
213    } else {
214        RankedDiscoveredModels {
215            recommended_count: 0,
216            models,
217        }
218    }
219}
220
221fn rank_aggregator_models(
222    provider_type: &DriverId,
223    models: Vec<DiscoveredProviderModel>,
224    current_model: Option<&str>,
225    curated: &[&str],
226) -> RankedDiscoveredModels {
227    let mut recommended_ids: Vec<String> = Vec::new();
228
229    for suggestion in curated {
230        let bare = bare_model_id(suggestion);
231        if models.iter().any(|model| model.model_id == bare) {
232            push_unique(&mut recommended_ids, bare.to_string());
233        }
234    }
235
236    if let Some(current) = current_model.map(bare_model_id)
237        && models.iter().any(|model| model.model_id == current)
238    {
239        push_unique(&mut recommended_ids, current.to_string());
240    }
241
242    let mut profile_candidates: Vec<String> = models
243        .iter()
244        .filter(|model| {
245            !recommended_ids.contains(&model.model_id)
246                && is_major_vendor_model(&model.model_id)
247                && get_model_profile(provider_type, &model.model_id).is_some()
248        })
249        .map(|model| model.model_id.clone())
250        .collect();
251    profile_candidates.sort();
252    for model_id in profile_candidates {
253        if recommended_ids.len() >= RECOMMENDED_CAP {
254            break;
255        }
256        push_unique(&mut recommended_ids, model_id);
257    }
258
259    let recommended_count = recommended_ids.len();
260    let mut ranked = Vec::with_capacity(models.len());
261    for model_id in &recommended_ids {
262        if let Some(index) = models.iter().position(|model| &model.model_id == model_id) {
263            ranked.push(models[index].clone());
264        }
265    }
266
267    let mut rest: Vec<DiscoveredProviderModel> = models
268        .into_iter()
269        .filter(|model| !recommended_ids.contains(&model.model_id))
270        .collect();
271    rest.sort_by(|a, b| a.model_id.cmp(&b.model_id));
272    ranked.extend(rest);
273
274    RankedDiscoveredModels {
275        models: ranked,
276        recommended_count,
277    }
278}
279
280/// Strip a trailing reasoning-effort suffix from a model spec
281/// (`"nvidia/nemotron-3 high"` → `"nvidia/nemotron-3"`).
282pub fn bare_model_id(spec: &str) -> &str {
283    spec.split_whitespace().next().unwrap_or(spec)
284}
285
286fn push_unique(ids: &mut Vec<String>, id: String) {
287    if !ids.contains(&id) {
288        ids.push(id);
289    }
290}
291
292fn is_major_vendor_model(model_id: &str) -> bool {
293    model_id.starts_with("openai/")
294        || model_id.starts_with("anthropic/")
295        || model_id.starts_with("google/")
296        || model_id.starts_with("nvidia/")
297}
298
299#[cfg(test)]
300mod tests {
301    use super::*;
302
303    fn bare_discovered(model_id: &str) -> DiscoveredModel {
304        DiscoveredModel {
305            model_id: model_id.to_string(),
306            display_name: None,
307            created_at: None,
308            owned_by: None,
309            discovered_profile: None,
310        }
311    }
312
313    fn model(id: &str) -> DiscoveredProviderModel {
314        DiscoveredProviderModel {
315            model_id: id.to_string(),
316            display_name: None,
317            description: None,
318        }
319    }
320
321    #[tokio::test]
322    async fn discovery_is_unsupported_for_llmsim() {
323        // The offline simulator has no models API; discovery must signal
324        // "unsupported" rather than erroring, so callers keep curated lists.
325        let registry = DriverRegistry::new();
326        let result = discover_provider_models(&registry, &ProviderConfig::new(DriverId::LlmSim))
327            .await
328            .expect("llmsim discovery should not error");
329        assert!(result.is_none());
330    }
331
332    #[test]
333    fn enrichment_fills_names_and_descriptions_from_profiles() {
334        let enriched = enrich_with_profiles(&DriverId::OpenAI, vec![bare_discovered("gpt-5.5")]);
335
336        assert_eq!(enriched.len(), 1);
337        assert_eq!(enriched[0].model_id, "gpt-5.5");
338        assert_eq!(enriched[0].display_name.as_deref(), Some("GPT-5.5"));
339        assert!(
340            enriched[0].description.is_some(),
341            "profile description should be carried over"
342        );
343    }
344
345    #[test]
346    fn enrichment_prefers_api_display_name_over_profile() {
347        let mut discovered = bare_discovered("gpt-5.5");
348        discovered.display_name = Some("GPT-5.5 (via gateway)".to_string());
349
350        let enriched = enrich_with_profiles(&DriverId::OpenAI, vec![discovered]);
351
352        assert_eq!(
353            enriched[0].display_name.as_deref(),
354            Some("GPT-5.5 (via gateway)")
355        );
356    }
357
358    #[test]
359    fn enrichment_keeps_unknown_models_with_bare_ids() {
360        let enriched = enrich_with_profiles(
361            &DriverId::OpenAI,
362            vec![bare_discovered("totally-new-model")],
363        );
364
365        assert_eq!(enriched[0].model_id, "totally-new-model");
366        assert!(enriched[0].display_name.is_none());
367        assert!(enriched[0].description.is_none());
368    }
369
370    #[test]
371    fn normalization_strips_gemini_style_prefixes_and_sorts_newest_first() {
372        let mut older = bare_discovered("models/qwen3");
373        older.created_at = chrono::DateTime::from_timestamp(1_600_000_000, 0);
374        let mut newer = bare_discovered("llama3.2:latest");
375        newer.created_at = chrono::DateTime::from_timestamp(1_700_000_000, 0);
376
377        let normalized = normalize_and_enrich(&DriverId::OpenAI, vec![older, newer]);
378
379        let ids: Vec<&str> = normalized.iter().map(|m| m.model_id.as_str()).collect();
380        assert_eq!(ids, &["llama3.2:latest", "qwen3"]);
381    }
382
383    #[test]
384    fn aggregator_ranking_puts_curated_and_current_first_then_sorts_rest() {
385        let ranked = rank_discovered_models(
386            &DriverId::OpenRouter,
387            vec![
388                model("zai/glm-5"),
389                model("openai/gpt-5.5"),
390                model("anthropic/claude-opus-4-8"),
391                model("moon/kimi-k3"),
392            ],
393            Some("moon/kimi-k3"),
394            &["openai/gpt-5.5", "anthropic/claude-opus-4-8"],
395        );
396
397        assert_eq!(ranked.recommended_count, 3);
398        let ids: Vec<&str> = ranked.models.iter().map(|m| m.model_id.as_str()).collect();
399        assert_eq!(
400            ids,
401            &[
402                "openai/gpt-5.5",
403                "anthropic/claude-opus-4-8",
404                "moon/kimi-k3",
405                "zai/glm-5",
406            ]
407        );
408    }
409
410    #[test]
411    fn curated_ids_absent_from_the_catalog_are_not_recommended() {
412        let ranked = rank_discovered_models(
413            &DriverId::OpenRouter,
414            vec![model("zai/glm-5")],
415            None,
416            &["openai/gpt-5.5"],
417        );
418
419        assert_eq!(ranked.recommended_count, 0);
420        assert_eq!(ranked.models.len(), 1);
421    }
422
423    #[test]
424    fn single_vendor_providers_keep_discovery_order() {
425        let ranked = rank_discovered_models(
426            &DriverId::OpenAI,
427            vec![model("gpt-5.5"), model("gpt-5.2")],
428            None,
429            &["gpt-5.2"],
430        );
431
432        assert_eq!(ranked.recommended_count, 0);
433        let ids: Vec<&str> = ranked.models.iter().map(|m| m.model_id.as_str()).collect();
434        assert_eq!(ids, &["gpt-5.5", "gpt-5.2"]);
435    }
436
437    #[test]
438    fn bare_model_id_strips_reasoning_effort_suffix() {
439        assert_eq!(
440            bare_model_id("nvidia/nemotron-3-super-120b-a12b high"),
441            "nvidia/nemotron-3-super-120b-a12b"
442        );
443    }
444
445    /// Drivers decline listing for unrecognized custom endpoints (here a
446    /// localhost "Ollama"); discovery must then query the OpenAI-compatible
447    /// `GET <base>/models` itself.
448    #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
449    async fn openai_compatible_fallback_lists_models() {
450        use std::io::{Read, Write};
451
452        let listener = std::net::TcpListener::bind("127.0.0.1:0").expect("bind mock server");
453        let addr = listener.local_addr().expect("mock server addr");
454        let server = std::thread::spawn(move || {
455            let (mut stream, _) = listener.accept().expect("accept");
456            let mut buf = [0u8; 4096];
457            let _ = stream.read(&mut buf);
458            let body = r#"{"object":"list","data":[
459                {"id":"llama3.2:latest","object":"model","created":1700000000,"owned_by":"library"},
460                {"id":"models/qwen3","object":"model"}
461            ]}"#;
462            let response = format!(
463                "HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{body}",
464                body.len(),
465            );
466            stream
467                .write_all(response.as_bytes())
468                .expect("write response");
469        });
470
471        let discovered = list_openai_compatible_models(&format!("http://{addr}/v1"), None)
472            .await
473            .expect("fallback discovery should succeed")
474            .expect("endpoint lists models");
475        server.join().expect("mock server thread");
476
477        let presented = normalize_and_enrich(&DriverId::OpenAI, discovered);
478        let ids: Vec<&str> = presented.iter().map(|m| m.model_id.as_str()).collect();
479        assert!(ids.contains(&"llama3.2:latest"), "ids: {ids:?}");
480        // Gemini-style `models/` prefixes are normalized to bare ids.
481        assert!(ids.contains(&"qwen3"), "ids: {ids:?}");
482    }
483}