1#![allow(clippy::cast_possible_truncation)]
9#![allow(clippy::cast_sign_loss)]
10
11use crate::core::{Buffer, BufferMetadata, Chunk, ChunkMetadata, Context};
12use crate::error::{Result, StorageError};
13use crate::storage::schema::{
14 CHECK_SCHEMA_SQL, CURRENT_SCHEMA_VERSION, GET_VERSION_SQL, SCHEMA_SQL, SET_VERSION_SQL,
15};
16use crate::storage::traits::{Storage, StorageStats};
17use rusqlite::{Connection, OptionalExtension, params};
18use std::path::{Path, PathBuf};
19
20pub struct SqliteStorage {
33 conn: Connection,
35 path: Option<PathBuf>,
37}
38
39impl SqliteStorage {
40 pub fn open<P: AsRef<Path>>(path: P) -> Result<Self> {
50 let path = path.as_ref().to_path_buf();
51
52 if let Some(parent) = path.parent()
54 && !parent.exists()
55 {
56 std::fs::create_dir_all(parent).map_err(|e| StorageError::Database(e.to_string()))?;
57 }
58
59 let conn = Connection::open(&path).map_err(StorageError::from)?;
60
61 conn.execute("PRAGMA foreign_keys = ON;", [])
63 .map_err(StorageError::from)?;
64
65 let _: String = conn
67 .query_row("PRAGMA journal_mode = WAL;", [], |row| row.get(0))
68 .map_err(StorageError::from)?;
69
70 Ok(Self {
71 conn,
72 path: Some(path),
73 })
74 }
75
76 pub fn in_memory() -> Result<Self> {
84 let conn = Connection::open_in_memory().map_err(StorageError::from)?;
85 conn.execute("PRAGMA foreign_keys = ON;", [])
86 .map_err(StorageError::from)?;
87
88 Ok(Self { conn, path: None })
89 }
90
91 #[must_use]
93 pub fn path(&self) -> Option<&Path> {
94 self.path.as_deref()
95 }
96
97 fn get_schema_version(&self) -> Result<Option<u32>> {
99 let version: Option<String> = self
100 .conn
101 .query_row(GET_VERSION_SQL, [], |row| row.get(0))
102 .optional()
103 .map_err(StorageError::from)?;
104
105 Ok(version.and_then(|v| v.parse().ok()))
106 }
107
108 fn set_schema_version(&self, version: u32) -> Result<()> {
110 self.conn
111 .execute(SET_VERSION_SQL, params![version.to_string()])
112 .map_err(StorageError::from)?;
113 Ok(())
114 }
115
116 #[allow(clippy::cast_possible_wrap)]
118 fn now() -> i64 {
119 std::time::SystemTime::now()
120 .duration_since(std::time::UNIX_EPOCH)
121 .map_or(0, |d| d.as_secs() as i64)
122 }
123}
124
125impl Storage for SqliteStorage {
126 fn init(&mut self) -> Result<()> {
127 let is_init: i64 = self
129 .conn
130 .query_row(CHECK_SCHEMA_SQL, [], |row| row.get(0))
131 .map_err(StorageError::from)?;
132
133 if is_init == 0 {
134 self.conn
136 .execute_batch(SCHEMA_SQL)
137 .map_err(StorageError::from)?;
138 self.set_schema_version(CURRENT_SCHEMA_VERSION)?;
139 } else if let Some(current) = self.get_schema_version()?
140 && current < CURRENT_SCHEMA_VERSION
141 {
142 let migrations = crate::storage::schema::get_migrations_from(current);
144 for migration in migrations {
145 self.conn
146 .execute_batch(migration.sql)
147 .map_err(|e| StorageError::Migration(e.to_string()))?;
148 }
149 self.set_schema_version(CURRENT_SCHEMA_VERSION)?;
150 }
151
152 Ok(())
153 }
154
155 fn is_initialized(&self) -> Result<bool> {
156 let count: i64 = self
157 .conn
158 .query_row(CHECK_SCHEMA_SQL, [], |row| row.get(0))
159 .map_err(StorageError::from)?;
160 Ok(count > 0)
161 }
162
163 fn reset(&mut self) -> Result<()> {
164 self.conn
165 .execute_batch(
166 r"
167 DELETE FROM chunk_embeddings;
168 DELETE FROM chunks;
169 DELETE FROM buffers;
170 DELETE FROM context;
171 DELETE FROM metadata;
172 ",
173 )
174 .map_err(StorageError::from)?;
175 Ok(())
176 }
177
178 fn save_context(&mut self, context: &Context) -> Result<()> {
181 let data = serde_json::to_string(context).map_err(StorageError::from)?;
182 let now = Self::now();
183
184 self.conn
185 .execute(
186 r"
187 INSERT OR REPLACE INTO context (id, data, created_at, updated_at)
188 VALUES (1, ?, COALESCE((SELECT created_at FROM context WHERE id = 1), ?), ?)
189 ",
190 params![data, now, now],
191 )
192 .map_err(StorageError::from)?;
193
194 Ok(())
195 }
196
197 fn load_context(&self) -> Result<Option<Context>> {
198 let data: Option<String> = self
199 .conn
200 .query_row("SELECT data FROM context WHERE id = 1", [], |row| {
201 row.get(0)
202 })
203 .optional()
204 .map_err(StorageError::from)?;
205
206 match data {
207 Some(json) => {
208 let context = serde_json::from_str(&json).map_err(StorageError::from)?;
209 Ok(Some(context))
210 }
211 None => Ok(None),
212 }
213 }
214
215 fn delete_context(&mut self) -> Result<()> {
216 self.conn
217 .execute("DELETE FROM context WHERE id = 1", [])
218 .map_err(StorageError::from)?;
219 Ok(())
220 }
221
222 #[allow(clippy::cast_possible_wrap)]
225 fn add_buffer(&mut self, buffer: &Buffer) -> Result<i64> {
226 let now = Self::now();
227
228 self.conn
229 .execute(
230 r"
231 INSERT INTO buffers (
232 name, source_path, content, content_type, content_hash,
233 size, line_count, chunk_count, created_at, updated_at
234 ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
235 ",
236 params![
237 buffer.name,
238 buffer
239 .source
240 .as_ref()
241 .map(|p| p.to_string_lossy().to_string()),
242 buffer.content,
243 buffer.metadata.content_type,
244 buffer.metadata.content_hash,
245 buffer.metadata.size as i64,
246 buffer.metadata.line_count.map(|c| c as i64),
247 buffer.metadata.chunk_count.map(|c| c as i64),
248 now,
249 now,
250 ],
251 )
252 .map_err(StorageError::from)?;
253
254 Ok(self.conn.last_insert_rowid())
255 }
256
257 fn get_buffer(&self, id: i64) -> Result<Option<Buffer>> {
258 let result = self
259 .conn
260 .query_row(
261 r"
262 SELECT id, name, source_path, content, content_type, content_hash,
263 size, line_count, chunk_count, created_at, updated_at
264 FROM buffers WHERE id = ?
265 ",
266 params![id],
267 |row| {
268 Ok(Buffer {
269 id: Some(row.get::<_, i64>(0)?),
270 name: row.get(1)?,
271 source: row.get::<_, Option<String>>(2)?.map(PathBuf::from),
272 content: row.get(3)?,
273 metadata: BufferMetadata {
274 content_type: row.get(4)?,
275 content_hash: row.get(5)?,
276 size: row.get::<_, i64>(6)? as usize,
277 line_count: row.get::<_, Option<i64>>(7)?.map(|c| c as usize),
278 chunk_count: row.get::<_, Option<i64>>(8)?.map(|c| c as usize),
279 created_at: row.get(9)?,
280 updated_at: row.get(10)?,
281 },
282 })
283 },
284 )
285 .optional()
286 .map_err(StorageError::from)?;
287
288 Ok(result)
289 }
290
291 fn get_buffer_by_name(&self, name: &str) -> Result<Option<Buffer>> {
292 let id: Option<i64> = self
293 .conn
294 .query_row(
295 "SELECT id FROM buffers WHERE name = ?",
296 params![name],
297 |row| row.get(0),
298 )
299 .optional()
300 .map_err(StorageError::from)?;
301
302 id.map_or(Ok(None), |id| self.get_buffer(id))
303 }
304
305 fn list_buffers(&self) -> Result<Vec<Buffer>> {
306 let mut stmt = self
307 .conn
308 .prepare(
309 r"
310 SELECT id, name, source_path, content, content_type, content_hash,
311 size, line_count, chunk_count, created_at, updated_at
312 FROM buffers ORDER BY id
313 ",
314 )
315 .map_err(StorageError::from)?;
316
317 let buffers = stmt
318 .query_map([], |row| {
319 Ok(Buffer {
320 id: Some(row.get::<_, i64>(0)?),
321 name: row.get(1)?,
322 source: row.get::<_, Option<String>>(2)?.map(PathBuf::from),
323 content: row.get(3)?,
324 metadata: BufferMetadata {
325 content_type: row.get(4)?,
326 content_hash: row.get(5)?,
327 size: row.get::<_, i64>(6)? as usize,
328 line_count: row.get::<_, Option<i64>>(7)?.map(|c| c as usize),
329 chunk_count: row.get::<_, Option<i64>>(8)?.map(|c| c as usize),
330 created_at: row.get(9)?,
331 updated_at: row.get(10)?,
332 },
333 })
334 })
335 .map_err(StorageError::from)?
336 .collect::<std::result::Result<Vec<_>, _>>()
337 .map_err(StorageError::from)?;
338
339 Ok(buffers)
340 }
341
342 #[allow(clippy::cast_possible_wrap)]
343 fn update_buffer(&mut self, buffer: &Buffer) -> Result<()> {
344 let id = buffer.id.ok_or_else(|| StorageError::BufferNotFound {
345 identifier: "no ID".to_string(),
346 })?;
347
348 let now = Self::now();
349
350 self.conn
351 .execute(
352 r"
353 UPDATE buffers SET
354 name = ?, source_path = ?, content = ?, content_type = ?,
355 content_hash = ?, size = ?, line_count = ?, chunk_count = ?,
356 updated_at = ?
357 WHERE id = ?
358 ",
359 params![
360 buffer.name,
361 buffer
362 .source
363 .as_ref()
364 .map(|p| p.to_string_lossy().to_string()),
365 buffer.content,
366 buffer.metadata.content_type,
367 buffer.metadata.content_hash,
368 buffer.metadata.size as i64,
369 buffer.metadata.line_count.map(|c| c as i64),
370 buffer.metadata.chunk_count.map(|c| c as i64),
371 now,
372 id,
373 ],
374 )
375 .map_err(StorageError::from)?;
376
377 Ok(())
378 }
379
380 fn delete_buffer(&mut self, id: i64) -> Result<()> {
381 self.conn
383 .execute("DELETE FROM buffers WHERE id = ?", params![id])
384 .map_err(StorageError::from)?;
385 Ok(())
386 }
387
388 fn buffer_count(&self) -> Result<usize> {
389 let count: i64 = self
390 .conn
391 .query_row("SELECT COUNT(*) FROM buffers", [], |row| row.get(0))
392 .map_err(StorageError::from)?;
393 Ok(count as usize)
394 }
395
396 #[allow(clippy::cast_possible_wrap)]
399 fn add_chunks(&mut self, buffer_id: i64, chunks: &[Chunk]) -> Result<()> {
400 let tx = self.conn.transaction().map_err(StorageError::from)?;
401 let now = Self::now();
402
403 {
404 let mut stmt = tx
405 .prepare(
406 r"
407 INSERT INTO chunks (
408 buffer_id, content, byte_start, byte_end, chunk_index,
409 strategy, token_count, line_start, line_end, has_overlap,
410 content_hash, custom_metadata, created_at
411 ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
412 ",
413 )
414 .map_err(StorageError::from)?;
415
416 for chunk in chunks {
417 let custom_meta = chunk.metadata.custom.clone();
418
419 let (line_start, line_end) = chunk
420 .metadata
421 .line_range
422 .as_ref()
423 .map_or((None, None), |r| (Some(r.start as i64), Some(r.end as i64)));
424
425 stmt.execute(params![
426 buffer_id,
427 chunk.content,
428 chunk.byte_range.start as i64,
429 chunk.byte_range.end as i64,
430 chunk.index as i64,
431 chunk.metadata.strategy,
432 chunk.metadata.token_count.map(|c| c as i64),
433 line_start,
434 line_end,
435 i64::from(chunk.metadata.has_overlap),
436 chunk.metadata.content_hash,
437 custom_meta,
438 now,
439 ])
440 .map_err(StorageError::from)?;
441 }
442 }
443
444 tx.commit().map_err(StorageError::from)?;
445
446 self.conn
448 .execute(
449 "UPDATE buffers SET chunk_count = ? WHERE id = ?",
450 params![chunks.len() as i64, buffer_id],
451 )
452 .map_err(StorageError::from)?;
453
454 Ok(())
455 }
456
457 fn get_chunks(&self, buffer_id: i64) -> Result<Vec<Chunk>> {
458 let mut stmt = self
459 .conn
460 .prepare(
461 r"
462 SELECT id, buffer_id, content, byte_start, byte_end, chunk_index,
463 strategy, token_count, line_start, line_end, has_overlap,
464 content_hash, custom_metadata, created_at
465 FROM chunks WHERE buffer_id = ? ORDER BY chunk_index
466 ",
467 )
468 .map_err(StorageError::from)?;
469
470 let chunks = stmt
471 .query_map(params![buffer_id], |row| {
472 let line_start: Option<i64> = row.get(8)?;
473 let line_end: Option<i64> = row.get(9)?;
474 let line_range = match (line_start, line_end) {
475 (Some(s), Some(e)) => Some((s as usize)..(e as usize)),
476 _ => None,
477 };
478
479 Ok(Chunk {
480 id: Some(row.get::<_, i64>(0)?),
481 buffer_id: row.get(1)?,
482 content: row.get(2)?,
483 byte_range: (row.get::<_, i64>(3)? as usize)..(row.get::<_, i64>(4)? as usize),
484 index: row.get::<_, i64>(5)? as usize,
485 metadata: ChunkMetadata {
486 strategy: row.get(6)?,
487 token_count: row.get::<_, Option<i64>>(7)?.map(|c| c as usize),
488 line_range,
489 has_overlap: row.get::<_, i64>(10)? != 0,
490 content_hash: row.get(11)?,
491 custom: row.get(12)?,
492 created_at: row.get(13)?,
493 },
494 })
495 })
496 .map_err(StorageError::from)?
497 .collect::<std::result::Result<Vec<_>, _>>()
498 .map_err(StorageError::from)?;
499
500 Ok(chunks)
501 }
502
503 fn get_chunk(&self, id: i64) -> Result<Option<Chunk>> {
504 let result = self
505 .conn
506 .query_row(
507 r"
508 SELECT id, buffer_id, content, byte_start, byte_end, chunk_index,
509 strategy, token_count, line_start, line_end, has_overlap,
510 content_hash, custom_metadata, created_at
511 FROM chunks WHERE id = ?
512 ",
513 params![id],
514 |row| {
515 let line_start: Option<i64> = row.get(8)?;
516 let line_end: Option<i64> = row.get(9)?;
517 let line_range = match (line_start, line_end) {
518 (Some(s), Some(e)) => Some((s as usize)..(e as usize)),
519 _ => None,
520 };
521
522 Ok(Chunk {
523 id: Some(row.get::<_, i64>(0)?),
524 buffer_id: row.get(1)?,
525 content: row.get(2)?,
526 byte_range: (row.get::<_, i64>(3)? as usize)
527 ..(row.get::<_, i64>(4)? as usize),
528 index: row.get::<_, i64>(5)? as usize,
529 metadata: ChunkMetadata {
530 strategy: row.get(6)?,
531 token_count: row.get::<_, Option<i64>>(7)?.map(|c| c as usize),
532 line_range,
533 has_overlap: row.get::<_, i64>(10)? != 0,
534 content_hash: row.get(11)?,
535 custom: row.get(12)?,
536 created_at: row.get(13)?,
537 },
538 })
539 },
540 )
541 .optional()
542 .map_err(StorageError::from)?;
543
544 Ok(result)
545 }
546
547 fn delete_chunks(&mut self, buffer_id: i64) -> Result<()> {
548 self.conn
549 .execute("DELETE FROM chunks WHERE buffer_id = ?", params![buffer_id])
550 .map_err(StorageError::from)?;
551
552 self.conn
554 .execute(
555 "UPDATE buffers SET chunk_count = 0 WHERE id = ?",
556 params![buffer_id],
557 )
558 .map_err(StorageError::from)?;
559
560 Ok(())
561 }
562
563 fn chunk_count(&self, buffer_id: i64) -> Result<usize> {
564 let count: i64 = self
565 .conn
566 .query_row(
567 "SELECT COUNT(*) FROM chunks WHERE buffer_id = ?",
568 params![buffer_id],
569 |row| row.get(0),
570 )
571 .map_err(StorageError::from)?;
572 Ok(count as usize)
573 }
574
575 fn export_buffers(&self) -> Result<String> {
578 let buffers = self.list_buffers()?;
579 let mut output = String::new();
580
581 for (i, buffer) in buffers.iter().enumerate() {
582 if i > 0 {
583 output.push_str("\n\n");
584 }
585 output.push_str(&buffer.content);
586 }
587
588 Ok(output)
589 }
590
591 fn stats(&self) -> Result<StorageStats> {
592 let buffer_count = self.buffer_count()?;
593
594 let chunk_count: i64 = self
595 .conn
596 .query_row("SELECT COUNT(*) FROM chunks", [], |row| row.get(0))
597 .map_err(StorageError::from)?;
598
599 let total_size: i64 = self
600 .conn
601 .query_row("SELECT COALESCE(SUM(size), 0) FROM buffers", [], |row| {
602 row.get(0)
603 })
604 .map_err(StorageError::from)?;
605
606 let has_context = self.load_context()?.is_some();
607
608 let schema_version = self.get_schema_version()?.unwrap_or(0);
609
610 let db_size = self
611 .path
612 .as_ref()
613 .and_then(|p| std::fs::metadata(p).ok().map(|m| m.len()));
614
615 Ok(StorageStats {
616 buffer_count,
617 chunk_count: chunk_count as usize,
618 total_content_size: total_size as usize,
619 has_context,
620 schema_version,
621 db_size,
622 })
623 }
624}
625
626impl SqliteStorage {
629 pub fn get_chunks_by_ids(&self, ids: &[i64]) -> Result<std::collections::HashMap<i64, Chunk>> {
640 if ids.is_empty() {
641 return Ok(std::collections::HashMap::new());
642 }
643
644 let placeholders = ids.iter().map(|_| "?").collect::<Vec<_>>().join(",");
645 let sql = format!(
646 "SELECT id, buffer_id, content, byte_start, byte_end, chunk_index, \
647 strategy, token_count, line_start, line_end, has_overlap, \
648 content_hash, custom_metadata, created_at \
649 FROM chunks WHERE id IN ({placeholders})"
650 );
651
652 let mut stmt = self.conn.prepare(&sql).map_err(StorageError::from)?;
653
654 let chunks = stmt
655 .query_map(rusqlite::params_from_iter(ids.iter().copied()), |row| {
656 let line_start: Option<i64> = row.get(8)?;
657 let line_end: Option<i64> = row.get(9)?;
658 let line_range = match (line_start, line_end) {
659 (Some(s), Some(e)) => Some((s as usize)..(e as usize)),
660 _ => None,
661 };
662
663 Ok(Chunk {
664 id: Some(row.get::<_, i64>(0)?),
665 buffer_id: row.get(1)?,
666 content: row.get(2)?,
667 byte_range: (row.get::<_, i64>(3)? as usize)..(row.get::<_, i64>(4)? as usize),
668 index: row.get::<_, i64>(5)? as usize,
669 metadata: ChunkMetadata {
670 strategy: row.get(6)?,
671 token_count: row.get::<_, Option<i64>>(7)?.map(|c| c as usize),
672 line_range,
673 has_overlap: row.get::<_, i64>(10)? != 0,
674 content_hash: row.get(11)?,
675 custom: row.get(12)?,
676 created_at: row.get(13)?,
677 },
678 })
679 })
680 .map_err(StorageError::from)?
681 .collect::<std::result::Result<Vec<_>, _>>()
682 .map_err(StorageError::from)?;
683
684 Ok(chunks
685 .into_iter()
686 .filter_map(|c| c.id.map(|id| (id, c)))
687 .collect())
688 }
689
690 #[allow(clippy::cast_possible_wrap)]
702 pub fn store_embedding(
703 &mut self,
704 chunk_id: i64,
705 embedding: &[f32],
706 model_name: Option<&str>,
707 ) -> Result<()> {
708 let now = Self::now();
709
710 let mut bytes = Vec::with_capacity(embedding.len() * 4);
712 for f in embedding {
713 bytes.extend_from_slice(&f.to_le_bytes());
714 }
715
716 self.conn
717 .execute(
718 r"
719 INSERT OR REPLACE INTO chunk_embeddings (chunk_id, embedding, dimensions, model_name, created_at)
720 VALUES (?, ?, ?, ?, ?)
721 ",
722 params![chunk_id, bytes, embedding.len() as i64, model_name, now],
723 )
724 .map_err(StorageError::from)?;
725
726 Ok(())
727 }
728
729 pub fn get_embedding(&self, chunk_id: i64) -> Result<Option<Vec<f32>>> {
735 let result: Option<Vec<u8>> = self
736 .conn
737 .query_row(
738 "SELECT embedding FROM chunk_embeddings WHERE chunk_id = ?",
739 params![chunk_id],
740 |row| row.get(0),
741 )
742 .optional()
743 .map_err(StorageError::from)?;
744
745 Ok(result.map(|bytes| {
746 bytes
747 .chunks_exact(4)
748 .map(|chunk| f32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]))
749 .collect()
750 }))
751 }
752
753 pub fn get_embedding_models(&self, buffer_id: i64) -> Result<Vec<String>> {
762 let mut stmt = self
763 .conn
764 .prepare(
765 r"
766 SELECT DISTINCT ce.model_name
767 FROM chunk_embeddings ce
768 JOIN chunks c ON ce.chunk_id = c.id
769 WHERE c.buffer_id = ? AND ce.model_name IS NOT NULL
770 ",
771 )
772 .map_err(StorageError::from)?;
773
774 let models = stmt
775 .query_map(params![buffer_id], |row| row.get::<_, String>(0))
776 .map_err(StorageError::from)?
777 .filter_map(std::result::Result::ok)
778 .collect();
779
780 Ok(models)
781 }
782
783 pub fn get_embedding_model_counts(&self, buffer_id: i64) -> Result<Vec<(Option<String>, i64)>> {
791 let mut stmt = self
792 .conn
793 .prepare(
794 r"
795 SELECT ce.model_name, COUNT(*) as count
796 FROM chunk_embeddings ce
797 JOIN chunks c ON ce.chunk_id = c.id
798 WHERE c.buffer_id = ?
799 GROUP BY ce.model_name
800 ",
801 )
802 .map_err(StorageError::from)?;
803
804 let counts = stmt
805 .query_map(params![buffer_id], |row| {
806 Ok((row.get::<_, Option<String>>(0)?, row.get::<_, i64>(1)?))
807 })
808 .map_err(StorageError::from)?
809 .filter_map(std::result::Result::ok)
810 .collect();
811
812 Ok(counts)
813 }
814
815 #[allow(clippy::cast_possible_wrap)]
821 pub fn store_embeddings_batch(
822 &mut self,
823 embeddings: &[(i64, Vec<f32>)],
824 model_name: Option<&str>,
825 ) -> Result<()> {
826 let tx = self.conn.transaction().map_err(StorageError::from)?;
827 let now = Self::now();
828
829 {
830 let mut stmt = tx
831 .prepare(
832 r"
833 INSERT OR REPLACE INTO chunk_embeddings (chunk_id, embedding, dimensions, model_name, created_at)
834 VALUES (?, ?, ?, ?, ?)
835 ",
836 )
837 .map_err(StorageError::from)?;
838
839 for (chunk_id, embedding) in embeddings {
840 let mut bytes = Vec::with_capacity(embedding.len() * 4);
841 bytes.extend(embedding.iter().flat_map(|f| f.to_le_bytes()));
842
843 stmt.execute(params![
844 chunk_id,
845 bytes,
846 embedding.len() as i64,
847 model_name,
848 now
849 ])
850 .map_err(StorageError::from)?;
851 }
852 }
853
854 tx.commit().map_err(StorageError::from)?;
855 Ok(())
856 }
857
858 pub fn delete_embedding(&mut self, chunk_id: i64) -> Result<()> {
864 self.conn
865 .execute(
866 "DELETE FROM chunk_embeddings WHERE chunk_id = ?",
867 params![chunk_id],
868 )
869 .map_err(StorageError::from)?;
870 Ok(())
871 }
872
873 #[allow(clippy::cast_possible_wrap)]
886 pub fn search_fts(&self, query: &str, limit: usize) -> Result<Vec<(i64, f64)>> {
887 self.search_fts_in_buffer(query, limit, None)
888 }
889
890 #[allow(clippy::cast_possible_wrap)]
896 pub fn search_fts_in_buffer(
897 &self,
898 query: &str,
899 limit: usize,
900 buffer_id: Option<i64>,
901 ) -> Result<Vec<(i64, f64)>> {
902 let fts_query = query
909 .split_whitespace()
910 .map(|term| format!("\"{}\"", term.replace('"', "\"\"")))
911 .collect::<Vec<_>>()
912 .join(" OR ");
913
914 let results = if let Some(buffer_id) = buffer_id {
915 let mut stmt = self
916 .conn
917 .prepare(
918 r"
919 SELECT chunks_fts.rowid, -bm25(chunks_fts) as score
920 FROM chunks_fts
921 INNER JOIN chunks ON chunks.id = chunks_fts.rowid
922 WHERE chunks_fts MATCH ?1
923 AND chunks.buffer_id = ?2
924 ORDER BY score DESC
925 LIMIT ?3
926 ",
927 )
928 .map_err(StorageError::from)?;
929
930 stmt.query_map(params![fts_query, buffer_id, limit as i64], |row| {
931 Ok((row.get::<_, i64>(0)?, row.get::<_, f64>(1)?))
932 })
933 .map_err(StorageError::from)?
934 .collect::<std::result::Result<Vec<_>, _>>()
935 .map_err(StorageError::from)?
936 } else {
937 let mut stmt = self
938 .conn
939 .prepare(
940 r"
941 SELECT rowid, -bm25(chunks_fts) as score
942 FROM chunks_fts
943 WHERE chunks_fts MATCH ?1
944 ORDER BY score DESC
945 LIMIT ?2
946 ",
947 )
948 .map_err(StorageError::from)?;
949
950 stmt.query_map(params![fts_query, limit as i64], |row| {
951 Ok((row.get::<_, i64>(0)?, row.get::<_, f64>(1)?))
952 })
953 .map_err(StorageError::from)?
954 .collect::<std::result::Result<Vec<_>, _>>()
955 .map_err(StorageError::from)?
956 };
957
958 Ok(results)
959 }
960
961 pub fn get_all_embeddings(&self) -> Result<Vec<(i64, Vec<f32>)>> {
967 let mut stmt = self
968 .conn
969 .prepare("SELECT chunk_id, embedding FROM chunk_embeddings")
970 .map_err(StorageError::from)?;
971
972 let results = stmt
973 .query_map([], |row| {
974 let chunk_id: i64 = row.get(0)?;
975 let bytes: Vec<u8> = row.get(1)?;
976 let embedding: Vec<f32> = bytes
977 .chunks_exact(4)
978 .map(|chunk| f32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]))
979 .collect();
980 Ok((chunk_id, embedding))
981 })
982 .map_err(StorageError::from)?
983 .collect::<std::result::Result<Vec<_>, _>>()
984 .map_err(StorageError::from)?;
985
986 Ok(results)
987 }
988
989 pub fn get_embeddings_in_buffer(&self, buffer_id: i64) -> Result<Vec<(i64, Vec<f32>)>> {
995 let mut stmt = self
996 .conn
997 .prepare(
998 r"
999 SELECT ce.chunk_id, ce.embedding
1000 FROM chunk_embeddings ce
1001 INNER JOIN chunks c ON c.id = ce.chunk_id
1002 WHERE c.buffer_id = ?
1003 ",
1004 )
1005 .map_err(StorageError::from)?;
1006
1007 let results = stmt
1008 .query_map(params![buffer_id], |row| {
1009 let chunk_id: i64 = row.get(0)?;
1010 let bytes: Vec<u8> = row.get(1)?;
1011 let embedding: Vec<f32> = bytes
1012 .chunks_exact(4)
1013 .map(|chunk| f32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]))
1014 .collect();
1015 Ok((chunk_id, embedding))
1016 })
1017 .map_err(StorageError::from)?
1018 .collect::<std::result::Result<Vec<_>, _>>()
1019 .map_err(StorageError::from)?;
1020
1021 Ok(results)
1022 }
1023
1024 pub fn embedding_count(&self) -> Result<usize> {
1030 let count: i64 = self
1031 .conn
1032 .query_row("SELECT COUNT(*) FROM chunk_embeddings", [], |row| {
1033 row.get(0)
1034 })
1035 .map_err(StorageError::from)?;
1036 Ok(count as usize)
1037 }
1038
1039 pub fn all_chunks_have_embeddings(&self, buffer_id: i64) -> Result<bool> {
1048 let result: i64 = self
1049 .conn
1050 .query_row(
1051 r"
1052 SELECT NOT EXISTS (
1053 SELECT 1 FROM chunks c
1054 LEFT JOIN chunk_embeddings e ON e.chunk_id = c.id
1055 WHERE c.buffer_id = ? AND e.chunk_id IS NULL
1056 )
1057 ",
1058 params![buffer_id],
1059 |row| row.get(0),
1060 )
1061 .map_err(StorageError::from)?;
1062 Ok(result != 0)
1063 }
1064
1065 pub fn has_embedding(&self, chunk_id: i64) -> Result<bool> {
1071 let count: i64 = self
1072 .conn
1073 .query_row(
1074 "SELECT COUNT(*) FROM chunk_embeddings WHERE chunk_id = ?",
1075 params![chunk_id],
1076 |row| row.get(0),
1077 )
1078 .map_err(StorageError::from)?;
1079 Ok(count > 0)
1080 }
1081
1082 pub fn get_chunks_needing_embedding(
1098 &self,
1099 buffer_id: i64,
1100 current_model: Option<&str>,
1101 ) -> Result<Vec<i64>> {
1102 let mut results = Vec::new();
1103
1104 let mut stmt = self
1106 .conn
1107 .prepare(
1108 r"
1109 SELECT c.id FROM chunks c
1110 LEFT JOIN chunk_embeddings e ON c.id = e.chunk_id
1111 WHERE c.buffer_id = ? AND e.chunk_id IS NULL
1112 ",
1113 )
1114 .map_err(StorageError::from)?;
1115
1116 let rows = stmt
1117 .query_map(params![buffer_id], |row| row.get(0))
1118 .map_err(StorageError::from)?;
1119
1120 for row in rows {
1121 results.push(row.map_err(StorageError::from)?);
1122 }
1123
1124 if let Some(model) = current_model {
1126 let mut stmt = self
1127 .conn
1128 .prepare(
1129 r"
1130 SELECT c.id FROM chunks c
1131 INNER JOIN chunk_embeddings e ON c.id = e.chunk_id
1132 WHERE c.buffer_id = ? AND (e.model_name IS NULL OR e.model_name != ?)
1133 ",
1134 )
1135 .map_err(StorageError::from)?;
1136
1137 let rows = stmt
1138 .query_map(params![buffer_id, model], |row| row.get(0))
1139 .map_err(StorageError::from)?;
1140
1141 for row in rows {
1142 results.push(row.map_err(StorageError::from)?);
1143 }
1144 }
1145
1146 results.sort_unstable();
1148 results.dedup();
1149 Ok(results)
1150 }
1151
1152 pub fn get_chunks_without_embedding(&self, buffer_id: i64) -> Result<Vec<i64>> {
1160 self.get_chunks_needing_embedding(buffer_id, None)
1161 }
1162
1163 pub fn delete_embeddings_by_model(
1180 &mut self,
1181 buffer_id: i64,
1182 model_name: Option<&str>,
1183 ) -> Result<usize> {
1184 let deleted = match model_name {
1185 Some(name) => self
1186 .conn
1187 .execute(
1188 r"
1189 DELETE FROM chunk_embeddings
1190 WHERE chunk_id IN (
1191 SELECT id FROM chunks WHERE buffer_id = ?
1192 ) AND model_name = ?
1193 ",
1194 params![buffer_id, name],
1195 )
1196 .map_err(StorageError::from)?,
1197 None => self
1198 .conn
1199 .execute(
1200 r"
1201 DELETE FROM chunk_embeddings
1202 WHERE chunk_id IN (
1203 SELECT id FROM chunks WHERE buffer_id = ?
1204 ) AND model_name IS NULL
1205 ",
1206 params![buffer_id],
1207 )
1208 .map_err(StorageError::from)?,
1209 };
1210 Ok(deleted)
1211 }
1212
1213 pub fn get_embedding_stats(&self, buffer_id: i64) -> Result<EmbeddingStats> {
1221 let total_chunks: i64 = self
1223 .conn
1224 .query_row(
1225 "SELECT COUNT(*) FROM chunks WHERE buffer_id = ?",
1226 params![buffer_id],
1227 |row| row.get(0),
1228 )
1229 .map_err(StorageError::from)?;
1230
1231 let embedded_chunks: i64 = self
1233 .conn
1234 .query_row(
1235 r"
1236 SELECT COUNT(*) FROM chunk_embeddings e
1237 INNER JOIN chunks c ON e.chunk_id = c.id
1238 WHERE c.buffer_id = ?
1239 ",
1240 params![buffer_id],
1241 |row| row.get(0),
1242 )
1243 .map_err(StorageError::from)?;
1244
1245 let model_counts = self.get_embedding_model_counts(buffer_id)?;
1247
1248 Ok(EmbeddingStats {
1249 total_chunks: total_chunks as usize,
1250 embedded_chunks: embedded_chunks as usize,
1251 model_counts,
1252 })
1253 }
1254}
1255
1256#[derive(Debug, Clone)]
1258pub struct EmbeddingStats {
1259 pub total_chunks: usize,
1261 pub embedded_chunks: usize,
1263 pub model_counts: Vec<(Option<String>, i64)>,
1265}
1266
1267#[cfg(test)]
1268mod tests {
1269 use super::*;
1270 use crate::core::ContextValue;
1271
1272 fn setup() -> SqliteStorage {
1273 let mut storage = SqliteStorage::in_memory().unwrap();
1274 storage.init().unwrap();
1275 storage
1276 }
1277
1278 #[test]
1279 fn test_init() {
1280 let mut storage = SqliteStorage::in_memory().unwrap();
1281 assert!(storage.init().is_ok());
1282 assert!(storage.is_initialized().unwrap());
1283 }
1284
1285 #[test]
1286 fn test_init_idempotent() {
1287 let mut storage = SqliteStorage::in_memory().unwrap();
1288 assert!(storage.init().is_ok());
1289 assert!(storage.init().is_ok()); }
1291
1292 #[test]
1293 fn test_context_crud() {
1294 let mut storage = setup();
1295
1296 assert!(storage.load_context().unwrap().is_none());
1298
1299 let mut ctx = Context::new();
1301 ctx.set_variable("key".to_string(), ContextValue::String("value".to_string()));
1302 storage.save_context(&ctx).unwrap();
1303
1304 let loaded = storage.load_context().unwrap().unwrap();
1306 assert_eq!(
1307 loaded.get_variable("key"),
1308 Some(&ContextValue::String("value".to_string()))
1309 );
1310
1311 storage.delete_context().unwrap();
1313 assert!(storage.load_context().unwrap().is_none());
1314 }
1315
1316 #[test]
1317 fn test_buffer_crud() {
1318 let mut storage = setup();
1319
1320 let buffer = Buffer::from_named("test".to_string(), "Hello, world!".to_string());
1322 let id = storage.add_buffer(&buffer).unwrap();
1323 assert!(id > 0);
1324
1325 let loaded = storage.get_buffer(id).unwrap().unwrap();
1327 assert_eq!(loaded.name, Some("test".to_string()));
1328 assert_eq!(loaded.content, "Hello, world!");
1329
1330 let by_name = storage.get_buffer_by_name("test").unwrap().unwrap();
1332 assert_eq!(by_name.id, Some(id));
1333
1334 let buffers = storage.list_buffers().unwrap();
1336 assert_eq!(buffers.len(), 1);
1337
1338 let mut updated = loaded;
1340 updated.content = "Updated content".to_string();
1341 storage.update_buffer(&updated).unwrap();
1342
1343 let reloaded = storage.get_buffer(id).unwrap().unwrap();
1344 assert_eq!(reloaded.content, "Updated content");
1345
1346 storage.delete_buffer(id).unwrap();
1348 assert!(storage.get_buffer(id).unwrap().is_none());
1349 }
1350
1351 #[test]
1352 fn test_chunk_crud() {
1353 let mut storage = setup();
1354
1355 let buffer = Buffer::from_content("Hello, world!".to_string());
1357 let buffer_id = storage.add_buffer(&buffer).unwrap();
1358
1359 let chunks = vec![
1361 Chunk::new(buffer_id, "Hello, ".to_string(), 0..7, 0),
1362 Chunk::new(buffer_id, "world!".to_string(), 7..13, 1),
1363 ];
1364 storage.add_chunks(buffer_id, &chunks).unwrap();
1365
1366 let loaded = storage.get_chunks(buffer_id).unwrap();
1368 assert_eq!(loaded.len(), 2);
1369 assert_eq!(loaded[0].content, "Hello, ");
1370 assert_eq!(loaded[1].content, "world!");
1371
1372 assert_eq!(storage.chunk_count(buffer_id).unwrap(), 2);
1374
1375 let chunk_id = loaded[0].id.unwrap();
1377 let single = storage.get_chunk(chunk_id).unwrap().unwrap();
1378 assert_eq!(single.content, "Hello, ");
1379
1380 storage.delete_chunks(buffer_id).unwrap();
1382 assert_eq!(storage.chunk_count(buffer_id).unwrap(), 0);
1383 }
1384
1385 #[test]
1386 fn test_cascade_delete() {
1387 let mut storage = setup();
1388
1389 let buffer = Buffer::from_content("Hello, world!".to_string());
1391 let buffer_id = storage.add_buffer(&buffer).unwrap();
1392
1393 let chunks = vec![Chunk::new(buffer_id, "Hello".to_string(), 0..5, 0)];
1394 storage.add_chunks(buffer_id, &chunks).unwrap();
1395
1396 assert_eq!(storage.chunk_count(buffer_id).unwrap(), 1);
1398
1399 storage.delete_buffer(buffer_id).unwrap();
1401
1402 let count: i64 = storage
1404 .conn
1405 .query_row("SELECT COUNT(*) FROM chunks", [], |row| row.get(0))
1406 .unwrap();
1407 assert_eq!(count, 0);
1408 }
1409
1410 #[test]
1411 fn test_reset() {
1412 let mut storage = setup();
1413
1414 let ctx = Context::new();
1416 storage.save_context(&ctx).unwrap();
1417
1418 let buffer = Buffer::from_content("test".to_string());
1419 storage.add_buffer(&buffer).unwrap();
1420
1421 storage.reset().unwrap();
1423
1424 assert!(storage.load_context().unwrap().is_none());
1426 assert_eq!(storage.buffer_count().unwrap(), 0);
1427 }
1428
1429 #[test]
1430 fn test_stats() {
1431 let mut storage = setup();
1432
1433 let stats = storage.stats().unwrap();
1435 assert_eq!(stats.buffer_count, 0);
1436 assert_eq!(stats.chunk_count, 0);
1437 assert!(!stats.has_context);
1438
1439 let ctx = Context::new();
1441 storage.save_context(&ctx).unwrap();
1442
1443 let buffer = Buffer::from_content("Hello, world!".to_string());
1444 let buffer_id = storage.add_buffer(&buffer).unwrap();
1445
1446 let chunks = vec![Chunk::new(buffer_id, "Hello".to_string(), 0..5, 0)];
1447 storage.add_chunks(buffer_id, &chunks).unwrap();
1448
1449 let stats = storage.stats().unwrap();
1451 assert_eq!(stats.buffer_count, 1);
1452 assert_eq!(stats.chunk_count, 1);
1453 assert!(stats.has_context);
1454 assert_eq!(stats.total_content_size, 13);
1455 }
1456
1457 #[test]
1458 fn test_export_buffers() {
1459 let mut storage = setup();
1460
1461 storage
1462 .add_buffer(&Buffer::from_content("First".to_string()))
1463 .unwrap();
1464 storage
1465 .add_buffer(&Buffer::from_content("Second".to_string()))
1466 .unwrap();
1467
1468 let exported = storage.export_buffers().unwrap();
1469 assert_eq!(exported, "First\n\nSecond");
1470 }
1471
1472 fn setup_buffer_with_chunk(storage: &mut SqliteStorage) -> (i64, i64) {
1474 let buffer = Buffer::from_content("test content".to_string());
1475 let buffer_id = storage.add_buffer(&buffer).unwrap();
1476 let chunks = vec![Chunk::new(buffer_id, "test content".to_string(), 0..12, 0)];
1477 storage.add_chunks(buffer_id, &chunks).unwrap();
1478 let chunk_id = storage.get_chunks(buffer_id).unwrap()[0].id.unwrap();
1479 (buffer_id, chunk_id)
1480 }
1481
1482 #[test]
1483 fn test_store_and_get_embedding() {
1484 let mut storage = setup();
1485 let (_buffer_id, chunk_id) = setup_buffer_with_chunk(&mut storage);
1486
1487 let embedding = vec![0.1_f32, 0.2, 0.3, 0.4];
1488 storage
1489 .store_embedding(chunk_id, &embedding, Some("test-model"))
1490 .unwrap();
1491
1492 let loaded = storage.get_embedding(chunk_id).unwrap().unwrap();
1493 assert_eq!(loaded.len(), 4);
1494 for (a, b) in loaded.iter().zip(embedding.iter()) {
1495 assert!((a - b).abs() < 1e-6, "expected {b}, got {a}");
1496 }
1497 }
1498
1499 #[test]
1500 fn test_get_embedding_nonexistent() {
1501 let storage = setup();
1502 let result = storage.get_embedding(9999).unwrap();
1503 assert!(result.is_none());
1504 }
1505
1506 #[test]
1507 fn test_store_embedding_upsert() {
1508 let mut storage = setup();
1509 let (_buffer_id, chunk_id) = setup_buffer_with_chunk(&mut storage);
1510
1511 let embedding1 = vec![0.1_f32, 0.2];
1512 storage
1513 .store_embedding(chunk_id, &embedding1, Some("model-a"))
1514 .unwrap();
1515
1516 let embedding2 = vec![0.9_f32, 0.8];
1518 storage
1519 .store_embedding(chunk_id, &embedding2, Some("model-a"))
1520 .unwrap();
1521
1522 let loaded = storage.get_embedding(chunk_id).unwrap().unwrap();
1523 for (a, b) in loaded.iter().zip(embedding2.iter()) {
1524 assert!((a - b).abs() < 1e-6);
1525 }
1526 }
1527
1528 #[test]
1529 fn test_has_embedding() {
1530 let mut storage = setup();
1531 let (_buffer_id, chunk_id) = setup_buffer_with_chunk(&mut storage);
1532
1533 assert!(!storage.has_embedding(chunk_id).unwrap());
1534
1535 storage.store_embedding(chunk_id, &[0.1_f32], None).unwrap();
1536
1537 assert!(storage.has_embedding(chunk_id).unwrap());
1538 }
1539
1540 #[test]
1541 fn test_embedding_count() {
1542 let mut storage = setup();
1543 assert_eq!(storage.embedding_count().unwrap(), 0);
1544
1545 let (buffer_id, chunk_id) = setup_buffer_with_chunk(&mut storage);
1546
1547 let chunks2 = vec![Chunk::new(buffer_id, "more".to_string(), 0..4, 1)];
1549 storage.add_chunks(buffer_id, &chunks2).unwrap();
1550 let chunk_id2 = storage
1551 .get_chunks(buffer_id)
1552 .unwrap()
1553 .into_iter()
1554 .find(|c| c.id != Some(chunk_id))
1555 .unwrap()
1556 .id
1557 .unwrap();
1558
1559 storage.store_embedding(chunk_id, &[0.1_f32], None).unwrap();
1560 assert_eq!(storage.embedding_count().unwrap(), 1);
1561
1562 storage
1563 .store_embedding(chunk_id2, &[0.2_f32], None)
1564 .unwrap();
1565 assert_eq!(storage.embedding_count().unwrap(), 2);
1566 }
1567
1568 #[test]
1569 fn test_delete_embedding() {
1570 let mut storage = setup();
1571 let (_buffer_id, chunk_id) = setup_buffer_with_chunk(&mut storage);
1572
1573 storage.store_embedding(chunk_id, &[0.1_f32], None).unwrap();
1574 assert!(storage.has_embedding(chunk_id).unwrap());
1575
1576 storage.delete_embedding(chunk_id).unwrap();
1577 assert!(!storage.has_embedding(chunk_id).unwrap());
1578 }
1579
1580 #[test]
1581 fn test_store_embeddings_batch() {
1582 let mut storage = setup();
1583 let (buffer_id, chunk_id1) = setup_buffer_with_chunk(&mut storage);
1584
1585 let chunks2 = vec![Chunk::new(buffer_id, "second".to_string(), 0..6, 1)];
1587 storage.add_chunks(buffer_id, &chunks2).unwrap();
1588 let chunk_id2 = storage
1589 .get_chunks(buffer_id)
1590 .unwrap()
1591 .into_iter()
1592 .find(|c| c.id != Some(chunk_id1))
1593 .unwrap()
1594 .id
1595 .unwrap();
1596
1597 let batch = vec![
1598 (chunk_id1, vec![0.1_f32, 0.2]),
1599 (chunk_id2, vec![0.3_f32, 0.4]),
1600 ];
1601 storage
1602 .store_embeddings_batch(&batch, Some("batch-model"))
1603 .unwrap();
1604
1605 assert!(storage.has_embedding(chunk_id1).unwrap());
1606 assert!(storage.has_embedding(chunk_id2).unwrap());
1607 assert_eq!(storage.embedding_count().unwrap(), 2);
1608 }
1609
1610 #[test]
1611 fn test_get_all_embeddings() {
1612 let mut storage = setup();
1613 let (buffer_id, chunk_id1) = setup_buffer_with_chunk(&mut storage);
1614
1615 let chunks2 = vec![Chunk::new(buffer_id, "second".to_string(), 0..6, 1)];
1616 storage.add_chunks(buffer_id, &chunks2).unwrap();
1617 let chunk_id2 = storage
1618 .get_chunks(buffer_id)
1619 .unwrap()
1620 .into_iter()
1621 .find(|c| c.id != Some(chunk_id1))
1622 .unwrap()
1623 .id
1624 .unwrap();
1625
1626 storage
1627 .store_embedding(chunk_id1, &[0.1_f32], Some("m"))
1628 .unwrap();
1629 storage
1630 .store_embedding(chunk_id2, &[0.2_f32], Some("m"))
1631 .unwrap();
1632 let (other_buffer_id, other_chunk_id) = setup_buffer_with_chunk(&mut storage);
1633 storage
1634 .store_embedding(other_chunk_id, &[0.3_f32], Some("m"))
1635 .unwrap();
1636
1637 let all = storage.get_all_embeddings().unwrap();
1638 assert_eq!(all.len(), 3);
1639
1640 let scoped = storage.get_embeddings_in_buffer(buffer_id).unwrap();
1641 assert_eq!(scoped.len(), 2);
1642 assert!(scoped.iter().any(|(chunk_id, _)| *chunk_id == chunk_id1));
1643 assert!(scoped.iter().any(|(chunk_id, _)| *chunk_id == chunk_id2));
1644 assert!(
1645 !scoped
1646 .iter()
1647 .any(|(chunk_id, _)| *chunk_id == other_chunk_id)
1648 );
1649 assert!(
1650 storage
1651 .get_embeddings_in_buffer(other_buffer_id)
1652 .unwrap()
1653 .iter()
1654 .any(|(chunk_id, _)| *chunk_id == other_chunk_id)
1655 );
1656 }
1657
1658 #[test]
1659 fn test_get_embedding_models() {
1660 let mut storage = setup();
1661 let (buffer_id, chunk_id) = setup_buffer_with_chunk(&mut storage);
1662
1663 assert!(storage.get_embedding_models(buffer_id).unwrap().is_empty());
1665
1666 storage
1667 .store_embedding(chunk_id, &[0.1_f32], Some("model-x"))
1668 .unwrap();
1669
1670 let models = storage.get_embedding_models(buffer_id).unwrap();
1671 assert_eq!(models.len(), 1);
1672 assert_eq!(models[0], "model-x");
1673 }
1674
1675 #[test]
1676 fn test_get_embedding_model_counts() {
1677 let mut storage = setup();
1678 let (buffer_id, chunk_id1) = setup_buffer_with_chunk(&mut storage);
1679
1680 let chunks2 = vec![Chunk::new(buffer_id, "extra".to_string(), 0..5, 1)];
1681 storage.add_chunks(buffer_id, &chunks2).unwrap();
1682 let chunk_id2 = storage
1683 .get_chunks(buffer_id)
1684 .unwrap()
1685 .into_iter()
1686 .find(|c| c.id != Some(chunk_id1))
1687 .unwrap()
1688 .id
1689 .unwrap();
1690
1691 storage
1692 .store_embedding(chunk_id1, &[0.1_f32], Some("model-a"))
1693 .unwrap();
1694 storage
1695 .store_embedding(chunk_id2, &[0.2_f32], Some("model-a"))
1696 .unwrap();
1697
1698 let counts = storage.get_embedding_model_counts(buffer_id).unwrap();
1699 assert_eq!(counts.len(), 1);
1700 assert_eq!(counts[0], (Some("model-a".to_string()), 2));
1701 }
1702
1703 #[test]
1704 fn test_get_chunks_needing_embedding_no_embeddings() {
1705 let mut storage = setup();
1706 let (buffer_id, _chunk_id) = setup_buffer_with_chunk(&mut storage);
1707
1708 let needing = storage
1709 .get_chunks_needing_embedding(buffer_id, None)
1710 .unwrap();
1711 assert_eq!(needing.len(), 1);
1712 }
1713
1714 #[test]
1715 fn test_get_chunks_needing_embedding_with_model() {
1716 let mut storage = setup();
1717 let (buffer_id, chunk_id) = setup_buffer_with_chunk(&mut storage);
1718
1719 let chunks2 = vec![Chunk::new(buffer_id, "b".to_string(), 0..1, 1)];
1721 storage.add_chunks(buffer_id, &chunks2).unwrap();
1722 let chunk_id2 = storage
1723 .get_chunks(buffer_id)
1724 .unwrap()
1725 .into_iter()
1726 .find(|c| c.id != Some(chunk_id))
1727 .unwrap()
1728 .id
1729 .unwrap();
1730
1731 storage
1733 .store_embedding(chunk_id, &[0.1_f32], Some("model-a"))
1734 .unwrap();
1735
1736 let needing = storage
1738 .get_chunks_needing_embedding(buffer_id, Some("model-a"))
1739 .unwrap();
1740 assert!(needing.contains(&chunk_id2));
1741 assert!(!needing.contains(&chunk_id));
1742
1743 let needing_b = storage
1745 .get_chunks_needing_embedding(buffer_id, Some("model-b"))
1746 .unwrap();
1747 assert!(needing_b.contains(&chunk_id));
1748 assert!(needing_b.contains(&chunk_id2));
1749 }
1750
1751 #[test]
1752 fn test_get_chunks_without_embedding() {
1753 let mut storage = setup();
1754 let (buffer_id, chunk_id) = setup_buffer_with_chunk(&mut storage);
1755
1756 let without = storage.get_chunks_without_embedding(buffer_id).unwrap();
1757 assert_eq!(without.len(), 1);
1758 assert!(without.contains(&chunk_id));
1759
1760 storage.store_embedding(chunk_id, &[0.1_f32], None).unwrap();
1761
1762 let without_after = storage.get_chunks_without_embedding(buffer_id).unwrap();
1763 assert!(without_after.is_empty());
1764 }
1765
1766 #[test]
1767 fn test_delete_embeddings_by_model_named() {
1768 let mut storage = setup();
1769 let (buffer_id, chunk_id1) = setup_buffer_with_chunk(&mut storage);
1770
1771 let chunks2 = vec![Chunk::new(buffer_id, "b".to_string(), 0..1, 1)];
1772 storage.add_chunks(buffer_id, &chunks2).unwrap();
1773 let chunk_id2 = storage
1774 .get_chunks(buffer_id)
1775 .unwrap()
1776 .into_iter()
1777 .find(|c| c.id != Some(chunk_id1))
1778 .unwrap()
1779 .id
1780 .unwrap();
1781
1782 storage
1783 .store_embedding(chunk_id1, &[0.1_f32], Some("model-a"))
1784 .unwrap();
1785 storage
1786 .store_embedding(chunk_id2, &[0.2_f32], Some("model-b"))
1787 .unwrap();
1788
1789 let deleted = storage
1790 .delete_embeddings_by_model(buffer_id, Some("model-a"))
1791 .unwrap();
1792 assert_eq!(deleted, 1);
1793 assert!(!storage.has_embedding(chunk_id1).unwrap());
1794 assert!(storage.has_embedding(chunk_id2).unwrap());
1795 }
1796
1797 #[test]
1798 fn test_delete_embeddings_by_model_null() {
1799 let mut storage = setup();
1800 let (buffer_id, chunk_id1) = setup_buffer_with_chunk(&mut storage);
1801
1802 let chunks2 = vec![Chunk::new(buffer_id, "b".to_string(), 0..1, 1)];
1803 storage.add_chunks(buffer_id, &chunks2).unwrap();
1804 let chunk_id2 = storage
1805 .get_chunks(buffer_id)
1806 .unwrap()
1807 .into_iter()
1808 .find(|c| c.id != Some(chunk_id1))
1809 .unwrap()
1810 .id
1811 .unwrap();
1812
1813 storage
1814 .store_embedding(chunk_id1, &[0.1_f32], None)
1815 .unwrap();
1816 storage
1817 .store_embedding(chunk_id2, &[0.2_f32], Some("model-b"))
1818 .unwrap();
1819
1820 let deleted = storage.delete_embeddings_by_model(buffer_id, None).unwrap();
1822 assert_eq!(deleted, 1);
1823 assert!(!storage.has_embedding(chunk_id1).unwrap());
1824 assert!(storage.has_embedding(chunk_id2).unwrap());
1825 }
1826
1827 #[test]
1828 fn test_get_embedding_stats() {
1829 let mut storage = setup();
1830 let (buffer_id, chunk_id) = setup_buffer_with_chunk(&mut storage);
1831
1832 let chunks2 = vec![Chunk::new(buffer_id, "extra".to_string(), 0..5, 1)];
1834 storage.add_chunks(buffer_id, &chunks2).unwrap();
1835
1836 let stats = storage.get_embedding_stats(buffer_id).unwrap();
1838 assert_eq!(stats.total_chunks, 2);
1839 assert_eq!(stats.embedded_chunks, 0);
1840
1841 storage
1843 .store_embedding(chunk_id, &[0.1_f32], Some("m1"))
1844 .unwrap();
1845
1846 let stats = storage.get_embedding_stats(buffer_id).unwrap();
1847 assert_eq!(stats.total_chunks, 2);
1848 assert_eq!(stats.embedded_chunks, 1);
1849 assert_eq!(stats.model_counts.len(), 1);
1850 }
1851
1852 #[test]
1853 fn test_search_fts_finds_match() {
1854 let mut storage = setup();
1855 let buffer = Buffer::from_content("The quick brown fox jumps".to_string());
1856 let buffer_id = storage.add_buffer(&buffer).unwrap();
1857
1858 let chunks = vec![
1859 Chunk::new(buffer_id, "The quick brown fox".to_string(), 0..19, 0),
1860 Chunk::new(buffer_id, "fox jumps".to_string(), 20..29, 1),
1861 ];
1862 storage.add_chunks(buffer_id, &chunks).unwrap();
1863
1864 let results = storage.search_fts("quick", 10).unwrap();
1865 assert!(!results.is_empty());
1867 for (_id, score) in &results {
1869 assert!(
1870 *score >= 0.0,
1871 "BM25 score should be non-negative, got {score}"
1872 );
1873 }
1874 }
1875
1876 #[test]
1877 fn test_search_fts_no_match() {
1878 let mut storage = setup();
1879 let buffer = Buffer::from_content("The quick brown fox".to_string());
1880 let buffer_id = storage.add_buffer(&buffer).unwrap();
1881
1882 let chunks = vec![Chunk::new(
1883 buffer_id,
1884 "The quick brown fox".to_string(),
1885 0..19,
1886 0,
1887 )];
1888 storage.add_chunks(buffer_id, &chunks).unwrap();
1889
1890 let results = storage.search_fts("zzzyyyxxx", 10).unwrap();
1891 assert!(results.is_empty());
1892 }
1893
1894 #[test]
1895 fn test_search_fts_respects_limit() {
1896 let mut storage = setup();
1897 let buffer = Buffer::from_content("hello world hello world hello".to_string());
1898 let buffer_id = storage.add_buffer(&buffer).unwrap();
1899
1900 let chunks = vec![
1902 Chunk::new(buffer_id, "hello world".to_string(), 0..11, 0),
1903 Chunk::new(buffer_id, "hello world".to_string(), 12..23, 1),
1904 Chunk::new(buffer_id, "hello".to_string(), 24..29, 2),
1905 ];
1906 storage.add_chunks(buffer_id, &chunks).unwrap();
1907
1908 let results = storage.search_fts("hello", 2).unwrap();
1909 assert!(results.len() <= 2);
1910 }
1911
1912 #[test]
1913 fn test_get_chunks_by_ids_empty() {
1914 let mut storage = SqliteStorage::in_memory().unwrap();
1915 storage.init().unwrap();
1916
1917 let result = storage.get_chunks_by_ids(&[]).unwrap();
1918 assert!(result.is_empty());
1919 }
1920
1921 #[test]
1922 fn test_get_chunks_by_ids_batch() {
1923 let mut storage = SqliteStorage::in_memory().unwrap();
1924 storage.init().unwrap();
1925
1926 let buffer_id = storage
1927 .add_buffer(&Buffer::from_content("abc def ghi".to_string()))
1928 .unwrap();
1929 let chunks = vec![
1930 Chunk::new(buffer_id, "abc".to_string(), 0..3, 0),
1931 Chunk::new(buffer_id, "def".to_string(), 4..7, 1),
1932 Chunk::new(buffer_id, "ghi".to_string(), 8..11, 2),
1933 ];
1934 storage.add_chunks(buffer_id, &chunks).unwrap();
1935
1936 let all = storage.get_chunks(buffer_id).unwrap();
1938 let ids: Vec<i64> = all.iter().filter_map(|c| c.id).take(2).collect();
1939 assert_eq!(ids.len(), 2);
1940
1941 let map = storage.get_chunks_by_ids(&ids).unwrap();
1942 assert_eq!(map.len(), 2);
1943 for id in &ids {
1944 assert!(map.contains_key(id));
1945 }
1946 }
1947
1948 #[test]
1949 fn test_get_chunks_by_ids_missing_id() {
1950 let mut storage = SqliteStorage::in_memory().unwrap();
1951 storage.init().unwrap();
1952
1953 let map = storage.get_chunks_by_ids(&[99999]).unwrap();
1955 assert!(map.is_empty());
1956 }
1957
1958 #[test]
1959 fn test_fts_and_embedding_pipeline() {
1960 let mut storage = setup();
1963 let buffer = Buffer::from_content("machine learning and neural networks".to_string());
1964 let buffer_id = storage.add_buffer(&buffer).unwrap();
1965
1966 let chunks = vec![
1967 Chunk::new(
1968 buffer_id,
1969 "machine learning algorithms".to_string(),
1970 0..27,
1971 0,
1972 ),
1973 Chunk::new(
1974 buffer_id,
1975 "neural networks architecture".to_string(),
1976 28..56,
1977 1,
1978 ),
1979 ];
1980 storage.add_chunks(buffer_id, &chunks).unwrap();
1981
1982 let chunk_ids: Vec<i64> = storage
1984 .get_chunks(buffer_id)
1985 .unwrap()
1986 .into_iter()
1987 .map(|c| c.id.unwrap())
1988 .collect();
1989 assert_eq!(chunk_ids.len(), 2);
1990
1991 let embeddings: &[&[f32]] = &[&[0.1, 0.2], &[0.2, 0.4]];
1992 for (&chunk_id, embedding) in chunk_ids.iter().zip(embeddings) {
1993 storage
1994 .store_embedding(chunk_id, embedding, Some("test-model"))
1995 .unwrap();
1996 }
1997
1998 let fts_results = storage.search_fts("machine", 10).unwrap();
2000 assert!(!fts_results.is_empty());
2001
2002 for (chunk_id, _score) in &fts_results {
2004 assert!(
2005 storage.has_embedding(*chunk_id).unwrap(),
2006 "FTS result chunk {chunk_id} should have an embedding"
2007 );
2008 }
2009
2010 assert_eq!(storage.embedding_count().unwrap(), 2);
2012 }
2013}