#[path = "vectors/provenance.rs"]
mod provenance;
use provenance::{provenance_read_sql, provenance_sidecar_exists};
use std::collections::HashSet;
use std::sync::{Arc, OnceLock};
use async_trait::async_trait;
use chrono::{DateTime, Utc};
use rusqlite::OptionalExtension;
use uuid::Uuid;
use khive_score::{cmp_desc_then_id, try_cosine_score_with_f32_tolerance, DeterministicScore};
use khive_storage::error::StorageError;
use khive_storage::types::{
BatchWriteErrorClass, BatchWriteRetryability, BatchWriteSummary, IndexRebuildScope,
OrphanSweepConfig, OrphanSweepResult, SqlStatement, SqlValue, VectorIndexKind,
VectorProvenance, VectorRecord, VectorSearchHit, VectorSearchRequest, VectorStoreCapabilities,
VectorStoreInfo,
};
use khive_storage::VectorStore;
use khive_storage::{encode_f32_native, StorageResult};
use khive_storage::{ContentRef, StorageCapability};
use khive_types::SubstrateKind;
use crate::error::SqliteError;
use crate::pool::ConnectionPool;
use crate::sql_bridge::bind_params;
use crate::writer_task::execute_wrapped_transaction;
pub(crate) fn delete_vector_statement(
table: &str,
subject_id: Uuid,
namespace: &str,
) -> SqlStatement {
SqlStatement {
sql: format!("DELETE FROM {table} WHERE subject_id = ?1 AND namespace = ?2"),
params: vec![
SqlValue::Text(subject_id.to_string()),
SqlValue::Text(namespace.to_string()),
],
label: Some(format!("vec-delete-{table}")),
}
}
#[cfg(test)]
mod failpoint {
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::cell::RefCell;
thread_local! {
pub(super) static CURRENT: RefCell<Option<Arc<AtomicBool>>> = const { RefCell::new(None) };
}
#[cfg(feature = "vectors")]
pub(super) fn arm() {
let flag = Arc::new(AtomicBool::new(true));
CURRENT.with(|c| *c.borrow_mut() = Some(flag));
}
#[cfg(feature = "vectors")]
pub(super) fn disarm() {
CURRENT.with(|c| *c.borrow_mut() = None);
}
pub(super) fn take(flag: &Arc<AtomicBool>) -> bool {
flag.compare_exchange(true, false, Ordering::SeqCst, Ordering::SeqCst)
.is_ok()
}
#[cfg(feature = "vectors")]
pub(super) struct FailpointGuard;
#[cfg(feature = "vectors")]
impl FailpointGuard {
pub(super) fn new() -> Self {
arm();
Self
}
}
#[cfg(feature = "vectors")]
impl Drop for FailpointGuard {
fn drop(&mut self) {
disarm();
}
}
}
#[cfg(test)]
fn current_failpoint() -> Option<std::sync::Arc<std::sync::atomic::AtomicBool>> {
failpoint::CURRENT.with(|c| c.borrow().clone())
}
#[cfg(not(test))]
fn current_failpoint() -> Option<std::sync::Arc<std::sync::atomic::AtomicBool>> {
None
}
fn map_err(e: rusqlite::Error, op: &'static str) -> StorageError {
StorageError::driver(StorageCapability::Vectors, op, e)
}
fn map_sqlite_err(e: SqliteError, op: &'static str) -> StorageError {
e.into_storage_error(StorageCapability::Vectors, op)
}
fn non_finite_index(data: &[f32]) -> Option<usize> {
data.iter().position(|v| !v.is_finite())
}
fn non_finite_vector_error(op: &'static str, idx: usize, value: f32) -> StorageError {
StorageError::InvalidInput {
capability: StorageCapability::Vectors,
operation: op.into(),
message: format!(
"non-finite value at index {idx}: {value} \
(NaN/Inf values corrupt distance computations)"
),
}
}
fn sqlite_cosine_score(distance: f64) -> Result<DeterministicScore, rusqlite::Error> {
let conversion_error = |error| {
rusqlite::Error::FromSqlConversionFailure(1, rusqlite::types::Type::Real, Box::new(error))
};
try_cosine_score_with_f32_tolerance(distance).map_err(conversion_error)
}
#[cfg(test)]
mod sqlite_cosine_score_tests;
fn validate_model_key(model_key: &str) -> Result<(), SqliteError> {
if model_key.is_empty()
|| !model_key
.chars()
.all(|c| c.is_ascii_alphanumeric() || c == '_')
{
return Err(SqliteError::InvalidData(format!(
"invalid model_key '{}': must be non-empty and contain only ASCII alphanumeric/underscore characters",
model_key
)));
}
Ok(())
}
pub struct SqliteVecStore {
pool: Arc<ConnectionPool>,
model_key: String,
embedding_model: String,
dimensions: usize,
table_name: String,
namespace: String,
writer_task: Option<crate::writer_task::WriterTaskHandle>,
}
impl SqliteVecStore {
pub fn new(
pool: Arc<ConnectionPool>,
_is_file_backed: bool,
model_key: String,
embedding_model: String,
dimensions: usize,
namespace: String,
) -> Result<Self, SqliteError> {
validate_model_key(&model_key)?;
let table_name = format!("vec_{}", model_key);
let writer_task = pool.writer_task_handle().ok().flatten();
Ok(Self {
pool,
model_key,
embedding_model,
dimensions,
table_name,
namespace,
writer_task,
})
}
fn current_writer_task(
&self,
operation: &'static str,
) -> Result<Option<crate::writer_task::WriterTaskHandle>, StorageError> {
self.pool
.writer_task_for_write(self.writer_task.as_ref(), operation)
}
async fn with_writer<F, R>(&self, op: &'static str, f: F) -> Result<R, StorageError>
where
F: FnOnce(&rusqlite::Connection) -> Result<R, rusqlite::Error> + Send + 'static,
R: Send + 'static,
{
if let Some(writer_task) = self.current_writer_task(op)? {
return writer_task
.send_bounded(move |conn| f(conn).map_err(|e| map_err(e, op)))
.await;
}
self.pool
.record_direct_route(crate::timeout_sink::Site::DirectRouteVecGeneralWrite);
self.with_writer_unmanaged(op, f).await
}
async fn with_writer_unmanaged<F, R>(&self, op: &'static str, f: F) -> Result<R, StorageError>
where
F: FnOnce(&rusqlite::Connection) -> Result<R, rusqlite::Error> + Send + 'static,
R: Send + 'static,
{
let pool = Arc::clone(&self.pool);
tokio::task::spawn_blocking(move || {
let guard = pool
.transaction_write_unit()
.map_err(|e| map_sqlite_err(e, op))
.inspect_err(|error| pool.record_direct_writer_error(error))?;
let conn = guard.conn();
let _tx_handle = khive_storage::tx_registry::register_scoped(
Some(format!("{op}_tx")),
pool.origin(),
);
let db_label = crate::timeout_sink::db_label(&pool);
let (result, terminal_state) = execute_wrapped_transaction(conn, op, move |conn| {
f(conn).map_err(|e| {
crate::timeout_sink::maybe_emit_sqlite_full(&db_label, &e);
map_err(e, op)
})
});
if terminal_state.is_some() {
pool.retire_pooled_writer(conn);
}
result.inspect_err(|error| pool.record_direct_writer_error(error))
})
.await
.map_err(|e| StorageError::driver(StorageCapability::Vectors, op, e))?
}
async fn with_reader<F, R>(&self, op: &'static str, f: F) -> Result<R, StorageError>
where
F: FnOnce(&rusqlite::Connection) -> Result<R, rusqlite::Error> + Send + 'static,
R: Send + 'static,
{
super::run_pooled_store_read(
Arc::clone(&self.pool),
StorageCapability::Vectors,
op,
move |conn| f(conn).map_err(|error| map_err(error, op)),
)
.await
}
#[allow(clippy::too_many_arguments)]
async fn insert_one(
&self,
subject_id: Uuid,
kind: SubstrateKind,
namespace: &str,
field: &str,
vectors: Vec<Vec<f32>>,
record_ann_delta: bool,
operation: &'static str,
savepoint_name: &'static str,
) -> Result<(), StorageError> {
if vectors.len() != 1 {
return Err(StorageError::Unsupported {
capability: StorageCapability::Vectors,
operation: operation.into(),
message: "sqlite-vec supports exactly one vector per record".into(),
});
}
let embedding = vectors.into_iter().next().expect("len checked");
let table = self.table_name.clone();
let dims = self.dimensions;
let namespace = namespace.to_string();
let field = field.to_string();
let kind_str = kind.to_string();
let embedding_model = self.embedding_model.clone();
if embedding.len() == dims {
if let Some(index) = non_finite_index(&embedding) {
return Err(non_finite_vector_error(operation, index, embedding[index]));
}
}
let failpoint_flag = current_failpoint();
if let Some(writer_task) = self.current_writer_task(operation)? {
let table_for_write = table.clone();
let namespace_for_write = namespace.clone();
let field_for_write = field.clone();
let kind_for_write = kind_str.clone();
let model_for_write = embedding_model.clone();
let embedding_for_write = embedding.clone();
return writer_task
.send_bounded(move |connection| {
vec_upsert_atomic_dml(
connection,
&table_for_write,
dims,
subject_id,
&kind_for_write,
&namespace_for_write,
&field_for_write,
&model_for_write,
&embedding_for_write,
savepoint_name,
record_ann_delta,
failpoint_flag,
)
.map_err(|error| map_err(error, operation))
})
.await;
}
self.with_writer(operation, move |connection| {
replace_vector_row_dml(
connection,
&table,
dims,
VectorRowRef {
subject_id,
namespace: &namespace,
kind: &kind_str,
field: &field,
embedding_model: &embedding_model,
embedding: &embedding,
text_fingerprint: None,
updated_at: None,
},
record_ann_delta,
failpoint_flag,
)
})
.await
}
}
mod dml;
use super::classify_batch_sqlite_error;
pub use dml::delete_subject_from_vector_tables;
use dml::{
batch_insert_vectors_dml, delete_vector_provenance, delete_vector_subjects_dml,
log_vector_deletes, orphan_sweep_dml, replace_vector_row_dml, vec_upsert_atomic_dml,
VectorRowRef,
};
#[cfg(all(test, feature = "vectors"))]
#[path = "orphan_sweep_dml_tests.rs"]
mod orphan_sweep_dml_tests;
#[cfg(all(test, feature = "vectors"))]
#[path = "vector_read_tests.rs"]
mod vector_read_tests;
mod vector_store_impl;
impl SqliteVecStore {
pub async fn score_candidates(
&self,
query_embedding: &[f32],
candidate_ids: &[Uuid],
) -> Result<Vec<VectorSearchHit>, StorageError> {
let dims = self.dimensions;
if query_embedding.len() != dims {
return Err(StorageError::InvalidInput {
capability: StorageCapability::Vectors,
operation: "score_candidates".into(),
message: format!(
"query has {} dims, expected {}",
query_embedding.len(),
dims
),
});
}
if candidate_ids.is_empty() {
return Ok(Vec::new());
}
if let Some(idx) = non_finite_index(query_embedding) {
return Err(non_finite_vector_error(
"score_candidates",
idx,
query_embedding[idx],
));
}
let table = self.table_name.clone();
let namespace = self.namespace.clone();
let embedding_model = self.embedding_model.clone();
let query_vec = query_embedding.to_vec();
let ids: Vec<String> = candidate_ids.iter().map(|id| id.to_string()).collect();
self.with_reader("score_candidates", move |conn| {
let mut all_hits: Vec<VectorSearchHit> = Vec::new();
let query_blob = encode_f32_native(&query_vec);
let sql = format!(
"SELECT e.subject_id, vec_distance_cosine(e.embedding, ?1) as distance \
FROM {table} e \
WHERE e.namespace = ?2 AND e.embedding_model = ?3 \
AND e.subject_id = ?4"
);
let mut stmt = conn.prepare(&sql)?;
for chunk in ids.chunks(399) {
let mut seen = HashSet::with_capacity(chunk.len());
for id in chunk.iter().filter(|id| seen.insert(*id)) {
let row: Option<(String, f64)> = stmt
.query_row(
rusqlite::params![query_blob, &namespace, &embedding_model, id],
|row| Ok((row.get(0)?, row.get(1)?)),
)
.optional()?;
let Some((id_str, distance)) = row else {
continue;
};
let subject_id = Uuid::parse_str(&id_str).map_err(|e| {
rusqlite::Error::FromSqlConversionFailure(
0,
rusqlite::types::Type::Text,
Box::new(e),
)
})?;
all_hits.push(VectorSearchHit {
subject_id,
score: sqlite_cosine_score(distance)?,
rank: 0,
});
}
}
all_hits
.sort_by(|a, b| cmp_desc_then_id(a.score, &a.subject_id, b.score, &b.subject_id));
for (i, hit) in all_hits.iter_mut().enumerate() {
hit.rank = (i + 1) as u32;
}
Ok(all_hits)
})
.await
}
}
#[cfg(all(test, feature = "vectors"))]
mod point_lookup_tests;
#[cfg(test)]
mod unmanaged_write_escalation_tests;
#[cfg(all(test, feature = "vectors"))]
mod batch_exists_tests;
#[cfg(test)]
mod first_error_tests;
#[cfg(test)]
mod capabilities_tests;
#[cfg(all(test, feature = "vectors"))]
mod delete_subjects_atomic_tests;
#[cfg(all(test, feature = "vectors"))]
mod atomic_replace_tests;
#[cfg(all(test, feature = "vectors"))]
mod orphan_sweep_tests;
#[cfg(all(test, feature = "vectors"))]
mod write_queue_tests;
#[cfg(all(test, feature = "vectors"))]
mod provenance_tests;
#[cfg(test)]
#[path = "vectors_busy_tests.rs"]
mod direct_busy_tests;