1use std::cmp::Reverse;
4use std::collections::BinaryHeap;
5use std::sync::Arc;
6
7use async_trait::async_trait;
8use uuid::Uuid;
9
10use khive_score::DeterministicScore;
11use khive_storage::error::StorageError;
12use khive_storage::types::{
13 BatchWriteErrorClass, BatchWriteRetryability, BatchWriteSummary, SparseRecord, SparseSearchHit,
14 SparseSearchRequest, SparseVector,
15};
16use khive_storage::{SparseStore, StorageCapability};
17use khive_types::SubstrateKind;
18
19use crate::error::SqliteError;
20use crate::pool::ConnectionPool;
21use crate::writer_task::WriterTaskHandle;
22
23fn map_err(e: rusqlite::Error, op: &'static str) -> StorageError {
24 StorageError::driver(StorageCapability::Sparse, op, e)
25}
26
27fn map_sqlite_err(e: SqliteError, op: &'static str) -> StorageError {
28 StorageError::driver(StorageCapability::Sparse, op, e)
29}
30
31fn validate_sparse_vector(vector: &SparseVector, op: &'static str) -> Result<(), StorageError> {
38 if vector.indices.len() != vector.values.len() {
39 return Err(StorageError::InvalidInput {
40 capability: StorageCapability::Sparse,
41 operation: op.into(),
42 message: format!(
43 "indices length ({}) != values length ({})",
44 vector.indices.len(),
45 vector.values.len()
46 ),
47 });
48 }
49 if vector.indices.is_empty() {
50 return Err(StorageError::InvalidInput {
51 capability: StorageCapability::Sparse,
52 operation: op.into(),
53 message: "sparse vector must have at least one element".into(),
54 });
55 }
56 for (i, v) in vector.values.iter().enumerate() {
57 if !v.is_finite() {
58 return Err(StorageError::InvalidInput {
59 capability: StorageCapability::Sparse,
60 operation: op.into(),
61 message: format!("non-finite value at position {i}: {v}"),
62 });
63 }
64 }
65 for window in vector.indices.windows(2) {
67 if window[0] >= window[1] {
68 return Err(StorageError::InvalidInput {
69 capability: StorageCapability::Sparse,
70 operation: op.into(),
71 message: format!(
72 "indices must be strictly increasing; found {} then {}",
73 window[0], window[1]
74 ),
75 });
76 }
77 }
78 Ok(())
79}
80
81fn f32_slice_as_bytes(data: &[f32]) -> &[u8] {
83 unsafe { std::slice::from_raw_parts(data.as_ptr() as *const u8, std::mem::size_of_val(data)) }
85}
86
87fn batch_insert_sparse_dml(
96 conn: &rusqlite::Connection,
97 table: &str,
98 records: &[SparseRecord],
99 attempted: u64,
100) -> Result<BatchWriteSummary, rusqlite::Error> {
101 let sql = format!(
102 "INSERT INTO {table} \
103 (subject_id, namespace, kind, field, indices_json, values_blob, updated_at) \
104 VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7) \
105 ON CONFLICT(subject_id, namespace, field) DO UPDATE SET \
106 indices_json = excluded.indices_json, \
107 values_blob = excluded.values_blob, \
108 updated_at = excluded.updated_at"
109 );
110
111 let mut summary = BatchWriteSummary {
112 attempted,
113 ..BatchWriteSummary::default()
114 };
115
116 for (index, record) in records.iter().enumerate() {
117 let item_id = Some(record.subject_id.to_string());
118 if record.vector.indices.len() != record.vector.values.len()
120 || record.vector.indices.is_empty()
121 || record.vector.values.iter().any(|v| !v.is_finite())
122 || record.vector.indices.windows(2).any(|w| w[0] >= w[1])
123 {
124 summary.record_failure(
125 index,
126 item_id,
127 BatchWriteErrorClass::InvalidInput,
128 BatchWriteRetryability::Permanent,
129 format!("invalid sparse vector for subject {}", record.subject_id),
130 );
131 continue;
132 }
133
134 let indices_json = match serde_json::to_string(&record.vector.indices) {
135 Ok(j) => j,
136 Err(e) => {
137 summary.record_failure(
138 index,
139 item_id,
140 BatchWriteErrorClass::Serialization,
141 BatchWriteRetryability::Permanent,
142 e.to_string(),
143 );
144 continue;
145 }
146 };
147 let values_blob = f32_slice_as_bytes(&record.vector.values);
148 let now = record.updated_at.timestamp();
149 let id_str = record.subject_id.to_string();
150 let kind_str = record.kind.to_string();
151
152 match conn.execute(
153 &sql,
154 rusqlite::params![
155 &id_str,
156 &record.namespace,
157 &kind_str,
158 &record.field,
159 &indices_json,
160 values_blob,
161 now
162 ],
163 ) {
164 Ok(_) => summary.affected = summary.affected.saturating_add(1),
165 Err(e) => {
166 let (class, retryability) = super::classify_batch_sqlite_error(&e);
167 summary.record_failure(index, item_id, class, retryability, e.to_string());
168 }
169 }
170 }
171
172 Ok(summary)
173}
174
175pub(crate) fn ensure_sparse_schema(
177 conn: &rusqlite::Connection,
178 model_key: &str,
179) -> Result<(), rusqlite::Error> {
180 let table = format!("sparse_{}", model_key);
181 let ddl = format!(
182 "CREATE TABLE IF NOT EXISTS {table} (\
183 subject_id TEXT NOT NULL, \
184 namespace TEXT NOT NULL, \
185 kind TEXT NOT NULL, \
186 field TEXT NOT NULL, \
187 indices_json TEXT NOT NULL, \
188 values_blob BLOB NOT NULL, \
189 updated_at INTEGER NOT NULL, \
190 PRIMARY KEY(subject_id, namespace, field)\
191 ); \
192 CREATE INDEX IF NOT EXISTS idx_{table}_namespace_kind \
193 ON {table}(namespace, kind);"
194 );
195 conn.execute_batch(&ddl)
196}
197
198pub struct SqliteSparseStore {
200 pool: Arc<ConnectionPool>,
201 table_name: String,
202 namespace: String,
203 writer_task: Option<WriterTaskHandle>,
204}
205
206impl SqliteSparseStore {
207 pub fn new(
209 pool: Arc<ConnectionPool>,
210 _is_file_backed: bool,
211 model_key: String,
212 namespace: String,
213 ) -> Result<Self, SqliteError> {
214 let table_name = format!("sparse_{}", model_key);
215 let writer_task = pool.writer_task_handle().ok().flatten();
220 Ok(Self {
221 pool,
222 table_name,
223 namespace,
224 writer_task,
225 })
226 }
227
228 fn current_writer_task(
229 &self,
230 operation: &'static str,
231 ) -> Result<Option<WriterTaskHandle>, StorageError> {
232 self.pool
233 .writer_task_for_write(self.writer_task.as_ref(), operation)
234 }
235
236 async fn with_writer<F, R>(&self, op: &'static str, f: F) -> Result<R, StorageError>
250 where
251 F: FnOnce(&rusqlite::Connection) -> Result<R, rusqlite::Error> + Send + 'static,
252 R: Send + 'static,
253 {
254 if let Some(writer_task) = self.current_writer_task(op)? {
255 return writer_task
256 .send_bounded(move |conn| f(conn).map_err(|e| map_err(e, op)))
257 .await;
258 }
259
260 self.pool
261 .record_direct_route(crate::timeout_sink::Site::DirectRouteSparseGeneralWrite);
262 let pool = Arc::clone(&self.pool);
263 tokio::task::spawn_blocking(move || {
264 let guard = pool.try_writer().map_err(|e| map_sqlite_err(e, op))?;
265 f(guard.conn()).map_err(|e| map_err(e, op))
266 })
267 .await
268 .map_err(|e| StorageError::driver(StorageCapability::Sparse, op, e))?
269 }
270
271 async fn with_reader<F, R>(&self, op: &'static str, f: F) -> Result<R, StorageError>
272 where
273 F: FnOnce(&rusqlite::Connection) -> Result<R, rusqlite::Error> + Send + 'static,
274 R: Send + 'static,
275 {
276 super::run_pooled_store_read(
277 Arc::clone(&self.pool),
278 StorageCapability::Sparse,
279 op,
280 move |conn| f(conn).map_err(|error| map_err(error, op)),
281 )
282 .await
283 }
284
285 async fn upsert_sparse_vector(
286 &self,
287 subject_id: Uuid,
288 kind: SubstrateKind,
289 namespace: &str,
290 field: &str,
291 vector: SparseVector,
292 ) -> Result<(), StorageError> {
293 let table = self.table_name.clone();
294 let ns = namespace.to_string();
295 let field = field.to_string();
296 let id_str = subject_id.to_string();
297 let kind_str = kind.to_string();
298
299 self.with_writer("sparse_upsert", move |conn| {
300 let indices_json = serde_json::to_string(&vector.indices).map_err(|e| {
301 rusqlite::Error::FromSqlConversionFailure(
302 0,
303 rusqlite::types::Type::Text,
304 Box::new(e),
305 )
306 })?;
307 let values_blob = f32_slice_as_bytes(&vector.values);
308 let now = chrono::Utc::now().timestamp();
309 let sql = format!(
310 "INSERT INTO {table} \
311 (subject_id, namespace, kind, field, indices_json, values_blob, updated_at) \
312 VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7) \
313 ON CONFLICT(subject_id, namespace, field) DO UPDATE SET \
314 kind = excluded.kind, \
315 indices_json = excluded.indices_json, \
316 values_blob = excluded.values_blob, \
317 updated_at = excluded.updated_at"
318 );
319 conn.execute(
320 &sql,
321 rusqlite::params![
322 &id_str,
323 &ns,
324 &kind_str,
325 &field,
326 &indices_json,
327 values_blob,
328 now
329 ],
330 )?;
331 Ok(())
332 })
333 .await
334 }
335
336 async fn insert_sparse_batch(
337 &self,
338 records: Vec<SparseRecord>,
339 ) -> Result<BatchWriteSummary, StorageError> {
340 let table = self.table_name.clone();
341 let attempted = records.len() as u64;
342
343 if let Some(writer_task) = self.current_writer_task("sparse_insert_batch")? {
348 let table2 = table.clone();
349 return writer_task
350 .send_bounded(move |conn| {
351 batch_insert_sparse_dml(conn, &table2, &records, attempted)
352 .map_err(|e| map_err(e, "sparse_insert_batch"))
353 })
354 .await;
355 }
356
357 let origin = self.pool.origin();
360 self.with_writer("sparse_insert_batch", move |conn| {
361 conn.execute_batch("BEGIN IMMEDIATE")?;
362 let _tx_handle = khive_storage::tx_registry::register_scoped(
363 Some("sparse_insert_batch".to_string()),
364 origin,
365 );
366
367 let summary = batch_insert_sparse_dml(conn, &table, &records, attempted)?;
368
369 conn.execute_batch("COMMIT")?;
370 Ok(summary)
371 })
372 .await
373 }
374
375 async fn delete_sparse_subject(&self, subject_id: Uuid) -> Result<bool, StorageError> {
376 let table = self.table_name.clone();
377 let namespace = self.namespace.clone();
378 let id_str = subject_id.to_string();
379
380 self.with_writer("sparse_delete", move |conn| {
381 let sql = format!("DELETE FROM {table} WHERE subject_id = ?1 AND namespace = ?2");
382 let deleted = conn.execute(&sql, rusqlite::params![&id_str, &namespace])?;
383 Ok(deleted > 0)
384 })
385 .await
386 }
387
388 async fn search_sparse_vectors(
389 &self,
390 request: SparseSearchRequest,
391 ) -> Result<Vec<SparseSearchHit>, StorageError> {
392 request
393 .validate()
394 .map_err(|message| StorageError::InvalidInput {
395 capability: StorageCapability::Sparse,
396 operation: "sparse_search".into(),
397 message,
398 })?;
399
400 let table = self.table_name.clone();
401 let ns = request
402 .namespace
403 .clone()
404 .unwrap_or_else(|| self.namespace.clone());
405 let kind_filter = request.kind.map(|k| k.to_string());
406 let query = request.query;
407 let top_k = usize::try_from(request.top_k).map_err(|_| StorageError::InvalidInput {
408 capability: StorageCapability::Sparse,
409 operation: "sparse_search".into(),
410 message: "SparseSearchRequest: top_k does not fit usize".into(),
411 })?;
412 let heap_capacity = top_k
413 .checked_add(1)
414 .ok_or_else(|| StorageError::InvalidInput {
415 capability: StorageCapability::Sparse,
416 operation: "sparse_search".into(),
417 message: "SparseSearchRequest: top_k capacity overflow".into(),
418 })?;
419
420 self.with_reader("sparse_search", move |conn| {
421 let (sql, kind_str_ref) = if let Some(ref kind_str) = kind_filter {
423 (
424 format!(
425 "SELECT subject_id, indices_json, values_blob \
426 FROM {table} WHERE namespace = ?1 AND kind = ?2"
427 ),
428 Some(kind_str.as_str()),
429 )
430 } else {
431 (
432 format!(
433 "SELECT subject_id, indices_json, values_blob \
434 FROM {table} WHERE namespace = ?1"
435 ),
436 None,
437 )
438 };
439
440 let mut stmt = conn.prepare(&sql)?;
441
442 let rows: Vec<rusqlite::Result<(String, String, Vec<u8>)>> =
444 if let Some(kind_str) = kind_str_ref {
445 stmt.query_map(rusqlite::params![&ns, kind_str], |row| {
446 Ok((row.get(0)?, row.get(1)?, row.get(2)?))
447 })?
448 .collect()
449 } else {
450 stmt.query_map(rusqlite::params![&ns], |row| {
451 Ok((row.get(0)?, row.get(1)?, row.get(2)?))
452 })?
453 .collect()
454 };
455
456 let mut heap: BinaryHeap<Reverse<ScoredCandidate>> =
458 BinaryHeap::with_capacity(heap_capacity);
459
460 for row_result in rows {
461 let (id_str, indices_json, values_blob) = row_result?;
462
463 let subject_id = Uuid::parse_str(&id_str).map_err(|e| {
464 rusqlite::Error::FromSqlConversionFailure(
465 0,
466 rusqlite::types::Type::Text,
467 Box::new(e),
468 )
469 })?;
470
471 let stored_indices: Vec<u32> =
473 serde_json::from_str(&indices_json).map_err(|e| {
474 rusqlite::Error::FromSqlConversionFailure(
475 0,
476 rusqlite::types::Type::Text,
477 Box::<dyn std::error::Error + Send + Sync>::from(format!(
478 "corrupt sparse row {id_str}: invalid indices JSON: {e}"
479 )),
480 )
481 })?;
482
483 if values_blob.len() % 4 != 0 {
484 return Err(rusqlite::Error::FromSqlConversionFailure(
485 0,
486 rusqlite::types::Type::Blob,
487 Box::<dyn std::error::Error + Send + Sync>::from(format!(
488 "corrupt sparse row {id_str}: values blob length {} not a multiple of 4",
489 values_blob.len()
490 )),
491 ));
492 }
493
494 #[allow(unknown_lints, clippy::chunks_exact_to_as_chunks)]
497 let stored_values: Vec<f32> = values_blob
498 .chunks_exact(4)
499 .map(|b| f32::from_le_bytes([b[0], b[1], b[2], b[3]]))
500 .collect();
501
502 validate_persisted_sparse(&id_str, &stored_indices, &stored_values)?;
503
504 let score = sparse_dot_product(
505 &query.indices,
506 &query.values,
507 &stored_indices,
508 &stored_values,
509 );
510
511 heap.push(Reverse(ScoredCandidate { score, subject_id }));
512 if heap.len() > top_k {
513 heap.pop();
514 }
515 }
516
517 let mut top: Vec<_> = heap.into_iter().map(|Reverse(c)| c).collect();
519 top.sort_by(|a, b| {
520 b.score
521 .partial_cmp(&a.score)
522 .unwrap_or(std::cmp::Ordering::Equal)
523 .then_with(|| a.subject_id.cmp(&b.subject_id))
524 });
525
526 let hits = top
527 .into_iter()
528 .enumerate()
529 .map(|(i, c)| SparseSearchHit {
530 subject_id: c.subject_id,
531 score: DeterministicScore::from_f64(c.score),
532 rank: (i + 1) as u32,
533 })
534 .collect();
535
536 Ok(hits)
537 })
538 .await
539 }
540
541 async fn count_sparse_rows(&self) -> Result<u64, StorageError> {
542 let table = self.table_name.clone();
543 let namespace = self.namespace.clone();
544 self.with_reader("sparse_count", move |conn| {
545 let sql = format!("SELECT COUNT(*) FROM {table} WHERE namespace = ?1");
546 let count: i64 =
547 conn.query_row(&sql, rusqlite::params![&namespace], |row| row.get(0))?;
548 Ok(count as u64)
549 })
550 .await
551 }
552}
553
554#[derive(PartialEq)]
557struct ScoredCandidate {
558 score: f64,
559 subject_id: Uuid,
560}
561
562impl Eq for ScoredCandidate {}
563
564impl PartialOrd for ScoredCandidate {
565 fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
566 Some(self.cmp(other))
567 }
568}
569
570impl Ord for ScoredCandidate {
571 fn cmp(&self, other: &Self) -> std::cmp::Ordering {
572 match self
575 .score
576 .partial_cmp(&other.score)
577 .unwrap_or(std::cmp::Ordering::Equal)
578 {
579 std::cmp::Ordering::Equal => other.subject_id.cmp(&self.subject_id),
580 ord => ord,
581 }
582 }
583}
584
585fn validate_persisted_sparse(
589 subject_id: &str,
590 indices: &[u32],
591 values: &[f32],
592) -> Result<(), rusqlite::Error> {
593 if indices.len() != values.len() {
594 return Err(rusqlite::Error::FromSqlConversionFailure(
595 0,
596 rusqlite::types::Type::Blob,
597 Box::<dyn std::error::Error + Send + Sync>::from(format!(
598 "corrupt sparse row {subject_id}: indices len {} != values len {}",
599 indices.len(),
600 values.len()
601 )),
602 ));
603 }
604 for (i, v) in values.iter().enumerate() {
605 if !v.is_finite() {
606 return Err(rusqlite::Error::FromSqlConversionFailure(
607 0,
608 rusqlite::types::Type::Blob,
609 Box::<dyn std::error::Error + Send + Sync>::from(format!(
610 "corrupt sparse row {subject_id}: non-finite value at position {i}: {v}"
611 )),
612 ));
613 }
614 }
615 for window in indices.windows(2) {
616 if window[0] >= window[1] {
617 return Err(rusqlite::Error::FromSqlConversionFailure(
618 0,
619 rusqlite::types::Type::Blob,
620 Box::<dyn std::error::Error + Send + Sync>::from(format!(
621 "corrupt sparse row {subject_id}: indices not strictly increasing at {} >= {}",
622 window[0], window[1]
623 )),
624 ));
625 }
626 }
627 Ok(())
628}
629
630fn sparse_dot_product(q_idx: &[u32], q_val: &[f32], s_idx: &[u32], s_val: &[f32]) -> f64 {
632 let mut dot = 0.0f64;
633 let mut qi = 0;
634 let mut si = 0;
635 while qi < q_idx.len() && si < s_idx.len() {
636 match q_idx[qi].cmp(&s_idx[si]) {
637 std::cmp::Ordering::Equal => {
638 dot += q_val[qi] as f64 * s_val[si] as f64;
639 qi += 1;
640 si += 1;
641 }
642 std::cmp::Ordering::Less => qi += 1,
643 std::cmp::Ordering::Greater => si += 1,
644 }
645 }
646 dot
647}
648
649#[async_trait]
650impl SparseStore for SqliteSparseStore {
651 async fn insert_sparse(
652 &self,
653 subject_id: Uuid,
654 kind: SubstrateKind,
655 namespace: &str,
656 field: &str,
657 vector: SparseVector,
658 ) -> Result<(), StorageError> {
659 validate_sparse_vector(&vector, "sparse_insert")?;
660 self.upsert_sparse_vector(subject_id, kind, namespace, field, vector)
661 .await
662 }
663
664 async fn insert_batch(
665 &self,
666 records: Vec<SparseRecord>,
667 ) -> Result<BatchWriteSummary, StorageError> {
668 self.insert_sparse_batch(records).await
669 }
670
671 async fn delete(&self, subject_id: Uuid) -> Result<bool, StorageError> {
672 self.delete_sparse_subject(subject_id).await
673 }
674
675 async fn search_sparse(
676 &self,
677 request: SparseSearchRequest,
678 ) -> Result<Vec<SparseSearchHit>, StorageError> {
679 validate_sparse_vector(&request.query, "sparse_search")?;
680 self.search_sparse_vectors(request).await
681 }
682
683 async fn count(&self) -> Result<u64, StorageError> {
684 self.count_sparse_rows().await
685 }
686}
687
688#[cfg(test)]
689mod tests {
690 use super::*;
691 use crate::pool::{ConnectionPool, PoolConfig};
692
693 fn make_store(model_key: &str) -> SqliteSparseStore {
694 let config = PoolConfig {
695 path: None,
696 ..PoolConfig::default()
697 };
698 let pool = Arc::new(ConnectionPool::new(config).expect("pool"));
699 {
701 let writer = pool.try_writer().expect("writer");
702 ensure_sparse_schema(writer.conn(), model_key).expect("schema");
703 }
704 SqliteSparseStore::new(pool, false, model_key.to_string(), "ns:test".to_string())
705 .expect("store")
706 }
707
708 fn sv(indices: Vec<u32>, values: Vec<f32>) -> SparseVector {
709 SparseVector { indices, values }
710 }
711
712 #[tokio::test]
713 async fn insert_and_count() {
714 let store = make_store("test_count");
715 let id = Uuid::new_v4();
716 store
717 .insert_sparse(
718 id,
719 SubstrateKind::Entity,
720 "ns:test",
721 "body",
722 sv(vec![0, 2], vec![1.0, 0.5]),
723 )
724 .await
725 .unwrap();
726 assert_eq!(store.count().await.unwrap(), 1);
727 }
728
729 #[tokio::test]
730 async fn insert_and_search() {
731 let store = make_store("test_search");
732 let id1 = Uuid::new_v4();
733 let id2 = Uuid::new_v4();
734 store
735 .insert_sparse(
736 id1,
737 SubstrateKind::Entity,
738 "ns:test",
739 "body",
740 sv(vec![0, 1], vec![1.0, 0.0]),
741 )
742 .await
743 .unwrap();
744 store
745 .insert_sparse(
746 id2,
747 SubstrateKind::Entity,
748 "ns:test",
749 "body",
750 sv(vec![0, 1], vec![0.0, 1.0]),
751 )
752 .await
753 .unwrap();
754
755 let hits = store
756 .search_sparse(SparseSearchRequest {
757 query: sv(vec![0], vec![1.0]),
758 top_k: 2,
759 namespace: Some("ns:test".into()),
760 kind: None,
761 })
762 .await
763 .unwrap();
764
765 assert!(!hits.is_empty());
766 assert_eq!(hits[0].subject_id, id1, "id1 should rank first");
767 assert_eq!(hits[0].rank, 1);
768 }
769
770 #[tokio::test]
773 async fn sparse_top_k_u32_max_rejected() {
774 let store = make_store("test_top_k_max");
775 let id = Uuid::new_v4();
776 store
777 .insert_sparse(
778 id,
779 SubstrateKind::Entity,
780 "ns:test",
781 "body",
782 sv(vec![0], vec![1.0]),
783 )
784 .await
785 .unwrap();
786
787 let result = store
788 .search_sparse(SparseSearchRequest {
789 query: sv(vec![0], vec![1.0]),
790 top_k: u32::MAX,
791 namespace: Some("ns:test".into()),
792 kind: None,
793 })
794 .await;
795
796 assert!(
797 matches!(result, Err(StorageError::InvalidInput { .. })),
798 "expected InvalidInput, got {result:?}"
799 );
800 }
801
802 #[tokio::test]
803 async fn delete_removes_row() {
804 let store = make_store("test_delete");
805 let id = Uuid::new_v4();
806 store
807 .insert_sparse(
808 id,
809 SubstrateKind::Entity,
810 "ns:test",
811 "body",
812 sv(vec![1], vec![1.0]),
813 )
814 .await
815 .unwrap();
816 assert_eq!(store.count().await.unwrap(), 1);
817
818 let deleted = store.delete(id).await.unwrap();
819 assert!(deleted);
820 assert_eq!(store.count().await.unwrap(), 0);
821 }
822
823 #[tokio::test]
824 async fn mismatched_lengths_rejected() {
825 let store = make_store("test_mismatch");
826 let result = store
827 .insert_sparse(
828 Uuid::new_v4(),
829 SubstrateKind::Entity,
830 "ns:test",
831 "body",
832 SparseVector {
833 indices: vec![0, 1],
834 values: vec![1.0],
835 },
836 )
837 .await;
838 assert!(matches!(result, Err(StorageError::InvalidInput { .. })));
839 }
840
841 #[tokio::test]
842 async fn non_finite_values_rejected() {
843 let store = make_store("test_nonfinite");
844 let result = store
845 .insert_sparse(
846 Uuid::new_v4(),
847 SubstrateKind::Entity,
848 "ns:test",
849 "body",
850 sv(vec![0], vec![f32::NAN]),
851 )
852 .await;
853 assert!(matches!(result, Err(StorageError::InvalidInput { .. })));
854 }
855
856 #[tokio::test]
857 async fn duplicate_indices_rejected() {
858 let store = make_store("test_dup_idx");
859 let result = store
860 .insert_sparse(
861 Uuid::new_v4(),
862 SubstrateKind::Entity,
863 "ns:test",
864 "body",
865 sv(vec![0, 0], vec![1.0, 2.0]),
866 )
867 .await;
868 assert!(matches!(result, Err(StorageError::InvalidInput { .. })));
869 }
870
871 #[tokio::test]
872 async fn empty_vector_rejected() {
873 let store = make_store("test_empty");
874 let result = store
875 .insert_sparse(
876 Uuid::new_v4(),
877 SubstrateKind::Entity,
878 "ns:test",
879 "body",
880 sv(vec![], vec![]),
881 )
882 .await;
883 assert!(matches!(result, Err(StorageError::InvalidInput { .. })));
884 }
885
886 #[tokio::test]
887 async fn namespace_isolation() {
888 let store = make_store("test_ns_iso");
889 let id = Uuid::new_v4();
890 store
891 .insert_sparse(
892 id,
893 SubstrateKind::Entity,
894 "ns:a",
895 "body",
896 sv(vec![0], vec![1.0]),
897 )
898 .await
899 .unwrap();
900
901 let hits = store
902 .search_sparse(SparseSearchRequest {
903 query: sv(vec![0], vec![1.0]),
904 top_k: 5,
905 namespace: Some("ns:b".into()),
906 kind: None,
907 })
908 .await
909 .unwrap();
910 assert!(hits.is_empty(), "ns:b should not see ns:a data");
911 }
912
913 #[tokio::test]
914 async fn insert_batch_happy_path() {
915 use chrono::Utc;
916 use khive_types::SubstrateKind;
917
918 let store = make_store("test_batch");
919 let id1 = Uuid::new_v4();
920 let id2 = Uuid::new_v4();
921 let records = vec![
922 SparseRecord {
923 subject_id: id1,
924 kind: SubstrateKind::Entity,
925 namespace: "ns:test".into(),
926 field: "body".into(),
927 vector: sv(vec![0, 3], vec![0.5, 0.8]),
928 updated_at: Utc::now(),
929 },
930 SparseRecord {
931 subject_id: id2,
932 kind: SubstrateKind::Entity,
933 namespace: "ns:test".into(),
934 field: "body".into(),
935 vector: sv(vec![1], vec![1.0]),
936 updated_at: Utc::now(),
937 },
938 ];
939 let summary = store.insert_batch(records).await.unwrap();
940 assert_eq!(summary.attempted, 2);
941 assert_eq!(summary.affected, 2);
942 assert_eq!(summary.failed, 0);
943 assert_eq!(store.count().await.unwrap(), 2);
944 }
945
946 #[tokio::test]
958 async fn insert_batch_routes_through_writer_task_when_flag_enabled() {
959 use chrono::Utc;
960 use khive_types::SubstrateKind;
961
962 let model_key = "write_queue_flag_test";
963 let dir = tempfile::tempdir().unwrap();
964 let path = dir.path().join("write_queue_sparse.db");
965 let pool_cfg = PoolConfig {
966 path: Some(path.clone()),
967 write_queue_enabled: Some(true),
968 ..PoolConfig::for_test()
969 };
970 let pool = Arc::new(ConnectionPool::new(pool_cfg).expect("pool"));
971 {
972 let writer = pool.writer().expect("writer");
973 ensure_sparse_schema(writer.conn(), model_key).expect("schema");
974 }
975
976 let store = SqliteSparseStore::new(
977 Arc::clone(&pool),
978 true,
979 model_key.to_string(),
980 "ns:test".to_string(),
981 )
982 .expect("store");
983
984 let id1 = Uuid::new_v4();
985 let id2 = Uuid::new_v4();
986 let records = vec![
987 SparseRecord {
988 subject_id: id1,
989 kind: SubstrateKind::Entity,
990 namespace: "ns:test".into(),
991 field: "body".into(),
992 vector: sv(vec![0, 3], vec![0.5, 0.8]),
993 updated_at: Utc::now(),
994 },
995 SparseRecord {
996 subject_id: id2,
997 kind: SubstrateKind::Entity,
998 namespace: "ns:test".into(),
999 field: "body".into(),
1000 vector: sv(vec![1], vec![1.0]),
1001 updated_at: Utc::now(),
1002 },
1003 ];
1004
1005 let summary = store.insert_batch(records).await.unwrap();
1006 assert_eq!(summary.attempted, 2);
1007 assert_eq!(summary.affected, 2);
1008 assert_eq!(summary.failed, 0);
1009 assert_eq!(store.count().await.unwrap(), 2);
1010 assert_eq!(
1011 pool.writer_task_spawn_count(),
1012 1,
1013 "the flag-ON path must actually spawn and use the writer task"
1014 );
1015 }
1016}