1use std::collections::BTreeMap;
33use std::path::{Path, PathBuf};
34
35use crate::catalog::{BuiltinModelEntry, BuiltinProviderEntry};
36
37#[derive(Debug, Default, Clone, serde::Deserialize)]
42pub struct OverrideFile {
43 #[serde(default)]
45 pub provider: Vec<BuiltinProviderEntry>,
46 #[serde(default)]
48 pub model: Vec<BuiltinModelEntry>,
49}
50
51pub fn find_override_files() -> Vec<(PathBuf, String)> {
59 find_override_files_at(crate::product_env::catalog_override_dir().as_deref())
60}
61
62fn find_override_files_at(global_dir: Option<&Path>) -> Vec<(PathBuf, String)> {
70 let mut out = Vec::new();
71
72 if let Ok(path) = std::env::var("OXICODE_CATALOG_OVERRIDE")
74 && let Some(pair) = read_override(&PathBuf::from(path))
75 {
76 out.push(pair);
77 }
78
79 if let Some(dir) = global_dir {
81 let path = dir.join("overrides.toml");
82 if let Some(pair) = read_override(&path) {
83 out.push(pair);
84 }
85 }
86
87 let path = PathBuf::from(".oxicode/catalog.local.toml");
89 if let Some(pair) = read_override(&path) {
90 out.push(pair);
91 }
92
93 out
94}
95
96fn read_override(path: &Path) -> Option<(PathBuf, String)> {
97 if !path.exists() {
98 return None;
99 }
100 match std::fs::read_to_string(path) {
101 Ok(content) => Some((path.to_path_buf(), content)),
102 Err(e) => {
103 tracing::warn!(?path, error = %e, "Failed to read override file");
104 None
105 }
106 }
107}
108
109pub fn load_overrides() -> Option<OverrideFile> {
117 let files = find_override_files();
118 if files.is_empty() {
119 return None;
120 }
121
122 let mut merged = OverrideFile::default();
123 for (path, content) in files {
124 match toml::from_str::<OverrideFile>(&content) {
125 Ok(file) => {
126 tracing::info!(
127 ?path,
128 providers = file.provider.len(),
129 models = file.model.len(),
130 "Loaded catalog override"
131 );
132 merged.provider.extend(file.provider);
133 merged.model.extend(file.model);
134 }
135 Err(e) => {
136 tracing::warn!(?path, error = %e, "Failed to parse override file; skipping");
137 }
138 }
139 }
140
141 if merged.provider.is_empty() && merged.model.is_empty() {
142 None
143 } else {
144 Some(merged)
145 }
146}
147
148pub fn apply_provider_overrides(
153 providers: &mut Vec<BuiltinProviderEntry>,
154 overrides: &[BuiltinProviderEntry],
155) {
156 for ov in overrides {
157 if let Some(existing) = providers.iter_mut().find(|p| p.id == ov.id) {
158 tracing::debug!(provider = %ov.id, "Replacing built-in provider with override");
159 *existing = ov.clone();
160 } else {
161 tracing::debug!(provider = %ov.id, "Adding new provider from override");
162 providers.push(ov.clone());
163 }
164 }
165}
166
167pub fn apply_model_overrides(
172 models: &mut BTreeMap<String, Vec<BuiltinModelEntry>>,
173 overrides: &[BuiltinModelEntry],
174) {
175 for ov in overrides {
176 let entry = models.entry(ov.provider.clone()).or_default();
177 if let Some(existing) = entry.iter_mut().find(|m| m.id == ov.id) {
178 tracing::debug!(provider = %ov.provider, model = %ov.id,
179 "Replacing built-in model with override");
180 *existing = ov.clone();
181 } else {
182 tracing::debug!(provider = %ov.provider, model = %ov.id,
183 "Adding new model from override");
184 entry.push(ov.clone());
185 }
186 }
187}
188
189#[cfg(test)]
190mod tests {
191 use super::*;
192
193 #[test]
194 fn parse_minimal_override() {
195 let toml = r#"
196 [[provider]]
197 id = "my-company-gateway"
198 display_name = "My Company AI Gateway"
199 env_key = "MY_GATEWAY_API_KEY"
200 api = "openai-completions"
201 auth_method = "bearer"
202 category = "enterprise"
203 description = "Internal AI gateway"
204
205 [[model]]
206 id = "my-company-gpt"
207 name = "Internal GPT-4 variant"
208 api = "openai-completions"
209 provider = "my-company-gateway"
210 context_window = 128000
211 max_tokens = 8192
212 cost_input = 1.0
213 cost_output = 2.0
214 "#;
215 let parsed: OverrideFile = toml::from_str(toml).expect("parse");
216 assert_eq!(parsed.provider.len(), 1);
217 assert_eq!(parsed.model.len(), 1);
218 assert_eq!(parsed.provider[0].id, "my-company-gateway");
219 assert_eq!(parsed.model[0].id, "my-company-gpt");
220 }
221
222 #[test]
223 fn apply_provider_override_replaces() {
224 let mut providers = vec![BuiltinProviderEntry {
225 id: "anthropic".into(),
226 display_name: "Anthropic".into(),
227 api: "anthropic-messages".into(),
228 env_key: "ANTHROPIC_API_KEY".into(),
229 category: "primary".into(),
230 description: "Old".into(),
231 auth_method: crate::catalog::AuthMethod::XApiKey,
232 aliases: vec![],
233 extra_env_keys: vec![],
234 base_url: "".into(),
235 extra_headers: vec![],
236 default_enabled: true,
237 }];
238 let overrides = vec![BuiltinProviderEntry {
239 id: "anthropic".into(),
240 display_name: "Anthropic (Custom Pricing)".into(),
241 api: "anthropic-messages".into(),
242 env_key: "ANTHROPIC_API_KEY".into(),
243 category: "primary".into(),
244 description: "New".into(),
245 auth_method: crate::catalog::AuthMethod::XApiKey,
246 aliases: vec![],
247 extra_env_keys: vec![],
248 base_url: "".into(),
249 extra_headers: vec![],
250 default_enabled: true,
251 }];
252 apply_provider_overrides(&mut providers, &overrides);
253 assert_eq!(providers.len(), 1);
254 assert_eq!(providers[0].display_name, "Anthropic (Custom Pricing)");
255 }
256
257 #[test]
258 fn apply_model_override_appends_new() {
259 let mut models: BTreeMap<String, Vec<BuiltinModelEntry>> = BTreeMap::new();
260 models.insert("anthropic".into(), vec![]);
261 let overrides = vec![BuiltinModelEntry {
262 id: "claude-test".into(),
263 name: "Test".into(),
264 api: "anthropic-messages".into(),
265 provider: "anthropic".into(),
266 reasoning: false,
267 input: vec!["text".into()],
268 cost_input: 1.0,
269 cost_output: 2.0,
270 cost_cache_read: 0.0,
271 cost_cache_write: 0.0,
272 context_window: 200000,
273 max_tokens: 8192,
274 auth_method: crate::catalog::provider::AuthMethod::Bearer,
275 base_url: None,
276 }];
277 apply_model_overrides(&mut models, &overrides);
278 assert_eq!(models.get("anthropic").unwrap().len(), 1);
279 }
280 #[test]
286 fn find_override_files_at_reads_global_dir() {
287 let tmp = tempfile::TempDir::new().expect("tempdir");
288 let catalog_dir = tmp.path().join("catalog");
289 std::fs::create_dir_all(&catalog_dir).expect("mkdir catalog");
290 let override_path = catalog_dir.join("overrides.toml");
291 std::fs::write(
292 &override_path,
293 "[[provider]]\nid = \"oxicode-home-regression\"\napi = \"openai-completions\"\n",
294 )
295 .expect("write override");
296
297 let files = find_override_files_at(Some(&catalog_dir));
298 let found = files.iter().any(|(p, _)| p == &override_path);
299 assert!(
300 found,
301 "global-dir override must be discovered; got {:?}",
302 files.iter().map(|(p, _)| p).collect::<Vec<_>>()
303 );
304 }
305
306 #[test]
309 fn find_override_files_at_none_skips_global() {
310 let files = find_override_files_at(None);
311 let _ = files;
314 }
315}