Skip to main content

xz_embed/store/
sqlite_vec.rs

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/// sqlite-vec 向量存储实现
10///
11/// 特性:
12/// - 零外部依赖(通过 sqlx + sqlite)
13/// - 余弦距离搜索
14/// - 元数据过滤通过 SQL WHERE 子句实现
15/// - WAL 模式支持并发读
16#[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    /// 创建新的 sqlite-vec 存储
26    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        // 启用 WAL 模式
45        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    /// 设置表名
54    pub fn with_table_name(mut self, name: &str) -> Self {
55        self.table_name = name.to_string();
56        self
57    }
58
59    /// 设置最大存储容量
60    pub fn with_max_capacity(mut self, capacity: Option<u64>) -> Self {
61        self.max_capacity = capacity;
62        self
63    }
64
65    /// 清理过期数据
66    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        // 检查维度
199        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 &params {
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 &params {
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        // sqlite-vec 不依赖传统索引,此操作仅做 WAL checkpoint
391        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}