Skip to main content

remem/retrieval/
vector.rs

1use anyhow::{Context, Result};
2use rusqlite::{params, Connection, OptionalExtension, Statement};
3use std::time::Instant;
4
5use super::embedding::TextEmbedding;
6pub use super::vector_candidates::VECTOR_SEARCH_CANDIDATE_LIMIT;
7
8mod coverage;
9mod reindex;
10
11pub use super::embedding::{
12    LOCAL_EMBEDDING_DIMENSIONS as EMBEDDING_DIMENSIONS,
13    LOCAL_EMBEDDING_MODEL as DEFAULT_EMBEDDING_MODEL,
14};
15pub use coverage::{
16    active_embedding_coverage, active_embedding_coverage_for_status,
17    prune_inactive_memory_embeddings, ActiveEmbeddingCoverage, InactiveEmbeddingPruneReport,
18};
19use reindex::select_memory_embedding_reindex_candidates;
20
21const EMBEDDING_REINDEX_WRITE_BATCH_SIZE: usize = 512;
22const UPSERT_EMBEDDING_SQL: &str = "INSERT INTO memory_embeddings
23         (memory_id, embedding, dimensions, model, content_hash, updated_at_epoch)
24         VALUES (?1, ?2, ?3, ?4, ?5, ?6)
25         ON CONFLICT(memory_id, model, dimensions) DO UPDATE SET
26             embedding = excluded.embedding,
27             content_hash = excluded.content_hash,
28             updated_at_epoch = excluded.updated_at_epoch";
29
30#[derive(Debug, Clone, PartialEq)]
31pub struct VectorHit {
32    pub memory_id: i64,
33    pub distance: f32,
34}
35
36#[derive(Debug, Clone, PartialEq)]
37pub struct VectorSearchOutcome {
38    pub hits: Vec<VectorHit>,
39    pub disabled_reason: Option<String>,
40    pub candidates_scanned: usize,
41    pub timings: Vec<crate::perf::PhaseTiming>,
42}
43
44impl VectorSearchOutcome {
45    pub fn disabled(reason: impl Into<String>) -> Self {
46        Self {
47            hits: vec![],
48            disabled_reason: Some(reason.into()),
49            candidates_scanned: 0,
50            timings: vec![],
51        }
52    }
53
54    fn disabled_with_timings(
55        reason: impl Into<String>,
56        timings: Vec<crate::perf::PhaseTiming>,
57    ) -> Self {
58        Self {
59            hits: vec![],
60            disabled_reason: Some(reason.into()),
61            candidates_scanned: 0,
62            timings,
63        }
64    }
65
66    pub fn ready(hits: Vec<VectorHit>) -> Self {
67        let candidates_scanned = hits.len();
68        Self::ready_with_scan_count(hits, candidates_scanned)
69    }
70
71    pub fn ready_with_scan_count(hits: Vec<VectorHit>, candidates_scanned: usize) -> Self {
72        Self::ready_with_scan_count_and_timings(hits, candidates_scanned, vec![])
73    }
74
75    fn ready_with_scan_count_and_timings(
76        hits: Vec<VectorHit>,
77        candidates_scanned: usize,
78        timings: Vec<crate::perf::PhaseTiming>,
79    ) -> Self {
80        Self {
81            hits,
82            disabled_reason: None,
83            candidates_scanned,
84            timings,
85        }
86    }
87}
88
89#[derive(Debug, Clone, Copy, Default)]
90pub struct VectorSearchFilters<'a> {
91    pub project: Option<&'a str>,
92    pub memory_type: Option<&'a str>,
93    pub branch: Option<&'a str>,
94    pub include_stale: bool,
95}
96
97/// Load a native vector extension when one is configured.
98///
99/// The current production path is a portable SQLite table plus in-process
100/// cosine scan. That keeps vector recall available in the single-binary build;
101/// sqlite-vec can replace the scan later without changing the search contract.
102pub fn load_vec_extension(_conn: &Connection) -> Result<()> {
103    Ok(())
104}
105
106pub fn ensure_vec_table(conn: &Connection) -> Result<()> {
107    create_embedding_table(conn)
108}
109
110pub fn upsert_embedding(conn: &Connection, memory_id: i64, embedding: &[f32]) -> Result<()> {
111    if super::embedding::provider_disabled_or_error()? {
112        return Ok(());
113    }
114    upsert_embedding_with_metadata(
115        conn,
116        memory_id,
117        DEFAULT_EMBEDDING_MODEL,
118        "",
119        embedding,
120        chrono::Utc::now().timestamp(),
121    )
122}
123
124pub fn upsert_memory_embedding(
125    conn: &Connection,
126    memory_id: i64,
127    title: &str,
128    content: &str,
129    memory_type: &str,
130    topic_key: Option<&str>,
131) -> Result<()> {
132    if super::embedding::provider_disabled_or_error()? {
133        return Ok(());
134    }
135    let embedding = match super::embedding::embed_memory(title, content, memory_type, topic_key) {
136        Ok(embedding) => embedding,
137        Err(error) if super::embedding::is_embedding_provider_off_error(&error) => return Ok(()),
138        Err(error) if super::embedding::is_local_embedding_model_unavailable_error(&error) => {
139            crate::log::error(
140                "embedding",
141                &format!("memory embedding deferred for memory id={memory_id}: {error}"),
142            );
143            return Ok(());
144        }
145        Err(error) => return Err(error),
146    };
147    let content_hash =
148        super::embedding::embedding_content_hash(title, content, memory_type, topic_key);
149    upsert_embedding_with_metadata(
150        conn,
151        memory_id,
152        embedding.model(),
153        &content_hash,
154        embedding.values(),
155        chrono::Utc::now().timestamp(),
156    )
157    .with_context(|| format!("memory embedding upsert failed for memory id={memory_id}"))
158}
159
160pub fn upsert_memory_embedding_for_row(conn: &Connection, memory_id: i64) -> Result<()> {
161    let (topic_key, title, content, memory_type): (Option<String>, String, String, String) = conn
162        .query_row(
163            "SELECT topic_key, title, content, memory_type
164             FROM memories
165             WHERE id = ?1",
166            [memory_id],
167            |row| Ok((row.get(0)?, row.get(1)?, row.get(2)?, row.get(3)?)),
168        )
169        .with_context(|| format!("load memory row for embedding id={memory_id}"))?;
170    upsert_memory_embedding(
171        conn,
172        memory_id,
173        &title,
174        &content,
175        &memory_type,
176        topic_key.as_deref(),
177    )
178}
179
180pub fn backfill_missing_memory_embeddings(conn: &Connection, limit: i64) -> Result<usize> {
181    reindex_memory_embeddings(conn, limit)
182}
183
184#[derive(Debug, Clone, PartialEq)]
185pub struct EmbeddingReindexReport {
186    pub selected: usize,
187    pub processed: usize,
188    pub model: String,
189    pub dimensions: usize,
190    pub timings: Vec<crate::perf::PhaseTiming>,
191}
192
193pub fn reindex_memory_embeddings(conn: &Connection, limit: i64) -> Result<usize> {
194    let mut remaining_limit = limit.max(0);
195    let mut processed = 0usize;
196    while remaining_limit > 0 {
197        let batch_limit = remaining_limit.min(EMBEDDING_REINDEX_WRITE_BATCH_SIZE as i64);
198        let report = reindex_memory_embeddings_with_report(conn, batch_limit)?;
199        if report.processed == 0 {
200            break;
201        }
202        processed += report.processed;
203        remaining_limit -= report.processed as i64;
204        if report.processed < batch_limit as usize {
205            break;
206        }
207    }
208    Ok(processed)
209}
210
211pub fn reindex_memory_embeddings_with_report(
212    conn: &Connection,
213    limit: i64,
214) -> Result<EmbeddingReindexReport> {
215    let total_start = Instant::now();
216    let mut timings = vec![];
217    if super::embedding::provider_disabled_or_error()? {
218        crate::perf::push_elapsed(&mut timings, "total", total_start);
219        return Ok(EmbeddingReindexReport {
220            selected: 0,
221            processed: 0,
222            model: "off".to_string(),
223            dimensions: 0,
224            timings,
225        });
226    }
227    if !table_exists(conn, "memories")? || !table_exists(conn, "memory_embeddings")? {
228        crate::perf::push_elapsed(&mut timings, "total", total_start);
229        return Ok(EmbeddingReindexReport {
230            selected: 0,
231            processed: 0,
232            model: String::new(),
233            dimensions: 0,
234            timings,
235        });
236    }
237    let limit = limit.max(0);
238    if limit == 0 {
239        crate::perf::push_elapsed(&mut timings, "total", total_start);
240        return Ok(EmbeddingReindexReport {
241            selected: 0,
242            processed: 0,
243            model: String::new(),
244            dimensions: 0,
245            timings,
246        });
247    }
248
249    let profile_start = Instant::now();
250    let mut fallback_cache = super::embedding::EmbeddingFallbackCache::default();
251    let mut target =
252        match super::embedding::configured_backfill_target_with_fallback_cache(&mut fallback_cache)
253        {
254            Ok(target) => target,
255            Err(error) if super::embedding::is_embedding_provider_off_error(&error) => {
256                crate::perf::push_elapsed(&mut timings, "total", total_start);
257                return Ok(EmbeddingReindexReport {
258                    selected: 0,
259                    processed: 0,
260                    model: "off".to_string(),
261                    dimensions: 0,
262                    timings,
263                });
264            }
265            Err(error) => return Err(error),
266        };
267    crate::perf::push_elapsed(&mut timings, "profile_probe", profile_start);
268
269    let select_start = Instant::now();
270    let mut pending = select_memory_embedding_reindex_candidates(conn, &target, limit)?;
271    crate::perf::push_elapsed(&mut timings, "select_pending", select_start);
272
273    let mut selected = pending.len();
274    if pending.is_empty() {
275        crate::perf::push_elapsed(&mut timings, "total", total_start);
276        return Ok(EmbeddingReindexReport {
277            selected,
278            processed: 0,
279            model: target.model,
280            dimensions: target.dimensions,
281            timings,
282        });
283    }
284
285    let mut prepared = prepare_memory_embedding_batch(&pending, &mut timings, &mut fallback_cache)?;
286    if let Some(fallback_target) = fallback_cache.call_failure_fallback_target() {
287        if fallback_target != target {
288            target = fallback_target;
289            let fallback_select_start = Instant::now();
290            pending = select_memory_embedding_reindex_candidates(conn, &target, limit)?;
291            crate::perf::push_elapsed(
292                &mut timings,
293                "select_pending_after_fallback",
294                fallback_select_start,
295            );
296            selected = pending.len();
297            if pending.is_empty() {
298                crate::perf::push_elapsed(&mut timings, "total", total_start);
299                return Ok(EmbeddingReindexReport {
300                    selected,
301                    processed: 0,
302                    model: target.model,
303                    dimensions: target.dimensions,
304                    timings,
305                });
306            }
307            prepared = prepare_memory_embedding_batch(&pending, &mut timings, &mut fallback_cache)?;
308        }
309    }
310
311    let processed = upsert_prepared_memory_embedding_batch(conn, &prepared, &mut timings)?;
312    crate::perf::push_elapsed(&mut timings, "total", total_start);
313    Ok(EmbeddingReindexReport {
314        selected,
315        processed,
316        model: target.model,
317        dimensions: target.dimensions,
318        timings,
319    })
320}
321
322pub fn pending_memory_embedding_count(conn: &Connection) -> Result<i64> {
323    if !table_exists(conn, "memories")? || !table_exists(conn, "memory_embeddings")? {
324        return Ok(0);
325    }
326    count_pending_memory_embedding_reindex(conn)
327}
328
329pub fn pending_memory_embedding_reindex_count(conn: &Connection) -> Result<i64> {
330    count_pending_memory_embedding_reindex(conn)
331}
332
333pub fn embedding_count(conn: &Connection) -> Result<i64> {
334    if !table_exists(conn, "memory_embeddings")? {
335        return Ok(0);
336    }
337    Ok(
338        conn.query_row("SELECT COUNT(*) FROM memory_embeddings", [], |row| {
339            row.get(0)
340        })?,
341    )
342}
343
344struct MemoryEmbeddingReindexCandidate {
345    id: i64,
346    topic_key: Option<String>,
347    title: String,
348    content: String,
349    memory_type: String,
350}
351
352struct PreparedMemoryEmbedding {
353    memory_id: i64,
354    model: String,
355    content_hash: String,
356    values: Vec<f32>,
357    updated_at_epoch: i64,
358}
359
360fn prepare_memory_embedding_batch(
361    batch: &[MemoryEmbeddingReindexCandidate],
362    timings: &mut Vec<crate::perf::PhaseTiming>,
363    fallback_cache: &mut super::embedding::EmbeddingFallbackCache,
364) -> Result<Vec<PreparedMemoryEmbedding>> {
365    if batch.is_empty() {
366        return Ok(Vec::new());
367    }
368
369    let embed_start = Instant::now();
370    let mut prepared = Vec::with_capacity(batch.len());
371    for candidate in batch {
372        prepared.push(
373            prepare_memory_embedding(candidate, fallback_cache).with_context(|| {
374                format!(
375                    "memory embedding preparation failed for memory id={}",
376                    candidate.id
377                )
378            })?,
379        );
380    }
381    crate::perf::push_elapsed(timings, "embed_memory", embed_start);
382    Ok(prepared)
383}
384
385fn upsert_prepared_memory_embedding_batch(
386    conn: &Connection,
387    prepared: &[PreparedMemoryEmbedding],
388    timings: &mut Vec<crate::perf::PhaseTiming>,
389) -> Result<usize> {
390    if prepared.is_empty() {
391        return Ok(0);
392    }
393    let prepared_count = prepared.len();
394
395    conn.execute_batch("SAVEPOINT remem_embedding_reindex_batch")
396        .context("start memory embedding reindex savepoint")?;
397    let result = (|| -> Result<()> {
398        let upsert_start = Instant::now();
399        {
400            let mut stmt = conn.prepare(UPSERT_EMBEDDING_SQL)?;
401            for embedding in prepared {
402                execute_embedding_upsert(
403                    &mut stmt,
404                    embedding.memory_id,
405                    &embedding.model,
406                    &embedding.content_hash,
407                    &embedding.values,
408                    embedding.updated_at_epoch,
409                )
410                .with_context(|| {
411                    format!(
412                        "memory embedding upsert failed for memory id={}",
413                        embedding.memory_id
414                    )
415                })?;
416            }
417        }
418        crate::perf::push_elapsed(timings, "upsert_embeddings", upsert_start);
419        Ok(())
420    })();
421
422    match result {
423        Ok(()) => {
424            let commit_start = Instant::now();
425            conn.execute_batch("RELEASE SAVEPOINT remem_embedding_reindex_batch")
426                .context("release memory embedding reindex savepoint")?;
427            crate::perf::push_elapsed(timings, "commit", commit_start);
428            Ok(prepared_count)
429        }
430        Err(error) => {
431            let rollback_result = conn.execute_batch(
432                "ROLLBACK TO SAVEPOINT remem_embedding_reindex_batch;
433                 RELEASE SAVEPOINT remem_embedding_reindex_batch",
434            );
435            match rollback_result {
436                Ok(()) => Err(error),
437                Err(rollback_error) => Err(error).context(format!(
438                    "memory embedding reindex failed and rollback failed: {rollback_error}"
439                )),
440            }
441        }
442    }
443}
444
445fn prepare_memory_embedding(
446    candidate: &MemoryEmbeddingReindexCandidate,
447    fallback_cache: &mut super::embedding::EmbeddingFallbackCache,
448) -> Result<PreparedMemoryEmbedding> {
449    let embedding = super::embedding::embed_memory_with_fallback_cache(
450        &candidate.title,
451        &candidate.content,
452        &candidate.memory_type,
453        candidate.topic_key.as_deref(),
454        fallback_cache,
455    )?;
456    let content_hash = super::embedding::embedding_content_hash(
457        &candidate.title,
458        &candidate.content,
459        &candidate.memory_type,
460        candidate.topic_key.as_deref(),
461    );
462    Ok(PreparedMemoryEmbedding {
463        memory_id: candidate.id,
464        model: embedding.model().to_string(),
465        content_hash,
466        values: embedding.values().to_vec(),
467        updated_at_epoch: chrono::Utc::now().timestamp(),
468    })
469}
470
471fn count_pending_memory_embedding_reindex(conn: &Connection) -> Result<i64> {
472    if super::embedding::provider_disabled_or_error()? {
473        return Ok(0);
474    }
475    if !table_exists(conn, "memories")? || !table_exists(conn, "memory_embeddings")? {
476        return Ok(0);
477    }
478    let target = match super::embedding::configured_backfill_target() {
479        Ok(target) => target,
480        Err(error) if super::embedding::is_embedding_provider_off_error(&error) => return Ok(0),
481        Err(error) => return Err(error),
482    };
483    let sql = "SELECT COUNT(*)
484               FROM memories m
485               LEFT JOIN memory_embeddings e
486                 ON e.memory_id = m.id
487                AND e.model = ?1
488                AND e.dimensions = ?2
489               WHERE (e.memory_id IS NULL
490                      OR e.updated_at_epoch < m.updated_at_epoch)
491                 AND m.status IN ('active', 'stale', 'archived')";
492    Ok(conn.query_row(
493        sql,
494        params![target.model.as_str(), target.dimensions as i64],
495        |row| row.get(0),
496    )?)
497}
498
499pub fn embed_query_text(query: &str) -> Vec<f32> {
500    super::embedding::embed_query_text_local(query)
501}
502
503pub fn embed_memory_text(
504    title: &str,
505    content: &str,
506    memory_type: &str,
507    topic_key: Option<&str>,
508) -> Vec<f32> {
509    super::embedding::embed_memory_text_local(title, content, memory_type, topic_key)
510}
511
512pub fn vector_search(
513    conn: &Connection,
514    query_embedding: &[f32],
515    limit: usize,
516) -> Result<Vec<(i64, f32)>> {
517    Ok(
518        vector_search_filtered(conn, query_embedding, VectorSearchFilters::default(), limit)?
519            .hits
520            .into_iter()
521            .map(|hit| (hit.memory_id, hit.distance))
522            .collect(),
523    )
524}
525
526pub fn vector_search_filtered(
527    conn: &Connection,
528    query_embedding: &[f32],
529    filters: VectorSearchFilters<'_>,
530    limit: usize,
531) -> Result<VectorSearchOutcome> {
532    if query_embedding.len() != EMBEDDING_DIMENSIONS {
533        anyhow::bail!(
534            "query embedding must be {} dimensions, got {}",
535            EMBEDDING_DIMENSIONS,
536            query_embedding.len()
537        );
538    }
539    let embedding = TextEmbedding::new(DEFAULT_EMBEDDING_MODEL, query_embedding.to_vec())?;
540    vector_search_embedding_filtered(conn, &embedding, filters, limit)
541}
542
543pub fn vector_search_embedding_filtered(
544    conn: &Connection,
545    query_embedding: &TextEmbedding,
546    filters: VectorSearchFilters<'_>,
547    limit: usize,
548) -> Result<VectorSearchOutcome> {
549    if limit == 0 {
550        return Ok(VectorSearchOutcome::ready(vec![]));
551    }
552    if super::embedding::provider_disabled_or_error()? {
553        return Ok(VectorSearchOutcome::disabled("embedding provider is off"));
554    }
555    if !table_exists(conn, "memory_embeddings")? {
556        return Ok(VectorSearchOutcome::disabled(
557            "memory_embeddings table is missing; run migrations/backfill",
558        ));
559    }
560    let mut timings = Vec::new();
561    let profile = query_embedding.profile();
562    let candidate_ids = crate::perf::time_result(&mut timings, "vector_select_candidates", || {
563        super::vector_candidates::select_candidate_ids(conn, filters, profile, limit)
564    })?;
565    let candidates_scanned = candidate_ids.len();
566    if candidate_ids.is_empty() {
567        if super::vector_candidates::matching_memory_count(conn, filters)? > 0 {
568            if embedding_count(conn)? == 0 {
569                return Ok(VectorSearchOutcome::disabled_with_timings(
570                    "memory_embeddings table is empty; run `remem reindex-embeddings --limit 1000`",
571                    timings,
572                ));
573            }
574            return Ok(VectorSearchOutcome::disabled_with_timings(
575                format!(
576                    "memory_embeddings has no rows for model={} dimensions={}; run `remem reindex-embeddings --limit 1000`",
577                    profile.model, profile.dimensions
578                ),
579                timings,
580            ));
581        }
582        return Ok(VectorSearchOutcome::ready_with_scan_count_and_timings(
583            vec![],
584            0,
585            timings,
586        ));
587    }
588    let placeholders = std::iter::repeat_n("?", candidate_ids.len())
589        .collect::<Vec<_>>()
590        .join(", ");
591    let sql = format!(
592        "SELECT memory_id, embedding, dimensions
593         FROM memory_embeddings INDEXED BY idx_memory_embeddings_profile_memory_id
594         WHERE model = ?
595           AND dimensions = ?
596           AND memory_id IN ({placeholders})"
597    );
598    let mut param_values: Vec<Box<dyn rusqlite::types::ToSql>> = vec![
599        Box::new(profile.model.to_string()),
600        Box::new(profile.dimensions as i64),
601    ];
602    param_values.extend(
603        candidate_ids
604            .iter()
605            .map(|id| Box::new(*id) as Box<dyn rusqlite::types::ToSql>),
606    );
607    let candidates = crate::perf::time_result(&mut timings, "vector_load_embeddings", || {
608        let refs = crate::db::to_sql_refs(&param_values);
609        let mut stmt = conn.prepare(&sql)?;
610        let rows = stmt.query_map(refs.as_slice(), |row| {
611            Ok((
612                row.get::<_, i64>(0)?,
613                row.get::<_, Vec<u8>>(1)?,
614                row.get::<_, i64>(2)?,
615            ))
616        })?;
617        crate::db::query::collect_rows(rows)
618    })?;
619    let mut hits = crate::perf::time_result(&mut timings, "vector_decode_cosine", || {
620        let mut hits = Vec::new();
621        for (memory_id, blob, dimensions) in candidates {
622            let embedding = decode_embedding(&blob, dimensions)
623                .with_context(|| format!("invalid embedding blob for memory id={memory_id}"))?;
624            let distance = cosine_distance(query_embedding.values(), &embedding)?;
625            hits.push(VectorHit {
626                memory_id,
627                distance,
628            });
629        }
630        Ok(hits)
631    })?;
632    crate::perf::time_value(&mut timings, "vector_sort_truncate", || {
633        hits.sort_by(|a, b| {
634            a.distance
635                .partial_cmp(&b.distance)
636                .unwrap_or(std::cmp::Ordering::Equal)
637                .then_with(|| a.memory_id.cmp(&b.memory_id))
638        });
639        hits.truncate(limit);
640    });
641    Ok(VectorSearchOutcome::ready_with_scan_count_and_timings(
642        hits,
643        candidates_scanned,
644        timings,
645    ))
646}
647
648pub fn find_similar_observations(
649    conn: &Connection,
650    query_embedding: &[f32],
651    threshold: f32,
652    limit: usize,
653) -> Result<Vec<i64>> {
654    let candidates = vector_search(conn, query_embedding, limit)?;
655    let distance_threshold = 1.0 - threshold;
656    let similar: Vec<i64> = candidates
657        .into_iter()
658        .filter(|(_, dist)| *dist < distance_threshold)
659        .map(|(id, _)| id)
660        .collect();
661
662    Ok(similar)
663}
664
665fn create_embedding_table(conn: &Connection) -> Result<()> {
666    conn.execute_batch(
667        "CREATE TABLE IF NOT EXISTS memory_embeddings (
668             memory_id INTEGER NOT NULL,
669             embedding BLOB NOT NULL,
670             dimensions INTEGER NOT NULL,
671             model TEXT NOT NULL,
672             content_hash TEXT NOT NULL,
673             updated_at_epoch INTEGER NOT NULL,
674             PRIMARY KEY(memory_id, model, dimensions),
675             FOREIGN KEY(memory_id) REFERENCES memories(id) ON DELETE CASCADE
676         );
677         CREATE INDEX IF NOT EXISTS idx_memory_embeddings_model
678             ON memory_embeddings(model, updated_at_epoch);
679         CREATE INDEX IF NOT EXISTS idx_memory_embeddings_profile_memory_id
680             ON memory_embeddings(model, dimensions, memory_id);",
681    )?;
682    Ok(())
683}
684
685fn upsert_embedding_with_metadata(
686    conn: &Connection,
687    memory_id: i64,
688    model: &str,
689    content_hash: &str,
690    embedding: &[f32],
691    updated_at_epoch: i64,
692) -> Result<()> {
693    let mut stmt = conn.prepare(UPSERT_EMBEDDING_SQL)?;
694    execute_embedding_upsert(
695        &mut stmt,
696        memory_id,
697        model,
698        content_hash,
699        embedding,
700        updated_at_epoch,
701    )
702}
703
704fn execute_embedding_upsert(
705    stmt: &mut Statement<'_>,
706    memory_id: i64,
707    model: &str,
708    content_hash: &str,
709    embedding: &[f32],
710    updated_at_epoch: i64,
711) -> Result<()> {
712    if model.trim().is_empty() {
713        anyhow::bail!("embedding model must not be empty");
714    }
715    if embedding.is_empty() {
716        anyhow::bail!("embedding vector must not be empty");
717    }
718    if embedding.iter().any(|value| !value.is_finite()) {
719        anyhow::bail!("embedding vector contains non-finite values");
720    }
721    let blob = encode_embedding(embedding);
722    let dimensions = embedding.len() as i64;
723    stmt.execute(params![
724        memory_id,
725        blob,
726        dimensions,
727        model,
728        content_hash,
729        updated_at_epoch
730    ])?;
731    Ok(())
732}
733
734fn encode_embedding(embedding: &[f32]) -> Vec<u8> {
735    let mut out = Vec::with_capacity(std::mem::size_of_val(embedding));
736    for value in embedding {
737        out.extend_from_slice(&value.to_le_bytes());
738    }
739    out
740}
741
742pub(crate) fn decode_embedding(blob: &[u8], dimensions: i64) -> Result<Vec<f32>> {
743    if dimensions <= 0 {
744        anyhow::bail!("embedding dimensions must be positive, got {dimensions}");
745    }
746    let dimensions = dimensions as usize;
747    let expected_bytes = dimensions * std::mem::size_of::<f32>();
748    if blob.len() != expected_bytes {
749        anyhow::bail!(
750            "embedding blob must be {} bytes, got {}",
751            expected_bytes,
752            blob.len()
753        );
754    }
755    Ok(blob
756        .chunks_exact(std::mem::size_of::<f32>())
757        .map(|chunk| f32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]))
758        .collect())
759}
760
761pub(crate) fn cosine_distance(a: &[f32], b: &[f32]) -> Result<f32> {
762    if a.len() != b.len() {
763        anyhow::bail!(
764            "embedding dimensions differ: query={} stored={}",
765            a.len(),
766            b.len()
767        );
768    }
769    let mut dot = 0.0f32;
770    let mut a_norm = 0.0f32;
771    let mut b_norm = 0.0f32;
772    for (left, right) in a.iter().zip(b) {
773        dot += left * right;
774        a_norm += left * left;
775        b_norm += right * right;
776    }
777    if a_norm == 0.0 || b_norm == 0.0 {
778        return Ok(1.0);
779    }
780    Ok((1.0 - dot / (a_norm.sqrt() * b_norm.sqrt())).clamp(0.0, 2.0))
781}
782
783fn table_exists(conn: &Connection, table: &str) -> Result<bool> {
784    Ok(conn
785        .query_row(
786            "SELECT 1 FROM sqlite_master WHERE type = 'table' AND name = ?1 LIMIT 1",
787            params![table],
788            |_| Ok(()),
789        )
790        .optional()?
791        .is_some())
792}
793
794#[cfg(test)]
795mod tests;