Skip to main content

docling_rag/store/
sqlite.rs

1//! SQLite vector store (feature `sqlite`, on by default).
2//!
3//! Uses the bundled SQLite that ships with `sqlx`, plus the statically-compiled
4//! [`sqlite-vec`](https://github.com/asg017/sqlite-vec) extension: embeddings live
5//! in a `vec0` virtual table (`chunks_vec`, cosine metric) keyed by the `chunks`
6//! rowid, and `vector_search` is a real KNN `MATCH` query instead of a full-table
7//! scan. The extension is registered process-wide with `sqlite3_auto_extension`
8//! before the first connection, so every pooled connection sees `vec0`.
9
10use super::VectorStore;
11use crate::math;
12use crate::model::{Chunk, Document, Scored};
13use crate::{RagError, Result};
14use async_trait::async_trait;
15use sqlx::sqlite::{SqliteConnectOptions, SqlitePool, SqlitePoolOptions};
16use sqlx::Row;
17use std::str::FromStr;
18use std::sync::Once;
19
20/// Register sqlite-vec for every SQLite connection opened by this process.
21/// Safe to call repeatedly; the registration itself happens once.
22fn register_sqlite_vec() {
23    static REGISTER: Once = Once::new();
24    REGISTER.call_once(|| unsafe {
25        // sqlite3_auto_extension expects `fn()`; sqlite3_vec_init's real signature
26        // (db, err, api) is what SQLite actually calls it with.
27        #[allow(clippy::missing_transmute_annotations)]
28        libsqlite3_sys::sqlite3_auto_extension(Some(std::mem::transmute(
29            sqlite_vec::sqlite3_vec_init as *const (),
30        )));
31    });
32}
33
34/// SQLite-backed store with sqlite-vec KNN search.
35pub struct SqliteStore {
36    pool: SqlitePool,
37    dim: usize,
38}
39
40impl SqliteStore {
41    /// Connect to (creating if missing) the SQLite database at `url`
42    /// (e.g. `sqlite://data/rag.db` or `sqlite::memory:`), expecting
43    /// `dim`-dimensional embeddings.
44    pub async fn connect(url: &str, dim: usize) -> Result<Self> {
45        register_sqlite_vec();
46        let opts = SqliteConnectOptions::from_str(url)
47            .map_err(|e| RagError::config(format!("invalid RAG_DATABASE_URL '{url}': {e}")))?
48            .create_if_missing(true);
49        // SQLite won't create parent directories for the DB file; do it ourselves.
50        let filename = opts.get_filename();
51        if let Some(parent) = filename.parent() {
52            if !parent.as_os_str().is_empty() {
53                std::fs::create_dir_all(parent).ok();
54            }
55        }
56        let pool = SqlitePoolOptions::new()
57            .max_connections(4)
58            .connect_with(opts)
59            .await?;
60        Ok(SqliteStore { pool, dim })
61    }
62}
63
64fn row_to_chunk(row: &sqlx::sqlite::SqliteRow) -> Result<Chunk> {
65    let metadata: String = row.try_get("metadata")?;
66    Ok(Chunk {
67        id: row.try_get("id")?,
68        doc_id: row.try_get("doc_id")?,
69        ordinal: row.try_get("ordinal")?,
70        text: row.try_get("text")?,
71        token_count: row.try_get("token_count")?,
72        metadata: serde_json::from_str(&metadata).unwrap_or(serde_json::Value::Null),
73        embedding: None,
74    })
75}
76
77#[async_trait]
78impl VectorStore for SqliteStore {
79    async fn migrate(&self) -> Result<()> {
80        sqlx::query(
81            "CREATE TABLE IF NOT EXISTS documents (
82                id TEXT PRIMARY KEY,
83                source_uri TEXT NOT NULL,
84                title TEXT NOT NULL,
85                hash TEXT NOT NULL,
86                metadata TEXT NOT NULL DEFAULT 'null',
87                created_at TEXT NOT NULL
88            )",
89        )
90        .execute(&self.pool)
91        .await?;
92        sqlx::query("CREATE INDEX IF NOT EXISTS idx_documents_hash ON documents(hash)")
93            .execute(&self.pool)
94            .await?;
95        sqlx::query(
96            "CREATE TABLE IF NOT EXISTS chunks (
97                id TEXT PRIMARY KEY,
98                doc_id TEXT NOT NULL REFERENCES documents(id) ON DELETE CASCADE,
99                ordinal INTEGER NOT NULL,
100                text TEXT NOT NULL,
101                token_count INTEGER NOT NULL,
102                metadata TEXT NOT NULL DEFAULT 'null'
103            )",
104        )
105        .execute(&self.pool)
106        .await?;
107        sqlx::query("CREATE INDEX IF NOT EXISTS idx_chunks_doc ON chunks(doc_id)")
108            .execute(&self.pool)
109            .await?;
110        // The vec0 virtual table holds one embedding per chunk, sharing the chunks
111        // table's rowid. Dimension and metric are fixed at creation time.
112        sqlx::query(&format!(
113            "CREATE VIRTUAL TABLE IF NOT EXISTS chunks_vec USING vec0(
114                embedding float[{}] distance_metric=cosine
115            )",
116            self.dim
117        ))
118        .execute(&self.pool)
119        .await?;
120        Ok(())
121    }
122
123    async fn upsert_document(&self, doc: &Document) -> Result<()> {
124        sqlx::query(
125            "INSERT INTO documents (id, source_uri, title, hash, metadata, created_at)
126             VALUES (?1, ?2, ?3, ?4, ?5, ?6)
127             ON CONFLICT(id) DO UPDATE SET
128                source_uri = excluded.source_uri,
129                title = excluded.title,
130                hash = excluded.hash,
131                metadata = excluded.metadata",
132        )
133        .bind(&doc.id)
134        .bind(&doc.source_uri)
135        .bind(&doc.title)
136        .bind(&doc.hash)
137        .bind(serde_json::to_string(&doc.metadata)?)
138        .bind(&doc.created_at)
139        .execute(&self.pool)
140        .await?;
141        Ok(())
142    }
143
144    async fn find_document_by_hash(&self, hash: &str) -> Result<Option<String>> {
145        let row = sqlx::query("SELECT id FROM documents WHERE hash = ?1 LIMIT 1")
146            .bind(hash)
147            .fetch_optional(&self.pool)
148            .await?;
149        Ok(row.map(|r| r.get::<String, _>("id")))
150    }
151
152    async fn insert_chunks(&self, chunks: &[Chunk]) -> Result<()> {
153        let mut tx = self.pool.begin().await?;
154        for c in chunks {
155            let emb = c
156                .embedding
157                .as_ref()
158                .ok_or_else(|| RagError::Store(format!("chunk {} has no embedding", c.id)))?;
159            if emb.len() != self.dim {
160                return Err(RagError::Store(format!(
161                    "chunk {} embedding has dim {}, store expects {}",
162                    c.id,
163                    emb.len(),
164                    self.dim
165                )));
166            }
167            let rowid: i64 = sqlx::query_scalar(
168                "INSERT INTO chunks (id, doc_id, ordinal, text, token_count, metadata)
169                 VALUES (?1, ?2, ?3, ?4, ?5, ?6) RETURNING rowid",
170            )
171            .bind(&c.id)
172            .bind(&c.doc_id)
173            .bind(c.ordinal)
174            .bind(&c.text)
175            .bind(c.token_count)
176            .bind(serde_json::to_string(&c.metadata)?)
177            .fetch_one(&mut *tx)
178            .await?;
179            // vec0 accepts a raw little-endian f32 blob as the vector value.
180            sqlx::query("INSERT INTO chunks_vec (rowid, embedding) VALUES (?1, ?2)")
181                .bind(rowid)
182                .bind(math::to_bytes(emb))
183                .execute(&mut *tx)
184                .await?;
185        }
186        tx.commit().await?;
187        Ok(())
188    }
189
190    async fn vector_search(&self, query: &[f32], k: usize) -> Result<Vec<Scored>> {
191        if query.len() != self.dim {
192            return Err(RagError::Store(format!(
193                "query embedding has dim {}, store expects {}",
194                query.len(),
195                self.dim
196            )));
197        }
198        // KNN via the vec0 MATCH operator; distance is cosine distance in [0, 2],
199        // so similarity = 1 - distance.
200        let rows = sqlx::query(
201            "SELECT c.id, c.doc_id, c.ordinal, c.text, c.token_count, c.metadata, v.distance
202             FROM (SELECT rowid, distance FROM chunks_vec
203                   WHERE embedding MATCH ?1 AND k = ?2) v
204             JOIN chunks c ON c.rowid = v.rowid
205             ORDER BY v.distance",
206        )
207        .bind(math::to_bytes(query))
208        .bind(k as i64)
209        .fetch_all(&self.pool)
210        .await?;
211        let mut out = Vec::with_capacity(rows.len());
212        for row in &rows {
213            let distance: f64 = row.try_get("distance")?;
214            out.push(Scored::new(row_to_chunk(row)?, 1.0 - distance as f32));
215        }
216        Ok(out)
217    }
218
219    async fn all_chunks(&self) -> Result<Vec<Chunk>> {
220        let rows =
221            sqlx::query("SELECT id, doc_id, ordinal, text, token_count, metadata FROM chunks")
222                .fetch_all(&self.pool)
223                .await?;
224        rows.iter().map(row_to_chunk).collect()
225    }
226
227    async fn count_chunks(&self) -> Result<usize> {
228        let row = sqlx::query("SELECT COUNT(*) AS n FROM chunks")
229            .fetch_one(&self.pool)
230            .await?;
231        Ok(row.get::<i64, _>("n") as usize)
232    }
233
234    async fn count_documents(&self) -> Result<usize> {
235        let row = sqlx::query("SELECT COUNT(*) AS n FROM documents")
236            .fetch_one(&self.pool)
237            .await?;
238        Ok(row.get::<i64, _>("n") as usize)
239    }
240
241    async fn list_documents(&self) -> Result<Vec<Document>> {
242        let rows =
243            sqlx::query("SELECT id, source_uri, title, hash, metadata, created_at FROM documents")
244                .fetch_all(&self.pool)
245                .await?;
246        rows.iter()
247            .map(|row| {
248                let metadata: String = row.try_get("metadata")?;
249                Ok(Document {
250                    id: row.try_get("id")?,
251                    source_uri: row.try_get("source_uri")?,
252                    title: row.try_get("title")?,
253                    hash: row.try_get("hash")?,
254                    metadata: serde_json::from_str(&metadata).unwrap_or(serde_json::Value::Null),
255                    created_at: row.try_get("created_at")?,
256                })
257            })
258            .collect()
259    }
260
261    async fn delete_document(&self, doc_id: &str) -> Result<()> {
262        let mut tx = self.pool.begin().await?;
263        // chunks_vec shares the chunks rowids; delete its rows first.
264        sqlx::query(
265            "DELETE FROM chunks_vec WHERE rowid IN (SELECT rowid FROM chunks WHERE doc_id = ?1)",
266        )
267        .bind(doc_id)
268        .execute(&mut *tx)
269        .await?;
270        sqlx::query("DELETE FROM chunks WHERE doc_id = ?1")
271            .bind(doc_id)
272            .execute(&mut *tx)
273            .await?;
274        sqlx::query("DELETE FROM documents WHERE id = ?1")
275            .bind(doc_id)
276            .execute(&mut *tx)
277            .await?;
278        tx.commit().await?;
279        Ok(())
280    }
281
282    async fn delete_documents_by_source(&self, source_uri: &str) -> Result<()> {
283        let mut tx = self.pool.begin().await?;
284        sqlx::query(
285            "DELETE FROM chunks_vec WHERE rowid IN (
286                SELECT c.rowid FROM chunks c
287                JOIN documents d ON c.doc_id = d.id
288                WHERE d.source_uri = ?1
289            )",
290        )
291        .bind(source_uri)
292        .execute(&mut *tx)
293        .await?;
294        sqlx::query(
295            "DELETE FROM chunks WHERE doc_id IN (SELECT id FROM documents WHERE source_uri = ?1)",
296        )
297        .bind(source_uri)
298        .execute(&mut *tx)
299        .await?;
300        sqlx::query("DELETE FROM documents WHERE source_uri = ?1")
301            .bind(source_uri)
302            .execute(&mut *tx)
303            .await?;
304        tx.commit().await?;
305        Ok(())
306    }
307
308    async fn clear(&self) -> Result<()> {
309        sqlx::query("DELETE FROM chunks_vec")
310            .execute(&self.pool)
311            .await?;
312        sqlx::query("DELETE FROM chunks")
313            .execute(&self.pool)
314            .await?;
315        sqlx::query("DELETE FROM documents")
316            .execute(&self.pool)
317            .await?;
318        Ok(())
319    }
320}
321
322#[cfg(test)]
323mod tests {
324    use super::*;
325    use crate::embed::{Embedder, HashEmbedder};
326    use crate::model::Document;
327
328    #[tokio::test]
329    async fn knn_search_via_sqlite_vec() {
330        let store = SqliteStore::connect("sqlite::memory:", 64).await.unwrap();
331        store.migrate().await.unwrap();
332
333        let embedder = HashEmbedder::new(64);
334        let doc = Document::new("mem://t", "T", "h1");
335        store.upsert_document(&doc).await.unwrap();
336
337        let texts = [
338            "vector database semantic search",
339            "banana smoothie recipe",
340            "tokio async runtime",
341        ];
342        let mut chunks = Vec::new();
343        for (i, t) in texts.iter().enumerate() {
344            let mut c = Chunk::new(&doc.id, i as i64, *t, 0);
345            c.embedding = Some(embedder.embed_one(t).await.unwrap());
346            chunks.push(c);
347        }
348        store.insert_chunks(&chunks).await.unwrap();
349        assert_eq!(store.count_chunks().await.unwrap(), 3);
350
351        let q = embedder
352            .embed_one("semantic search in a vector database")
353            .await
354            .unwrap();
355        let hits = store.vector_search(&q, 2).await.unwrap();
356        assert_eq!(hits.len(), 2);
357        assert!(
358            hits[0].chunk.text.contains("vector database"),
359            "got: {}",
360            hits[0].chunk.text
361        );
362        assert!(hits[0].score >= hits[1].score);
363
364        // Dedup lookup and clear.
365        assert_eq!(
366            store.find_document_by_hash("h1").await.unwrap(),
367            Some(doc.id.clone())
368        );
369        store.clear().await.unwrap();
370        assert_eq!(store.count_chunks().await.unwrap(), 0);
371        assert!(store.vector_search(&q, 2).await.unwrap().is_empty());
372    }
373
374    #[tokio::test]
375    async fn rejects_wrong_dimension() {
376        let store = SqliteStore::connect("sqlite::memory:", 8).await.unwrap();
377        store.migrate().await.unwrap();
378        assert!(store.vector_search(&[0.0; 4], 1).await.is_err());
379    }
380}