Skip to main content

smriti/services/
semantic.rs

1//! Semantic image search over local CLIP-style embeddings.
2//!
3//! The database stores only indexing state and vector offsets. The
4//! high-volume embedding payload lives in `.photovault/semantic/...`
5//! beside thumbnails and other per-library cache data.
6
7use std::collections::HashMap;
8use std::fs::{File, OpenOptions};
9use std::io::{Read, Seek, SeekFrom, Write};
10use std::path::{Path, PathBuf};
11use std::sync::atomic::{AtomicBool, Ordering};
12
13use image::{DynamicImage, ImageBuffer, Rgb};
14use ndarray::Array1;
15use rusqlite::{params, Connection, OptionalExtension};
16use serde::{Deserialize, Serialize};
17use tokenizers::Tokenizer;
18
19use crate::db::connection::library_metadata_dir;
20use crate::ml::OnnxRuntime;
21use crate::services::image_io;
22use crate::services::path_util::safe_join_relative;
23
24pub const SEMANTIC_MODEL_KEY: &str = "immich-app/ViT-B-32-SigLIP2-256__webli";
25pub const SEMANTIC_MODEL_DISPLAY: &str = "ViT-B-32 SigLIP2 256";
26pub const SEMANTIC_MODEL_REVISION: &str = "762c736d366fc253e9453021144f9fe71789b075";
27pub const SEMANTIC_DIM: usize = 768;
28pub const SEMANTIC_CONTEXT_LEN: usize = 64;
29pub const SEMANTIC_TEXT_SEARCH_LIMIT: usize = 250;
30pub const SEMANTIC_TEXT_RESULT_CAP: usize = 80;
31
32const SEMANTIC_TEXT_MIN_SCORE: f32 = 0.06;
33const SEMANTIC_TEXT_MAX_SCORE_DROP: f32 = 0.02;
34const SEMANTIC_TEXT_MIN_SCORE_RATIO: f32 = 0.75;
35
36const MODEL_DIR_NAME: &str = "vit-b-32-siglip2-256-webli";
37const VECTOR_FILE: &str = "vectors.f32";
38const MANIFEST_FILE: &str = "manifest.json";
39
40const VISUAL_MODEL_URL: &str =
41    "https://huggingface.co/immich-app/ViT-B-32-SigLIP2-256__webli/resolve/main/visual/model.onnx";
42const TEXTUAL_MODEL_URL: &str =
43    "https://huggingface.co/immich-app/ViT-B-32-SigLIP2-256__webli/resolve/main/textual/model.onnx";
44const TOKENIZER_URL: &str = "https://huggingface.co/immich-app/ViT-B-32-SigLIP2-256__webli/resolve/main/textual/tokenizer.json";
45const PREPROCESS_URL: &str = "https://huggingface.co/immich-app/ViT-B-32-SigLIP2-256__webli/resolve/main/visual/preprocess_cfg.json";
46const CONFIG_URL: &str =
47    "https://huggingface.co/immich-app/ViT-B-32-SigLIP2-256__webli/resolve/main/config.json";
48
49const VISUAL_MODEL_BYTES: u64 = 378_359_772;
50const TEXTUAL_MODEL_BYTES: u64 = 1_129_435_819;
51const TOKENIZER_BYTES: u64 = 34_362_885;
52const PREPROCESS_BYTES: u64 = 154;
53const CONFIG_BYTES: u64 = 551;
54
55/// Total network payload for the optional local visual-search model pack.
56/// Exposed so setup surfaces can accurately disclose the one-time download.
57pub const SEMANTIC_MODEL_DOWNLOAD_BYTES: u64 =
58    VISUAL_MODEL_BYTES + TEXTUAL_MODEL_BYTES + TOKENIZER_BYTES + PREPROCESS_BYTES + CONFIG_BYTES;
59
60#[derive(Debug, Clone, Serialize, Deserialize)]
61pub struct SemanticStatus {
62    pub model_key: String,
63    pub display_name: String,
64    pub model_dir: String,
65    pub assets_installed: bool,
66    pub onnx_runtime_installed: bool,
67    pub indexed_photos: u64,
68    pub pending_photos: u64,
69    pub failed_photos: u64,
70    pub vector_bytes: u64,
71}
72
73#[derive(Debug, Clone, Copy, Default, Serialize, Deserialize)]
74pub struct SemanticIndexStats {
75    pub indexed: u64,
76    pub pending: u64,
77    pub failed: u64,
78}
79
80#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
81pub struct SemanticIndexBatchOutcome {
82    pub processed: u64,
83    pub indexed: u64,
84    pub failed: u64,
85    pub done: bool,
86}
87
88#[derive(Debug, Clone)]
89pub struct SemanticCandidate {
90    pub photo_id: i64,
91    pub score: f32,
92}
93
94#[derive(Debug, Clone, Serialize, Deserialize)]
95struct VectorManifest {
96    model_key: String,
97    revision: String,
98    dim: usize,
99    vector_count: u64,
100}
101
102#[derive(Debug, Clone)]
103struct SemanticAssetPaths {
104    root: PathBuf,
105    visual_model: PathBuf,
106    textual_model: PathBuf,
107    tokenizer: PathBuf,
108    preprocess: PathBuf,
109    config: PathBuf,
110}
111
112impl SemanticAssetPaths {
113    fn in_root(root: PathBuf) -> Self {
114        let model_root = root.join("models").join("semantic").join(MODEL_DIR_NAME);
115        Self {
116            visual_model: model_root.join("visual").join("model.onnx"),
117            textual_model: model_root.join("textual").join("model.onnx"),
118            tokenizer: model_root.join("textual").join("tokenizer.json"),
119            preprocess: model_root.join("visual").join("preprocess_cfg.json"),
120            config: model_root.join("config.json"),
121            root: model_root,
122        }
123    }
124
125    fn installed(&self) -> bool {
126        self.visual_model.exists()
127            && self.textual_model.exists()
128            && self.tokenizer.exists()
129            && self.preprocess.exists()
130            && self.config.exists()
131    }
132}
133
134pub struct SemanticSearchService {
135    drive_root: PathBuf,
136}
137
138#[derive(Default)]
139pub struct SemanticIndexCache {
140    indexed_count: u64,
141    #[cfg(feature = "hnsw_clustering")]
142    index: Option<SemanticHnswIndex>,
143}
144
145#[cfg(feature = "hnsw_clustering")]
146struct SemanticHnswIndex {
147    photo_ids: Vec<i64>,
148    exact_vectors: Option<Vec<Vec<f32>>>,
149    hnsw: hnsw_rs::prelude::Hnsw<'static, f32, hnsw_rs::prelude::DistCosine>,
150}
151
152#[cfg(feature = "hnsw_clustering")]
153impl SemanticHnswIndex {
154    fn search(&self, query: &[f32], limit: usize) -> Vec<SemanticCandidate> {
155        if self.photo_ids.is_empty() || limit == 0 {
156            return Vec::new();
157        }
158        if let Some(vectors) = &self.exact_vectors {
159            let mut candidates: Vec<_> = vectors
160                .iter()
161                .zip(&self.photo_ids)
162                .map(|(vector, photo_id)| SemanticCandidate {
163                    photo_id: *photo_id,
164                    score: cosine(query, vector).clamp(-1.0, 1.0),
165                })
166                .collect();
167            candidates.sort_by(|a, b| {
168                b.score
169                    .total_cmp(&a.score)
170                    .then_with(|| a.photo_id.cmp(&b.photo_id))
171            });
172            candidates.truncate(limit);
173            return candidates;
174        }
175        self.hnsw
176            .search(query, limit.min(self.photo_ids.len()).max(1), 200)
177            .into_iter()
178            .filter_map(|nb| {
179                self.photo_ids
180                    .get(nb.d_id)
181                    .map(|photo_id| SemanticCandidate {
182                        photo_id: *photo_id,
183                        score: (1.0 - nb.distance).clamp(-1.0, 1.0),
184                    })
185            })
186            .collect()
187    }
188}
189
190impl SemanticSearchService {
191    pub fn new(drive_root: impl Into<PathBuf>) -> Self {
192        Self {
193            drive_root: drive_root.into(),
194        }
195    }
196
197    pub fn status(&self, conn: &Connection) -> rusqlite::Result<SemanticStatus> {
198        let stats = self.index_stats(conn)?;
199        let store = VectorStore::new(&self.drive_root)?;
200        let assets = Self::find_assets();
201        Ok(SemanticStatus {
202            model_key: SEMANTIC_MODEL_KEY.to_string(),
203            display_name: SEMANTIC_MODEL_DISPLAY.to_string(),
204            model_dir: assets
205                .as_ref()
206                .map(|a| a.root.display().to_string())
207                .unwrap_or_else(|| Self::default_asset_paths().root.display().to_string()),
208            assets_installed: assets.as_ref().is_some_and(SemanticAssetPaths::installed),
209            onnx_runtime_installed: crate::bootstrap::onnx_runtime_exists(),
210            indexed_photos: stats.indexed,
211            pending_photos: stats.pending,
212            failed_photos: stats.failed,
213            vector_bytes: std::fs::metadata(store.vector_path())
214                .map(|m| m.len())
215                .unwrap_or(0),
216        })
217    }
218
219    /// Whether every file needed for local visual search is present.
220    pub fn model_assets_installed() -> bool {
221        Self::find_assets().is_some_and(|paths| paths.installed())
222    }
223
224    pub async fn install_model_assets<F>(
225        cancel: Option<&AtomicBool>,
226        mut progress: F,
227    ) -> Result<(), String>
228    where
229        F: FnMut(&str, u64, Option<u64>) + Send,
230    {
231        let paths = Self::default_asset_paths();
232        let assets = [
233            SemanticDownload {
234                url: VISUAL_MODEL_URL,
235                destination: paths.visual_model,
236                stage: "visual-model",
237                expected_size: VISUAL_MODEL_BYTES,
238            },
239            SemanticDownload {
240                url: TEXTUAL_MODEL_URL,
241                destination: paths.textual_model,
242                stage: "text-model",
243                expected_size: TEXTUAL_MODEL_BYTES,
244            },
245            SemanticDownload {
246                url: TOKENIZER_URL,
247                destination: paths.tokenizer,
248                stage: "tokenizer",
249                expected_size: TOKENIZER_BYTES,
250            },
251            SemanticDownload {
252                url: PREPROCESS_URL,
253                destination: paths.preprocess,
254                stage: "preprocess",
255                expected_size: PREPROCESS_BYTES,
256            },
257            SemanticDownload {
258                url: CONFIG_URL,
259                destination: paths.config,
260                stage: "config",
261                expected_size: CONFIG_BYTES,
262            },
263        ];
264        let total = SEMANTIC_MODEL_DOWNLOAD_BYTES;
265        let mut completed = 0;
266        for asset in assets {
267            completed = download_asset(asset, completed, total, cancel, &mut progress).await?;
268        }
269        Ok(())
270    }
271
272    pub fn index_stats(&self, conn: &Connection) -> rusqlite::Result<SemanticIndexStats> {
273        let indexed = count_state(conn, "indexed")?;
274        let failed = count_state(conn, "failed")?;
275        let total_active: u64 = conn.query_row(
276            "SELECT COUNT(*) FROM photos WHERE is_trashed = FALSE",
277            [],
278            |r| r.get::<_, i64>(0),
279        )? as u64;
280        let pending = total_active.saturating_sub(indexed + failed);
281        Ok(SemanticIndexStats {
282            indexed,
283            pending,
284            failed,
285        })
286    }
287
288    pub fn next_pending_batch(
289        &self,
290        conn: &Connection,
291        limit: usize,
292    ) -> rusqlite::Result<Vec<SemanticPhotoInput>> {
293        let mut stmt = conn.prepare(
294            "SELECT p.id, p.file_path, p.thumbnail_path, p.media_type
295             FROM photos p
296             LEFT JOIN semantic_index_state s
297               ON s.photo_id = p.id AND s.model_key = ?1
298             WHERE p.is_trashed = FALSE
299               AND COALESCE(s.status, 'pending') = 'pending'
300             ORDER BY p.date_taken IS NULL ASC, p.date_taken DESC, p.id DESC
301             LIMIT ?2",
302        )?;
303        let rows = stmt.query_map(params![SEMANTIC_MODEL_KEY, limit as i64], |row| {
304            Ok(SemanticPhotoInput {
305                photo_id: row.get(0)?,
306                file_path: row.get(1)?,
307                thumbnail_path: row.get(2)?,
308                media_type: row.get(3)?,
309            })
310        })?;
311        rows.collect()
312    }
313
314    pub fn mark_failed(conn: &Connection, photo_id: i64, error: &str) -> rusqlite::Result<()> {
315        conn.execute(
316            "INSERT INTO semantic_index_state
317                (photo_id, model_key, status, attempts, last_error)
318             VALUES (?1, ?2, 'failed', 1, ?3)
319             ON CONFLICT(photo_id, model_key) DO UPDATE SET
320                status = 'failed',
321                attempts = attempts + 1,
322                last_error = excluded.last_error,
323                updated_at = CURRENT_TIMESTAMP",
324            params![photo_id, SEMANTIC_MODEL_KEY, truncate_error(error)],
325        )?;
326        Ok(())
327    }
328
329    pub fn mark_indexed(
330        &self,
331        conn: &mut Connection,
332        photo_id: i64,
333        vector: &[f32],
334    ) -> rusqlite::Result<()> {
335        self.record_index_batch(conn, &[(photo_id, vector.to_vec())], &[])
336    }
337
338    pub fn index_next_batch(
339        &self,
340        conn: &mut Connection,
341        runner: &mut SemanticImageRunner,
342        limit: usize,
343        cancel: &AtomicBool,
344    ) -> Result<SemanticIndexBatchOutcome, String> {
345        let batch = self
346            .next_pending_batch(conn, limit)
347            .map_err(|e| e.to_string())?;
348        if batch.is_empty() {
349            return Ok(SemanticIndexBatchOutcome {
350                done: true,
351                ..Default::default()
352            });
353        }
354
355        let mut indexed = Vec::new();
356        let mut failed = Vec::new();
357        for photo in &batch {
358            if cancel.load(Ordering::Relaxed) {
359                return Err("Semantic indexing cancelled".into());
360            }
361            match photo
362                .source_path(&self.drive_root)
363                .and_then(|path| runner.embed_image_path(&path))
364            {
365                Ok(vector) => indexed.push((photo.photo_id, vector)),
366                Err(err) => failed.push((photo.photo_id, err)),
367            }
368        }
369
370        self.record_index_batch(conn, &indexed, &failed)
371            .map_err(|e| e.to_string())?;
372
373        Ok(SemanticIndexBatchOutcome {
374            processed: batch.len() as u64,
375            indexed: indexed.len() as u64,
376            failed: failed.len() as u64,
377            done: false,
378        })
379    }
380
381    fn record_index_batch(
382        &self,
383        conn: &mut Connection,
384        indexed: &[(i64, Vec<f32>)],
385        failed: &[(i64, String)],
386    ) -> rusqlite::Result<()> {
387        let offsets = if indexed.is_empty() {
388            Vec::new()
389        } else {
390            let mut store = VectorStore::new(&self.drive_root)?;
391            store.append_many(indexed.iter().map(|(_, vector)| vector.as_slice()))?
392        };
393
394        let tx = conn.transaction()?;
395        for ((photo_id, vector), offset) in indexed.iter().zip(offsets.iter()) {
396            tx.execute(
397                "INSERT INTO semantic_index_state
398                    (photo_id, model_key, status, vector_offset, vector_dim, attempts, last_error, indexed_at)
399                 VALUES (?1, ?2, 'indexed', ?3, ?4, 0, NULL, CURRENT_TIMESTAMP)
400                 ON CONFLICT(photo_id, model_key) DO UPDATE SET
401                    status = 'indexed',
402                    vector_offset = excluded.vector_offset,
403                    vector_dim = excluded.vector_dim,
404                    attempts = 0,
405                    last_error = NULL,
406                    indexed_at = CURRENT_TIMESTAMP,
407                    updated_at = CURRENT_TIMESTAMP",
408                params![
409                    photo_id,
410                    SEMANTIC_MODEL_KEY,
411                    *offset as i64,
412                    vector.len() as i64
413                ],
414            )?;
415        }
416        for (photo_id, err) in failed {
417            tx.execute(
418                "INSERT INTO semantic_index_state
419                    (photo_id, model_key, status, attempts, last_error)
420                 VALUES (?1, ?2, 'failed', 1, ?3)
421                 ON CONFLICT(photo_id, model_key) DO UPDATE SET
422                    status = 'failed',
423                    attempts = attempts + 1,
424                    last_error = excluded.last_error,
425                    updated_at = CURRENT_TIMESTAMP",
426                params![photo_id, SEMANTIC_MODEL_KEY, truncate_error(err)],
427            )?;
428        }
429        tx.commit()
430    }
431
432    pub fn search_text(
433        &self,
434        conn: &Connection,
435        runner: &mut SemanticModelRunner,
436        query: &str,
437        limit: usize,
438    ) -> Result<Vec<SemanticCandidate>, String> {
439        let vector = runner.embed_text(query)?;
440        self.search_vector(conn, &vector, limit)
441    }
442
443    pub fn search_text_cached(
444        &self,
445        conn: &Connection,
446        cache: &mut SemanticIndexCache,
447        runner: &mut SemanticModelRunner,
448        query: &str,
449        limit: usize,
450    ) -> Result<Vec<SemanticCandidate>, String> {
451        let vector = runner.embed_text(query)?;
452        self.search_vector_cached(conn, cache, &vector, limit)
453    }
454
455    pub fn similar_to_photo(
456        &self,
457        conn: &Connection,
458        photo_id: i64,
459        limit: usize,
460    ) -> Result<Vec<SemanticCandidate>, String> {
461        let Some(vector) = self.vector_for_photo(conn, photo_id)? else {
462            return Ok(Vec::new());
463        };
464        let mut out = self.search_vector(conn, &vector, limit + 1)?;
465        out.retain(|c| c.photo_id != photo_id);
466        out.truncate(limit);
467        Ok(out)
468    }
469
470    pub fn similar_to_photo_cached(
471        &self,
472        conn: &Connection,
473        cache: &mut SemanticIndexCache,
474        photo_id: i64,
475        limit: usize,
476    ) -> Result<Vec<SemanticCandidate>, String> {
477        let Some(vector) = self.vector_for_photo(conn, photo_id)? else {
478            return Ok(Vec::new());
479        };
480        let mut out = self.search_vector_cached(conn, cache, &vector, limit + 1)?;
481        out.retain(|c| c.photo_id != photo_id);
482        out.truncate(limit);
483        Ok(out)
484    }
485
486    pub fn search_vector_cached(
487        &self,
488        conn: &Connection,
489        cache: &mut SemanticIndexCache,
490        query: &[f32],
491        limit: usize,
492    ) -> Result<Vec<SemanticCandidate>, String> {
493        #[cfg(not(feature = "hnsw_clustering"))]
494        {
495            let _ = (conn, cache, query, limit);
496            return Err("HNSW semantic search requires the hnsw_clustering feature".into());
497        }
498
499        #[cfg(feature = "hnsw_clustering")]
500        {
501            if query.len() != SEMANTIC_DIM {
502                return Ok(Vec::new());
503            }
504            let indexed_count = self.index_stats(conn).map_err(|e| e.to_string())?.indexed;
505            if cache.index.is_none() || cache.indexed_count != indexed_count {
506                cache.index = Some(self.build_hnsw_index(conn)?);
507                cache.indexed_count = indexed_count;
508            }
509            Ok(cache
510                .index
511                .as_ref()
512                .map(|idx| idx.search(query, limit))
513                .unwrap_or_default())
514        }
515    }
516
517    pub fn search_vector(
518        &self,
519        conn: &Connection,
520        query: &[f32],
521        limit: usize,
522    ) -> Result<Vec<SemanticCandidate>, String> {
523        #[cfg(not(feature = "hnsw_clustering"))]
524        {
525            let _ = (conn, query, limit);
526            return Err("HNSW semantic search requires the hnsw_clustering feature".into());
527        }
528
529        #[cfg(feature = "hnsw_clustering")]
530        {
531            if query.len() != SEMANTIC_DIM {
532                return Ok(Vec::new());
533            }
534            let index = self.build_hnsw_index(conn)?;
535            Ok(index.search(query, limit))
536        }
537    }
538
539    #[cfg(feature = "hnsw_clustering")]
540    fn build_hnsw_index(&self, conn: &Connection) -> Result<SemanticHnswIndex, String> {
541        use hnsw_rs::prelude::*;
542
543        let rows = self.load_index_rows(conn).map_err(|e| e.to_string())?;
544        if rows.is_empty() {
545            return Ok(SemanticHnswIndex {
546                photo_ids: Vec::new(),
547                exact_vectors: None,
548                hnsw: Hnsw::new(16, 1, 1, 200, DistCosine {}),
549            });
550        }
551
552        let hnsw: Hnsw<f32, DistCosine> = Hnsw::new(
553            16,
554            rows.len(),
555            16.min(rows.len().max(1)),
556            200,
557            DistCosine {},
558        );
559        let data: Vec<(&[f32], usize)> = rows
560            .iter()
561            .enumerate()
562            .map(|(idx, row)| (row.vector.as_slice(), idx))
563            .collect();
564        hnsw.parallel_insert_slice(&data);
565        let exact_vectors =
566            (rows.len() <= 256).then(|| rows.iter().map(|row| row.vector.clone()).collect());
567        let photo_ids = rows.into_iter().map(|row| row.photo_id).collect();
568        Ok(SemanticHnswIndex {
569            photo_ids,
570            exact_vectors,
571            hnsw,
572        })
573    }
574
575    fn vector_for_photo(
576        &self,
577        conn: &Connection,
578        photo_id: i64,
579    ) -> Result<Option<Vec<f32>>, String> {
580        let row: Option<(i64, i64)> = conn
581            .query_row(
582                "SELECT vector_offset, vector_dim
583                 FROM semantic_index_state
584                 WHERE photo_id = ?1 AND model_key = ?2 AND status = 'indexed'",
585                params![photo_id, SEMANTIC_MODEL_KEY],
586                |r| Ok((r.get(0)?, r.get(1)?)),
587            )
588            .optional()
589            .map_err(|e| e.to_string())?;
590        let Some((offset, dim)) = row else {
591            return Ok(None);
592        };
593        if dim != SEMANTIC_DIM as i64 || offset < 0 {
594            return Ok(None);
595        }
596        let store = VectorStore::new(&self.drive_root).map_err(|e| e.to_string())?;
597        store
598            .read(offset as u64, dim as usize)
599            .map(Some)
600            .map_err(|e| e.to_string())
601    }
602
603    fn load_index_rows(&self, conn: &Connection) -> rusqlite::Result<Vec<IndexRow>> {
604        let mut stmt = conn.prepare(
605            "SELECT s.photo_id, s.vector_offset, s.vector_dim
606             FROM semantic_index_state s
607             JOIN photos p ON p.id = s.photo_id
608             WHERE s.model_key = ?1
609               AND s.status = 'indexed'
610               AND s.vector_dim = ?2
611               AND p.is_trashed = FALSE",
612        )?;
613        let rows = stmt.query_map(params![SEMANTIC_MODEL_KEY, SEMANTIC_DIM as i64], |row| {
614            Ok((
615                row.get::<_, i64>(0)?,
616                row.get::<_, i64>(1)?,
617                row.get::<_, i64>(2)?,
618            ))
619        })?;
620        let store = VectorStore::new(&self.drive_root)?;
621        let mut out = Vec::new();
622        for row in rows {
623            let (photo_id, offset, dim) = row?;
624            if let Ok(vector) = store.read(offset as u64, dim as usize) {
625                out.push(IndexRow { photo_id, vector });
626            }
627        }
628        Ok(out)
629    }
630
631    fn find_assets() -> Option<SemanticAssetPaths> {
632        crate::bootstrap::asset_roots()
633            .into_iter()
634            .map(SemanticAssetPaths::in_root)
635            .find(SemanticAssetPaths::installed)
636    }
637
638    pub fn image_runner() -> Result<SemanticImageRunner, String> {
639        let paths = Self::find_assets().ok_or_else(|| {
640            format!(
641                "Semantic search model is not installed. Install {} from Settings.",
642                SEMANTIC_MODEL_DISPLAY
643            )
644        })?;
645        if !crate::bootstrap::onnx_runtime_exists() {
646            return Err(
647                "ONNX Runtime is missing. Use Settings -> Assets -> Download assets before indexing visual search."
648                    .into(),
649            );
650        }
651        let rt = OnnxRuntime::init().map_err(|e| e.to_string())?;
652        SemanticImageRunner::new(&rt, paths)
653    }
654
655    pub fn model_runner() -> Result<SemanticModelRunner, String> {
656        let paths = Self::find_assets().ok_or_else(|| {
657            format!(
658                "Semantic search model is not installed. Install {} from Settings.",
659                SEMANTIC_MODEL_DISPLAY
660            )
661        })?;
662        if !crate::bootstrap::onnx_runtime_exists() {
663            return Err(
664                "ONNX Runtime is missing. Use Settings -> Assets -> Download assets before indexing visual search."
665                    .into(),
666            );
667        }
668        let rt = OnnxRuntime::init().map_err(|e| e.to_string())?;
669        SemanticModelRunner::new(&rt, paths)
670    }
671
672    fn default_asset_paths() -> SemanticAssetPaths {
673        SemanticAssetPaths::in_root(crate::bootstrap::default_asset_install_dir())
674    }
675}
676
677#[derive(Debug, Clone)]
678pub struct SemanticPhotoInput {
679    pub photo_id: i64,
680    pub file_path: String,
681    pub thumbnail_path: Option<String>,
682    pub media_type: String,
683}
684
685impl SemanticPhotoInput {
686    pub fn source_path(&self, drive_root: &Path) -> Result<PathBuf, String> {
687        if let Some(thumbnail) = &self.thumbnail_path {
688            match safe_join_relative(drive_root, thumbnail) {
689                Ok(path) if path.exists() => return Ok(path),
690                Ok(_) if self.media_type == "video" => {
691                    return Err("video poster thumbnail is not ready".into());
692                }
693                Err(e) if self.media_type == "video" => {
694                    return Err(format!("invalid video thumbnail path: {e}"));
695                }
696                _ => {}
697            }
698        }
699        if self.media_type == "video" {
700            return Err("video poster thumbnail is not ready".into());
701        }
702        safe_join_relative(drive_root, &self.file_path)
703            .map_err(|e| format!("invalid photo path: {e}"))
704    }
705}
706
707struct IndexRow {
708    photo_id: i64,
709    vector: Vec<f32>,
710}
711
712pub struct SemanticImageRunner {
713    visual: ort::session::Session,
714}
715
716impl SemanticImageRunner {
717    fn new(rt: &OnnxRuntime, paths: SemanticAssetPaths) -> Result<Self, String> {
718        let visual = rt
719            .load_model_with_threads(&paths.visual_model, 1)
720            .map_err(|e| format!("visual model load failed: {e}"))?;
721        Ok(Self { visual })
722    }
723
724    pub fn embed_image_path(&mut self, path: &Path) -> Result<Vec<f32>, String> {
725        let img = image_io::open_image(path)?;
726        self.embed_image(&img)
727    }
728
729    pub fn embed_image(&mut self, img: &DynamicImage) -> Result<Vec<f32>, String> {
730        let tensor = preprocess_image(img);
731        let input = ort::value::TensorRef::<f32>::from_array_view((
732            vec![1, 3, 256, 256],
733            tensor.as_slice(),
734        ))
735        .map_err(|e| e.to_string())?;
736        let outputs = self
737            .visual
738            .run(ort::inputs![input])
739            .map_err(|e| format!("visual inference failed: {e}"))?;
740        extract_normalized_output(outputs)
741    }
742}
743
744pub struct SemanticModelRunner {
745    textual: ort::session::Session,
746    tokenizer: Tokenizer,
747}
748
749impl SemanticModelRunner {
750    fn new(rt: &OnnxRuntime, paths: SemanticAssetPaths) -> Result<Self, String> {
751        let textual = rt
752            .load_model_with_threads(&paths.textual_model, 1)
753            .map_err(|e| format!("text model load failed: {e}"))?;
754        let tokenizer = Tokenizer::from_file(&paths.tokenizer)
755            .map_err(|e| format!("tokenizer load failed: {e}"))?;
756        Ok(Self { textual, tokenizer })
757    }
758
759    pub fn embed_text(&mut self, text: &str) -> Result<Vec<f32>, String> {
760        let encoding = self
761            .tokenizer
762            .encode(text, true)
763            .map_err(|e| format!("tokenization failed: {e}"))?;
764        let ids = padded_text_context(encoding.get_ids());
765        let input_ids = ort::value::TensorRef::<i32>::from_array_view((
766            vec![1, SEMANTIC_CONTEXT_LEN as i64],
767            ids.as_slice(),
768        ))
769        .map_err(|e| e.to_string())?;
770        let outputs = self
771            .textual
772            .run(ort::inputs![input_ids])
773            .map_err(|e| format!("text inference failed: {e}"))?;
774        extract_normalized_output(outputs)
775    }
776}
777
778fn padded_text_context(token_ids: &[u32]) -> Vec<i32> {
779    let mut ids = vec![0i32; SEMANTIC_CONTEXT_LEN];
780    for (idx, id) in token_ids.iter().take(SEMANTIC_CONTEXT_LEN).enumerate() {
781        ids[idx] = *id as i32;
782    }
783    ids
784}
785
786fn preprocess_image(img: &DynamicImage) -> Vec<f32> {
787    let resized = img.resize_exact(256, 256, image::imageops::FilterType::CatmullRom);
788    let rgb: ImageBuffer<Rgb<u8>, Vec<u8>> = resized.to_rgb8();
789    let mut out = vec![0.0f32; 3 * 256 * 256];
790    let hw = 256 * 256;
791    for y in 0..256u32 {
792        for x in 0..256u32 {
793            let p = rgb.get_pixel(x, y);
794            let idx = (y * 256 + x) as usize;
795            out[idx] = (p[0] as f32 / 255.0 - 0.5) / 0.5;
796            out[hw + idx] = (p[1] as f32 / 255.0 - 0.5) / 0.5;
797            out[2 * hw + idx] = (p[2] as f32 / 255.0 - 0.5) / 0.5;
798        }
799    }
800    out
801}
802
803fn extract_normalized_output(outputs: ort::session::SessionOutputs) -> Result<Vec<f32>, String> {
804    let (_name, output) = outputs
805        .iter()
806        .next()
807        .ok_or_else(|| "model produced no output tensor".to_string())?;
808    let (_shape, data) = output
809        .try_extract_tensor::<f32>()
810        .map_err(|e| e.to_string())?;
811    let mut vector = data.to_vec();
812    if vector.len() != SEMANTIC_DIM {
813        return Err(format!(
814            "unexpected semantic embedding dimension: expected {}, got {}",
815            SEMANTIC_DIM,
816            vector.len()
817        ));
818    }
819    normalize_in_place(&mut vector);
820    Ok(vector)
821}
822
823fn normalize_in_place(v: &mut [f32]) {
824    let norm = v.iter().map(|x| x * x).sum::<f32>().sqrt();
825    if norm > 0.0 {
826        for x in v {
827            *x /= norm;
828        }
829    }
830}
831
832struct VectorStore {
833    root: PathBuf,
834}
835
836impl VectorStore {
837    fn new(drive_root: &Path) -> rusqlite::Result<Self> {
838        let root = library_metadata_dir(drive_root)
839            .join("semantic")
840            .join(MODEL_DIR_NAME);
841        std::fs::create_dir_all(&root)
842            .map_err(|e| rusqlite::Error::ToSqlConversionFailure(Box::new(e)))?;
843        let manifest = root.join(MANIFEST_FILE);
844        if !manifest.exists() {
845            let data = serde_json::to_vec_pretty(&VectorManifest {
846                model_key: SEMANTIC_MODEL_KEY.to_string(),
847                revision: SEMANTIC_MODEL_REVISION.to_string(),
848                dim: SEMANTIC_DIM,
849                vector_count: 0,
850            })
851            .map_err(|e| rusqlite::Error::ToSqlConversionFailure(Box::new(e)))?;
852            std::fs::write(&manifest, data)
853                .map_err(|e| rusqlite::Error::ToSqlConversionFailure(Box::new(e)))?;
854        }
855        Ok(Self { root })
856    }
857
858    fn vector_path(&self) -> PathBuf {
859        self.root.join(VECTOR_FILE)
860    }
861
862    fn append_many<'a, I>(&mut self, vectors: I) -> rusqlite::Result<Vec<u64>>
863    where
864        I: IntoIterator<Item = &'a [f32]>,
865    {
866        let vectors = vectors.into_iter().collect::<Vec<_>>();
867        if vectors.is_empty() {
868            return Ok(Vec::new());
869        }
870        for vector in &vectors {
871            if vector.len() != SEMANTIC_DIM {
872                return Err(rusqlite::Error::InvalidParameterName(format!(
873                    "semantic vector dimension {} != {}",
874                    vector.len(),
875                    SEMANTIC_DIM
876                )));
877            }
878        }
879        let path = self.vector_path();
880        let mut file = OpenOptions::new()
881            .create(true)
882            .append(true)
883            .read(true)
884            .open(&path)
885            .map_err(|e| rusqlite::Error::ToSqlConversionFailure(Box::new(e)))?;
886        let mut offset = file
887            .seek(SeekFrom::End(0))
888            .map_err(|e| rusqlite::Error::ToSqlConversionFailure(Box::new(e)))?;
889        let mut offsets = Vec::with_capacity(vectors.len());
890        for vector in vectors {
891            offsets.push(offset);
892            for value in vector {
893                file.write_all(&value.to_le_bytes())
894                    .map_err(|e| rusqlite::Error::ToSqlConversionFailure(Box::new(e)))?;
895            }
896            offset += (SEMANTIC_DIM * 4) as u64;
897        }
898        self.bump_manifest_by(offsets.len() as u64)?;
899        Ok(offsets)
900    }
901
902    fn read(&self, offset: u64, dim: usize) -> std::io::Result<Vec<f32>> {
903        let mut file = File::open(self.vector_path())?;
904        file.seek(SeekFrom::Start(offset))?;
905        let mut bytes = vec![0u8; dim * 4];
906        file.read_exact(&mut bytes)?;
907        Ok(bytes
908            .chunks_exact(4)
909            .map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
910            .collect())
911    }
912
913    fn bump_manifest_by(&self, count: u64) -> rusqlite::Result<()> {
914        let path = self.root.join(MANIFEST_FILE);
915        let mut manifest: VectorManifest = std::fs::read(&path)
916            .ok()
917            .and_then(|b| serde_json::from_slice(&b).ok())
918            .unwrap_or(VectorManifest {
919                model_key: SEMANTIC_MODEL_KEY.to_string(),
920                revision: SEMANTIC_MODEL_REVISION.to_string(),
921                dim: SEMANTIC_DIM,
922                vector_count: 0,
923            });
924        manifest.vector_count += count;
925        let data = serde_json::to_vec_pretty(&manifest)
926            .map_err(|e| rusqlite::Error::ToSqlConversionFailure(Box::new(e)))?;
927        std::fs::write(path, data)
928            .map_err(|e| rusqlite::Error::ToSqlConversionFailure(Box::new(e)))?;
929        Ok(())
930    }
931}
932
933fn count_state(conn: &Connection, status: &str) -> rusqlite::Result<u64> {
934    conn.query_row(
935        "SELECT COUNT(*)
936         FROM semantic_index_state s
937         JOIN photos p ON p.id = s.photo_id
938         WHERE s.model_key = ?1
939           AND s.status = ?2
940           AND p.is_trashed = FALSE",
941        params![SEMANTIC_MODEL_KEY, status],
942        |r| r.get::<_, i64>(0),
943    )
944    .map(|v| v as u64)
945}
946
947fn truncate_error(error: &str) -> String {
948    error.chars().take(500).collect()
949}
950
951struct SemanticDownload {
952    url: &'static str,
953    stage: &'static str,
954    destination: PathBuf,
955    expected_size: u64,
956}
957
958async fn download_asset<F>(
959    asset: SemanticDownload,
960    completed_before: u64,
961    total_bytes: u64,
962    cancel: Option<&AtomicBool>,
963    progress: &mut F,
964) -> Result<u64, String>
965where
966    F: FnMut(&str, u64, Option<u64>) + Send,
967{
968    if asset
969        .destination
970        .metadata()
971        .is_ok_and(|metadata| metadata.is_file() && metadata.len() == asset.expected_size)
972    {
973        let completed = completed_before + asset.expected_size;
974        progress(asset.stage, completed, Some(total_bytes));
975        return Ok(completed);
976    }
977    if asset.destination.exists() {
978        tokio::fs::remove_file(&asset.destination)
979            .await
980            .map_err(|e| format!("failed replacing {}: {e}", asset.destination.display()))?;
981    }
982    if let Some(parent) = asset.destination.parent() {
983        tokio::fs::create_dir_all(parent)
984            .await
985            .map_err(|e| format!("failed creating {}: {e}", parent.display()))?;
986    }
987    if cancel.is_some_and(|flag| flag.load(Ordering::Relaxed)) {
988        return Err("Semantic model install cancelled".into());
989    }
990
991    let response = reqwest::get(asset.url)
992        .await
993        .map_err(|e| format!("download request failed for {}: {e}", asset.url))?;
994    if !response.status().is_success() {
995        return Err(format!(
996            "download failed for {}: HTTP {}",
997            asset.url,
998            response.status()
999        ));
1000    }
1001    let expected = response.content_length().unwrap_or(asset.expected_size);
1002    let tmp = asset.destination.with_extension("tmp");
1003    let mut file = tokio::fs::File::create(&tmp)
1004        .await
1005        .map_err(|e| format!("failed writing {}: {e}", tmp.display()))?;
1006    let mut downloaded = 0u64;
1007    let mut last_emit = 0u64;
1008    let mut stream = response.bytes_stream();
1009
1010    use futures::StreamExt;
1011    while let Some(chunk) = stream.next().await {
1012        if cancel.is_some_and(|flag| flag.load(Ordering::Relaxed)) {
1013            let _ = tokio::fs::remove_file(&tmp).await;
1014            return Err("Semantic model install cancelled".into());
1015        }
1016        let chunk = chunk.map_err(|e| format!("download body failed for {}: {e}", asset.url))?;
1017        tokio::io::AsyncWriteExt::write_all(&mut file, &chunk)
1018            .await
1019            .map_err(|e| format!("failed writing {}: {e}", tmp.display()))?;
1020        downloaded += chunk.len() as u64;
1021        if downloaded.saturating_sub(last_emit) >= 1_048_576 || downloaded >= expected {
1022            last_emit = downloaded;
1023            progress(
1024                asset.stage,
1025                completed_before + downloaded.min(asset.expected_size),
1026                Some(total_bytes),
1027            );
1028        }
1029    }
1030    tokio::io::AsyncWriteExt::flush(&mut file)
1031        .await
1032        .map_err(|e| format!("failed flushing {}: {e}", tmp.display()))?;
1033    drop(file);
1034    if downloaded != asset.expected_size {
1035        let _ = tokio::fs::remove_file(&tmp).await;
1036        return Err(format!(
1037            "download size mismatch for {}: expected {} bytes, got {}",
1038            asset.url, asset.expected_size, downloaded
1039        ));
1040    }
1041    tokio::fs::rename(&tmp, &asset.destination)
1042        .await
1043        .map_err(|e| format!("failed moving {}: {e}", asset.destination.display()))?;
1044    let completed = completed_before + asset.expected_size;
1045    progress(asset.stage, completed, Some(total_bytes));
1046    Ok(completed)
1047}
1048pub fn semantic_ids_by_score(candidates: &[SemanticCandidate]) -> HashMap<i64, usize> {
1049    candidates
1050        .iter()
1051        .enumerate()
1052        .map(|(idx, c)| (c.photo_id, idx))
1053        .collect()
1054}
1055
1056pub fn relevant_text_search_candidates(
1057    mut candidates: Vec<SemanticCandidate>,
1058) -> Vec<SemanticCandidate> {
1059    candidates.sort_by(|a, b| {
1060        b.score
1061            .partial_cmp(&a.score)
1062            .unwrap_or(std::cmp::Ordering::Equal)
1063    });
1064    let Some(top) = candidates.first().map(|c| c.score) else {
1065        return Vec::new();
1066    };
1067    if top < SEMANTIC_TEXT_MIN_SCORE {
1068        return Vec::new();
1069    }
1070
1071    let threshold = SEMANTIC_TEXT_MIN_SCORE
1072        .max(top - SEMANTIC_TEXT_MAX_SCORE_DROP)
1073        .max(top * SEMANTIC_TEXT_MIN_SCORE_RATIO);
1074    candidates
1075        .into_iter()
1076        .filter(|c| c.score >= threshold)
1077        .take(SEMANTIC_TEXT_RESULT_CAP)
1078        .collect()
1079}
1080
1081pub fn cosine(a: &[f32], b: &[f32]) -> f32 {
1082    let av = Array1::from_vec(a.to_vec());
1083    let bv = Array1::from_vec(b.to_vec());
1084    let dot = av.dot(&bv);
1085    let na = av.dot(&av).sqrt();
1086    let nb = bv.dot(&bv).sqrt();
1087    if na > 0.0 && nb > 0.0 {
1088        dot / (na * nb)
1089    } else {
1090        0.0
1091    }
1092}
1093
1094#[cfg(test)]
1095mod tests {
1096    use super::*;
1097    use tempfile::tempdir;
1098
1099    fn setup_semantic_test_conn() -> Connection {
1100        let conn = Connection::open_in_memory().unwrap();
1101        conn.execute_batch(
1102            "CREATE TABLE photos (
1103                id INTEGER PRIMARY KEY,
1104                file_path TEXT NOT NULL,
1105                thumbnail_path TEXT,
1106                media_type TEXT NOT NULL DEFAULT 'photo',
1107                date_taken TEXT,
1108                is_trashed BOOLEAN NOT NULL DEFAULT FALSE
1109            );
1110            CREATE TABLE semantic_index_state (
1111                photo_id INTEGER NOT NULL,
1112                model_key TEXT NOT NULL,
1113                status TEXT NOT NULL DEFAULT 'pending',
1114                vector_offset INTEGER,
1115                vector_dim INTEGER,
1116                attempts INTEGER NOT NULL DEFAULT 0,
1117                last_error TEXT,
1118                indexed_at TEXT,
1119                updated_at TEXT,
1120                PRIMARY KEY(photo_id, model_key)
1121            );",
1122        )
1123        .unwrap();
1124        conn
1125    }
1126
1127    #[test]
1128    fn vector_store_round_trips_fixed_width_vectors() {
1129        let dir = tempdir().unwrap();
1130        let mut store = VectorStore::new(dir.path()).unwrap();
1131        let mut first = vec![0.0f32; SEMANTIC_DIM];
1132        first[3] = 1.0;
1133        let mut second = vec![0.0f32; SEMANTIC_DIM];
1134        second[9] = 1.0;
1135
1136        let offsets = store
1137            .append_many([first.as_slice(), second.as_slice()])
1138            .unwrap();
1139        let off_a = offsets[0];
1140        let off_b = offsets[1];
1141
1142        assert_eq!(off_a, 0);
1143        assert_eq!(off_b, (SEMANTIC_DIM * 4) as u64);
1144        assert_eq!(store.read(off_a, SEMANTIC_DIM).unwrap(), first);
1145        assert_eq!(store.read(off_b, SEMANTIC_DIM).unwrap(), second);
1146    }
1147
1148    #[test]
1149    fn vector_for_photo_ignores_corrupt_vector_dimension() {
1150        let conn = setup_semantic_test_conn();
1151        conn.execute(
1152            "INSERT INTO photos (id, file_path, media_type, is_trashed) VALUES
1153                (1, 'a.jpg', 'photo', FALSE)",
1154            [],
1155        )
1156        .unwrap();
1157        conn.execute(
1158            "INSERT INTO semantic_index_state
1159                (photo_id, model_key, status, vector_offset, vector_dim)
1160             VALUES (?1, ?2, 'indexed', 0, 999999999)",
1161            rusqlite::params![1_i64, SEMANTIC_MODEL_KEY],
1162        )
1163        .unwrap();
1164        let svc = SemanticSearchService::new(tempdir().unwrap().path());
1165
1166        assert!(svc.vector_for_photo(&conn, 1).unwrap().is_none());
1167    }
1168
1169    #[test]
1170    fn pending_batch_does_not_retry_failed_rows() {
1171        let conn = setup_semantic_test_conn();
1172        conn.execute(
1173            "INSERT INTO photos (id, file_path, media_type, is_trashed) VALUES
1174                (1, 'a.jpg', 'photo', FALSE),
1175                (2, 'b.jpg', 'photo', FALSE)",
1176            [],
1177        )
1178        .unwrap();
1179        SemanticSearchService::mark_failed(&conn, 1, "bad image").unwrap();
1180
1181        let svc = SemanticSearchService::new(tempdir().unwrap().path());
1182        let batch = svc.next_pending_batch(&conn, 10).unwrap();
1183
1184        assert_eq!(
1185            batch.iter().map(|p| p.photo_id).collect::<Vec<_>>(),
1186            vec![2]
1187        );
1188    }
1189
1190    #[test]
1191    fn index_stats_ignore_trashed_index_state_rows() {
1192        let conn = setup_semantic_test_conn();
1193        conn.execute(
1194            "INSERT INTO photos (id, file_path, media_type, is_trashed) VALUES
1195                (1, 'a.jpg', 'photo', FALSE),
1196                (2, 'b.jpg', 'photo', TRUE),
1197                (3, 'c.jpg', 'photo', FALSE)",
1198            [],
1199        )
1200        .unwrap();
1201        conn.execute(
1202            "INSERT INTO semantic_index_state
1203                (photo_id, model_key, status, vector_offset, vector_dim)
1204             VALUES
1205                (1, ?1, 'indexed', 0, ?2),
1206                (2, ?1, 'indexed', 0, ?2),
1207                (3, ?1, 'failed', NULL, NULL)",
1208            params![SEMANTIC_MODEL_KEY, SEMANTIC_DIM as i64],
1209        )
1210        .unwrap();
1211
1212        let svc = SemanticSearchService::new(tempdir().unwrap().path());
1213        let stats = svc.index_stats(&conn).unwrap();
1214
1215        assert_eq!(stats.indexed, 1);
1216        assert_eq!(stats.failed, 1);
1217        assert_eq!(stats.pending, 0);
1218    }
1219
1220    #[test]
1221    fn record_index_batch_persists_vectors_and_failures_once() {
1222        let mut conn = setup_semantic_test_conn();
1223        conn.execute(
1224            "INSERT INTO photos (id, file_path, media_type, is_trashed) VALUES
1225                (1, 'a.jpg', 'photo', FALSE),
1226                (2, 'b.jpg', 'photo', FALSE),
1227                (3, 'c.jpg', 'photo', FALSE)",
1228            [],
1229        )
1230        .unwrap();
1231        let dir = tempdir().unwrap();
1232        let svc = SemanticSearchService::new(dir.path());
1233        let mut first = vec![0.0f32; SEMANTIC_DIM];
1234        first[0] = 1.0;
1235        let mut second = vec![0.0f32; SEMANTIC_DIM];
1236        second[1] = 1.0;
1237
1238        svc.record_index_batch(
1239            &mut conn,
1240            &[(1, first.clone()), (2, second.clone())],
1241            &[(3, "decode failed".into())],
1242        )
1243        .unwrap();
1244
1245        let rows = conn
1246            .prepare(
1247                "SELECT photo_id, status, vector_offset, vector_dim, attempts, COALESCE(last_error, '')
1248                 FROM semantic_index_state
1249                 ORDER BY photo_id",
1250            )
1251            .unwrap()
1252            .query_map([], |row| {
1253                Ok((
1254                    row.get::<_, i64>(0)?,
1255                    row.get::<_, String>(1)?,
1256                    row.get::<_, Option<i64>>(2)?,
1257                    row.get::<_, Option<i64>>(3)?,
1258                    row.get::<_, i64>(4)?,
1259                    row.get::<_, String>(5)?,
1260                ))
1261            })
1262            .unwrap()
1263            .collect::<rusqlite::Result<Vec<_>>>()
1264            .unwrap();
1265
1266        assert_eq!(rows.len(), 3);
1267        assert_eq!(rows[0].0, 1);
1268        assert_eq!(rows[0].1, "indexed");
1269        assert_eq!(rows[0].2, Some(0));
1270        assert_eq!(rows[0].3, Some(SEMANTIC_DIM as i64));
1271        assert_eq!(rows[1].0, 2);
1272        assert_eq!(rows[1].1, "indexed");
1273        assert_eq!(rows[1].2, Some((SEMANTIC_DIM * 4) as i64));
1274        assert_eq!(rows[2].0, 3);
1275        assert_eq!(rows[2].1, "failed");
1276        assert_eq!(rows[2].4, 1);
1277        assert_eq!(rows[2].5, "decode failed");
1278
1279        let store = VectorStore::new(dir.path()).unwrap();
1280        assert_eq!(store.read(0, SEMANTIC_DIM).unwrap(), first);
1281        assert_eq!(
1282            store.read((SEMANTIC_DIM * 4) as u64, SEMANTIC_DIM).unwrap(),
1283            second
1284        );
1285        assert_eq!(
1286            std::fs::metadata(store.vector_path()).unwrap().len(),
1287            (2 * SEMANTIC_DIM * 4) as u64
1288        );
1289    }
1290
1291    #[test]
1292    fn photo_source_prefers_existing_thumbnail() {
1293        let dir = tempdir().unwrap();
1294        std::fs::create_dir_all(dir.path().join(".photovault/thumbs")).unwrap();
1295        std::fs::write(dir.path().join("photo.jpg"), b"original").unwrap();
1296        std::fs::write(dir.path().join(".photovault/thumbs/photo.jpg"), b"thumb").unwrap();
1297        let input = SemanticPhotoInput {
1298            photo_id: 1,
1299            file_path: "photo.jpg".into(),
1300            thumbnail_path: Some(".photovault/thumbs/photo.jpg".into()),
1301            media_type: "photo".into(),
1302        };
1303
1304        assert_eq!(
1305            input.source_path(dir.path()).unwrap(),
1306            dir.path().join(".photovault/thumbs/photo.jpg")
1307        );
1308    }
1309
1310    #[test]
1311    fn search_vector_returns_indexed_candidates() {
1312        let mut conn = setup_semantic_test_conn();
1313        conn.execute(
1314            "INSERT INTO photos (id, file_path, media_type, is_trashed) VALUES
1315                (1, 'a.jpg', 'photo', FALSE),
1316                (2, 'b.jpg', 'photo', FALSE),
1317                (3, 'c.jpg', 'photo', TRUE)",
1318            [],
1319        )
1320        .unwrap();
1321        let dir = tempdir().unwrap();
1322        let svc = SemanticSearchService::new(dir.path());
1323        let mut first = vec![0.0f32; SEMANTIC_DIM];
1324        first[0] = 1.0;
1325        let mut second = vec![0.0f32; SEMANTIC_DIM];
1326        second[1] = 1.0;
1327        let mut trashed = vec![0.0f32; SEMANTIC_DIM];
1328        trashed[0] = 1.0;
1329
1330        svc.record_index_batch(
1331            &mut conn,
1332            &[(1, first.clone()), (2, second), (3, trashed)],
1333            &[],
1334        )
1335        .unwrap();
1336
1337        let matches = svc.search_vector(&conn, &first, 5).unwrap();
1338
1339        assert_eq!(matches.first().map(|c| c.photo_id), Some(1));
1340        assert!(!matches.iter().any(|c| c.photo_id == 3));
1341    }
1342
1343    #[test]
1344    fn cosine_handles_normal_vectors() {
1345        assert!((cosine(&[1.0, 0.0], &[1.0, 0.0]) - 1.0).abs() < 0.001);
1346        assert!(cosine(&[1.0, 0.0], &[0.0, 1.0]).abs() < 0.001);
1347    }
1348
1349    #[test]
1350    fn text_context_is_fixed_width_int32_and_padded() {
1351        let ids = padded_text_context(&[2, 101, 102, 1]);
1352
1353        assert_eq!(ids.len(), SEMANTIC_CONTEXT_LEN);
1354        assert_eq!(&ids[..5], &[2, 101, 102, 1, 0]);
1355
1356        let long = (0..(SEMANTIC_CONTEXT_LEN as u32 + 10)).collect::<Vec<_>>();
1357        let truncated = padded_text_context(&long);
1358        assert_eq!(truncated.len(), SEMANTIC_CONTEXT_LEN);
1359        assert_eq!(truncated[0], 0);
1360        assert_eq!(
1361            truncated[SEMANTIC_CONTEXT_LEN - 1],
1362            (SEMANTIC_CONTEXT_LEN - 1) as i32
1363        );
1364    }
1365
1366    #[test]
1367    fn text_search_gate_rejects_weak_absent_queries() {
1368        let kept = relevant_text_search_candidates(vec![
1369            SemanticCandidate {
1370                photo_id: 1,
1371                score: 0.035,
1372            },
1373            SemanticCandidate {
1374                photo_id: 2,
1375                score: 0.030,
1376            },
1377        ]);
1378
1379        assert!(kept.is_empty());
1380    }
1381
1382    #[test]
1383    fn text_search_gate_keeps_only_standout_matches() {
1384        let kept = relevant_text_search_candidates(vec![
1385            SemanticCandidate {
1386                photo_id: 1,
1387                score: 0.095,
1388            },
1389            SemanticCandidate {
1390                photo_id: 2,
1391                score: 0.070,
1392            },
1393            SemanticCandidate {
1394                photo_id: 3,
1395                score: 0.040,
1396            },
1397        ]);
1398
1399        assert_eq!(kept.iter().map(|c| c.photo_id).collect::<Vec<_>>(), vec![1]);
1400    }
1401
1402    #[test]
1403    fn text_search_gate_keeps_dense_relevant_clusters() {
1404        let kept = relevant_text_search_candidates(vec![
1405            SemanticCandidate {
1406                photo_id: 1,
1407                score: 0.078,
1408            },
1409            SemanticCandidate {
1410                photo_id: 2,
1411                score: 0.074,
1412            },
1413            SemanticCandidate {
1414                photo_id: 3,
1415                score: 0.048,
1416            },
1417        ]);
1418
1419        assert_eq!(
1420            kept.iter().map(|c| c.photo_id).collect::<Vec<_>>(),
1421            vec![1, 2]
1422        );
1423    }
1424}