1use anyhow::Result;
4use async_trait::async_trait;
5use serde::{Deserialize, Serialize};
6use std::path::Path;
7use surrealdb::engine::local::RocksDb;
8use surrealdb::Surreal;
9
10use super::traits::*;
11
12#[derive(Clone)]
13pub struct SurrealMemory {
14 db: Surreal<surrealdb::engine::local::Db>,
15}
16
17#[derive(Debug, Clone, Serialize, Deserialize)]
18struct MemoryRow {
19 namespace: String,
20 key: String,
21 value: String,
22 metadata: Option<serde_json::Value>,
23 created_at: String,
24}
25
26#[derive(Debug, Clone, Serialize, Deserialize)]
27struct ConversationRow {
28 chat_id: String,
29 sender_id: String,
30 role: String,
31 content: String,
32 seq: i64,
33 created_at: String,
34}
35
36#[derive(Debug, Clone, Serialize, Deserialize)]
37struct StickerRow {
38 sticker_id: String,
39 file_id: String,
40 description: String,
41 analyzed_at: String,
42}
43
44#[derive(Debug, Clone, Serialize, Deserialize)]
45struct EmbeddingRow {
46 namespace: String,
47 key: String,
48 vector: Vec<f32>,
49 text: String,
50 created_at: String,
51}
52
53#[derive(Debug, Clone, Serialize, Deserialize)]
54struct FileIndexRow {
55 path: String,
56 hash: String,
57 last_indexed: String,
58}
59
60#[derive(Debug, Clone, Serialize, Deserialize)]
61struct ChunkRow {
62 file_path: String,
63 start_line: u32,
64 end_line: u32,
65 content: String,
66 embedding: Option<Vec<f32>>,
67 created_at: String,
68}
69
70impl SurrealMemory {
71 pub fn db(&self) -> Surreal<surrealdb::engine::local::Db> {
72 self.db.clone()
73 }
74
75 pub async fn new<P: AsRef<Path>>(path: P) -> Result<Self> {
76 let db = Surreal::new::<RocksDb>(path.as_ref()).await?;
77 db.use_ns("claw").use_db("memory").await?;
78 db.query(SCHEMA_SQL).await?;
79 Ok(Self { db })
80 }
81
82 fn memory_id(namespace: &str, key: &str) -> String {
83 format!("{namespace}::{key}")
84 }
85}
86
87const SCHEMA_SQL: &str = r#"
88 DEFINE TABLE IF NOT EXISTS memories SCHEMALESS;
89 DEFINE FIELD IF NOT EXISTS namespace ON memories TYPE string;
90 DEFINE FIELD IF NOT EXISTS key ON memories TYPE string;
91 DEFINE FIELD IF NOT EXISTS value ON memories TYPE string;
92 DEFINE FIELD IF NOT EXISTS metadata ON memories TYPE option<object>;
93 DEFINE FIELD IF NOT EXISTS created_at ON memories TYPE string;
94 DEFINE INDEX IF NOT EXISTS memory_lookup_idx ON memories FIELDS namespace, key UNIQUE;
95 DEFINE INDEX IF NOT EXISTS memory_namespace_idx ON memories FIELDS namespace;
96
97 DEFINE TABLE IF NOT EXISTS conversations SCHEMALESS;
98 DEFINE FIELD IF NOT EXISTS chat_id ON conversations TYPE string;
99 DEFINE FIELD IF NOT EXISTS sender_id ON conversations TYPE string;
100 DEFINE FIELD IF NOT EXISTS role ON conversations TYPE string;
101 DEFINE FIELD IF NOT EXISTS content ON conversations TYPE string;
102 DEFINE FIELD IF NOT EXISTS seq ON conversations TYPE int;
103 DEFINE FIELD IF NOT EXISTS created_at ON conversations TYPE string;
104 DEFINE INDEX IF NOT EXISTS conversation_chat_idx ON conversations FIELDS chat_id, seq;
105
106 DEFINE TABLE IF NOT EXISTS sticker_cache SCHEMALESS;
107 DEFINE FIELD IF NOT EXISTS sticker_id ON sticker_cache TYPE string;
108 DEFINE FIELD IF NOT EXISTS file_id ON sticker_cache TYPE string;
109 DEFINE FIELD IF NOT EXISTS description ON sticker_cache TYPE string;
110 DEFINE FIELD IF NOT EXISTS analyzed_at ON sticker_cache TYPE string;
111 DEFINE INDEX IF NOT EXISTS sticker_id_idx ON sticker_cache FIELDS sticker_id UNIQUE;
112
113 DEFINE TABLE IF NOT EXISTS embeddings SCHEMALESS;
114 DEFINE FIELD IF NOT EXISTS namespace ON embeddings TYPE string;
115 DEFINE FIELD IF NOT EXISTS key ON embeddings TYPE string;
116 DEFINE FIELD IF NOT EXISTS vector ON embeddings TYPE array;
117 DEFINE FIELD IF NOT EXISTS text ON embeddings TYPE string;
118 DEFINE FIELD IF NOT EXISTS created_at ON embeddings TYPE string;
119 DEFINE INDEX IF NOT EXISTS embedding_lookup_idx ON embeddings FIELDS namespace, key UNIQUE;
120 DEFINE INDEX IF NOT EXISTS embedding_namespace_idx ON embeddings FIELDS namespace;
121 -- TODO: define an MTREE index on `vector` for ANN-accelerated KNN when
122 -- the embedding dimension is fixed at deploy time. MTREE requires a
123 -- fixed DIMENSION, so it cannot be used while the table stores vectors
124 -- from multiple providers (e.g. OpenAI 1536-dim, Gemini 768-dim).
125 -- Example: DEFINE INDEX embedding_vector_mtree ON embeddings FIELDS vector MTREE DIMENSION 1536 DISTANCE COSINE;
126
127 DEFINE TABLE IF NOT EXISTS files SCHEMALESS;
128 DEFINE FIELD IF NOT EXISTS path ON files TYPE string;
129 DEFINE FIELD IF NOT EXISTS hash ON files TYPE string;
130 DEFINE FIELD IF NOT EXISTS last_indexed ON files TYPE string;
131 DEFINE INDEX IF NOT EXISTS file_path_idx ON files FIELDS path UNIQUE;
132
133 DEFINE TABLE IF NOT EXISTS chunks SCHEMALESS;
134 DEFINE FIELD IF NOT EXISTS file_path ON chunks TYPE string;
135 DEFINE FIELD IF NOT EXISTS start_line ON chunks TYPE int;
136 DEFINE FIELD IF NOT EXISTS end_line ON chunks TYPE int;
137 DEFINE FIELD IF NOT EXISTS content ON chunks TYPE string;
138 DEFINE FIELD IF NOT EXISTS embedding ON chunks TYPE option<array>;
139 DEFINE FIELD IF NOT EXISTS created_at ON chunks TYPE string;
140 DEFINE INDEX IF NOT EXISTS chunk_file_idx ON chunks FIELDS file_path;
141
142 DEFINE TABLE IF NOT EXISTS cron_jobs SCHEMALESS;
143 DEFINE FIELD IF NOT EXISTS name ON cron_jobs TYPE string;
144 DEFINE FIELD IF NOT EXISTS schedule ON cron_jobs TYPE string;
145 DEFINE FIELD IF NOT EXISTS task ON cron_jobs TYPE string;
146 DEFINE FIELD IF NOT EXISTS channel ON cron_jobs TYPE string;
147 DEFINE FIELD IF NOT EXISTS model ON cron_jobs TYPE string;
148 DEFINE FIELD IF NOT EXISTS enabled ON cron_jobs TYPE bool;
149 DEFINE FIELD IF NOT EXISTS last_run ON cron_jobs TYPE option<string>;
150 DEFINE FIELD IF NOT EXISTS next_run ON cron_jobs TYPE option<string>;
151 DEFINE INDEX IF NOT EXISTS cron_name_idx ON cron_jobs FIELDS name UNIQUE;
152
153 DEFINE TABLE IF NOT EXISTS memory_nodes SCHEMALESS;
154 DEFINE FIELD IF NOT EXISTS id ON memory_nodes TYPE string;
155 DEFINE FIELD IF NOT EXISTS kind ON memory_nodes TYPE string;
156 DEFINE FIELD IF NOT EXISTS text ON memory_nodes TYPE string;
157 DEFINE FIELD IF NOT EXISTS confidence ON memory_nodes TYPE float;
158 DEFINE FIELD IF NOT EXISTS status ON memory_nodes TYPE string;
159 DEFINE FIELD IF NOT EXISTS created_at ON memory_nodes TYPE string;
160 DEFINE INDEX IF NOT EXISTS memory_node_id_idx ON memory_nodes FIELDS id UNIQUE;
161
162 DEFINE TABLE IF NOT EXISTS memory_edges SCHEMALESS;
163 DEFINE FIELD IF NOT EXISTS from_id ON memory_edges TYPE string;
164 DEFINE FIELD IF NOT EXISTS to_id ON memory_edges TYPE string;
165 DEFINE FIELD IF NOT EXISTS rel ON memory_edges TYPE string;
166 DEFINE FIELD IF NOT EXISTS created_at ON memory_edges TYPE string;
167
168 DEFINE ANALYZER IF NOT EXISTS memory_analyzer TOKENIZERS blank, class FILTERS lowercase, snowball(english);
169 DEFINE INDEX IF NOT EXISTS memory_fts_idx ON memories FIELDS value
170 SEARCH ANALYZER memory_analyzer BM25;
171"#;
172
173fn parse_timestamp(value: &str) -> chrono::DateTime<chrono::Utc> {
174 chrono::DateTime::parse_from_rfc3339(value)
175 .map(|dt| dt.with_timezone(&chrono::Utc))
176 .unwrap_or_else(|_| chrono::Utc::now())
177}
178
179#[async_trait]
180impl MemoryBackend for SurrealMemory {
181 fn as_any(&self) -> &dyn std::any::Any {
182 self
183 }
184
185 async fn store(
186 &self,
187 namespace: &str,
188 key: &str,
189 value: &str,
190 metadata: Option<serde_json::Value>,
191 ) -> Result<()> {
192 let created_at = chrono::Utc::now().to_rfc3339();
193 let row = MemoryRow {
194 namespace: namespace.to_string(),
195 key: key.to_string(),
196 value: value.to_string(),
197 metadata,
198 created_at,
199 };
200 let _: Option<MemoryRow> = self
201 .db
202 .upsert(("memories", Self::memory_id(namespace, key)))
203 .content(row)
204 .await?;
205 Ok(())
206 }
207
208 async fn recall(&self, namespace: &str, key: &str) -> Result<Option<MemoryEntry>> {
209 let row: Option<MemoryRow> = self
210 .db
211 .select(("memories", Self::memory_id(namespace, key)))
212 .await?;
213 Ok(row.map(|entry| MemoryEntry {
214 key: entry.key,
215 value: entry.value,
216 metadata: entry.metadata,
217 created_at: parse_timestamp(&entry.created_at),
218 }))
219 }
220
221 async fn search(&self, namespace: &str, query: &str, limit: usize) -> Result<Vec<MemoryEntry>> {
222 let mut result = self
224 .db
225 .query(
226 "SELECT *, search::score(1) AS score
227 FROM memories
228 WHERE namespace = $namespace
229 AND value @1@ $query
230 ORDER BY score DESC
231 LIMIT $limit",
232 )
233 .bind(("namespace", namespace.to_string()))
234 .bind(("query", query.to_string()))
235 .bind(("limit", limit as i64))
236 .await?;
237 let rows: Vec<MemoryRow> = result.take(0)?;
238
239 if !rows.is_empty() {
240 return Ok(rows
241 .into_iter()
242 .map(|entry| MemoryEntry {
243 key: entry.key,
244 value: entry.value,
245 metadata: entry.metadata,
246 created_at: parse_timestamp(&entry.created_at),
247 })
248 .collect());
249 }
250
251 let query_lower = query.to_lowercase();
253 let mut result = self.db
254 .query(
255 "SELECT * FROM memories
256 WHERE namespace = $namespace
257 AND (string::lowercase(key) CONTAINS $query OR string::lowercase(value) CONTAINS $query)
258 ORDER BY created_at DESC
259 LIMIT $limit"
260 )
261 .bind(("namespace", namespace.to_string()))
262 .bind(("query", query_lower))
263 .bind(("limit", limit as i64))
264 .await?;
265 let rows: Vec<MemoryRow> = result.take(0)?;
266 Ok(rows
267 .into_iter()
268 .map(|entry| MemoryEntry {
269 key: entry.key,
270 value: entry.value,
271 metadata: entry.metadata,
272 created_at: parse_timestamp(&entry.created_at),
273 })
274 .collect())
275 }
276
277 async fn forget(&self, namespace: &str, key: &str) -> Result<()> {
278 let _: Option<MemoryRow> = self
279 .db
280 .delete(("memories", Self::memory_id(namespace, key)))
281 .await?;
282 Ok(())
283 }
284
285 async fn list(&self, namespace: &str) -> Result<Vec<MemoryEntry>> {
286 let mut result = self.db
287 .query("SELECT key, value, metadata, created_at FROM memories WHERE namespace = $namespace ORDER BY created_at DESC")
288 .bind(("namespace", namespace.to_string()))
289 .await?;
290 let rows: Vec<MemoryRow> = result.take(0)?;
291 Ok(rows
292 .into_iter()
293 .map(|entry| MemoryEntry {
294 key: entry.key,
295 value: entry.value,
296 metadata: entry.metadata,
297 created_at: parse_timestamp(&entry.created_at),
298 })
299 .collect())
300 }
301
302 async fn store_conversation(
303 &self,
304 chat_id: &str,
305 sender_id: &str,
306 role: &str,
307 content: &str,
308 ) -> Result<()> {
309 let now = chrono::Utc::now();
310 let row = ConversationRow {
311 chat_id: chat_id.to_string(),
312 sender_id: sender_id.to_string(),
313 role: role.to_string(),
314 content: content.to_string(),
315 seq: now.timestamp_millis(),
316 created_at: now.to_rfc3339(),
317 };
318 let _: Option<ConversationRow> = self.db.create("conversations").content(row).await?;
319 Ok(())
320 }
321
322 async fn store_conversation_batch(&self, entries: &[(&str, &str, &str, &str)]) -> Result<()> {
323 for (offset, (chat_id, sender_id, role, content)) in entries.iter().enumerate() {
324 let now = chrono::Utc::now();
325 let row = ConversationRow {
326 chat_id: (*chat_id).to_string(),
327 sender_id: (*sender_id).to_string(),
328 role: (*role).to_string(),
329 content: (*content).to_string(),
330 seq: now.timestamp_millis() + offset as i64,
331 created_at: now.to_rfc3339(),
332 };
333 let _: Option<ConversationRow> = self.db.create("conversations").content(row).await?;
334 }
335 Ok(())
336 }
337
338 async fn get_conversation_history(
339 &self,
340 chat_id: &str,
341 limit: usize,
342 ) -> Result<Vec<(String, String)>> {
343 let mut result = self
344 .db
345 .query(
346 "SELECT * FROM conversations
347 WHERE chat_id = $chat_id
348 ORDER BY seq DESC
349 LIMIT $limit",
350 )
351 .bind(("chat_id", chat_id.to_string()))
352 .bind(("limit", limit as i64))
353 .await?;
354 let mut rows: Vec<ConversationRow> = result.take(0)?;
355 rows.reverse();
356 Ok(rows
357 .into_iter()
358 .map(|row| (row.role, row.content))
359 .collect())
360 }
361
362 async fn search_conversations(
363 &self,
364 query: &str,
365 limit: usize,
366 chat_id: Option<&str>,
367 ) -> Result<Vec<ConversationSearchHit>> {
368 let query = query.to_lowercase();
369 let mut result = if let Some(chat_id) = chat_id {
370 self.db
371 .query(
372 "SELECT * FROM conversations
373 WHERE chat_id = $chat_id
374 AND string::lowercase(content) CONTAINS $query
375 ORDER BY seq DESC
376 LIMIT $limit",
377 )
378 .bind(("chat_id", chat_id.to_string()))
379 .bind(("query", query))
380 .bind(("limit", limit as i64))
381 .await?
382 } else {
383 self.db
384 .query(
385 "SELECT * FROM conversations
386 WHERE string::lowercase(content) CONTAINS $query
387 ORDER BY seq DESC
388 LIMIT $limit",
389 )
390 .bind(("query", query))
391 .bind(("limit", limit as i64))
392 .await?
393 };
394 let rows: Vec<ConversationRow> = result.take(0)?;
395 Ok(rows
396 .into_iter()
397 .map(|row| ConversationSearchHit {
398 chat_id: row.chat_id,
399 role: row.role,
400 content: row.content,
401 created_at: parse_timestamp(&row.created_at),
402 })
403 .collect())
404 }
405
406 async fn get_sticker_cache(&self, sticker_id: &str) -> Result<Option<String>> {
407 let row: Option<StickerRow> = self.db.select(("sticker_cache", sticker_id)).await?;
408 Ok(row.map(|entry| entry.description))
409 }
410
411 async fn store_sticker_cache(
412 &self,
413 sticker_id: &str,
414 file_id: &str,
415 description: &str,
416 ) -> Result<()> {
417 let row = StickerRow {
418 sticker_id: sticker_id.to_string(),
419 file_id: file_id.to_string(),
420 description: description.to_string(),
421 analyzed_at: chrono::Utc::now().to_rfc3339(),
422 };
423 let _: Option<StickerRow> = self
424 .db
425 .upsert(("sticker_cache", sticker_id))
426 .content(row)
427 .await?;
428 Ok(())
429 }
430
431 async fn store_embedding(
434 &self,
435 namespace: &str,
436 key: &str,
437 vector: &[f32],
438 text: &str,
439 ) -> Result<()> {
440 let row = EmbeddingRow {
441 namespace: namespace.to_string(),
442 key: key.to_string(),
443 vector: vector.to_vec(),
444 text: text.to_string(),
445 created_at: chrono::Utc::now().to_rfc3339(),
446 };
447 let id = Self::memory_id(namespace, key);
448 let _: Option<EmbeddingRow> = self.db.upsert(("embeddings", &id)).content(row).await?;
449 Ok(())
450 }
451
452 async fn search_embeddings(
453 &self,
454 namespace: &str,
455 query_vector: &[f32],
456 limit: usize,
457 ) -> Result<Vec<EmbeddingEntry>> {
458 let knn_result = self
467 .db
468 .query(
469 "SELECT * FROM embeddings
470 WHERE namespace = $namespace
471 AND vector <| $query_vector |>
472 LIMIT $limit",
473 )
474 .bind(("namespace", namespace.to_string()))
475 .bind(("query_vector", query_vector.to_vec()))
476 .bind(("limit", limit as i64))
477 .await;
478
479 if let Ok(mut result) = knn_result {
480 let rows: std::result::Result<Vec<EmbeddingRow>, _> = result.take(0);
481 if let Ok(rows) = rows {
482 if !rows.is_empty() {
483 return Ok(rows
484 .into_iter()
485 .map(|row| EmbeddingEntry {
486 namespace: row.namespace,
487 key: row.key,
488 vector: row.vector,
489 text: row.text,
490 created_at: parse_timestamp(&row.created_at),
491 })
492 .collect());
493 }
494 }
495 }
496
497 let mut result = self
501 .db
502 .query("SELECT * FROM embeddings WHERE namespace = $namespace")
503 .bind(("namespace", namespace.to_string()))
504 .await?;
505 let rows: Vec<EmbeddingRow> = result.take(0)?;
506
507 let mut scored: Vec<(f32, EmbeddingRow)> = rows
508 .into_iter()
509 .map(|row| {
510 let sim = cosine_similarity(query_vector, &row.vector);
511 (sim, row)
512 })
513 .collect();
514 scored.sort_by(|a, b| b.0.partial_cmp(&a.0).unwrap_or(std::cmp::Ordering::Equal));
515 scored.truncate(limit);
516
517 Ok(scored
518 .into_iter()
519 .map(|(_, row)| EmbeddingEntry {
520 namespace: row.namespace,
521 key: row.key,
522 vector: row.vector,
523 text: row.text,
524 created_at: parse_timestamp(&row.created_at),
525 })
526 .collect())
527 }
528
529 async fn store_file_index(&self, path: &str, hash: &str) -> Result<()> {
532 let row = FileIndexRow {
533 path: path.to_string(),
534 hash: hash.to_string(),
535 last_indexed: chrono::Utc::now().to_rfc3339(),
536 };
537 let id = format!("{:x}", md5_hash(path));
539 let _: Option<FileIndexRow> = self.db.upsert(("files", &id)).content(row).await?;
540 Ok(())
541 }
542
543 async fn get_file_index(&self, path: &str) -> Result<Option<FileIndex>> {
544 let id = format!("{:x}", md5_hash(path));
545 let row: Option<FileIndexRow> = self.db.select(("files", &id)).await?;
546 Ok(row.map(|r| FileIndex {
547 path: r.path,
548 hash: r.hash,
549 last_indexed: parse_timestamp(&r.last_indexed),
550 }))
551 }
552
553 async fn store_chunk(
556 &self,
557 file_path: &str,
558 start_line: u32,
559 end_line: u32,
560 content: &str,
561 embedding: Option<&[f32]>,
562 ) -> Result<()> {
563 let row = ChunkRow {
564 file_path: file_path.to_string(),
565 start_line,
566 end_line,
567 content: content.to_string(),
568 embedding: embedding.map(|e| e.to_vec()),
569 created_at: chrono::Utc::now().to_rfc3339(),
570 };
571 let _: Option<ChunkRow> = self.db.create("chunks").content(row).await?;
572 Ok(())
573 }
574
575 async fn get_chunks_for_file(&self, file_path: &str) -> Result<Vec<Chunk>> {
576 let mut result = self
577 .db
578 .query("SELECT * FROM chunks WHERE file_path = $file_path ORDER BY start_line ASC")
579 .bind(("file_path", file_path.to_string()))
580 .await?;
581 let rows: Vec<ChunkRow> = result.take(0)?;
582 Ok(rows
583 .into_iter()
584 .map(|r| Chunk {
585 file_path: r.file_path,
586 start_line: r.start_line,
587 end_line: r.end_line,
588 content: r.content,
589 embedding: r.embedding,
590 created_at: parse_timestamp(&r.created_at),
591 })
592 .collect())
593 }
594
595 async fn delete_chunks_for_file(&self, file_path: &str) -> Result<()> {
596 self.db
597 .query("DELETE FROM chunks WHERE file_path = $file_path")
598 .bind(("file_path", file_path.to_string()))
599 .await?;
600 Ok(())
601 }
602}
603
604fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
605 if a.len() != b.len() || a.is_empty() {
606 return 0.0;
607 }
608
609 let dot: f32 = a.iter().zip(b).map(|(x, y)| x * y).sum();
610 let ma: f32 = a.iter().map(|x| x * x).sum::<f32>().sqrt();
611 let mb: f32 = b.iter().map(|x| x * x).sum::<f32>().sqrt();
612
613 if ma == 0.0 || mb == 0.0 {
614 return 0.0;
615 }
616
617 dot / (ma * mb)
618}
619
620fn md5_hash(input: &str) -> u64 {
622 use std::hash::{Hash, Hasher};
623 let mut hasher = std::collections::hash_map::DefaultHasher::new();
624 input.hash(&mut hasher);
625 hasher.finish()
626}
627
628#[cfg(test)]
629mod tests {
630 use super::*;
631
632 #[test]
633 fn test_md5_hash_deterministic() {
634 let h1 = md5_hash("test-path");
635 let h2 = md5_hash("test-path");
636 assert_eq!(h1, h2);
637 }
638
639 #[test]
640 fn test_md5_hash_different_inputs() {
641 let h1 = md5_hash("path-a");
642 let h2 = md5_hash("path-b");
643 assert_ne!(h1, h2);
644 }
645
646 #[tokio::test]
647 async fn test_surreal_store_and_recall() {
648 let dir = tempfile::tempdir().unwrap();
649 let mem = SurrealMemory::new(dir.path()).await.unwrap();
650
651 mem.store("test", "greeting", "hello world", None)
652 .await
653 .unwrap();
654 let val = mem.recall("test", "greeting").await.unwrap();
655 assert!(val.is_some());
656 assert_eq!(val.unwrap().value, "hello world");
657 }
658
659 #[tokio::test]
660 async fn test_surreal_recall_missing() {
661 let dir = tempfile::tempdir().unwrap();
662 let mem = SurrealMemory::new(dir.path()).await.unwrap();
663
664 let val = mem.recall("test", "missing-key").await.unwrap();
665 assert!(val.is_none());
666 }
667
668 #[tokio::test]
669 async fn test_surreal_search() {
670 let dir = tempfile::tempdir().unwrap();
671 let mem = SurrealMemory::new(dir.path()).await.unwrap();
672
673 mem.store("ns", "k1", "the quick brown fox", None)
674 .await
675 .unwrap();
676 mem.store("ns", "k2", "lazy dog sleeps", None)
677 .await
678 .unwrap();
679 mem.store("ns", "k3", "fox runs fast", None).await.unwrap();
680
681 let results: Vec<MemoryEntry> = mem.search("ns", "fox", 10).await.unwrap();
682 assert!(!results.is_empty());
683 assert!(results.iter().any(|e| e.value.contains("fox")));
684 }
685
686 #[tokio::test]
687 async fn test_surreal_delete() {
688 let dir = tempfile::tempdir().unwrap();
689 let mem = SurrealMemory::new(dir.path()).await.unwrap();
690
691 mem.store("ns", "del-key", "to delete", None).await.unwrap();
692 assert!(mem.recall("ns", "del-key").await.unwrap().is_some());
693
694 mem.forget("ns", "del-key").await.unwrap();
695 assert!(mem.recall("ns", "del-key").await.unwrap().is_none());
696 }
697
698 #[tokio::test]
699 async fn test_surreal_conversation_history() {
700 let dir = tempfile::tempdir().unwrap();
701 let mem = SurrealMemory::new(dir.path()).await.unwrap();
702
703 mem.store_conversation("chat-1", "user-1", "user", "Hello")
704 .await
705 .unwrap();
706 tokio::time::sleep(std::time::Duration::from_millis(5)).await;
708 mem.store_conversation("chat-1", "assistant", "assistant", "Hi there")
709 .await
710 .unwrap();
711
712 let history: Vec<(String, String)> =
713 mem.get_conversation_history("chat-1", 10).await.unwrap();
714 assert_eq!(history.len(), 2);
715 assert_eq!(history[0].1, "Hello");
716 assert_eq!(history[1].1, "Hi there");
717 }
718
719 #[tokio::test]
720 async fn test_surreal_embeddings() {
721 let dir = tempfile::tempdir().unwrap();
722 let mem = SurrealMemory::new(dir.path()).await.unwrap();
723
724 let vec1 = vec![1.0, 0.0, 0.0];
725 let vec2 = vec![0.0, 1.0, 0.0];
726 let vec3 = vec![0.9, 0.1, 0.0];
727
728 mem.store_embedding("ns", "e1", &vec1, "first")
729 .await
730 .unwrap();
731 mem.store_embedding("ns", "e2", &vec2, "second")
732 .await
733 .unwrap();
734 mem.store_embedding("ns", "e3", &vec3, "third")
735 .await
736 .unwrap();
737
738 let results = mem.search_embeddings("ns", &vec1, 2).await.unwrap();
739 assert!(!results.is_empty());
740 assert!(results[0].key == "e1" || results[0].key == "e3");
742 }
743
744 #[tokio::test]
745 async fn test_surreal_file_index() {
746 let dir = tempfile::tempdir().unwrap();
747 let mem = SurrealMemory::new(dir.path()).await.unwrap();
748
749 mem.store_file_index("/src/main.rs", "abc123")
750 .await
751 .unwrap();
752 let idx = mem.get_file_index("/src/main.rs").await.unwrap();
753 assert!(idx.is_some());
754 assert_eq!(idx.unwrap().hash, "abc123");
755
756 let missing = mem.get_file_index("/src/nonexistent.rs").await.unwrap();
757 assert!(missing.is_none());
758 }
759
760 #[tokio::test]
761 async fn test_surreal_chunks() {
762 let dir = tempfile::tempdir().unwrap();
763 let mem = SurrealMemory::new(dir.path()).await.unwrap();
764
765 mem.store_chunk("/src/lib.rs", 1, 10, "fn main() {}", None)
766 .await
767 .unwrap();
768 mem.store_chunk("/src/lib.rs", 11, 20, "fn helper() {}", None)
769 .await
770 .unwrap();
771
772 let chunks = mem.get_chunks_for_file("/src/lib.rs").await.unwrap();
773 assert_eq!(chunks.len(), 2);
774 assert_eq!(chunks[0].start_line, 1);
775 assert_eq!(chunks[1].start_line, 11);
776
777 mem.delete_chunks_for_file("/src/lib.rs").await.unwrap();
778 let empty = mem.get_chunks_for_file("/src/lib.rs").await.unwrap();
779 assert!(empty.is_empty());
780 }
781
782 #[tokio::test]
783 async fn test_surreal_sticker_cache() {
784 let dir = tempfile::tempdir().unwrap();
785 let mem = SurrealMemory::new(dir.path()).await.unwrap();
786
787 mem.store_sticker_cache("stk-1", "file-1", "A happy cat")
788 .await
789 .unwrap();
790 let desc: Option<String> = mem.get_sticker_cache("stk-1").await.unwrap();
791 assert_eq!(desc, Some("A happy cat".to_string()));
792
793 let missing: Option<String> = mem.get_sticker_cache("stk-999").await.unwrap();
794 assert!(missing.is_none());
795 }
796}