Skip to main content

agent_memo/
lib.rs

1//! `agent-memo` — the platform's unified **context memory** (memo).
2//!
3//! Design boundary (see `docs/adr/0004-memo-storage.md`):
4//! * memo is the **only** source of conversational / long-term context.
5//! * memo uses a **local / embedded** store (sled). It does **NOT** use Postgres.
6//! * Postgres is used exclusively by `agent-cloud` for structured metadata
7//!   (agents / sessions / runs). Context text never lives in Postgres.
8
9use async_trait::async_trait;
10use serde::{Deserialize, Serialize};
11use std::path::Path;
12use std::sync::Arc;
13use thiserror::Error;
14use uuid::Uuid;
15
16/// Kinds of context fragments memo can hold.
17#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq)]
18pub enum FragmentKind {
19    /// A turn of the conversation (user message or assistant reply).
20    Message,
21    /// A tool call result that should be remembered across turns.
22    ToolResult,
23    /// An explicitly externalized long-term memory written via `memorize`.
24    LongTerm,
25    /// A system / scratch note.
26    Note,
27}
28
29impl FragmentKind {
30    pub fn as_str(&self) -> &'static str {
31        match self {
32            FragmentKind::Message => "message",
33            FragmentKind::ToolResult => "tool_result",
34            FragmentKind::LongTerm => "long_term",
35            FragmentKind::Note => "note",
36        }
37    }
38}
39
40impl std::str::FromStr for FragmentKind {
41    type Err = String;
42
43    fn from_str(s: &str) -> Result<Self, Self::Err> {
44        match s {
45            "message" => Ok(FragmentKind::Message),
46            "tool_result" => Ok(FragmentKind::ToolResult),
47            "long_term" => Ok(FragmentKind::LongTerm),
48            "note" => Ok(FragmentKind::Note),
49            other => Err(format!("unknown fragment kind: {other}")),
50        }
51    }
52}
53
54/// A single unit of remembered context.
55#[derive(Debug, Clone, Serialize, Deserialize)]
56pub struct ContextFragment {
57    pub id: String,
58    pub session: String,
59    /// Optional explicit key (used by `Session::memorize`/`recall`).
60    pub key: Option<String>,
61    pub kind: FragmentKind,
62    pub content: String,
63    /// Unix epoch milliseconds.
64    pub created_at: i64,
65    /// Optional dense vector for semantic (vector) recall. Populated by the
66    /// local [`embed::LocalEmbedder`] when a fragment is memorized without one.
67    pub embedding: Option<Vec<f32>>,
68}
69
70impl ContextFragment {
71    pub fn new(session: &str, kind: FragmentKind, content: impl Into<String>) -> Self {
72        Self {
73            id: Uuid::new_v4().to_string(),
74            session: session.to_string(),
75            key: None,
76            kind,
77            content: content.into(),
78            created_at: now_ms(),
79            embedding: None,
80        }
81    }
82
83    pub fn with_key(mut self, key: impl Into<String>) -> Self {
84        self.key = Some(key.into());
85        self
86    }
87}
88
89/// Query used by [`MemoStore::recall`].
90#[derive(Debug, Clone)]
91pub struct RecallQuery {
92    pub session: String,
93    pub text: String,
94    pub top_k: usize,
95    pub kind: Option<FragmentKind>,
96}
97
98impl RecallQuery {
99    pub fn new(session: &str, text: impl Into<String>) -> Self {
100        Self {
101            session: session.to_string(),
102            text: text.into(),
103            top_k: 8,
104            kind: None,
105        }
106    }
107
108    pub fn with_kind(mut self, kind: FragmentKind) -> Self {
109        self.kind = Some(kind);
110        self
111    }
112}
113
114#[derive(Debug, Error)]
115pub enum MemoError {
116    #[error("storage error: {0}")]
117    Storage(String),
118    #[error("serialization error: {0}")]
119    Serialization(#[from] serde_json::Error),
120    #[error("not found: {0}")]
121    NotFound(String),
122    #[error("embedding error: {0}")]
123    Embedding(String),
124}
125
126/// The unified context-memory contract. Everything that needs conversational
127/// or long-term context goes through this trait.
128#[async_trait]
129pub trait MemoStore: Send + Sync {
130    /// Persist a fragment.
131    async fn memorize(&self, frag: ContextFragment) -> Result<(), MemoError>;
132    /// Retrieve the most relevant fragments for a query.
133    async fn recall(&self, query: &RecallQuery) -> Result<Vec<ContextFragment>, MemoError>;
134    /// Merge a session's fragments into a single compact fragment.
135    async fn compact(&self, session: &str) -> Result<ContextFragment, MemoError>;
136    /// Lookup an explicit long-term memory by key (used by SDK `recall`).
137    async fn get_by_key(
138        &self,
139        session: &str,
140        key: &str,
141    ) -> Result<Option<ContextFragment>, MemoError>;
142}
143
144/// sled-backed implementation. Self-contained, no Postgres. Uses a local
145/// [`embed::LocalEmbedder`] to compute dense vectors for semantic (vector) recall.
146pub struct SledMemoStore {
147    fragments: sled::Tree,
148    embedder: Arc<dyn embed::Embedder>,
149}
150
151impl SledMemoStore {
152    /// Open (or create) a memo store at `path`. Use `":memory:"` for a temp db.
153    pub fn open(path: &Path) -> Result<Arc<Self>, MemoError> {
154        let db = sled::open(path).map_err(|e| MemoError::Storage(e.to_string()))?;
155        let fragments = db
156            .open_tree("fragments")
157            .map_err(|e| MemoError::Storage(e.to_string()))?;
158        Ok(Arc::new(Self {
159            fragments,
160            embedder: Arc::new(embed::LocalEmbedder::new(64)),
161        }))
162    }
163
164    /// In-memory variant (handy for tests and for the SDK default).
165    pub fn memory() -> Result<Arc<Self>, MemoError> {
166        let db = sled::Config::new()
167            .temporary(true)
168            .open()
169            .map_err(|e| MemoError::Storage(e.to_string()))?;
170        let fragments = db
171            .open_tree("fragments")
172            .map_err(|e| MemoError::Storage(e.to_string()))?;
173        Ok(Arc::new(Self {
174            fragments,
175            embedder: Arc::new(embed::LocalEmbedder::new(64)),
176        }))
177    }
178}
179
180#[async_trait]
181impl MemoStore for SledMemoStore {
182    async fn memorize(&self, mut frag: ContextFragment) -> Result<(), MemoError> {
183        // Compute a dense vector for semantic recall when one isn't supplied.
184        if frag.embedding.is_none() && !frag.content.trim().is_empty() {
185            if let Ok(v) = self.embedder.embed(&frag.content) {
186                frag.embedding = Some(v);
187            }
188        }
189        let key = frag.id.as_bytes().to_vec();
190        let value = serde_json::to_vec(&frag)?;
191        self.fragments
192            .insert(key, value)
193            .map_err(|e| MemoError::Storage(e.to_string()))?;
194        Ok(())
195    }
196
197    async fn recall(&self, query: &RecallQuery) -> Result<Vec<ContextFragment>, MemoError> {
198        let q = query.text.to_lowercase();
199        // Query vector for semantic (vector) recall; None when empty/unsupported.
200        let q_emb = if q.trim().is_empty() {
201            None
202        } else {
203            self.embedder.embed(&query.text).ok()
204        };
205        let mut scored: Vec<(f32, ContextFragment)> = Vec::new();
206        for item in self.fragments.iter() {
207            let (_k, v) = item.map_err(|e| MemoError::Storage(e.to_string()))?;
208            let frag: ContextFragment = serde_json::from_slice(&v)?;
209            if frag.session != query.session {
210                continue;
211            }
212            if let Some(kind) = query.kind {
213                if frag.kind != kind {
214                    continue;
215                }
216            }
217            // Keyword + length-weighted relevance (embeddings are optional).
218            // An exact key match scores highest; otherwise we fall back to
219            // substring/keyword matching against the content.
220            let mut kw = 0.0f32;
221            if frag.key.as_deref() == Some(query.text.as_str()) {
222                kw = 1.0;
223            }
224            let hay = frag.content.to_lowercase();
225            if hay.contains(&q) {
226                let hits = q
227                    .split_whitespace()
228                    .filter(|w| !w.is_empty() && hay.contains(*w))
229                    .count() as f32;
230                kw = kw.max(hits / (hay.len() as f32).max(1.0).log10());
231            }
232            // Blend semantic (vector) similarity with the keyword score.
233            // Pure keyword when a vector isn't available on either side.
234            let score = match (&q_emb, &frag.embedding) {
235                (Some(a), Some(b)) => match embed::cosine(a, b) {
236                    Some(v) => 0.7 * v + 0.3 * kw,
237                    None => kw,
238                },
239                _ => kw,
240            };
241            if score > 0.0 {
242                scored.push((score, frag));
243            }
244        }
245        scored.sort_by(|a, b| b.0.partial_cmp(&a.0).unwrap_or(std::cmp::Ordering::Equal));
246        scored.truncate(query.top_k);
247        Ok(scored.into_iter().map(|(_, f)| f).collect())
248    }
249
250    async fn compact(&self, session: &str) -> Result<ContextFragment, MemoError> {
251        let mut parts: Vec<String> = Vec::new();
252        for item in self.fragments.iter() {
253            let (_k, v) = item.map_err(|e| MemoError::Storage(e.to_string()))?;
254            let frag: ContextFragment = serde_json::from_slice(&v)?;
255            if frag.session == session {
256                parts.push(format!("[{}] {}", frag.kind.as_str(), frag.content));
257            }
258        }
259        if parts.is_empty() {
260            return Err(MemoError::NotFound(session.to_string()));
261        }
262        let merged = ContextFragment::new(session, FragmentKind::Note, parts.join("\n---\n"))
263            .with_key(format!("__compact__{}", session));
264        self.memorize(merged.clone()).await?;
265        Ok(merged)
266    }
267
268    async fn get_by_key(
269        &self,
270        session: &str,
271        key: &str,
272    ) -> Result<Option<ContextFragment>, MemoError> {
273        for item in self.fragments.iter() {
274            let (_k, v) = item.map_err(|e| MemoError::Storage(e.to_string()))?;
275            let frag: ContextFragment = serde_json::from_slice(&v)?;
276            if frag.session == session && frag.key.as_deref() == Some(key) {
277                return Ok(Some(frag));
278            }
279        }
280        Ok(None)
281    }
282}
283
284pub fn now_ms() -> i64 {
285    std::time::SystemTime::now()
286        .duration_since(std::time::UNIX_EPOCH)
287        .map(|d| d.as_millis() as i64)
288        .unwrap_or(0)
289}
290
291/// Local lightweight embedding + cosine similarity for vector recall.
292///
293/// Uses the hashing trick (word 1~2-grams + char 2-grams via FNV-1a) into a
294/// fixed-dim vector, TF-normalized then L2-normalized — the same approach as
295/// the `memo` product's `memo-embed::LocalEmbedder`. Zero external/ML deps.
296pub mod embed {
297    use super::MemoError;
298    use std::collections::HashMap;
299
300    /// Text embedding seam for the local memo store.
301    pub trait Embedder: Send + Sync {
302        fn embed(&self, text: &str) -> Result<Vec<f32>, MemoError>;
303        fn dim(&self) -> usize;
304    }
305
306    /// Hashing-trick local embedder.
307    pub struct LocalEmbedder {
308        dim: usize,
309    }
310
311    impl LocalEmbedder {
312        pub fn new(dim: usize) -> Self {
313            Self { dim: dim.max(1) }
314        }
315
316        fn vectorize(&self, text: &str) -> Result<Vec<f32>, MemoError> {
317            let toks = tokenize(text);
318            if toks.is_empty() {
319                return Err(MemoError::Embedding("empty embedding text".into()));
320            }
321            let mut vec = vec![0.0f32; self.dim];
322            let mut counts: HashMap<usize, f32> = HashMap::new();
323            for t in &toks {
324                let h = hash_dim(t, self.dim);
325                *counts.entry(h).or_insert(0.0) += 1.0;
326            }
327            let max = counts.values().cloned().fold(1.0f32, f32::max);
328            for (h, c) in counts {
329                vec[h] = (c / max).sqrt();
330            }
331            let norm = vec.iter().map(|v| v * v).sum::<f32>().sqrt();
332            if norm == 0.0 {
333                return Err(MemoError::Embedding("zero-magnitude vector".into()));
334            }
335            for v in vec.iter_mut() {
336                *v /= norm;
337            }
338            Ok(vec)
339        }
340    }
341
342    impl Embedder for LocalEmbedder {
343        fn embed(&self, text: &str) -> Result<Vec<f32>, MemoError> {
344            self.vectorize(text)
345        }
346        fn dim(&self) -> usize {
347            self.dim
348        }
349    }
350
351    /// Cosine similarity; `None` on empty/mismatched dims.
352    pub fn cosine(a: &[f32], b: &[f32]) -> Option<f32> {
353        if a.is_empty() || b.is_empty() || a.len() != b.len() {
354            return None;
355        }
356        let dot = a.iter().zip(b).map(|(x, y)| x * y).sum::<f32>();
357        let na = a.iter().map(|x| x * x).sum::<f32>().sqrt();
358        let nb = b.iter().map(|x| x * x).sum::<f32>().sqrt();
359        if na == 0.0 || nb == 0.0 {
360            return Some(0.0);
361        }
362        Some(dot / (na * nb))
363    }
364
365    fn tokenize(text: &str) -> Vec<String> {
366        let lower = text.to_lowercase();
367        let mut toks: Vec<String> = Vec::new();
368        let words: Vec<&str> = lower
369            .split(|c: char| !c.is_alphanumeric())
370            .filter(|w| !w.is_empty())
371            .collect();
372        for w in &words {
373            toks.push((*w).to_string());
374        }
375        for pair in words.windows(2) {
376            toks.push(format!("{} {}", pair[0], pair[1]));
377        }
378        let chars: Vec<char> = lower.chars().filter(|c| c.is_alphanumeric()).collect();
379        for pair in chars.windows(2) {
380            toks.push(pair.iter().collect());
381        }
382        toks
383    }
384
385    fn hash_dim(s: &str, dim: usize) -> usize {
386        let mut h: u64 = 0xcbf29ce484222325;
387        for b in s.bytes() {
388            h ^= b as u64;
389            h = h.wrapping_mul(0x100000001b3);
390        }
391        (h as usize) % dim
392    }
393
394    #[cfg(test)]
395    mod tests {
396        use super::*;
397
398        #[test]
399        fn similar_text_close_vectors() {
400            let e = LocalEmbedder::new(64);
401            let a = e.embed("user prefers rust programming language").unwrap();
402            let b = e.embed("user likes rust programming language").unwrap();
403            let c = e.embed("banana smoothie recipe with ice").unwrap();
404            assert!(cosine(&a, &b).unwrap() > cosine(&a, &c).unwrap());
405        }
406
407        #[test]
408        fn deterministic_and_dim() {
409            let e = LocalEmbedder::new(64);
410            let a = e.embed("the quick brown fox").unwrap();
411            let b = e.embed("the quick brown fox").unwrap();
412            assert_eq!(a.len(), 64);
413            assert_eq!(a, b);
414        }
415
416        #[test]
417        fn empty_text_is_error() {
418            let e = LocalEmbedder::new(64);
419            assert!(e.embed("   ").is_err());
420        }
421    }
422}
423
424#[cfg(test)]
425mod tests {
426    use super::*;
427
428    #[tokio::test]
429    async fn memorize_recall_roundtrip() {
430        let store = SledMemoStore::memory().unwrap();
431        let f =
432            ContextFragment::new("s1", FragmentKind::Message, "the sky is blue").with_key("fact1");
433        store.memorize(f).await.unwrap();
434        let out = store.recall(&RecallQuery::new("s1", "sky")).await.unwrap();
435        assert!(out.iter().any(|f| f.content.contains("sky")));
436
437        let by_key = store.get_by_key("s1", "fact1").await.unwrap();
438        assert!(by_key.is_some());
439    }
440
441    #[tokio::test]
442    async fn compact_merges_session() {
443        let store = SledMemoStore::memory().unwrap();
444        store
445            .memorize(ContextFragment::new("s2", FragmentKind::Message, "a"))
446            .await
447            .unwrap();
448        store
449            .memorize(ContextFragment::new("s2", FragmentKind::Message, "b"))
450            .await
451            .unwrap();
452        let c = store.compact("s2").await.unwrap();
453        assert!(c.content.contains("a") && c.content.contains("b"));
454    }
455
456    #[tokio::test]
457    async fn vector_recall_prefers_similar() {
458        let store = SledMemoStore::memory().unwrap();
459        // Memorize without an explicit embedding; the store computes it.
460        let rust = ContextFragment::new(
461            "s3",
462            FragmentKind::Message,
463            "user prefers rust for systems programming",
464        );
465        let food = ContextFragment::new(
466            "s3",
467            FragmentKind::Message,
468            "banana smoothie recipe with ice",
469        );
470        store.memorize(rust).await.unwrap();
471        store.memorize(food).await.unwrap();
472
473        let frags = store
474            .recall(&RecallQuery::new("s3", "rust programming language"))
475            .await
476            .unwrap();
477        assert!(!frags.is_empty());
478        // The semantically closer rust memory must rank first.
479        assert!(frags[0].content.contains("rust"));
480        // Fragments are persisted with their dense vectors.
481        assert!(frags[0].embedding.is_some());
482    }
483
484    #[test]
485    fn fragment_kind_roundtrip() {
486        for k in [
487            FragmentKind::Message,
488            FragmentKind::ToolResult,
489            FragmentKind::LongTerm,
490            FragmentKind::Note,
491        ] {
492            let s = k.as_str();
493            let back: FragmentKind = s.parse().unwrap();
494            assert_eq!(k, back, "roundtrip failed for {k:?}");
495        }
496        assert!("bogus".parse::<FragmentKind>().is_err());
497    }
498
499    #[test]
500    fn recall_query_defaults() {
501        let q = RecallQuery::new("s", "x");
502        assert_eq!(q.session, "s");
503        assert_eq!(q.text, "x");
504        assert_eq!(q.top_k, 8);
505        assert!(q.kind.is_none());
506        let q = q.with_kind(FragmentKind::Note);
507        assert_eq!(q.kind, Some(FragmentKind::Note));
508    }
509
510    #[tokio::test]
511    async fn recall_filters_by_kind() {
512        let store = SledMemoStore::memory().unwrap();
513        store
514            .memorize(ContextFragment::new("s", FragmentKind::Message, "alpha"))
515            .await
516            .unwrap();
517        store
518            .memorize(ContextFragment::new("s", FragmentKind::Note, "beta"))
519            .await
520            .unwrap();
521        let msgs = store
522            .recall(&RecallQuery::new("s", "alpha").with_kind(FragmentKind::Message))
523            .await
524            .unwrap();
525        assert_eq!(msgs.len(), 1);
526        assert_eq!(msgs[0].kind, FragmentKind::Message);
527        assert!(msgs[0].content.contains("alpha"));
528    }
529
530    #[tokio::test]
531    async fn recall_respects_session() {
532        let store = SledMemoStore::memory().unwrap();
533        store
534            .memorize(ContextFragment::new("a", FragmentKind::Message, "from a"))
535            .await
536            .unwrap();
537        store
538            .memorize(ContextFragment::new("b", FragmentKind::Message, "from b"))
539            .await
540            .unwrap();
541        let out = store.recall(&RecallQuery::new("a", "from")).await.unwrap();
542        assert!(!out.is_empty());
543        assert!(out.iter().all(|f| f.session == "a"));
544        assert!(out.iter().any(|f| f.content == "from a"));
545        assert!(!out.iter().any(|f| f.content == "from b"));
546    }
547
548    #[tokio::test]
549    async fn recall_empty_query_returns_empty() {
550        let store = SledMemoStore::memory().unwrap();
551        store
552            .memorize(ContextFragment::new("s", FragmentKind::Message, "hello"))
553            .await
554            .unwrap();
555        let out = store.recall(&RecallQuery::new("s", "")).await.unwrap();
556        assert!(out.is_empty());
557    }
558
559    #[tokio::test]
560    async fn get_by_key_roundtrip_and_missing() {
561        let store = SledMemoStore::memory().unwrap();
562        let f = ContextFragment::new("s", FragmentKind::LongTerm, "fact").with_key("k1");
563        store.memorize(f).await.unwrap();
564        let got = store.get_by_key("s", "k1").await.unwrap();
565        assert!(got.is_some());
566        assert_eq!(got.unwrap().content, "fact");
567        assert!(store.get_by_key("s", "nope").await.unwrap().is_none());
568    }
569
570    #[tokio::test]
571    async fn compact_missing_session_is_not_found() {
572        let store = SledMemoStore::memory().unwrap();
573        assert!(matches!(
574            store.compact("ghost").await,
575            Err(MemoError::NotFound(_))
576        ));
577    }
578
579    #[tokio::test]
580    async fn memorize_populates_embedding() {
581        let store = SledMemoStore::memory().unwrap();
582        let f = ContextFragment::new("s", FragmentKind::Message, "rust programming").with_key("ke");
583        store.memorize(f).await.unwrap();
584        let got = store.get_by_key("s", "ke").await.unwrap().unwrap();
585        assert!(got.embedding.is_some());
586        assert_eq!(got.embedding.unwrap().len(), 64);
587    }
588
589    #[tokio::test]
590    async fn recall_exact_key_match_ranks_first() {
591        let store = SledMemoStore::memory().unwrap();
592        store
593            .memorize(
594                ContextFragment::new("s", FragmentKind::Message, "unrelated content")
595                    .with_key("fact1"),
596            )
597            .await
598            .unwrap();
599        let out = store.recall(&RecallQuery::new("s", "fact1")).await.unwrap();
600        assert!(!out.is_empty());
601        assert_eq!(out[0].key.as_deref(), Some("fact1"));
602    }
603}