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
97pub 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(¶m_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;