Skip to main content

rs_agent/memory/
mod.rs

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// Memory backend implementations
14#[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// Re-export backends
27#[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/// Memory record storing a piece of information
40#[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/// Memory store trait for different backends
55#[async_trait::async_trait]
56pub trait MemoryStore: Send + Sync {
57    /// Stores a memory record
58    async fn store(&self, record: MemoryRecord) -> Result<()>;
59
60    /// Retrieves memories for a session
61    async fn retrieve(&self, session_id: &str, limit: usize) -> Result<Vec<MemoryRecord>>;
62
63    /// Searches for similar memories using embeddings
64    async fn search(
65        &self,
66        session_id: &str,
67        query_embedding: Vec<f32>,
68        limit: usize,
69    ) -> Result<Vec<MemoryRecord>>;
70
71    /// Embeds text using the store's embedding model
72    async fn embed(&self, text: &str) -> Result<Vec<f32>>;
73
74    /// Flushes all pending writes
75    async fn flush(&self) -> Result<()>;
76}
77
78/// In-memory store implementation
79pub 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
169/// Calculates cosine similarity between two vectors
170fn 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    // Select first item with highest similarity to query
201    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    // Iteratively select items that maximize MMR score
217    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                // Relevance: similarity to query
225                let relevance = cosine_similarity(query_embedding, emb);
226
227                // Diversity: max similarity to already selected items
228                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                // MMR score: λ * relevance - (1-λ) * max_similarity_to_selected
235                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
256/// Maximal Marginal Relevance (MMR) for diverse retrieval
257///
258/// Balances relevance to query with diversity in results.
259/// Lambda controls the trade-off: 1.0 = pure relevance, 0.0 = pure diversity
260pub 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    // Select first item with highest similarity to query
275    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    // Iteratively select items that maximize MMR score
289    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                // Relevance: similarity to query
297                let relevance = cosine_similarity(query_embedding, emb);
298
299                // Diversity: max similarity to already selected items
300                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                // MMR score: λ * relevance - (1-λ) * max_similarity_to_selected
307                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
324/// Session memory manages short-term and long-term memory for a session
325pub struct SessionMemory {
326    store: Box<dyn MemoryStore>,
327    // Short-term cache of recent messages
328    short_term: parking_lot::RwLock<HashMap<String, Vec<MemoryRecord>>>,
329    context_window: usize,
330}
331
332impl SessionMemory {
333    /// Creates a new session memory with the given store
334    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    /// Stores a memory record
343    pub async fn store(&self, record: MemoryRecord) -> Result<()> {
344        let session_id = record.session_id.clone();
345
346        // Add to short-term cache
347        {
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            // Trim to context window
353            if session_records.len() > self.context_window {
354                session_records.drain(0..session_records.len() - self.context_window);
355            }
356        }
357
358        // Generate embedding if not present
359        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        // Store in long-term
369        self.store.store(record).await
370    }
371
372    /// Retrieves recent memories from short-term cache
373    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    /// Embeds text manually (exposed for testing/utils)
392    pub async fn embed(&self, text: &str) -> Result<Vec<f32>> {
393        self.store.embed(text).await
394    }
395
396    /// Flushes all pending writes
397    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        // Mock store that doesn't actually embed but stores records
429        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}