1use async_trait::async_trait;
2use std::collections::HashMap;
3use std::fmt::Debug;
4
5use crate::error::StoreError;
6use crate::traits::{StoreLifecycle, VectorStore};
7use crate::types::{MetadataFilter, SearchResult, StoreStats, VectorEntry};
8
9#[derive(Debug)]
17pub struct SqliteVecStore {
18 pool: sqlx::SqlitePool,
19 dimensions: usize,
20 table_name: String,
21 max_capacity: Option<u64>,
22}
23
24impl SqliteVecStore {
25 pub async fn new(
27 path: &str,
28 dimensions: usize,
29 max_pool_size: Option<usize>,
30 ) -> Result<Self, StoreError> {
31 let pool_size = max_pool_size.unwrap_or(5);
32 let conn_str = if path == ":memory:" {
33 "sqlite::memory:".to_string()
34 } else {
35 format!("sqlite://{path}")
36 };
37
38 let pool = sqlx::sqlite::SqlitePoolOptions::new()
39 .max_connections(pool_size as u32)
40 .connect(&conn_str)
41 .await
42 .map_err(|e| StoreError::Database(e.to_string()))?;
43
44 sqlx::query("PRAGMA journal_mode=WAL")
46 .execute(&pool)
47 .await
48 .map_err(|e| StoreError::Database(e.to_string()))?;
49
50 Ok(Self { pool, dimensions, table_name: "embeddings".into(), max_capacity: None })
51 }
52
53 pub fn with_table_name(mut self, name: &str) -> Self {
55 self.table_name = name.to_string();
56 self
57 }
58
59 pub fn with_max_capacity(mut self, capacity: Option<u64>) -> Self {
61 self.max_capacity = capacity;
62 self
63 }
64
65 pub async fn purge_expired(&self) -> Result<usize, StoreError> {
67 let now_ms = std::time::SystemTime::now()
68 .duration_since(std::time::UNIX_EPOCH)
69 .unwrap_or_default()
70 .as_millis() as u64;
71
72 let result = sqlx::query(&format!(
73 "DELETE FROM {} WHERE expires_at IS NOT NULL AND expires_at < ?",
74 self.table_name
75 ))
76 .bind(now_ms as i64)
77 .execute(&self.pool)
78 .await
79 .map_err(|e| StoreError::Database(e.to_string()))?;
80
81 Ok(result.rows_affected() as usize)
82 }
83
84 fn build_filter_clause(filter: &MetadataFilter) -> (String, Vec<String>) {
85 match filter {
86 MetadataFilter::Eq { key, value } => {
87 (format!("json_extract(metadata_json, '$.{key}') = ?"), vec![value.clone()])
88 }
89 MetadataFilter::Ne { key, value } => {
90 (format!("json_extract(metadata_json, '$.{key}') != ?"), vec![value.clone()])
91 }
92 MetadataFilter::In { key, values } => {
93 let placeholders: Vec<String> = values.iter().map(|_| "?".to_string()).collect();
94 (
95 format!(
96 "json_extract(metadata_json, '$.{key}') IN ({})",
97 placeholders.join(", ")
98 ),
99 values.clone(),
100 )
101 }
102 MetadataFilter::NotIn { key, values } => {
103 let placeholders: Vec<String> = values.iter().map(|_| "?".to_string()).collect();
104 (
105 format!(
106 "json_extract(metadata_json, '$.{key}') NOT IN ({})",
107 placeholders.join(", ")
108 ),
109 values.clone(),
110 )
111 }
112 MetadataFilter::Exists { key } => {
113 (format!("json_extract(metadata_json, '$.{key}') IS NOT NULL"), vec![])
114 }
115 MetadataFilter::Contains { key, value } => (
116 format!("json_extract(metadata_json, '$.{key}') LIKE ?"),
117 vec![format!("%{value}%")],
118 ),
119 MetadataFilter::Range { key, min, max } => {
120 let mut clauses = Vec::new();
121 let mut params = Vec::new();
122 if let Some(min_val) = min {
123 clauses
124 .push(format!("CAST(json_extract(metadata_json, '$.{key}') AS REAL) >= ?"));
125 params.push(min_val.to_string());
126 }
127 if let Some(max_val) = max {
128 clauses
129 .push(format!("CAST(json_extract(metadata_json, '$.{key}') AS REAL) <= ?"));
130 params.push(max_val.to_string());
131 }
132 (clauses.join(" AND "), params)
133 }
134 MetadataFilter::And(filters) => {
135 let mut clauses = Vec::new();
136 let mut all_params = Vec::new();
137 for f in filters {
138 let (clause, mut params) = Self::build_filter_clause(f);
139 if !clause.is_empty() {
140 clauses.push(format!("({clause})"));
141 all_params.append(&mut params);
142 }
143 }
144 (clauses.join(" AND "), all_params)
145 }
146 MetadataFilter::Or(filters) => {
147 let mut clauses = Vec::new();
148 let mut all_params = Vec::new();
149 for f in filters {
150 let (clause, mut params) = Self::build_filter_clause(f);
151 if !clause.is_empty() {
152 clauses.push(format!("({clause})"));
153 all_params.append(&mut params);
154 }
155 }
156 (clauses.join(" OR "), all_params)
157 }
158 MetadataFilter::Not(filter) => {
159 let (inner, params) = Self::build_filter_clause(filter);
160 (format!("NOT ({inner})"), params)
161 }
162 }
163 }
164
165 fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
166 if a.len() != b.len() {
167 return 0.0;
168 }
169 let dot: f32 = a.iter().zip(b).map(|(x, y)| x * y).sum();
170 let norm_a: f32 = a.iter().map(|x| x * x).sum::<f32>().sqrt();
171 let norm_b: f32 = b.iter().map(|x| x * x).sum::<f32>().sqrt();
172 if norm_a == 0.0 || norm_b == 0.0 {
173 return 0.0;
174 }
175 dot / (norm_a * norm_b)
176 }
177
178 fn vector_to_blob(v: &[f32]) -> Vec<u8> {
179 v.iter().flat_map(|f| f.to_le_bytes()).collect()
180 }
181
182 fn blob_to_vector(b: &[u8]) -> Vec<f32> {
183 b.chunks_exact(4).map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]])).collect()
184 }
185}
186
187#[async_trait]
188impl VectorStore for SqliteVecStore {
189 async fn insert(&self, entry: VectorEntry) -> Result<(), StoreError> {
190 self.insert_batch(vec![entry]).await
191 }
192
193 async fn insert_batch(&self, entries: Vec<VectorEntry>) -> Result<(), StoreError> {
194 if entries.is_empty() {
195 return Ok(());
196 }
197
198 for entry in &entries {
200 if entry.vector.len() != self.dimensions {
201 return Err(StoreError::DimensionMismatch {
202 expected: self.dimensions,
203 actual: entry.vector.len(),
204 });
205 }
206 }
207
208 for entry in entries {
209 let vector_blob = Self::vector_to_blob(&entry.vector);
210 let metadata_json = serde_json::to_string(&entry.metadata)
211 .map_err(|e| StoreError::Serialization(e.to_string()))?;
212
213 sqlx::query(&format!(
214 "INSERT OR REPLACE INTO {} (id, content, metadata_json, channel, created_at, expires_at, embedding) VALUES (?, ?, ?, ?, ?, ?, ?)",
215 self.table_name
216 ))
217 .bind(&entry.id)
218 .bind(&entry.content)
219 .bind(&metadata_json)
220 .bind(&entry.channel)
221 .bind(entry.created_at as i64)
222 .bind(entry.expires_at.map(|t| t as i64))
223 .bind(&vector_blob)
224 .execute(&self.pool)
225 .await
226 .map_err(|e| StoreError::Database(e.to_string()))?;
227 }
228
229 Ok(())
230 }
231
232 async fn search(&self, query: &[f32], limit: usize) -> Result<Vec<SearchResult>, StoreError> {
233 if query.len() != self.dimensions {
234 return Err(StoreError::DimensionMismatch {
235 expected: self.dimensions,
236 actual: query.len(),
237 });
238 }
239
240 let rows = sqlx::query_as::<_, EmbeddingRow>(&format!(
241 "SELECT id, content, metadata_json, channel, embedding FROM {}",
242 self.table_name
243 ))
244 .fetch_all(&self.pool)
245 .await
246 .map_err(|e| StoreError::Database(e.to_string()))?;
247
248 let mut scored: Vec<(SearchResult, f32)> = rows
249 .iter()
250 .map(|row| {
251 let vector = Self::blob_to_vector(&row.embedding);
252 let similarity = Self::cosine_similarity(query, &vector);
253 let metadata: HashMap<String, String> =
254 serde_json::from_str(&row.metadata_json).unwrap_or_default();
255
256 (
257 SearchResult {
258 id: row.id.clone(),
259 score: similarity,
260 metadata,
261 content: row.content.clone(),
262 channel: row.channel.clone(),
263 },
264 similarity,
265 )
266 })
267 .collect();
268
269 scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
270 scored.truncate(limit);
271
272 Ok(scored.into_iter().map(|(r, _)| r).collect())
273 }
274
275 async fn search_with_filter(
276 &self,
277 query: &[f32],
278 filter: &MetadataFilter,
279 limit: usize,
280 ) -> Result<Vec<SearchResult>, StoreError> {
281 if query.len() != self.dimensions {
282 return Err(StoreError::DimensionMismatch {
283 expected: self.dimensions,
284 actual: query.len(),
285 });
286 }
287
288 let (filter_clause, params) = Self::build_filter_clause(filter);
289 let sql = if filter_clause.is_empty() {
290 format!(
291 "SELECT id, content, metadata_json, channel, embedding FROM {}",
292 self.table_name
293 )
294 } else {
295 format!(
296 "SELECT id, content, metadata_json, channel, embedding FROM {} WHERE {}",
297 self.table_name, filter_clause
298 )
299 };
300
301 let mut query_builder = sqlx::query_as::<_, EmbeddingRow>(&sql);
302 for param in ¶ms {
303 query_builder = query_builder.bind(param);
304 }
305
306 let rows = query_builder
307 .fetch_all(&self.pool)
308 .await
309 .map_err(|e| StoreError::Database(e.to_string()))?;
310
311 let mut scored: Vec<(SearchResult, f32)> = rows
312 .iter()
313 .map(|row| {
314 let vector = Self::blob_to_vector(&row.embedding);
315 let similarity = Self::cosine_similarity(query, &vector);
316 let metadata: HashMap<String, String> =
317 serde_json::from_str(&row.metadata_json).unwrap_or_default();
318
319 (
320 SearchResult {
321 id: row.id.clone(),
322 score: similarity,
323 metadata,
324 content: row.content.clone(),
325 channel: row.channel.clone(),
326 },
327 similarity,
328 )
329 })
330 .collect();
331
332 scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
333 scored.truncate(limit);
334
335 Ok(scored.into_iter().map(|(r, _)| r).collect())
336 }
337
338 async fn delete(&self, ids: &[String]) -> Result<usize, StoreError> {
339 if ids.is_empty() {
340 return Ok(0);
341 }
342 let placeholders: Vec<String> = ids.iter().map(|_| "?".to_string()).collect();
343 let sql =
344 format!("DELETE FROM {} WHERE id IN ({})", self.table_name, placeholders.join(", "));
345
346 let mut query = sqlx::query(&sql);
347 for id in ids {
348 query = query.bind(id);
349 }
350
351 let result =
352 query.execute(&self.pool).await.map_err(|e| StoreError::Database(e.to_string()))?;
353
354 Ok(result.rows_affected() as usize)
355 }
356
357 async fn delete_by_filter(&self, filter: &MetadataFilter) -> Result<usize, StoreError> {
358 let (filter_clause, params) = Self::build_filter_clause(filter);
359 let sql = format!("DELETE FROM {} WHERE {}", self.table_name, filter_clause);
360
361 let mut query = sqlx::query(&sql);
362 for param in ¶ms {
363 query = query.bind(param);
364 }
365
366 let result =
367 query.execute(&self.pool).await.map_err(|e| StoreError::Database(e.to_string()))?;
368
369 Ok(result.rows_affected() as usize)
370 }
371
372 async fn clear(&self) -> Result<(), StoreError> {
373 sqlx::query(&format!("DELETE FROM {}", self.table_name))
374 .execute(&self.pool)
375 .await
376 .map_err(|e| StoreError::Database(e.to_string()))?;
377 Ok(())
378 }
379
380 async fn count(&self) -> Result<usize, StoreError> {
381 let (count,): (i64,) = sqlx::query_as(&format!("SELECT COUNT(*) FROM {}", self.table_name))
382 .fetch_one(&self.pool)
383 .await
384 .map_err(|e| StoreError::Database(e.to_string()))?;
385
386 Ok(count as usize)
387 }
388
389 async fn rebuild_index(&self) -> Result<(), StoreError> {
390 sqlx::query("PRAGMA wal_checkpoint(FULL)")
392 .execute(&self.pool)
393 .await
394 .map_err(|e| StoreError::Database(e.to_string()))?;
395 Ok(())
396 }
397
398 async fn stats(&self) -> Result<StoreStats, StoreError> {
399 let count = self.count().await?;
400 Ok(StoreStats {
401 total_vectors: count,
402 total_dimensions: self.dimensions,
403 index_size_bytes: 0,
404 data_size_bytes: 0,
405 last_indexed_at: None,
406 })
407 }
408}
409
410#[async_trait]
411impl StoreLifecycle for SqliteVecStore {
412 async fn initialize(&self) -> Result<(), StoreError> {
413 sqlx::query(&format!(
414 "CREATE TABLE IF NOT EXISTS {} (
415 id TEXT PRIMARY KEY,
416 content TEXT,
417 metadata_json TEXT,
418 channel TEXT,
419 created_at INTEGER NOT NULL,
420 expires_at INTEGER,
421 embedding BLOB NOT NULL
422 )",
423 self.table_name
424 ))
425 .execute(&self.pool)
426 .await
427 .map_err(|e| StoreError::Database(e.to_string()))?;
428
429 Ok(())
430 }
431
432 async fn close(&self) -> Result<(), StoreError> {
433 self.pool.close().await;
434 Ok(())
435 }
436
437 async fn checkpoint(&self) -> Result<(), StoreError> {
438 sqlx::query("PRAGMA wal_checkpoint(FULL)")
439 .execute(&self.pool)
440 .await
441 .map_err(|e| StoreError::Database(e.to_string()))?;
442 Ok(())
443 }
444
445 async fn health_check(&self) -> Result<bool, StoreError> {
446 sqlx::query("SELECT 1")
447 .execute(&self.pool)
448 .await
449 .map_err(|e| StoreError::Database(e.to_string()))?;
450 Ok(true)
451 }
452}
453
454#[derive(Debug, sqlx::FromRow)]
455struct EmbeddingRow {
456 id: String,
457 content: Option<String>,
458 metadata_json: String,
459 channel: Option<String>,
460 embedding: Vec<u8>,
461}
462
463#[cfg(test)]
464mod tests {
465 use super::*;
466 use crate::store::sqlite_vec::SqliteVecStore;
467
468 #[test]
469 fn build_filter_clause_empty_and() {
470 let filter = MetadataFilter::And(vec![]);
471 let (clause, _params) = SqliteVecStore::build_filter_clause(&filter);
472 assert!(clause.is_empty(), "empty And filter should produce empty clause");
473 }
474
475 #[test]
476 fn build_filter_clause_empty_or() {
477 let filter = MetadataFilter::Or(vec![]);
478 let (clause, _params) = SqliteVecStore::build_filter_clause(&filter);
479 assert!(clause.is_empty(), "empty Or filter should produce empty clause");
480 }
481}