1use chrono::{DateTime, Utc};
2use serde::{Deserialize, Serialize};
3use std::collections::HashMap;
4use uuid::Uuid;
5
6use crate::error::Result;
7
8#[cfg(feature = "fastembed")]
9use fastembed::{InitOptions, TextEmbedding};
10#[cfg(feature = "fastembed")]
11use tokio::sync::OnceCell;
12
13#[cfg(feature = "postgres")]
15pub mod postgres;
16
17#[cfg(feature = "qdrant")]
18pub mod qdrant;
19
20#[cfg(feature = "mongodb")]
21pub mod mongodb;
22
23#[cfg(feature = "turbovec")]
24pub mod turbovec;
25
26#[cfg(feature = "postgres")]
28pub use postgres::PostgresStore;
29
30#[cfg(feature = "qdrant")]
31pub use qdrant::QdrantStore;
32
33#[cfg(feature = "mongodb")]
34pub use mongodb::MongoStore;
35
36#[cfg(feature = "turbovec")]
37pub use turbovec::TurboVecStore;
38
39#[derive(Debug, Clone, Serialize, Deserialize)]
41pub struct MemoryRecord {
42 pub id: Uuid,
43 pub session_id: String,
44 pub role: String,
45 pub content: String,
46 pub importance: f32,
47 pub timestamp: DateTime<Utc>,
48 #[serde(skip_serializing_if = "Option::is_none")]
49 pub metadata: Option<HashMap<String, String>>,
50 #[serde(skip_serializing_if = "Option::is_none")]
51 pub embedding: Option<Vec<f32>>,
52}
53
54#[async_trait::async_trait]
56pub trait MemoryStore: Send + Sync {
57 async fn store(&self, record: MemoryRecord) -> Result<()>;
59
60 async fn retrieve(&self, session_id: &str, limit: usize) -> Result<Vec<MemoryRecord>>;
62
63 async fn search(
65 &self,
66 session_id: &str,
67 query_embedding: Vec<f32>,
68 limit: usize,
69 ) -> Result<Vec<MemoryRecord>>;
70
71 async fn embed(&self, text: &str) -> Result<Vec<f32>>;
73
74 async fn flush(&self) -> Result<()>;
76}
77
78pub struct InMemoryStore {
80 records: parking_lot::RwLock<Vec<MemoryRecord>>,
81 #[cfg(feature = "fastembed")]
82 embedder: OnceCell<TextEmbedding>,
83}
84
85impl InMemoryStore {
86 pub fn new() -> Self {
87 Self {
88 records: parking_lot::RwLock::new(Vec::new()),
89 #[cfg(feature = "fastembed")]
90 embedder: OnceCell::new(),
91 }
92 }
93}
94
95impl Default for InMemoryStore {
96 fn default() -> Self {
97 Self::new()
98 }
99}
100
101#[async_trait::async_trait]
102impl MemoryStore for InMemoryStore {
103 async fn store(&self, record: MemoryRecord) -> Result<()> {
104 let mut records = self.records.write();
105 records.push(record);
106 Ok(())
107 }
108
109 async fn retrieve(&self, session_id: &str, limit: usize) -> Result<Vec<MemoryRecord>> {
110 let records = self.records.read();
111 let filtered: Vec<MemoryRecord> = records
112 .iter()
113 .filter(|r| r.session_id == session_id)
114 .rev()
115 .take(limit)
116 .cloned()
117 .collect();
118 Ok(filtered)
119 }
120
121 async fn search(
122 &self,
123 session_id: &str,
124 query_embedding: Vec<f32>,
125 limit: usize,
126 ) -> Result<Vec<MemoryRecord>> {
127 let records = self.records.read();
128 let mut scored: Vec<(f32, MemoryRecord)> = records
129 .iter()
130 .filter(|r| r.session_id == session_id && r.embedding.is_some())
131 .map(|r| {
132 let embedding = r.embedding.as_ref().unwrap();
133 let similarity = cosine_similarity(&query_embedding, embedding);
134 (similarity, r.clone())
135 })
136 .collect();
137
138 scored.sort_by(|a, b| b.0.partial_cmp(&a.0).unwrap());
139 Ok(scored.into_iter().take(limit).map(|(_, r)| r).collect())
140 }
141
142 async fn flush(&self) -> Result<()> {
143 Ok(())
144 }
145
146 async fn embed(&self, _text: &str) -> Result<Vec<f32>> {
147 #[cfg(feature = "fastembed")]
148 {
149 let embedder = self
150 .embedder
151 .get_or_try_init(|| async {
152 TextEmbedding::try_new(InitOptions::default())
153 .map_err(|e| crate::error::AgentError::MemoryError(e.to_string()))
154 })
155 .await?;
156
157 let embeddings = embedder
158 .embed(vec![_text], None)
159 .map_err(|e| crate::error::AgentError::MemoryError(e.to_string()))?;
160
161 Ok(embeddings[0].clone())
162 }
163
164 #[cfg(not(feature = "fastembed"))]
165 Ok(vec![])
166 }
167}
168
169fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
171 if a.len() != b.len() {
172 return 0.0;
173 }
174
175 let dot: f32 = a.iter().zip(b.iter()).map(|(x, y)| x * y).sum();
176 let mag_a: f32 = a.iter().map(|x| x * x).sum::<f32>().sqrt();
177 let mag_b: f32 = b.iter().map(|x| x * x).sum::<f32>().sqrt();
178
179 if mag_a == 0.0 || mag_b == 0.0 {
180 0.0
181 } else {
182 dot / (mag_a * mag_b)
183 }
184}
185
186pub fn mmr_rerank_records(
187 query_embedding: &[f32],
188 candidates: Vec<MemoryRecord>,
189 k: usize,
190 lambda: f32,
191) -> Vec<MemoryRecord> {
192 if candidates.is_empty() {
193 return Vec::new();
194 }
195
196 let k = k.min(candidates.len());
197 let mut selected_indices = Vec::with_capacity(k);
198 let mut remaining_indices: Vec<usize> = (0..candidates.len()).collect();
199
200 if let Some((idx, _)) = remaining_indices
202 .iter()
203 .enumerate()
204 .filter_map(|(i, &r_idx)| {
205 candidates[r_idx]
206 .embedding
207 .as_ref()
208 .map(|emb| (i, cosine_similarity(query_embedding, emb)))
209 })
210 .max_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap())
211 {
212 let selected_idx = remaining_indices.remove(idx);
213 selected_indices.push(selected_idx);
214 }
215
216 while selected_indices.len() < k && !remaining_indices.is_empty() {
218 let next_idx = remaining_indices
219 .iter()
220 .enumerate()
221 .filter_map(|(i, &r_idx)| {
222 let emb = candidates[r_idx].embedding.as_ref()?;
223
224 let relevance = cosine_similarity(query_embedding, emb);
226
227 let max_sim_selected = selected_indices
229 .iter()
230 .filter_map(|&s_idx| candidates[s_idx].embedding.as_ref())
231 .map(|s_emb| cosine_similarity(emb, s_emb))
232 .fold(f32::NEG_INFINITY, f32::max);
233
234 let mmr_score = lambda * relevance - (1.0 - lambda) * max_sim_selected;
236
237 Some((i, mmr_score))
238 })
239 .max_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap())
240 .map(|(i, _)| i);
241
242 if let Some(idx) = next_idx {
243 let selected_idx = remaining_indices.remove(idx);
244 selected_indices.push(selected_idx);
245 } else {
246 break;
247 }
248 }
249
250 selected_indices
251 .into_iter()
252 .map(|i| candidates[i].clone())
253 .collect()
254}
255
256pub fn mmr_rerank(
261 query_embedding: &[f32],
262 candidates: Vec<MemoryRecord>,
263 k: usize,
264 lambda: f32,
265) -> Vec<MemoryRecord> {
266 if candidates.is_empty() {
267 return Vec::new();
268 }
269
270 let k = k.min(candidates.len());
271 let mut selected = Vec::with_capacity(k);
272 let mut remaining = candidates;
273
274 if let Some((idx, _)) = remaining
276 .iter()
277 .enumerate()
278 .filter_map(|(i, r)| {
279 r.embedding
280 .as_ref()
281 .map(|emb| (i, cosine_similarity(query_embedding, emb)))
282 })
283 .max_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap())
284 {
285 selected.push(remaining.swap_remove(idx));
286 }
287
288 while selected.len() < k && !remaining.is_empty() {
290 let next_idx = remaining
291 .iter()
292 .enumerate()
293 .filter_map(|(i, r)| {
294 let emb = r.embedding.as_ref()?;
295
296 let relevance = cosine_similarity(query_embedding, emb);
298
299 let max_sim_selected = selected
301 .iter()
302 .filter_map(|s| s.embedding.as_ref())
303 .map(|s_emb| cosine_similarity(emb, s_emb))
304 .fold(f32::NEG_INFINITY, f32::max);
305
306 let mmr_score = lambda * relevance - (1.0 - lambda) * max_sim_selected;
308
309 Some((i, mmr_score))
310 })
311 .max_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap())
312 .map(|(i, _)| i);
313
314 if let Some(idx) = next_idx {
315 selected.push(remaining.swap_remove(idx));
316 } else {
317 break;
318 }
319 }
320
321 selected
322}
323
324pub struct SessionMemory {
326 store: Box<dyn MemoryStore>,
327 short_term: parking_lot::RwLock<HashMap<String, Vec<MemoryRecord>>>,
329 context_window: usize,
330}
331
332impl SessionMemory {
333 pub fn new(store: Box<dyn MemoryStore>, context_window: usize) -> Self {
335 Self {
336 store,
337 short_term: parking_lot::RwLock::new(HashMap::new()),
338 context_window,
339 }
340 }
341
342 pub async fn store(&self, record: MemoryRecord) -> Result<()> {
344 let session_id = record.session_id.clone();
345
346 {
348 let mut short_term = self.short_term.write();
349 let session_records = short_term.entry(session_id).or_insert_with(Vec::new);
350 session_records.push(record.clone());
351
352 if session_records.len() > self.context_window {
354 session_records.drain(0..session_records.len() - self.context_window);
355 }
356 }
357
358 let mut record = record;
360 if record.embedding.is_none() && !record.content.is_empty() {
361 if let Ok(embedding) = self.store.embed(&record.content).await {
362 if !embedding.is_empty() {
363 record.embedding = Some(embedding);
364 }
365 }
366 }
367
368 self.store.store(record).await
370 }
371
372 pub async fn retrieve_recent(&self, session_id: &str) -> Result<Vec<MemoryRecord>> {
374 let short_term = self.short_term.read();
375 Ok(short_term.get(session_id).cloned().unwrap_or_default())
376 }
377
378 pub async fn search(
379 &self,
380 session_id: &str,
381 query: &str,
382 limit: usize,
383 ) -> Result<Vec<MemoryRecord>> {
384 let query_embedding = self.store.embed(query).await?;
385 if query_embedding.is_empty() {
386 return Ok(Vec::new());
387 }
388 self.store.search(session_id, query_embedding, limit).await
389 }
390
391 pub async fn embed(&self, text: &str) -> Result<Vec<f32>> {
393 self.store.embed(text).await
394 }
395
396 pub async fn flush(&self) -> Result<()> {
398 self.store.flush().await
399 }
400}
401
402#[cfg(test)]
403mod tests {
404 use super::*;
405
406 #[tokio::test]
407 async fn test_in_memory_store() {
408 let store = InMemoryStore::new();
409 let record = MemoryRecord {
410 id: Uuid::new_v4(),
411 session_id: "test".to_string(),
412 role: "user".to_string(),
413 content: "Hello".to_string(),
414 importance: 0.8,
415 timestamp: Utc::now(),
416 metadata: None,
417 embedding: None,
418 };
419
420 store.store(record.clone()).await.unwrap();
421 let retrieved = store.retrieve("test", 10).await.unwrap();
422 assert_eq!(retrieved.len(), 1);
423 assert_eq!(retrieved[0].content, "Hello");
424 }
425
426 #[tokio::test]
427 async fn test_session_memory() {
428 let store = Box::new(InMemoryStore::new());
430 let memory = SessionMemory::new(store, 5);
431
432 let record = MemoryRecord {
433 id: Uuid::new_v4(),
434 session_id: "test".to_string(),
435 role: "user".to_string(),
436 content: "Test message".to_string(),
437 importance: 0.9,
438 timestamp: Utc::now(),
439 metadata: None,
440 embedding: None,
441 };
442
443 memory.store(record).await.unwrap();
444 let recent = memory.retrieve_recent("test").await.unwrap();
445 assert_eq!(recent.len(), 1);
446 }
447}