Skip to main content

rlm_rs/storage/
sqlite.rs

1//! `SQLite` storage implementation.
2//!
3//! Provides persistent storage using `SQLite` with proper transaction
4//! management and migration support.
5
6// SQLite stores all integers as i64. These casts are intentional and safe
7// because we only store non-negative values that fit in usize.
8#![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
20/// SQLite-based storage implementation.
21///
22/// Provides persistent storage for RLM state with full ACID guarantees.
23///
24/// # Examples
25///
26/// ```no_run
27/// use rlm_rs::storage::{SqliteStorage, Storage};
28///
29/// let mut storage = SqliteStorage::open("rlm-state.db").unwrap();
30/// storage.init().unwrap();
31/// ```
32pub struct SqliteStorage {
33    /// `SQLite` connection.
34    conn: Connection,
35    /// Path to the database file (None for in-memory).
36    path: Option<PathBuf>,
37}
38
39impl SqliteStorage {
40    /// Opens or creates a `SQLite` database at the given path.
41    ///
42    /// # Arguments
43    ///
44    /// * `path` - Path to the database file. Parent directory must exist.
45    ///
46    /// # Errors
47    ///
48    /// Returns an error if the database cannot be opened or initialized.
49    pub fn open<P: AsRef<Path>>(path: P) -> Result<Self> {
50        let path = path.as_ref().to_path_buf();
51
52        // Ensure parent directory exists
53        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        // Enable foreign keys
62        conn.execute("PRAGMA foreign_keys = ON;", [])
63            .map_err(StorageError::from)?;
64
65        // Use WAL mode for better concurrent access (returns result, use query_row)
66        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    /// Creates an in-memory `SQLite` database.
77    ///
78    /// Useful for testing.
79    ///
80    /// # Errors
81    ///
82    /// Returns an error if the database cannot be created.
83    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    /// Returns the database path (None for in-memory).
92    #[must_use]
93    pub fn path(&self) -> Option<&Path> {
94        self.path.as_deref()
95    }
96
97    /// Gets the current schema version.
98    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    /// Sets the schema version.
109    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    /// Returns current Unix timestamp.
117    #[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        // Check if already initialized
128        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            // Fresh install - create schema
135            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            // Run migrations
143            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    // ==================== Context Operations ====================
179
180    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    // ==================== Buffer Operations ====================
223
224    #[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        // Chunks are deleted automatically via CASCADE
382        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    // ==================== Chunk Operations ====================
397
398    #[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        // Update chunk count on buffer
447        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        // Update chunk count on buffer
553        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    // ==================== Utility Operations ====================
576
577    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
626// ==================== Embedding & Search Operations ====================
627
628impl SqliteStorage {
629    /// Retrieves multiple chunks by their IDs in a single query.
630    ///
631    /// More efficient than calling [`Storage::get_chunk`] repeatedly when fetching
632    /// several chunks at once (e.g., populating search result previews).
633    ///
634    /// Returns a map from chunk ID to [`Chunk`].
635    ///
636    /// # Errors
637    ///
638    /// Returns an error if the query fails.
639    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    /// Stores an embedding for a chunk.
691    ///
692    /// # Arguments
693    ///
694    /// * `chunk_id` - The chunk ID to associate the embedding with.
695    /// * `embedding` - The embedding vector (f32 array).
696    /// * `model_name` - Optional name of the model that generated the embedding.
697    ///
698    /// # Errors
699    ///
700    /// Returns an error if the embedding cannot be stored.
701    #[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        // Serialize f32 array to bytes (little-endian)
711        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    /// Retrieves the embedding for a chunk.
730    ///
731    /// # Errors
732    ///
733    /// Returns an error if the query fails.
734    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    /// Gets the distinct model names used for embeddings in a buffer.
754    ///
755    /// Returns the set of model names used to generate embeddings for
756    /// chunks belonging to the specified buffer.
757    ///
758    /// # Errors
759    ///
760    /// Returns an error if the query fails.
761    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    /// Gets the count of embeddings by model name for a buffer.
784    ///
785    /// Returns a list of (`model_name`, count) pairs.
786    ///
787    /// # Errors
788    ///
789    /// Returns an error if the query fails.
790    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    /// Stores embeddings for multiple chunks in a batch.
816    ///
817    /// # Errors
818    ///
819    /// Returns an error if any embedding cannot be stored.
820    #[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    /// Deletes the embedding for a chunk.
859    ///
860    /// # Errors
861    ///
862    /// Returns an error if deletion fails.
863    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    /// Performs FTS5 BM25 full-text search.
874    ///
875    /// Returns chunk IDs and their BM25 scores (higher is better match).
876    ///
877    /// # Arguments
878    ///
879    /// * `query` - The search query (supports FTS5 query syntax).
880    /// * `limit` - Maximum number of results to return.
881    ///
882    /// # Errors
883    ///
884    /// Returns an error if the search fails.
885    #[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    /// Performs FTS5 BM25 full-text search, optionally restricted to a buffer.
891    ///
892    /// # Errors
893    ///
894    /// Returns an error if the search fails.
895    #[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        // FTS5 bm25() returns negative scores, more negative = better match
903        // We negate it so higher scores = better match
904
905        // Convert space-separated terms to OR query for more forgiving search
906        // Each term is quoted to escape FTS5 special characters (?, *, ^, etc.)
907        // "CLI tool?" becomes '"CLI" OR "tool?"' so special chars are treated as literals
908        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    /// Returns all chunk embeddings for vector similarity search.
962    ///
963    /// # Errors
964    ///
965    /// Returns an error if the query fails.
966    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    /// Returns chunk embeddings belonging to one buffer for vector similarity search.
990    ///
991    /// # Errors
992    ///
993    /// Returns an error if the query fails.
994    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    /// Counts chunks with embeddings.
1025    ///
1026    /// # Errors
1027    ///
1028    /// Returns an error if the count fails.
1029    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    /// Returns `true` if every chunk in the buffer has an embedding (or the buffer has no chunks).
1040    ///
1041    /// Uses a single `NOT EXISTS` query instead of per-chunk lookups, making it O(1) in terms
1042    /// of round-trips regardless of how many chunks the buffer contains.
1043    ///
1044    /// # Errors
1045    ///
1046    /// Returns an error if the query fails.
1047    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    /// Checks if a chunk has an embedding.
1066    ///
1067    /// # Errors
1068    ///
1069    /// Returns an error if the query fails.
1070    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    /// Gets chunk IDs that need embedding (either no embedding or wrong model).
1083    ///
1084    /// This is used for incremental embedding updates. Returns chunks that:
1085    /// - Have no embedding at all, OR
1086    /// - Have an embedding with a different model name (if `current_model` is provided)
1087    ///
1088    /// # Arguments
1089    ///
1090    /// * `buffer_id` - The buffer to check.
1091    /// * `current_model` - Optional model name to check against. If provided,
1092    ///   chunks with different models are included.
1093    ///
1094    /// # Errors
1095    ///
1096    /// Returns an error if the query fails.
1097    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        // Get chunks without any embedding
1105        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 model specified, also get chunks with different model
1125        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        // Deduplicate (in case of overlap, though shouldn't happen)
1147        results.sort_unstable();
1148        results.dedup();
1149        Ok(results)
1150    }
1151
1152    /// Gets chunks without any embedding for a buffer.
1153    ///
1154    /// Simpler version of `get_chunks_needing_embedding` when model doesn't matter.
1155    ///
1156    /// # Errors
1157    ///
1158    /// Returns an error if the query fails.
1159    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    /// Deletes embeddings with a specific model name.
1164    ///
1165    /// Useful for cleaning up embeddings from old models before re-embedding.
1166    ///
1167    /// # Arguments
1168    ///
1169    /// * `buffer_id` - The buffer to clean.
1170    /// * `model_name` - The model name to match (or None to match NULL).
1171    ///
1172    /// # Returns
1173    ///
1174    /// The number of embeddings deleted.
1175    ///
1176    /// # Errors
1177    ///
1178    /// Returns an error if deletion fails.
1179    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    /// Gets embedding statistics for a buffer.
1214    ///
1215    /// Returns counts of embedded vs total chunks, and model breakdown.
1216    ///
1217    /// # Errors
1218    ///
1219    /// Returns an error if the query fails.
1220    pub fn get_embedding_stats(&self, buffer_id: i64) -> Result<EmbeddingStats> {
1221        // Total chunks
1222        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        // Embedded chunks
1232        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        // Model counts
1246        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/// Statistics about embeddings for a buffer.
1257#[derive(Debug, Clone)]
1258pub struct EmbeddingStats {
1259    /// Total number of chunks in the buffer.
1260    pub total_chunks: usize,
1261    /// Number of chunks with embeddings.
1262    pub embedded_chunks: usize,
1263    /// Count of embeddings by model (`model_name`, count).
1264    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()); // Second init should be fine
1290    }
1291
1292    #[test]
1293    fn test_context_crud() {
1294        let mut storage = setup();
1295
1296        // No context initially
1297        assert!(storage.load_context().unwrap().is_none());
1298
1299        // Save context
1300        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        // Load context
1305        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        // Delete context
1312        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        // Add buffer
1321        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        // Get buffer
1326        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        // Get by name
1331        let by_name = storage.get_buffer_by_name("test").unwrap().unwrap();
1332        assert_eq!(by_name.id, Some(id));
1333
1334        // List buffers
1335        let buffers = storage.list_buffers().unwrap();
1336        assert_eq!(buffers.len(), 1);
1337
1338        // Update buffer
1339        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        // Delete buffer
1347        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        // Create buffer first
1356        let buffer = Buffer::from_content("Hello, world!".to_string());
1357        let buffer_id = storage.add_buffer(&buffer).unwrap();
1358
1359        // Add chunks
1360        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        // Get chunks
1367        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        // Chunk count
1373        assert_eq!(storage.chunk_count(buffer_id).unwrap(), 2);
1374
1375        // Get single chunk
1376        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        // Delete chunks
1381        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        // Create buffer with chunks
1390        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        // Verify chunk exists
1397        assert_eq!(storage.chunk_count(buffer_id).unwrap(), 1);
1398
1399        // Delete buffer - chunks should be deleted too
1400        storage.delete_buffer(buffer_id).unwrap();
1401
1402        // Verify no orphan chunks (query all chunks)
1403        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        // Add some data
1415        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        // Reset
1422        storage.reset().unwrap();
1423
1424        // Verify empty
1425        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        // Empty stats
1434        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        // Add data
1440        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        // Stats with data
1450        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    // Helper: create a buffer with one chunk and return (buffer_id, chunk_id).
1473    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        // Upsert with new values
1517        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        // Add a second chunk
1548        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        // Add second chunk
1586        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        // No models initially
1664        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        // Add second chunk
1720        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        // chunk_id has model-a, chunk_id2 has no embedding
1732        storage
1733            .store_embedding(chunk_id, &[0.1_f32], Some("model-a"))
1734            .unwrap();
1735
1736        // When checking for model-a: chunk_id2 needs one (no embedding)
1737        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        // When checking for model-b: chunk_id has wrong model, chunk_id2 has none
1744        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        // Delete only the NULL-model embedding
1821        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        // Add second chunk
1833        let chunks2 = vec![Chunk::new(buffer_id, "extra".to_string(), 0..5, 1)];
1834        storage.add_chunks(buffer_id, &chunks2).unwrap();
1835
1836        // Stats before embedding
1837        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        // Embed one chunk
1842        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        // At least the first chunk should match
1866        assert!(!results.is_empty());
1867        // Scores should be positive
1868        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        // Add 3 chunks all containing "hello"
1901        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        // Fetch by known IDs
1937        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        // Query for an ID that doesn't exist
1954        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        // Integration: index chunks, store embeddings, then verify that FTS
1961        // hits match chunks that also have embeddings.
1962        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        // Store embeddings for both chunks
1983        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        // FTS search should find the relevant chunk
1999        let fts_results = storage.search_fts("machine", 10).unwrap();
2000        assert!(!fts_results.is_empty());
2001
2002        // Every FTS result should also have an embedding stored
2003        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        // Verify embedding count matches what was stored
2011        assert_eq!(storage.embedding_count().unwrap(), 2);
2012    }
2013}