1use 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
20fn register_sqlite_vec() {
23 static REGISTER: Once = Once::new();
24 REGISTER.call_once(|| unsafe {
25 #[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
34pub struct SqliteStore {
36 pool: SqlitePool,
37 dim: usize,
38}
39
40impl SqliteStore {
41 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 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 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 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 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_chunks_for(&self, doc_id: &str) -> Result<usize> {
235 let row = sqlx::query("SELECT COUNT(*) AS n FROM chunks WHERE doc_id = ?")
236 .bind(doc_id)
237 .fetch_one(&self.pool)
238 .await?;
239 Ok(row.get::<i64, _>("n") as usize)
240 }
241
242 async fn chunk_neighborhood(&self, doc_id: &str, ordinal: i64) -> Result<Vec<Chunk>> {
243 let rows = sqlx::query(
244 "SELECT * FROM chunks WHERE doc_id = ? AND ordinal BETWEEN ? - 1 AND ? + 1 \
245 ORDER BY ordinal",
246 )
247 .bind(doc_id)
248 .bind(ordinal)
249 .bind(ordinal) .fetch_all(&self.pool)
251 .await?;
252 rows.iter().map(row_to_chunk).collect()
253 }
254
255 async fn count_documents(&self) -> Result<usize> {
256 let row = sqlx::query("SELECT COUNT(*) AS n FROM documents")
257 .fetch_one(&self.pool)
258 .await?;
259 Ok(row.get::<i64, _>("n") as usize)
260 }
261
262 async fn list_documents(&self) -> Result<Vec<Document>> {
263 let rows =
264 sqlx::query("SELECT id, source_uri, title, hash, metadata, created_at FROM documents")
265 .fetch_all(&self.pool)
266 .await?;
267 rows.iter()
268 .map(|row| {
269 let metadata: String = row.try_get("metadata")?;
270 Ok(Document {
271 id: row.try_get("id")?,
272 source_uri: row.try_get("source_uri")?,
273 title: row.try_get("title")?,
274 hash: row.try_get("hash")?,
275 metadata: serde_json::from_str(&metadata).unwrap_or(serde_json::Value::Null),
276 created_at: row.try_get("created_at")?,
277 })
278 })
279 .collect()
280 }
281
282 async fn delete_document(&self, doc_id: &str) -> Result<()> {
283 let mut tx = self.pool.begin().await?;
284 sqlx::query(
286 "DELETE FROM chunks_vec WHERE rowid IN (SELECT rowid FROM chunks WHERE doc_id = ?1)",
287 )
288 .bind(doc_id)
289 .execute(&mut *tx)
290 .await?;
291 sqlx::query("DELETE FROM chunks WHERE doc_id = ?1")
292 .bind(doc_id)
293 .execute(&mut *tx)
294 .await?;
295 sqlx::query("DELETE FROM documents WHERE id = ?1")
296 .bind(doc_id)
297 .execute(&mut *tx)
298 .await?;
299 tx.commit().await?;
300 Ok(())
301 }
302
303 async fn delete_documents_by_source(&self, source_uri: &str) -> Result<()> {
304 let mut tx = self.pool.begin().await?;
305 sqlx::query(
306 "DELETE FROM chunks_vec WHERE rowid IN (
307 SELECT c.rowid FROM chunks c
308 JOIN documents d ON c.doc_id = d.id
309 WHERE d.source_uri = ?1
310 )",
311 )
312 .bind(source_uri)
313 .execute(&mut *tx)
314 .await?;
315 sqlx::query(
316 "DELETE FROM chunks WHERE doc_id IN (SELECT id FROM documents WHERE source_uri = ?1)",
317 )
318 .bind(source_uri)
319 .execute(&mut *tx)
320 .await?;
321 sqlx::query("DELETE FROM documents WHERE source_uri = ?1")
322 .bind(source_uri)
323 .execute(&mut *tx)
324 .await?;
325 tx.commit().await?;
326 Ok(())
327 }
328
329 async fn clear(&self) -> Result<()> {
330 sqlx::query("DELETE FROM chunks_vec")
331 .execute(&self.pool)
332 .await?;
333 sqlx::query("DELETE FROM chunks")
334 .execute(&self.pool)
335 .await?;
336 sqlx::query("DELETE FROM documents")
337 .execute(&self.pool)
338 .await?;
339 Ok(())
340 }
341}
342
343#[cfg(test)]
344mod tests {
345 use super::*;
346 use crate::embed::{Embedder, HashEmbedder};
347 use crate::model::Document;
348
349 #[tokio::test]
350 async fn knn_search_via_sqlite_vec() {
351 let store = SqliteStore::connect("sqlite::memory:", 64).await.unwrap();
352 store.migrate().await.unwrap();
353
354 let embedder = HashEmbedder::new(64);
355 let doc = Document::new("mem://t", "T", "h1");
356 store.upsert_document(&doc).await.unwrap();
357
358 let texts = [
359 "vector database semantic search",
360 "banana smoothie recipe",
361 "tokio async runtime",
362 ];
363 let mut chunks = Vec::new();
364 for (i, t) in texts.iter().enumerate() {
365 let mut c = Chunk::new(&doc.id, i as i64, *t, 0);
366 c.embedding = Some(embedder.embed_one(t).await.unwrap());
367 chunks.push(c);
368 }
369 store.insert_chunks(&chunks).await.unwrap();
370 assert_eq!(store.count_chunks().await.unwrap(), 3);
371
372 let q = embedder
373 .embed_one("semantic search in a vector database")
374 .await
375 .unwrap();
376 let hits = store.vector_search(&q, 2).await.unwrap();
377 assert_eq!(hits.len(), 2);
378 assert!(
379 hits[0].chunk.text.contains("vector database"),
380 "got: {}",
381 hits[0].chunk.text
382 );
383 assert!(hits[0].score >= hits[1].score);
384
385 assert_eq!(
387 store.find_document_by_hash("h1").await.unwrap(),
388 Some(doc.id.clone())
389 );
390 store.clear().await.unwrap();
391 assert_eq!(store.count_chunks().await.unwrap(), 0);
392 assert!(store.vector_search(&q, 2).await.unwrap().is_empty());
393 }
394
395 #[tokio::test]
396 async fn rejects_wrong_dimension() {
397 let store = SqliteStore::connect("sqlite::memory:", 8).await.unwrap();
398 store.migrate().await.unwrap();
399 assert!(store.vector_search(&[0.0; 4], 1).await.is_err());
400 }
401}