#[cfg(test)]
use super::failpoint;
use super::{
f32_slice_as_bytes, non_finite_index, provenance_sidecar_exists, BatchWriteErrorClass,
BatchWriteRetryability, BatchWriteSummary, ContentRef, DateTime, OrphanSweepResult, Utc, Uuid,
VectorRecord,
};
pub(super) struct VectorRowRef<'a> {
pub(super) subject_id: Uuid,
pub(super) namespace: &'a str,
pub(super) kind: &'a str,
pub(super) field: &'a str,
pub(super) embedding_model: &'a str,
pub(super) embedding: &'a [f32],
pub(super) text_fingerprint: Option<&'a ContentRef>,
pub(super) updated_at: Option<&'a DateTime<Utc>>,
}
pub(super) fn replace_vector_row_dml(
conn: &rusqlite::Connection,
table: &str,
dims: usize,
row: VectorRowRef<'_>,
record_ann_delta: bool,
failpoint_flag: Option<std::sync::Arc<std::sync::atomic::AtomicBool>>,
) -> Result<(), rusqlite::Error> {
if row.embedding.len() != dims {
return Err(rusqlite::Error::InvalidParameterCount(
row.embedding.len(),
dims,
));
}
let subject_id = row.subject_id.to_string();
let delete_same_identity_sql = format!(
"DELETE FROM {table} WHERE subject_id = ?1 AND namespace = ?2 \
AND embedding_model = ?3 AND kind = ?4 AND field = ?5"
);
let deleted_same_identity = conn.execute(
&delete_same_identity_sql,
rusqlite::params![
&subject_id,
row.namespace,
row.embedding_model,
row.kind,
row.field
],
)?;
if deleted_same_identity == 0 {
if record_ann_delta {
let logged = log_vector_deletes(conn, table, "subject_id = ?1", &[&subject_id])?;
if logged > 0 {
let delete_prior_identity_sql =
format!("DELETE FROM {table} WHERE subject_id = ?1");
conn.execute(&delete_prior_identity_sql, rusqlite::params![&subject_id])?;
}
} else {
let delete_prior_identity_sql = format!("DELETE FROM {table} WHERE subject_id = ?1");
conn.execute(&delete_prior_identity_sql, rusqlite::params![&subject_id])?;
}
}
#[cfg(test)]
if let Some(ref fp) = failpoint_flag {
if failpoint::take(fp) {
return Err(rusqlite::Error::InvalidParameterName(
"__test_failpoint_after_delete__".into(),
));
}
}
#[cfg(not(test))]
let _ = failpoint_flag;
let ins_sql = format!(
"INSERT INTO {table} (subject_id, namespace, kind, field, embedding_model, embedding) \
VALUES (?1, ?2, ?3, ?4, ?5, ?6)"
);
let blob = f32_slice_as_bytes(row.embedding);
conn.execute(
&ins_sql,
rusqlite::params![
&subject_id,
row.namespace,
row.kind,
row.field,
row.embedding_model,
blob
],
)?;
if provenance_sidecar_exists(conn)? {
let model_key = table
.strip_prefix("vec_")
.expect("vector table names use the vec_ prefix");
let stored_embedding: Vec<u8> = conn.query_row(
&format!("SELECT embedding FROM {table} WHERE subject_id = ?1 AND namespace = ?2"),
rusqlite::params![&subject_id, row.namespace],
|stored| stored.get(0),
)?;
let embedding_digest = blake3::hash(&stored_embedding).to_hex().to_string();
let updated_at = row.updated_at.map(DateTime::to_rfc3339);
conn.execute(
"INSERT INTO vector_provenance \
(model_key, subject_id, namespace, embedding_digest, text_fingerprint, updated_at) \
VALUES (?1, ?2, ?3, ?4, ?5, ?6) \
ON CONFLICT(model_key, subject_id) DO UPDATE SET \
namespace = excluded.namespace, \
embedding_digest = excluded.embedding_digest, \
text_fingerprint = excluded.text_fingerprint, \
updated_at = excluded.updated_at",
rusqlite::params![
model_key,
&subject_id,
row.namespace,
embedding_digest,
row.text_fingerprint.map(ContentRef::as_str),
updated_at,
],
)?;
}
if record_ann_delta {
conn.execute(
"INSERT INTO ann_write_log (namespace, embedding_model, kind, field, subject_id, op) \
VALUES (?1, ?2, ?3, ?4, ?5, 'upsert')",
rusqlite::params![
row.namespace,
row.embedding_model,
row.kind,
row.field,
&subject_id
],
)?;
}
Ok(())
}
pub(super) fn log_vector_deletes(
conn: &rusqlite::Connection,
table: &str,
where_clause: &str,
params: &[&dyn rusqlite::ToSql],
) -> Result<usize, rusqlite::Error> {
let sql = format!(
"INSERT INTO ann_write_log (namespace, embedding_model, kind, field, subject_id, op) \
SELECT namespace, embedding_model, kind, field, subject_id, 'delete' \
FROM {table} WHERE {where_clause}"
);
conn.execute(&sql, params)
}
pub(super) fn delete_vector_provenance(
conn: &rusqlite::Connection,
table: &str,
subject_ids: &[String],
) -> Result<(), rusqlite::Error> {
if subject_ids.is_empty() {
return Ok(());
}
if !provenance_sidecar_exists(conn)? {
return Ok(());
}
let model_key = table
.strip_prefix("vec_")
.expect("vector table names use the vec_ prefix");
let placeholders = (2..=subject_ids.len() + 1)
.map(|i| format!("?{i}"))
.collect::<Vec<_>>()
.join(", ");
let sql = format!(
"DELETE FROM vector_provenance \
WHERE model_key = ?1 AND subject_id IN ({placeholders})"
);
let mut statement = conn.prepare(&sql)?;
statement.raw_bind_parameter(1, model_key)?;
for (index, subject_id) in subject_ids.iter().enumerate() {
statement.raw_bind_parameter(index + 2, subject_id.as_str())?;
}
statement.raw_execute()?;
Ok(())
}
pub(super) fn delete_vector_subjects_dml(
conn: &rusqlite::Connection,
table: &str,
id_strings: &[String],
) -> Result<u64, rusqlite::Error> {
let mut total_deleted = 0u64;
let log_sql = format!(
"INSERT INTO ann_write_log (namespace, embedding_model, kind, field, subject_id, op) \
SELECT namespace, embedding_model, kind, field, subject_id, 'delete' \
FROM {table} WHERE subject_id = ?1"
);
let mut log_stmt = conn.prepare(&log_sql)?;
let mut delete_stmt = conn.prepare(&format!("DELETE FROM {table} WHERE subject_id = ?1"))?;
for chunk in id_strings.chunks(400) {
for id in chunk {
log_stmt.execute([id.as_str()])?;
total_deleted += delete_stmt.execute([id.as_str()])? as u64;
}
delete_vector_provenance(conn, table, chunk)?;
}
Ok(total_deleted)
}
pub fn delete_subject_from_vector_tables(
conn: &rusqlite::Connection,
tables: &[String],
subject_id: Uuid,
namespace: &str,
) -> Result<(), rusqlite::Error> {
let subject_id = subject_id.to_string();
for table in tables {
log_vector_deletes(
conn,
table,
"subject_id = ?1 AND namespace = ?2",
&[&subject_id, &namespace],
)?;
let sql = format!("DELETE FROM {table} WHERE subject_id = ?1 AND namespace = ?2");
if conn.execute(&sql, rusqlite::params![&subject_id, namespace])? > 0 {
delete_vector_provenance(conn, table, std::slice::from_ref(&subject_id))?;
}
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub(super) fn batch_insert_vectors_dml(
conn: &rusqlite::Connection,
table: &str,
dims: usize,
store_embedding_model: &str,
records: &[VectorRecord],
attempted: u64,
failpoint_flag: Option<std::sync::Arc<std::sync::atomic::AtomicBool>>,
) -> Result<BatchWriteSummary, rusqlite::Error> {
let mut summary = BatchWriteSummary {
attempted,
..BatchWriteSummary::default()
};
for (index, record) in records.iter().enumerate() {
let item_id = Some(record.subject_id.to_string());
if record.vectors.len() != 1 {
summary.record_failure(
index,
item_id,
BatchWriteErrorClass::InvalidInput,
BatchWriteRetryability::Permanent,
format!("expected 1 vector per record, got {}", record.vectors.len()),
);
continue;
}
let embedding = &record.vectors[0];
if embedding.len() != dims {
summary.record_failure(
index,
item_id,
BatchWriteErrorClass::InvalidInput,
BatchWriteRetryability::Permanent,
format!(
"wrong vector dimension: expected {dims}, got {}",
embedding.len()
),
);
continue;
}
if non_finite_index(embedding).is_some() {
summary.record_failure(
index,
item_id,
BatchWriteErrorClass::InvalidInput,
BatchWriteRetryability::Permanent,
"embedding contains non-finite values (NaN or Inf)",
);
continue;
}
let kind_str = record.kind.to_string();
conn.execute_batch("SAVEPOINT vec_batch_record")?;
let result = replace_vector_row_dml(
conn,
table,
dims,
VectorRowRef {
subject_id: record.subject_id,
namespace: &record.namespace,
kind: &kind_str,
field: &record.field,
embedding_model: store_embedding_model,
embedding,
text_fingerprint: record.text_fingerprint.as_ref(),
updated_at: Some(&record.updated_at),
},
true,
failpoint_flag.clone(),
);
match result {
Ok(()) => {
conn.execute_batch("RELEASE SAVEPOINT vec_batch_record")?;
summary.affected = summary.affected.saturating_add(1);
}
Err(e) => {
let _ = conn.execute_batch("ROLLBACK TO SAVEPOINT vec_batch_record");
let _ = conn.execute_batch("RELEASE SAVEPOINT vec_batch_record");
let (class, retryability) = super::classify_batch_sqlite_error(&e);
summary.record_failure(index, item_id, class, retryability, e.to_string());
}
}
}
Ok(summary)
}
#[allow(clippy::too_many_arguments)]
pub(super) fn vec_upsert_atomic_dml(
conn: &rusqlite::Connection,
table: &str,
dims: usize,
subject_id: Uuid,
kind_str: &str,
namespace: &str,
field: &str,
embedding_model: &str,
embedding: &[f32],
savepoint_name: &'static str,
record_ann_delta: bool,
failpoint_flag: Option<std::sync::Arc<std::sync::atomic::AtomicBool>>,
) -> Result<(), rusqlite::Error> {
conn.execute_batch(&format!("SAVEPOINT {savepoint_name}"))?;
let result = replace_vector_row_dml(
conn,
table,
dims,
VectorRowRef {
subject_id,
namespace,
kind: kind_str,
field,
embedding_model,
embedding,
text_fingerprint: None,
updated_at: None,
},
record_ann_delta,
failpoint_flag,
);
match result {
Ok(()) => {
conn.execute_batch(&format!("RELEASE SAVEPOINT {savepoint_name}"))?;
Ok(())
}
Err(e) => {
let _ = conn.execute_batch(&format!("ROLLBACK TO SAVEPOINT {savepoint_name}"));
let _ = conn.execute_batch(&format!("RELEASE SAVEPOINT {savepoint_name}"));
Err(e)
}
}
}
pub(super) fn orphan_sweep_dml(
conn: &rusqlite::Connection,
table: &str,
ns_json: Option<&str>,
kind_json: Option<&str>,
allow_json: Option<&str>,
max_delete: i64,
dry_run: bool,
) -> Result<OrphanSweepResult, rusqlite::Error> {
let filter_pred = "(?1 IS NULL OR namespace IN (SELECT value FROM json_each(?1))) \
AND (?2 IS NULL OR kind IN (SELECT value FROM json_each(?2))) \
AND (?3 IS NULL OR subject_id IN (SELECT value FROM json_each(?3)))";
let live_subq = "SELECT id FROM entities WHERE deleted_at IS NULL \
UNION ALL \
SELECT id FROM notes WHERE deleted_at IS NULL \
UNION ALL \
SELECT id FROM knowledge_atoms WHERE deleted_at IS NULL";
let scan_sql = format!(
"SELECT COUNT(*) FROM {t} WHERE {f}",
t = table,
f = filter_pred
);
let scanned: i64 = conn.query_row(
&scan_sql,
rusqlite::params![ns_json, kind_json, allow_json],
|row| row.get(0),
)?;
conn.execute_batch("CREATE TEMP TABLE khive_orphan_sweep_live_ids(id TEXT)")?;
conn.execute_batch(&format!(
"INSERT INTO temp.khive_orphan_sweep_live_ids(id) {live_subq}"
))?;
conn.execute_batch(
"CREATE INDEX temp.khive_orphan_sweep_live_ids_idx \
ON khive_orphan_sweep_live_ids(id)",
)?;
let orphan_pred = format!(
"subject_id NOT IN (SELECT id FROM temp.khive_orphan_sweep_live_ids) AND {filter_pred}"
);
let count_sql = format!(
"SELECT COUNT(*) FROM {t} WHERE {p}",
t = table,
p = orphan_pred,
);
let would_delete: i64 = conn.query_row(
&count_sql,
rusqlite::params![ns_json, kind_json, allow_json],
|row| row.get(0),
)?;
let max_delete_hit = would_delete > max_delete;
let deleted: i64 = if dry_run {
0
} else {
let select_sql = format!(
"SELECT subject_id FROM {t} WHERE {p} LIMIT ?4",
t = table,
p = orphan_pred,
);
let mut select_stmt = conn.prepare(&select_sql)?;
let log_sql = format!(
"INSERT INTO ann_write_log (namespace, embedding_model, kind, field, subject_id, op) \
SELECT namespace, embedding_model, kind, field, subject_id, 'delete' \
FROM {table} WHERE subject_id = ?1"
);
let mut log_stmt = conn.prepare(&log_sql)?;
let del_sql = format!("DELETE FROM {table} WHERE subject_id = ?1");
let mut delete_stmt = conn.prepare(&del_sql)?;
let mut total: i64 = 0;
let mut remaining = max_delete;
while remaining > 0 {
let batch_limit = remaining.min(400);
let victim_ids: Vec<String> = select_stmt
.query_map(
rusqlite::params![ns_json, kind_json, allow_json, batch_limit],
|row| row.get::<_, String>(0),
)?
.collect::<Result<_, _>>()?;
if victim_ids.is_empty() {
break;
}
for id in &victim_ids {
log_stmt.execute([id.as_str()])?;
total += delete_stmt.execute([id.as_str()])? as i64;
}
delete_vector_provenance(conn, table, &victim_ids)?;
remaining -= victim_ids.len() as i64;
}
total
};
conn.execute_batch("DROP TABLE temp.khive_orphan_sweep_live_ids")?;
Ok(OrphanSweepResult {
scanned: scanned as u64,
would_delete: would_delete as u64,
deleted: deleted as u64,
max_delete_hit,
})
}