1use anyhow::Context;
10use chrono::Utc;
11use hashbrown::HashSet;
12use std::path::{Path, PathBuf};
13use std::sync::Arc;
14use std::sync::atomic::{AtomicU64, Ordering};
15use std::time::Duration;
16use tokio::sync::{Mutex, RwLock};
17use tracing::{debug, error, info};
18
19use super::cache::{self, ModelsCache};
20use super::model_family::{ModelFamily, find_family_for_model};
21use super::model_presets::{
22 ModelInfo, ModelPreset, ReasoningEffortPreset, builtin_model_presets, presets_for_provider,
23};
24use crate::config::models::Provider;
25use crate::llm::providers::{
26 MergeCatalogAvailability, MergeCatalogFilters, MergeCatalogModel, MergeGatewayCatalogClient,
27 llamacpp::fetch_llamacpp_models,
28};
29use vtcode_commons::VtCodePaths;
30use vtcode_config::constants::{env_vars, urls};
31
32const LEGACY_MODEL_CACHE_FILE: &str = "models_cache.json";
34
35const DEFAULT_MODEL_CACHE_TTL: Duration = Duration::from_secs(120);
37
38const GEMINI_DEFAULT_MODEL: &str = "gemini-3-flash-preview";
40
41const OPENAI_DEFAULT_MODEL: &str = "gpt-5.6-sol";
43
44const ANTHROPIC_DEFAULT_MODEL: &str = "claude-sonnet-5";
46
47#[derive(Debug)]
49pub struct ModelsManager {
50 local_models: Vec<ModelPreset>,
52 remote_models: RwLock<Vec<ModelInfo>>,
54 etag: RwLock<Option<String>>,
56 vtcode_home: PathBuf,
58 cache_ttl: Duration,
60 current_provider: RwLock<Provider>,
62 refresh_generation: AtomicU64,
64 remote_models_enabled: bool,
66 cache_write_lock: Mutex<()>,
68}
69
70impl Default for ModelsManager {
71 fn default() -> Self {
72 Self::new()
73 }
74}
75
76impl ModelsManager {
77 pub fn new() -> Self {
79 let vtcode_home = Self::default_vtcode_home();
80 Self {
81 local_models: builtin_model_presets(),
82 remote_models: RwLock::new(Vec::new()),
83 etag: RwLock::new(None),
84 vtcode_home,
85 cache_ttl: DEFAULT_MODEL_CACHE_TTL,
86 current_provider: RwLock::new(Provider::default()),
87 refresh_generation: AtomicU64::new(0),
88 remote_models_enabled: true,
89 cache_write_lock: Mutex::new(()),
90 }
91 }
92
93 pub fn with_home(vtcode_home: PathBuf) -> Self {
95 Self {
96 local_models: builtin_model_presets(),
97 remote_models: RwLock::new(Vec::new()),
98 etag: RwLock::new(None),
99 vtcode_home,
100 cache_ttl: DEFAULT_MODEL_CACHE_TTL,
101 current_provider: RwLock::new(Provider::default()),
102 refresh_generation: AtomicU64::new(0),
103 remote_models_enabled: true,
104 cache_write_lock: Mutex::new(()),
105 }
106 }
107
108 pub fn with_provider(provider: Provider) -> Self {
110 let vtcode_home = Self::default_vtcode_home();
111 Self {
112 local_models: presets_for_provider(provider),
113 remote_models: RwLock::new(Vec::new()),
114 etag: RwLock::new(None),
115 vtcode_home,
116 cache_ttl: DEFAULT_MODEL_CACHE_TTL,
117 current_provider: RwLock::new(provider),
118 refresh_generation: AtomicU64::new(0),
119 remote_models_enabled: true,
120 cache_write_lock: Mutex::new(()),
121 }
122 }
123
124 pub fn with_home_and_provider(vtcode_home: PathBuf, provider: Provider) -> Self {
126 Self {
127 local_models: presets_for_provider(provider),
128 remote_models: RwLock::new(Vec::new()),
129 etag: RwLock::new(None),
130 vtcode_home,
131 cache_ttl: DEFAULT_MODEL_CACHE_TTL,
132 current_provider: RwLock::new(provider),
133 refresh_generation: AtomicU64::new(0),
134 remote_models_enabled: true,
135 cache_write_lock: Mutex::new(()),
136 }
137 }
138
139 pub fn set_remote_models_enabled(&mut self, enabled: bool) {
141 self.remote_models_enabled = enabled;
142 }
143
144 pub fn set_cache_ttl(&mut self, ttl: Duration) {
146 self.cache_ttl = ttl;
147 }
148
149 fn default_vtcode_home() -> PathBuf {
151 VtCodePaths::resolve()
152 .map(|paths| paths.cache_dir().to_path_buf())
153 .unwrap_or_else(|_| PathBuf::from(".cache/vtcode"))
154 }
155
156 pub async fn refresh_available_models(&self) -> anyhow::Result<()> {
158 if !self.remote_models_enabled {
159 debug!("Remote model fetching is disabled");
160 return Ok(());
161 }
162
163 let provider = *self.current_provider.read().await;
164 let generation = self.refresh_generation.fetch_add(1, Ordering::AcqRel) + 1;
165
166 if self.try_load_cache_for_at(provider, generation).await {
168 debug!("Using cached models");
169 return Ok(());
170 }
171
172 match provider {
173 Provider::Ollama => {
174 debug!("Fetching remote models for Ollama...");
175 match self.fetch_ollama_models().await {
176 Ok(models) => {
177 info!("Fetched {} models from Ollama", models.len());
178 if self
179 .apply_remote_state_for_provider_at(provider, generation, models.clone(), None)
180 .await
181 {
182 self.persist_cache_for_at(provider, generation, &models, None).await;
183 }
184 Ok(())
185 }
186 Err(e) => {
187 error!("Failed to fetch Ollama models: {e}");
188 Ok(())
190 }
191 }
192 }
193 Provider::LlamaCpp => {
194 debug!("Fetching remote models for llama.cpp...");
195 match self.fetch_llamacpp_models().await {
196 Ok(models) => {
197 info!("Fetched {} models from llama.cpp", models.len());
198 if self
199 .apply_remote_state_for_provider_at(provider, generation, models.clone(), None)
200 .await
201 {
202 self.persist_cache_for_at(provider, generation, &models, None).await;
203 }
204 Ok(())
205 }
206 Err(e) => {
207 error!("Failed to fetch llama.cpp models: {e}");
208 Ok(())
209 }
210 }
211 }
212 Provider::MergeGateway => self.refresh_merge_gateway_models(generation).await,
213 _ => {
214 info!("Remote model discovery for {:?} not implemented, using local presets", provider);
216 Ok(())
217 }
218 }
219 }
220
221 async fn refresh_merge_gateway_models(&self, generation: u64) -> anyhow::Result<()> {
222 let api_key = std::env::var(env_vars::MERGE_GATEWAY_API_KEY)
223 .ok()
224 .filter(|value| !value.trim().is_empty());
225 let Some(api_key) = api_key else {
226 debug!("Merge Gateway API key is unavailable; using cached or static models");
227 let _ = self.try_load_stale_cache_for_at(Provider::MergeGateway, generation).await;
228 return Ok(());
229 };
230
231 let base_url = std::env::var(env_vars::MERGE_GATEWAY_BASE_URL)
232 .ok()
233 .filter(|value| !value.trim().is_empty())
234 .unwrap_or_else(|| urls::MERGE_GATEWAY_NATIVE_API_BASE.to_string());
235 let client = MergeGatewayCatalogClient::try_with_timeouts(api_key, base_url, None)
236 .context("failed to initialize Merge Gateway catalog client")?;
237 let cached = self.load_cache_for(Provider::MergeGateway).await;
238 let etag = self
239 .etag
240 .read()
241 .await
242 .clone()
243 .or_else(|| cached.as_ref().and_then(|cache| cache.etag.clone()));
244
245 match client.fetch_snapshot(&MergeCatalogFilters::default(), etag.as_deref()).await {
246 Ok(None) => {
247 debug!("Merge Gateway catalog is unchanged");
248 if let Some(cache) = cached {
249 let _ = self
250 .apply_remote_state_for_provider_at(
251 Provider::MergeGateway,
252 generation,
253 cache.models,
254 cache.etag,
255 )
256 .await;
257 }
258 Ok(())
259 }
260 Ok(Some(snapshot)) => {
261 let models = snapshot
262 .models
263 .into_iter()
264 .filter_map(Self::merge_catalog_model_info)
265 .collect::<Vec<_>>();
266 if self
267 .apply_remote_state_for_provider_at(
268 Provider::MergeGateway,
269 generation,
270 models.clone(),
271 snapshot.etag.clone(),
272 )
273 .await
274 {
275 self.persist_cache_for_at(Provider::MergeGateway, generation, &models, snapshot.etag)
276 .await;
277 }
278 Ok(())
279 }
280 Err(error) => {
281 error!("Failed to fetch Merge Gateway model catalog: {error}");
282 let _ = self.try_load_stale_cache_for_at(Provider::MergeGateway, generation).await;
283 Ok(())
284 }
285 }
286 }
287
288 fn merge_catalog_model_info(model: MergeCatalogModel) -> Option<ModelInfo> {
289 if model.availability != MergeCatalogAvailability::Available {
290 return None;
291 }
292
293 let display_name = model.display_name.unwrap_or_else(|| model.model.clone());
294 let supported_reasoning_levels = Self::merge_supported_reasoning_presets(model.supports_reasoning);
295 Some(ModelInfo {
296 slug: model.model.clone(),
297 display_name,
298 description: format!("Merge Gateway model: {}", model.model),
299 provider: Provider::MergeGateway,
300 default_reasoning_level: crate::config::types::ReasoningEffortLevel::Medium,
301 supported_reasoning_levels,
302 context_window: model.context_window.map(i64::from),
303 supports_tool_use: model.supports_tool_use,
304 supports_streaming: model.supports_streaming,
305 supports_vision: model.supports_vision,
306 supports_structured_output: model.supports_structured_output,
307 supports_reasoning: model.supports_reasoning,
308 max_output_tokens: model.max_output_tokens.map(i64::from),
309 priority: 50,
310 visibility: "list".to_string(),
311 supported_in_api: true,
312 upgrade: None,
313 })
314 }
315
316 fn merge_supported_reasoning_presets(supports_reasoning: bool) -> Vec<ReasoningEffortPreset> {
317 if !supports_reasoning {
318 return Vec::new();
319 }
320
321 use crate::config::types::ReasoningEffortLevel;
322
323 vec![
324 ReasoningEffortPreset {
325 effort: ReasoningEffortLevel::Minimal,
326 description: "Minimal reasoning depth".to_string(),
327 },
328 ReasoningEffortPreset {
329 effort: ReasoningEffortLevel::Low,
330 description: "Fast responses with lightweight reasoning".to_string(),
331 },
332 ReasoningEffortPreset {
333 effort: ReasoningEffortLevel::Medium,
334 description: "Balanced depth and speed".to_string(),
335 },
336 ReasoningEffortPreset {
337 effort: ReasoningEffortLevel::High,
338 description: "Deep reasoning for complex problems".to_string(),
339 },
340 ReasoningEffortPreset {
341 effort: ReasoningEffortLevel::XHigh,
342 description: "Extra reasoning for the hardest long-running tasks".to_string(),
343 },
344 ReasoningEffortPreset {
345 effort: ReasoningEffortLevel::Max,
346 description: "Maximum reasoning depth".to_string(),
347 },
348 ]
349 }
350
351 async fn fetch_ollama_models(&self) -> anyhow::Result<Vec<ModelInfo>> {
353 let client = reqwest::Client::new();
354 let resp = client.get("http://localhost:11434/api/tags").send().await?;
355
356 if !resp.status().is_success() {
357 return Err(anyhow::anyhow!("Ollama API returned {}", resp.status()));
358 }
359
360 let json: serde_json::Value = resp.json().await?;
361 let mut models = Vec::new();
362
363 if let Some(ollama_models) = json.get("models").and_then(|m| m.as_array()) {
364 for m in ollama_models {
365 if let Some(name) = m.get("name").and_then(|s| s.as_str()) {
366 models.push(ModelInfo {
367 slug: name.to_string(),
368 display_name: format!("{name} (Ollama)"),
369 description: format!("Ollama model: {name}"),
370 provider: Provider::Ollama,
371 default_reasoning_level: crate::config::types::ReasoningEffortLevel::Medium,
372 supported_reasoning_levels: vec![],
373 context_window: Some(32_000), supports_tool_use: true,
375 supports_streaming: true,
376 supports_vision: false,
377 supports_structured_output: false,
378 supports_reasoning: false,
379 max_output_tokens: None,
380 priority: 100,
381 visibility: "list".to_string(),
382 supported_in_api: true,
383 upgrade: None,
384 });
385 }
386 }
387 }
388
389 Ok(models)
390 }
391
392 async fn fetch_llamacpp_models(&self) -> anyhow::Result<Vec<ModelInfo>> {
393 let mut models = Vec::new();
394 for model in fetch_llamacpp_models(None).await? {
395 models.push(ModelInfo {
396 slug: model.clone(),
397 display_name: format!("{model} (llama.cpp)"),
398 description: format!("llama.cpp model: {model}"),
399 provider: Provider::LlamaCpp,
400 default_reasoning_level: crate::config::types::ReasoningEffortLevel::Medium,
401 supported_reasoning_levels: vec![],
402 context_window: Some(131_072),
403 supports_tool_use: true,
404 supports_streaming: true,
405 supports_vision: false,
406 supports_structured_output: false,
407 supports_reasoning: true,
408 max_output_tokens: None,
409 priority: 100,
410 visibility: "list".to_string(),
411 supported_in_api: true,
412 upgrade: None,
413 });
414 }
415
416 Ok(models)
417 }
418
419 pub async fn list_models(&self) -> Vec<ModelPreset> {
421 if let Err(err) = self.refresh_available_models().await {
422 error!("Failed to refresh available models: {err}");
423 }
424 let remote_models = self.remote_models.read().await;
425 self.build_available_models(remote_models.clone())
426 }
427
428 pub async fn list_models_for_provider(&self, provider: Provider) -> Vec<ModelPreset> {
430 let all_models = self.list_models().await;
431 all_models.into_iter().filter(|m| m.provider == provider).collect()
432 }
433
434 pub fn try_list_models(&self) -> Result<Vec<ModelPreset>, tokio::sync::TryLockError> {
436 let remote_models = self.remote_models.try_read()?;
437 Ok(self.build_available_models(remote_models.clone()))
438 }
439
440 pub async fn construct_model_family(&self, model: &str) -> ModelFamily {
442 find_family_for_model(model)
443 }
444
445 pub async fn get_model(&self, model: Option<&str>) -> String {
447 if let Some(m) = model {
448 return m.to_string();
449 }
450
451 if let Err(err) = self.refresh_available_models().await {
453 error!("Failed to refresh available models: {err}");
454 }
455
456 let provider = *self.current_provider.read().await;
458 self.get_default_model_for_provider(provider)
459 }
460
461 pub fn get_default_model_for_provider(&self, provider: Provider) -> String {
463 if let Some(preset) = self.local_models.iter().find(|p| p.provider == provider && p.is_default) {
465 return preset.model.clone();
466 }
467
468 match provider {
470 Provider::Gemini => GEMINI_DEFAULT_MODEL.to_string(),
471 Provider::OpenAI => OPENAI_DEFAULT_MODEL.to_string(),
472 Provider::Anthropic => ANTHROPIC_DEFAULT_MODEL.to_string(),
473 Provider::Copilot => crate::config::constants::models::copilot::DEFAULT_MODEL.to_string(),
474 Provider::DeepSeek => "deepseek-reasoner".to_string(),
475 Provider::Meta => crate::config::constants::models::meta::DEFAULT_MODEL.to_string(),
476 Provider::ZAI => "glm-5.3".to_string(),
477 Provider::Minimax => crate::config::constants::models::minimax::DEFAULT_MODEL.to_string(),
478 Provider::Mistral => crate::config::constants::models::mistral::MISTRAL_LARGE_3.to_string(),
479 Provider::OpenRouter => "xiaomi/mimo-v2.6-pro".to_string(),
480 Provider::Ollama => "gpt-oss:20b".to_string(),
481 Provider::OllamaCloud => crate::config::constants::models::ollama::DEFAULT_CLOUD_MODEL.to_string(),
482 Provider::LmStudio => crate::config::constants::models::lmstudio::DEFAULT_MODEL.to_string(),
483 Provider::LlamaCpp => crate::config::constants::models::llamacpp::DEFAULT_MODEL.to_string(),
484 Provider::Moonshot => crate::config::constants::models::moonshot::DEFAULT_MODEL.to_string(),
485 Provider::HuggingFace => "deepseek-ai/DeepSeek-V3-0324".to_string(),
486 Provider::OpenCodeZen => crate::config::constants::models::opencode_zen::DEFAULT_MODEL.to_string(),
487 Provider::OpenCodeGo => crate::config::constants::models::opencode_go::DEFAULT_MODEL.to_string(),
488 Provider::MiMo => crate::config::constants::models::mimo::DEFAULT_MODEL.to_string(),
489 Provider::Qwen => crate::config::constants::models::qwen::DEFAULT_MODEL.to_string(),
490 Provider::StepFun => crate::config::constants::models::stepfun::DEFAULT_MODEL.to_string(),
491 Provider::Evolink => crate::config::constants::models::evolink::DEFAULT_MODEL.to_string(),
492 Provider::Poolside => crate::config::constants::models::poolside::DEFAULT_MODEL.to_string(),
493 Provider::XAI => crate::config::constants::models::xai::DEFAULT_MODEL.to_string(),
494 Provider::NVIDIA => crate::config::constants::models::nvidia::DEFAULT_MODEL.to_string(),
495 Provider::MergeGateway => crate::config::constants::models::merge_gateway::DEFAULT_MODEL.to_string(),
496 Provider::Vercel => crate::config::constants::models::vercel::DEFAULT_MODEL.to_string(),
497 }
498 }
499
500 #[cfg(test)]
502 pub fn get_model_offline(model: Option<&str>) -> String {
503 model.unwrap_or(GEMINI_DEFAULT_MODEL).to_string()
504 }
505
506 #[cfg(test)]
508 pub fn construct_model_family_offline(model: &str) -> ModelFamily {
509 find_family_for_model(model)
510 }
511
512 #[cfg(test)]
514 async fn apply_remote_state_for_provider(
515 &self,
516 provider: Provider,
517 models: Vec<ModelInfo>,
518 etag: Option<String>,
519 ) -> bool {
520 let generation = self.refresh_generation.load(Ordering::Acquire);
521 self.apply_remote_state_for_provider_at(provider, generation, models, etag)
522 .await
523 }
524
525 async fn apply_remote_state_for_provider_at(
526 &self,
527 provider: Provider,
528 generation: u64,
529 models: Vec<ModelInfo>,
530 etag: Option<String>,
531 ) -> bool {
532 let current_provider = self.current_provider.read().await;
533 if *current_provider != provider || self.refresh_generation.load(Ordering::Acquire) != generation {
534 debug!(
535 requested_provider = %provider,
536 current_provider = %*current_provider,
537 "Discarding stale model metadata for a provider that is no longer active"
538 );
539 return false;
540 }
541
542 *self.remote_models.write().await = models;
543 *self.etag.write().await = etag;
544 true
545 }
546
547 #[cfg(test)]
549 async fn try_load_cache(&self) -> bool {
550 let provider = *self.current_provider.read().await;
551 self.try_load_cache_for(provider).await
552 }
553
554 #[cfg(test)]
555 async fn try_load_cache_for(&self, provider: Provider) -> bool {
556 let generation = self.refresh_generation.load(Ordering::Acquire);
557 self.try_load_cache_for_at(provider, generation).await
558 }
559
560 async fn try_load_cache_for_at(&self, provider: Provider, generation: u64) -> bool {
561 let Some(cache) = self.load_cache_for(provider).await else {
562 return false;
563 };
564 if !cache.is_fresh(self.cache_ttl) {
565 debug!("Cache is stale (age: {:?})", cache.age());
566 return false;
567 }
568 self.apply_remote_state_for_provider_at(provider, generation, cache.models.into_iter().collect(), cache.etag)
569 .await
570 }
571
572 async fn try_load_stale_cache_for_at(&self, provider: Provider, generation: u64) -> bool {
573 let Some(cache) = self.load_cache_for(provider).await else {
574 return false;
575 };
576 self.apply_remote_state_for_provider_at(provider, generation, cache.models.into_iter().collect(), cache.etag)
577 .await
578 }
579
580 async fn load_cache_for(&self, provider: Provider) -> Option<ModelsCache> {
581 let cache_path = self.cache_path_for(provider);
582 load_cache_from_paths(&cache_path, &self.legacy_cache_paths(provider), provider).await
583 }
584
585 #[cfg(test)]
587 async fn persist_cache(&self, models: &[ModelInfo], etag: Option<String>) {
588 let provider = *self.current_provider.read().await;
589 let generation = self.refresh_generation.load(Ordering::Acquire);
590 self.persist_cache_for_at(provider, generation, models, etag).await;
591 }
592
593 #[cfg(test)]
594 async fn persist_cache_for(&self, provider: Provider, models: &[ModelInfo], etag: Option<String>) {
595 let cache = ModelsCache {
596 fetched_at: Utc::now(),
597 etag,
598 provider: provider.to_string(),
599 models: models.to_vec(),
600 };
601 let cache_path = self.cache_path_for(provider);
602 cache::save_cache(&cache_path, &cache).await.expect("test cache write");
603 }
604
605 async fn persist_cache_for_at(
606 &self,
607 provider: Provider,
608 generation: u64,
609 models: &[ModelInfo],
610 etag: Option<String>,
611 ) {
612 let _write_guard = self.cache_write_lock.lock().await;
613 let current_provider = self.current_provider.read().await;
614 if *current_provider != provider || self.refresh_generation.load(Ordering::Acquire) != generation {
615 return;
616 }
617 let cache = ModelsCache {
618 fetched_at: Utc::now(),
619 etag,
620 provider: provider.to_string(),
621 models: models.to_vec(),
622 };
623 let cache_path = self.cache_path_for(provider);
624 if let Err(err) = cache::save_cache(&cache_path, &cache).await {
625 error!("Failed to write models cache: {err}");
626 }
627 }
628
629 fn build_available_models(&self, mut remote_models: Vec<ModelInfo>) -> Vec<ModelPreset> {
631 remote_models.sort_by_key(|a| a.priority);
633
634 let remote_presets: Vec<ModelPreset> = remote_models.into_iter().map(Into::into).collect();
636 let existing_presets = self.local_models.clone();
637 let mut merged_presets = Self::merge_presets(remote_presets, existing_presets);
638 merged_presets = self.filter_visible_models(merged_presets);
639
640 self.ensure_defaults(&mut merged_presets);
642
643 merged_presets
644 }
645
646 fn filter_visible_models(&self, models: Vec<ModelPreset>) -> Vec<ModelPreset> {
648 models
649 .into_iter()
650 .filter(|model| model.show_in_picker && model.supported_in_api)
651 .collect()
652 }
653
654 fn merge_presets(remote_presets: Vec<ModelPreset>, existing_presets: Vec<ModelPreset>) -> Vec<ModelPreset> {
656 if remote_presets.is_empty() {
657 return existing_presets;
658 }
659
660 let remote_slugs: HashSet<String> = remote_presets.iter().map(|preset| preset.model.clone()).collect();
661
662 let mut merged_presets = remote_presets;
663 for mut preset in existing_presets {
664 if remote_slugs.contains(&preset.model) {
665 continue;
666 }
667 preset.is_default = false;
668 merged_presets.push(preset);
669 }
670
671 merged_presets
672 }
673
674 fn ensure_defaults(&self, presets: &mut [ModelPreset]) {
676 let has_default = presets.iter().any(|p| p.is_default);
677 if !has_default && let Some(first) = presets.first_mut() {
678 first.is_default = true;
679 }
680 }
681
682 fn cache_path_for(&self, provider: Provider) -> PathBuf {
684 self.vtcode_home.join(format!("models_cache_{provider}.json"))
685 }
686
687 fn legacy_cache_paths(&self, provider: Provider) -> Vec<PathBuf> {
688 let provider_cache_file = format!("models_cache_{provider}.json");
689 let mut candidates = if let Ok(paths) = VtCodePaths::resolve()
690 && self.vtcode_home == paths.cache_dir()
691 {
692 vec![
693 paths.legacy_dir().join(&provider_cache_file),
694 paths.legacy_dir().join(LEGACY_MODEL_CACHE_FILE),
695 ]
696 } else {
697 vec![
698 self.vtcode_home.join(&provider_cache_file),
699 self.vtcode_home.join(LEGACY_MODEL_CACHE_FILE),
700 ]
701 };
702 candidates.dedup();
703 candidates
704 }
705
706 pub async fn set_provider(&self, provider: Provider) {
708 *self.current_provider.write().await = provider;
709 self.refresh_generation.fetch_add(1, Ordering::AcqRel);
710 self.remote_models.write().await.clear();
711 *self.etag.write().await = None;
712 }
713
714 pub async fn get_provider(&self) -> Provider {
716 *self.current_provider.read().await
717 }
718
719 pub async fn find_model(&self, model_id: &str) -> Option<ModelPreset> {
721 let models = self.list_models().await;
722 models.into_iter().find(|m| m.model == model_id || m.id == model_id)
723 }
724
725 pub async fn model_exists(&self, model_id: &str) -> bool {
727 self.find_model(model_id).await.is_some()
728 }
729
730 pub fn model_exists_sync(&self, model_id: &str) -> bool {
735 self.local_models.iter().any(|m| m.model == model_id || m.id == model_id)
736 }
737
738 pub fn supported_providers() -> Vec<Provider> {
740 Provider::all_providers()
741 }
742
743 pub fn client_version() -> String {
745 format!(
746 "{}.{}.{}",
747 env!("CARGO_PKG_VERSION_MAJOR"),
748 env!("CARGO_PKG_VERSION_MINOR"),
749 env!("CARGO_PKG_VERSION_PATCH")
750 )
751 }
752}
753
754async fn load_cache_from_paths(
755 canonical_path: &Path,
756 legacy_paths: &[PathBuf],
757 provider: Provider,
758) -> Option<ModelsCache> {
759 let expected_provider = provider.to_string();
760 let canonical_missing = match cache::load_cache(canonical_path).await {
761 Ok(Some(cache)) => {
762 if cache.provider == expected_provider {
763 return Some(cache);
764 }
765 debug!(
766 cached_provider = %cache.provider,
767 requested_provider = %provider,
768 "Ignoring model cache for a different provider"
769 );
770 false
771 }
772 Ok(None) => true,
773 Err(err) if err.kind() == std::io::ErrorKind::InvalidData => {
774 debug!(path = %canonical_path.display(), "Ignoring malformed models cache");
775 false
776 }
777 Err(err) => {
778 error!(path = %canonical_path.display(), "Failed to load models cache: {err}");
779 return None;
780 }
781 };
782
783 for legacy_path in legacy_paths {
784 if legacy_path == canonical_path {
785 continue;
786 }
787 let cache = match cache::load_cache(legacy_path).await {
788 Ok(Some(cache)) => cache,
789 Ok(None) => continue,
790 Err(err) => {
791 debug!(path = %legacy_path.display(), "Ignoring unreadable legacy models cache: {err}");
792 continue;
793 }
794 };
795 if cache.provider != expected_provider {
796 debug!(
797 path = %legacy_path.display(),
798 cached_provider = %cache.provider,
799 requested_provider = %provider,
800 "Ignoring legacy model cache for a different provider"
801 );
802 continue;
803 }
804
805 if canonical_missing {
806 if let Err(err) = cache::save_cache_if_absent(canonical_path, &cache).await {
807 debug!(path = %canonical_path.display(), "Failed to republish legacy models cache: {err}");
808 }
809 }
810 return Some(cache);
811 }
812
813 None
814}
815
816pub type SharedModelsManager = Arc<ModelsManager>;
818
819pub fn new_shared_models_manager() -> SharedModelsManager {
821 Arc::new(ModelsManager::new())
822}
823
824pub fn new_shared_models_manager_with_provider(provider: Provider) -> SharedModelsManager {
826 Arc::new(ModelsManager::with_provider(provider))
827}
828
829#[cfg(test)]
830mod tests {
831 use super::*;
832 use tempfile::tempdir;
833
834 #[tokio::test]
835 async fn test_new_manager() {
836 let manager = ModelsManager::new();
837 assert!(!manager.local_models.is_empty());
838 }
839
840 #[tokio::test]
841 async fn test_list_models() {
842 let manager = ModelsManager::new();
843 let models = manager.list_models().await;
844 assert!(!models.is_empty());
845 }
846
847 #[tokio::test]
848 async fn test_list_models_for_provider() {
849 let manager = ModelsManager::new();
850 let gemini_models = manager.list_models_for_provider(Provider::Gemini).await;
851 assert!(!gemini_models.is_empty());
852 assert!(gemini_models.iter().all(|m| m.provider == Provider::Gemini));
853 }
854
855 #[tokio::test]
856 async fn test_get_model_with_default() {
857 let manager = ModelsManager::with_provider(Provider::Gemini);
858 let model = manager.get_model(None).await;
859 assert!(!model.is_empty());
860 }
861
862 #[tokio::test]
863 async fn test_get_model_with_explicit() {
864 let manager = ModelsManager::new();
865 let model = manager.get_model(Some("custom-model")).await;
866 assert_eq!(model, "custom-model");
867 }
868
869 #[tokio::test]
870 async fn test_construct_model_family() {
871 let manager = ModelsManager::new();
872 let family = manager.construct_model_family("gemini-3-flash-preview").await;
873 assert_eq!(family.family, "gemini-3");
874 assert_eq!(family.provider, Provider::Gemini);
875 }
876
877 #[tokio::test]
878 async fn test_find_model() {
879 let manager = ModelsManager::new();
880 let model = manager.find_model("gemini-3-flash-preview").await;
881 assert!(model.is_some());
882 }
883
884 #[tokio::test]
885 async fn test_model_exists() {
886 let manager = ModelsManager::new();
887 assert!(manager.model_exists("gemini-3-flash-preview").await);
888 assert!(!manager.model_exists("nonexistent-model").await);
889 }
890
891 #[tokio::test]
892 async fn test_set_provider() {
893 let manager = ModelsManager::new();
894 manager.set_provider(Provider::Anthropic).await;
895 assert_eq!(manager.get_provider().await, Provider::Anthropic);
896 }
897
898 #[test]
899 fn merge_catalog_models_map_to_picker_metadata_and_hide_deprecated_routes() {
900 let model = MergeCatalogModel {
901 model: "anthropic/claude-sonnet-5".to_string(),
902 provider: "anthropic".to_string(),
903 display_name: Some("Claude Sonnet 5".to_string()),
904 availability: MergeCatalogAvailability::Available,
905 context_window: Some(200_000),
906 max_output_tokens: Some(16_384),
907 supports_tool_use: true,
908 supports_streaming: true,
909 supports_vision: true,
910 supports_structured_output: true,
911 service_tiers: Vec::new(),
912 supports_reasoning: true,
913 reasoning_disable_supported: true,
914 reasoning_controls: vec!["thinking.budget_tokens".to_string()],
915 };
916
917 let info = ModelsManager::merge_catalog_model_info(model).expect("available model should map");
918 assert_eq!(info.slug, "anthropic/claude-sonnet-5");
919 assert_eq!(info.display_name, "Claude Sonnet 5");
920 assert_eq!(info.provider, Provider::MergeGateway);
921 assert_eq!(info.context_window, Some(200_000));
922 assert_eq!(info.max_output_tokens, Some(16_384));
923 assert!(info.supports_vision);
924 assert!(info.supports_structured_output);
925 assert!(info.supports_reasoning);
926 assert_eq!(info.supported_reasoning_levels.len(), ModelsManager::merge_supported_reasoning_presets(true).len());
927
928 let deprecated = MergeCatalogModel {
929 availability: MergeCatalogAvailability::Deprecated,
930 ..MergeCatalogModel {
931 model: "openai/old".to_string(),
932 provider: "openai".to_string(),
933 display_name: None,
934 availability: MergeCatalogAvailability::Available,
935 context_window: None,
936 max_output_tokens: None,
937 supports_tool_use: false,
938 supports_streaming: false,
939 supports_vision: false,
940 supports_structured_output: false,
941 service_tiers: Vec::new(),
942 supports_reasoning: false,
943 reasoning_disable_supported: false,
944 reasoning_controls: Vec::new(),
945 }
946 };
947 assert!(ModelsManager::merge_catalog_model_info(deprecated).is_none());
948 }
949
950 #[tokio::test]
951 async fn test_cache_operations() {
952 let dir = tempdir().expect("create temp dir");
953 let manager = ModelsManager::with_home(dir.path().to_path_buf());
954
955 let cached = manager.try_load_cache().await;
957 assert!(!cached);
958
959 let models = vec![ModelInfo {
961 slug: "test-model".to_string(),
962 display_name: "Test Model".to_string(),
963 description: "A test".to_string(),
964 provider: Provider::Gemini,
965 default_reasoning_level: crate::config::types::ReasoningEffortLevel::Medium,
966 supported_reasoning_levels: vec![],
967 context_window: Some(128_000),
968 supports_tool_use: true,
969 supports_streaming: true,
970 supports_vision: false,
971 supports_structured_output: false,
972 supports_reasoning: false,
973 max_output_tokens: None,
974 priority: 0,
975 visibility: "list".to_string(),
976 supported_in_api: true,
977 upgrade: None,
978 }];
979 manager.persist_cache(&models, None).await;
980
981 let cached = manager.try_load_cache().await;
983 assert!(cached);
984 }
985
986 #[tokio::test]
987 async fn loads_provider_scoped_legacy_cache_and_republishes_it() {
988 let temp_dir = tempdir().expect("create temp dir");
989 let canonical_path = temp_dir.path().join("current/models_cache_gemini.json");
990 let legacy_path = temp_dir.path().join("legacy/models_cache_gemini.json");
991 let legacy_cache = ModelsCache::new(Provider::Gemini.to_string(), Vec::new());
992 cache::save_cache(&legacy_path, &legacy_cache).await.expect("save legacy cache");
993
994 let loaded = load_cache_from_paths(&canonical_path, std::slice::from_ref(&legacy_path), Provider::Gemini)
995 .await
996 .expect("load legacy cache");
997
998 assert_eq!(loaded.provider, Provider::Gemini.to_string());
999 assert!(
1000 cache::load_cache(&canonical_path)
1001 .await
1002 .expect("load republished cache")
1003 .is_some()
1004 );
1005 }
1006
1007 #[tokio::test]
1008 async fn malformed_canonical_cache_recovers_legacy_without_replacing_it() {
1009 let temp_dir = tempdir().expect("create temp dir");
1010 let canonical_path = temp_dir.path().join("current/models_cache_gemini.json");
1011 let legacy_path = temp_dir.path().join("legacy/models_cache_gemini.json");
1012 std::fs::create_dir_all(canonical_path.parent().expect("canonical parent")).expect("canonical directory");
1013 std::fs::write(&canonical_path, b"not json").expect("malformed canonical cache");
1014 let legacy_cache = ModelsCache::new(Provider::Gemini.to_string(), Vec::new());
1015 cache::save_cache(&legacy_path, &legacy_cache).await.expect("save legacy cache");
1016
1017 let loaded = load_cache_from_paths(&canonical_path, std::slice::from_ref(&legacy_path), Provider::Gemini)
1018 .await
1019 .expect("recover legacy cache");
1020
1021 assert_eq!(loaded.provider, Provider::Gemini.to_string());
1022 assert_eq!(std::fs::read(&canonical_path).expect("read canonical cache"), b"not json");
1023 }
1024
1025 #[tokio::test]
1026 async fn mismatched_canonical_cache_falls_back_to_matching_legacy_cache() {
1027 let temp_dir = tempdir().expect("create temp dir");
1028 let canonical_path = temp_dir.path().join("current/models_cache_gemini.json");
1029 let legacy_path = temp_dir.path().join("legacy/models_cache_gemini.json");
1030 let canonical_cache = ModelsCache::new(Provider::OpenAI.to_string(), Vec::new());
1031 cache::save_cache(&canonical_path, &canonical_cache)
1032 .await
1033 .expect("save mismatched cache");
1034 let legacy_cache = ModelsCache::new(Provider::Gemini.to_string(), Vec::new());
1035 cache::save_cache(&legacy_path, &legacy_cache).await.expect("save legacy cache");
1036
1037 let loaded = load_cache_from_paths(&canonical_path, std::slice::from_ref(&legacy_path), Provider::Gemini)
1038 .await
1039 .expect("recover matching legacy cache");
1040
1041 assert_eq!(loaded.provider, Provider::Gemini.to_string());
1042 assert_eq!(
1043 cache::load_cache(&canonical_path)
1044 .await
1045 .expect("read canonical cache")
1046 .unwrap()
1047 .provider,
1048 Provider::OpenAI.to_string()
1049 );
1050 }
1051
1052 #[test]
1053 fn legacy_cache_paths_preserve_provider_scoped_and_unscoped_names() {
1054 let temp_dir = tempdir().expect("create temp dir");
1055 let manager = ModelsManager::with_home(temp_dir.path().to_path_buf());
1056
1057 assert_eq!(
1058 manager.legacy_cache_paths(Provider::Gemini),
1059 vec![
1060 temp_dir.path().join("models_cache_gemini.json"),
1061 temp_dir.path().join(LEGACY_MODEL_CACHE_FILE),
1062 ]
1063 );
1064 }
1065
1066 #[tokio::test]
1067 async fn ignores_fresh_cache_for_another_provider() {
1068 let dir = tempdir().expect("create temp dir");
1069 let manager = ModelsManager::with_home_and_provider(dir.path().to_path_buf(), Provider::Gemini);
1070 let cache = ModelsCache::new("openai", Vec::new());
1071 cache::save_cache(&manager.cache_path_for(Provider::Gemini), &cache)
1072 .await
1073 .expect("save mismatched cache");
1074
1075 assert!(!manager.try_load_cache().await);
1076 assert!(manager.remote_models.read().await.is_empty());
1077 }
1078
1079 #[tokio::test]
1080 async fn ignores_stale_inflight_models_after_provider_switch() {
1081 let dir = tempdir().expect("create temp dir");
1082 let manager = ModelsManager::with_home_and_provider(dir.path().to_path_buf(), Provider::Gemini);
1083 manager.set_provider(Provider::Anthropic).await;
1084
1085 let models = vec![ModelInfo {
1086 slug: "gemini-stale".to_string(),
1087 display_name: "Gemini Stale".to_string(),
1088 description: "Stale fetch result".to_string(),
1089 provider: Provider::Gemini,
1090 default_reasoning_level: crate::config::types::ReasoningEffortLevel::Medium,
1091 supported_reasoning_levels: vec![],
1092 context_window: Some(128_000),
1093 supports_tool_use: true,
1094 supports_streaming: true,
1095 supports_vision: false,
1096 supports_structured_output: false,
1097 supports_reasoning: false,
1098 max_output_tokens: None,
1099 priority: 0,
1100 visibility: "list".to_string(),
1101 supported_in_api: true,
1102 upgrade: None,
1103 }];
1104
1105 assert!(
1106 !manager
1107 .apply_remote_state_for_provider(Provider::Gemini, models, Some("etag-1".to_string()))
1108 .await
1109 );
1110 assert!(manager.remote_models.read().await.is_empty());
1111 assert!(manager.etag.read().await.is_none());
1112 }
1113
1114 #[tokio::test]
1115 async fn persists_cache_under_requested_provider_scope() {
1116 let dir = tempdir().expect("create temp dir");
1117 let manager = ModelsManager::with_home_and_provider(dir.path().to_path_buf(), Provider::Gemini);
1118 manager.set_provider(Provider::Anthropic).await;
1119
1120 let models = vec![ModelInfo {
1121 slug: "gemini-cached".to_string(),
1122 display_name: "Gemini Cached".to_string(),
1123 description: "Cached fetch result".to_string(),
1124 provider: Provider::Gemini,
1125 default_reasoning_level: crate::config::types::ReasoningEffortLevel::Medium,
1126 supported_reasoning_levels: vec![],
1127 context_window: Some(128_000),
1128 supports_tool_use: true,
1129 supports_streaming: true,
1130 supports_vision: false,
1131 supports_structured_output: false,
1132 supports_reasoning: false,
1133 max_output_tokens: None,
1134 priority: 0,
1135 visibility: "list".to_string(),
1136 supported_in_api: true,
1137 upgrade: None,
1138 }];
1139
1140 manager.persist_cache_for(Provider::Gemini, &models, None).await;
1141
1142 let gemini_cache = cache::load_cache(&manager.cache_path_for(Provider::Gemini))
1143 .await
1144 .expect("load gemini cache");
1145 let anthropic_cache = cache::load_cache(&manager.cache_path_for(Provider::Anthropic))
1146 .await
1147 .expect("load anthropic cache");
1148
1149 assert!(gemini_cache.is_some());
1150 assert!(anthropic_cache.is_none());
1151 }
1152
1153 #[test]
1154 fn test_client_version() {
1155 let version = ModelsManager::client_version();
1156 assert!(!version.is_empty());
1157 assert!(version.contains('.'));
1158 }
1159
1160 #[test]
1161 fn test_supported_providers() {
1162 let providers = ModelsManager::supported_providers();
1163 assert!(!providers.is_empty());
1164 assert!(providers.contains(&Provider::Gemini));
1165 assert!(providers.contains(&Provider::OpenAI));
1166 }
1167
1168 #[test]
1169 fn moonshot_default_model_uses_curated_default() {
1170 let manager = ModelsManager::new();
1171 assert_eq!(
1172 manager.get_default_model_for_provider(Provider::Moonshot),
1173 crate::config::constants::models::moonshot::DEFAULT_MODEL
1174 );
1175 }
1176}