1use crate::config::types::{
4 ModelTaskSize, ProviderConfig, ProviderKind, ProviderModelConfig, ToolCallingMode,
5};
6use anyhow::{Context, Result};
7use rusqlite::{Connection, params};
8use std::path::Path;
9use std::sync::Mutex;
10
11use super::types::{
12 ModelCapability, ModelPricing, Profile, RankedModel, RegistryAttachments, RegistryManifest,
13 RegistryModel, RegistryProvider, RegistryTranscriptionProvider,
14};
15
16pub const LOCAL_API_SYNC_SHA: &str = "local-api-sync";
21
22fn remove_registry_db_files(db_path: &Path) {
24 let path_str = db_path.as_os_str().to_string_lossy();
25 for suffix in ["", "-wal", "-shm", "-journal"] {
26 let path = Path::new(&format!("{path_str}{suffix}")).to_path_buf();
27 match std::fs::remove_file(&path) {
28 Ok(()) => tracing::info!(path = %path.display(), "removed broken registry DB file"),
29 Err(err) if err.kind() == std::io::ErrorKind::NotFound => {}
30 Err(err) => tracing::warn!(
31 path = %path.display(),
32 error = %err,
33 "failed to remove broken registry DB file"
34 ),
35 }
36 }
37}
38
39pub struct RegistryStore {
44 conn: Mutex<Connection>,
45}
46
47impl RegistryStore {
48 pub fn open(data_dir: &Path) -> Result<Self> {
59 std::fs::create_dir_all(data_dir)
60 .with_context(|| format!("failed to create data dir {}", data_dir.display()))?;
61 let db_path = data_dir.join("registry.db");
62 match Self::open_at_path(&db_path) {
63 Ok(store) => Ok(store),
64 Err(first_err) => {
65 tracing::warn!(
66 error = %first_err,
67 path = %db_path.display(),
68 "registry DB unreadable; recreating cache from embedded snapshot"
69 );
70 remove_registry_db_files(&db_path);
71 Self::open_at_path(&db_path).with_context(|| {
72 format!(
73 "failed to open registry DB at {} after recreate (original error: {first_err})",
74 db_path.display()
75 )
76 })
77 }
78 }
79 }
80
81 fn open_at_path(db_path: &Path) -> Result<Self> {
82 let conn = Connection::open(db_path)
83 .with_context(|| format!("failed to open registry DB at {}", db_path.display()))?;
84
85 conn.query_row("PRAGMA schema_version", [], |row| row.get::<_, i64>(0))
88 .with_context(|| {
89 format!(
90 "registry DB at {} failed integrity probe",
91 db_path.display()
92 )
93 })?;
94
95 conn.pragma_update(None, "journal_mode", "WAL")?;
97 conn.pragma_update(None, "foreign_keys", "ON")?;
98
99 let store = Self {
100 conn: Mutex::new(conn),
101 };
102 store.init_schema()?;
103 store.seed_if_empty()?;
104 Ok(store)
105 }
106
107 fn seed_if_empty(&self) -> Result<()> {
109 if self.is_empty()? {
111 if let Ok(providers) = super::embedded::embedded_providers() {
112 tracing::info!(
113 providers = providers.len(),
114 "seeding registry cache from embedded snapshot"
115 );
116 self.replace_all(&providers)?;
117 }
118 if let Ok(manifest) = super::embedded::embedded_manifest() {
119 let _ = self.save_manifest_meta(&manifest);
120 let manifest_json = serde_json::to_string(&manifest).ok();
123 if let Some(json) = manifest_json {
124 let _ = self.meta_set("registry_manifest_json", &json);
125 }
126 if let Ok(catalog) = super::embedded::embedded_model_catalog() {
129 for (id, model) in catalog {
130 let sha = manifest.models.get(&id).map(|e| e.sha256.as_str());
131 let _ = self.upsert_canonical_model(&id, &model, sha);
132 }
133 }
134 }
135 }
136
137 self.seed_transcription_from_embedded_if_empty()?;
140 self.seed_canonical_models_from_embedded_if_empty()?;
141 Ok(())
142 }
143
144 #[cfg(test)]
146 pub fn open_memory() -> Result<Self> {
147 let conn = Connection::open_in_memory()?;
148 conn.pragma_update(None, "foreign_keys", "ON")?;
149 let store = Self {
150 conn: Mutex::new(conn),
151 };
152 store.init_schema()?;
153 Ok(store)
154 }
155
156 fn init_schema(&self) -> Result<()> {
157 let conn = self.conn.lock().unwrap_or_else(|e| e.into_inner());
158 conn.execute_batch(
159 "
160 CREATE TABLE IF NOT EXISTS registry_meta (
161 key TEXT PRIMARY KEY,
162 value TEXT NOT NULL
163 );
164
165 CREATE TABLE IF NOT EXISTS providers (
166 id TEXT PRIMARY KEY,
167 label TEXT NOT NULL,
168 description TEXT NOT NULL DEFAULT '',
169 kind TEXT NOT NULL,
170 api_key_env TEXT NOT NULL,
171 base_url TEXT,
172 tool_calling_mode TEXT,
173 request_options TEXT NOT NULL DEFAULT '{}',
174 sha256 TEXT,
175 aggregator INTEGER NOT NULL DEFAULT 0,
176 updated_at TEXT NOT NULL DEFAULT (datetime('now'))
177 );
178
179 CREATE TABLE IF NOT EXISTS models (
180 provider_id TEXT NOT NULL,
181 name TEXT NOT NULL,
182 task_size TEXT,
183 context_window_tokens INTEGER,
184 max_output_tokens INTEGER,
185 recommended_temperature REAL,
186 supports_thinking INTEGER,
187 supports_images INTEGER,
188 supports_audio INTEGER,
189 supports_video INTEGER,
190 supports_documents INTEGER,
191 tool_prompt_manifest INTEGER,
192 reasoning_levels TEXT NOT NULL DEFAULT '[]',
193 default_reasoning_effort TEXT,
194 PRIMARY KEY (provider_id, name),
195 FOREIGN KEY (provider_id) REFERENCES providers(id) ON DELETE CASCADE
196 );
197
198 CREATE TABLE IF NOT EXISTS model_capabilities (
199 model_id TEXT NOT NULL,
200 provider_id TEXT NOT NULL,
201 capability TEXT NOT NULL,
202 value TEXT NOT NULL,
203 PRIMARY KEY (model_id, capability),
204 FOREIGN KEY (provider_id) REFERENCES providers(id) ON DELETE CASCADE
205 );
206
207 CREATE TABLE IF NOT EXISTS model_pricing (
208 model_id TEXT PRIMARY KEY,
209 provider_id TEXT NOT NULL,
210 input_price REAL,
211 output_price REAL,
212 currency TEXT NOT NULL DEFAULT 'USD',
213 FOREIGN KEY (provider_id) REFERENCES providers(id) ON DELETE CASCADE
214 );
215
216 CREATE TABLE IF NOT EXISTS model_profiles (
217 model_id TEXT NOT NULL,
218 provider_id TEXT NOT NULL,
219 profile_id TEXT NOT NULL,
220 score REAL NOT NULL DEFAULT 0.0,
221 PRIMARY KEY (model_id, profile_id),
222 FOREIGN KEY (provider_id) REFERENCES providers(id) ON DELETE CASCADE
223 );
224
225 CREATE TABLE IF NOT EXISTS profiles (
226 id TEXT PRIMARY KEY,
227 description TEXT NOT NULL DEFAULT '',
228 min_context INTEGER,
229 max_input_price REAL,
230 requires_tools INTEGER NOT NULL DEFAULT 0
231 );
232
233 -- Remote speech-to-text / dictation providers (JSON blob + integrity hash).
234 CREATE TABLE IF NOT EXISTS transcription_providers (
235 id TEXT PRIMARY KEY,
236 json TEXT NOT NULL,
237 sha256 TEXT,
238 updated_at TEXT NOT NULL DEFAULT (datetime('now'))
239 );
240
241 -- Canonical model catalog (models/<id>.json), used for ref resolution.
242 CREATE TABLE IF NOT EXISTS canonical_models (
243 id TEXT PRIMARY KEY,
244 json TEXT NOT NULL,
245 sha256 TEXT,
246 updated_at TEXT NOT NULL DEFAULT (datetime('now'))
247 );
248 ",
249 )?;
250 ensure_provider_request_options_column(&conn)?;
251 ensure_model_output_columns(&conn)?;
252 ensure_provider_tool_calling_mode_column(&conn)?;
253 ensure_provider_sha256_column(&conn)?;
254 ensure_provider_aggregator_column(&conn)?;
255 relax_models_task_size_not_null(&conn)?;
256 Ok(())
257 }
258
259 pub fn meta_get(&self, key: &str) -> Result<Option<String>> {
263 let conn = self.conn.lock().unwrap_or_else(|e| e.into_inner());
264 let mut stmt = conn
265 .prepare("SELECT value FROM registry_meta WHERE key = ?1")
266 .context("prepare meta_get")?;
267 let mut rows = stmt.query_map(params![key], |row| row.get(0))?;
268 match rows.next() {
269 Some(Ok(v)) => Ok(Some(v)),
270 _ => Ok(None),
271 }
272 }
273
274 pub fn meta_set(&self, key: &str, value: &str) -> Result<()> {
276 let conn = self.conn.lock().unwrap_or_else(|e| e.into_inner());
277 conn.execute(
278 "INSERT OR REPLACE INTO registry_meta (key, value) VALUES (?1, ?2)",
279 params![key, value],
280 )?;
281 Ok(())
282 }
283
284 pub fn is_empty(&self) -> Result<bool> {
288 let conn = self.conn.lock().unwrap_or_else(|e| e.into_inner());
289 let count: i64 = conn.query_row("SELECT COUNT(*) FROM providers", [], |row| row.get(0))?;
290 Ok(count == 0)
291 }
292
293 pub fn provider_count(&self) -> Result<usize> {
295 let conn = self.conn.lock().unwrap_or_else(|e| e.into_inner());
296 let count: i64 = conn.query_row("SELECT COUNT(*) FROM providers", [], |row| row.get(0))?;
297 Ok(count as usize)
298 }
299
300 pub fn model_count(&self) -> Result<usize> {
302 let conn = self.conn.lock().unwrap_or_else(|e| e.into_inner());
303 let count: i64 = conn.query_row("SELECT COUNT(*) FROM models", [], |row| row.get(0))?;
304 Ok(count as usize)
305 }
306
307 pub fn provider_sha256(&self, provider_id: &str) -> Result<Option<String>> {
309 let conn = self.conn.lock().unwrap_or_else(|e| e.into_inner());
310 let mut stmt = conn.prepare("SELECT sha256 FROM providers WHERE id = ?1")?;
311 let mut rows =
312 stmt.query_map(params![provider_id], |row| row.get::<_, Option<String>>(0))?;
313 match rows.next() {
314 Some(Ok(v)) => Ok(v),
315 _ => Ok(None),
316 }
317 }
318
319 pub fn provider_ids(&self) -> Result<std::collections::HashSet<String>> {
321 let conn = self.conn.lock().unwrap_or_else(|e| e.into_inner());
322 let mut stmt = conn.prepare("SELECT id FROM providers")?;
323 let rows = stmt.query_map([], |row| row.get::<_, String>(0))?;
324 let mut ids = std::collections::HashSet::new();
325 for row in rows {
326 ids.insert(row?);
327 }
328 Ok(ids)
329 }
330
331 pub fn load_provider_models(
335 &self,
336 provider_id: &str,
337 ) -> Result<std::collections::HashMap<String, RegistryModel>> {
338 let conn = self.conn.lock().unwrap_or_else(|e| e.into_inner());
339 let mut stmt = conn.prepare(
340 "SELECT name, task_size, context_window_tokens, max_output_tokens, recommended_temperature, supports_thinking, supports_images, supports_audio, supports_video, supports_documents, reasoning_levels, default_reasoning_effort
341 FROM models WHERE provider_id = ?1",
342 )?;
343 let rows = stmt.query_map(params![provider_id], |row| {
344 let name: String = row.get(0)?;
345 let task_size_str: Option<String> = row.get(1)?;
346 let ctx: Option<i64> = row.get(2)?;
347 let max_out: Option<i64> = row.get(3)?;
348 let temp: Option<f64> = row.get(4)?;
349 let thinking: Option<i64> = row.get(5)?;
350 let images: Option<i64> = row.get(6)?;
351 let audio: Option<i64> = row.get(7)?;
352 let video: Option<i64> = row.get(8)?;
353 let documents: Option<i64> = row.get(9)?;
354 let levels_json: Option<String> = row.get(10)?;
355 let default_effort: Option<String> = row.get(11)?;
356
357 Ok(RegistryModel {
358 model_ref: None,
359 api_name: None,
360 name: name.clone(),
361 task_size: task_size_str,
362 context_window_tokens: ctx.map(|v| v as u64),
363 max_output_tokens: max_out.map(|v| v as u64),
364 recommended_temperature: temp,
365 supports_thinking: thinking.map(|v| v != 0),
366 reasoning_levels: parse_reasoning_levels_json(levels_json.as_deref()),
367 default_reasoning_effort: default_effort,
368 supports_images: images.map(|v| v != 0),
369 supports_audio: audio.map(|v| v != 0),
370 supports_video: video.map(|v| v != 0),
371 supports_documents: documents.map(|v| v != 0),
372 supports_attachments: None,
373 attachments: RegistryAttachments::default(),
374 capabilities: Vec::new(),
375 pricing: None,
376 })
377 })?;
378
379 let mut map = std::collections::HashMap::new();
380 for row in rows {
381 let model = row?;
382 map.insert(model.name.clone(), model);
383 }
384 Ok(map)
385 }
386
387 pub fn delete_providers_not_in(&self, keep: &std::collections::HashSet<&str>) -> Result<()> {
390 let conn = self.conn.lock().unwrap_or_else(|e| e.into_inner());
391 let mut stmt = conn.prepare("SELECT id FROM providers")?;
392 let to_delete: Vec<String> = stmt
393 .query_map([], |row| row.get::<_, String>(0))?
394 .filter_map(|r| r.ok())
395 .filter(|id| !keep.contains(id.as_str()))
396 .collect();
397 drop(stmt);
398 for id in &to_delete {
399 conn.execute("DELETE FROM providers WHERE id = ?1", params![id])?;
400 }
401 if !to_delete.is_empty() {
402 tracing::info!(
403 removed = to_delete.len(),
404 "removed stale providers from cache"
405 );
406 }
407 Ok(())
408 }
409
410 pub fn upsert_provider(&self, provider: &RegistryProvider) -> Result<()> {
412 self.upsert_provider_with_sha256(provider, None)
413 }
414
415 pub fn upsert_transcription_provider(
419 &self,
420 provider: &RegistryTranscriptionProvider,
421 sha256: Option<&str>,
422 ) -> Result<()> {
423 let json = serde_json::to_string(provider)
424 .context("serialize transcription provider for cache")?;
425 let conn = self.conn.lock().unwrap_or_else(|e| e.into_inner());
426 conn.execute(
427 "INSERT OR REPLACE INTO transcription_providers (id, json, sha256, updated_at)
428 VALUES (?1, ?2, ?3, datetime('now'))",
429 params![provider.id, json, sha256],
430 )?;
431 Ok(())
432 }
433
434 pub fn transcription_provider_sha256(&self, id: &str) -> Result<Option<String>> {
436 let conn = self.conn.lock().unwrap_or_else(|e| e.into_inner());
437 let mut stmt = conn.prepare("SELECT sha256 FROM transcription_providers WHERE id = ?1")?;
438 let mut rows = stmt.query_map(params![id], |row| row.get::<_, Option<String>>(0))?;
439 match rows.next() {
440 Some(Ok(v)) => Ok(v),
441 _ => Ok(None),
442 }
443 }
444
445 pub fn load_transcription_providers(&self) -> Result<Vec<RegistryTranscriptionProvider>> {
447 let conn = self.conn.lock().unwrap_or_else(|e| e.into_inner());
448 let mut stmt = conn.prepare("SELECT json FROM transcription_providers ORDER BY id")?;
449 let rows = stmt.query_map([], |row| row.get::<_, String>(0))?;
450 let mut out = Vec::new();
451 for row in rows {
452 let json = row?;
453 match serde_json::from_str::<RegistryTranscriptionProvider>(&json) {
454 Ok(p) => out.push(p),
455 Err(err) => {
456 tracing::warn!(error = %err, "skip corrupt transcription provider cache row");
457 }
458 }
459 }
460 Ok(out)
461 }
462
463 pub fn delete_transcription_providers_not_in(
465 &self,
466 keep: &std::collections::HashSet<&str>,
467 ) -> Result<()> {
468 let conn = self.conn.lock().unwrap_or_else(|e| e.into_inner());
469 let mut stmt = conn.prepare("SELECT id FROM transcription_providers")?;
470 let to_delete: Vec<String> = stmt
471 .query_map([], |row| row.get::<_, String>(0))?
472 .filter_map(|r| r.ok())
473 .filter(|id| !keep.contains(id.as_str()))
474 .collect();
475 drop(stmt);
476 for id in &to_delete {
477 conn.execute(
478 "DELETE FROM transcription_providers WHERE id = ?1",
479 params![id],
480 )?;
481 }
482 Ok(())
483 }
484
485 pub fn upsert_canonical_model(
489 &self,
490 id: &str,
491 model: &super::types::CanonicalModel,
492 sha256: Option<&str>,
493 ) -> Result<()> {
494 let json = serde_json::to_string(model).context("serialize canonical model")?;
495 let conn = self.conn.lock().unwrap_or_else(|e| e.into_inner());
496 conn.execute(
497 "INSERT OR REPLACE INTO canonical_models (id, json, sha256, updated_at)
498 VALUES (?1, ?2, ?3, datetime('now'))",
499 params![id, json, sha256],
500 )?;
501 Ok(())
502 }
503
504 pub fn canonical_model_sha256(&self, id: &str) -> Result<Option<String>> {
506 let conn = self.conn.lock().unwrap_or_else(|e| e.into_inner());
507 let mut stmt = conn.prepare("SELECT sha256 FROM canonical_models WHERE id = ?1")?;
508 let mut rows = stmt.query_map(params![id], |row| row.get::<_, Option<String>>(0))?;
509 match rows.next() {
510 Some(Ok(v)) => Ok(v),
511 _ => Ok(None),
512 }
513 }
514
515 pub fn load_canonical_model_catalog(&self) -> Result<super::resolve::ModelCatalog> {
517 let conn = self.conn.lock().unwrap_or_else(|e| e.into_inner());
518 let mut stmt = conn.prepare("SELECT id, json FROM canonical_models ORDER BY id")?;
519 let rows = stmt.query_map([], |row| {
520 let id: String = row.get(0)?;
521 let json: String = row.get(1)?;
522 Ok((id, json))
523 })?;
524 let mut catalog = std::collections::HashMap::new();
525 for row in rows {
526 let (id, json) = row?;
527 match serde_json::from_str::<super::types::CanonicalModel>(&json) {
528 Ok(model) => {
529 catalog.insert(id, model);
530 }
531 Err(err) => {
532 tracing::warn!(id = %id, error = %err, "skipping corrupt canonical model row");
533 }
534 }
535 }
536 Ok(catalog)
537 }
538
539 pub fn delete_canonical_models_not_in(
541 &self,
542 keep: &std::collections::HashSet<&str>,
543 ) -> Result<()> {
544 let conn = self.conn.lock().unwrap_or_else(|e| e.into_inner());
545 let mut stmt = conn.prepare("SELECT id FROM canonical_models")?;
546 let to_delete: Vec<String> = stmt
547 .query_map([], |row| row.get::<_, String>(0))?
548 .filter_map(|r| r.ok())
549 .filter(|id| !keep.contains(id.as_str()))
550 .collect();
551 for id in to_delete {
552 conn.execute("DELETE FROM canonical_models WHERE id = ?1", params![id])?;
553 }
554 Ok(())
555 }
556
557 pub fn canonical_model_count(&self) -> Result<usize> {
559 let conn = self.conn.lock().unwrap_or_else(|e| e.into_inner());
560 let count: i64 = conn.query_row("SELECT COUNT(*) FROM canonical_models", [], |row| {
561 row.get(0)
562 })?;
563 Ok(count as usize)
564 }
565
566 pub fn transcription_provider_count(&self) -> Result<usize> {
568 let conn = self.conn.lock().unwrap_or_else(|e| e.into_inner());
569 let count: i64 =
570 conn.query_row("SELECT COUNT(*) FROM transcription_providers", [], |row| {
571 row.get(0)
572 })?;
573 Ok(count as usize)
574 }
575
576 fn seed_canonical_models_from_embedded_if_empty(&self) -> Result<()> {
578 if self.canonical_model_count().unwrap_or(0) > 0 {
579 return Ok(());
580 }
581 let catalog = match super::embedded::embedded_model_catalog() {
582 Ok(c) if !c.is_empty() => c,
583 _ => return Ok(()),
584 };
585 let manifest = super::embedded::embedded_manifest().ok();
586 for (id, model) in catalog {
587 let sha = manifest
588 .as_ref()
589 .and_then(|m| m.models.get(&id))
590 .map(|e| e.sha256.as_str());
591 self.upsert_canonical_model(&id, &model, sha)?;
592 }
593 Ok(())
594 }
595
596 pub fn seed_transcription_from_embedded_if_empty(&self) -> Result<()> {
598 if self.transcription_provider_count()? > 0 {
599 return Ok(());
600 }
601 let providers = match super::embedded::embedded_transcription_providers() {
602 Ok(p) if !p.is_empty() => p,
603 _ => return Ok(()),
604 };
605 let manifest = super::embedded::embedded_manifest().ok();
606 for p in &providers {
607 let sha = manifest
608 .as_ref()
609 .and_then(|m| m.transcription_providers.get(&p.id))
610 .map(|e| e.sha256.as_str());
611 self.upsert_transcription_provider(p, sha)?;
612 }
613 tracing::info!(
614 providers = providers.len(),
615 "seeded transcription providers from embedded snapshot"
616 );
617 Ok(())
618 }
619
620 pub fn upsert_provider_union_models(
628 &self,
629 provider: &RegistryProvider,
630 sha256: Option<&str>,
631 ) -> Result<()> {
632 let existing = self.load_provider_models(&provider.id).unwrap_or_default();
633 if existing.is_empty() {
634 return self.upsert_provider_with_sha256(provider, sha256);
635 }
636
637 let mut models = provider.models.clone();
638 let incoming: std::collections::HashSet<String> =
639 models.iter().map(|m| m.name.to_ascii_lowercase()).collect();
640 for (name, model) in existing {
641 if !incoming.contains(&name.to_ascii_lowercase()) {
642 models.push(model);
643 }
644 }
645
646 let mut merged = provider.clone();
647 merged.models = models;
648 self.upsert_provider_with_sha256(&merged, sha256)
649 }
650
651 pub fn rehydrate_provider_models_from_catalog(&self) -> Result<usize> {
659 let catalog = self.load_canonical_model_catalog().unwrap_or_default();
660 if catalog.is_empty() {
661 return Ok(0);
662 }
663
664 let mut updated = 0usize;
665 let conn = self.conn.lock().unwrap_or_else(|e| e.into_inner());
666 let mut stmt = conn.prepare(
667 "UPDATE models SET
668 context_window_tokens = COALESCE(?1, context_window_tokens),
669 max_output_tokens = COALESCE(?2, max_output_tokens),
670 recommended_temperature = COALESCE(?3, recommended_temperature),
671 supports_thinking = COALESCE(?4, supports_thinking),
672 reasoning_levels = CASE
673 WHEN ?5 = '[]' THEN reasoning_levels
674 ELSE ?5
675 END,
676 default_reasoning_effort = COALESCE(?6, default_reasoning_effort)
677 WHERE lower(name) = lower(?7)",
678 )?;
679
680 for (id, canonical) in &catalog {
681 let levels_json =
682 serde_json::to_string(&canonical.reasoning_levels).unwrap_or_else(|_| "[]".into());
683 let mut names = vec![id.clone()];
684 names.extend(canonical.aliases.iter().cloned());
685 names.sort();
687 names.dedup();
688 for name in names {
689 updated += stmt.execute(params![
690 canonical.context_window_tokens.map(|v| v as i64),
691 canonical.max_output_tokens.map(|v| v as i64),
692 canonical.recommended_temperature,
693 canonical.supports_thinking.map(|v| v as i64),
694 levels_json,
695 canonical.default_reasoning_effort,
696 name,
697 ])?;
698 }
699 }
700
701 if updated > 0 {
702 tracing::info!(
703 models_updated = updated,
704 "rehydrated provider model metadata from canonical catalog"
705 );
706 }
707 Ok(updated)
708 }
709
710 pub fn upsert_provider_with_sha256(
712 &self,
713 provider: &RegistryProvider,
714 sha256: Option<&str>,
715 ) -> Result<()> {
716 let kind = parse_provider_kind(&provider.kind);
717
718 let conn = self.conn.lock().unwrap_or_else(|e| e.into_inner());
719 let tx = conn.unchecked_transaction()?;
720
721 tx.execute(
722 "INSERT OR REPLACE INTO providers (id, label, description, kind, api_key_env, base_url, tool_calling_mode, request_options, sha256, aggregator, updated_at)
723 VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, datetime('now'))",
724 params![
725 provider.id,
726 provider.label,
727 provider.description,
728 provider.kind,
729 provider.api_key_env,
730 provider.base_url,
731 provider.tool_calling_mode,
732 serde_json::to_string(&provider.request_options)?,
733 sha256,
734 provider.aggregator as i64,
735 ],
736 )?;
737
738 tx.execute(
740 "DELETE FROM models WHERE provider_id = ?1",
741 params![provider.id],
742 )?;
743
744 {
745 let mut stmt = tx.prepare(
746 "INSERT INTO models (provider_id, name, task_size, context_window_tokens, max_output_tokens, recommended_temperature, supports_thinking, supports_images, supports_audio, supports_video, supports_documents, tool_prompt_manifest, reasoning_levels, default_reasoning_effort)
747 VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, NULL, ?12, ?13)",
748 )?;
749
750 let attachment_defaults = &provider.defaults.attachments;
751 for model in &provider.models {
752 let levels_json =
753 serde_json::to_string(&model.reasoning_levels).unwrap_or_else(|_| "[]".into());
754 stmt.execute(params![
755 provider.id,
756 model.name,
757 model.task_size,
758 model.context_window_tokens.map(|v| v as i64),
759 model.max_output_tokens.map(|v| v as i64),
760 model.recommended_temperature,
761 model.supports_thinking.map(|v| v as i64),
762 registry_model_supports_images(model, attachment_defaults).map(|v| v as i64),
763 registry_model_supports_audio(model, attachment_defaults).map(|v| v as i64),
764 registry_model_supports_video(model, attachment_defaults).map(|v| v as i64),
765 registry_model_supports_documents(model, attachment_defaults).map(|v| v as i64),
766 levels_json,
767 model.default_reasoning_effort,
768 ])?;
769 }
770 }
771
772 tx.execute(
774 "DELETE FROM model_pricing WHERE provider_id = ?1",
775 params![provider.id],
776 )?;
777 {
778 let mut price_stmt = tx.prepare(
779 "INSERT OR REPLACE INTO model_pricing (model_id, provider_id, input_price, output_price, currency)
780 VALUES (?1, ?2, ?3, ?4, ?5)",
781 )?;
782 for model in &provider.models {
783 let Some(pricing) = model.pricing.as_ref() else {
784 continue;
785 };
786 if pricing.is_empty() {
787 continue;
788 }
789 let model_id = format!("{}:{}", provider.id, model.name);
790 price_stmt.execute(params![
791 model_id,
792 provider.id,
793 pricing.input_per_1m,
794 pricing.output_per_1m,
795 pricing.currency.as_deref().unwrap_or("USD"),
796 ])?;
797 }
798 }
799
800 tx.commit()?;
801 let _ = kind; Ok(())
803 }
804
805 pub fn load_all_providers(&self) -> Result<Vec<ProviderConfig>> {
808 let conn = self.conn.lock().unwrap_or_else(|e| e.into_inner());
809
810 let mut stmt = conn.prepare(
811 "SELECT id, label, description, kind, api_key_env, base_url, tool_calling_mode, request_options, aggregator FROM providers ORDER BY id",
812 )?;
813
814 let provider_rows = stmt.query_map([], |row| {
815 Ok((
816 row.get::<_, String>(0)?,
817 row.get::<_, String>(1)?,
818 row.get::<_, String>(2)?,
819 row.get::<_, String>(3)?,
820 row.get::<_, String>(4)?,
821 row.get::<_, Option<String>>(5)?,
822 row.get::<_, Option<String>>(6)?,
823 row.get::<_, String>(7)?,
824 row.get::<_, Option<i64>>(8)?,
825 ))
826 })?;
827
828 let mut providers = Vec::new();
829
830 for row in provider_rows {
831 let (
832 id,
833 label,
834 description,
835 kind_str,
836 api_key_env,
837 base_url,
838 tool_calling_mode_str,
839 request_options_json,
840 aggregator_val,
841 ) = row?;
842 let kind = parse_provider_kind(&kind_str);
843 let request_options = serde_json::from_str(&request_options_json).ok();
844 let tool_calling_mode = tool_calling_mode_str
845 .as_deref()
846 .map(parse_tool_calling_mode);
847 let aggregator = aggregator_val.unwrap_or(0) != 0;
848
849 let mut model_stmt = conn.prepare(
850 "SELECT m.name, m.task_size, m.context_window_tokens, m.max_output_tokens,
851 m.recommended_temperature, m.supports_thinking, m.supports_images,
852 m.supports_audio, m.supports_video, m.supports_documents,
853 m.tool_prompt_manifest, pr.input_price, pr.output_price,
854 m.reasoning_levels, m.default_reasoning_effort
855 FROM models m
856 LEFT JOIN model_pricing pr
857 ON pr.model_id = (m.provider_id || ':' || m.name)
858 WHERE m.provider_id = ?1
859 ORDER BY m.rowid",
860 )?;
861
862 let models = model_stmt
863 .query_map(params![id], |row| {
864 let name: String = row.get(0)?;
865 let task_size_str: Option<String> = row.get(1)?;
866 let ctx: Option<i64> = row.get(2)?;
867 let max_out: Option<i64> = row.get(3)?;
868 let temp: Option<f64> = row.get(4)?;
869 let thinking: Option<i64> = row.get(5)?;
870 let images: Option<i64> = row.get(6)?;
871 let audio: Option<i64> = row.get(7)?;
872 let video: Option<i64> = row.get(8)?;
873 let documents: Option<i64> = row.get(9)?;
874 let tpm: Option<i64> = row.get(10)?;
875 let input_price: Option<f64> = row.get(11)?;
876 let output_price: Option<f64> = row.get(12)?;
877 let levels_json: Option<String> = row.get(13)?;
878 let default_effort: Option<String> = row.get(14)?;
879
880 Ok(ProviderModelConfig {
881 name,
882 task_size: task_size_str.as_deref().and_then(|s| match s {
883 "small" => Some(ModelTaskSize::Small),
884 "large" => Some(ModelTaskSize::Large),
885 _ => None,
886 }),
887 context_window_tokens: ctx.map(|v| v as u64),
888 max_output_tokens: max_out.map(|v| v as u64),
889 recommended_temperature: temp,
890 supports_thinking: thinking.map(|v| v != 0),
891 reasoning_levels: parse_reasoning_levels_json(levels_json.as_deref()),
892 default_reasoning_effort: default_effort,
893 supports_images: images.map(|v| v != 0),
894 supports_audio: audio.map(|v| v != 0),
895 supports_video: video.map(|v| v != 0),
896 supports_documents: documents.map(|v| v != 0),
897 tool_prompt_manifest: tpm.map(|v| v != 0),
898 pricing_input_per_1m: input_price,
899 pricing_output_per_1m: output_price,
900 })
901 })?
902 .collect::<std::result::Result<Vec<_>, _>>()?;
903
904 providers.push(ProviderConfig {
905 id,
906 label,
907 description,
908 kind,
909 api_key_env,
910 base_url,
911 models,
912 request_options,
913 tool_calling_mode,
914 aggregator,
915 ..Default::default()
916 });
917 }
918
919 Ok(providers)
920 }
921
922 pub fn replace_all(&self, providers: &[RegistryProvider]) -> Result<()> {
924 let conn = self.conn.lock().unwrap_or_else(|e| e.into_inner());
925
926 conn.execute("DELETE FROM models", [])?;
928 conn.execute("DELETE FROM providers", [])?;
929
930 drop(conn); for provider in providers {
933 self.upsert_provider(provider)?;
934 }
935
936 Ok(())
937 }
938
939 pub fn save_manifest_meta(&self, manifest: &RegistryManifest) -> Result<()> {
941 self.meta_set("manifest_version", &manifest.version.to_string())?;
942 self.meta_set("manifest_updated_at", &manifest.updated_at)?;
943 self.meta_set(
944 "manifest_provider_count",
945 &manifest.providers.len().to_string(),
946 )?;
947 Ok(())
948 }
949
950 pub fn manifest_version(&self) -> Result<Option<u32>> {
952 match self.meta_get("manifest_version")? {
953 Some(v) => Ok(v.parse().ok()),
954 None => Ok(None),
955 }
956 }
957
958 pub fn manifest_updated_at(&self) -> Result<Option<String>> {
960 self.meta_get("manifest_updated_at")
961 }
962
963 pub fn upsert_capabilities(
967 &self,
968 model_id: &str,
969 provider_id: &str,
970 capabilities: &[(String, String)],
971 ) -> Result<()> {
972 let conn = self.conn.lock().unwrap_or_else(|e| e.into_inner());
973 let tx = conn.unchecked_transaction()?;
974 tx.execute(
975 "DELETE FROM model_capabilities WHERE model_id = ?1",
976 params![model_id],
977 )?;
978 {
979 let mut stmt = tx.prepare(
980 "INSERT INTO model_capabilities (model_id, provider_id, capability, value)
981 VALUES (?1, ?2, ?3, ?4)",
982 )?;
983 for (cap, value) in capabilities {
984 stmt.execute(params![model_id, provider_id, cap, value])?;
985 }
986 }
987 tx.commit()?;
988 Ok(())
989 }
990
991 pub fn load_capabilities(&self, model_id: &str) -> Result<Vec<ModelCapability>> {
993 let conn = self.conn.lock().unwrap_or_else(|e| e.into_inner());
994 let mut stmt = conn.prepare(
995 "SELECT model_id, provider_id, capability, value
996 FROM model_capabilities WHERE model_id = ?1",
997 )?;
998 let rows = stmt
999 .query_map(params![model_id], |row| {
1000 Ok(ModelCapability {
1001 model_id: row.get(0)?,
1002 provider_id: row.get(1)?,
1003 capability: row.get(2)?,
1004 value: row.get(3)?,
1005 })
1006 })?
1007 .collect::<std::result::Result<Vec<_>, _>>()?;
1008 Ok(rows)
1009 }
1010
1011 pub fn upsert_pricing(
1015 &self,
1016 model_id: &str,
1017 provider_id: &str,
1018 input_price: Option<f64>,
1019 output_price: Option<f64>,
1020 ) -> Result<()> {
1021 let conn = self.conn.lock().unwrap_or_else(|e| e.into_inner());
1022 conn.execute(
1023 "INSERT OR REPLACE INTO model_pricing (model_id, provider_id, input_price, output_price)
1024 VALUES (?1, ?2, ?3, ?4)",
1025 params![model_id, provider_id, input_price, output_price],
1026 )?;
1027 Ok(())
1028 }
1029
1030 pub fn load_pricing(&self, model_id: &str) -> Result<Option<ModelPricing>> {
1032 let conn = self.conn.lock().unwrap_or_else(|e| e.into_inner());
1033 let mut stmt = conn.prepare(
1034 "SELECT model_id, provider_id, input_price, output_price, currency
1035 FROM model_pricing WHERE model_id = ?1",
1036 )?;
1037 let mut rows = stmt.query_map(params![model_id], |row| {
1038 Ok(ModelPricing {
1039 model_id: row.get(0)?,
1040 provider_id: row.get(1)?,
1041 input_price: row.get(2)?,
1042 output_price: row.get(3)?,
1043 currency: row.get(4)?,
1044 })
1045 })?;
1046 match rows.next() {
1047 Some(Ok(p)) => Ok(Some(p)),
1048 _ => Ok(None),
1049 }
1050 }
1051
1052 pub fn upsert_profile(&self, profile: &Profile) -> Result<()> {
1056 let conn = self.conn.lock().unwrap_or_else(|e| e.into_inner());
1057 conn.execute(
1058 "INSERT OR REPLACE INTO profiles (id, description, min_context, max_input_price, requires_tools)
1059 VALUES (?1, ?2, ?3, ?4, ?5)",
1060 params![
1061 profile.id,
1062 profile.description,
1063 profile.min_context.map(|v| v as i64),
1064 profile.max_input_price,
1065 profile.requires_tools as i64,
1066 ],
1067 )?;
1068 Ok(())
1069 }
1070
1071 pub fn upsert_model_profile(
1073 &self,
1074 model_id: &str,
1075 provider_id: &str,
1076 profile_id: &str,
1077 score: f64,
1078 ) -> Result<()> {
1079 let conn = self.conn.lock().unwrap_or_else(|e| e.into_inner());
1080 conn.execute(
1081 "INSERT OR REPLACE INTO model_profiles (model_id, provider_id, profile_id, score)
1082 VALUES (?1, ?2, ?3, ?4)",
1083 params![model_id, provider_id, profile_id, score],
1084 )?;
1085 Ok(())
1086 }
1087
1088 pub fn query_models_by_profile(&self, profile_id: &str) -> Result<Vec<RankedModel>> {
1093 let conn = self.conn.lock().unwrap_or_else(|e| e.into_inner());
1094 let mut stmt = conn.prepare(
1095 "SELECT
1096 mp.model_id,
1097 mp.provider_id,
1098 m.name,
1099 mp.score,
1100 pr.input_price,
1101 pr.output_price,
1102 m.context_window_tokens
1103 FROM model_profiles mp
1104 JOIN models m ON m.provider_id = mp.provider_id AND m.name = (
1105 SELECT SUBSTR(mp.model_id, INSTR(mp.model_id, ':') + 1)
1106 )
1107 LEFT JOIN model_pricing pr ON pr.model_id = mp.model_id
1108 LEFT JOIN profiles p ON p.id = mp.profile_id
1109 WHERE mp.profile_id = ?1
1110 AND (p.min_context IS NULL OR m.context_window_tokens >= p.min_context)
1111 AND (p.max_input_price IS NULL OR pr.input_price IS NULL OR pr.input_price <= p.max_input_price)
1112 AND (p.requires_tools = 0 OR m.supports_thinking IS NOT NULL)
1113 ORDER BY mp.score DESC, pr.input_price ASC, pr.output_price ASC",
1114 )?;
1115 let rows = stmt
1116 .query_map(params![profile_id], |row| {
1117 Ok(RankedModel {
1118 model_id: row.get(0)?,
1119 provider_id: row.get(1)?,
1120 model_name: row.get(2)?,
1121 score: row.get(3)?,
1122 input_price: row.get(4)?,
1123 output_price: row.get(5)?,
1124 context_window_tokens: row.get::<_, Option<i64>>(6)?.map(|v| v as u64),
1125 })
1126 })?
1127 .collect::<std::result::Result<Vec<_>, _>>()?;
1128 Ok(rows)
1129 }
1130
1131 pub fn seed_default_profiles(&self) -> Result<()> {
1133 let defaults = vec![
1134 Profile {
1135 id: "cheap_general".to_string(),
1136 description: "General-purpose cheap model".to_string(),
1137 min_context: Some(32_000),
1138 max_input_price: Some(0.50),
1139 requires_tools: false,
1140 },
1141 Profile {
1142 id: "cheap_code".to_string(),
1143 description: "Cheap code-focused model with tool support".to_string(),
1144 min_context: Some(64_000),
1145 max_input_price: Some(1.00),
1146 requires_tools: true,
1147 },
1148 Profile {
1149 id: "repo_search".to_string(),
1150 description: "Fast repository exploration".to_string(),
1151 min_context: Some(64_000),
1152 max_input_price: Some(0.50),
1153 requires_tools: true,
1154 },
1155 Profile {
1156 id: "naming".to_string(),
1157 description: "Session title generation".to_string(),
1158 min_context: Some(8_000),
1159 max_input_price: Some(0.20),
1160 requires_tools: false,
1161 },
1162 Profile {
1163 id: "long_context_cheap".to_string(),
1164 description: "Compaction and summarization".to_string(),
1165 min_context: Some(128_000),
1166 max_input_price: Some(1.00),
1167 requires_tools: false,
1168 },
1169 Profile {
1170 id: "research_synthesis".to_string(),
1171 description: "Research subagent with tool access".to_string(),
1172 min_context: Some(64_000),
1173 max_input_price: Some(1.00),
1174 requires_tools: true,
1175 },
1176 ];
1177 for profile in &defaults {
1178 self.upsert_profile(profile)?;
1179 }
1180 Ok(())
1181 }
1182
1183 pub fn delete_provider_metadata(&self, provider_id: &str) -> Result<()> {
1185 let conn = self.conn.lock().unwrap_or_else(|e| e.into_inner());
1186 conn.execute(
1187 "DELETE FROM model_capabilities WHERE provider_id = ?1",
1188 params![provider_id],
1189 )?;
1190 conn.execute(
1191 "DELETE FROM model_pricing WHERE provider_id = ?1",
1192 params![provider_id],
1193 )?;
1194 conn.execute(
1195 "DELETE FROM model_profiles WHERE provider_id = ?1",
1196 params![provider_id],
1197 )?;
1198 Ok(())
1199 }
1200}
1201
1202fn parse_provider_kind(s: &str) -> ProviderKind {
1203 match s {
1204 "openai-responses" => ProviderKind::OpenAiResponses,
1205 "openai-chat-completions" => ProviderKind::OpenAiChatCompletions,
1206 "anthropic-messages" => ProviderKind::AnthropicMessages,
1207 "gemini-generate-content" => ProviderKind::GeminiGenerateContent,
1208 _ => ProviderKind::OpenAiChatCompletions,
1209 }
1210}
1211
1212fn parse_tool_calling_mode(s: &str) -> ToolCallingMode {
1213 match s {
1214 "native" => ToolCallingMode::Native,
1215 "text-extracted" => ToolCallingMode::TextExtracted,
1216 "manifest-only" => ToolCallingMode::ManifestOnly,
1217 "disabled" => ToolCallingMode::Disabled,
1218 _ => ToolCallingMode::Native,
1219 }
1220}
1221
1222pub fn registry_provider_to_config(rp: RegistryProvider) -> ProviderConfig {
1227 let kind = parse_provider_kind(&rp.kind);
1228 let tool_calling_mode = rp.tool_calling_mode.as_deref().map(parse_tool_calling_mode);
1229 let attachment_defaults = rp.defaults.attachments;
1230
1231 let models = rp
1232 .models
1233 .into_iter()
1234 .map(|m| {
1235 let supports_images = registry_model_supports_images(&m, &attachment_defaults);
1236 let supports_audio = registry_model_supports_audio(&m, &attachment_defaults);
1237 let supports_video = registry_model_supports_video(&m, &attachment_defaults);
1238 let supports_documents = registry_model_supports_documents(&m, &attachment_defaults);
1239 ProviderModelConfig {
1240 name: m.name,
1241 task_size: m.task_size.as_deref().and_then(|s| match s {
1242 "small" => Some(ModelTaskSize::Small),
1243 "large" => Some(ModelTaskSize::Large),
1244 _ => None,
1245 }),
1246 context_window_tokens: m.context_window_tokens,
1247 max_output_tokens: m.max_output_tokens,
1248 recommended_temperature: m.recommended_temperature,
1249 supports_thinking: m.supports_thinking,
1250 reasoning_levels: m.reasoning_levels,
1251 default_reasoning_effort: m.default_reasoning_effort,
1252 supports_images,
1253 supports_audio,
1254 supports_video,
1255 supports_documents,
1256 tool_prompt_manifest: None,
1257 pricing_input_per_1m: m.pricing.as_ref().and_then(|p| p.input_per_1m),
1258 pricing_output_per_1m: m.pricing.as_ref().and_then(|p| p.output_per_1m),
1259 }
1260 })
1261 .collect();
1262
1263 ProviderConfig {
1264 id: rp.id,
1265 label: rp.label,
1266 description: rp.description,
1267 kind,
1268 api_key_env: rp.api_key_env,
1269 base_url: rp.base_url,
1270 models,
1271 tool_calling_mode,
1272 request_options: if rp.request_options.is_empty() {
1273 None
1274 } else {
1275 Some(rp.request_options)
1276 },
1277 aggregator: rp.aggregator,
1278 ..Default::default()
1279 }
1280}
1281
1282fn registry_model_has_capability(model: &super::types::RegistryModel, names: &[&str]) -> bool {
1283 model.capabilities.iter().any(|capability| {
1284 let normalized = capability.trim().to_ascii_lowercase();
1285 names.iter().any(|name| normalized == *name)
1286 })
1287}
1288
1289fn registry_model_supports_images(
1290 model: &super::types::RegistryModel,
1291 defaults: &super::types::RegistryAttachments,
1292) -> Option<bool> {
1293 model
1294 .attachments
1295 .images
1296 .or(model.supports_images)
1297 .or_else(|| {
1298 (model.supports_attachments == Some(true)
1299 || registry_model_has_capability(model, &["image", "images", "vision"]))
1300 .then_some(true)
1301 })
1302 .or(defaults.images)
1303}
1304
1305fn registry_model_supports_audio(
1306 model: &super::types::RegistryModel,
1307 defaults: &super::types::RegistryAttachments,
1308) -> Option<bool> {
1309 model
1310 .attachments
1311 .audio
1312 .or(model.supports_audio)
1313 .or_else(|| {
1314 registry_model_has_capability(model, &["audio", "sound", "speech"]).then_some(true)
1315 })
1316 .or(defaults.audio)
1317}
1318
1319fn registry_model_supports_video(
1320 model: &super::types::RegistryModel,
1321 defaults: &super::types::RegistryAttachments,
1322) -> Option<bool> {
1323 model
1324 .attachments
1325 .video
1326 .or(model.supports_video)
1327 .or_else(|| registry_model_has_capability(model, &["video"]).then_some(true))
1328 .or(defaults.video)
1329}
1330
1331fn registry_model_supports_documents(
1332 model: &super::types::RegistryModel,
1333 defaults: &super::types::RegistryAttachments,
1334) -> Option<bool> {
1335 model
1336 .attachments
1337 .documents
1338 .or(model.supports_documents)
1339 .or_else(|| {
1340 (model.supports_attachments == Some(true)
1341 || registry_model_has_capability(
1342 model,
1343 &["document", "documents", "pdf", "file", "files"],
1344 ))
1345 .then_some(true)
1346 })
1347 .or(defaults.documents)
1348}
1349
1350fn ensure_provider_request_options_column(conn: &Connection) -> Result<()> {
1351 let mut stmt = conn.prepare("PRAGMA table_info(providers)")?;
1352 let has_column = stmt
1353 .query_map([], |row| row.get::<_, String>(1))?
1354 .any(|name| matches!(name, Ok(name) if name == "request_options"));
1355
1356 if !has_column {
1357 conn.execute(
1358 "ALTER TABLE providers ADD COLUMN request_options TEXT NOT NULL DEFAULT '{}'",
1359 [],
1360 )?;
1361 }
1362
1363 Ok(())
1364}
1365
1366fn parse_reasoning_levels_json(raw: Option<&str>) -> Vec<String> {
1367 let Some(raw) = raw.map(str::trim).filter(|s| !s.is_empty()) else {
1368 return Vec::new();
1369 };
1370 serde_json::from_str::<Vec<String>>(raw).unwrap_or_default()
1371}
1372
1373fn ensure_model_output_columns(conn: &Connection) -> Result<()> {
1374 let mut stmt = conn.prepare("PRAGMA table_info(models)")?;
1375 let columns: Vec<String> = stmt
1376 .query_map([], |row| row.get::<_, String>(1))?
1377 .filter_map(|r| r.ok())
1378 .collect();
1379
1380 if !columns.contains(&"max_output_tokens".to_string()) {
1381 conn.execute(
1382 "ALTER TABLE models ADD COLUMN max_output_tokens INTEGER",
1383 [],
1384 )?;
1385 }
1386 if !columns.contains(&"recommended_temperature".to_string()) {
1387 conn.execute(
1388 "ALTER TABLE models ADD COLUMN recommended_temperature REAL",
1389 [],
1390 )?;
1391 }
1392 if !columns.contains(&"supports_thinking".to_string()) {
1393 conn.execute(
1394 "ALTER TABLE models ADD COLUMN supports_thinking INTEGER",
1395 [],
1396 )?;
1397 }
1398 if !columns.contains(&"supports_images".to_string()) {
1399 conn.execute("ALTER TABLE models ADD COLUMN supports_images INTEGER", [])?;
1400 }
1401 if !columns.contains(&"supports_audio".to_string()) {
1402 conn.execute("ALTER TABLE models ADD COLUMN supports_audio INTEGER", [])?;
1403 }
1404 if !columns.contains(&"supports_video".to_string()) {
1405 conn.execute("ALTER TABLE models ADD COLUMN supports_video INTEGER", [])?;
1406 }
1407 if !columns.contains(&"supports_documents".to_string()) {
1408 conn.execute(
1409 "ALTER TABLE models ADD COLUMN supports_documents INTEGER",
1410 [],
1411 )?;
1412 }
1413 if !columns.contains(&"reasoning_levels".to_string()) {
1414 conn.execute(
1415 "ALTER TABLE models ADD COLUMN reasoning_levels TEXT NOT NULL DEFAULT '[]'",
1416 [],
1417 )?;
1418 }
1419 if !columns.contains(&"default_reasoning_effort".to_string()) {
1420 conn.execute(
1421 "ALTER TABLE models ADD COLUMN default_reasoning_effort TEXT",
1422 [],
1423 )?;
1424 }
1425
1426 Ok(())
1427}
1428
1429fn ensure_provider_tool_calling_mode_column(conn: &Connection) -> Result<()> {
1430 let mut stmt = conn.prepare("PRAGMA table_info(providers)")?;
1431 let has_column = stmt
1432 .query_map([], |row| row.get::<_, String>(1))?
1433 .any(|name| matches!(name, Ok(name) if name == "tool_calling_mode"));
1434
1435 if !has_column {
1436 conn.execute(
1437 "ALTER TABLE providers ADD COLUMN tool_calling_mode TEXT",
1438 [],
1439 )?;
1440 }
1441
1442 Ok(())
1443}
1444
1445fn ensure_provider_sha256_column(conn: &Connection) -> Result<()> {
1446 let mut stmt = conn.prepare("PRAGMA table_info(providers)")?;
1447 let has_column = stmt
1448 .query_map([], |row| row.get::<_, String>(1))?
1449 .any(|name| matches!(name, Ok(name) if name == "sha256"));
1450
1451 if !has_column {
1452 conn.execute("ALTER TABLE providers ADD COLUMN sha256 TEXT", [])?;
1453 }
1454
1455 Ok(())
1456}
1457
1458fn ensure_provider_aggregator_column(conn: &Connection) -> Result<()> {
1459 let mut stmt = conn.prepare("PRAGMA table_info(providers)")?;
1460 let has_column = stmt
1461 .query_map([], |row| row.get::<_, String>(1))?
1462 .any(|name| matches!(name, Ok(name) if name == "aggregator"));
1463
1464 if !has_column {
1465 conn.execute(
1466 "ALTER TABLE providers ADD COLUMN aggregator INTEGER NOT NULL DEFAULT 0",
1467 [],
1468 )?;
1469 }
1470
1471 Ok(())
1472}
1473
1474fn relax_models_task_size_not_null(conn: &Connection) -> Result<()> {
1478 let mut stmt = conn.prepare("PRAGMA table_info(models)")?;
1480 let has_not_null: bool = stmt
1481 .query_map([], |row| {
1482 let name: String = row.get(1)?;
1483 let notnull: i64 = row.get(3)?;
1484 Ok((name, notnull))
1485 })?
1486 .filter_map(|r| r.ok())
1487 .any(|(name, notnull)| name == "task_size" && notnull != 0);
1488
1489 if !has_not_null {
1490 return Ok(());
1491 }
1492
1493 tracing::info!("migrating models table: relaxing task_size NOT NULL constraint");
1494
1495 conn.execute_batch(
1496 "
1497 CREATE TABLE IF NOT EXISTS models_new (
1498 provider_id TEXT NOT NULL,
1499 name TEXT NOT NULL,
1500 task_size TEXT,
1501 context_window_tokens INTEGER,
1502 max_output_tokens INTEGER,
1503 recommended_temperature REAL,
1504 supports_thinking INTEGER,
1505 supports_images INTEGER,
1506 supports_audio INTEGER,
1507 supports_video INTEGER,
1508 supports_documents INTEGER,
1509 tool_prompt_manifest INTEGER,
1510 PRIMARY KEY (provider_id, name),
1511 FOREIGN KEY (provider_id) REFERENCES providers(id) ON DELETE CASCADE
1512 );
1513
1514 INSERT INTO models_new (provider_id, name, task_size, context_window_tokens, max_output_tokens, recommended_temperature, supports_thinking, supports_images, supports_audio, supports_video, supports_documents, tool_prompt_manifest)
1515 SELECT provider_id, name, task_size, context_window_tokens, max_output_tokens, recommended_temperature, supports_thinking, supports_images, supports_audio, supports_video, supports_documents, tool_prompt_manifest
1516 FROM models;
1517
1518 DROP TABLE models;
1519 ALTER TABLE models_new RENAME TO models;
1520 ",
1521 )?;
1522
1523 Ok(())
1524}
1525
1526#[cfg(test)]
1527mod tests {
1528 use crate::config::types::ProviderRequestOptions;
1529
1530 use super::*;
1531 use crate::registry::types::RegistryModel;
1532
1533 fn sample_provider() -> RegistryProvider {
1534 RegistryProvider {
1535 id: "test-provider".to_string(),
1536 label: "Test Provider".to_string(),
1537 description: "A test".to_string(),
1538 kind: "openai-chat-completions".to_string(),
1539 api_key_env: "TEST_API_KEY".to_string(),
1540 base_url: Some("https://api.test.com/v1".to_string()),
1541 extends: None,
1542 tool_calling_mode: None,
1543 aggregator: false,
1544 defaults: Default::default(),
1545 request_options: Default::default(),
1546 models: vec![
1547 RegistryModel {
1548 model_ref: None,
1549 api_name: None,
1550 name: "test-model-large".to_string(),
1551 task_size: Some("large".to_string()),
1552 context_window_tokens: Some(200_000),
1553 max_output_tokens: Some(8_192),
1554 recommended_temperature: Some(0.7),
1555 supports_thinking: None,
1556 reasoning_levels: Vec::new(),
1557 default_reasoning_effort: None,
1558 supports_attachments: None,
1559 supports_images: None,
1560 supports_audio: None,
1561 supports_video: None,
1562 supports_documents: None,
1563 attachments: Default::default(),
1564 capabilities: Vec::new(),
1565 pricing: None,
1566 },
1567 RegistryModel {
1568 model_ref: None,
1569 api_name: None,
1570 name: "test-model-small".to_string(),
1571 task_size: Some("small".to_string()),
1572 context_window_tokens: Some(128_000),
1573 max_output_tokens: Some(4_096),
1574 recommended_temperature: Some(0.5),
1575 supports_thinking: None,
1576 reasoning_levels: Vec::new(),
1577 default_reasoning_effort: None,
1578 supports_attachments: None,
1579 supports_images: None,
1580 supports_audio: None,
1581 supports_video: None,
1582 supports_documents: None,
1583 attachments: Default::default(),
1584 capabilities: Vec::new(),
1585 pricing: None,
1586 },
1587 ],
1588 }
1589 }
1590
1591 #[test]
1592 fn open_and_init_schema() {
1593 let store = RegistryStore::open_memory().expect("open");
1594 assert!(store.is_empty().unwrap());
1595 }
1596
1597 #[test]
1598 fn canonical_model_roundtrip() {
1599 let store = RegistryStore::open_memory().expect("open");
1600 let model = super::super::types::CanonicalModel {
1601 id: "gpt-test".into(),
1602 vendor: Some("openai".into()),
1603 family: None,
1604 label: None,
1605 description: None,
1606 context_window_tokens: Some(128_000),
1607 max_output_tokens: Some(8_192),
1608 recommended_temperature: None,
1609 supports_thinking: Some(true),
1610 reasoning_levels: vec!["low".into(), "high".into()],
1611 default_reasoning_effort: Some("low".into()),
1612 attachments: Default::default(),
1613 capabilities: Vec::new(),
1614 status: Some("active".into()),
1615 aliases: vec!["gpt-test-alias".into()],
1616 };
1617 store
1618 .upsert_canonical_model("gpt-test", &model, Some("abc123"))
1619 .expect("upsert");
1620 assert_eq!(store.canonical_model_count().unwrap(), 1);
1621 assert_eq!(
1622 store.canonical_model_sha256("gpt-test").unwrap().as_deref(),
1623 Some("abc123")
1624 );
1625 let catalog = store.load_canonical_model_catalog().expect("load");
1626 assert_eq!(catalog["gpt-test"].context_window_tokens, Some(128_000));
1627 assert_eq!(
1628 catalog["gpt-test"].aliases,
1629 vec!["gpt-test-alias".to_string()]
1630 );
1631
1632 let mut keep = std::collections::HashSet::new();
1633 keep.insert("other");
1634 store.delete_canonical_models_not_in(&keep).expect("delete");
1635 assert_eq!(store.canonical_model_count().unwrap(), 0);
1636 }
1637
1638 #[test]
1639 fn tool_calling_mode_roundtrips_through_store() {
1640 let store = RegistryStore::open_memory().expect("open");
1641 let mut provider = sample_provider();
1642 provider.tool_calling_mode = Some("native".to_string());
1643 store.upsert_provider(&provider).expect("upsert");
1644
1645 let loaded = store.load_all_providers().expect("load");
1646 assert_eq!(loaded.len(), 1);
1647 assert_eq!(loaded[0].tool_calling_mode, Some(ToolCallingMode::Native));
1648 }
1649
1650 #[test]
1651 fn upsert_and_load_provider() {
1652 let store = RegistryStore::open_memory().expect("open");
1653 let provider = sample_provider();
1654 store.upsert_provider(&provider).expect("upsert");
1655
1656 assert_eq!(store.provider_count().unwrap(), 1);
1657 assert_eq!(store.model_count().unwrap(), 2);
1658
1659 let loaded = store.load_all_providers().expect("load");
1660 assert_eq!(loaded.len(), 1);
1661 assert_eq!(loaded[0].id, "test-provider");
1662 assert_eq!(loaded[0].models.len(), 2);
1663 assert_eq!(loaded[0].models[0].name, "test-model-large");
1664 assert_eq!(loaded[0].models[0].context_window_tokens, Some(200_000));
1665 assert_eq!(loaded[0].models[0].max_output_tokens, Some(8_192));
1666 assert_eq!(loaded[0].models[0].recommended_temperature, Some(0.7));
1667 assert_eq!(loaded[0].models[0].task_size, Some(ModelTaskSize::Large));
1668 assert_eq!(loaded[0].models[1].task_size, Some(ModelTaskSize::Small));
1669 assert_eq!(loaded[0].kind, ProviderKind::OpenAiChatCompletions);
1670 assert_eq!(
1671 loaded[0].base_url,
1672 Some("https://api.test.com/v1".to_string())
1673 );
1674 }
1675
1676 #[test]
1677 fn upsert_replaces_models() {
1678 let store = RegistryStore::open_memory().expect("open");
1679 let mut provider = sample_provider();
1680 store.upsert_provider(&provider).expect("upsert");
1681 assert_eq!(store.model_count().unwrap(), 2);
1682
1683 provider.models = vec![RegistryModel {
1685 model_ref: None,
1686 api_name: None,
1687 name: "new-model".to_string(),
1688 task_size: Some("large".to_string()),
1689 context_window_tokens: Some(500_000),
1690 max_output_tokens: Some(16_384),
1691 recommended_temperature: Some(0.8),
1692 supports_thinking: None,
1693 reasoning_levels: Vec::new(),
1694 default_reasoning_effort: None,
1695 supports_attachments: None,
1696 supports_images: None,
1697 supports_audio: None,
1698 supports_video: None,
1699 supports_documents: None,
1700 attachments: Default::default(),
1701 capabilities: Vec::new(),
1702 pricing: None,
1703 }];
1704 store.upsert_provider(&provider).expect("upsert again");
1705
1706 assert_eq!(store.model_count().unwrap(), 1);
1707 let loaded = store.load_all_providers().expect("load");
1708 assert_eq!(loaded[0].models[0].name, "new-model");
1709 assert_eq!(loaded[0].models[0].context_window_tokens, Some(500_000));
1710 assert_eq!(loaded[0].models[0].max_output_tokens, Some(16_384));
1711 assert_eq!(loaded[0].models[0].recommended_temperature, Some(0.8));
1712 }
1713
1714 #[test]
1715 fn transcription_provider_roundtrip() {
1716 let store = RegistryStore::open_memory().expect("open");
1717 let provider = RegistryTranscriptionProvider {
1718 id: "openai".to_string(),
1719 label: "OpenAI Whisper".to_string(),
1720 description: "test".to_string(),
1721 kind: "openai-audio-transcriptions".to_string(),
1722 api_key_env: "OPENAI_API_KEY".to_string(),
1723 base_url: "https://api.openai.com/v1".to_string(),
1724 transcription_path: Some("/audio/transcriptions".to_string()),
1725 default_model: Some("whisper-1".to_string()),
1726 supports_streaming: false,
1727 models: vec![super::super::types::RegistryTranscriptionModel {
1728 name: "whisper-1".to_string(),
1729 label: Some("Whisper v1".to_string()),
1730 description: None,
1731 languages: vec![],
1732 sample_rate_hz: Some(16_000),
1733 max_duration_seconds: None,
1734 max_file_bytes: Some(25_000_000),
1735 pricing: None,
1736 }],
1737 };
1738 store
1739 .upsert_transcription_provider(&provider, Some("abc123"))
1740 .expect("upsert");
1741 assert_eq!(store.transcription_provider_count().unwrap(), 1);
1742 assert_eq!(
1743 store
1744 .transcription_provider_sha256("openai")
1745 .unwrap()
1746 .as_deref(),
1747 Some("abc123")
1748 );
1749 let loaded = store.load_transcription_providers().expect("load");
1750 assert_eq!(loaded.len(), 1);
1751 assert_eq!(loaded[0].id, "openai");
1752 assert_eq!(loaded[0].models[0].name, "whisper-1");
1753 assert_eq!(loaded[0].resolved_default_model(), Some("whisper-1"));
1754 }
1755
1756 #[test]
1757 fn upsert_union_preserves_api_synced_extras() {
1758 let store = RegistryStore::open_memory().expect("open");
1759 let mut provider = sample_provider();
1760 provider.models.push(RegistryModel {
1762 model_ref: None,
1763 api_name: None,
1764 name: "api-only-model".to_string(),
1765 task_size: Some("large".to_string()),
1766 context_window_tokens: Some(200_000),
1767 max_output_tokens: None,
1768 recommended_temperature: None,
1769 supports_thinking: Some(true),
1770 reasoning_levels: Vec::new(),
1771 default_reasoning_effort: None,
1772 supports_attachments: None,
1773 supports_images: None,
1774 supports_audio: None,
1775 supports_video: None,
1776 supports_documents: None,
1777 attachments: Default::default(),
1778 capabilities: Vec::new(),
1779 pricing: None,
1780 });
1781 store
1782 .upsert_provider_with_sha256(&provider, Some(LOCAL_API_SYNC_SHA))
1783 .expect("api sync upsert");
1784 assert_eq!(store.model_count().unwrap(), 3);
1785 assert_eq!(
1786 store.provider_sha256("test-provider").unwrap().as_deref(),
1787 Some(LOCAL_API_SYNC_SHA)
1788 );
1789
1790 let catalog = sample_provider(); store
1793 .upsert_provider_union_models(&catalog, Some("catalog-sha"))
1794 .expect("catalog union");
1795
1796 let loaded = store.load_provider_models("test-provider").expect("load");
1797 assert_eq!(loaded.len(), 3, "union must keep api-only-model");
1798 assert!(loaded.contains_key("api-only-model"));
1799 assert!(loaded.contains_key("test-model-large"));
1800 assert!(loaded.contains_key("test-model-small"));
1801 }
1802
1803 #[test]
1804 fn rehydrate_from_catalog_fixes_stale_context_and_efforts() {
1805 let store = RegistryStore::open_memory().expect("open");
1806
1807 let mut catalog_model = crate::registry::types::CanonicalModel {
1809 id: "grok-4.5".into(),
1810 vendor: Some("xai".into()),
1811 family: Some("grok".into()),
1812 label: None,
1813 description: None,
1814 context_window_tokens: Some(500_000),
1815 max_output_tokens: Some(131_072),
1816 recommended_temperature: Some(1.0),
1817 supports_thinking: Some(true),
1818 reasoning_levels: vec!["low".into(), "medium".into(), "high".into()],
1819 default_reasoning_effort: Some("medium".into()),
1820 attachments: Default::default(),
1821 capabilities: Vec::new(),
1822 status: Some("active".into()),
1823 aliases: Vec::new(),
1824 };
1825 store
1826 .upsert_canonical_model("grok-4.5", &catalog_model, Some("canon-sha"))
1827 .expect("upsert canonical");
1828
1829 let provider = RegistryProvider {
1831 id: "xai".into(),
1832 label: "xAI".into(),
1833 description: String::new(),
1834 kind: "openai-responses".into(),
1835 api_key_env: "XAI_API_KEY".into(),
1836 base_url: Some("https://api.x.ai/v1".into()),
1837 extends: None,
1838 tool_calling_mode: None,
1839 aggregator: false,
1840 defaults: Default::default(),
1841 request_options: Default::default(),
1842 models: vec![RegistryModel {
1843 model_ref: None,
1844 api_name: None,
1845 name: "grok-4.5".into(),
1846 task_size: None,
1847 context_window_tokens: Some(1_000_000),
1848 max_output_tokens: None,
1849 recommended_temperature: None,
1850 supports_thinking: Some(true),
1851 reasoning_levels: vec!["low".into(), "high".into()],
1852 default_reasoning_effort: Some("high".into()),
1853 supports_attachments: None,
1854 supports_images: Some(true),
1855 supports_audio: None,
1856 supports_video: None,
1857 supports_documents: None,
1858 attachments: Default::default(),
1859 capabilities: Vec::new(),
1860 pricing: None,
1861 }],
1862 };
1863 store
1864 .upsert_provider_with_sha256(&provider, Some(LOCAL_API_SYNC_SHA))
1865 .expect("api sync upsert");
1866
1867 let n = store
1868 .rehydrate_provider_models_from_catalog()
1869 .expect("rehydrate");
1870 assert!(n >= 1);
1871
1872 let models = store.load_provider_models("xai").expect("load");
1873 let grok = models.get("grok-4.5").expect("grok-4.5 present");
1874 assert_eq!(grok.context_window_tokens, Some(500_000));
1875 assert_eq!(
1876 grok.reasoning_levels,
1877 vec!["low".to_string(), "medium".to_string(), "high".to_string()]
1878 );
1879 assert_eq!(grok.default_reasoning_effort.as_deref(), Some("medium"));
1880 assert_eq!(grok.max_output_tokens, Some(131_072));
1881 assert_eq!(
1883 store.provider_sha256("xai").unwrap().as_deref(),
1884 Some(LOCAL_API_SYNC_SHA)
1885 );
1886
1887 let _ = &mut catalog_model; }
1889
1890 #[test]
1891 fn replace_all_clears_old_data() {
1892 let store = RegistryStore::open_memory().expect("open");
1893 store.upsert_provider(&sample_provider()).expect("upsert");
1894 assert_eq!(store.provider_count().unwrap(), 1);
1895
1896 let new_providers = vec![RegistryProvider {
1897 id: "other".to_string(),
1898 label: "Other".to_string(),
1899 description: String::new(),
1900 kind: "anthropic-messages".to_string(),
1901 api_key_env: "OTHER_KEY".to_string(),
1902 base_url: None,
1903 extends: None,
1904 tool_calling_mode: None,
1905 aggregator: false,
1906 defaults: Default::default(),
1907 request_options: Default::default(),
1908 models: vec![],
1909 }];
1910 store.replace_all(&new_providers).expect("replace");
1911
1912 assert_eq!(store.provider_count().unwrap(), 1);
1913 assert_eq!(store.model_count().unwrap(), 0);
1914 let loaded = store.load_all_providers().expect("load");
1915 assert_eq!(loaded[0].id, "other");
1916 assert_eq!(loaded[0].kind, ProviderKind::AnthropicMessages);
1917 }
1918
1919 #[test]
1920 fn meta_get_set() {
1921 let store = RegistryStore::open_memory().expect("open");
1922 assert_eq!(store.meta_get("foo").unwrap(), None);
1923
1924 store.meta_set("foo", "bar").expect("set");
1925 assert_eq!(store.meta_get("foo").unwrap(), Some("bar".to_string()));
1926
1927 store.meta_set("foo", "baz").expect("set");
1929 assert_eq!(store.meta_get("foo").unwrap(), Some("baz".to_string()));
1930 }
1931
1932 #[test]
1933 fn context_window_none_survives_roundtrip() {
1934 let store = RegistryStore::open_memory().expect("open");
1935 let provider = RegistryProvider {
1936 id: "p".to_string(),
1937 label: "P".to_string(),
1938 description: String::new(),
1939 kind: "openai-chat-completions".to_string(),
1940 api_key_env: "K".to_string(),
1941 base_url: None,
1942 extends: None,
1943 tool_calling_mode: None,
1944 aggregator: false,
1945 defaults: Default::default(),
1946 request_options: Default::default(),
1947 models: vec![RegistryModel {
1948 model_ref: None,
1949 api_name: None,
1950 name: "m".to_string(),
1951 task_size: Some("small".to_string()),
1952 context_window_tokens: None,
1953 max_output_tokens: None,
1954 recommended_temperature: None,
1955 supports_thinking: None,
1956 reasoning_levels: Vec::new(),
1957 default_reasoning_effort: None,
1958 supports_attachments: None,
1959 supports_images: None,
1960 supports_audio: None,
1961 supports_video: None,
1962 supports_documents: None,
1963 attachments: Default::default(),
1964 capabilities: Vec::new(),
1965 pricing: None,
1966 }],
1967 };
1968 store.upsert_provider(&provider).expect("upsert");
1969 let loaded = store.load_all_providers().expect("load");
1970 assert_eq!(loaded[0].models[0].context_window_tokens, None);
1971 }
1972
1973 #[test]
1974 fn request_options_survive_roundtrip() {
1975 let store = RegistryStore::open_memory().expect("open");
1976 let mut provider = sample_provider();
1977 provider.request_options = ProviderRequestOptions {
1978 prompt_cache_key: Some("openai".to_string()),
1979 prompt_cache_retention: Some("24h".to_string()),
1980 anthropic_cache_control: Some(serde_json::json!({
1981 "type": "ephemeral",
1982 "ttl": "1h"
1983 })),
1984 };
1985
1986 store.upsert_provider(&provider).expect("upsert");
1987 let loaded = store.load_all_providers().expect("load");
1988
1989 let opts = loaded[0]
1990 .request_options
1991 .as_ref()
1992 .expect("request_options roundtripped");
1993 assert_eq!(opts.prompt_cache_key.as_deref(), Some("openai"));
1994 assert_eq!(opts.prompt_cache_retention.as_deref(), Some("24h"));
1995 assert_eq!(
1996 opts.anthropic_cache_control
1997 .as_ref()
1998 .and_then(|value| value.get("ttl"))
1999 .and_then(serde_json::Value::as_str),
2000 Some("1h")
2001 );
2002 }
2003
2004 #[test]
2005 fn capabilities_upsert_and_load() {
2006 let store = RegistryStore::open_memory().expect("open");
2007 store.upsert_provider(&sample_provider()).expect("upsert");
2008
2009 let model_id = "test-provider:test-model-large";
2010 let caps = vec![
2011 ("tool_calling".to_string(), "true".to_string()),
2012 ("fast".to_string(), "true".to_string()),
2013 ("cheap".to_string(), "true".to_string()),
2014 ];
2015 store
2016 .upsert_capabilities(model_id, "test-provider", &caps)
2017 .expect("upsert caps");
2018
2019 let loaded = store.load_capabilities(model_id).expect("load caps");
2020 assert_eq!(loaded.len(), 3);
2021 assert!(loaded.iter().any(|c| c.capability == "tool_calling"));
2022
2023 let caps2 = vec![("fast".to_string(), "true".to_string())];
2025 store
2026 .upsert_capabilities(model_id, "test-provider", &caps2)
2027 .expect("replace caps");
2028 let loaded2 = store.load_capabilities(model_id).expect("load caps 2");
2029 assert_eq!(loaded2.len(), 1);
2030 assert_eq!(loaded2[0].capability, "fast");
2031 }
2032
2033 #[test]
2034 fn pricing_upsert_and_load() {
2035 let store = RegistryStore::open_memory().expect("open");
2036 store.upsert_provider(&sample_provider()).expect("upsert");
2037
2038 let model_id = "test-provider:test-model-large";
2039 store
2040 .upsert_pricing(model_id, "test-provider", Some(0.10), Some(0.30))
2041 .expect("upsert pricing");
2042
2043 let loaded = store.load_pricing(model_id).expect("load pricing");
2044 let pricing = loaded.expect("pricing exists");
2045 assert_eq!(pricing.input_price, Some(0.10));
2046 assert_eq!(pricing.output_price, Some(0.30));
2047 assert_eq!(pricing.currency, "USD");
2048 }
2049
2050 #[test]
2051 fn pricing_returns_none_for_missing() {
2052 let store = RegistryStore::open_memory().expect("open");
2053 let loaded = store.load_pricing("nonexistent:model").expect("load");
2054 assert!(loaded.is_none());
2055 }
2056
2057 #[test]
2058 fn profiles_seed_and_query() {
2059 let store = RegistryStore::open_memory().expect("open");
2060 store.upsert_provider(&sample_provider()).expect("upsert");
2061
2062 let model_id = "test-provider:test-model-large";
2063 store
2064 .upsert_pricing(model_id, "test-provider", Some(0.10), Some(0.30))
2065 .expect("pricing");
2066 store
2067 .upsert_model_profile(model_id, "test-provider", "cheap_general", 0.9)
2068 .expect("profile");
2069
2070 store.seed_default_profiles().expect("seed profiles");
2071
2072 let ranked = store
2073 .query_models_by_profile("cheap_general")
2074 .expect("query");
2075 assert!(!ranked.is_empty());
2076 assert_eq!(ranked[0].model_id, model_id);
2077 assert_eq!(ranked[0].score, 0.9);
2078 }
2079
2080 #[test]
2081 fn query_respects_min_context_filter() {
2082 let store = RegistryStore::open_memory().expect("open");
2083
2084 let provider = RegistryProvider {
2086 id: "tiny".to_string(),
2087 label: "Tiny".to_string(),
2088 description: String::new(),
2089 kind: "openai-chat-completions".to_string(),
2090 api_key_env: "TINY_KEY".to_string(),
2091 base_url: None,
2092 extends: None,
2093 tool_calling_mode: None,
2094 aggregator: false,
2095 defaults: Default::default(),
2096 request_options: Default::default(),
2097 models: vec![RegistryModel {
2098 model_ref: None,
2099 api_name: None,
2100 name: "tiny-model".to_string(),
2101 task_size: Some("small".to_string()),
2102 context_window_tokens: Some(4_000),
2103 max_output_tokens: None,
2104 recommended_temperature: None,
2105 supports_thinking: None,
2106 reasoning_levels: Vec::new(),
2107 default_reasoning_effort: None,
2108 supports_attachments: None,
2109 supports_images: None,
2110 supports_audio: None,
2111 supports_video: None,
2112 supports_documents: None,
2113 attachments: Default::default(),
2114 capabilities: Vec::new(),
2115 pricing: None,
2116 }],
2117 };
2118 store.upsert_provider(&provider).expect("upsert");
2119
2120 let model_id = "tiny:tiny-model";
2121 store
2122 .upsert_model_profile(model_id, "tiny", "cheap_general", 1.0)
2123 .expect("profile");
2124
2125 store.seed_default_profiles().expect("seed");
2126
2127 let ranked = store
2129 .query_models_by_profile("cheap_general")
2130 .expect("query");
2131 assert!(
2132 ranked.is_empty(),
2133 "tiny-model should be filtered out by min_context"
2134 );
2135 }
2136
2137 #[test]
2138 fn delete_provider_metadata_cascades() {
2139 let store = RegistryStore::open_memory().expect("open");
2140 store.upsert_provider(&sample_provider()).expect("upsert");
2141
2142 let model_id = "test-provider:test-model-large";
2143 store
2144 .upsert_capabilities(model_id, "test-provider", &[("fast".into(), "true".into())])
2145 .expect("caps");
2146 store
2147 .upsert_pricing(model_id, "test-provider", Some(0.10), Some(0.30))
2148 .expect("pricing");
2149 store
2150 .upsert_model_profile(model_id, "test-provider", "cheap_general", 0.9)
2151 .expect("profile");
2152
2153 store
2154 .delete_provider_metadata("test-provider")
2155 .expect("delete");
2156
2157 assert!(store.load_capabilities(model_id).unwrap().is_empty());
2158 assert!(store.load_pricing(model_id).unwrap().is_none());
2159 let ranked = store.query_models_by_profile("cheap_general").unwrap();
2160 assert!(ranked.is_empty());
2161 }
2162
2163 #[test]
2164 fn open_recreates_corrupt_on_disk_registry_db() {
2165 let dir = tempfile::tempdir().expect("tempdir");
2166 let db_path = dir.path().join("registry.db");
2167 std::fs::write(&db_path, b"not a sqlite database").expect("write corrupt db");
2169 std::fs::write(dir.path().join("registry.db-wal"), []).expect("wal");
2170 std::fs::write(dir.path().join("registry.db-shm"), vec![0u8; 32_768]).expect("shm");
2171
2172 let store = RegistryStore::open(dir.path()).expect("open should recreate");
2173 assert!(
2174 !store.is_empty().expect("is_empty"),
2175 "recreated store should be seeded from embedded snapshot"
2176 );
2177 assert!(db_path.exists(), "registry.db should exist after recreate");
2178 }
2179}