navi_core/
background_model.rs1use crate::credentials::{CredentialStore, resolve_provider_api_key};
5use crate::registry::RegistryStore;
6use crate::{ProviderConfig, resolve_provider_config};
7use std::sync::{Arc, RwLock};
8
9#[derive(Debug, Clone)]
11pub struct ResolvedBackgroundModel {
12 pub provider_id: String,
14 pub model_name: String,
16 pub provider_config: ProviderConfig,
18}
19
20pub struct BackgroundModelResolver {
23 registry: Option<Arc<RegistryStore>>,
24 config: Arc<RwLock<crate::config::NaviConfig>>,
25 credential_store: CredentialStore,
26}
27
28impl BackgroundModelResolver {
29 pub fn new(
31 registry: Option<Arc<RegistryStore>>,
32 config: Arc<RwLock<crate::config::NaviConfig>>,
33 credential_store: CredentialStore,
34 ) -> Self {
35 Self {
36 registry,
37 config,
38 credential_store,
39 }
40 }
41
42 pub fn resolve(&self, task: &str) -> ResolvedBackgroundModel {
50 let config = self.config.read().unwrap_or_else(|e| e.into_inner());
51 let bg_config = &config.background_models;
52
53 if let Some(entry) = bg_config.resolve(task) {
55 if let (Some(provider), Some(model)) = (&entry.provider, &entry.model)
56 && let Some(resolved) = self.try_explicit(provider, model, &config)
57 {
58 return resolved;
59 }
60 if let Some(profile) = &entry.profile
62 && let Some(resolved) = self.resolve_from_profile(profile)
63 {
64 return resolved;
65 }
66 }
67
68 let default_profile = match task {
70 "naming" => "naming",
71 "memory_extraction" => "cheap_general",
72 "repo_search" => "repo_search",
73 "compaction" => "long_context_cheap",
74 "subagent_research" => "research_synthesis",
75 "simple_code_edit" => "cheap_code",
76 _ => "cheap_general",
77 };
78 if let Some(resolved) = self.resolve_from_profile(default_profile) {
79 return resolved;
80 }
81
82 self.main_model_fallback(&config)
84 }
85
86 fn resolve_from_profile(&self, profile_id: &str) -> Option<ResolvedBackgroundModel> {
88 let registry = self.registry.as_ref()?;
89 let ranked = registry.query_models_by_profile(profile_id).ok()?;
90
91 for candidate in &ranked {
92 if self.has_credential(&candidate.provider_id) {
93 let config = self.config.read().unwrap_or_else(|e| e.into_inner());
94 if let Some(resolved) =
95 self.try_explicit(&candidate.provider_id, &candidate.model_name, &config)
96 {
97 return Some(resolved);
98 }
99 }
100 }
101
102 None
103 }
104
105 fn has_credential(&self, provider_id: &str) -> bool {
107 let config = self.config.read().unwrap_or_else(|e| e.into_inner());
108 let Some(provider_config) = resolve_provider_config(&config, provider_id) else {
109 return false;
110 };
111 resolve_provider_api_key(&self.credential_store, &provider_config, provider_id).is_some()
112 }
113
114 fn try_explicit(
116 &self,
117 provider_id: &str,
118 model_name: &str,
119 config: &crate::config::NaviConfig,
120 ) -> Option<ResolvedBackgroundModel> {
121 let provider_config = resolve_provider_config(config, provider_id)?;
122 Some(ResolvedBackgroundModel {
123 provider_id: provider_id.to_string(),
124 model_name: model_name.to_string(),
125 provider_config,
126 })
127 }
128
129 fn main_model_fallback(&self, config: &crate::config::NaviConfig) -> ResolvedBackgroundModel {
131 let provider_id = config.model.provider.clone();
132 let model_name = config.model.name.clone();
133 let provider_config =
134 resolve_provider_config(config, &provider_id).unwrap_or_else(|| ProviderConfig {
135 id: provider_id.clone(),
136 ..ProviderConfig::default()
137 });
138 ResolvedBackgroundModel {
139 provider_id,
140 model_name,
141 provider_config,
142 }
143 }
144}
145
146#[cfg(test)]
147mod tests {
148 use super::*;
149 use crate::NaviConfig;
150 use crate::config::types::{BackgroundModelEntry, ModelConfig};
151
152 fn test_resolver(config: NaviConfig) -> BackgroundModelResolver {
153 let tempdir = tempfile::tempdir().unwrap();
154 let registry = RegistryStore::open_memory().ok().map(Arc::new);
155 let config = Arc::new(RwLock::new(config));
156 let cred_store = CredentialStore::new(tempdir.path().to_path_buf());
157 BackgroundModelResolver::new(registry, config, cred_store)
158 }
159
160 #[test]
161 fn resolve_falls_back_to_main_model() {
162 let config = NaviConfig {
163 model: ModelConfig {
164 provider: "openai".to_string(),
165 name: "gpt-5.5".to_string(),
166 },
167 ..Default::default()
168 };
169 let resolver = test_resolver(config);
170 let resolved = resolver.resolve("naming");
171 assert_eq!(resolved.provider_id, "openai");
172 assert_eq!(resolved.model_name, "gpt-5.5");
173 }
174
175 #[test]
176 fn resolve_explicit_override() {
177 let mut config = NaviConfig {
178 model: ModelConfig {
179 provider: "openai".to_string(),
180 name: "gpt-5.5".to_string(),
181 },
182 ..Default::default()
183 };
184 config.background_models.naming = Some(BackgroundModelEntry {
185 profile: None,
186 provider: Some("anthropic".to_string()),
187 model: Some("claude-haiku".to_string()),
188 fallback: None,
189 });
190 config.providers.push(crate::ProviderConfig {
192 id: "anthropic".to_string(),
193 label: "Anthropic".to_string(),
194 kind: crate::ProviderKind::AnthropicMessages,
195 api_key_env: "ANTHROPIC_API_KEY".to_string(),
196 ..Default::default()
197 });
198
199 let resolver = test_resolver(config);
200 let resolved = resolver.resolve("naming");
201 assert_eq!(resolved.provider_id, "anthropic");
205 assert_eq!(resolved.model_name, "claude-haiku");
206 }
207
208 #[test]
209 fn task_to_default_profile_mapping() {
210 let config = NaviConfig::default();
211 let resolver = test_resolver(config);
212 for task in &["naming", "repo_search", "compaction", "subagent_research"] {
214 let resolved = resolver.resolve(task);
215 assert_eq!(resolved.provider_id, "openai");
217 }
218 }
219}