Skip to main content

navi_core/registry/
store.rs

1//! Local SQLite cache for the provider registry.
2
3use 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
16/// Marker written to `providers.sha256` when models were populated from a
17/// live provider API (`sync models`). Embedded/remote catalog updates must
18/// union-merge against these rows instead of treating a missing/different
19/// hash as permission to replace the list.
20pub const LOCAL_API_SYNC_SHA: &str = "local-api-sync";
21
22/// Removes `registry.db` plus SQLite sidecar files (`-wal`, `-shm`, `-journal`).
23fn 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
39/// SQLite-backed registry store.
40///
41/// Thread-safe via internal `Mutex<Connection>` — registry operations are
42/// short-lived and infrequent so contention is negligible.
43pub struct RegistryStore {
44    conn: Mutex<Connection>,
45}
46
47impl RegistryStore {
48    /// Opens (or creates) the registry database at `<data_dir>/registry.db`.
49    ///
50    /// On first run (empty database), seeds the cache from the embedded registry
51    /// snapshot so the provider catalog is immediately available without a
52    /// network fetch.
53    ///
54    /// If an existing database is unreadable (corrupt file, truncated WAL/SHM,
55    /// or SQLite disk I/O errors on open), the broken files are removed and a
56    /// fresh cache is recreated from the embedded snapshot. Callers that need
57    /// the remote catalog should follow up with `sync_registry`.
58    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        // Fail fast on truncated/corrupt caches before schema init. SQLite can
86        // open a header-only file and only error later on the first query.
87        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        // WAL mode for concurrent reads, faster writes.
96        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    /// Seeds providers/manifest/transcription catalog when the cache is empty.
108    fn seed_if_empty(&self) -> Result<()> {
109        // Seed from the embedded snapshot if the cache is empty.
110        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                // Also persist the full manifest JSON so load_cached_registry
121                // and check_registry_manifest can find it.
122                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                // Seed canonical model catalog with hashes from the embedded
127                // manifest so remote sync can skip unchanged models.
128                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        // Always ensure transcription providers are seeded (may be empty on
138        // DBs created before STT catalog support).
139        self.seed_transcription_from_embedded_if_empty()?;
140        self.seed_canonical_models_from_embedded_if_empty()?;
141        Ok(())
142    }
143
144    /// Opens an in-memory database (for testing).
145    #[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    // ── Metadata helpers ──────────────────────────────────────────────────
260
261    /// Returns the value of a metadata key, or `None`.
262    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    /// Sets a metadata key-value pair (upsert).
275    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    // ── Provider / model CRUD ─────────────────────────────────────────────
285
286    /// Returns `true` if the providers table is empty.
287    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    /// Returns the number of providers in the cache.
294    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    /// Returns the total number of models across all providers.
301    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    /// Returns the stored SHA-256 hash for a provider, or `None`.
308    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    /// Returns the set of provider ids currently in the cache.
320    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    /// Loads existing models for a provider from the cache, keyed by model name.
332    /// Used by aggregator sync to preserve metadata (context_window, etc) for
333    /// models that the API returns without rich metadata.
334    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    /// Deletes providers that are not in the given set of ids.
388    /// Used during sync to remove providers that were deleted from the remote registry.
389    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    /// Upserts a provider and its models from a [`RegistryProvider`].
411    pub fn upsert_provider(&self, provider: &RegistryProvider) -> Result<()> {
412        self.upsert_provider_with_sha256(provider, None)
413    }
414
415    // ── Transcription / dictation providers ───────────────────────────────
416
417    /// Upserts a remote transcription provider (stored as JSON for simplicity).
418    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    /// SHA-256 of a cached transcription provider, if known.
435    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    /// Loads all cached transcription providers.
446    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    /// Deletes transcription providers not in `keep`.
464    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    // ── Canonical model catalog ─────────────────────────────────────────
486
487    /// Upserts a canonical model JSON blob with its integrity hash.
488    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    /// Returns the SHA-256 of a cached canonical model, if any.
505    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    /// Loads the full canonical model catalog from the cache.
516    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    /// Deletes canonical models not present in `keep`.
540    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    /// Number of cached canonical models.
558    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    /// Number of cached transcription providers.
567    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    /// Seeds the canonical model catalog from the embedded snapshot when empty.
577    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    /// Seeds the transcription cache from the embedded snapshot when empty.
597    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    /// Upserts a provider while **preserving** any local models that are not in
621    /// `provider.models`.
622    ///
623    /// Used by embedded/remote registry updates so a smaller catalog snapshot
624    /// cannot wipe models discovered via `sync models` / provider `/models`.
625    /// Incoming models win on name conflicts (refresh metadata); local-only
626    /// names are kept.
627    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    /// Re-apply canonical model catalog metadata onto every cached provider model.
652    ///
653    /// Fixes stale rows written by API sync sibling inheritance (e.g. grok-4.5
654    /// wrongly holding grok-4.3's 1M context / low+high efforts) without wiping
655    /// API-discovered model names or changing provider `sha256` markers.
656    ///
657    /// Returns the number of model rows touched by the UPDATE.
658    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            // Dedup names (id may equal an alias).
686            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    /// Upserts a provider with its SHA-256 hash for diff-based sync.
711    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        // Delete existing models for this provider, then re-insert.
739        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        // Seed/refresh pricing from registry JSON (per 1M token rates).
773        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; // used above for validation if needed
802        Ok(())
803    }
804
805    /// Loads all providers from the cache as [`ProviderConfig`] values
806    /// compatible with the existing catalog.
807    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    /// Replaces the entire cache with the given providers (full refresh).
923    pub fn replace_all(&self, providers: &[RegistryProvider]) -> Result<()> {
924        let conn = self.conn.lock().unwrap_or_else(|e| e.into_inner());
925
926        // Wipe existing data.
927        conn.execute("DELETE FROM models", [])?;
928        conn.execute("DELETE FROM providers", [])?;
929
930        drop(conn); // release lock before per-provider upsert
931
932        for provider in providers {
933            self.upsert_provider(provider)?;
934        }
935
936        Ok(())
937    }
938
939    /// Saves the manifest metadata.
940    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    /// Returns the stored manifest version, if any.
951    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    /// Returns the stored manifest `updated_at`, if any.
959    pub fn manifest_updated_at(&self) -> Result<Option<String>> {
960        self.meta_get("manifest_updated_at")
961    }
962
963    // ── Capabilities CRUD ───────────────────────────────────────────────
964
965    /// Upserts capabilities for a model. Replaces all existing capabilities for that model.
966    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    /// Loads all capabilities for a model.
992    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    // ── Pricing CRUD ────────────────────────────────────────────────────
1012
1013    /// Upserts pricing for a model.
1014    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    /// Loads pricing for a model.
1031    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    // ── Profiles CRUD ───────────────────────────────────────────────────
1053
1054    /// Upserts a profile definition.
1055    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    /// Upserts a model-profile association.
1072    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    /// Queries for models matching a profile, ranked by score and price.
1089    ///
1090    /// Returns models that satisfy the profile's constraints (min context, max price,
1091    /// tool support) ordered by score descending, then input price ascending.
1092    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    /// Seeds the default built-in profile definitions.
1132    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    /// Deletes all capabilities, pricing, and model-profile entries for a provider.
1184    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
1222/// Converts a [`RegistryProvider`] into a [`ProviderConfig`].
1223///
1224/// This is the single conversion path used by both the SQLite cache loader and
1225/// the embedded snapshot fallback, ensuring consistent field mapping.
1226pub 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
1474/// Migrates the `models` table from an older schema where `task_size` was
1475/// `NOT NULL` to the current nullable version. SQLite doesn't support
1476/// `ALTER COLUMN`, so we rebuild the table.
1477fn relax_models_task_size_not_null(conn: &Connection) -> Result<()> {
1478    // Check if task_size has a NOT NULL constraint.
1479    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        // Update with fewer models.
1684        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        // Simulate API sync with an extra model beyond the catalog snapshot.
1761        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        // Smaller catalog snapshot must not wipe the API-only model.
1791        let catalog = sample_provider(); // 2 models
1792        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        // Canonical truth: grok-4.5 is 500k with low/medium/high.
1808        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        // Stale API-sync row: wrong 1M / low+high inherited from sibling.
1830        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        // API sync marker must survive rehydrate.
1882        assert_eq!(
1883            store.provider_sha256("xai").unwrap().as_deref(),
1884            Some(LOCAL_API_SYNC_SHA)
1885        );
1886
1887        let _ = &mut catalog_model; // silence unused mut if any
1888    }
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        // Overwrite.
1928        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        // Replace with fewer caps.
2024        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        // Insert a provider with a small-context model.
2085        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        // cheap_general requires min_context 32k, tiny-model has 4k.
2128        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        // Non-SQLite payload that fails on open with "file is not a database".
2168        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}