1use rig_core::Embed;
11use rig_core::embeddings::{Embedding, EmbeddingModel};
12use rig_core::vector_store::request::{FilterError, SearchFilter, VectorSearchRequest};
13use rig_core::vector_store::{InsertDocuments, VectorStoreError, VectorStoreIndex};
14use rig_core::wasm_compat::{WasmCompatSend, WasmCompatSync};
15use rusqlite::OptionalExtension;
16use rusqlite::types::{Type, Value, ValueRef};
17use serde::{Deserialize, Serialize};
18use std::marker::PhantomData;
19use std::ops::RangeInclusive;
20use tokio_rusqlite::Connection;
21use tracing::{debug, info};
22use zerocopy::IntoBytes;
23
24const SQLITE_VEC_MAX_K: u64 = 4096;
32
33#[derive(Debug)]
34pub enum SqliteError {
35 DatabaseError(Box<dyn std::error::Error + Send + Sync>),
36 SerializationError(Box<dyn std::error::Error + Send + Sync>),
37 InvalidColumnType(String),
38}
39
40pub trait ColumnValue: Send + Sync {
44 fn to_sql_value(&self) -> Value;
46
47 fn column_type(&self) -> &'static str;
49}
50
51#[derive(Clone, Debug)]
52pub struct Column {
53 name: &'static str,
54 col_type: &'static str,
55 indexed: bool,
56}
57
58impl Column {
59 pub fn new(name: &'static str, col_type: &'static str) -> Self {
60 Self {
61 name,
62 col_type,
63 indexed: false,
64 }
65 }
66
67 pub fn indexed(mut self) -> Self {
75 self.indexed = true;
76 self
77 }
78}
79
80pub trait SqliteVectorStoreTable: Send + Sync + Clone {
118 fn name() -> &'static str;
119 fn schema() -> Vec<Column>;
120 fn id(&self) -> String;
121 fn column_values(&self) -> Vec<(&'static str, Box<dyn ColumnValue>)>;
122}
123
124#[derive(Clone, Copy, Debug, Default, Deserialize, Eq, PartialEq, Serialize)]
132pub enum SqliteDistanceMetric {
133 #[default]
135 Cosine,
136 L2,
138 L1,
140}
141
142impl SqliteDistanceMetric {
143 fn vec0_name(self) -> &'static str {
144 match self {
145 Self::Cosine => "cosine",
146 Self::L2 => "l2",
147 Self::L1 => "l1",
148 }
149 }
150
151 fn score_expression(self, query_param: &str, embedding_expr: &str) -> String {
152 match self {
153 Self::Cosine => {
154 format!("(1 - vec_distance_cosine({query_param}, {embedding_expr}))")
155 }
156 Self::L2 => format!("(-vec_distance_l2({query_param}, {embedding_expr}))"),
157 Self::L1 => format!("(-vec_distance_l1({query_param}, {embedding_expr}))"),
158 }
159 }
160}
161
162#[derive(Debug, thiserror::Error)]
163enum SqliteInternalError {
164 #[error(
165 "SQLite vector table `{table_name}` uses {configured:?}, but {requested:?} was requested"
166 )]
167 DistanceMetricMismatch {
168 table_name: String,
169 requested: SqliteDistanceMetric,
170 configured: SqliteDistanceMetric,
171 },
172 #[error("SQLite vector table `{0}` was created but is missing from sqlite_schema")]
173 VectorTableMissingSchema(String),
174 #[error("SQLite metadata column `{column_name}` has unsupported type `{column_type}`")]
175 UnsupportedMetadataColumn {
176 column_name: &'static str,
177 column_type: &'static str,
178 },
179 #[error("SQLite vector table `{table_name}` is missing metadata column `{column_name} {}`", column_type.vec0_name())]
180 MetadataSchemaMismatch {
181 table_name: String,
182 column_name: &'static str,
183 column_type: SqliteMetadataType,
184 },
185 #[error("could not convert SQLite value type `{value_type:?}` for metadata column `{column_name} {}`", column_type.vec0_name())]
186 MetadataValueError {
187 column_name: &'static str,
188 column_type: SqliteMetadataType,
189 value_type: Type,
190 },
191 #[error("SQLite vector store table `{0}` is missing an `id` column")]
192 MissingIdColumn(String),
193 #[error(
194 "could not convert SQLite column `{column_name}` with declared type `{column_type}`: {message}"
195 )]
196 ColumnValueError {
197 column_name: &'static str,
198 column_type: &'static str,
199 message: String,
200 },
201}
202
203#[derive(Clone, Copy, Debug, Eq, PartialEq)]
204enum SqliteMetadataType {
205 Text,
206 Integer,
207 Float,
208 Boolean,
209}
210
211impl SqliteMetadataType {
212 fn from_column_type(column_type: &str) -> Option<Self> {
213 let first_type_token = column_type
214 .split_whitespace()
215 .next()
216 .unwrap_or_default()
217 .to_ascii_uppercase();
218
219 match first_type_token.as_str() {
220 "TEXT" => Some(Self::Text),
221 "INTEGER" | "INT" | "INT64" | "INTEGER64" => Some(Self::Integer),
222 "FLOAT" | "REAL" | "DOUBLE" | "FLOAT64" | "F64" => Some(Self::Float),
223 "BOOLEAN" | "BOOL" => Some(Self::Boolean),
224 _ => match SqliteColumnAffinity::from_column_type(column_type) {
225 SqliteColumnAffinity::Text => Some(Self::Text),
226 SqliteColumnAffinity::Integer => Some(Self::Integer),
227 SqliteColumnAffinity::Float => Some(Self::Float),
228 SqliteColumnAffinity::Boolean => Some(Self::Boolean),
229 SqliteColumnAffinity::Numeric | SqliteColumnAffinity::Blob => None,
230 },
231 }
232 }
233
234 fn vec0_name(self) -> &'static str {
235 match self {
236 Self::Text => "TEXT",
237 Self::Integer => "INTEGER",
238 Self::Float => "FLOAT",
239 Self::Boolean => "BOOLEAN",
240 }
241 }
242
243 fn supports_native_comparison(self, op: SqliteComparisonOp) -> bool {
244 !matches!(
245 (self, op),
246 (
247 Self::Boolean,
248 SqliteComparisonOp::Gt
249 | SqliteComparisonOp::Lt
250 | SqliteComparisonOp::Gte
251 | SqliteComparisonOp::Lte
252 )
253 )
254 }
255}
256
257#[derive(Clone, Copy, Debug, Eq, PartialEq)]
258enum SqliteColumnAffinity {
259 Text,
260 Integer,
261 Float,
262 Boolean,
263 Numeric,
264 Blob,
265}
266
267impl SqliteColumnAffinity {
268 fn from_column_type(column_type: &str) -> Self {
269 let column_type = column_type.to_ascii_uppercase();
270
271 if column_type.contains("INT") {
272 Self::Integer
273 } else if column_type.contains("CHAR")
274 || column_type.contains("CLOB")
275 || column_type.contains("TEXT")
276 {
277 Self::Text
278 } else if column_type.contains("BLOB") || column_type.trim().is_empty() {
279 Self::Blob
280 } else if column_type.contains("REAL")
281 || column_type.contains("FLOA")
282 || column_type.contains("DOUB")
283 {
284 Self::Float
285 } else if column_type.contains("BOOL") {
286 Self::Boolean
287 } else {
288 Self::Numeric
289 }
290 }
291}
292
293#[derive(Clone, Debug, Eq, PartialEq)]
294struct SqliteMetadataColumn {
295 name: &'static str,
296 metadata_type: SqliteMetadataType,
297}
298
299fn sqlite_metadata_columns(
300 schema: &[Column],
301) -> Result<Vec<SqliteMetadataColumn>, VectorStoreError> {
302 schema
303 .iter()
304 .filter(|column| column.indexed)
305 .map(|column| {
306 let metadata_type =
307 SqliteMetadataType::from_column_type(column.col_type).ok_or_else(|| {
308 VectorStoreError::datastore(SqliteInternalError::UnsupportedMetadataColumn {
309 column_name: column.name,
310 column_type: column.col_type,
311 })
312 })?;
313
314 Ok(SqliteMetadataColumn {
315 name: column.name,
316 metadata_type,
317 })
318 })
319 .collect()
320}
321
322fn sqlite_metadata_value(
323 values: &[(&'static str, Box<dyn ColumnValue>)],
324 column: &SqliteMetadataColumn,
325) -> rusqlite::Result<Value> {
326 let value = values
327 .iter()
328 .find(|(name, _)| *name == column.name)
329 .ok_or_else(|| rusqlite::Error::InvalidParameterName(column.name.to_string()))?
330 .1
331 .to_sql_value();
332
333 match (column.metadata_type, value) {
334 (SqliteMetadataType::Text, Value::Text(value)) => Ok(Value::Text(value)),
335 (SqliteMetadataType::Integer, Value::Integer(value)) => Ok(Value::Integer(value)),
336 (SqliteMetadataType::Float, Value::Real(value)) => Ok(Value::Real(value)),
337 (SqliteMetadataType::Float, Value::Integer(value)) => Ok(Value::Real(value as f64)),
338 (SqliteMetadataType::Boolean, Value::Integer(value @ (0 | 1))) => Ok(Value::Integer(value)),
339 (_, value) => Err(rusqlite::Error::ToSqlConversionFailure(Box::new(
340 SqliteInternalError::MetadataValueError {
341 column_name: column.name,
342 column_type: column.metadata_type,
343 value_type: value.data_type(),
344 },
345 ))),
346 }
347}
348
349#[derive(Clone)]
350pub struct SqliteVectorStore<E, T>
351where
352 E: EmbeddingModel + 'static,
353 T: SqliteVectorStoreTable + 'static,
354{
355 conn: Connection,
356 distance_metric: SqliteDistanceMetric,
357 metadata_columns: Vec<SqliteMetadataColumn>,
358 _phantom: PhantomData<(E, T)>,
359}
360
361impl<E, T> SqliteVectorStore<E, T>
362where
363 E: EmbeddingModel + 'static,
364 T: SqliteVectorStoreTable + 'static,
365{
366 async fn candidate_limit(&self, samples: u64, exhaustive: bool) -> Result<u64, VectorStoreError>
367 where
368 Self: 'static,
369 {
370 if samples == 0 {
371 return Ok(0);
372 }
373
374 let embedding_map_table_name = format!("{}_embedding_map", T::name());
375 let (embedding_count, document_count) = self
376 .conn
377 .call(move |conn| {
378 Ok(conn.query_row(
379 &format!(
380 "SELECT COUNT(*), COUNT(DISTINCT document_rowid) FROM {embedding_map_table_name}"
381 ),
382 [],
383 |row| Ok((row.get::<_, i64>(0)?, row.get::<_, i64>(1)?)),
384 )?)
385 })
386 .await
387 .map_err(VectorStoreError::datastore)?;
388
389 let embedding_count = u64::try_from(embedding_count).unwrap_or(0);
390 let document_count = u64::try_from(document_count).unwrap_or(0);
391
392 if exhaustive {
393 Ok(embedding_count.max(samples))
397 } else if embedding_count > document_count {
398 Ok(samples
406 .saturating_add(embedding_count - document_count)
407 .min(embedding_count))
408 } else {
409 Ok(samples)
410 }
411 }
412}
413
414impl<E, T> SqliteVectorStore<E, T>
415where
416 E: EmbeddingModel + Clone + 'static,
417 T: SqliteVectorStoreTable + 'static,
418{
419 pub async fn new(conn: Connection, embedding_model: &E) -> Result<Self, VectorStoreError> {
421 Self::with_distance_metric(conn, embedding_model, SqliteDistanceMetric::default()).await
422 }
423
424 pub async fn with_distance_metric(
430 conn: Connection,
431 embedding_model: &E,
432 distance_metric: SqliteDistanceMetric,
433 ) -> Result<Self, VectorStoreError> {
434 let dims = embedding_model.ndims();
435 let table_name = T::name();
436 let embeddings_table_name = format!("{table_name}_embeddings");
437 let embeddings_table_name_for_sql = embeddings_table_name.clone();
438 let embedding_map_table_name_for_sql = format!("{table_name}_embedding_map");
439 let schema = T::schema();
440 let metadata_columns = sqlite_metadata_columns(&schema)?;
441 let metadata_columns_for_schema_check = metadata_columns.clone();
442 let distance_metric_name = distance_metric.vec0_name();
443 let mut embeddings_columns =
444 format!("embedding float[{dims}] distance_metric={distance_metric_name}");
445 for column in &metadata_columns {
446 embeddings_columns.push_str(&format!(
447 ", {} {}",
448 column.name,
449 column.metadata_type.vec0_name()
450 ));
451 }
452
453 let mut create_table = format!("CREATE TABLE IF NOT EXISTS {table_name} (");
455
456 let mut first = true;
458 for column in &schema {
459 if !first {
460 create_table.push(',');
461 }
462 create_table.push_str(&format!("\n {} {}", column.name, column.col_type));
463 first = false;
464 }
465
466 create_table.push_str("\n)");
467
468 let mut create_indexes = vec![format!(
470 "CREATE INDEX IF NOT EXISTS idx_{}_id ON {}(id)",
471 table_name, table_name
472 )];
473
474 for column in schema {
476 if column.indexed {
477 create_indexes.push(format!(
478 "CREATE INDEX IF NOT EXISTS idx_{}_{} ON {}({})",
479 table_name, column.name, table_name, column.name
480 ));
481 }
482 }
483
484 let embeddings_table_sql = conn
485 .call(move |conn| {
486 conn.execute_batch("BEGIN")?;
487
488 conn.execute_batch(&create_table)?;
490
491 for index_stmt in create_indexes {
493 conn.execute_batch(&index_stmt)?;
494 }
495
496 conn.execute_batch(&format!(
498 "CREATE VIRTUAL TABLE IF NOT EXISTS {embeddings_table_name_for_sql} USING vec0({embeddings_columns})"
499 ))?;
500 conn.execute_batch(&format!(
501 "CREATE TABLE IF NOT EXISTS {embedding_map_table_name_for_sql} (
502 embedding_rowid INTEGER PRIMARY KEY AUTOINCREMENT,
503 document_rowid INTEGER NOT NULL
504 )"
505 ))?;
506 conn.execute_batch(&format!(
507 "CREATE INDEX IF NOT EXISTS idx_{table_name}_embedding_map_document_rowid ON {embedding_map_table_name_for_sql}(document_rowid)"
508 ))?;
509
510 conn.execute_batch("COMMIT")?;
511
512 let schema_sql = conn
513 .query_row(
514 "SELECT sql FROM sqlite_schema WHERE name = ?1",
515 [&embeddings_table_name_for_sql],
516 |row| row.get::<_, String>(0),
517 )
518 .optional()?;
519
520 Ok(schema_sql)
521 })
522 .await
523 .map_err(VectorStoreError::datastore)?;
524
525 let schema_sql = embeddings_table_sql.ok_or_else(|| {
526 VectorStoreError::datastore(SqliteInternalError::VectorTableMissingSchema(
527 embeddings_table_name.clone(),
528 ))
529 })?;
530
531 let configured = sqlite_distance_metric_from_schema(&schema_sql);
532 if configured != distance_metric {
533 return Err(VectorStoreError::datastore(
534 SqliteInternalError::DistanceMetricMismatch {
535 table_name: embeddings_table_name,
536 requested: distance_metric,
537 configured,
538 },
539 ));
540 }
541 for column in metadata_columns_for_schema_check {
542 if !sqlite_schema_contains_metadata_column(&schema_sql, &column) {
543 return Err(VectorStoreError::datastore(
544 SqliteInternalError::MetadataSchemaMismatch {
545 table_name: embeddings_table_name.clone(),
546 column_name: column.name,
547 column_type: column.metadata_type,
548 },
549 ));
550 }
551 }
552
553 Ok(Self {
554 conn,
555 distance_metric,
556 metadata_columns,
557 _phantom: PhantomData,
558 })
559 }
560
561 pub fn index(self, model: E) -> SqliteVectorIndex<E, T> {
562 SqliteVectorIndex::new(model, self)
563 }
564
565 pub fn add_rows_with_txn(
566 &self,
567 txn: &rusqlite::Transaction<'_>,
568 documents: Vec<(T, Vec<Embedding>)>,
569 ) -> Result<i64, tokio_rusqlite::Error> {
570 info!("Adding {} documents to store", documents.len());
571 let table_name = T::name();
572 let embeddings_table_name = format!("{table_name}_embeddings");
573 let embedding_map_table_name = format!("{table_name}_embedding_map");
574 let mut last_id = 0;
575 let embedding_columns = std::iter::once("rowid")
576 .chain(std::iter::once("embedding"))
577 .chain(self.metadata_columns.iter().map(|column| column.name))
578 .collect::<Vec<_>>();
579 let embedding_placeholders = (1..=embedding_columns.len())
580 .map(|i| format!("?{i}"))
581 .collect::<Vec<_>>();
582 let embeddings_sql = format!(
583 "INSERT INTO {embeddings_table_name} ({}) VALUES ({})",
584 embedding_columns.join(", "),
585 embedding_placeholders.join(", ")
586 );
587 let existing_rowid_sql = format!("SELECT rowid FROM {table_name} WHERE id = ?1");
588 let existing_embedding_rowids_sql = format!(
589 "SELECT embedding_rowid FROM {embedding_map_table_name} WHERE document_rowid = ?1"
590 );
591 let insert_embedding_map_sql =
592 format!("INSERT INTO {embedding_map_table_name}(document_rowid) VALUES (?1)");
593 let delete_embedding_map_sql =
594 format!("DELETE FROM {embedding_map_table_name} WHERE document_rowid = ?1");
595 let delete_embeddings_sql = format!("DELETE FROM {embeddings_table_name} WHERE rowid = ?1");
596
597 for (doc, embeddings) in &documents {
598 debug!("Storing document with id {}", doc.id());
599
600 let values = doc.column_values();
601 let id_value = values
602 .iter()
603 .find(|(name, _)| *name == "id")
604 .map(|(_, value)| value.to_sql_value())
605 .unwrap_or_else(|| Value::Text(doc.id()));
606 if let Some(existing_rowid) = txn
607 .query_row(&existing_rowid_sql, rusqlite::params![id_value], |row| {
608 row.get::<_, i64>(0)
609 })
610 .optional()?
611 {
612 let existing_embedding_rowids = txn
613 .prepare(&existing_embedding_rowids_sql)?
614 .query_map([existing_rowid], |row| row.get::<_, i64>(0))?
615 .collect::<rusqlite::Result<Vec<_>>>()?;
616 for embedding_rowid in existing_embedding_rowids {
617 txn.execute(&delete_embeddings_sql, [embedding_rowid])?;
618 }
619 txn.execute(&delete_embedding_map_sql, [existing_rowid])?;
620 }
621
622 let columns = values.iter().map(|(col, _)| *col).collect::<Vec<_>>();
623
624 let placeholders = (1..=values.len())
625 .map(|i| format!("?{i}"))
626 .collect::<Vec<_>>();
627
628 let insert_sql = format!(
629 "INSERT OR REPLACE INTO {} ({}) VALUES ({})",
630 table_name,
631 columns.join(", "),
632 placeholders.join(", ")
633 );
634
635 txn.execute(
636 &insert_sql,
637 rusqlite::params_from_iter(values.iter().map(|(_, val)| val.to_sql_value())),
638 )?;
639 last_id = txn.last_insert_rowid();
640
641 let metadata_values = self
642 .metadata_columns
643 .iter()
644 .map(|column| sqlite_metadata_value(&values, column))
645 .collect::<rusqlite::Result<Vec<_>>>()?;
646
647 let mut stmt = txn.prepare(&embeddings_sql)?;
648 for (i, embedding) in embeddings.iter().enumerate() {
649 let vec = serialize_embedding(embedding);
650 debug!(
651 "Storing embedding {} of {} (size: {} bytes)",
652 i + 1,
653 embeddings.len(),
654 vec.len() * 4
655 );
656 txn.execute(&insert_embedding_map_sql, [last_id])?;
657 let embedding_rowid = txn.last_insert_rowid();
658 let mut params = Vec::with_capacity(2 + metadata_values.len());
659 params.push(Value::Integer(embedding_rowid));
660 params.push(Value::Blob(vec.as_bytes().to_vec()));
661 params.extend(metadata_values.iter().cloned());
662 stmt.execute(rusqlite::params_from_iter(params))?;
663 }
664 }
665
666 Ok(last_id)
667 }
668
669 pub async fn add_rows(
670 &self,
671 documents: Vec<(T, Vec<Embedding>)>,
672 ) -> Result<i64, VectorStoreError>
673 where
674 T: 'static,
675 Self: 'static,
676 {
677 let cloned = self.clone();
678
679 self.conn
680 .call(move |conn| {
681 let tx = conn.transaction()?;
682 let result = cloned.add_rows_with_txn(&tx, documents)?;
683 tx.commit()?;
684
685 Ok(result)
686 })
687 .await
688 .map_err(VectorStoreError::datastore)
689 }
690}
691
692impl<E, T> InsertDocuments for SqliteVectorStore<E, T>
693where
694 E: EmbeddingModel + Clone + WasmCompatSend + WasmCompatSync + 'static,
695 T: SqliteVectorStoreTable
696 + for<'de> Deserialize<'de>
697 + WasmCompatSend
698 + WasmCompatSync
699 + 'static,
700{
701 async fn insert_documents<Doc: Serialize + Embed + WasmCompatSend>(
702 &self,
703 documents: Vec<(Doc, Vec<Embedding>)>,
704 ) -> Result<(), VectorStoreError> {
705 if documents.is_empty() {
706 return Ok(());
707 }
708
709 let rows = documents
710 .into_iter()
711 .map(|(document, embeddings)| {
712 let document = serde_json::to_value(document)?;
713 let row = serde_json::from_value::<T>(document)?;
714
715 Ok((row, embeddings))
716 })
717 .collect::<Result<Vec<_>, VectorStoreError>>()?;
718
719 self.add_rows(rows).await?;
720
721 Ok(())
722 }
723}
724
725#[derive(Clone, Deserialize, Serialize, Debug)]
737pub struct SqliteSearchFilter {
738 expr: SqliteSearchFilterExpr,
739}
740
741impl Default for SqliteSearchFilter {
742 fn default() -> Self {
743 Self {
744 expr: SqliteSearchFilterExpr::Noop,
745 }
746 }
747}
748
749#[derive(Clone, Deserialize, Serialize, Debug)]
750enum SqliteSearchFilterExpr {
751 Comparison {
752 key: String,
753 op: SqliteComparisonOp,
754 value: serde_json::Value,
755 },
756 And(Box<SqliteSearchFilterExpr>, Box<SqliteSearchFilterExpr>),
757 Or(Box<SqliteSearchFilterExpr>, Box<SqliteSearchFilterExpr>),
758 Not(Box<SqliteSearchFilterExpr>),
759 Between {
760 key: String,
761 lo: serde_json::Value,
762 hi: serde_json::Value,
763 },
764 NullCheck {
765 key: String,
766 negated: bool,
767 },
768 Pattern {
769 key: String,
770 op: SqlitePatternOp,
771 pattern: String,
772 },
773 Noop,
775}
776
777#[derive(Clone, Copy, Deserialize, Eq, PartialEq, Serialize, Debug)]
778enum SqliteComparisonOp {
779 Eq,
780 Ne,
781 Gt,
782 Gte,
783 Lt,
784 Lte,
785}
786
787impl SqliteComparisonOp {
788 fn as_sql(self) -> &'static str {
789 match self {
790 Self::Eq => "=",
791 Self::Ne => "!=",
792 Self::Gt => ">",
793 Self::Gte => ">=",
794 Self::Lt => "<",
795 Self::Lte => "<=",
796 }
797 }
798
799 fn negate(self) -> Self {
800 match self {
801 Self::Eq => Self::Ne,
802 Self::Ne => Self::Eq,
803 Self::Gt => Self::Lte,
804 Self::Gte => Self::Lt,
805 Self::Lt => Self::Gte,
806 Self::Lte => Self::Gt,
807 }
808 }
809}
810
811#[derive(Clone, Copy, Deserialize, Serialize, Debug)]
812enum SqlitePatternOp {
813 Glob,
814 Like,
815}
816
817impl SqlitePatternOp {
818 fn as_sql(self) -> &'static str {
819 match self {
820 Self::Glob => "glob",
821 Self::Like => "like",
822 }
823 }
824}
825
826#[derive(Debug, Default)]
827struct SqliteRenderedFilters {
828 native: Vec<SqliteRenderedFilter>,
829 post: Vec<SqliteRenderedFilter>,
830}
831
832impl SqliteRenderedFilters {
833 fn post_only(filter: SqliteRenderedFilter) -> Self {
834 Self {
835 native: Vec::new(),
836 post: vec![filter],
837 }
838 }
839
840 fn extend(&mut self, rhs: Self) {
841 self.native.extend(rhs.native);
842 self.post.extend(rhs.post);
843 }
844
845 fn has_post_filters(&self) -> bool {
846 !self.post.is_empty()
847 }
848}
849
850#[derive(Debug)]
851struct SqliteRenderedFilter {
852 condition: String,
853 params: Vec<Value>,
854}
855
856impl SqliteRenderedFilter {
857 fn combine(joiner: &str, lhs: Self, rhs: Self) -> Self {
858 Self {
859 condition: format!("({}) {joiner} ({})", lhs.condition, rhs.condition),
860 params: lhs.params.into_iter().chain(rhs.params).collect(),
861 }
862 }
863}
864
865#[derive(Clone, Copy, Debug, Eq, PartialEq)]
866enum SqliteDocumentValueMode {
867 Sql,
868 JsonText,
869}
870
871#[derive(Debug)]
872struct SqliteQualifiedDocumentKey {
873 expression: String,
874 value_mode: SqliteDocumentValueMode,
875 plain_column: Option<String>,
876}
877
878impl SqliteSearchFilter {
879 fn cmp(key: impl AsRef<str>, op: SqliteComparisonOp, value: serde_json::Value) -> Self {
880 Self {
881 expr: SqliteSearchFilterExpr::Comparison {
882 key: key.as_ref().to_string(),
883 op,
884 value,
885 },
886 }
887 }
888
889 fn pattern(key: String, op: SqlitePatternOp, pattern: impl Into<String>) -> Self {
890 Self {
891 expr: SqliteSearchFilterExpr::Pattern {
892 key,
893 op,
894 pattern: pattern.into(),
895 },
896 }
897 }
898
899 fn null_check(key: String, negated: bool) -> Self {
900 Self {
901 expr: SqliteSearchFilterExpr::NullCheck { key, negated },
902 }
903 }
904}
905
906impl SearchFilter for SqliteSearchFilter {
907 type Value = serde_json::Value;
908
909 fn eq(key: impl AsRef<str>, value: Self::Value) -> Self {
910 Self::cmp(key, SqliteComparisonOp::Eq, value)
911 }
912
913 fn gt(key: impl AsRef<str>, value: Self::Value) -> Self {
914 Self::cmp(key, SqliteComparisonOp::Gt, value)
915 }
916
917 fn lt(key: impl AsRef<str>, value: Self::Value) -> Self {
918 Self::cmp(key, SqliteComparisonOp::Lt, value)
919 }
920
921 fn and(self, rhs: Self) -> Self {
922 Self {
923 expr: SqliteSearchFilterExpr::And(Box::new(self.expr), Box::new(rhs.expr)),
924 }
925 }
926
927 fn or(self, rhs: Self) -> Self {
928 Self {
929 expr: SqliteSearchFilterExpr::Or(Box::new(self.expr), Box::new(rhs.expr)),
930 }
931 }
932}
933
934impl SqliteSearchFilter {
935 #[allow(clippy::should_implement_trait)]
936 pub fn not(self) -> Self {
943 Self {
944 expr: SqliteSearchFilterExpr::Not(Box::new(self.expr)),
945 }
946 }
947
948 pub fn between<N>(key: String, range: RangeInclusive<N>) -> Self
954 where
955 N: Into<serde_json::Value>,
956 {
957 let (lo, hi) = range.into_inner();
958
959 Self {
960 expr: SqliteSearchFilterExpr::Between {
961 key,
962 lo: lo.into(),
963 hi: hi.into(),
964 },
965 }
966 }
967
968 pub fn is_null(key: String) -> Self {
970 Self::null_check(key, false)
971 }
972
973 pub fn is_not_null(key: String) -> Self {
974 Self::null_check(key, true)
975 }
976
977 pub fn glob(key: String, pattern: impl Into<String>) -> Self {
982 Self::pattern(key, SqlitePatternOp::Glob, pattern)
983 }
984
985 pub fn like(key: String, pattern: impl Into<String>) -> Self {
990 Self::pattern(key, SqlitePatternOp::Like, pattern)
991 }
992}
993
994impl SqliteSearchFilter {
995 fn render_split(
996 &self,
997 metadata_columns: &[SqliteMetadataColumn],
998 ) -> Result<SqliteRenderedFilters, FilterError> {
999 self.expr.render_split(metadata_columns)
1000 }
1001}
1002
1003impl SqliteSearchFilterExpr {
1004 fn render_native_comparison(
1005 key: &str,
1006 op: SqliteComparisonOp,
1007 value: serde_json::Value,
1008 metadata_columns: &[SqliteMetadataColumn],
1009 ) -> Result<SqliteRenderedFilters, FilterError> {
1010 let Some(metadata_column) = sqlite_native_metadata_column(key, metadata_columns) else {
1011 return Ok(SqliteRenderedFilters::post_only(
1012 Self::render_document_comparison(key, op, value, metadata_columns)?,
1013 ));
1014 };
1015
1016 if !metadata_column.metadata_type.supports_native_comparison(op) {
1017 return Err(sqlite_unsupported_filter(format!(
1018 "`{key}` is a BOOLEAN metadata column, and sqlite-vec only supports `=` and `!=` filters for booleans"
1019 )));
1020 }
1021
1022 Ok(SqliteRenderedFilters {
1023 native: vec![SqliteRenderedFilter {
1024 condition: format!("e.{key} {} ?", op.as_sql()),
1025 params: vec![sqlite_metadata_filter_param(metadata_column, value)?],
1026 }],
1027 post: Vec::new(),
1028 })
1029 }
1030
1031 fn render_document_comparison(
1032 key: &str,
1033 op: SqliteComparisonOp,
1034 value: serde_json::Value,
1035 metadata_columns: &[SqliteMetadataColumn],
1036 ) -> Result<SqliteRenderedFilter, FilterError> {
1037 let key = sqlite_qualify_document_key(key)?;
1038 Ok(SqliteRenderedFilter {
1039 condition: format!("{} {} ?", key.expression, op.as_sql()),
1040 params: vec![sqlite_document_filter_param(&key, metadata_columns, value)?],
1041 })
1042 }
1043
1044 fn render_split(
1045 &self,
1046 metadata_columns: &[SqliteMetadataColumn],
1047 ) -> Result<SqliteRenderedFilters, FilterError> {
1048 match self {
1049 Self::Comparison { key, op, value } => {
1050 Self::render_native_comparison(key, *op, value.clone(), metadata_columns)
1051 }
1052 Self::And(lhs, rhs) => {
1053 let mut rendered = lhs.render_split(metadata_columns)?;
1054 rendered.extend(rhs.render_split(metadata_columns)?);
1055 Ok(rendered)
1056 }
1057 Self::Between { key, lo, hi } => {
1058 let Some(metadata_column) = sqlite_native_metadata_column(key, metadata_columns)
1059 else {
1060 return Ok(SqliteRenderedFilters::post_only(
1061 self.render_document(metadata_columns)?,
1062 ));
1063 };
1064
1065 if metadata_column.metadata_type == SqliteMetadataType::Boolean {
1066 return Err(sqlite_unsupported_filter(format!(
1067 "`{key}` is a BOOLEAN metadata column, and sqlite-vec does not support range filters for booleans"
1068 )));
1069 }
1070
1071 Ok(SqliteRenderedFilters {
1072 native: vec![SqliteRenderedFilter {
1073 condition: format!("e.{key} >= ? AND e.{key} <= ?"),
1074 params: vec![
1075 sqlite_metadata_filter_param(metadata_column, lo.clone())?,
1076 sqlite_metadata_filter_param(metadata_column, hi.clone())?,
1077 ],
1078 }],
1079 post: Vec::new(),
1080 })
1081 }
1082 Self::Noop => Ok(SqliteRenderedFilters::default()),
1083 Self::Or(_, _) | Self::NullCheck { .. } | Self::Pattern { .. } => Ok(
1084 SqliteRenderedFilters::post_only(self.render_document(metadata_columns)?),
1085 ),
1086 Self::Not(expr) => expr.render_negated_split(metadata_columns),
1087 }
1088 }
1089
1090 fn render_negated_split(
1091 &self,
1092 metadata_columns: &[SqliteMetadataColumn],
1093 ) -> Result<SqliteRenderedFilters, FilterError> {
1094 match self {
1095 Self::Comparison { key, op, value } => {
1096 Self::render_native_comparison(key, op.negate(), value.clone(), metadata_columns)
1097 }
1098 Self::Not(expr) => expr.render_split(metadata_columns),
1099 _ => {
1100 let rendered = self.render_document(metadata_columns)?;
1101 Ok(SqliteRenderedFilters::post_only(SqliteRenderedFilter {
1102 condition: format!("NOT ({})", rendered.condition),
1103 params: rendered.params,
1104 }))
1105 }
1106 }
1107 }
1108
1109 fn render_document(
1110 &self,
1111 metadata_columns: &[SqliteMetadataColumn],
1112 ) -> Result<SqliteRenderedFilter, FilterError> {
1113 match self {
1114 Self::Comparison { key, op, value } => {
1115 Self::render_document_comparison(key, *op, value.clone(), metadata_columns)
1116 }
1117 Self::And(lhs, rhs) => Ok(SqliteRenderedFilter::combine(
1118 "AND",
1119 lhs.render_document(metadata_columns)?,
1120 rhs.render_document(metadata_columns)?,
1121 )),
1122 Self::Or(lhs, rhs) => Ok(SqliteRenderedFilter::combine(
1123 "OR",
1124 lhs.render_document(metadata_columns)?,
1125 rhs.render_document(metadata_columns)?,
1126 )),
1127 Self::Not(expr) => {
1128 let expr = expr.render_document(metadata_columns)?;
1129 Ok(SqliteRenderedFilter {
1130 condition: format!("NOT ({})", expr.condition),
1131 params: expr.params,
1132 })
1133 }
1134 Self::Between { key, lo, hi } => {
1135 let key = sqlite_qualify_document_key(key)?;
1136 Ok(SqliteRenderedFilter {
1137 condition: format!("{} between ? and ?", key.expression),
1138 params: vec![
1139 sqlite_document_filter_param(&key, metadata_columns, lo.clone())?,
1140 sqlite_document_filter_param(&key, metadata_columns, hi.clone())?,
1141 ],
1142 })
1143 }
1144 Self::NullCheck { key, negated } => {
1145 let key = sqlite_qualify_document_key(key)?;
1146 let operator = if *negated { "is not null" } else { "is null" };
1147 Ok(SqliteRenderedFilter {
1148 condition: format!("{} {operator}", key.expression),
1149 params: Vec::new(),
1150 })
1151 }
1152 Self::Pattern { key, op, pattern } => {
1153 let key = sqlite_qualify_document_key(key)?;
1154 Ok(SqliteRenderedFilter {
1155 condition: format!("{} {} ?", key.expression, op.as_sql()),
1156 params: vec![Value::Text(pattern.clone())],
1157 })
1158 }
1159 Self::Noop => Ok(SqliteRenderedFilter {
1162 condition: "1 = 1".to_owned(),
1163 params: Vec::new(),
1164 }),
1165 }
1166 }
1167}
1168
1169fn sqlite_native_metadata_column<'a>(
1170 key: &str,
1171 metadata_columns: &'a [SqliteMetadataColumn],
1172) -> Option<&'a SqliteMetadataColumn> {
1173 if !sqlite_is_plain_identifier(key) {
1174 return None;
1175 }
1176
1177 metadata_columns.iter().find(|column| column.name == key)
1178}
1179
1180fn sqlite_is_plain_identifier(key: &str) -> bool {
1181 let mut chars = key.chars();
1182 let Some(first) = chars.next() else {
1183 return false;
1184 };
1185
1186 (first == '_' || first.is_ascii_alphabetic())
1187 && chars.all(|c| c == '_' || c.is_ascii_alphanumeric())
1188}
1189
1190fn sqlite_leading_identifier_len(key: &str) -> Option<usize> {
1191 let mut chars = key.char_indices();
1192 let (_, first) = chars.next()?;
1193 if !(first == '_' || first.is_ascii_alphabetic()) {
1194 return None;
1195 }
1196
1197 let mut end = first.len_utf8();
1198 for (index, c) in chars {
1199 if c == '_' || c.is_ascii_alphanumeric() {
1200 end = index + c.len_utf8();
1201 } else {
1202 break;
1203 }
1204 }
1205
1206 Some(end)
1207}
1208
1209fn sqlite_unsupported_filter(reason: impl Into<String>) -> FilterError {
1210 FilterError::TypeError(format!(
1211 "SQLite filter cannot be safely lowered; {}",
1212 reason.into()
1213 ))
1214}
1215
1216fn sqlite_json_type_name(value: &serde_json::Value) -> &'static str {
1217 match value {
1218 serde_json::Value::Null => "null",
1219 serde_json::Value::Bool(_) => "boolean",
1220 serde_json::Value::Number(_) => "number",
1221 serde_json::Value::String(_) => "string",
1222 serde_json::Value::Array(_) => "array",
1223 serde_json::Value::Object(_) => "object",
1224 }
1225}
1226
1227fn sqlite_metadata_filter_type_error(
1228 column: &SqliteMetadataColumn,
1229 value: &serde_json::Value,
1230 expected: &str,
1231) -> FilterError {
1232 sqlite_unsupported_filter(format!(
1233 "`{}` is a {} metadata column and requires {expected}; got {}",
1234 column.name,
1235 column.metadata_type.vec0_name(),
1236 sqlite_json_type_name(value)
1237 ))
1238}
1239
1240fn sqlite_metadata_filter_param(
1241 column: &SqliteMetadataColumn,
1242 value: serde_json::Value,
1243) -> Result<Value, FilterError> {
1244 let expected = match column.metadata_type {
1245 SqliteMetadataType::Text => "a string filter value",
1246 SqliteMetadataType::Integer => "an integer filter value",
1247 SqliteMetadataType::Float => "a finite number filter value",
1248 SqliteMetadataType::Boolean => "a boolean filter value",
1249 };
1250
1251 match (column.metadata_type, value) {
1252 (SqliteMetadataType::Text, serde_json::Value::String(value)) => Ok(Value::Text(value)),
1253 (SqliteMetadataType::Integer, serde_json::Value::Number(number)) => {
1254 if let Some(value) = number.as_i64() {
1255 Ok(Value::Integer(value))
1256 } else if let Some(value) = number.as_u64() {
1257 i64::try_from(value).map(Value::Integer).map_err(|_| {
1258 FilterError::TypeError(format!(
1259 "SQLite integer filter value `{number}` exceeds i64::MAX"
1260 ))
1261 })
1262 } else {
1263 Err(sqlite_metadata_filter_type_error(
1264 column,
1265 &serde_json::Value::Number(number),
1266 expected,
1267 ))
1268 }
1269 }
1270 (SqliteMetadataType::Float, serde_json::Value::Number(number)) => {
1271 number.as_f64().map(Value::Real).ok_or_else(|| {
1272 sqlite_metadata_filter_type_error(
1273 column,
1274 &serde_json::Value::Number(number),
1275 expected,
1276 )
1277 })
1278 }
1279 (SqliteMetadataType::Boolean, serde_json::Value::Bool(value)) => {
1280 Ok(Value::Integer(value as i64))
1281 }
1282 (_, value) => Err(sqlite_metadata_filter_type_error(column, &value, expected)),
1283 }
1284}
1285
1286fn sqlite_filter_param(value: serde_json::Value) -> Result<Value, FilterError> {
1287 use serde_json::Value::*;
1288
1289 match value {
1290 Null => Ok(Value::Null),
1291 Bool(b) => Ok(Value::Integer(b as i64)),
1292 String(s) => Ok(Value::Text(s)),
1293 Number(n) => Ok(if let Some(value) = n.as_i64() {
1294 Value::Integer(value)
1295 } else if let Some(value) = n.as_u64() {
1296 let value = i64::try_from(value).map_err(|_| {
1297 FilterError::TypeError(format!(
1298 "SQLite integer filter value `{n}` exceeds i64::MAX"
1299 ))
1300 })?;
1301 Value::Integer(value)
1302 } else if let Some(float) = n.as_f64() {
1303 Value::Real(float)
1304 } else {
1305 Value::Text(n.to_string())
1306 }),
1307 Array(arr) => {
1308 let blob =
1309 serde_json::to_vec(&arr).map_err(|e| FilterError::Serialization(e.to_string()))?;
1310
1311 Ok(Value::Blob(blob))
1312 }
1313 Object(obj) => {
1314 let blob =
1315 serde_json::to_vec(&obj).map_err(|e| FilterError::Serialization(e.to_string()))?;
1316
1317 Ok(Value::Blob(blob))
1318 }
1319 }
1320}
1321
1322fn sqlite_qualify_document_key(key: &str) -> Result<SqliteQualifiedDocumentKey, FilterError> {
1323 if let Some(key_without_alias) = key.strip_prefix("d.") {
1324 if sqlite_is_plain_identifier(key_without_alias) {
1325 return Ok(SqliteQualifiedDocumentKey {
1326 expression: key.to_string(),
1327 value_mode: SqliteDocumentValueMode::Sql,
1328 plain_column: Some(key_without_alias.to_string()),
1329 });
1330 }
1331
1332 if let Some(value_mode) = sqlite_json_operator_value_mode(key_without_alias) {
1333 return Ok(SqliteQualifiedDocumentKey {
1334 expression: key.to_string(),
1335 value_mode,
1336 plain_column: None,
1337 });
1338 }
1339
1340 return Err(sqlite_unsupported_filter(format!(
1341 "`{key}` is not a supported SQLite document filter expression"
1342 )));
1343 }
1344
1345 if sqlite_is_plain_identifier(key) {
1346 return Ok(SqliteQualifiedDocumentKey {
1347 expression: format!("d.{key}"),
1348 value_mode: SqliteDocumentValueMode::Sql,
1349 plain_column: Some(key.to_string()),
1350 });
1351 }
1352
1353 if let Some(value_mode) = sqlite_json_operator_value_mode(key) {
1354 return Ok(SqliteQualifiedDocumentKey {
1355 expression: format!("d.{key}"),
1356 value_mode,
1357 plain_column: None,
1358 });
1359 }
1360
1361 Err(sqlite_unsupported_filter(format!(
1362 "`{key}` is not a supported SQLite document filter expression"
1363 )))
1364}
1365
1366fn sqlite_document_filter_param(
1367 key: &SqliteQualifiedDocumentKey,
1368 metadata_columns: &[SqliteMetadataColumn],
1369 value: serde_json::Value,
1370) -> Result<Value, FilterError> {
1371 match key.value_mode {
1372 SqliteDocumentValueMode::Sql => {
1373 if let Some(column_name) = key.plain_column.as_deref()
1374 && let Some(metadata_column) = metadata_columns
1375 .iter()
1376 .find(|column| column.name == column_name)
1377 {
1378 return sqlite_metadata_filter_param(metadata_column, value);
1379 }
1380
1381 sqlite_filter_param(value)
1382 }
1383 SqliteDocumentValueMode::JsonText => serde_json::to_string(&value)
1384 .map(Value::Text)
1385 .map_err(|e| FilterError::Serialization(e.to_string())),
1386 }
1387}
1388
1389fn sqlite_json_operator_value_mode(expr: &str) -> Option<SqliteDocumentValueMode> {
1390 let mut index = sqlite_leading_identifier_len(expr)?;
1391
1392 if index == expr.len() {
1393 return None;
1394 }
1395
1396 let mut value_mode = None;
1397 while index < expr.len() {
1398 let remaining = &expr[index..];
1399 let (operator_len, next_value_mode) = if remaining.starts_with("->>") {
1400 (3, SqliteDocumentValueMode::Sql)
1401 } else if remaining.starts_with("->") {
1402 (2, SqliteDocumentValueMode::JsonText)
1403 } else {
1404 return None;
1405 };
1406 value_mode = Some(next_value_mode);
1407 index += operator_len;
1408
1409 let operand_len = sqlite_json_operator_operand_len(&expr[index..])?;
1410 index += operand_len;
1411 }
1412
1413 value_mode
1414}
1415
1416fn sqlite_json_operator_operand_len(operand: &str) -> Option<usize> {
1417 if operand.is_empty() {
1418 return None;
1419 }
1420
1421 if let Some(operand) = operand.strip_prefix('\'') {
1422 let closing_quote = operand.find('\'')?;
1423 let literal = &operand[..closing_quote];
1424 if literal.chars().any(char::is_control) {
1425 return None;
1426 }
1427
1428 return Some(closing_quote + 2);
1429 }
1430
1431 let mut chars = operand.char_indices();
1432 let mut end = 0;
1433 if let Some((_, '-')) = chars.clone().next() {
1434 end = 1;
1435 chars.next();
1436 }
1437
1438 let mut has_digit = false;
1439 for (index, c) in chars {
1440 if c.is_ascii_digit() {
1441 has_digit = true;
1442 end = index + c.len_utf8();
1443 } else {
1444 break;
1445 }
1446 }
1447
1448 has_digit.then_some(end)
1449}
1450
1451pub struct SqliteVectorIndex<E, T>
1550where
1551 E: EmbeddingModel + 'static,
1552 T: SqliteVectorStoreTable + 'static,
1553{
1554 store: SqliteVectorStore<E, T>,
1555 embedding_model: E,
1556}
1557
1558impl<E, T> SqliteVectorIndex<E, T>
1559where
1560 E: EmbeddingModel + 'static,
1561 T: SqliteVectorStoreTable,
1562{
1563 pub fn new(embedding_model: E, store: SqliteVectorStore<E, T>) -> Self {
1564 Self {
1565 store,
1566 embedding_model,
1567 }
1568 }
1569}
1570
1571impl<E, T> SqliteVectorIndex<E, T>
1572where
1573 E: EmbeddingModel + 'static,
1574 T: SqliteVectorStoreTable + 'static,
1575{
1576 async fn search_rows<R, F>(
1582 &self,
1583 req: &VectorSearchRequest<SqliteSearchFilter>,
1584 outer_select_cols: String,
1585 map_row: F,
1586 ) -> Result<Vec<R>, VectorStoreError>
1587 where
1588 R: Send + 'static,
1589 F: Fn(&rusqlite::Row<'_>) -> rusqlite::Result<R> + Send + 'static,
1590 {
1591 let embedding = self.embedding_model.embed_text(req.query()).await?;
1592 let query_vec: Vec<f32> = serialize_embedding(&embedding);
1593 let table_name = T::name();
1594 let embedding_map_table_name = format!("{table_name}_embedding_map");
1595
1596 let distance_metric = self.store.distance_metric;
1597 let score_expression = distance_metric.score_expression("?1", "e.embedding");
1598 let filters = render_search_filters(req, distance_metric, &self.store.metadata_columns)?;
1599 let candidate_limit = self
1600 .store
1601 .candidate_limit(req.samples(), filters.has_post_filters())
1602 .await?;
1603 let search_query = build_search_query(query_vec, filters, candidate_limit)?;
1604 let where_clause = search_query.vector_where_clause;
1605 let document_filter_clause = search_query.document_filter_clause;
1606 let mut params = search_query.params;
1607 params.push(sqlite_limit_param(req.samples(), "result limit")?);
1608
1609 self.store
1610 .conn
1611 .call(move |conn| {
1612 let mut stmt = conn.prepare(&format!(
1613 "WITH scored AS (
1614 SELECT m.document_rowid AS __rig_document_rowid,
1615 {score_expression} AS __rig_score,
1616 ROW_NUMBER() OVER (
1617 PARTITION BY m.document_rowid
1618 ORDER BY {score_expression} DESC, e.rowid ASC
1619 ) AS __rig_rank
1620 FROM {table_name}_embeddings e
1621 JOIN {embedding_map_table_name} m ON e.rowid = m.embedding_rowid
1622 {where_clause}
1623 )
1624 SELECT {outer_select_cols}, scored.__rig_score
1625 FROM scored
1626 JOIN {table_name} d ON scored.__rig_document_rowid = d.rowid
1627 WHERE scored.__rig_rank = 1
1628 {document_filter_clause}
1629 ORDER BY scored.__rig_score DESC, d.id ASC
1630 LIMIT ?"
1631 ))?;
1632
1633 let rows = stmt
1634 .query_map(rusqlite::params_from_iter(params), |row| map_row(row))?
1635 .collect::<Result<Vec<_>, _>>()?;
1636 Ok(rows)
1637 })
1638 .await
1639 .map_err(VectorStoreError::datastore)
1640 }
1641}
1642
1643fn sqlite_distance_metric_from_schema(schema_sql: &str) -> SqliteDistanceMetric {
1644 let normalized = sqlite_normalized_schema(schema_sql);
1645
1646 if normalized.contains("distance_metric=cosine") {
1647 SqliteDistanceMetric::Cosine
1648 } else if normalized.contains("distance_metric=l1") {
1649 SqliteDistanceMetric::L1
1650 } else {
1651 SqliteDistanceMetric::L2
1652 }
1653}
1654
1655fn sqlite_normalized_schema(schema_sql: &str) -> String {
1656 schema_sql
1657 .chars()
1658 .filter(|c| !c.is_whitespace())
1659 .flat_map(char::to_lowercase)
1660 .collect()
1661}
1662
1663fn sqlite_schema_contains_metadata_column(schema_sql: &str, column: &SqliteMetadataColumn) -> bool {
1664 let normalized = sqlite_normalized_schema(schema_sql);
1665 let column_sql = format!(
1666 ",{}{}",
1667 column.name.to_ascii_lowercase(),
1668 column.metadata_type.vec0_name().to_ascii_lowercase()
1669 );
1670
1671 normalized.contains(&column_sql)
1672}
1673
1674struct SqliteSearchQuery {
1675 vector_where_clause: String,
1676 document_filter_clause: String,
1677 params: Vec<Value>,
1678}
1679
1680fn render_search_filters(
1681 req: &VectorSearchRequest<SqliteSearchFilter>,
1682 distance_metric: SqliteDistanceMetric,
1683 metadata_columns: &[SqliteMetadataColumn],
1684) -> Result<SqliteRenderedFilters, FilterError> {
1685 let score_expression = distance_metric.score_expression("?1", "e.embedding");
1686
1687 let mut filters = SqliteRenderedFilters::default();
1688 if let Some(threshold) = req.threshold() {
1689 filters.native.push(SqliteRenderedFilter {
1690 condition: format!("{score_expression} >= ?"),
1691 params: vec![Value::Real(threshold)],
1692 });
1693 }
1694 if let Some(filter) = req.filter() {
1695 filters.extend(filter.render_split(metadata_columns)?);
1696 }
1697
1698 Ok(filters)
1699}
1700
1701fn build_search_query(
1702 query_vec: Vec<f32>,
1703 filters: SqliteRenderedFilters,
1704 candidate_limit: u64,
1705) -> Result<SqliteSearchQuery, FilterError> {
1706 let brute_force = candidate_limit > SQLITE_VEC_MAX_K;
1713
1714 let mut conditions = Vec::new();
1715 if !brute_force {
1716 conditions.push("e.embedding MATCH ?".to_string());
1717 conditions.push("k = ?".to_string());
1718 }
1719 conditions.extend(
1720 filters
1721 .native
1722 .iter()
1723 .map(|filter| format!("({})", filter.condition)),
1724 );
1725
1726 let vector_where_clause = if conditions.is_empty() {
1729 String::new()
1730 } else {
1731 format!("WHERE {}", conditions.join(" AND "))
1732 };
1733 let document_filter_clause = if filters.post.is_empty() {
1734 String::new()
1735 } else {
1736 format!(
1737 "AND {}",
1738 filters
1739 .post
1740 .iter()
1741 .map(|filter| format!("({})", filter.condition))
1742 .collect::<Vec<_>>()
1743 .join(" AND ")
1744 )
1745 };
1746
1747 let query_vec = query_vec.into_iter().flat_map(f32::to_le_bytes).collect();
1748 let query_vec = Value::Blob(query_vec);
1749
1750 let mut params = if brute_force {
1758 vec![query_vec]
1759 } else {
1760 let candidate_limit = sqlite_limit_param(candidate_limit, "candidate limit")?;
1761 vec![query_vec.clone(), query_vec, candidate_limit]
1762 };
1763 params.extend(filters.native.into_iter().flat_map(|filter| filter.params));
1764 params.extend(filters.post.into_iter().flat_map(|filter| filter.params));
1765
1766 Ok(SqliteSearchQuery {
1767 vector_where_clause,
1768 document_filter_clause,
1769 params,
1770 })
1771}
1772
1773#[cfg(test)]
1774fn build_where_clause(
1775 req: &VectorSearchRequest<SqliteSearchFilter>,
1776 query_vec: Vec<f32>,
1777 distance_metric: SqliteDistanceMetric,
1778 metadata_columns: &[SqliteMetadataColumn],
1779 candidate_limit: u64,
1780) -> Result<(String, Vec<Value>), FilterError> {
1781 let filters = render_search_filters(req, distance_metric, metadata_columns)?;
1782 let query = build_search_query(query_vec, filters, candidate_limit)?;
1783 Ok((query.vector_where_clause, query.params))
1784}
1785
1786fn sqlite_limit_param(value: u64, name: &str) -> Result<Value, FilterError> {
1787 i64::try_from(value)
1788 .map(Value::Integer)
1789 .map_err(|_| FilterError::TypeError(format!("SQLite {name} `{value}` exceeds i64::MAX")))
1790}
1791
1792fn sqlite_column_value_error(
1793 index: usize,
1794 value_type: Type,
1795 column: &Column,
1796 message: impl Into<String>,
1797) -> rusqlite::Error {
1798 rusqlite::Error::FromSqlConversionFailure(
1799 index,
1800 value_type,
1801 Box::new(SqliteInternalError::ColumnValueError {
1802 column_name: column.name,
1803 column_type: column.col_type,
1804 message: message.into(),
1805 }),
1806 )
1807}
1808
1809fn sqlite_number_value(
1810 index: usize,
1811 value_type: Type,
1812 column: &Column,
1813 value: f64,
1814) -> rusqlite::Result<serde_json::Value> {
1815 let number = serde_json::Number::from_f64(value).ok_or_else(|| {
1816 sqlite_column_value_error(index, value_type, column, "non-finite float value")
1817 })?;
1818
1819 Ok(serde_json::Value::Number(number))
1820}
1821
1822fn sqlite_utf8_value<'a>(
1823 index: usize,
1824 value_type: Type,
1825 column: &Column,
1826 value: &'a [u8],
1827 label: &str,
1828) -> rusqlite::Result<&'a str> {
1829 std::str::from_utf8(value).map_err(|e| {
1830 sqlite_column_value_error(
1831 index,
1832 value_type,
1833 column,
1834 format!("invalid UTF-8 {label}: {e}"),
1835 )
1836 })
1837}
1838
1839fn sqlite_text_value(
1840 index: usize,
1841 value_type: Type,
1842 column: &Column,
1843 value: &[u8],
1844) -> rusqlite::Result<serde_json::Value> {
1845 let value = sqlite_utf8_value(index, value_type, column, value, "text")?;
1846
1847 Ok(serde_json::Value::String(value.to_string()))
1848}
1849
1850fn sqlite_column_declares_json(column_type: &str) -> bool {
1851 column_type
1852 .split_whitespace()
1853 .next()
1854 .is_some_and(|token| token.eq_ignore_ascii_case("JSON"))
1855}
1856
1857fn sqlite_json_text_value(
1858 index: usize,
1859 value_type: Type,
1860 column: &Column,
1861 value: &[u8],
1862) -> rusqlite::Result<serde_json::Value> {
1863 let value = sqlite_utf8_value(index, value_type, column, value, "JSON text")?;
1864
1865 serde_json::from_str(value).map_err(|e| {
1866 sqlite_column_value_error(index, value_type, column, format!("invalid JSON text: {e}"))
1867 })
1868}
1869
1870fn sqlite_column_value_to_json(
1871 index: usize,
1872 column: &Column,
1873 value: ValueRef<'_>,
1874) -> rusqlite::Result<serde_json::Value> {
1875 let value_type = value.data_type();
1876
1877 if sqlite_column_declares_json(column.col_type) {
1878 return match value {
1879 ValueRef::Null => Ok(serde_json::Value::Null),
1880 ValueRef::Text(value) => sqlite_json_text_value(index, value_type, column, value),
1881 ValueRef::Integer(value) => Ok(serde_json::Value::Number(value.into())),
1882 ValueRef::Real(value) => sqlite_number_value(index, value_type, column, value),
1883 ValueRef::Blob(value) => sqlite_json_text_value(index, value_type, column, value),
1884 };
1885 }
1886
1887 let column_affinity = SqliteColumnAffinity::from_column_type(column.col_type);
1888
1889 match (column_affinity, value) {
1890 (_, ValueRef::Null) => Ok(serde_json::Value::Null),
1891 (SqliteColumnAffinity::Boolean, ValueRef::Integer(0)) => Ok(serde_json::Value::Bool(false)),
1892 (SqliteColumnAffinity::Boolean, ValueRef::Integer(1)) => Ok(serde_json::Value::Bool(true)),
1893 (SqliteColumnAffinity::Boolean, _) => Err(sqlite_column_value_error(
1894 index,
1895 value_type,
1896 column,
1897 "stored SQLite boolean value must be 0 or 1",
1898 )),
1899 (_, ValueRef::Text(value)) => sqlite_text_value(index, value_type, column, value),
1900 (_, ValueRef::Integer(value)) => Ok(serde_json::Value::Number(value.into())),
1901 (_, ValueRef::Real(value)) => sqlite_number_value(index, value_type, column, value),
1902 (_, ValueRef::Blob(value)) => Ok(serde_json::to_value(value)
1903 .map_err(|e| sqlite_column_value_error(index, value_type, column, e.to_string()))?),
1904 }
1905}
1906
1907fn sqlite_id_value_to_string(index: usize, value: ValueRef<'_>) -> rusqlite::Result<String> {
1908 match value {
1909 ValueRef::Integer(value) => Ok(value.to_string()),
1910 ValueRef::Real(value) => Ok(value.to_string()),
1911 ValueRef::Text(value) => std::str::from_utf8(value)
1912 .map(ToString::to_string)
1913 .map_err(|e| {
1914 rusqlite::Error::FromSqlConversionFailure(
1915 index,
1916 Type::Text,
1917 Box::new(SqliteInternalError::ColumnValueError {
1918 column_name: "id",
1919 column_type: "TEXT",
1920 message: format!("invalid UTF-8 text: {e}"),
1921 }),
1922 )
1923 }),
1924 value => Err(rusqlite::Error::FromSqlConversionFailure(
1925 index,
1926 value.data_type(),
1927 Box::new(SqliteInternalError::ColumnValueError {
1928 column_name: "id",
1929 column_type: "TEXT or INTEGER",
1930 message: "id cannot be NULL or BLOB".to_string(),
1931 }),
1932 )),
1933 }
1934}
1935
1936impl<E: EmbeddingModel + std::marker::Sync, T: SqliteVectorStoreTable> VectorStoreIndex
1937 for SqliteVectorIndex<E, T>
1938{
1939 type Filter = SqliteSearchFilter;
1940
1941 async fn top_n<D>(
1942 &self,
1943 req: VectorSearchRequest<SqliteSearchFilter>,
1944 ) -> Result<Vec<(f64, String, D)>, VectorStoreError>
1945 where
1946 D: for<'de> Deserialize<'de>,
1947 {
1948 tracing::debug!("Finding top {} matches for query", req.samples() as usize);
1949 if req.samples() == 0 {
1950 return Ok(Vec::new());
1951 }
1952
1953 let columns = T::schema();
1954 let id_column_index = columns
1955 .iter()
1956 .position(|column| column.name == "id")
1957 .ok_or_else(|| {
1958 VectorStoreError::datastore(SqliteInternalError::MissingIdColumn(
1959 T::name().to_string(),
1960 ))
1961 })?;
1962
1963 let outer_select_cols = columns
1964 .iter()
1965 .map(|column| format!("d.{} AS {}", column.name, column.name))
1966 .collect::<Vec<_>>()
1967 .join(", ");
1968
1969 let rows = self
1970 .search_rows(&req, outer_select_cols, move |row| {
1971 let mut map = serde_json::Map::new();
1973 for (i, column) in columns.iter().enumerate() {
1974 let value = sqlite_column_value_to_json(i, column, row.get_ref(i)?)?;
1975 map.insert(column.name.to_string(), value);
1976 }
1977 let score: f64 = row.get(columns.len())?;
1978 let id = sqlite_id_value_to_string(id_column_index, row.get_ref(id_column_index)?)?;
1979
1980 Ok((id, serde_json::Value::Object(map), score))
1981 })
1982 .await?;
1983
1984 debug!("Found {} potential matches", rows.len());
1985 let mut top_n = Vec::new();
1986 for (id, doc_value, score) in rows {
1987 match serde_json::from_value::<D>(doc_value) {
1988 Ok(doc) => {
1989 top_n.push((score, id, doc));
1990 }
1991 Err(e) => {
1992 debug!("Failed to deserialize document {}: {}", id, e);
1993 continue;
1994 }
1995 }
1996 }
1997
1998 debug!("Returning {} matches", top_n.len());
1999 Ok(top_n)
2000 }
2001
2002 async fn top_n_ids(
2003 &self,
2004 req: VectorSearchRequest<SqliteSearchFilter>,
2005 ) -> Result<Vec<(f64, String)>, VectorStoreError> {
2006 tracing::debug!(
2007 "Finding top {} document IDs for query",
2008 req.samples() as usize
2009 );
2010 if req.samples() == 0 {
2011 return Ok(Vec::new());
2012 }
2013
2014 let results = self
2015 .search_rows(&req, "d.id".to_string(), |row| {
2016 Ok((
2017 row.get::<_, f64>(1)?,
2018 sqlite_id_value_to_string(0, row.get_ref(0)?)?,
2019 ))
2020 })
2021 .await?;
2022
2023 debug!("Found {} matching document IDs", results.len());
2024 Ok(results)
2025 }
2026}
2027
2028fn serialize_embedding(embedding: &Embedding) -> Vec<f32> {
2029 embedding.vec.iter().map(|x| *x as f32).collect()
2030}
2031
2032macro_rules! impl_column_value {
2033 ($($ty:ty => $col_type:literal, |$value:ident| $to_sql:expr;)*) => {$(
2034 impl ColumnValue for $ty {
2035 fn to_sql_value(&self) -> Value {
2036 let $value = self;
2037 $to_sql
2038 }
2039
2040 fn column_type(&self) -> &'static str {
2041 $col_type
2042 }
2043 }
2044 )*};
2045}
2046
2047impl_column_value! {
2048 String => "TEXT", |value| Value::Text(value.clone());
2049 i64 => "INTEGER", |value| Value::Integer(*value);
2050 i32 => "INTEGER", |value| Value::Integer(i64::from(*value));
2051 f64 => "FLOAT", |value| Value::Real(*value);
2052 f32 => "FLOAT", |value| Value::Real(f64::from(*value));
2053 bool => "BOOLEAN", |value| Value::Integer(if *value { 1 } else { 0 });
2054 serde_json::Value => "JSON", |value| Value::Text(value.to_string());
2055}
2056
2057#[cfg(test)]
2058mod tests {
2059 use super::*;
2060 use rig_core::embeddings::EmbeddingError;
2061 use rusqlite::ffi::{sqlite3, sqlite3_api_routines, sqlite3_auto_extension};
2062 use sqlite_vec::sqlite3_vec_init;
2063 use std::cmp::Ordering;
2064 use std::os::raw::c_char;
2065 use std::sync::Once;
2066 use tokio_rusqlite::Connection;
2067
2068 const SCORE_EPSILON: f64 = 1e-5;
2069
2070 fn test_metadata_columns() -> Vec<SqliteMetadataColumn> {
2071 vec![SqliteMetadataColumn {
2072 name: "category",
2073 metadata_type: SqliteMetadataType::Text,
2074 }]
2075 }
2076
2077 fn typed_metadata_columns() -> Vec<SqliteMetadataColumn> {
2078 vec![
2079 SqliteMetadataColumn {
2080 name: "priority",
2081 metadata_type: SqliteMetadataType::Integer,
2082 },
2083 SqliteMetadataColumn {
2084 name: "rating",
2085 metadata_type: SqliteMetadataType::Float,
2086 },
2087 SqliteMetadataColumn {
2088 name: "published",
2089 metadata_type: SqliteMetadataType::Boolean,
2090 },
2091 ]
2092 }
2093
2094 #[test]
2095 fn json_column_text_decodes_to_json_object() -> anyhow::Result<()> {
2096 let column = Column::new("metadata", "JSON");
2097 let value = sqlite_column_value_to_json(
2098 0,
2099 &column,
2100 ValueRef::Text(br#"{"knowledge_doc_id":361,"knowledge_id":1,"user_id":1}"#),
2101 )?;
2102
2103 let expected = serde_json::json!({
2104 "knowledge_doc_id": 361,
2105 "knowledge_id": 1,
2106 "user_id": 1
2107 });
2108 anyhow::ensure!(
2109 value == expected,
2110 "JSON column text should decode to a JSON object, got {value:?}"
2111 );
2112
2113 Ok(())
2114 }
2115
2116 #[test]
2117 fn text_column_json_looking_text_stays_string() -> anyhow::Result<()> {
2118 let column = Column::new("metadata", "TEXT");
2119 let value = sqlite_column_value_to_json(
2120 0,
2121 &column,
2122 ValueRef::Text(br#"{"knowledge_doc_id":361,"knowledge_id":1,"user_id":1}"#),
2123 )?;
2124
2125 let expected =
2126 serde_json::json!(r#"{"knowledge_doc_id":361,"knowledge_id":1,"user_id":1}"#);
2127 anyhow::ensure!(
2128 value == expected,
2129 "TEXT column should preserve JSON-looking text as a string, got {value:?}"
2130 );
2131
2132 Ok(())
2133 }
2134
2135 #[test]
2136 fn json_column_invalid_text_returns_conversion_error() -> anyhow::Result<()> {
2137 let column = Column::new("metadata", "JSON");
2138 let err = match sqlite_column_value_to_json(0, &column, ValueRef::Text(b"not json")) {
2139 Ok(value) => anyhow::bail!("invalid JSON column text should fail, got {value:?}"),
2140 Err(err) => err,
2141 };
2142
2143 anyhow::ensure!(
2144 matches!(
2145 err,
2146 rusqlite::Error::FromSqlConversionFailure(0, Type::Text, _)
2147 ),
2148 "invalid JSON column text should return a conversion error, got {err}"
2149 );
2150
2151 Ok(())
2152 }
2153
2154 #[test]
2155 fn serde_json_value_column_value_round_trips_json_column() -> anyhow::Result<()> {
2156 let value = serde_json::json!({
2157 "knowledge_doc_id": 361,
2158 "knowledge_id": 1,
2159 "user_id": 1
2160 });
2161 anyhow::ensure!(
2162 value.column_type() == "JSON",
2163 "serde_json::Value should declare JSON column type"
2164 );
2165
2166 let text = match value.to_sql_value() {
2167 Value::Text(text) => text,
2168 value => {
2169 anyhow::bail!("serde_json::Value should serialize as JSON text, got {value:?}")
2170 }
2171 };
2172
2173 let column = Column::new("metadata", "JSON");
2174 let round_trip = sqlite_column_value_to_json(0, &column, ValueRef::Text(text.as_bytes()))?;
2175 anyhow::ensure!(
2176 round_trip == value,
2177 "serde_json::Value should round-trip through a JSON column, got {round_trip:?}"
2178 );
2179
2180 Ok(())
2181 }
2182
2183 fn filter_error<T: std::fmt::Debug>(
2184 result: Result<T, FilterError>,
2185 context: &str,
2186 ) -> anyhow::Result<FilterError> {
2187 match result {
2188 Ok(value) => anyhow::bail!("{context} should have failed, got {value:?}"),
2189 Err(err) => Ok(err),
2190 }
2191 }
2192
2193 fn ensure_vector_store_filter_error<T: std::fmt::Debug>(
2194 result: Result<T, VectorStoreError>,
2195 context: &str,
2196 ) -> anyhow::Result<()> {
2197 match result {
2198 Err(VectorStoreError::FilterError(_)) => Ok(()),
2199 Err(err) => anyhow::bail!("{context} returned unexpected error: {err}"),
2200 Ok(value) => anyhow::bail!("{context} should have failed, got {value:?}"),
2201 }
2202 }
2203
2204 #[test]
2205 fn threshold_filter_uses_computed_similarity_expression() -> anyhow::Result<()> {
2206 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
2207 .query("needle")
2208 .samples(5)
2209 .threshold(0.95)
2210 .build();
2211
2212 let (where_clause, params) =
2213 build_where_clause(&req, vec![1.0, 0.0], SqliteDistanceMetric::Cosine, &[], 5)?;
2214
2215 anyhow::ensure!(
2216 where_clause.contains("e.embedding MATCH ?"),
2217 "missing vector match constraint: {where_clause}"
2218 );
2219 anyhow::ensure!(
2220 where_clause.contains("k = ?"),
2221 "missing vector k constraint: {where_clause}"
2222 );
2223 anyhow::ensure!(
2224 where_clause.contains("(1 - vec_distance_cosine(?1, e.embedding)) >= ?"),
2225 "threshold should use computed similarity expression: {where_clause}"
2226 );
2227 anyhow::ensure!(params.len() == 4, "unexpected params: {params:?}");
2228 anyhow::ensure!(
2229 params.get(3) == Some(&Value::Real(0.95)),
2230 "unexpected threshold param: {params:?}"
2231 );
2232
2233 Ok(())
2234 }
2235
2236 #[test]
2237 fn l2_threshold_filter_uses_l2_score_expression() -> anyhow::Result<()> {
2238 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
2239 .query("needle")
2240 .samples(5)
2241 .threshold(-1.5)
2242 .build();
2243
2244 let (where_clause, params) =
2245 build_where_clause(&req, vec![1.0, 0.0], SqliteDistanceMetric::L2, &[], 5)?;
2246
2247 anyhow::ensure!(
2248 where_clause.contains("(-vec_distance_l2(?1, e.embedding)) >= ?"),
2249 "threshold should use L2 score expression: {where_clause}"
2250 );
2251 anyhow::ensure!(params.len() == 4, "unexpected params: {params:?}");
2252 anyhow::ensure!(
2253 params.get(3) == Some(&Value::Real(-1.5)),
2254 "unexpected threshold param: {params:?}"
2255 );
2256
2257 Ok(())
2258 }
2259
2260 #[test]
2261 fn no_threshold_does_not_add_similarity_predicate() -> anyhow::Result<()> {
2262 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
2263 .query("needle")
2264 .samples(5)
2265 .build();
2266
2267 let (where_clause, params) =
2268 build_where_clause(&req, vec![1.0, 0.0], SqliteDistanceMetric::Cosine, &[], 5)?;
2269
2270 anyhow::ensure!(
2271 where_clause == "WHERE e.embedding MATCH ? AND k = ?",
2272 "unexpected where clause: {where_clause}"
2273 );
2274 anyhow::ensure!(params.len() == 3, "unexpected params: {params:?}");
2275
2276 Ok(())
2277 }
2278
2279 #[test]
2280 fn candidate_limit_at_k_cap_still_uses_knn_path() -> anyhow::Result<()> {
2281 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
2282 .query("needle")
2283 .samples(5)
2284 .build();
2285
2286 let (where_clause, params) = build_where_clause(
2287 &req,
2288 vec![1.0, 0.0],
2289 SqliteDistanceMetric::Cosine,
2290 &[],
2291 SQLITE_VEC_MAX_K,
2292 )?;
2293
2294 anyhow::ensure!(
2295 where_clause == "WHERE e.embedding MATCH ? AND k = ?",
2296 "candidate limit at the cap should keep the KNN path: {where_clause}"
2297 );
2298 anyhow::ensure!(params.len() == 3, "unexpected params: {params:?}");
2300 anyhow::ensure!(
2301 params.get(2) == Some(&Value::Integer(SQLITE_VEC_MAX_K as i64)),
2302 "k param should be the candidate limit: {params:?}"
2303 );
2304
2305 Ok(())
2306 }
2307
2308 #[test]
2309 fn candidate_limit_above_k_cap_falls_back_to_brute_force_scan() -> anyhow::Result<()> {
2310 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
2311 .query("needle")
2312 .samples(5)
2313 .build();
2314
2315 let (where_clause, params) = build_where_clause(
2316 &req,
2317 vec![1.0, 0.0],
2318 SqliteDistanceMetric::Cosine,
2319 &[],
2320 SQLITE_VEC_MAX_K + 1,
2321 )?;
2322
2323 anyhow::ensure!(
2327 !where_clause.contains("MATCH") && !where_clause.contains("k = ?"),
2328 "brute-force scan should drop the KNN constraints: {where_clause}"
2329 );
2330 anyhow::ensure!(
2331 where_clause.is_empty(),
2332 "brute-force scan without filters should emit no WHERE clause: {where_clause:?}"
2333 );
2334 anyhow::ensure!(params.len() == 1, "unexpected params: {params:?}");
2337 anyhow::ensure!(
2338 matches!(params.first(), Some(Value::Blob(_))),
2339 "remaining param should be the query vector: {params:?}"
2340 );
2341
2342 Ok(())
2343 }
2344
2345 #[test]
2346 fn brute_force_scan_keeps_filter_params_aligned() -> anyhow::Result<()> {
2347 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
2348 .query("needle")
2349 .samples(5)
2350 .threshold(0.95)
2351 .build();
2352
2353 let (where_clause, params) = build_where_clause(
2354 &req,
2355 vec![1.0, 0.0],
2356 SqliteDistanceMetric::Cosine,
2357 &[],
2358 SQLITE_VEC_MAX_K + 1,
2359 )?;
2360
2361 anyhow::ensure!(
2364 where_clause == "WHERE ((1 - vec_distance_cosine(?1, e.embedding)) >= ?)",
2365 "brute-force scan should keep native filters: {where_clause}"
2366 );
2367 anyhow::ensure!(params.len() == 2, "unexpected params: {params:?}");
2368 anyhow::ensure!(
2369 matches!(params.first(), Some(Value::Blob(_))),
2370 "first param should be the query vector: {params:?}"
2371 );
2372 anyhow::ensure!(
2373 params.get(1) == Some(&Value::Real(0.95)),
2374 "threshold param should follow the query vector: {params:?}"
2375 );
2376
2377 Ok(())
2378 }
2379
2380 #[test]
2381 fn default_filter_composes_under_or_as_a_tautology() -> anyhow::Result<()> {
2382 let filter = SqliteSearchFilter::default().or(SqliteSearchFilter::eq(
2383 "category",
2384 serde_json::json!("docs"),
2385 ));
2386
2387 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
2388 .query("needle")
2389 .samples(5)
2390 .filter(filter)
2391 .build();
2392
2393 let filters =
2394 render_search_filters(&req, SqliteDistanceMetric::Cosine, &test_metadata_columns())?;
2395 let query = build_search_query(vec![1.0, 0.0], filters, 5)?;
2396 anyhow::ensure!(
2397 query.document_filter_clause == "AND ((1 = 1) OR (d.category = ?))",
2398 "default() under OR should render as a tautology: {}",
2399 query.document_filter_clause
2400 );
2401
2402 Ok(())
2403 }
2404
2405 #[test]
2406 fn or_filter_uses_document_filter_to_preserve_boolean_semantics() -> anyhow::Result<()> {
2407 let filter = SqliteSearchFilter::eq("category", serde_json::json!("docs")).or(
2408 SqliteSearchFilter::eq("title", serde_json::json!("archive")),
2409 );
2410
2411 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
2412 .query("needle")
2413 .samples(5)
2414 .filter(filter)
2415 .build();
2416
2417 let filters =
2418 render_search_filters(&req, SqliteDistanceMetric::Cosine, &test_metadata_columns())?;
2419 anyhow::ensure!(
2420 filters.has_post_filters(),
2421 "OR filters should be applied after vector candidate search"
2422 );
2423 let query = build_search_query(vec![1.0, 0.0], filters, 5)?;
2424
2425 anyhow::ensure!(
2426 query.vector_where_clause == "WHERE e.embedding MATCH ? AND k = ?",
2427 "OR filters should not be partially pushed into sqlite-vec: {}",
2428 query.vector_where_clause
2429 );
2430 anyhow::ensure!(
2431 query.document_filter_clause == "AND ((d.category = ?) OR (d.title = ?))",
2432 "unexpected document filter clause: {}",
2433 query.document_filter_clause
2434 );
2435 anyhow::ensure!(
2436 query.params.get(3) == Some(&Value::Text("docs".to_string()))
2437 && query.params.get(4) == Some(&Value::Text("archive".to_string())),
2438 "unexpected OR filter params: {:?}",
2439 query.params
2440 );
2441
2442 Ok(())
2443 }
2444
2445 #[test]
2446 fn indexed_filter_uses_vec0_metadata_constraint() -> anyhow::Result<()> {
2447 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
2448 .query("needle")
2449 .samples(5)
2450 .filter(SqliteSearchFilter::eq(
2451 "category",
2452 serde_json::json!("docs"),
2453 ))
2454 .build();
2455
2456 let (where_clause, params) = build_where_clause(
2457 &req,
2458 vec![1.0, 0.0],
2459 SqliteDistanceMetric::Cosine,
2460 &test_metadata_columns(),
2461 5,
2462 )?;
2463
2464 anyhow::ensure!(
2465 where_clause == "WHERE e.embedding MATCH ? AND k = ? AND (e.category = ?)",
2466 "unexpected where clause: {where_clause}"
2467 );
2468 anyhow::ensure!(params.len() == 4, "unexpected params: {params:?}");
2469 anyhow::ensure!(
2470 params.get(3) == Some(&Value::Text("docs".to_string())),
2471 "unexpected filter param: {params:?}"
2472 );
2473
2474 Ok(())
2475 }
2476
2477 #[test]
2478 fn negated_eq_filter_uses_vec0_metadata_inequality() -> anyhow::Result<()> {
2479 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
2480 .query("needle")
2481 .samples(5)
2482 .filter(SqliteSearchFilter::eq("category", serde_json::json!("docs")).not())
2483 .build();
2484
2485 let (where_clause, params) = build_where_clause(
2486 &req,
2487 vec![1.0, 0.0],
2488 SqliteDistanceMetric::Cosine,
2489 &test_metadata_columns(),
2490 5,
2491 )?;
2492
2493 anyhow::ensure!(
2494 where_clause == "WHERE e.embedding MATCH ? AND k = ? AND (e.category != ?)",
2495 "unexpected where clause: {where_clause}"
2496 );
2497 anyhow::ensure!(params.len() == 4, "unexpected params: {params:?}");
2498 anyhow::ensure!(
2499 params.get(3) == Some(&Value::Text("docs".to_string())),
2500 "unexpected filter param: {params:?}"
2501 );
2502
2503 Ok(())
2504 }
2505
2506 #[test]
2507 fn negated_range_comparison_uses_vec0_metadata_boundary() -> anyhow::Result<()> {
2508 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
2509 .query("needle")
2510 .samples(5)
2511 .filter(SqliteSearchFilter::gt("priority", serde_json::json!(10)).not())
2512 .build();
2513
2514 let (where_clause, params) = build_where_clause(
2515 &req,
2516 vec![1.0, 0.0],
2517 SqliteDistanceMetric::Cosine,
2518 &typed_metadata_columns(),
2519 5,
2520 )?;
2521
2522 anyhow::ensure!(
2523 where_clause == "WHERE e.embedding MATCH ? AND k = ? AND (e.priority <= ?)",
2524 "unexpected where clause: {where_clause}"
2525 );
2526 anyhow::ensure!(params.len() == 4, "unexpected params: {params:?}");
2527 anyhow::ensure!(
2528 params.get(3) == Some(&Value::Integer(10)),
2529 "unexpected filter param: {params:?}"
2530 );
2531
2532 Ok(())
2533 }
2534
2535 #[test]
2536 fn negated_boolean_eq_filter_uses_vec0_metadata_inequality() -> anyhow::Result<()> {
2537 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
2538 .query("needle")
2539 .samples(5)
2540 .filter(SqliteSearchFilter::eq("published", serde_json::json!(true)).not())
2541 .build();
2542
2543 let (where_clause, params) = build_where_clause(
2544 &req,
2545 vec![1.0, 0.0],
2546 SqliteDistanceMetric::Cosine,
2547 &typed_metadata_columns(),
2548 5,
2549 )?;
2550
2551 anyhow::ensure!(
2552 where_clause == "WHERE e.embedding MATCH ? AND k = ? AND (e.published != ?)",
2553 "unexpected where clause: {where_clause}"
2554 );
2555 anyhow::ensure!(
2556 params.get(3) == Some(&Value::Integer(1)),
2557 "unexpected boolean filter param: {params:?}"
2558 );
2559
2560 Ok(())
2561 }
2562
2563 #[test]
2564 fn negated_between_filter_uses_document_filter() -> anyhow::Result<()> {
2565 let filter = SqliteSearchFilter::between("priority".to_string(), 1_i64..=10_i64).not();
2566 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
2567 .query("needle")
2568 .samples(5)
2569 .filter(filter)
2570 .build();
2571
2572 let filters = render_search_filters(
2573 &req,
2574 SqliteDistanceMetric::Cosine,
2575 &typed_metadata_columns(),
2576 )?;
2577 anyhow::ensure!(
2578 filters.has_post_filters(),
2579 "negated range filters should be applied after vector candidate search"
2580 );
2581 let query = build_search_query(vec![1.0, 0.0], filters, 5)?;
2582
2583 anyhow::ensure!(
2584 query.vector_where_clause == "WHERE e.embedding MATCH ? AND k = ?",
2585 "negated range filters should not be partially pushed into sqlite-vec: {}",
2586 query.vector_where_clause
2587 );
2588 anyhow::ensure!(
2589 query.document_filter_clause == "AND (NOT (d.priority between ? and ?))",
2590 "unexpected document filter clause: {}",
2591 query.document_filter_clause
2592 );
2593 anyhow::ensure!(
2594 query.params.get(3) == Some(&Value::Integer(1))
2595 && query.params.get(4) == Some(&Value::Integer(10)),
2596 "unexpected negated between params: {:?}",
2597 query.params
2598 );
2599
2600 Ok(())
2601 }
2602
2603 #[test]
2604 fn boolean_range_filter_is_rejected() -> anyhow::Result<()> {
2605 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
2606 .query("needle")
2607 .samples(5)
2608 .filter(SqliteSearchFilter::gt(
2609 "published",
2610 serde_json::json!(false),
2611 ))
2612 .build();
2613
2614 let err = filter_error(
2615 build_where_clause(
2616 &req,
2617 vec![1.0, 0.0],
2618 SqliteDistanceMetric::Cosine,
2619 &typed_metadata_columns(),
2620 5,
2621 ),
2622 "boolean range filters",
2623 )?;
2624
2625 anyhow::ensure!(
2626 err.to_string().contains("BOOLEAN"),
2627 "unexpected error for boolean range filter: {err}"
2628 );
2629
2630 Ok(())
2631 }
2632
2633 #[test]
2634 fn indexed_between_filter_uses_vec0_metadata_constraints() -> anyhow::Result<()> {
2635 let filter = SqliteSearchFilter::between("priority".to_string(), 1_i64..=10_i64);
2636 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
2637 .query("needle")
2638 .samples(5)
2639 .filter(filter)
2640 .build();
2641
2642 let (where_clause, params) = build_where_clause(
2643 &req,
2644 vec![1.0, 0.0],
2645 SqliteDistanceMetric::Cosine,
2646 &typed_metadata_columns(),
2647 5,
2648 )?;
2649
2650 anyhow::ensure!(
2651 where_clause
2652 == "WHERE e.embedding MATCH ? AND k = ? AND (e.priority >= ? AND e.priority <= ?)",
2653 "unexpected where clause: {where_clause}"
2654 );
2655 anyhow::ensure!(params.len() == 5, "unexpected params: {params:?}");
2656 anyhow::ensure!(
2657 params.get(3) == Some(&Value::Integer(1)) && params.get(4) == Some(&Value::Integer(10)),
2658 "between bounds should be bound as parameters: {params:?}"
2659 );
2660
2661 Ok(())
2662 }
2663
2664 #[test]
2665 fn mismatched_metadata_filter_value_types_are_rejected() -> anyhow::Result<()> {
2666 let cases = [
2667 (
2668 SqliteSearchFilter::eq("published", serde_json::json!("true")),
2669 "boolean filter value",
2670 ),
2671 (
2672 SqliteSearchFilter::gt("priority", serde_json::json!(1.5)),
2673 "integer filter value",
2674 ),
2675 (
2676 SqliteSearchFilter::eq("category", serde_json::json!({ "name": "docs" })),
2677 "string filter value",
2678 ),
2679 (
2680 SqliteSearchFilter::between(
2681 "priority".to_string(),
2682 "1".to_string()..="10".to_string(),
2683 ),
2684 "integer filter value",
2685 ),
2686 ];
2687
2688 for (filter, expected) in cases {
2689 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
2690 .query("needle")
2691 .samples(5)
2692 .filter(filter)
2693 .build();
2694
2695 let err = filter_error(
2696 build_where_clause(
2697 &req,
2698 vec![1.0, 0.0],
2699 SqliteDistanceMetric::Cosine,
2700 &typed_metadata_columns()
2701 .into_iter()
2702 .chain(test_metadata_columns())
2703 .collect::<Vec<_>>(),
2704 5,
2705 ),
2706 "mismatched metadata filter value",
2707 )?;
2708
2709 anyhow::ensure!(
2710 err.to_string().contains(expected),
2711 "unexpected error for mismatched metadata filter value: {err}"
2712 );
2713 }
2714
2715 Ok(())
2716 }
2717
2718 #[test]
2719 fn pattern_and_null_filters_use_document_filter() -> anyhow::Result<()> {
2720 let filter = SqliteSearchFilter::like("title".to_string(), "%O'Reilly%")
2721 .and(SqliteSearchFilter::glob("category".to_string(), "doc*"))
2722 .and(SqliteSearchFilter::is_null(
2723 "metadata->>'$.missing'".to_string(),
2724 ));
2725 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
2726 .query("needle")
2727 .samples(5)
2728 .filter(filter)
2729 .build();
2730
2731 let filters =
2732 render_search_filters(&req, SqliteDistanceMetric::Cosine, &test_metadata_columns())?;
2733 anyhow::ensure!(
2734 filters.has_post_filters(),
2735 "pattern and null filters should be applied after vector candidate search"
2736 );
2737 let query = build_search_query(vec![1.0, 0.0], filters, 5)?;
2738
2739 anyhow::ensure!(
2740 query.vector_where_clause == "WHERE e.embedding MATCH ? AND k = ?",
2741 "pattern filters should not be pushed into sqlite-vec: {}",
2742 query.vector_where_clause
2743 );
2744 anyhow::ensure!(
2745 query.document_filter_clause
2746 == "AND (d.title like ?) AND (d.category glob ?) AND (d.metadata->>'$.missing' is null)",
2747 "unexpected document filter clause: {}",
2748 query.document_filter_clause
2749 );
2750 anyhow::ensure!(
2751 query.params.get(3) == Some(&Value::Text("%O'Reilly%".to_string()))
2752 && query.params.get(4) == Some(&Value::Text("doc*".to_string())),
2753 "unexpected pattern filter params: {:?}",
2754 query.params
2755 );
2756
2757 Ok(())
2758 }
2759
2760 #[test]
2761 fn nonindexed_filters_use_document_filter() -> anyhow::Result<()> {
2762 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
2763 .query("needle")
2764 .samples(5)
2765 .filter(SqliteSearchFilter::eq("title", serde_json::json!("docs")))
2766 .build();
2767
2768 let filters =
2769 render_search_filters(&req, SqliteDistanceMetric::Cosine, &test_metadata_columns())?;
2770 anyhow::ensure!(
2771 filters.has_post_filters(),
2772 "non-indexed filters should be applied after vector candidate search"
2773 );
2774 let query = build_search_query(vec![1.0, 0.0], filters, 5)?;
2775
2776 anyhow::ensure!(
2777 query.vector_where_clause == "WHERE e.embedding MATCH ? AND k = ?",
2778 "unexpected vector where clause: {}",
2779 query.vector_where_clause
2780 );
2781 anyhow::ensure!(
2782 query.document_filter_clause == "AND (d.title = ?)",
2783 "unexpected document filter clause: {}",
2784 query.document_filter_clause
2785 );
2786 anyhow::ensure!(
2787 query.params.get(3) == Some(&Value::Text("docs".to_string())),
2788 "unexpected document filter param: {:?}",
2789 query.params
2790 );
2791
2792 Ok(())
2793 }
2794
2795 #[test]
2796 fn json_metadata_expression_uses_document_filter() -> anyhow::Result<()> {
2797 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
2798 .query("needle")
2799 .samples(5)
2800 .filter(SqliteSearchFilter::eq(
2801 "metadata->>'$.xxx'",
2802 serde_json::json!("vvv"),
2803 ))
2804 .build();
2805
2806 let filters =
2807 render_search_filters(&req, SqliteDistanceMetric::Cosine, &test_metadata_columns())?;
2808 anyhow::ensure!(
2809 filters.has_post_filters(),
2810 "JSON metadata expressions should be applied after vector candidate search"
2811 );
2812 let query = build_search_query(vec![1.0, 0.0], filters, 5)?;
2813
2814 anyhow::ensure!(
2815 query.vector_where_clause == "WHERE e.embedding MATCH ? AND k = ?",
2816 "unexpected vector where clause: {}",
2817 query.vector_where_clause
2818 );
2819 anyhow::ensure!(
2820 query.document_filter_clause == "AND (d.metadata->>'$.xxx' = ?)",
2821 "unexpected document filter clause: {}",
2822 query.document_filter_clause
2823 );
2824 anyhow::ensure!(
2825 query.params.get(3) == Some(&Value::Text("vvv".to_string())),
2826 "unexpected JSON metadata filter param: {:?}",
2827 query.params
2828 );
2829
2830 Ok(())
2831 }
2832
2833 #[test]
2834 fn json_metadata_arrow_expression_binds_rhs_as_json_text() -> anyhow::Result<()> {
2835 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
2836 .query("needle")
2837 .samples(5)
2838 .filter(SqliteSearchFilter::eq(
2839 "metadata->'$.xxx'",
2840 serde_json::json!("vvv"),
2841 ))
2842 .build();
2843
2844 let filters =
2845 render_search_filters(&req, SqliteDistanceMetric::Cosine, &test_metadata_columns())?;
2846 let query = build_search_query(vec![1.0, 0.0], filters, 5)?;
2847
2848 anyhow::ensure!(
2849 query.document_filter_clause == "AND (d.metadata->'$.xxx' = ?)",
2850 "unexpected document filter clause: {}",
2851 query.document_filter_clause
2852 );
2853 anyhow::ensure!(
2854 query.params.get(3) == Some(&Value::Text("\"vvv\"".to_string())),
2855 "SQLite `->` should compare against JSON text: {:?}",
2856 query.params
2857 );
2858
2859 Ok(())
2860 }
2861
2862 #[test]
2863 fn chained_json_metadata_expression_uses_final_operator_for_param_mode() -> anyhow::Result<()> {
2864 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
2865 .query("needle")
2866 .samples(5)
2867 .filter(SqliteSearchFilter::eq(
2868 "metadata->'$.nested'->>'$.xxx'",
2869 serde_json::json!("vvv"),
2870 ))
2871 .build();
2872
2873 let filters =
2874 render_search_filters(&req, SqliteDistanceMetric::Cosine, &test_metadata_columns())?;
2875 let query = build_search_query(vec![1.0, 0.0], filters, 5)?;
2876
2877 anyhow::ensure!(
2878 query.document_filter_clause == "AND (d.metadata->'$.nested'->>'$.xxx' = ?)",
2879 "unexpected document filter clause: {}",
2880 query.document_filter_clause
2881 );
2882 anyhow::ensure!(
2883 query.params.get(3) == Some(&Value::Text("vvv".to_string())),
2884 "final `->>` should compare against SQL scalar text: {:?}",
2885 query.params
2886 );
2887
2888 Ok(())
2889 }
2890
2891 #[test]
2892 fn unsupported_document_filter_expressions_are_rejected() -> anyhow::Result<()> {
2893 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
2894 .query("needle")
2895 .samples(5)
2896 .filter(SqliteSearchFilter::eq(
2897 "metadata) OR 1 = 1 --",
2898 serde_json::json!("vvv"),
2899 ))
2900 .build();
2901
2902 let err = filter_error(
2903 render_search_filters(&req, SqliteDistanceMetric::Cosine, &test_metadata_columns()),
2904 "unsupported document filter expressions",
2905 )?;
2906
2907 anyhow::ensure!(
2908 err.to_string()
2909 .contains("supported SQLite document filter expression"),
2910 "unexpected error for unsupported document filter expression: {err}"
2911 );
2912
2913 Ok(())
2914 }
2915
2916 #[tokio::test]
2917 async fn live_search_orders_by_similarity_and_applies_threshold() -> anyhow::Result<()> {
2918 let index = live_test_index(
2919 "live_search_orders_by_similarity_and_applies_threshold",
2920 vec![
2921 row("exact", "docs", "exact match", vec![1.0, 0.0]),
2922 row("close", "docs", "close match", vec![0.8, 0.6]),
2923 row("opposite", "docs", "opposite match", vec![-1.0, 0.0]),
2924 ],
2925 )
2926 .await?;
2927
2928 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
2929 .query("needle")
2930 .samples(3)
2931 .threshold(0.75)
2932 .build();
2933
2934 let results = index.top_n::<TestDocument>(req.clone()).await?;
2935 let ids = results
2936 .iter()
2937 .map(|(_, id, _)| id.as_str())
2938 .collect::<Vec<_>>();
2939 let exact_score = results.first().map(|(score, _, _)| *score);
2940 let close_score = results.get(1).map(|(score, _, _)| *score);
2941
2942 anyhow::ensure!(
2943 ids.as_slice() == ["exact", "close"],
2944 "unexpected ids: {ids:?}"
2945 );
2946 anyhow::ensure!(
2947 exact_score
2948 .zip(close_score)
2949 .is_some_and(|(exact, close)| exact > close),
2950 "expected exact score to be greater than close score: {results:?}"
2951 );
2952 anyhow::ensure!(
2953 results.iter().all(|(score, _, _)| *score > 0.75),
2954 "threshold should remove low-scoring rows: {results:?}"
2955 );
2956
2957 let id_results = index.top_n_ids(req).await?;
2958 let result_ids = id_results
2959 .iter()
2960 .map(|(_, id)| id.as_str())
2961 .collect::<Vec<_>>();
2962
2963 anyhow::ensure!(
2964 result_ids.as_slice() == ["exact", "close"],
2965 "unexpected top_n_ids ids: {id_results:?}"
2966 );
2967 anyhow::ensure!(
2968 id_results.iter().all(|(score, _)| *score > 0.75),
2969 "top_n_ids threshold should remove low-scoring rows: {id_results:?}"
2970 );
2971
2972 Ok(())
2973 }
2974
2975 #[tokio::test]
2976 async fn live_reinsert_same_document_id_removes_stale_vec0_candidates() -> anyhow::Result<()> {
2977 register_sqlite_vec_extension();
2978
2979 let conn = Connection::open(
2980 "file:live_reinsert_same_document_id_removes_stale_vec0_candidates?mode=memory",
2981 )
2982 .await?;
2983 let model = TestEmbeddingModel;
2984 let vector_store: SqliteVectorStore<_, TestDocument> =
2985 SqliteVectorStore::new(conn, &model).await?;
2986
2987 vector_store
2988 .add_rows(vec![row(
2989 "replace",
2990 "docs",
2991 "original near vector",
2992 vec![1.0, 0.0],
2993 )])
2994 .await?;
2995 vector_store
2996 .add_rows(vec![
2997 row("replace", "docs", "replacement far vector", vec![-1.0, 0.0]),
2998 row("fresh", "docs", "fresh near vector", vec![0.9, 0.1]),
2999 ])
3000 .await?;
3001
3002 let index = vector_store.index(model);
3003 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
3004 .query("needle")
3005 .samples(1)
3006 .build();
3007
3008 let results = index.top_n::<TestDocument>(req.clone()).await?;
3009 let ids = results
3010 .iter()
3011 .map(|(_, id, _)| id.as_str())
3012 .collect::<Vec<_>>();
3013 anyhow::ensure!(
3014 ids.as_slice() == ["fresh"],
3015 "stale replaced vectors should not consume sqlite-vec candidates: {results:?}"
3016 );
3017
3018 let id_results = index.top_n_ids(req).await?;
3019 let result_ids = id_results
3020 .iter()
3021 .map(|(_, id)| id.as_str())
3022 .collect::<Vec<_>>();
3023 anyhow::ensure!(
3024 result_ids.as_slice() == ["fresh"],
3025 "top_n_ids should not return or be starved by stale replaced vectors: {id_results:?}"
3026 );
3027
3028 Ok(())
3029 }
3030
3031 #[tokio::test]
3032 async fn live_reinsert_preserves_unrelated_multivector_embeddings() -> anyhow::Result<()> {
3033 register_sqlite_vec_extension();
3034
3035 let conn = Connection::open(
3036 "file:live_reinsert_preserves_unrelated_multivector_embeddings?mode=memory",
3037 )
3038 .await?;
3039 let model = TestEmbeddingModel;
3040 let vector_store: SqliteVectorStore<_, TestDocument> =
3041 SqliteVectorStore::new(conn, &model).await?;
3042
3043 let multi_document = TestDocument {
3044 id: "multi".to_string(),
3045 category: "docs".to_string(),
3046 title: "multi-vector document".to_string(),
3047 };
3048 vector_store
3049 .add_rows(vec![
3050 (
3051 multi_document.clone(),
3052 vec![
3053 Embedding {
3054 document: "far chunk".to_string(),
3055 vec: vec![-1.0, 0.0],
3056 },
3057 Embedding {
3058 document: "exact chunk".to_string(),
3059 vec: vec![1.0, 0.0],
3060 },
3061 ],
3062 ),
3063 row(
3064 "replace",
3065 "docs",
3066 "initial replacement vector",
3067 vec![0.8, 0.2],
3068 ),
3069 ])
3070 .await?;
3071 vector_store
3072 .add_rows(vec![row(
3073 "replace",
3074 "docs",
3075 "replacement far vector",
3076 vec![-1.0, 0.0],
3077 )])
3078 .await?;
3079
3080 let index = vector_store.index(model);
3081 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
3082 .query("needle")
3083 .samples(1)
3084 .threshold(0.9)
3085 .build();
3086
3087 let results = index.top_n::<TestDocument>(req.clone()).await?;
3088 let ids = results
3089 .iter()
3090 .map(|(_, id, _)| id.as_str())
3091 .collect::<Vec<_>>();
3092 anyhow::ensure!(
3093 ids.as_slice() == ["multi"],
3094 "reinsert should not delete another document's best embedding: {results:?}"
3095 );
3096
3097 let id_results = index.top_n_ids(req).await?;
3098 let result_ids = id_results
3099 .iter()
3100 .map(|(_, id)| id.as_str())
3101 .collect::<Vec<_>>();
3102 anyhow::ensure!(
3103 result_ids.as_slice() == ["multi"],
3104 "top_n_ids should preserve unrelated multivector embeddings after reinsert: {id_results:?}"
3105 );
3106
3107 Ok(())
3108 }
3109
3110 #[tokio::test]
3111 async fn live_multiple_embeddings_per_document_use_best_embedding() -> anyhow::Result<()> {
3112 let multi_document = TestDocument {
3113 id: "multi".to_string(),
3114 category: "docs".to_string(),
3115 title: "multi-vector document".to_string(),
3116 };
3117 let index = live_test_index(
3118 "live_multiple_embeddings_per_document_use_best_embedding",
3119 vec![
3120 (
3121 multi_document.clone(),
3122 vec![
3123 Embedding {
3124 document: "far chunk".to_string(),
3125 vec: vec![-1.0, 0.0],
3126 },
3127 Embedding {
3128 document: "exact chunk".to_string(),
3129 vec: vec![1.0, 0.0],
3130 },
3131 ],
3132 ),
3133 row("single", "docs", "single close chunk", vec![0.8, 0.6]),
3134 ],
3135 )
3136 .await?;
3137
3138 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
3139 .query("needle")
3140 .samples(2)
3141 .build();
3142 let results = index.top_n::<TestDocument>(req.clone()).await?;
3143 let ids = results
3144 .iter()
3145 .map(|(_, id, _)| id.as_str())
3146 .collect::<Vec<_>>();
3147 anyhow::ensure!(
3148 ids.as_slice() == ["multi", "single"],
3149 "top_n should return each document once using its best embedding: {results:?}"
3150 );
3151
3152 let id_results = index.top_n_ids(req).await?;
3153 let result_ids = id_results
3154 .iter()
3155 .map(|(_, id)| id.as_str())
3156 .collect::<Vec<_>>();
3157 anyhow::ensure!(
3158 result_ids.as_slice() == ["multi", "single"],
3159 "top_n_ids should return each document once using its best embedding: {id_results:?}"
3160 );
3161
3162 let threshold_req = VectorSearchRequest::<SqliteSearchFilter>::builder()
3163 .query("needle")
3164 .samples(2)
3165 .threshold(1.0)
3166 .build();
3167 let threshold_results = index.top_n::<TestDocument>(threshold_req.clone()).await?;
3168 let threshold_ids = threshold_results
3169 .iter()
3170 .map(|(_, id, _)| id.as_str())
3171 .collect::<Vec<_>>();
3172 anyhow::ensure!(
3173 threshold_ids.as_slice() == ["multi"],
3174 "threshold should include scores equal to the minimum and filter lower scores: {threshold_results:?}"
3175 );
3176
3177 let threshold_id_results = index.top_n_ids(threshold_req).await?;
3178 let threshold_result_ids = threshold_id_results
3179 .iter()
3180 .map(|(_, id)| id.as_str())
3181 .collect::<Vec<_>>();
3182 anyhow::ensure!(
3183 threshold_result_ids.as_slice() == ["multi"],
3184 "top_n_ids threshold should include scores equal to the minimum: {threshold_id_results:?}"
3185 );
3186
3187 Ok(())
3188 }
3189
3190 #[tokio::test]
3196 async fn live_multivector_search_beyond_knn_k_cap_succeeds() -> anyhow::Result<()> {
3197 let filler_chunks = (0..4100)
3202 .map(|i| Embedding {
3203 document: format!("filler chunk {i}"),
3204 vec: vec![0.0, 1.0],
3205 })
3206 .collect::<Vec<_>>();
3207 let filler_document = TestDocument {
3208 id: "filler".to_string(),
3209 category: "docs".to_string(),
3210 title: "many-embedding document".to_string(),
3211 };
3212
3213 let index = live_test_index(
3214 "live_multivector_search_beyond_knn_k_cap_succeeds",
3215 vec![
3216 (filler_document, filler_chunks),
3217 row("best", "docs", "best", vec![1.0, 0.0]),
3218 row("mid", "docs", "mid", vec![0.5, 0.5]),
3219 row("worst", "docs", "worst", vec![-1.0, 0.0]),
3220 ],
3221 )
3222 .await?;
3223
3224 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
3225 .query("needle")
3226 .samples(3)
3227 .build();
3228
3229 let results = index.top_n::<TestDocument>(req.clone()).await?;
3230 let ids = results
3231 .iter()
3232 .map(|(_, id, _)| id.as_str())
3233 .collect::<Vec<_>>();
3234 anyhow::ensure!(
3235 ids.as_slice() == ["best", "mid", "filler"],
3236 "brute-force scan should return the exact top-n past the knn k cap: {results:?}"
3237 );
3238
3239 let id_results = index.top_n_ids(req).await?;
3240 let result_ids = id_results
3241 .iter()
3242 .map(|(_, id)| id.as_str())
3243 .collect::<Vec<_>>();
3244 anyhow::ensure!(
3245 result_ids.as_slice() == ["best", "mid", "filler"],
3246 "top_n_ids should also brute-force past the knn k cap: {id_results:?}"
3247 );
3248
3249 Ok(())
3250 }
3251
3252 #[tokio::test]
3258 async fn live_post_filter_search_beyond_knn_k_cap_succeeds() -> anyhow::Result<()> {
3259 let mut rows = (0..4096)
3260 .map(|i| row(format!("noise-{i}"), "docs", "noise title", vec![1.0, 0.0]))
3261 .collect::<Vec<_>>();
3262 rows.push(row("wanted", "docs", "wanted title", vec![-1.0, 0.0]));
3266
3267 let index =
3268 live_test_index("live_post_filter_search_beyond_knn_k_cap_succeeds", rows).await?;
3269
3270 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
3271 .query("needle")
3272 .samples(1)
3273 .filter(SqliteSearchFilter::eq(
3274 "title",
3275 serde_json::json!("wanted title"),
3276 ))
3277 .build();
3278
3279 let results = index.top_n::<TestDocument>(req.clone()).await?;
3280 let ids = results
3281 .iter()
3282 .map(|(_, id, _)| id.as_str())
3283 .collect::<Vec<_>>();
3284 anyhow::ensure!(
3285 ids.as_slice() == ["wanted"],
3286 "exhaustive non-indexed filter past the knn k cap should still find the match: {results:?}"
3287 );
3288
3289 let id_results = index.top_n_ids(req).await?;
3290 let result_ids = id_results
3291 .iter()
3292 .map(|(_, id)| id.as_str())
3293 .collect::<Vec<_>>();
3294 anyhow::ensure!(
3295 result_ids.as_slice() == ["wanted"],
3296 "top_n_ids should also apply the exhaustive filter past the knn k cap: {id_results:?}"
3297 );
3298
3299 Ok(())
3300 }
3301
3302 #[tokio::test]
3303 async fn live_equal_score_results_are_ordered_by_document_id() -> anyhow::Result<()> {
3304 let index = live_test_index(
3305 "live_equal_score_results_are_ordered_by_document_id",
3306 vec![
3307 row("b", "docs", "second id exact match", vec![1.0, 0.0]),
3308 row("a", "docs", "first id exact match", vec![1.0, 0.0]),
3309 ],
3310 )
3311 .await?;
3312
3313 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
3314 .query("needle")
3315 .samples(2)
3316 .build();
3317
3318 let results = index.top_n::<TestDocument>(req.clone()).await?;
3319 let ids = results
3320 .iter()
3321 .map(|(_, id, _)| id.as_str())
3322 .collect::<Vec<_>>();
3323 anyhow::ensure!(
3324 ids.as_slice() == ["a", "b"],
3325 "equal-score top_n results should use document id as a stable tie-breaker: {results:?}"
3326 );
3327
3328 let id_results = index.top_n_ids(req).await?;
3329 let result_ids = id_results
3330 .iter()
3331 .map(|(_, id)| id.as_str())
3332 .collect::<Vec<_>>();
3333 anyhow::ensure!(
3334 result_ids.as_slice() == ["a", "b"],
3335 "equal-score top_n_ids results should use document id as a stable tie-breaker: {id_results:?}"
3336 );
3337
3338 Ok(())
3339 }
3340
3341 #[tokio::test]
3342 async fn live_common_sqlite_text_types_round_trip_in_top_n() -> anyhow::Result<()> {
3343 let index = live_common_type_test_index(
3344 "live_common_sqlite_text_types_round_trip_in_top_n",
3345 vec![common_type_row(
3346 "common",
3347 "varchar name",
3348 "clob notes",
3349 7,
3350 vec![1.0, 0.0],
3351 )],
3352 )
3353 .await?;
3354
3355 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
3356 .query("needle")
3357 .samples(1)
3358 .build();
3359 let results = index.top_n::<CommonTypeDocument>(req).await?;
3360
3361 let Some((_, id, doc)) = results.first() else {
3362 anyhow::bail!("expected common type document result");
3363 };
3364 anyhow::ensure!(id == "common", "unexpected id: {id}");
3365 anyhow::ensure!(
3366 doc.name == "varchar name",
3367 "VARCHAR value should round-trip: {doc:?}"
3368 );
3369 anyhow::ensure!(
3370 doc.notes == "clob notes",
3371 "CLOB value should round-trip: {doc:?}"
3372 );
3373 anyhow::ensure!(doc.rank == 7, "NUMERIC value should round-trip: {doc:?}");
3374
3375 Ok(())
3376 }
3377
3378 #[tokio::test]
3379 async fn live_json_column_structured_metadata_round_trips_in_top_n() -> anyhow::Result<()> {
3380 let metadata = StructuredMetadata {
3381 user_id: 1,
3382 knowledge_id: 1,
3383 knowledge_doc_id: 361,
3384 };
3385 let index = live_structured_json_metadata_test_index(
3386 "live_json_column_structured_metadata_round_trips_in_top_n",
3387 vec![structured_json_metadata_row(
3388 "structured",
3389 metadata.clone(),
3390 "metadata document",
3391 vec![1.0, 0.0],
3392 )],
3393 )
3394 .await?;
3395
3396 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
3397 .query("needle")
3398 .samples(1)
3399 .build();
3400 let results = index
3401 .top_n::<StructuredJsonMetadataDocument>(req.clone())
3402 .await?;
3403
3404 let Some((_, id, doc)) = results.first() else {
3405 anyhow::bail!("expected structured JSON metadata document result");
3406 };
3407 anyhow::ensure!(id == "structured", "unexpected id: {id}");
3408 anyhow::ensure!(
3409 doc.metadata == metadata,
3410 "JSON column should deserialize into structured metadata: {doc:?}"
3411 );
3412
3413 let id_results = index.top_n_ids(req).await?;
3414 anyhow::ensure!(
3415 id_results.first().is_some_and(|(_, id)| id == "structured"),
3416 "top_n_ids should still return the structured metadata document id: {id_results:?}"
3417 );
3418
3419 Ok(())
3420 }
3421
3422 #[tokio::test]
3423 async fn live_text_affinity_metadata_filters_during_candidate_search() -> anyhow::Result<()> {
3424 let index = live_common_type_test_index(
3425 "live_text_affinity_metadata_filters_during_candidate_search",
3426 vec![
3427 common_type_row("nearest", "misc", "nearest excluded", 1, vec![1.0, 0.0]),
3428 common_type_row("docs", "docs", "docs match", 2, vec![0.0, 1.0]),
3429 ],
3430 )
3431 .await?;
3432
3433 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
3434 .query("needle")
3435 .samples(1)
3436 .filter(SqliteSearchFilter::eq("name", serde_json::json!("docs")))
3437 .build();
3438
3439 let results = index.top_n::<CommonTypeDocument>(req.clone()).await?;
3440 let ids = results
3441 .iter()
3442 .map(|(_, id, _)| id.as_str())
3443 .collect::<Vec<_>>();
3444
3445 anyhow::ensure!(
3446 ids.as_slice() == ["docs"],
3447 "VARCHAR metadata filters should constrain sqlite-vec candidate search: {results:?}"
3448 );
3449
3450 let id_results = index.top_n_ids(req).await?;
3451 let result_ids = id_results
3452 .iter()
3453 .map(|(_, id)| id.as_str())
3454 .collect::<Vec<_>>();
3455
3456 anyhow::ensure!(
3457 result_ids.as_slice() == ["docs"],
3458 "top_n_ids should use VARCHAR metadata filters during candidate search: {id_results:?}"
3459 );
3460
3461 Ok(())
3462 }
3463
3464 #[tokio::test]
3465 async fn live_l2_metric_is_consistent() -> anyhow::Result<()> {
3466 let index = live_test_index_with_metric(
3467 "live_l2_metric_is_consistent",
3468 vec![
3469 row("exact", "docs", "exact match", vec![1.0, 0.0]),
3470 row("l2-close", "docs", "l2 close match", vec![1.0, 1.0]),
3471 row(
3472 "same-direction-far",
3473 "docs",
3474 "same direction far away",
3475 vec![10.0, 0.0],
3476 ),
3477 ],
3478 SqliteDistanceMetric::L2,
3479 )
3480 .await?;
3481
3482 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
3483 .query("needle")
3484 .samples(2)
3485 .threshold(-2.0)
3486 .build();
3487
3488 let results = index.top_n::<TestDocument>(req.clone()).await?;
3489 let ids = results
3490 .iter()
3491 .map(|(_, id, _)| id.as_str())
3492 .collect::<Vec<_>>();
3493 let exact_score = results
3494 .iter()
3495 .find(|(_, id, _)| id == "exact")
3496 .map(|(score, _, _)| *score);
3497 let close_score = results
3498 .iter()
3499 .find(|(_, id, _)| id == "l2-close")
3500 .map(|(score, _, _)| *score);
3501
3502 anyhow::ensure!(
3503 ids.as_slice() == ["exact", "l2-close"],
3504 "L2 search should return the nearest L2 candidates: {results:?}"
3505 );
3506 anyhow::ensure!(
3507 exact_score
3508 .zip(close_score)
3509 .is_some_and(|(exact, close)| exact > close && close > -2.0),
3510 "expected L2 scores to be ordered and thresholded: {results:?}"
3511 );
3512 anyhow::ensure!(
3513 results.iter().all(|(score, _, _)| *score > -2.0),
3514 "threshold should be applied to L2 scores: {results:?}"
3515 );
3516
3517 let id_results = index.top_n_ids(req).await?;
3518 let result_ids = id_results
3519 .iter()
3520 .map(|(_, id)| id.as_str())
3521 .collect::<Vec<_>>();
3522
3523 anyhow::ensure!(
3524 result_ids.as_slice() == ["exact", "l2-close"],
3525 "top_n_ids should use the same L2 metric: {id_results:?}"
3526 );
3527
3528 Ok(())
3529 }
3530
3531 #[tokio::test]
3532 async fn live_indexed_filter_is_applied_during_candidate_search() -> anyhow::Result<()> {
3533 let index = live_test_index(
3534 "live_indexed_filter_is_applied_during_candidate_search",
3535 vec![
3536 row(
3537 "nearest",
3538 "misc",
3539 "nearest excluded category",
3540 vec![1.0, 0.0],
3541 ),
3542 row("docs", "docs", "docs match", vec![0.0, 1.0]),
3543 ],
3544 )
3545 .await?;
3546
3547 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
3548 .query("needle")
3549 .samples(1)
3550 .filter(SqliteSearchFilter::eq(
3551 "category",
3552 serde_json::json!("docs"),
3553 ))
3554 .build();
3555
3556 let results = index.top_n::<TestDocument>(req.clone()).await?;
3557 let ids = results
3558 .iter()
3559 .map(|(_, id, _)| id.as_str())
3560 .collect::<Vec<_>>();
3561
3562 anyhow::ensure!(
3563 ids.as_slice() == ["docs"],
3564 "indexed filters should constrain sqlite-vec candidate search: {results:?}"
3565 );
3566
3567 let id_results = index.top_n_ids(req).await?;
3568 let result_ids = id_results
3569 .iter()
3570 .map(|(_, id)| id.as_str())
3571 .collect::<Vec<_>>();
3572
3573 anyhow::ensure!(
3574 result_ids.as_slice() == ["docs"],
3575 "top_n_ids should use indexed filters during candidate search: {id_results:?}"
3576 );
3577
3578 Ok(())
3579 }
3580
3581 #[tokio::test]
3582 async fn live_nonindexed_filter_is_applied_after_candidate_search() -> anyhow::Result<()> {
3583 let index = live_test_index(
3584 "live_nonindexed_filter_is_applied_after_candidate_search",
3585 vec![
3586 row("nearest", "docs", "nearest excluded title", vec![1.0, 0.0]),
3587 row("wanted", "docs", "wanted title", vec![0.0, 1.0]),
3588 ],
3589 )
3590 .await?;
3591
3592 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
3593 .query("needle")
3594 .samples(1)
3595 .filter(SqliteSearchFilter::eq(
3596 "title",
3597 serde_json::json!("wanted title"),
3598 ))
3599 .build();
3600
3601 let results = index.top_n::<TestDocument>(req.clone()).await?;
3602 let ids = results
3603 .iter()
3604 .map(|(_, id, _)| id.as_str())
3605 .collect::<Vec<_>>();
3606 anyhow::ensure!(
3607 ids.as_slice() == ["wanted"],
3608 "non-indexed filters should not be starved by the initial candidate limit: {results:?}"
3609 );
3610
3611 let id_results = index.top_n_ids(req).await?;
3612 let result_ids = id_results
3613 .iter()
3614 .map(|(_, id)| id.as_str())
3615 .collect::<Vec<_>>();
3616 anyhow::ensure!(
3617 result_ids.as_slice() == ["wanted"],
3618 "top_n_ids should apply non-indexed filters after candidate search: {id_results:?}"
3619 );
3620
3621 Ok(())
3622 }
3623
3624 #[tokio::test]
3625 async fn live_json_metadata_filter_is_applied_after_candidate_search() -> anyhow::Result<()> {
3626 let index = live_json_metadata_test_index(
3627 "live_json_metadata_filter_is_applied_after_candidate_search",
3628 vec![
3629 json_metadata_row("nearest", "docs", "skip", "nearest skipped", vec![1.0, 0.0]),
3630 json_metadata_row("matched", "docs", "vvv", "metadata match", vec![0.0, 1.0]),
3631 ],
3632 )
3633 .await?;
3634
3635 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
3636 .query("needle")
3637 .samples(1)
3638 .filter(SqliteSearchFilter::eq(
3639 "metadata->>'$.xxx'",
3640 serde_json::json!("vvv"),
3641 ))
3642 .build();
3643
3644 let results = index.top_n::<JsonMetadataDocument>(req.clone()).await?;
3645 let ids = results
3646 .iter()
3647 .map(|(_, id, _)| id.as_str())
3648 .collect::<Vec<_>>();
3649 anyhow::ensure!(
3650 ids.as_slice() == ["matched"],
3651 "JSON metadata filters should not be starved by the initial candidate limit: {results:?}"
3652 );
3653
3654 let id_results = index.top_n_ids(req).await?;
3655 let result_ids = id_results
3656 .iter()
3657 .map(|(_, id)| id.as_str())
3658 .collect::<Vec<_>>();
3659 anyhow::ensure!(
3660 result_ids.as_slice() == ["matched"],
3661 "top_n_ids should apply JSON metadata filters after candidate search: {id_results:?}"
3662 );
3663
3664 Ok(())
3665 }
3666
3667 #[tokio::test]
3668 async fn live_json_arrow_filter_compares_against_json_text() -> anyhow::Result<()> {
3669 let index = live_json_metadata_test_index(
3670 "live_json_arrow_filter_compares_against_json_text",
3671 vec![
3672 json_metadata_row("nearest", "docs", "skip", "nearest skipped", vec![1.0, 0.0]),
3673 json_metadata_row("matched", "docs", "vvv", "metadata match", vec![0.0, 1.0]),
3674 ],
3675 )
3676 .await?;
3677
3678 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
3679 .query("needle")
3680 .samples(1)
3681 .filter(SqliteSearchFilter::eq(
3682 "metadata->'$.xxx'",
3683 serde_json::json!("vvv"),
3684 ))
3685 .build();
3686
3687 let results = index.top_n::<JsonMetadataDocument>(req.clone()).await?;
3688 let ids = results
3689 .iter()
3690 .map(|(_, id, _)| id.as_str())
3691 .collect::<Vec<_>>();
3692 anyhow::ensure!(
3693 ids.as_slice() == ["matched"],
3694 "SQLite `->` JSON filters should compare against JSON text: {results:?}"
3695 );
3696
3697 let id_results = index.top_n_ids(req).await?;
3698 let result_ids = id_results
3699 .iter()
3700 .map(|(_, id)| id.as_str())
3701 .collect::<Vec<_>>();
3702 anyhow::ensure!(
3703 result_ids.as_slice() == ["matched"],
3704 "top_n_ids should apply SQLite `->` JSON filters: {id_results:?}"
3705 );
3706
3707 Ok(())
3708 }
3709
3710 #[tokio::test]
3711 async fn live_mixed_indexed_and_json_metadata_filters_are_applied() -> anyhow::Result<()> {
3712 let index = live_json_metadata_test_index(
3713 "live_mixed_indexed_and_json_metadata_filters_are_applied",
3714 vec![
3715 json_metadata_row(
3716 "nearest-docs",
3717 "docs",
3718 "skip",
3719 "nearest docs skipped by JSON metadata",
3720 vec![1.0, 0.0],
3721 ),
3722 json_metadata_row(
3723 "nearest-json",
3724 "misc",
3725 "vvv",
3726 "nearest JSON match skipped by category",
3727 vec![0.9, 0.1],
3728 ),
3729 json_metadata_row(
3730 "matched",
3731 "docs",
3732 "vvv",
3733 "matching category and JSON metadata",
3734 vec![0.0, 1.0],
3735 ),
3736 ],
3737 )
3738 .await?;
3739
3740 let filter = SqliteSearchFilter::eq("category", serde_json::json!("docs")).and(
3741 SqliteSearchFilter::eq("metadata->>'$.xxx'", serde_json::json!("vvv")),
3742 );
3743 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
3744 .query("needle")
3745 .samples(1)
3746 .filter(filter)
3747 .build();
3748
3749 let results = index.top_n::<JsonMetadataDocument>(req.clone()).await?;
3750 let ids = results
3751 .iter()
3752 .map(|(_, id, _)| id.as_str())
3753 .collect::<Vec<_>>();
3754 anyhow::ensure!(
3755 ids.as_slice() == ["matched"],
3756 "indexed and JSON metadata filters should both be applied: {results:?}"
3757 );
3758
3759 let id_results = index.top_n_ids(req).await?;
3760 let result_ids = id_results
3761 .iter()
3762 .map(|(_, id)| id.as_str())
3763 .collect::<Vec<_>>();
3764 anyhow::ensure!(
3765 result_ids.as_slice() == ["matched"],
3766 "top_n_ids should apply both indexed and JSON metadata filters: {id_results:?}"
3767 );
3768
3769 Ok(())
3770 }
3771
3772 #[tokio::test]
3773 async fn live_negated_eq_filter_is_applied_during_candidate_search() -> anyhow::Result<()> {
3774 let index = live_test_index(
3775 "live_negated_eq_filter_is_applied_during_candidate_search",
3776 vec![
3777 row(
3778 "nearest",
3779 "misc",
3780 "nearest excluded category",
3781 vec![1.0, 0.0],
3782 ),
3783 row("docs", "docs", "docs match", vec![0.0, 1.0]),
3784 ],
3785 )
3786 .await?;
3787
3788 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
3789 .query("needle")
3790 .samples(1)
3791 .filter(SqliteSearchFilter::eq("category", serde_json::json!("misc")).not())
3792 .build();
3793
3794 let results = index.top_n::<TestDocument>(req.clone()).await?;
3795 let ids = results
3796 .iter()
3797 .map(|(_, id, _)| id.as_str())
3798 .collect::<Vec<_>>();
3799
3800 anyhow::ensure!(
3801 ids.as_slice() == ["docs"],
3802 "negated filters should constrain sqlite-vec candidate search: {results:?}"
3803 );
3804
3805 let id_results = index.top_n_ids(req).await?;
3806 let result_ids = id_results
3807 .iter()
3808 .map(|(_, id)| id.as_str())
3809 .collect::<Vec<_>>();
3810
3811 anyhow::ensure!(
3812 result_ids.as_slice() == ["docs"],
3813 "top_n_ids should use negated filters during candidate search: {id_results:?}"
3814 );
3815
3816 Ok(())
3817 }
3818
3819 #[tokio::test]
3820 async fn live_top_n_reads_id_by_column_name_not_schema_position() -> anyhow::Result<()> {
3821 register_sqlite_vec_extension();
3822
3823 let conn = Connection::open(
3824 "file:live_top_n_reads_id_by_column_name_not_schema_position?mode=memory",
3825 )
3826 .await?;
3827 let model = TestEmbeddingModel;
3828 let vector_store: SqliteVectorStore<_, ReorderedIdDocument> =
3829 SqliteVectorStore::new(conn, &model).await?;
3830
3831 vector_store
3832 .add_rows(vec![
3833 reordered_id_row("winner", "winner title", "docs", vec![1.0, 0.0]),
3834 reordered_id_row("other", "other title", "docs", vec![0.0, 1.0]),
3835 ])
3836 .await?;
3837
3838 let index = vector_store.index(model);
3839 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
3840 .query("needle")
3841 .samples(1)
3842 .build();
3843
3844 let results = index.top_n::<ReorderedIdDocument>(req.clone()).await?;
3845 let Some((_, id, doc)) = results.first() else {
3846 anyhow::bail!("expected reordered-id result");
3847 };
3848 anyhow::ensure!(
3849 id == "winner",
3850 "top_n should return the id column, not the first schema column: {results:?}"
3851 );
3852 anyhow::ensure!(
3853 doc.id == "winner" && doc.title == "winner title",
3854 "document columns should still deserialize in schema order: {doc:?}"
3855 );
3856
3857 let id_results = index.top_n_ids(req).await?;
3858 anyhow::ensure!(
3859 id_results.first().map(|(_, id)| id.as_str()) == Some("winner"),
3860 "top_n_ids should agree with top_n id handling: {id_results:?}"
3861 );
3862
3863 Ok(())
3864 }
3865
3866 #[tokio::test]
3867 async fn live_internal_score_and_rank_column_names_do_not_shadow_search_columns()
3868 -> anyhow::Result<()> {
3869 register_sqlite_vec_extension();
3870
3871 let conn = Connection::open(
3872 "file:live_internal_score_and_rank_column_names_do_not_shadow_search_columns?mode=memory",
3873 )
3874 .await?;
3875 let model = TestEmbeddingModel;
3876 let vector_store: SqliteVectorStore<_, InternalAliasDocument> =
3877 SqliteVectorStore::new(conn, &model).await?;
3878
3879 vector_store
3880 .add_rows(vec![
3881 internal_alias_row(
3882 "winner",
3883 "payload score",
3884 "payload rank",
3885 "winner title",
3886 vec![1.0, 0.0],
3887 ),
3888 internal_alias_row(
3889 "other",
3890 "other score",
3891 "other rank",
3892 "other title",
3893 vec![0.0, 1.0],
3894 ),
3895 ])
3896 .await?;
3897
3898 let index = vector_store.index(model);
3899 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
3900 .query("needle")
3901 .samples(1)
3902 .threshold(0.9)
3903 .build();
3904
3905 let results = index.top_n::<InternalAliasDocument>(req.clone()).await?;
3906 let Some((score, id, doc)) = results.first() else {
3907 anyhow::bail!("expected internal-alias document result");
3908 };
3909
3910 anyhow::ensure!(id == "winner", "unexpected id: {results:?}");
3911 anyhow::ensure!(
3912 (*score - 1.0).abs() <= SCORE_EPSILON,
3913 "top_n should return computed score, not the document __rig_score column: {results:?}"
3914 );
3915 anyhow::ensure!(
3916 doc.rig_score == "payload score" && doc.rig_rank == "payload rank",
3917 "document columns with internal-looking names should still deserialize: {doc:?}"
3918 );
3919
3920 let id_results = index.top_n_ids(req).await?;
3921 anyhow::ensure!(
3922 id_results
3923 .first()
3924 .map(|(score, id)| ((*score - 1.0).abs() <= SCORE_EPSILON, id.as_str()))
3925 == Some((true, "winner")),
3926 "top_n_ids should agree with top_n despite internal-looking document columns: {id_results:?}"
3927 );
3928
3929 Ok(())
3930 }
3931
3932 #[tokio::test]
3933 async fn live_typed_columns_round_trip_and_filter_during_candidate_search() -> anyhow::Result<()>
3934 {
3935 let index = live_typed_test_index(
3936 "live_typed_columns_round_trip_and_filter_during_candidate_search",
3937 vec![
3938 typed_row(
3939 1,
3940 "misc",
3941 100,
3942 0.99,
3943 true,
3944 "nearest excluded by typed metadata",
3945 vec![1.0, 0.0],
3946 ),
3947 typed_row(2, "docs", 5, 0.95, true, "typed docs match", vec![0.0, 1.0]),
3948 typed_row(
3949 3,
3950 "docs",
3951 5,
3952 0.97,
3953 false,
3954 "unpublished docs match",
3955 vec![0.0, 0.9],
3956 ),
3957 ],
3958 )
3959 .await?;
3960
3961 let filter = SqliteSearchFilter::lt("priority", serde_json::json!(10))
3962 .and(SqliteSearchFilter::gt("rating", serde_json::json!(0.9)))
3963 .and(SqliteSearchFilter::eq("published", serde_json::json!(true)));
3964 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
3965 .query("needle")
3966 .samples(1)
3967 .filter(filter)
3968 .build();
3969
3970 let results = index.top_n::<TypedTestDocument>(req.clone()).await?;
3971 anyhow::ensure!(
3972 results.len() == 1,
3973 "expected one typed document result: {results:?}"
3974 );
3975
3976 let Some((_, id, doc)) = results.first() else {
3977 anyhow::bail!("expected one typed document result");
3978 };
3979 anyhow::ensure!(id == "2", "expected integer id to be returned as string");
3980 anyhow::ensure!(doc.id == 2, "typed integer id should round-trip: {doc:?}");
3981 anyhow::ensure!(
3982 doc.priority == 5,
3983 "typed integer field should round-trip: {doc:?}"
3984 );
3985 anyhow::ensure!(
3986 (doc.rating - 0.95).abs() < f64::EPSILON,
3987 "typed float field should round-trip: {doc:?}"
3988 );
3989 anyhow::ensure!(
3990 doc.published,
3991 "typed boolean field should round-trip: {doc:?}"
3992 );
3993
3994 let id_results = index.top_n_ids(req).await?;
3995 let result_ids = id_results
3996 .iter()
3997 .map(|(_, id)| id.as_str())
3998 .collect::<Vec<_>>();
3999 anyhow::ensure!(
4000 result_ids.as_slice() == ["2"],
4001 "top_n_ids should use the same typed metadata filters: {id_results:?}"
4002 );
4003
4004 Ok(())
4005 }
4006
4007 #[tokio::test]
4008 async fn live_boolean_range_filter_is_rejected() -> anyhow::Result<()> {
4009 let index = live_typed_test_index(
4010 "live_boolean_range_filter_is_rejected",
4011 vec![
4012 typed_row(
4013 1,
4014 "misc",
4015 1,
4016 0.5,
4017 false,
4018 "nearest unpublished doc",
4019 vec![1.0, 0.0],
4020 ),
4021 typed_row(2, "docs", 2, 0.7, true, "published doc", vec![0.0, 1.0]),
4022 ],
4023 )
4024 .await?;
4025
4026 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
4027 .query("needle")
4028 .samples(2)
4029 .filter(SqliteSearchFilter::gt(
4030 "published",
4031 serde_json::json!(false),
4032 ))
4033 .build();
4034
4035 ensure_vector_store_filter_error(
4036 index.top_n::<TypedTestDocument>(req.clone()).await,
4037 "top_n boolean range filter",
4038 )?;
4039 ensure_vector_store_filter_error(
4040 index.top_n_ids(req).await,
4041 "top_n_ids boolean range filter",
4042 )?;
4043
4044 Ok(())
4045 }
4046
4047 #[tokio::test]
4048 async fn live_mismatched_metadata_filter_value_type_is_rejected() -> anyhow::Result<()> {
4049 let index = live_typed_test_index(
4050 "live_mismatched_metadata_filter_value_type_is_rejected",
4051 vec![typed_row(
4052 1,
4053 "docs",
4054 1,
4055 0.95,
4056 true,
4057 "published doc",
4058 vec![1.0, 0.0],
4059 )],
4060 )
4061 .await?;
4062
4063 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
4064 .query("needle")
4065 .samples(1)
4066 .filter(SqliteSearchFilter::eq(
4067 "published",
4068 serde_json::json!("true"),
4069 ))
4070 .build();
4071
4072 ensure_vector_store_filter_error(
4073 index.top_n::<TypedTestDocument>(req.clone()).await,
4074 "top_n mismatched metadata filter value type",
4075 )?;
4076 ensure_vector_store_filter_error(
4077 index.top_n_ids(req).await,
4078 "top_n_ids mismatched metadata filter value type",
4079 )?;
4080
4081 Ok(())
4082 }
4083
4084 #[tokio::test]
4085 async fn live_matches_exact_oracle_for_metrics_filters_and_thresholds() -> anyhow::Result<()> {
4086 let query = vec![1.0, 0.0];
4087 let rows = oracle_test_rows();
4088 let filter = SqliteSearchFilter::eq("category", serde_json::json!("docs"))
4089 .and(SqliteSearchFilter::lt("priority", serde_json::json!(10)))
4090 .and(SqliteSearchFilter::gt("rating", serde_json::json!(0.8)))
4091 .and(SqliteSearchFilter::eq("published", serde_json::json!(true)));
4092
4093 for distance_metric in [
4094 SqliteDistanceMetric::Cosine,
4095 SqliteDistanceMetric::L2,
4096 SqliteDistanceMetric::L1,
4097 ] {
4098 let threshold = oracle_threshold(distance_metric);
4099 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
4100 .query("needle")
4101 .samples(u64::try_from(rows.len())?)
4102 .threshold(threshold)
4103 .filter(filter.clone())
4104 .build();
4105 let expected = exact_oracle_results(
4106 &rows,
4107 &query,
4108 distance_metric,
4109 threshold,
4110 rows.len(),
4111 |row| {
4112 row.category == "docs" && row.priority < 10 && row.rating > 0.8 && row.published
4113 },
4114 )?;
4115 let test_name =
4116 format!("live_matches_exact_oracle_for_{distance_metric:?}").to_ascii_lowercase();
4117 let index = live_typed_test_index_with_metric(
4118 &test_name,
4119 sqlite_oracle_rows(&rows),
4120 distance_metric,
4121 )
4122 .await?;
4123
4124 let results = index.top_n::<TypedTestDocument>(req.clone()).await?;
4125 let scored_ids = results
4126 .iter()
4127 .map(|(score, id, doc)| {
4128 anyhow::ensure!(
4129 id == &doc.id.to_string(),
4130 "top_n returned mismatched id and document: id={id}, doc={doc:?}"
4131 );
4132 Ok((*score, id.clone()))
4133 })
4134 .collect::<anyhow::Result<Vec<_>>>()?;
4135 assert_scored_ids_match(&scored_ids, &expected, distance_metric, "top_n")?;
4136
4137 let id_results = index.top_n_ids(req).await?;
4138 assert_scored_ids_match(&id_results, &expected, distance_metric, "top_n_ids")?;
4139 }
4140
4141 Ok(())
4142 }
4143
4144 #[tokio::test]
4145 async fn live_or_filter_preserves_mixed_document_semantics() -> anyhow::Result<()> {
4146 let index = live_test_index(
4147 "live_or_filter_preserves_mixed_document_semantics",
4148 vec![
4149 row(
4150 "nearest",
4151 "misc",
4152 "nearest excluded category",
4153 vec![1.0, 0.0],
4154 ),
4155 row("special", "misc", "special title", vec![0.9, 0.1]),
4156 row("docs", "docs", "far docs match", vec![0.0, 1.0]),
4157 ],
4158 )
4159 .await?;
4160
4161 let filter = SqliteSearchFilter::eq("category", serde_json::json!("docs")).or(
4162 SqliteSearchFilter::eq("title", serde_json::json!("special title")),
4163 );
4164
4165 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
4166 .query("needle")
4167 .samples(1)
4168 .filter(filter)
4169 .build();
4170
4171 let results = index.top_n::<TestDocument>(req.clone()).await?;
4172 let ids = results
4173 .iter()
4174 .map(|(_, id, _)| id.as_str())
4175 .collect::<Vec<_>>();
4176 anyhow::ensure!(
4177 ids.as_slice() == ["special"],
4178 "OR filters should be applied as a whole document predicate: {results:?}"
4179 );
4180
4181 let id_results = index.top_n_ids(req).await?;
4182 let result_ids = id_results
4183 .iter()
4184 .map(|(_, id)| id.as_str())
4185 .collect::<Vec<_>>();
4186 anyhow::ensure!(
4187 result_ids.as_slice() == ["special"],
4188 "top_n_ids should preserve OR document semantics: {id_results:?}"
4189 );
4190
4191 Ok(())
4192 }
4193
4194 #[tokio::test]
4195 async fn live_pattern_and_null_filters_are_applied_after_candidate_search() -> anyhow::Result<()>
4196 {
4197 let index = live_json_metadata_test_index(
4198 "live_pattern_and_null_filters_are_applied_after_candidate_search",
4199 vec![
4200 json_metadata_row("nearest", "docs", "skip", "skip this", vec![1.0, 0.0]),
4201 json_metadata_row("matched", "docs", "vvv", "metadata match", vec![0.0, 1.0]),
4202 ],
4203 )
4204 .await?;
4205
4206 let filter = SqliteSearchFilter::is_null("metadata->>'$.missing'".to_string())
4207 .and(SqliteSearchFilter::like("title".to_string(), "metadata%"))
4208 .and(SqliteSearchFilter::glob("category".to_string(), "doc*"));
4209
4210 let req = VectorSearchRequest::<SqliteSearchFilter>::builder()
4211 .query("needle")
4212 .samples(1)
4213 .filter(filter)
4214 .build();
4215
4216 let results = index.top_n::<JsonMetadataDocument>(req.clone()).await?;
4217 let ids = results
4218 .iter()
4219 .map(|(_, id, _)| id.as_str())
4220 .collect::<Vec<_>>();
4221 anyhow::ensure!(
4222 ids.as_slice() == ["matched"],
4223 "pattern and null filters should not be starved by the initial candidate limit: {results:?}"
4224 );
4225
4226 let id_results = index.top_n_ids(req).await?;
4227 let result_ids = id_results
4228 .iter()
4229 .map(|(_, id)| id.as_str())
4230 .collect::<Vec<_>>();
4231 anyhow::ensure!(
4232 result_ids.as_slice() == ["matched"],
4233 "top_n_ids should apply pattern and null filters after candidate search: {id_results:?}"
4234 );
4235
4236 Ok(())
4237 }
4238
4239 type SqliteExtensionFn =
4240 unsafe extern "C" fn(*mut sqlite3, *mut *mut c_char, *const sqlite3_api_routines) -> i32;
4241
4242 fn register_sqlite_vec_extension() {
4243 static REGISTER_SQLITE_VEC: Once = Once::new();
4244
4245 REGISTER_SQLITE_VEC.call_once(|| unsafe {
4246 sqlite3_auto_extension(Some(std::mem::transmute::<*const (), SqliteExtensionFn>(
4247 sqlite3_vec_init as *const (),
4248 )));
4249 });
4250 }
4251
4252 async fn live_test_index(
4253 name: &str,
4254 rows: Vec<(TestDocument, Vec<Embedding>)>,
4255 ) -> anyhow::Result<SqliteVectorIndex<TestEmbeddingModel, TestDocument>> {
4256 live_test_index_with_metric(name, rows, SqliteDistanceMetric::Cosine).await
4257 }
4258
4259 async fn live_test_index_with_metric(
4260 name: &str,
4261 rows: Vec<(TestDocument, Vec<Embedding>)>,
4262 distance_metric: SqliteDistanceMetric,
4263 ) -> anyhow::Result<SqliteVectorIndex<TestEmbeddingModel, TestDocument>> {
4264 register_sqlite_vec_extension();
4265
4266 let conn = Connection::open(format!("file:{name}?mode=memory")).await?;
4267 let model = TestEmbeddingModel;
4268 let vector_store =
4269 SqliteVectorStore::with_distance_metric(conn, &model, distance_metric).await?;
4270
4271 vector_store.add_rows(rows).await?;
4272
4273 Ok(vector_store.index(model))
4274 }
4275
4276 async fn live_typed_test_index(
4277 name: &str,
4278 rows: Vec<(TypedTestDocument, Vec<Embedding>)>,
4279 ) -> anyhow::Result<SqliteVectorIndex<TestEmbeddingModel, TypedTestDocument>> {
4280 live_typed_test_index_with_metric(name, rows, SqliteDistanceMetric::Cosine).await
4281 }
4282
4283 async fn live_typed_test_index_with_metric(
4284 name: &str,
4285 rows: Vec<(TypedTestDocument, Vec<Embedding>)>,
4286 distance_metric: SqliteDistanceMetric,
4287 ) -> anyhow::Result<SqliteVectorIndex<TestEmbeddingModel, TypedTestDocument>> {
4288 register_sqlite_vec_extension();
4289
4290 let conn = Connection::open(format!("file:{name}?mode=memory")).await?;
4291 let model = TestEmbeddingModel;
4292 let vector_store: SqliteVectorStore<_, TypedTestDocument> =
4293 SqliteVectorStore::with_distance_metric(conn, &model, distance_metric).await?;
4294
4295 vector_store.add_rows(rows).await?;
4296
4297 Ok(vector_store.index(model))
4298 }
4299
4300 async fn live_common_type_test_index(
4301 name: &str,
4302 rows: Vec<(CommonTypeDocument, Vec<Embedding>)>,
4303 ) -> anyhow::Result<SqliteVectorIndex<TestEmbeddingModel, CommonTypeDocument>> {
4304 register_sqlite_vec_extension();
4305
4306 let conn = Connection::open(format!("file:{name}?mode=memory")).await?;
4307 let model = TestEmbeddingModel;
4308 let vector_store: SqliteVectorStore<_, CommonTypeDocument> =
4309 SqliteVectorStore::new(conn, &model).await?;
4310
4311 vector_store.add_rows(rows).await?;
4312
4313 Ok(vector_store.index(model))
4314 }
4315
4316 async fn live_json_metadata_test_index(
4317 name: &str,
4318 rows: Vec<(JsonMetadataDocument, Vec<Embedding>)>,
4319 ) -> anyhow::Result<SqliteVectorIndex<TestEmbeddingModel, JsonMetadataDocument>> {
4320 register_sqlite_vec_extension();
4321
4322 let conn = Connection::open(format!("file:{name}?mode=memory")).await?;
4323 let model = TestEmbeddingModel;
4324 let vector_store: SqliteVectorStore<_, JsonMetadataDocument> =
4325 SqliteVectorStore::new(conn, &model).await?;
4326
4327 vector_store.add_rows(rows).await?;
4328
4329 Ok(vector_store.index(model))
4330 }
4331
4332 async fn live_structured_json_metadata_test_index(
4333 name: &str,
4334 rows: Vec<(StructuredJsonMetadataDocument, Vec<Embedding>)>,
4335 ) -> anyhow::Result<SqliteVectorIndex<TestEmbeddingModel, StructuredJsonMetadataDocument>> {
4336 register_sqlite_vec_extension();
4337
4338 let conn = Connection::open(format!("file:{name}?mode=memory")).await?;
4339 let model = TestEmbeddingModel;
4340 let vector_store: SqliteVectorStore<_, StructuredJsonMetadataDocument> =
4341 SqliteVectorStore::new(conn, &model).await?;
4342
4343 vector_store.add_rows(rows).await?;
4344
4345 Ok(vector_store.index(model))
4346 }
4347
4348 fn row(
4349 id: impl Into<String>,
4350 category: impl Into<String>,
4351 title: impl Into<String>,
4352 vec: Vec<f64>,
4353 ) -> (TestDocument, Vec<Embedding>) {
4354 let document = TestDocument {
4355 id: id.into(),
4356 category: category.into(),
4357 title: title.into(),
4358 };
4359
4360 (
4361 document.clone(),
4362 vec![Embedding {
4363 document: document.title,
4364 vec,
4365 }],
4366 )
4367 }
4368
4369 fn common_type_row(
4370 id: impl Into<String>,
4371 name: impl Into<String>,
4372 notes: impl Into<String>,
4373 rank: i64,
4374 vec: Vec<f64>,
4375 ) -> (CommonTypeDocument, Vec<Embedding>) {
4376 let document = CommonTypeDocument {
4377 id: id.into(),
4378 name: name.into(),
4379 notes: notes.into(),
4380 rank,
4381 };
4382
4383 (
4384 document.clone(),
4385 vec![Embedding {
4386 document: document.name.clone(),
4387 vec,
4388 }],
4389 )
4390 }
4391
4392 fn json_metadata_row(
4393 id: impl Into<String>,
4394 category: impl Into<String>,
4395 xxx: impl AsRef<str>,
4396 title: impl Into<String>,
4397 vec: Vec<f64>,
4398 ) -> (JsonMetadataDocument, Vec<Embedding>) {
4399 let document = JsonMetadataDocument {
4400 id: id.into(),
4401 category: category.into(),
4402 metadata: serde_json::json!({ "xxx": xxx.as_ref() }).to_string(),
4403 title: title.into(),
4404 };
4405
4406 (
4407 document.clone(),
4408 vec![Embedding {
4409 document: document.title.clone(),
4410 vec,
4411 }],
4412 )
4413 }
4414
4415 fn structured_json_metadata_row(
4416 id: impl Into<String>,
4417 metadata: StructuredMetadata,
4418 title: impl Into<String>,
4419 vec: Vec<f64>,
4420 ) -> (StructuredJsonMetadataDocument, Vec<Embedding>) {
4421 let document = StructuredJsonMetadataDocument {
4422 id: id.into(),
4423 metadata,
4424 title: title.into(),
4425 };
4426
4427 (
4428 document.clone(),
4429 vec![Embedding {
4430 document: document.title.clone(),
4431 vec,
4432 }],
4433 )
4434 }
4435
4436 fn reordered_id_row(
4437 id: impl Into<String>,
4438 title: impl Into<String>,
4439 category: impl Into<String>,
4440 vec: Vec<f64>,
4441 ) -> (ReorderedIdDocument, Vec<Embedding>) {
4442 let document = ReorderedIdDocument {
4443 title: title.into(),
4444 id: id.into(),
4445 category: category.into(),
4446 };
4447
4448 (
4449 document.clone(),
4450 vec![Embedding {
4451 document: document.title.clone(),
4452 vec,
4453 }],
4454 )
4455 }
4456
4457 fn internal_alias_row(
4458 id: impl Into<String>,
4459 rig_score: impl Into<String>,
4460 rig_rank: impl Into<String>,
4461 title: impl Into<String>,
4462 vec: Vec<f64>,
4463 ) -> (InternalAliasDocument, Vec<Embedding>) {
4464 let document = InternalAliasDocument {
4465 id: id.into(),
4466 rig_score: rig_score.into(),
4467 rig_rank: rig_rank.into(),
4468 title: title.into(),
4469 };
4470
4471 (
4472 document.clone(),
4473 vec![Embedding {
4474 document: document.title.clone(),
4475 vec,
4476 }],
4477 )
4478 }
4479
4480 fn typed_row(
4481 id: i64,
4482 category: impl Into<String>,
4483 priority: i64,
4484 rating: f64,
4485 published: bool,
4486 title: impl Into<String>,
4487 vec: Vec<f64>,
4488 ) -> (TypedTestDocument, Vec<Embedding>) {
4489 let document = TypedTestDocument {
4490 id,
4491 category: category.into(),
4492 priority,
4493 rating,
4494 published,
4495 title: title.into(),
4496 };
4497
4498 (
4499 document.clone(),
4500 vec![Embedding {
4501 document: document.title,
4502 vec,
4503 }],
4504 )
4505 }
4506
4507 #[derive(Clone, Debug)]
4508 struct OracleRow {
4509 document: TypedTestDocument,
4510 embedding: Vec<f64>,
4511 }
4512
4513 #[derive(Debug)]
4514 struct ExpectedScoredId {
4515 id: String,
4516 score: f64,
4517 }
4518
4519 fn oracle_test_rows() -> Vec<OracleRow> {
4520 vec![
4521 oracle_row(1, "docs", 1, 0.95, true, "exact match", vec![1.0, 0.0]),
4522 oracle_row(2, "docs", 2, 0.90, true, "close match", vec![0.8, 0.6]),
4523 oracle_row(3, "docs", 3, 0.81, true, "borderline match", vec![0.5, 0.5]),
4524 oracle_row(
4525 4,
4526 "docs",
4527 4,
4528 0.70,
4529 true,
4530 "filtered by rating",
4531 vec![0.95, 0.05],
4532 ),
4533 oracle_row(
4534 5,
4535 "docs",
4536 15,
4537 0.99,
4538 true,
4539 "filtered by priority",
4540 vec![1.0, 0.0],
4541 ),
4542 oracle_row(
4543 6,
4544 "docs",
4545 5,
4546 0.99,
4547 false,
4548 "filtered by published",
4549 vec![1.0, 0.0],
4550 ),
4551 oracle_row(
4552 7,
4553 "misc",
4554 1,
4555 0.99,
4556 true,
4557 "filtered by category",
4558 vec![1.0, 0.0],
4559 ),
4560 oracle_row(8, "docs", 5, 0.95, true, "far match", vec![0.0, 1.0]),
4561 ]
4562 }
4563
4564 fn oracle_row(
4565 id: i64,
4566 category: impl Into<String>,
4567 priority: i64,
4568 rating: f64,
4569 published: bool,
4570 title: impl Into<String>,
4571 embedding: Vec<f64>,
4572 ) -> OracleRow {
4573 OracleRow {
4574 document: TypedTestDocument {
4575 id,
4576 category: category.into(),
4577 priority,
4578 rating,
4579 published,
4580 title: title.into(),
4581 },
4582 embedding,
4583 }
4584 }
4585
4586 fn sqlite_oracle_rows(rows: &[OracleRow]) -> Vec<(TypedTestDocument, Vec<Embedding>)> {
4587 rows.iter()
4588 .map(|row| {
4589 (
4590 row.document.clone(),
4591 vec![Embedding {
4592 document: row.document.title.clone(),
4593 vec: row.embedding.clone(),
4594 }],
4595 )
4596 })
4597 .collect()
4598 }
4599
4600 fn oracle_threshold(distance_metric: SqliteDistanceMetric) -> f64 {
4601 match distance_metric {
4602 SqliteDistanceMetric::Cosine => 0.75,
4603 SqliteDistanceMetric::L2 => -0.8,
4604 SqliteDistanceMetric::L1 => -0.9,
4605 }
4606 }
4607
4608 fn exact_oracle_results(
4609 rows: &[OracleRow],
4610 query: &[f64],
4611 distance_metric: SqliteDistanceMetric,
4612 threshold: f64,
4613 samples: usize,
4614 filter: impl Fn(&TypedTestDocument) -> bool,
4615 ) -> anyhow::Result<Vec<ExpectedScoredId>> {
4616 let mut expected = Vec::new();
4617 for row in rows {
4618 if !filter(&row.document) {
4619 continue;
4620 }
4621
4622 let score = oracle_score(distance_metric, query, &row.embedding)?;
4623 if score >= threshold {
4624 expected.push(ExpectedScoredId {
4625 id: row.document.id.to_string(),
4626 score,
4627 });
4628 }
4629 }
4630
4631 sort_expected_scores(&mut expected);
4632 expected.truncate(samples);
4633 Ok(expected)
4634 }
4635
4636 fn sort_expected_scores(expected: &mut [ExpectedScoredId]) {
4637 expected.sort_by(|lhs, rhs| {
4638 rhs.score
4639 .partial_cmp(&lhs.score)
4640 .unwrap_or(Ordering::Equal)
4641 .then_with(|| lhs.id.cmp(&rhs.id))
4642 });
4643 }
4644
4645 fn oracle_score(
4646 distance_metric: SqliteDistanceMetric,
4647 query: &[f64],
4648 embedding: &[f64],
4649 ) -> anyhow::Result<f64> {
4650 anyhow::ensure!(
4651 query.len() == embedding.len(),
4652 "query and embedding dimensions differ: query={}, embedding={}",
4653 query.len(),
4654 embedding.len()
4655 );
4656
4657 let query = query.iter().map(|value| *value as f32).collect::<Vec<_>>();
4658 let embedding = embedding
4659 .iter()
4660 .map(|value| *value as f32)
4661 .collect::<Vec<_>>();
4662
4663 let score = match distance_metric {
4664 SqliteDistanceMetric::Cosine => {
4665 let dot = query
4666 .iter()
4667 .zip(&embedding)
4668 .map(|(lhs, rhs)| lhs * rhs)
4669 .sum::<f32>();
4670 let query_norm = query.iter().map(|value| value * value).sum::<f32>().sqrt();
4671 let embedding_norm = embedding
4672 .iter()
4673 .map(|value| value * value)
4674 .sum::<f32>()
4675 .sqrt();
4676 anyhow::ensure!(
4677 query_norm > 0.0 && embedding_norm > 0.0,
4678 "cosine oracle requires non-zero vectors"
4679 );
4680 dot / (query_norm * embedding_norm)
4681 }
4682 SqliteDistanceMetric::L2 => -query
4683 .iter()
4684 .zip(&embedding)
4685 .map(|(lhs, rhs)| {
4686 let delta = lhs - rhs;
4687 delta * delta
4688 })
4689 .sum::<f32>()
4690 .sqrt(),
4691 SqliteDistanceMetric::L1 => -query
4692 .iter()
4693 .zip(&embedding)
4694 .map(|(lhs, rhs)| (lhs - rhs).abs())
4695 .sum::<f32>(),
4696 };
4697
4698 Ok(f64::from(score))
4699 }
4700
4701 fn assert_scored_ids_match(
4702 actual: &[(f64, String)],
4703 expected: &[ExpectedScoredId],
4704 distance_metric: SqliteDistanceMetric,
4705 context: &str,
4706 ) -> anyhow::Result<()> {
4707 let actual_ids = actual.iter().map(|(_, id)| id.as_str()).collect::<Vec<_>>();
4708 let expected_ids = expected
4709 .iter()
4710 .map(|expected| expected.id.as_str())
4711 .collect::<Vec<_>>();
4712 anyhow::ensure!(
4713 actual_ids == expected_ids,
4714 "{context} ids for {distance_metric:?} did not match exact oracle: actual={actual:?}, expected={expected:?}"
4715 );
4716
4717 for ((actual_score, actual_id), expected) in actual.iter().zip(expected) {
4718 anyhow::ensure!(
4719 (actual_score - expected.score).abs() <= SCORE_EPSILON,
4720 "{context} score for {distance_metric:?} id `{actual_id}` did not match exact oracle: actual={actual_score}, expected={}",
4721 expected.score
4722 );
4723 }
4724
4725 Ok(())
4726 }
4727
4728 #[derive(Clone, Debug, Deserialize, Serialize)]
4729 struct TestDocument {
4730 id: String,
4731 category: String,
4732 title: String,
4733 }
4734
4735 impl SqliteVectorStoreTable for TestDocument {
4736 fn name() -> &'static str {
4737 "live_test_documents"
4738 }
4739
4740 fn schema() -> Vec<Column> {
4741 vec![
4742 Column::new("id", "TEXT PRIMARY KEY"),
4743 Column::new("category", "TEXT").indexed(),
4744 Column::new("title", "TEXT"),
4745 ]
4746 }
4747
4748 fn id(&self) -> String {
4749 self.id.clone()
4750 }
4751
4752 fn column_values(&self) -> Vec<(&'static str, Box<dyn ColumnValue>)> {
4753 vec![
4754 ("id", Box::new(self.id.clone())),
4755 ("category", Box::new(self.category.clone())),
4756 ("title", Box::new(self.title.clone())),
4757 ]
4758 }
4759 }
4760
4761 #[derive(Clone, Debug, Deserialize, Serialize)]
4762 struct ReorderedIdDocument {
4763 title: String,
4764 id: String,
4765 category: String,
4766 }
4767
4768 impl SqliteVectorStoreTable for ReorderedIdDocument {
4769 fn name() -> &'static str {
4770 "live_reordered_id_test_documents"
4771 }
4772
4773 fn schema() -> Vec<Column> {
4774 vec![
4775 Column::new("title", "TEXT"),
4776 Column::new("id", "TEXT PRIMARY KEY"),
4777 Column::new("category", "TEXT").indexed(),
4778 ]
4779 }
4780
4781 fn id(&self) -> String {
4782 self.id.clone()
4783 }
4784
4785 fn column_values(&self) -> Vec<(&'static str, Box<dyn ColumnValue>)> {
4786 vec![
4787 ("title", Box::new(self.title.clone())),
4788 ("id", Box::new(self.id.clone())),
4789 ("category", Box::new(self.category.clone())),
4790 ]
4791 }
4792 }
4793
4794 #[derive(Clone, Debug, Deserialize, Serialize)]
4795 struct InternalAliasDocument {
4796 id: String,
4797 #[serde(rename = "__rig_score")]
4798 rig_score: String,
4799 #[serde(rename = "__rig_rank")]
4800 rig_rank: String,
4801 title: String,
4802 }
4803
4804 impl SqliteVectorStoreTable for InternalAliasDocument {
4805 fn name() -> &'static str {
4806 "live_internal_alias_test_documents"
4807 }
4808
4809 fn schema() -> Vec<Column> {
4810 vec![
4811 Column::new("id", "TEXT PRIMARY KEY"),
4812 Column::new("__rig_score", "TEXT"),
4813 Column::new("__rig_rank", "TEXT"),
4814 Column::new("title", "TEXT"),
4815 ]
4816 }
4817
4818 fn id(&self) -> String {
4819 self.id.clone()
4820 }
4821
4822 fn column_values(&self) -> Vec<(&'static str, Box<dyn ColumnValue>)> {
4823 vec![
4824 ("id", Box::new(self.id.clone())),
4825 ("__rig_score", Box::new(self.rig_score.clone())),
4826 ("__rig_rank", Box::new(self.rig_rank.clone())),
4827 ("title", Box::new(self.title.clone())),
4828 ]
4829 }
4830 }
4831
4832 #[derive(Clone, Debug, Deserialize, Serialize)]
4833 struct CommonTypeDocument {
4834 id: String,
4835 name: String,
4836 notes: String,
4837 rank: i64,
4838 }
4839
4840 impl SqliteVectorStoreTable for CommonTypeDocument {
4841 fn name() -> &'static str {
4842 "live_common_type_test_documents"
4843 }
4844
4845 fn schema() -> Vec<Column> {
4846 vec![
4847 Column::new("id", "TEXT PRIMARY KEY"),
4848 Column::new("name", "VARCHAR(255)").indexed(),
4849 Column::new("notes", "CLOB"),
4850 Column::new("rank", "NUMERIC"),
4851 ]
4852 }
4853
4854 fn id(&self) -> String {
4855 self.id.clone()
4856 }
4857
4858 fn column_values(&self) -> Vec<(&'static str, Box<dyn ColumnValue>)> {
4859 vec![
4860 ("id", Box::new(self.id.clone())),
4861 ("name", Box::new(self.name.clone())),
4862 ("notes", Box::new(self.notes.clone())),
4863 ("rank", Box::new(self.rank)),
4864 ]
4865 }
4866 }
4867
4868 #[derive(Clone, Debug, Deserialize, Serialize)]
4869 struct JsonMetadataDocument {
4870 id: String,
4871 category: String,
4872 metadata: String,
4873 title: String,
4874 }
4875
4876 impl SqliteVectorStoreTable for JsonMetadataDocument {
4877 fn name() -> &'static str {
4878 "live_json_metadata_test_documents"
4879 }
4880
4881 fn schema() -> Vec<Column> {
4882 vec![
4883 Column::new("id", "TEXT PRIMARY KEY"),
4884 Column::new("category", "TEXT").indexed(),
4885 Column::new("metadata", "TEXT"),
4886 Column::new("title", "TEXT"),
4887 ]
4888 }
4889
4890 fn id(&self) -> String {
4891 self.id.clone()
4892 }
4893
4894 fn column_values(&self) -> Vec<(&'static str, Box<dyn ColumnValue>)> {
4895 vec![
4896 ("id", Box::new(self.id.clone())),
4897 ("category", Box::new(self.category.clone())),
4898 ("metadata", Box::new(self.metadata.clone())),
4899 ("title", Box::new(self.title.clone())),
4900 ]
4901 }
4902 }
4903
4904 #[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
4905 struct StructuredMetadata {
4906 user_id: i64,
4907 knowledge_id: i64,
4908 knowledge_doc_id: i64,
4909 }
4910
4911 #[derive(Clone, Debug, Deserialize, Serialize)]
4912 struct StructuredJsonMetadataDocument {
4913 id: String,
4914 metadata: StructuredMetadata,
4915 title: String,
4916 }
4917
4918 impl SqliteVectorStoreTable for StructuredJsonMetadataDocument {
4919 fn name() -> &'static str {
4920 "live_structured_json_metadata_test_documents"
4921 }
4922
4923 fn schema() -> Vec<Column> {
4924 vec![
4925 Column::new("id", "TEXT PRIMARY KEY"),
4926 Column::new("metadata", "JSON"),
4927 Column::new("title", "TEXT"),
4928 ]
4929 }
4930
4931 fn id(&self) -> String {
4932 self.id.clone()
4933 }
4934
4935 fn column_values(&self) -> Vec<(&'static str, Box<dyn ColumnValue>)> {
4936 vec![
4937 ("id", Box::new(self.id.clone())),
4938 (
4939 "metadata",
4940 Box::new(serde_json::json!({
4941 "user_id": self.metadata.user_id,
4942 "knowledge_id": self.metadata.knowledge_id,
4943 "knowledge_doc_id": self.metadata.knowledge_doc_id,
4944 })),
4945 ),
4946 ("title", Box::new(self.title.clone())),
4947 ]
4948 }
4949 }
4950
4951 #[derive(Clone, Debug, Deserialize, Serialize)]
4952 struct TypedTestDocument {
4953 id: i64,
4954 category: String,
4955 priority: i64,
4956 rating: f64,
4957 published: bool,
4958 title: String,
4959 }
4960
4961 impl SqliteVectorStoreTable for TypedTestDocument {
4962 fn name() -> &'static str {
4963 "live_typed_test_documents"
4964 }
4965
4966 fn schema() -> Vec<Column> {
4967 vec![
4968 Column::new("id", "INTEGER PRIMARY KEY"),
4969 Column::new("category", "TEXT").indexed(),
4970 Column::new("priority", "INTEGER").indexed(),
4971 Column::new("rating", "FLOAT").indexed(),
4972 Column::new("published", "BOOLEAN").indexed(),
4973 Column::new("title", "TEXT"),
4974 ]
4975 }
4976
4977 fn id(&self) -> String {
4978 self.id.to_string()
4979 }
4980
4981 fn column_values(&self) -> Vec<(&'static str, Box<dyn ColumnValue>)> {
4982 vec![
4983 ("id", Box::new(self.id)),
4984 ("category", Box::new(self.category.clone())),
4985 ("priority", Box::new(self.priority)),
4986 ("rating", Box::new(self.rating)),
4987 ("published", Box::new(self.published)),
4988 ("title", Box::new(self.title.clone())),
4989 ]
4990 }
4991 }
4992
4993 #[derive(Clone)]
4994 struct TestEmbeddingModel;
4995
4996 impl EmbeddingModel for TestEmbeddingModel {
4997 const MAX_DOCUMENTS: usize = 16;
4998
4999 type Client = ();
5000
5001 fn make(_: &Self::Client, _: impl Into<String>, _: Option<usize>) -> Self {
5002 Self
5003 }
5004
5005 fn ndims(&self) -> usize {
5006 2
5007 }
5008
5009 async fn embed_texts(
5010 &self,
5011 texts: impl IntoIterator<Item = String> + WasmCompatSend,
5012 ) -> Result<Vec<Embedding>, EmbeddingError> {
5013 Ok(texts
5014 .into_iter()
5015 .map(|text| Embedding {
5016 document: text,
5017 vec: vec![1.0, 0.0],
5018 })
5019 .collect())
5020 }
5021 }
5022}