1use crate::driver_helpers::shared_request_http_client;
15use crate::driver_registry::{DiscoveredModel, DriverId, DriverRegistry, ProviderConfig};
16use crate::error::{AgentLoopError, Result};
17use crate::model_profiles::get_model_profile;
18
19#[derive(Clone, Debug, PartialEq, Eq)]
23pub struct DiscoveredProviderModel {
24 pub model_id: String,
26 pub display_name: Option<String>,
28 pub description: Option<String>,
30}
31
32pub async fn discover_provider_models(
39 registry: &DriverRegistry,
40 config: &ProviderConfig,
41) -> Result<Option<Vec<DiscoveredProviderModel>>> {
42 let driver = registry.create_chat_driver(config)?;
43 let Some(models) = driver
44 .list_models(&crate::runtime_provider::ProviderEndpoint::default())
45 .await?
46 else {
47 return Ok(None);
48 };
49
50 Ok(Some(normalize_and_enrich(&config.provider_type, models)))
51}
52
53pub fn normalize_and_enrich(
59 provider_type: &DriverId,
60 mut models: Vec<DiscoveredModel>,
61) -> Vec<DiscoveredProviderModel> {
62 for model in models.iter_mut() {
63 if let Some(bare) = model.model_id.strip_prefix("models/") {
66 model.model_id = bare.to_string();
67 }
68 }
69 models.sort_by(|a, b| {
70 b.created_at
71 .cmp(&a.created_at)
72 .then_with(|| a.model_id.cmp(&b.model_id))
73 });
74 enrich_with_profiles(provider_type, models)
75}
76
77pub fn enrich_with_profiles(
84 provider_type: &DriverId,
85 models: Vec<DiscoveredModel>,
86) -> Vec<DiscoveredProviderModel> {
87 models
88 .into_iter()
89 .map(|model| {
90 let core_profile = get_model_profile(provider_type, &model.model_id);
91 let api_profile = model.discovered_profile;
92 let display_name = model
93 .display_name
94 .filter(|name| !name.is_empty() && *name != model.model_id)
95 .or_else(|| core_profile.as_ref().map(|profile| profile.name.clone()));
96 let description = core_profile
97 .as_ref()
98 .and_then(|profile| profile.description.clone())
99 .or_else(|| {
100 api_profile
101 .as_ref()
102 .and_then(|profile| profile.description.clone())
103 });
104 DiscoveredProviderModel {
105 model_id: model.model_id,
106 display_name,
107 description,
108 }
109 })
110 .collect()
111}
112
113#[derive(serde::Deserialize)]
114struct OpenAiCompatibleModelsResponse {
115 data: Vec<OpenAiCompatibleModel>,
116}
117
118#[derive(serde::Deserialize)]
119struct OpenAiCompatibleModel {
120 id: String,
121 #[serde(default)]
122 created: Option<i64>,
123 #[serde(default)]
124 owned_by: Option<String>,
125}
126
127pub async fn list_openai_compatible_models(
130 endpoint: &crate::runtime_provider::ProviderEndpoint,
131) -> Result<Option<Vec<DiscoveredModel>>> {
132 let url = endpoint
133 .url("models")
134 .ok_or_else(|| AgentLoopError::config("provider endpoint is not configured"))?;
135 let resolved = endpoint.resolve("GET", url, &[]).await?;
136 list_openai_compatible_models_with_client(&shared_request_http_client(), &resolved).await
140}
141
142async fn list_openai_compatible_models_with_client(
143 client: &reqwest::Client,
144 resolved: &crate::runtime_provider::ResolvedProviderRequest,
145) -> Result<Option<Vec<DiscoveredModel>>> {
146 let mut request = client.get(&resolved.url);
147 for (name, value) in &resolved.headers {
148 request = request.header(name, value);
149 }
150 let response = request
151 .send()
152 .await
153 .map_err(|error| AgentLoopError::llm(format!("fetch models: {error}")))?;
154 if !response.status().is_success() {
155 return Err(AgentLoopError::llm(format!(
156 "models API returned {}",
157 response.status()
158 )));
159 }
160 let parsed: OpenAiCompatibleModelsResponse = response
161 .json()
162 .await
163 .map_err(|error| AgentLoopError::llm(format!("parse models response: {error}")))?;
164 let models = parsed
165 .data
166 .into_iter()
167 .map(|model| DiscoveredModel {
168 capabilities: vec!["chat".to_string()],
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 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#[derive(Clone, Debug, PartialEq, Eq)]
302pub struct ModelSearchMatch {
303 pub provider: String,
305 pub model_id: String,
307 pub display_name: Option<String>,
309}
310
311#[derive(Clone, Debug, Default, PartialEq, Eq)]
317pub struct ModelSearchResult {
318 pub matches: Vec<ModelSearchMatch>,
320 pub providers_searched: Vec<String>,
323 pub provider_errors: Vec<String>,
325}
326
327pub fn match_models(
333 provider: &str,
334 models: &[DiscoveredProviderModel],
335 query: &str,
336) -> Vec<ModelSearchMatch> {
337 let needle = query.trim().to_lowercase();
338 if needle.is_empty() {
339 return Vec::new();
340 }
341 models
342 .iter()
343 .filter(|model| {
344 model.model_id.to_lowercase().contains(&needle)
345 || model
346 .display_name
347 .as_deref()
348 .is_some_and(|name| name.to_lowercase().contains(&needle))
349 })
350 .map(|model| ModelSearchMatch {
351 provider: provider.to_string(),
352 model_id: model.model_id.clone(),
353 display_name: model.display_name.clone(),
354 })
355 .collect()
356}
357
358pub async fn search_provider_models(
370 registry: &DriverRegistry,
371 providers: &[(String, ProviderConfig)],
372 query: &str,
373) -> ModelSearchResult {
374 let mut result = ModelSearchResult::default();
375 if query.trim().is_empty() {
376 return result;
377 }
378
379 for (label, config) in providers {
380 match discover_provider_models(registry, config).await {
381 Ok(Some(models)) => {
382 result.providers_searched.push(label.clone());
383 result.matches.extend(match_models(label, &models, query));
384 }
385 Ok(None) => {}
386 Err(error) => result.provider_errors.push(format!("{label}: {error}")),
387 }
388 }
389
390 result
391 .matches
392 .sort_by(|a, b| (&a.provider, &a.model_id).cmp(&(&b.provider, &b.model_id)));
393 result
394}
395
396#[cfg(test)]
397mod tests {
398 use super::*;
399
400 struct NoCatalog;
401
402 #[async_trait::async_trait]
403 impl crate::ChatDriver for NoCatalog {
404 async fn chat_completion_stream(
405 &self,
406 _endpoint: &crate::ProviderEndpoint,
407 _messages: Vec<crate::LlmMessage>,
408 _config: &crate::LlmCallConfig,
409 ) -> crate::Result<crate::LlmResponseStream> {
410 unreachable!()
411 }
412 }
413
414 fn bare_discovered(model_id: &str) -> DiscoveredModel {
415 DiscoveredModel {
416 capabilities: vec!["chat".to_string()],
417 model_id: model_id.to_string(),
418 display_name: None,
419 created_at: None,
420 owned_by: None,
421 discovered_profile: None,
422 }
423 }
424
425 fn model(id: &str) -> DiscoveredProviderModel {
426 DiscoveredProviderModel {
427 model_id: id.to_string(),
428 display_name: None,
429 description: None,
430 }
431 }
432
433 #[tokio::test]
434 async fn discovery_is_unsupported_for_llmsim() {
435 let mut registry = DriverRegistry::new();
438 registry.register(DriverId::LlmSim, |_| Box::new(NoCatalog));
439 let result = discover_provider_models(®istry, &ProviderConfig::new(DriverId::LlmSim))
440 .await
441 .expect("llmsim discovery should not error");
442 assert!(result.is_none());
443 }
444
445 #[test]
446 fn enrichment_fills_names_and_descriptions_from_profiles() {
447 let enriched = enrich_with_profiles(&DriverId::OpenAI, vec![bare_discovered("gpt-5.5")]);
448
449 assert_eq!(enriched.len(), 1);
450 assert_eq!(enriched[0].model_id, "gpt-5.5");
451 assert_eq!(enriched[0].display_name.as_deref(), Some("GPT-5.5"));
452 assert!(
453 enriched[0].description.is_some(),
454 "profile description should be carried over"
455 );
456 }
457
458 #[test]
459 fn enrichment_prefers_api_display_name_over_profile() {
460 let mut discovered = bare_discovered("gpt-5.5");
461 discovered.display_name = Some("GPT-5.5 (via gateway)".to_string());
462
463 let enriched = enrich_with_profiles(&DriverId::OpenAI, vec![discovered]);
464
465 assert_eq!(
466 enriched[0].display_name.as_deref(),
467 Some("GPT-5.5 (via gateway)")
468 );
469 }
470
471 #[test]
472 fn enrichment_keeps_unknown_models_with_bare_ids() {
473 let enriched = enrich_with_profiles(
474 &DriverId::OpenAI,
475 vec![bare_discovered("totally-new-model")],
476 );
477
478 assert_eq!(enriched[0].model_id, "totally-new-model");
479 assert!(enriched[0].display_name.is_none());
480 assert!(enriched[0].description.is_none());
481 }
482
483 #[test]
484 fn normalization_strips_gemini_style_prefixes_and_sorts_newest_first() {
485 let mut older = bare_discovered("models/qwen3");
486 older.created_at = chrono::DateTime::from_timestamp(1_600_000_000, 0);
487 let mut newer = bare_discovered("llama3.2:latest");
488 newer.created_at = chrono::DateTime::from_timestamp(1_700_000_000, 0);
489
490 let normalized = normalize_and_enrich(&DriverId::OpenAI, vec![older, newer]);
491
492 let ids: Vec<&str> = normalized.iter().map(|m| m.model_id.as_str()).collect();
493 assert_eq!(ids, &["llama3.2:latest", "qwen3"]);
494 }
495
496 #[test]
497 fn aggregator_ranking_puts_curated_and_current_first_then_sorts_rest() {
498 let ranked = rank_discovered_models(
499 &DriverId::OpenRouter,
500 vec![
501 model("zai/glm-5"),
502 model("openai/gpt-5.5"),
503 model("anthropic/claude-opus-4-8"),
504 model("moon/kimi-k3"),
505 ],
506 Some("moon/kimi-k3"),
507 &["openai/gpt-5.5", "anthropic/claude-opus-4-8"],
508 );
509
510 assert_eq!(ranked.recommended_count, 3);
511 let ids: Vec<&str> = ranked.models.iter().map(|m| m.model_id.as_str()).collect();
512 assert_eq!(
513 ids,
514 &[
515 "openai/gpt-5.5",
516 "anthropic/claude-opus-4-8",
517 "moon/kimi-k3",
518 "zai/glm-5",
519 ]
520 );
521 }
522
523 #[test]
524 fn curated_ids_absent_from_the_catalog_are_not_recommended() {
525 let ranked = rank_discovered_models(
526 &DriverId::OpenRouter,
527 vec![model("zai/glm-5")],
528 None,
529 &["openai/gpt-5.5"],
530 );
531
532 assert_eq!(ranked.recommended_count, 0);
533 assert_eq!(ranked.models.len(), 1);
534 }
535
536 #[test]
537 fn single_vendor_providers_keep_discovery_order() {
538 let ranked = rank_discovered_models(
539 &DriverId::OpenAI,
540 vec![model("gpt-5.5"), model("gpt-5.2")],
541 None,
542 &["gpt-5.2"],
543 );
544
545 assert_eq!(ranked.recommended_count, 0);
546 let ids: Vec<&str> = ranked.models.iter().map(|m| m.model_id.as_str()).collect();
547 assert_eq!(ids, &["gpt-5.5", "gpt-5.2"]);
548 }
549
550 #[test]
551 fn bare_model_id_strips_reasoning_effort_suffix() {
552 assert_eq!(
553 bare_model_id("nvidia/nemotron-3-super-120b-a12b high"),
554 "nvidia/nemotron-3-super-120b-a12b"
555 );
556 }
557
558 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
561 async fn openai_compatible_fallback_lists_models() {
562 use std::io::{Read, Write};
563
564 let listener = std::net::TcpListener::bind("127.0.0.1:0").expect("bind mock server");
565 let addr = listener.local_addr().expect("mock server addr");
566 let server = std::thread::spawn(move || {
567 let (mut stream, _) = listener.accept().expect("accept");
568 let mut buf = [0u8; 4096];
569 let _ = stream.read(&mut buf);
570 let body = r#"{"object":"list","data":[
571 {"id":"llama3.2:latest","object":"model","created":1700000000,"owned_by":"library"},
572 {"id":"models/qwen3","object":"model"}
573 ]}"#;
574 let response = format!(
575 "HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{body}",
576 body.len(),
577 );
578 stream
579 .write_all(response.as_bytes())
580 .expect("write response");
581 });
582
583 let discovered = list_openai_compatible_models_with_client(
584 &reqwest::Client::builder()
585 .no_proxy()
586 .build()
587 .expect("build mock HTTP client"),
588 &crate::ResolvedProviderRequest {
589 url: format!("http://{addr}/v1/models"),
590 headers: Vec::new(),
591 },
592 )
593 .await
594 .expect("fallback discovery should succeed")
595 .expect("endpoint lists models");
596 server.join().expect("mock server thread");
597
598 let presented = normalize_and_enrich(&DriverId::OpenAI, discovered);
599 let ids: Vec<&str> = presented.iter().map(|m| m.model_id.as_str()).collect();
600 assert!(ids.contains(&"llama3.2:latest"), "ids: {ids:?}");
601 assert!(ids.contains(&"qwen3"), "ids: {ids:?}");
603 }
604
605 #[tokio::test]
606 async fn openai_compatible_fallback_blocks_internal_addresses() {
607 list_openai_compatible_models_with_client(
608 &shared_request_http_client(),
609 &crate::ResolvedProviderRequest {
610 url: "http://127.0.0.1:9/v1/models".into(),
611 headers: Vec::new(),
612 },
613 )
614 .await
615 .expect_err("model discovery must use the provider SSRF guard");
616 }
617
618 fn presented(id: &str, display_name: Option<&str>) -> DiscoveredProviderModel {
619 DiscoveredProviderModel {
620 model_id: id.to_string(),
621 display_name: display_name.map(str::to_string),
622 description: None,
623 }
624 }
625
626 #[test]
627 fn matching_is_case_insensitive_over_id_and_display_name() {
628 let catalog = vec![
629 presented("openai/gpt-5.5", Some("GPT-5.5")),
630 presented("moon/luna-1", None),
631 presented("acme/nebula", Some("Luna Nebula")),
632 ];
633
634 let by_id = match_models("openrouter", &catalog, "LUNA");
635 let ids: Vec<&str> = by_id.iter().map(|m| m.model_id.as_str()).collect();
636 assert_eq!(
637 ids,
638 vec!["moon/luna-1", "acme/nebula"],
639 "a display-name hit counts as much as an id hit"
640 );
641 assert!(by_id.iter().all(|m| m.provider == "openrouter"));
642 }
643
644 #[test]
645 fn an_empty_query_matches_nothing_rather_than_everything() {
646 let catalog = vec![presented("openai/gpt-5.5", None)];
649 assert!(match_models("openrouter", &catalog, " ").is_empty());
650 }
651
652 #[tokio::test]
653 async fn search_skips_catalog_less_providers_and_records_failures() {
654 let mut registry = DriverRegistry::new();
655 registry.register(DriverId::LlmSim, |_| Box::new(NoCatalog));
656 let providers = vec![
657 ("sim".to_string(), ProviderConfig::new(DriverId::LlmSim)),
659 (
662 "broken".to_string(),
663 ProviderConfig::new(DriverId::OpenAI).with_api_key("k"),
664 ),
665 ];
666
667 let result = search_provider_models(®istry, &providers, "gpt").await;
668
669 assert!(result.matches.is_empty());
670 assert!(
671 !result.providers_searched.contains(&"sim".to_string()),
672 "a provider with no catalog was not searched"
673 );
674 assert_eq!(result.provider_errors.len(), 1);
675 assert!(result.provider_errors[0].starts_with("broken: "));
676 }
677
678 #[tokio::test]
679 async fn an_empty_query_short_circuits_before_any_provider_call() {
680 let registry = DriverRegistry::new();
681 let providers = vec![(
682 "broken".to_string(),
683 ProviderConfig::new(DriverId::OpenAI).with_api_key("k"),
684 )];
685
686 let result = search_provider_models(®istry, &providers, "").await;
687
688 assert_eq!(result, ModelSearchResult::default());
689 }
690}