1use crate::driver_registry::{DiscoveredModel, DriverId, DriverRegistry, ProviderConfig};
15use crate::error::{AgentLoopError, Result};
16use crate::model_profiles::get_model_profile;
17
18#[derive(Clone, Debug, PartialEq, Eq)]
22pub struct DiscoveredProviderModel {
23 pub model_id: String,
25 pub display_name: Option<String>,
27 pub description: Option<String>,
29}
30
31pub async fn discover_provider_models(
41 registry: &DriverRegistry,
42 config: &ProviderConfig,
43) -> Result<Option<Vec<DiscoveredProviderModel>>> {
44 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
67pub fn normalize_and_enrich(
73 provider_type: &DriverId,
74 mut models: Vec<DiscoveredModel>,
75) -> Vec<DiscoveredProviderModel> {
76 for model in models.iter_mut() {
77 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
91pub 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
141pub 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#[derive(Clone, Debug, PartialEq, Eq)]
184pub struct RankedDiscoveredModels {
185 pub models: Vec<DiscoveredProviderModel>,
187 pub recommended_count: usize,
189}
190
191const RECOMMENDED_CAP: usize = 20;
194
195pub 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
280pub 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 let registry = DriverRegistry::new();
326 let result = discover_provider_models(®istry, &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 #[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 assert!(ids.contains(&"qwen3"), "ids: {ids:?}");
482 }
483}