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#[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}