Skip to main content

agent_core/context/
mod.rs

1//! Context memory contract shared by every runtime tier.
2//!
3//! Historically this lived in `aria-agent-memo` as `MemoStore` (local/embedded).
4//! The contract is lifted into `agent-core` so the cloud can ship a
5//! **Postgres + pgvector** implementation (see
6//! `crates/aria-agent-cloud/src/context/pg.rs`) while the on-device SDK keeps
7//! the **aria memo** implementation (`aria-agent-memo`). The two share:
8//!
9//! * [`ContextStore`] — the storage contract (`memorize` / `recall` / `compact`
10//!   / `get_by_key`),
11//! * [`ContextFragment`] / [`FragmentKind`] / [`RecallQuery`] — the data shapes,
12//! * [`embed`] — the zero-dependency local embedder used for semantic recall,
13//! * [`score_fragment`] / [`rank`] — the keyword + vector blending rules.
14//!
15//! See `docs/adr/0010-context-storage-pgvector.md`.
16
17use async_trait::async_trait;
18use embed::Embedder as _;
19use serde::{Deserialize, Serialize};
20use std::collections::HashMap;
21use std::str::FromStr;
22use std::sync::{Arc, RwLock};
23use thiserror::Error;
24use uuid::Uuid;
25
26/// Dimension of the dense embedding produced by [`embed::LocalEmbedder`].
27///
28/// Fixed so every backend (Postgres `vector(256)`, aria memo, in-memory)
29/// agrees on the vector width; [`embed::cosine`] returns `None` on a width
30/// mismatch.
31pub const EMBED_DIM: usize = 256;
32
33/// Kinds of context fragments a store can hold.
34#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq)]
35pub enum FragmentKind {
36    /// A turn of the conversation (user message or assistant reply).
37    Message,
38    /// A tool call result that should be remembered across turns.
39    ToolResult,
40    /// An explicitly externalized long-term memory written via `memorize`.
41    LongTerm,
42    /// A system / scratch note.
43    Note,
44}
45
46impl FragmentKind {
47    pub fn as_str(&self) -> &'static str {
48        match self {
49            FragmentKind::Message => "message",
50            FragmentKind::ToolResult => "tool_result",
51            FragmentKind::LongTerm => "long_term",
52            FragmentKind::Note => "note",
53        }
54    }
55}
56
57impl FromStr for FragmentKind {
58    type Err = String;
59
60    fn from_str(s: &str) -> Result<Self, Self::Err> {
61        match s {
62            "message" => Ok(FragmentKind::Message),
63            "tool_result" => Ok(FragmentKind::ToolResult),
64            "long_term" => Ok(FragmentKind::LongTerm),
65            "note" => Ok(FragmentKind::Note),
66            other => Err(format!("unknown fragment kind: {other}")),
67        }
68    }
69}
70
71/// A single unit of remembered context.
72#[derive(Debug, Clone, Serialize, Deserialize)]
73pub struct ContextFragment {
74    pub id: String,
75    pub session: String,
76    /// Optional explicit key (used by `Session::memorize`/`recall`).
77    pub key: Option<String>,
78    pub kind: FragmentKind,
79    pub content: String,
80    /// Unix epoch milliseconds.
81    pub created_at: i64,
82    /// Optional dense vector for semantic (vector) recall. Populated by the
83    /// [`embed::LocalEmbedder`] when a fragment is memorized without one.
84    pub embedding: Option<Vec<f32>>,
85}
86
87impl ContextFragment {
88    pub fn new(session: &str, kind: FragmentKind, content: impl Into<String>) -> Self {
89        Self {
90            id: Uuid::new_v4().to_string(),
91            session: session.to_string(),
92            key: None,
93            kind,
94            content: content.into(),
95            created_at: now_ms(),
96            embedding: None,
97        }
98    }
99
100    pub fn with_key(mut self, key: impl Into<String>) -> Self {
101        self.key = Some(key.into());
102        self
103    }
104}
105
106/// Query used by [`ContextStore::recall`].
107#[derive(Debug, Clone)]
108pub struct RecallQuery {
109    pub session: String,
110    pub text: String,
111    pub top_k: usize,
112    pub kind: Option<FragmentKind>,
113}
114
115impl RecallQuery {
116    pub fn new(session: &str, text: impl Into<String>) -> Self {
117        Self {
118            session: session.to_string(),
119            text: text.into(),
120            top_k: 8,
121            kind: None,
122        }
123    }
124
125    pub fn with_kind(mut self, kind: FragmentKind) -> Self {
126        self.kind = Some(kind);
127        self
128    }
129}
130
131#[derive(Debug, Error)]
132pub enum ContextError {
133    #[error("storage error: {0}")]
134    Storage(String),
135    #[error("serialization error: {0}")]
136    Serialization(#[from] serde_json::Error),
137    #[error("not found: {0}")]
138    NotFound(String),
139    #[error("embedding error: {0}")]
140    Embedding(String),
141}
142
143/// Where a memory context lives.
144///
145/// * `Cloud` — the agent-cloud store (Postgres + pgvector), shared across
146///   devices and instances.
147/// * `Local` — the on-device **aria memo** store (SQLite), works offline.
148/// * `Both` — write to both, read merged (see [`CompositeContextStore`]).
149#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
150pub enum MemoryBackend {
151    #[default]
152    Cloud,
153    Local,
154    Both,
155}
156
157impl MemoryBackend {
158    pub fn as_str(&self) -> &'static str {
159        match self {
160            MemoryBackend::Cloud => "cloud",
161            MemoryBackend::Local => "local",
162            MemoryBackend::Both => "both",
163        }
164    }
165}
166
167impl std::str::FromStr for MemoryBackend {
168    type Err = String;
169
170    fn from_str(s: &str) -> Result<Self, Self::Err> {
171        match s.trim().to_ascii_lowercase().as_str() {
172            "cloud" => Ok(MemoryBackend::Cloud),
173            "local" => Ok(MemoryBackend::Local),
174            "both" => Ok(MemoryBackend::Both),
175            other => Err(format!("unknown memory backend: {other}")),
176        }
177    }
178}
179
180impl std::fmt::Display for MemoryBackend {
181    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
182        f.write_str(self.as_str())
183    }
184}
185
186/// The unified context-memory contract. Everything that needs conversational
187/// or long-term context goes through this trait.
188///
189/// Implementations:
190/// * `aria-agent-memo::MemoContextStore` — on-device / local (aria memo, SQLite),
191/// * `aria-agent-cloud::context::PgContextStore` — cloud (Postgres + pgvector),
192/// * [`CompositeContextStore`] — `both`: writes to two stores, reads merged.
193#[async_trait]
194pub trait ContextStore: Send + Sync {
195    /// Persist a fragment.
196    async fn memorize(&self, frag: ContextFragment) -> Result<(), ContextError>;
197    /// Retrieve the most relevant fragments for a query.
198    async fn recall(&self, query: &RecallQuery) -> Result<Vec<ContextFragment>, ContextError>;
199    /// Merge a session's fragments into a single compact fragment.
200    async fn compact(&self, session: &str) -> Result<ContextFragment, ContextError>;
201    /// Lookup an explicit long-term memory by key (used by SDK `recall`).
202    async fn get_by_key(
203        &self,
204        session: &str,
205        key: &str,
206    ) -> Result<Option<ContextFragment>, ContextError>;
207    /// Every fragment of a session, oldest first (used by `compact` and by the
208    /// session memory listing).
209    async fn list_session(
210        &self,
211        session: &str,
212        top_k: usize,
213    ) -> Result<Vec<ContextFragment>, ContextError>;
214}
215
216/// Embed a fragment's content when it carries no vector yet. Shared by every
217/// backend so stored vectors stay comparable.
218pub fn ensure_embedding(frag: &mut ContextFragment) {
219    if frag.embedding.is_none() && !frag.content.trim().is_empty() {
220        if let Ok(v) = embed::LocalEmbedder::new(EMBED_DIM).embed(&frag.content) {
221            frag.embedding = Some(v);
222        }
223    }
224}
225
226/// Keyword relevance of `frag` against `query_text`:
227/// * `1.0` when the fragment's explicit key matches the query exactly,
228/// * otherwise a length-normalized count of query words found in the content.
229pub fn keyword_score(query_text: &str, frag: &ContextFragment) -> f32 {
230    let q = query_text.to_lowercase();
231    let mut kw = 0.0f32;
232    if frag.key.as_deref() == Some(query_text) {
233        kw = 1.0;
234    }
235    let hay = frag.content.to_lowercase();
236    if !q.is_empty() && hay.contains(&q) {
237        let hits = q
238            .split_whitespace()
239            .filter(|w| !w.is_empty() && hay.contains(*w))
240            .count() as f32;
241        kw = kw.max(hits / (hay.len() as f32).max(1.0).log10());
242    }
243    kw
244}
245
246/// Blend semantic (vector) similarity with the keyword score: `0.7 * cosine +
247/// 0.3 * keyword`. Pure keyword when a vector is missing on either side.
248pub fn score_fragment(
249    query_text: &str,
250    query_embedding: Option<&[f32]>,
251    frag: &ContextFragment,
252) -> f32 {
253    let kw = keyword_score(query_text, frag);
254    match (query_embedding, frag.embedding.as_deref()) {
255        (Some(a), Some(b)) => match embed::cosine(a, b) {
256            Some(v) => 0.7 * v + 0.3 * kw,
257            None => kw,
258        },
259        _ => kw,
260    }
261}
262
263/// Score, sort (descending) and truncate candidates to `top_k`. Fragments that
264/// score `0.0` are dropped — an empty query therefore recalls nothing.
265pub fn rank(
266    query_text: &str,
267    query_embedding: Option<&[f32]>,
268    mut frags: Vec<ContextFragment>,
269    top_k: usize,
270) -> Vec<ContextFragment> {
271    let mut scored: Vec<(f32, ContextFragment)> = frags
272        .drain(..)
273        .map(|f| (score_fragment(query_text, query_embedding, &f), f))
274        .filter(|(s, _)| *s > 0.0)
275        .collect();
276    scored.sort_by(|a, b| b.0.partial_cmp(&a.0).unwrap_or(std::cmp::Ordering::Equal));
277    scored.truncate(top_k);
278    scored.into_iter().map(|(_, f)| f).collect()
279}
280
281pub fn now_ms() -> i64 {
282    std::time::SystemTime::now()
283        .duration_since(std::time::UNIX_EPOCH)
284        .map(|d| d.as_millis() as i64)
285        .unwrap_or(0)
286}
287
288/// Merge results from two backends: a fragment is a duplicate when **either**
289/// its `id` or its `key` was already seen; the freshest copy wins.
290pub fn merge_fragments(a: Vec<ContextFragment>, b: Vec<ContextFragment>) -> Vec<ContextFragment> {
291    let mut out: Vec<ContextFragment> = Vec::with_capacity(a.len() + b.len());
292    let mut index: HashMap<String, usize> = HashMap::new();
293    for frag in a.into_iter().chain(b) {
294        let ids: Vec<String> = [
295            Some(frag.id.clone()).filter(|s| !s.is_empty()),
296            frag.key.clone(),
297        ]
298        .into_iter()
299        .flatten()
300        .collect();
301        let existing = ids.iter().find_map(|k| index.get(k).copied());
302        match existing {
303            Some(i) if frag.created_at > out[i].created_at => {
304                for k in &ids {
305                    index.insert(k.clone(), i);
306                }
307                out[i] = frag;
308            }
309            Some(_) => {}
310            None => {
311                let i = out.len();
312                for k in &ids {
313                    index.insert(k.clone(), i);
314                }
315                out.push(frag);
316            }
317        }
318    }
319    out
320}
321
322/// The `both` backend: a local (aria memo) store plus a cloud store.
323///
324/// * `memorize` writes to both (local first, so an offline write still lands);
325///   a single-side failure is logged and tolerated — only a total failure
326///   returns an error.
327/// * `recall` queries both, merges, dedupes and re-ranks with the shared
328///   [`rank`] scoring so ordering matches a single-backend run.
329/// * `get_by_key` prefers the cloud copy and falls back to local.
330pub struct CompositeContextStore {
331    local: Arc<dyn ContextStore>,
332    cloud: Arc<dyn ContextStore>,
333}
334
335impl CompositeContextStore {
336    pub fn new(local: Arc<dyn ContextStore>, cloud: Arc<dyn ContextStore>) -> Arc<Self> {
337        Arc::new(Self { local, cloud })
338    }
339
340    pub fn local(&self) -> &Arc<dyn ContextStore> {
341        &self.local
342    }
343
344    pub fn cloud(&self) -> &Arc<dyn ContextStore> {
345        &self.cloud
346    }
347}
348
349#[async_trait]
350impl ContextStore for CompositeContextStore {
351    async fn memorize(&self, frag: ContextFragment) -> Result<(), ContextError> {
352        let local = self.local.memorize(frag.clone()).await;
353        let cloud = self.cloud.memorize(frag).await;
354        match (local, cloud) {
355            (Ok(()), Ok(())) => Ok(()),
356            (Ok(()), Err(e)) => {
357                tracing::warn!("context: cloud memorize failed, kept local copy: {e}");
358                Ok(())
359            }
360            (Err(e), Ok(())) => {
361                tracing::warn!("context: local memorize failed, kept cloud copy: {e}");
362                Ok(())
363            }
364            (Err(e), Err(_)) => Err(e),
365        }
366    }
367
368    async fn recall(&self, query: &RecallQuery) -> Result<Vec<ContextFragment>, ContextError> {
369        let local = self.local.recall(query).await;
370        let cloud = self.cloud.recall(query).await;
371        let (local, cloud) = match (local, cloud) {
372            (Ok(l), Ok(c)) => (l, c),
373            (Ok(l), Err(e)) => {
374                tracing::warn!("context: cloud recall failed, using local only: {e}");
375                (l, Vec::new())
376            }
377            (Err(e), Ok(c)) => {
378                tracing::warn!("context: local recall failed, using cloud only: {e}");
379                (Vec::new(), c)
380            }
381            (Err(e), Err(_)) => return Err(e),
382        };
383        let merged = merge_fragments(local, cloud);
384        let q_emb = embed::LocalEmbedder::new(EMBED_DIM).embed(&query.text).ok();
385        Ok(rank(&query.text, q_emb.as_deref(), merged, query.top_k))
386    }
387
388    async fn compact(&self, session: &str) -> Result<ContextFragment, ContextError> {
389        // Compact from whichever side has the session; the cloud copy wins when
390        // both do, and the result is persisted back to both.
391        let compacted = match self.cloud.compact(session).await {
392            Ok(c) => c,
393            Err(e) => {
394                tracing::warn!("context: cloud compact failed, using local: {e}");
395                match self.local.compact(session).await {
396                    Ok(c) => c,
397                    Err(e) => return Err(e),
398                }
399            }
400        };
401        self.memorize(compacted.clone()).await?;
402        Ok(compacted)
403    }
404
405    async fn get_by_key(
406        &self,
407        session: &str,
408        key: &str,
409    ) -> Result<Option<ContextFragment>, ContextError> {
410        match self.cloud.get_by_key(session, key).await {
411            Ok(Some(f)) => Ok(Some(f)),
412            Ok(None) => self.local.get_by_key(session, key).await,
413            Err(e) => {
414                tracing::warn!("context: cloud get_by_key failed, trying local: {e}");
415                self.local.get_by_key(session, key).await
416            }
417        }
418    }
419
420    async fn list_session(
421        &self,
422        session: &str,
423        top_k: usize,
424    ) -> Result<Vec<ContextFragment>, ContextError> {
425        let local = self.local.list_session(session, top_k).await;
426        let cloud = self.cloud.list_session(session, top_k).await;
427        let (local, cloud) = match (local, cloud) {
428            (Ok(l), Ok(c)) => (l, c),
429            (Ok(l), Err(e)) => {
430                tracing::warn!("context: cloud list failed, using local only: {e}");
431                (l, Vec::new())
432            }
433            (Err(e), Ok(c)) => {
434                tracing::warn!("context: local list failed, using cloud only: {e}");
435                (Vec::new(), c)
436            }
437            (Err(e), Err(_)) => return Err(e),
438        };
439        let mut merged = merge_fragments(local, cloud);
440        merged.sort_by_key(|f| f.created_at);
441        merged.truncate(top_k);
442        Ok(merged)
443    }
444}
445
446/// In-process [`ContextStore`] used by tests and by the SDK default.
447///
448/// Kept in `agent-core` so core's tests need no on-device database (and no
449/// dependency on the embedded store crate).
450pub struct MemoryContextStore {
451    inner: RwLock<HashMap<String, ContextFragment>>,
452}
453
454impl MemoryContextStore {
455    pub fn new() -> Arc<Self> {
456        Arc::new(Self {
457            inner: RwLock::new(HashMap::new()),
458        })
459    }
460}
461
462impl Default for MemoryContextStore {
463    fn default() -> Self {
464        Self {
465            inner: RwLock::new(HashMap::new()),
466        }
467    }
468}
469
470#[async_trait]
471impl ContextStore for MemoryContextStore {
472    async fn memorize(&self, mut frag: ContextFragment) -> Result<(), ContextError> {
473        ensure_embedding(&mut frag);
474        let mut g = self
475            .inner
476            .write()
477            .map_err(|e| ContextError::Storage(format!("memory store lock poisoned: {e}")))?;
478        g.insert(frag.id.clone(), frag);
479        Ok(())
480    }
481
482    async fn recall(&self, query: &RecallQuery) -> Result<Vec<ContextFragment>, ContextError> {
483        let g = self
484            .inner
485            .read()
486            .map_err(|e| ContextError::Storage(format!("memory store lock poisoned: {e}")))?;
487        if query.text.trim().is_empty() {
488            return Ok(Vec::new());
489        }
490        let q_emb = if query.text.trim().is_empty() {
491            None
492        } else {
493            embed::LocalEmbedder::new(EMBED_DIM).embed(&query.text).ok()
494        };
495        let candidates: Vec<ContextFragment> = g
496            .values()
497            .filter(|f| f.session == query.session)
498            .filter(|f| query.kind.map(|k| f.kind == k).unwrap_or(true))
499            .cloned()
500            .collect();
501        Ok(rank(&query.text, q_emb.as_deref(), candidates, query.top_k))
502    }
503
504    async fn compact(&self, session: &str) -> Result<ContextFragment, ContextError> {
505        let parts: Vec<String> = {
506            let g = self
507                .inner
508                .read()
509                .map_err(|e| ContextError::Storage(format!("memory store lock poisoned: {e}")))?;
510            g.values()
511                .filter(|f| f.session == session)
512                .map(|f| format!("[{}] {}", f.kind.as_str(), f.content))
513                .collect()
514        };
515        if parts.is_empty() {
516            return Err(ContextError::NotFound(session.to_string()));
517        }
518        let merged = ContextFragment::new(session, FragmentKind::Note, parts.join("\n---\n"))
519            .with_key(format!("__compact__{session}"));
520        self.memorize(merged.clone()).await?;
521        Ok(merged)
522    }
523
524    async fn get_by_key(
525        &self,
526        session: &str,
527        key: &str,
528    ) -> Result<Option<ContextFragment>, ContextError> {
529        let g = self
530            .inner
531            .read()
532            .map_err(|e| ContextError::Storage(format!("memory store lock poisoned: {e}")))?;
533        Ok(g.values()
534            .find(|f| f.session == session && f.key.as_deref() == Some(key))
535            .cloned())
536    }
537
538    async fn list_session(
539        &self,
540        session: &str,
541        top_k: usize,
542    ) -> Result<Vec<ContextFragment>, ContextError> {
543        let g = self
544            .inner
545            .read()
546            .map_err(|e| ContextError::Storage(format!("memory store lock poisoned: {e}")))?;
547        let mut out: Vec<ContextFragment> = g
548            .values()
549            .filter(|f| f.session == session)
550            .cloned()
551            .collect();
552        out.sort_by_key(|f| f.created_at);
553        out.truncate(top_k);
554        Ok(out)
555    }
556}
557
558/// Local lightweight embedding + cosine similarity for vector recall.
559///
560/// Uses the hashing trick (word 1~2-grams + char 2-grams via FNV-1a) into a
561/// fixed-dim vector, TF-normalized then L2-normalized. Zero external/ML deps.
562pub mod embed {
563    use super::ContextError;
564    use std::collections::HashMap;
565
566    /// Text embedding seam.
567    pub trait Embedder: Send + Sync {
568        fn embed(&self, text: &str) -> Result<Vec<f32>, ContextError>;
569        fn dim(&self) -> usize;
570    }
571
572    /// Hashing-trick local embedder.
573    pub struct LocalEmbedder {
574        dim: usize,
575    }
576
577    impl LocalEmbedder {
578        pub fn new(dim: usize) -> Self {
579            Self { dim: dim.max(1) }
580        }
581
582        fn vectorize(&self, text: &str) -> Result<Vec<f32>, ContextError> {
583            let toks = tokenize(text);
584            if toks.is_empty() {
585                return Err(ContextError::Embedding("empty embedding text".into()));
586            }
587            let mut vec = vec![0.0f32; self.dim];
588            let mut counts: HashMap<usize, f32> = HashMap::new();
589            for t in &toks {
590                let h = hash_dim(t, self.dim);
591                *counts.entry(h).or_insert(0.0) += 1.0;
592            }
593            let max = counts.values().cloned().fold(1.0f32, f32::max);
594            for (h, c) in counts {
595                vec[h] = (c / max).sqrt();
596            }
597            let norm = vec.iter().map(|v| v * v).sum::<f32>().sqrt();
598            if norm == 0.0 {
599                return Err(ContextError::Embedding("zero-magnitude vector".into()));
600            }
601            for v in vec.iter_mut() {
602                *v /= norm;
603            }
604            Ok(vec)
605        }
606    }
607
608    impl Embedder for LocalEmbedder {
609        fn embed(&self, text: &str) -> Result<Vec<f32>, ContextError> {
610            self.vectorize(text)
611        }
612        fn dim(&self) -> usize {
613            self.dim
614        }
615    }
616
617    /// Cosine similarity; `None` on empty/mismatched dims.
618    pub fn cosine(a: &[f32], b: &[f32]) -> Option<f32> {
619        if a.is_empty() || b.is_empty() || a.len() != b.len() {
620            return None;
621        }
622        let dot = a.iter().zip(b).map(|(x, y)| x * y).sum::<f32>();
623        let na = a.iter().map(|x| x * x).sum::<f32>().sqrt();
624        let nb = b.iter().map(|x| x * x).sum::<f32>().sqrt();
625        if na == 0.0 || nb == 0.0 {
626            return Some(0.0);
627        }
628        Some(dot / (na * nb))
629    }
630
631    fn tokenize(text: &str) -> Vec<String> {
632        let lower = text.to_lowercase();
633        let mut toks: Vec<String> = Vec::new();
634        let words: Vec<&str> = lower
635            .split(|c: char| !c.is_alphanumeric())
636            .filter(|w| !w.is_empty())
637            .collect();
638        for w in &words {
639            toks.push((*w).to_string());
640        }
641        for pair in words.windows(2) {
642            toks.push(format!("{} {}", pair[0], pair[1]));
643        }
644        let chars: Vec<char> = lower.chars().filter(|c| c.is_alphanumeric()).collect();
645        for pair in chars.windows(2) {
646            toks.push(pair.iter().collect());
647        }
648        toks
649    }
650
651    fn hash_dim(s: &str, dim: usize) -> usize {
652        let mut h: u64 = 0xcbf29ce484222325;
653        for b in s.bytes() {
654            h ^= b as u64;
655            h = h.wrapping_mul(0x100000001b3);
656        }
657        (h as usize) % dim
658    }
659
660    #[cfg(test)]
661    mod tests {
662        use super::*;
663
664        #[test]
665        fn similar_text_close_vectors() {
666            let e = LocalEmbedder::new(super::super::EMBED_DIM);
667            let a = e.embed("user prefers rust programming language").unwrap();
668            let b = e.embed("user likes rust programming language").unwrap();
669            let c = e.embed("banana smoothie recipe with ice").unwrap();
670            assert!(cosine(&a, &b).unwrap() > cosine(&a, &c).unwrap());
671        }
672
673        #[test]
674        fn deterministic_and_dim() {
675            let e = LocalEmbedder::new(64);
676            let a = e.embed("the quick brown fox").unwrap();
677            let b = e.embed("the quick brown fox").unwrap();
678            assert_eq!(a.len(), 64);
679            assert_eq!(a, b);
680        }
681
682        #[test]
683        fn empty_text_is_error() {
684            let e = LocalEmbedder::new(64);
685            assert!(e.embed("   ").is_err());
686        }
687    }
688}
689
690#[cfg(test)]
691mod tests {
692    use super::*;
693
694    #[test]
695    fn fragment_kind_roundtrip() {
696        for k in [
697            FragmentKind::Message,
698            FragmentKind::ToolResult,
699            FragmentKind::LongTerm,
700            FragmentKind::Note,
701        ] {
702            let s = k.as_str();
703            let back: FragmentKind = s.parse().unwrap();
704            assert_eq!(k, back, "roundtrip failed for {k:?}");
705        }
706        assert!("bogus".parse::<FragmentKind>().is_err());
707    }
708
709    #[test]
710    fn recall_query_defaults() {
711        let q = RecallQuery::new("s", "x");
712        assert_eq!(q.session, "s");
713        assert_eq!(q.text, "x");
714        assert_eq!(q.top_k, 8);
715        assert!(q.kind.is_none());
716        let q = q.with_kind(FragmentKind::Note);
717        assert_eq!(q.kind, Some(FragmentKind::Note));
718    }
719
720    #[test]
721    fn ensure_embedding_fills_vector_of_fixed_dim() {
722        let mut f = ContextFragment::new("s", FragmentKind::Message, "rust programming");
723        ensure_embedding(&mut f);
724        let v = f.embedding.expect("embedding populated");
725        assert_eq!(v.len(), EMBED_DIM);
726    }
727
728    #[test]
729    fn ensure_embedding_leaves_empty_content_without_vector() {
730        let mut f = ContextFragment::new("s", FragmentKind::Message, "   ");
731        ensure_embedding(&mut f);
732        assert!(f.embedding.is_none());
733    }
734
735    #[test]
736    fn score_prefers_exact_key_then_vector() {
737        let keyed =
738            ContextFragment::new("s", FragmentKind::LongTerm, "unrelated").with_key("fact1");
739        assert_eq!(keyword_score("fact1", &keyed), 1.0);
740
741        let a = ContextFragment::new("s", FragmentKind::Message, "rust programming language");
742        let b = ContextFragment::new("s", FragmentKind::Message, "banana smoothie recipe");
743        let q = embed::LocalEmbedder::new(EMBED_DIM)
744            .embed("rust programming")
745            .unwrap();
746        let sa = score_fragment("rust programming", Some(&q), &a);
747        let sb = score_fragment("rust programming", Some(&q), &b);
748        // Keyword-only scoring (no vectors on the fragments) still ranks the
749        // substring hit above the unrelated text.
750        assert!(score_fragment("rust", None, &a) > score_fragment("rust", None, &b));
751        assert!(sa > sb);
752    }
753
754    #[test]
755    fn rank_drops_zero_scores_and_truncates() {
756        let a = ContextFragment::new("s", FragmentKind::Message, "alpha text");
757        let b = ContextFragment::new("s", FragmentKind::Message, "beta text");
758        let c = ContextFragment::new("s", FragmentKind::Message, "unrelated");
759        let out = rank("alpha", None, vec![a, b, c], 8);
760        assert_eq!(out.len(), 1);
761        assert!(out[0].content.contains("alpha"));
762
763        let many: Vec<ContextFragment> = (0..10)
764            .map(|i| ContextFragment::new("s", FragmentKind::Message, format!("common {i}")))
765            .collect();
766        assert_eq!(rank("common", None, many, 3).len(), 3);
767    }
768
769    #[tokio::test]
770    async fn memory_store_roundtrip_and_compact() {
771        let store = MemoryContextStore::new();
772        store
773            .memorize(
774                ContextFragment::new("s1", FragmentKind::Message, "hello world").with_key("k"),
775            )
776            .await
777            .unwrap();
778        let out = store
779            .recall(&RecallQuery::new("s1", "hello"))
780            .await
781            .unwrap();
782        assert!(out.iter().any(|f| f.content.contains("hello")));
783        assert!(store.get_by_key("s1", "k").await.unwrap().is_some());
784        assert!(store.get_by_key("s1", "missing").await.unwrap().is_none());
785
786        // Empty query recalls nothing (abnormal path).
787        assert!(store
788            .recall(&RecallQuery::new("s1", ""))
789            .await
790            .unwrap()
791            .is_empty());
792
793        let c = store.compact("s1").await.unwrap();
794        assert!(c.content.contains("hello world"));
795        assert!(matches!(
796            store.compact("ghost").await,
797            Err(ContextError::NotFound(_))
798        ));
799    }
800
801    #[tokio::test]
802    async fn memory_store_isolates_sessions_and_kinds() {
803        let store = MemoryContextStore::new();
804        store
805            .memorize(ContextFragment::new("a", FragmentKind::Message, "from a"))
806            .await
807            .unwrap();
808        store
809            .memorize(ContextFragment::new("b", FragmentKind::Message, "from b"))
810            .await
811            .unwrap();
812        let out = store.recall(&RecallQuery::new("a", "from")).await.unwrap();
813        assert!(out.iter().all(|f| f.session == "a"));
814
815        let notes = store
816            .recall(&RecallQuery::new("a", "from").with_kind(FragmentKind::Note))
817            .await
818            .unwrap();
819        assert!(notes.is_empty());
820    }
821
822    // --- MemoryBackend / composite ("both") backend ---
823
824    struct FailingStore;
825
826    #[async_trait]
827    impl ContextStore for FailingStore {
828        async fn memorize(&self, _frag: ContextFragment) -> Result<(), ContextError> {
829            Err(ContextError::Storage("boom".into()))
830        }
831        async fn recall(&self, _query: &RecallQuery) -> Result<Vec<ContextFragment>, ContextError> {
832            Err(ContextError::Storage("boom".into()))
833        }
834        async fn compact(&self, session: &str) -> Result<ContextFragment, ContextError> {
835            Err(ContextError::NotFound(session.to_string()))
836        }
837        async fn get_by_key(
838            &self,
839            _session: &str,
840            _key: &str,
841        ) -> Result<Option<ContextFragment>, ContextError> {
842            Err(ContextError::Storage("boom".into()))
843        }
844        async fn list_session(
845            &self,
846            _session: &str,
847            _top_k: usize,
848        ) -> Result<Vec<ContextFragment>, ContextError> {
849            Err(ContextError::Storage("boom".into()))
850        }
851    }
852
853    #[test]
854    fn memory_backend_parses_and_roundtrips() {
855        assert_eq!(MemoryBackend::default(), MemoryBackend::Cloud);
856        for (raw, expected) in [
857            ("cloud", MemoryBackend::Cloud),
858            ("local", MemoryBackend::Local),
859            ("both", MemoryBackend::Both),
860            (" LOCAL ", MemoryBackend::Local),
861        ] {
862            let parsed: MemoryBackend = raw.parse().unwrap();
863            assert_eq!(parsed, expected);
864            assert_eq!(parsed.as_str(), expected.as_str());
865            assert_eq!(parsed.to_string(), expected.as_str());
866        }
867        assert!("nope".parse::<MemoryBackend>().is_err());
868    }
869
870    #[test]
871    fn merge_fragments_dedupes_by_id_keeping_the_freshest() {
872        let mut older = ContextFragment::new("s", FragmentKind::Message, "old");
873        older.id = "id-1".into();
874        older.created_at = 1;
875        let mut newer = ContextFragment::new("s", FragmentKind::Message, "new");
876        newer.id = "id-1".into();
877        newer.created_at = 2;
878        let mut other = ContextFragment::new("s", FragmentKind::Note, "other");
879        other.id = "id-2".into();
880
881        let merged = merge_fragments(vec![older, other.clone()], vec![newer]);
882        assert_eq!(merged.len(), 2);
883        assert!(merged.iter().any(|f| f.id == "id-1" && f.content == "new"));
884        assert!(merged
885            .iter()
886            .any(|f| f.id == "id-2" && f.content == "other"));
887    }
888
889    #[test]
890    fn merge_fragments_falls_back_to_key_when_id_missing() {
891        let a = ContextFragment::new("s", FragmentKind::LongTerm, "v1").with_key("k");
892        let mut b = ContextFragment::new("s", FragmentKind::LongTerm, "v2").with_key("k");
893        b.created_at = a.created_at + 5;
894        let merged = merge_fragments(vec![a], vec![b]);
895        assert_eq!(merged.len(), 1);
896        assert_eq!(merged[0].content, "v2");
897    }
898
899    #[tokio::test]
900    async fn composite_writes_to_both_backends() {
901        let local = MemoryContextStore::new();
902        let cloud = MemoryContextStore::new();
903        let composite = CompositeContextStore::new(local.clone(), cloud.clone());
904
905        composite
906            .memorize(ContextFragment::new("s", FragmentKind::LongTerm, "fact"))
907            .await
908            .unwrap();
909
910        let q = RecallQuery::new("s", "fact");
911        assert_eq!(local.recall(&q).await.unwrap().len(), 1);
912        assert_eq!(cloud.recall(&q).await.unwrap().len(), 1);
913    }
914
915    #[tokio::test]
916    async fn composite_tolerates_one_failing_write() {
917        let local = MemoryContextStore::new();
918        let composite = CompositeContextStore::new(local.clone(), Arc::new(FailingStore));
919        composite
920            .memorize(ContextFragment::new("s", FragmentKind::Message, "kept"))
921            .await
922            .unwrap();
923        assert_eq!(
924            local
925                .recall(&RecallQuery::new("s", "kept"))
926                .await
927                .unwrap()
928                .len(),
929            1
930        );
931
932        let cloud_only =
933            CompositeContextStore::new(Arc::new(FailingStore), MemoryContextStore::new());
934        cloud_only
935            .memorize(ContextFragment::new("s", FragmentKind::Message, "kept"))
936            .await
937            .unwrap();
938    }
939
940    #[tokio::test]
941    async fn composite_errors_when_every_write_fails() {
942        let composite = CompositeContextStore::new(Arc::new(FailingStore), Arc::new(FailingStore));
943        let res = composite
944            .memorize(ContextFragment::new("s", FragmentKind::Message, "x"))
945            .await;
946        assert!(matches!(res, Err(ContextError::Storage(_))));
947    }
948
949    #[tokio::test]
950    async fn composite_recall_merges_and_dedupes() {
951        let local = MemoryContextStore::new();
952        let cloud = MemoryContextStore::new();
953        let composite = CompositeContextStore::new(local.clone(), cloud.clone());
954
955        // Same fragment id written on both sides must appear once.
956        let mut frag = ContextFragment::new("s", FragmentKind::LongTerm, "shared memory");
957        frag.id = "shared".into();
958        local.memorize(frag.clone()).await.unwrap();
959        cloud.memorize(frag).await.unwrap();
960        // Unique to each side (both match the query so both are recalled).
961        local
962            .memorize(ContextFragment::new(
963                "s",
964                FragmentKind::Note,
965                "local memory",
966            ))
967            .await
968            .unwrap();
969        cloud
970            .memorize(ContextFragment::new(
971                "s",
972                FragmentKind::Note,
973                "cloud memory",
974            ))
975            .await
976            .unwrap();
977
978        let out = composite
979            .recall(&RecallQuery {
980                session: "s".into(),
981                text: "memory".into(),
982                top_k: 10,
983                kind: None,
984            })
985            .await
986            .unwrap();
987        let shared_hits = out.iter().filter(|f| f.id == "shared").count();
988        assert_eq!(shared_hits, 1, "duplicate fragments must be collapsed");
989        assert!(out.iter().any(|f| f.content == "local memory"));
990        assert!(out.iter().any(|f| f.content == "cloud memory"));
991    }
992
993    #[tokio::test]
994    async fn composite_recall_survives_one_failing_backend() {
995        let local = MemoryContextStore::new();
996        local
997            .memorize(ContextFragment::new(
998                "s",
999                FragmentKind::Message,
1000                "offline fact",
1001            ))
1002            .await
1003            .unwrap();
1004        let composite = CompositeContextStore::new(local, Arc::new(FailingStore));
1005        let out = composite
1006            .recall(&RecallQuery::new("s", "offline fact"))
1007            .await
1008            .unwrap();
1009        assert_eq!(out.len(), 1);
1010
1011        let both_broken =
1012            CompositeContextStore::new(Arc::new(FailingStore), Arc::new(FailingStore));
1013        assert!(both_broken
1014            .recall(&RecallQuery::new("s", "x"))
1015            .await
1016            .is_err());
1017    }
1018
1019    #[tokio::test]
1020    async fn composite_get_by_key_prefers_cloud_then_local() {
1021        let local = MemoryContextStore::new();
1022        let cloud = MemoryContextStore::new();
1023        local
1024            .memorize(ContextFragment::new("s", FragmentKind::LongTerm, "from local").with_key("k"))
1025            .await
1026            .unwrap();
1027        let composite = CompositeContextStore::new(local.clone(), cloud.clone());
1028        assert_eq!(
1029            composite
1030                .get_by_key("s", "k")
1031                .await
1032                .unwrap()
1033                .unwrap()
1034                .content,
1035            "from local"
1036        );
1037
1038        cloud
1039            .memorize(ContextFragment::new("s", FragmentKind::LongTerm, "from cloud").with_key("k"))
1040            .await
1041            .unwrap();
1042        assert_eq!(
1043            composite
1044                .get_by_key("s", "k")
1045                .await
1046                .unwrap()
1047                .unwrap()
1048                .content,
1049            "from cloud"
1050        );
1051    }
1052
1053    #[tokio::test]
1054    async fn memory_store_vector_recall_prefers_similar() {
1055        let store = MemoryContextStore::new();
1056        store
1057            .memorize(ContextFragment::new(
1058                "s",
1059                FragmentKind::Message,
1060                "user prefers rust for systems programming",
1061            ))
1062            .await
1063            .unwrap();
1064        store
1065            .memorize(ContextFragment::new(
1066                "s",
1067                FragmentKind::Message,
1068                "banana smoothie recipe with ice",
1069            ))
1070            .await
1071            .unwrap();
1072        let out = store
1073            .recall(&RecallQuery::new("s", "rust programming language"))
1074            .await
1075            .unwrap();
1076        assert!(!out.is_empty());
1077        assert!(out[0].content.contains("rust"));
1078        assert!(out[0].embedding.is_some());
1079    }
1080}