Skip to main content

apollo/memory/
surreal.rs

1//! SurrealDB-backed memory storage using the local RocksDB engine.
2
3use 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        // Try full-text search with BM25 ranking first
223        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        // Fallback to CONTAINS for partial matches
252        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    // ── Embeddings ──
432
433    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        // Use SurrealDB's built-in KNN vector search (<| |> operator).
459        // This performs a server-side nearest-neighbour scan and avoids
460        // loading every embedding into the client. When an MTREE index is
461        // defined on the vector field the query planner uses it; otherwise
462        // it falls back to a brute-force scan inside the database engine.
463        //
464        // If the KNN query fails (e.g. dimension mismatch, empty table) we
465        // fall back to the in-process cosine similarity scan below.
466        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        // Fallback: load all embeddings for the namespace and do cosine
498        // search in-process. This handles mixed-dimension vectors or any
499        // case where the KNN operator is unavailable.
500        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    // ── File indexing ──
530
531    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        // Use path hash as record ID to avoid special chars
538        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    // ── Code chunks ──
554
555    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
620/// Simple hash for file path → record ID
621fn 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        // Small delay to ensure different seq timestamps
707        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        // The closest to [1,0,0] should be e1 or e3
741        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}