1use async_trait::async_trait;
10use serde::{Deserialize, Serialize};
11use std::path::Path;
12use std::sync::Arc;
13use thiserror::Error;
14use uuid::Uuid;
15
16#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq)]
18pub enum FragmentKind {
19 Message,
21 ToolResult,
23 LongTerm,
25 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#[derive(Debug, Clone, Serialize, Deserialize)]
56pub struct ContextFragment {
57 pub id: String,
58 pub session: String,
59 pub key: Option<String>,
61 pub kind: FragmentKind,
62 pub content: String,
63 pub created_at: i64,
65 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#[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#[async_trait]
129pub trait MemoStore: Send + Sync {
130 async fn memorize(&self, frag: ContextFragment) -> Result<(), MemoError>;
132 async fn recall(&self, query: &RecallQuery) -> Result<Vec<ContextFragment>, MemoError>;
134 async fn compact(&self, session: &str) -> Result<ContextFragment, MemoError>;
136 async fn get_by_key(
138 &self,
139 session: &str,
140 key: &str,
141 ) -> Result<Option<ContextFragment>, MemoError>;
142}
143
144pub struct SledMemoStore {
147 fragments: sled::Tree,
148 embedder: Arc<dyn embed::Embedder>,
149}
150
151impl SledMemoStore {
152 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 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 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 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 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 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
291pub mod embed {
297 use super::MemoError;
298 use std::collections::HashMap;
299
300 pub trait Embedder: Send + Sync {
302 fn embed(&self, text: &str) -> Result<Vec<f32>, MemoError>;
303 fn dim(&self) -> usize;
304 }
305
306 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 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 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 assert!(frags[0].content.contains("rust"));
480 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}