1use 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
26pub const EMBED_DIM: usize = 256;
32
33#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq)]
35pub enum FragmentKind {
36 Message,
38 ToolResult,
40 LongTerm,
42 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#[derive(Debug, Clone, Serialize, Deserialize)]
73pub struct ContextFragment {
74 pub id: String,
75 pub session: String,
76 pub key: Option<String>,
78 pub kind: FragmentKind,
79 pub content: String,
80 pub created_at: i64,
82 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#[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#[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#[async_trait]
194pub trait ContextStore: Send + Sync {
195 async fn memorize(&self, frag: ContextFragment) -> Result<(), ContextError>;
197 async fn recall(&self, query: &RecallQuery) -> Result<Vec<ContextFragment>, ContextError>;
199 async fn compact(&self, session: &str) -> Result<ContextFragment, ContextError>;
201 async fn get_by_key(
203 &self,
204 session: &str,
205 key: &str,
206 ) -> Result<Option<ContextFragment>, ContextError>;
207 async fn list_session(
210 &self,
211 session: &str,
212 top_k: usize,
213 ) -> Result<Vec<ContextFragment>, ContextError>;
214}
215
216pub 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
226pub 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
246pub 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
263pub 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
288pub 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
322pub 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 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
446pub 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
558pub mod embed {
563 use super::ContextError;
564 use std::collections::HashMap;
565
566 pub trait Embedder: Send + Sync {
568 fn embed(&self, text: &str) -> Result<Vec<f32>, ContextError>;
569 fn dim(&self) -> usize;
570 }
571
572 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 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 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 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 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 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 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}