Skip to main content

vtcode_core/models_manager/
manager.rs

1//! Models Manager - Coordinates model discovery, caching, and selection.
2//!
3//! This module provides the main `ModelsManager` struct that coordinates:
4//! - Local model presets (built-in configurations)
5//! - Remote model discovery (fetching from provider APIs)
6//! - Disk caching with TTL
7//! - Model family resolution
8
9use 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
32/// Legacy cache file name retained for one-time compatible reads.
33const LEGACY_MODEL_CACHE_FILE: &str = "models_cache.json";
34
35/// Default cache TTL (2 minutes)
36const DEFAULT_MODEL_CACHE_TTL: Duration = Duration::from_secs(120);
37
38/// Default model for Gemini provider
39const GEMINI_DEFAULT_MODEL: &str = "gemini-3-flash-preview";
40
41/// Default model for OpenAI provider
42const OPENAI_DEFAULT_MODEL: &str = "gpt-5.6-sol";
43
44/// Default model for Anthropic provider
45const ANTHROPIC_DEFAULT_MODEL: &str = "claude-sonnet-5";
46
47/// Coordinates remote model discovery plus cached metadata on disk.
48#[derive(Debug)]
49pub struct ModelsManager {
50    /// Local built-in model presets
51    local_models: Vec<ModelPreset>,
52    /// Remote models fetched from provider APIs
53    remote_models: RwLock<Vec<ModelInfo>>,
54    /// ETag for conditional requests
55    etag: RwLock<Option<String>>,
56    /// VT Code home directory for cache storage
57    vtcode_home: PathBuf,
58    /// Cache TTL
59    cache_ttl: Duration,
60    /// Current active provider
61    current_provider: RwLock<Provider>,
62    /// Monotonically increasing generation used to reject late refresh results.
63    refresh_generation: AtomicU64,
64    /// Whether remote model fetching is enabled
65    remote_models_enabled: bool,
66    /// Serializes cache writes so stale concurrent refreshes cannot overwrite newer snapshots.
67    cache_write_lock: Mutex<()>,
68}
69
70impl Default for ModelsManager {
71    fn default() -> Self {
72        Self::new()
73    }
74}
75
76impl ModelsManager {
77    /// Construct a new ModelsManager with default settings
78    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    /// Construct with a specific home directory
94    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    /// Construct with a specific provider
109    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    /// Construct with specific home directory and provider
125    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    /// Enable or disable remote model fetching
140    pub fn set_remote_models_enabled(&mut self, enabled: bool) {
141        self.remote_models_enabled = enabled;
142    }
143
144    /// Set the cache TTL
145    pub fn set_cache_ttl(&mut self, ttl: Duration) {
146        self.cache_ttl = ttl;
147    }
148
149    /// Get the default VT Code home directory
150    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    /// Refresh available models, using cache if fresh
157    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        // Try to load from cache first
167        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                        // Fall back to local presets if fetch fails
189                        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                // For other providers, we don't have remote discovery yet
215                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    /// Fetch models from Ollama API
352    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), // Default for most Ollama models
374                        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    /// List available models for the current provider
420    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    /// List available models for a specific provider
429    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    /// Try to list models without async refresh (uses cache only)
435    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    /// Get the model family for a given model slug
441    pub async fn construct_model_family(&self, model: &str) -> ModelFamily {
442        find_family_for_model(model)
443    }
444
445    /// Get the model to use, resolving defaults if not specified
446    pub async fn get_model(&self, model: Option<&str>) -> String {
447        if let Some(m) = model {
448            return m.to_string();
449        }
450
451        // Refresh models to ensure we have the latest
452        if let Err(err) = self.refresh_available_models().await {
453            error!("Failed to refresh available models: {err}");
454        }
455
456        // Return default for current provider
457        let provider = *self.current_provider.read().await;
458        self.get_default_model_for_provider(provider)
459    }
460
461    /// Get the default model for a specific provider
462    pub fn get_default_model_for_provider(&self, provider: Provider) -> String {
463        // First check if there's a default in local presets
464        if let Some(preset) = self.local_models.iter().find(|p| p.provider == provider && p.is_default) {
465            return preset.model.clone();
466        }
467
468        // Fall back to hardcoded defaults
469        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    /// Get model offline (without network) for testing
501    #[cfg(test)]
502    pub fn get_model_offline(model: Option<&str>) -> String {
503        model.unwrap_or(GEMINI_DEFAULT_MODEL).to_string()
504    }
505
506    /// Construct model family offline for testing
507    #[cfg(test)]
508    pub fn construct_model_family_offline(model: &str) -> ModelFamily {
509        find_family_for_model(model)
510    }
511
512    /// Apply remote models (replace cached state)
513    #[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    /// Try to load from cache
548    #[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    /// Persist cache to disk
586    #[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    /// Build available models by merging remote and local presets
630    fn build_available_models(&self, mut remote_models: Vec<ModelInfo>) -> Vec<ModelPreset> {
631        // Sort by priority
632        remote_models.sort_by_key(|a| a.priority);
633
634        // Convert remote models to presets
635        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        // Ensure one default per provider
641        self.ensure_defaults(&mut merged_presets);
642
643        merged_presets
644    }
645
646    /// Filter to only visible models
647    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    /// Merge remote and local presets, preferring remote when duplicates exist
655    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    /// Ensure there's at least one default model
675    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    /// Get the provider-scoped cache file path.
683    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    /// Set the current provider
707    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    /// Get the current provider
715    pub async fn get_provider(&self) -> Provider {
716        *self.current_provider.read().await
717    }
718
719    /// Find a model preset by ID
720    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    /// Check if a model exists
726    pub async fn model_exists(&self, model_id: &str) -> bool {
727        self.find_model(model_id).await.is_some()
728    }
729
730    /// Check if a model exists (sync, uses local presets only)
731    ///
732    /// This is a fast, non-blocking check that only looks at local presets.
733    /// Use `model_exists` for the async version that includes remote models.
734    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    /// Get all supported providers
739    pub fn supported_providers() -> Vec<Provider> {
740        Provider::all_providers()
741    }
742
743    /// Get version string for API requests
744    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
816/// Thread-safe reference-counted ModelsManager
817pub type SharedModelsManager = Arc<ModelsManager>;
818
819/// Create a new shared ModelsManager
820pub fn new_shared_models_manager() -> SharedModelsManager {
821    Arc::new(ModelsManager::new())
822}
823
824/// Create a new shared ModelsManager with specific provider
825pub 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        // Initially no cache
956        let cached = manager.try_load_cache().await;
957        assert!(!cached);
958
959        // Persist some models
960        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        // Now cache should load
982        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}