Skip to main content

xz_embed/store/
memory.rs

1use async_trait::async_trait;
2use std::collections::HashMap;
3use std::fmt::Debug;
4use tokio::sync::RwLock;
5
6use crate::error::StoreError;
7use crate::traits::{StoreLifecycle, VectorStore};
8use crate::types::{MetadataFilter, SearchResult, StoreStats, VectorEntry};
9
10/// 内存向量存储(测试用)
11#[derive(Debug)]
12pub struct InMemoryVectorStore {
13    entries: RwLock<Vec<VectorEntry>>,
14    dimensions: usize,
15    closed: RwLock<bool>,
16}
17
18impl InMemoryVectorStore {
19    pub fn new(dimensions: usize) -> Self {
20        Self { entries: RwLock::new(Vec::new()), dimensions, closed: RwLock::new(false) }
21    }
22
23    async fn check_closed(&self) -> Result<(), StoreError> {
24        if *self.closed.read().await {
25            return Err(StoreError::Closed);
26        }
27        Ok(())
28    }
29
30    fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
31        let dot: f32 = a.iter().zip(b).map(|(x, y)| x * y).sum();
32        let norm_a: f32 = a.iter().map(|x| x * x).sum::<f32>().sqrt();
33        let norm_b: f32 = b.iter().map(|x| x * x).sum::<f32>().sqrt();
34        if norm_a == 0.0 || norm_b == 0.0 {
35            return 0.0;
36        }
37        dot / (norm_a * norm_b)
38    }
39
40    fn matches_filter(entry: &VectorEntry, filter: &MetadataFilter) -> bool {
41        match filter {
42            MetadataFilter::Eq { key, value } => {
43                entry.metadata.get(key).map(|v| v == value).unwrap_or(false)
44            }
45            MetadataFilter::Ne { key, value } => {
46                entry.metadata.get(key).map(|v| v != value).unwrap_or(true)
47            }
48            MetadataFilter::In { key, values } => {
49                entry.metadata.get(key).map(|v| values.contains(v)).unwrap_or(false)
50            }
51            MetadataFilter::NotIn { key, values } => {
52                entry.metadata.get(key).map(|v| !values.contains(v)).unwrap_or(true)
53            }
54            MetadataFilter::Exists { key } => entry.metadata.contains_key(key),
55            MetadataFilter::Contains { key, value } => {
56                entry.metadata.get(key).map(|v| v.contains(value)).unwrap_or(false)
57            }
58            MetadataFilter::Range { key, min, max } => {
59                if let Some(v) = entry.metadata.get(key) {
60                    if let Ok(num) = v.parse::<f64>() {
61                        return min.map_or(true, |m| num >= m) && max.map_or(true, |m| num <= m);
62                    }
63                }
64                false
65            }
66            MetadataFilter::And(filters) => filters.iter().all(|f| Self::matches_filter(entry, f)),
67            MetadataFilter::Or(filters) => filters.iter().any(|f| Self::matches_filter(entry, f)),
68            MetadataFilter::Not(filter) => !Self::matches_filter(entry, filter),
69        }
70    }
71}
72
73#[async_trait]
74impl VectorStore for InMemoryVectorStore {
75    async fn insert(&self, entry: VectorEntry) -> Result<(), StoreError> {
76        self.insert_batch(vec![entry]).await
77    }
78
79    async fn insert_batch(&self, entries: Vec<VectorEntry>) -> Result<(), StoreError> {
80        self.check_closed().await?;
81        for entry in &entries {
82            if entry.vector.len() != self.dimensions {
83                return Err(StoreError::DimensionMismatch {
84                    expected: self.dimensions,
85                    actual: entry.vector.len(),
86                });
87            }
88        }
89        self.entries.write().await.extend(entries);
90        Ok(())
91    }
92
93    async fn search(&self, query: &[f32], limit: usize) -> Result<Vec<SearchResult>, StoreError> {
94        self.check_closed().await?;
95        if query.len() != self.dimensions {
96            return Err(StoreError::DimensionMismatch {
97                expected: self.dimensions,
98                actual: query.len(),
99            });
100        }
101
102        let entries = self.entries.read().await;
103        let mut scored: Vec<(SearchResult, f32)> = entries
104            .iter()
105            .map(|entry| {
106                let similarity = Self::cosine_similarity(query, &entry.vector);
107                (
108                    SearchResult {
109                        id: entry.id.clone(),
110                        score: similarity,
111                        metadata: entry.metadata.clone(),
112                        content: entry.content.clone(),
113                        channel: entry.channel.clone(),
114                    },
115                    similarity,
116                )
117            })
118            .collect();
119
120        scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
121        scored.truncate(limit);
122
123        Ok(scored.into_iter().map(|(r, _)| r).collect())
124    }
125
126    async fn search_with_filter(
127        &self,
128        query: &[f32],
129        filter: &MetadataFilter,
130        limit: usize,
131    ) -> Result<Vec<SearchResult>, StoreError> {
132        self.check_closed().await?;
133        if query.len() != self.dimensions {
134            return Err(StoreError::DimensionMismatch {
135                expected: self.dimensions,
136                actual: query.len(),
137            });
138        }
139
140        let entries = self.entries.read().await;
141        let mut scored: Vec<(SearchResult, f32)> = entries
142            .iter()
143            .filter(|entry| Self::matches_filter(entry, filter))
144            .map(|entry| {
145                let similarity = Self::cosine_similarity(query, &entry.vector);
146                (
147                    SearchResult {
148                        id: entry.id.clone(),
149                        score: similarity,
150                        metadata: entry.metadata.clone(),
151                        content: entry.content.clone(),
152                        channel: entry.channel.clone(),
153                    },
154                    similarity,
155                )
156            })
157            .collect();
158
159        scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
160        scored.truncate(limit);
161
162        Ok(scored.into_iter().map(|(r, _)| r).collect())
163    }
164
165    async fn delete(&self, ids: &[String]) -> Result<usize, StoreError> {
166        self.check_closed().await?;
167        let mut entries = self.entries.write().await;
168        let before = entries.len();
169        entries.retain(|e| !ids.contains(&e.id));
170        Ok(before - entries.len())
171    }
172
173    async fn delete_by_filter(&self, filter: &MetadataFilter) -> Result<usize, StoreError> {
174        self.check_closed().await?;
175        let mut entries = self.entries.write().await;
176        let before = entries.len();
177        entries.retain(|e| !Self::matches_filter(e, filter));
178        Ok(before - entries.len())
179    }
180
181    async fn clear(&self) -> Result<(), StoreError> {
182        self.check_closed().await?;
183        self.entries.write().await.clear();
184        Ok(())
185    }
186
187    async fn count(&self) -> Result<usize, StoreError> {
188        self.check_closed().await?;
189        Ok(self.entries.read().await.len())
190    }
191
192    async fn rebuild_index(&self) -> Result<(), StoreError> {
193        Ok(())
194    }
195
196    async fn stats(&self) -> Result<StoreStats, StoreError> {
197        self.check_closed().await?;
198        let count = self.entries.read().await.len();
199        Ok(StoreStats {
200            total_vectors: count,
201            total_dimensions: self.dimensions,
202            index_size_bytes: 0,
203            data_size_bytes: 0,
204            last_indexed_at: None,
205        })
206    }
207}
208
209#[async_trait]
210impl StoreLifecycle for InMemoryVectorStore {
211    async fn initialize(&self) -> Result<(), StoreError> {
212        Ok(())
213    }
214
215    async fn close(&self) -> Result<(), StoreError> {
216        let mut closed = self.closed.write().await;
217        *closed = true;
218        Ok(())
219    }
220
221    async fn checkpoint(&self) -> Result<(), StoreError> {
222        Ok(())
223    }
224
225    async fn health_check(&self) -> Result<bool, StoreError> {
226        Ok(!*self.closed.read().await)
227    }
228}