Skip to main content

adk_memory/
inmemory.rs

1use crate::service::*;
2use adk_core::Result;
3use async_trait::async_trait;
4use std::collections::{HashMap, HashSet};
5use std::sync::{Arc, RwLock};
6
7#[derive(Clone, Debug, PartialEq, Eq, Hash)]
8struct MemoryKey {
9    app_name: String,
10    user_id: String,
11}
12
13#[derive(Clone)]
14struct StoredEntry {
15    entry: MemoryEntry,
16    words: HashSet<String>,
17    project_id: Option<String>,
18}
19
20type MemoryStore = HashMap<MemoryKey, HashMap<String, Vec<StoredEntry>>>;
21
22pub struct InMemoryMemoryService {
23    store: Arc<RwLock<MemoryStore>>,
24}
25
26impl InMemoryMemoryService {
27    pub fn new() -> Self {
28        Self { store: Arc::new(RwLock::new(HashMap::new())) }
29    }
30
31    fn has_intersection(set1: &HashSet<String>, set2: &HashSet<String>) -> bool {
32        if set1.is_empty() || set2.is_empty() {
33            return false;
34        }
35        set1.iter().any(|word| set2.contains(word))
36    }
37}
38
39impl Default for InMemoryMemoryService {
40    fn default() -> Self {
41        Self::new()
42    }
43}
44
45#[async_trait]
46impl MemoryService for InMemoryMemoryService {
47    async fn add_session(
48        &self,
49        app_name: &str,
50        user_id: &str,
51        session_id: &str,
52        entries: Vec<MemoryEntry>,
53    ) -> Result<()> {
54        let key = MemoryKey { app_name: app_name.to_string(), user_id: user_id.to_string() };
55
56        let stored_entries: Vec<StoredEntry> = entries
57            .into_iter()
58            .map(|entry| {
59                let words = crate::text::extract_words_from_content(&entry.content);
60                StoredEntry { entry, words, project_id: None }
61            })
62            .filter(|e| !e.words.is_empty())
63            .collect();
64
65        if stored_entries.is_empty() {
66            return Ok(());
67        }
68
69        let mut store = self.store.write().unwrap();
70        let sessions = store.entry(key).or_default();
71        sessions.insert(session_id.to_string(), stored_entries);
72
73        Ok(())
74    }
75
76    fn supports_project_scoping(&self) -> bool {
77        true
78    }
79
80    async fn add_session_to_project(
81        &self,
82        app_name: &str,
83        user_id: &str,
84        session_id: &str,
85        project_id: &str,
86        entries: Vec<MemoryEntry>,
87    ) -> Result<()> {
88        validate_project_id(project_id)?;
89
90        let key = MemoryKey { app_name: app_name.to_string(), user_id: user_id.to_string() };
91
92        let stored_entries: Vec<StoredEntry> = entries
93            .into_iter()
94            .map(|entry| {
95                let words = crate::text::extract_words_from_content(&entry.content);
96                StoredEntry { entry, words, project_id: Some(project_id.to_string()) }
97            })
98            .filter(|e| !e.words.is_empty())
99            .collect();
100
101        if stored_entries.is_empty() {
102            return Ok(());
103        }
104
105        let mut store = self.store.write().unwrap();
106        let sessions = store.entry(key).or_default();
107        sessions.insert(session_id.to_string(), stored_entries);
108
109        Ok(())
110    }
111
112    async fn add_entry(&self, app_name: &str, user_id: &str, entry: MemoryEntry) -> Result<()> {
113        let key = MemoryKey { app_name: app_name.to_string(), user_id: user_id.to_string() };
114        let words = crate::text::extract_words_from_content(&entry.content);
115        let stored = StoredEntry { entry, words, project_id: None };
116
117        let mut store = self.store.write().unwrap();
118        let sessions = store.entry(key).or_default();
119        sessions.entry("__direct__".to_string()).or_default().push(stored);
120
121        Ok(())
122    }
123
124    async fn add_entry_to_project(
125        &self,
126        app_name: &str,
127        user_id: &str,
128        project_id: &str,
129        entry: MemoryEntry,
130    ) -> Result<()> {
131        validate_project_id(project_id)?;
132
133        let key = MemoryKey { app_name: app_name.to_string(), user_id: user_id.to_string() };
134        let words = crate::text::extract_words_from_content(&entry.content);
135        let stored = StoredEntry { entry, words, project_id: Some(project_id.to_string()) };
136
137        let mut store = self.store.write().unwrap();
138        let sessions = store.entry(key).or_default();
139        sessions.entry("__direct__".to_string()).or_default().push(stored);
140
141        Ok(())
142    }
143
144    async fn delete_entries(&self, app_name: &str, user_id: &str, query: &str) -> Result<u64> {
145        let query_words = crate::text::extract_words(query);
146        if query_words.is_empty() {
147            return Ok(0);
148        }
149
150        let key = MemoryKey { app_name: app_name.to_string(), user_id: user_id.to_string() };
151
152        let mut store = self.store.write().unwrap();
153        let sessions = match store.get_mut(&key) {
154            Some(s) => s,
155            None => return Ok(0),
156        };
157
158        let mut removed: u64 = 0;
159        for entries in sessions.values_mut() {
160            let before = entries.len();
161            entries.retain(|stored| {
162                // Only delete global entries (project_id is None)
163                stored.project_id.is_some() || !Self::has_intersection(&stored.words, &query_words)
164            });
165            removed += (before - entries.len()) as u64;
166        }
167
168        Ok(removed)
169    }
170
171    async fn list_recent(
172        &self,
173        app_name: &str,
174        user_id: &str,
175        limit: usize,
176    ) -> Result<Vec<MemoryEntry>> {
177        let key = MemoryKey { app_name: app_name.to_string(), user_id: user_id.to_string() };
178        let store = self.store.read().unwrap();
179        let mut entries: Vec<MemoryEntry> = store
180            .get(&key)
181            .map(|sessions| {
182                sessions.values().flatten().map(|stored| stored.entry.clone()).collect()
183            })
184            .unwrap_or_default();
185        // Newest first. `Reverse` keeps this a key-based sort, which clippy prefers, without
186        // flipping the comparison by hand.
187        entries.sort_by_key(|entry| std::cmp::Reverse(entry.timestamp));
188        entries.truncate(limit);
189        Ok(entries)
190    }
191
192    async fn delete_entries_in_project(
193        &self,
194        app_name: &str,
195        user_id: &str,
196        project_id: &str,
197        query: &str,
198    ) -> Result<u64> {
199        let query_words = crate::text::extract_words(query);
200        if query_words.is_empty() {
201            return Ok(0);
202        }
203
204        let key = MemoryKey { app_name: app_name.to_string(), user_id: user_id.to_string() };
205
206        let mut store = self.store.write().unwrap();
207        let sessions = match store.get_mut(&key) {
208            Some(s) => s,
209            None => return Ok(0),
210        };
211
212        let mut removed: u64 = 0;
213        for entries in sessions.values_mut() {
214            let before = entries.len();
215            entries.retain(|stored| {
216                // Only delete entries matching the given project
217                stored.project_id.as_deref() != Some(project_id)
218                    || !Self::has_intersection(&stored.words, &query_words)
219            });
220            removed += (before - entries.len()) as u64;
221        }
222
223        Ok(removed)
224    }
225
226    async fn delete_project(&self, app_name: &str, user_id: &str, project_id: &str) -> Result<u64> {
227        let key = MemoryKey { app_name: app_name.to_string(), user_id: user_id.to_string() };
228
229        let mut store = self.store.write().unwrap();
230        let sessions = match store.get_mut(&key) {
231            Some(s) => s,
232            None => return Ok(0),
233        };
234
235        let mut removed: u64 = 0;
236        for entries in sessions.values_mut() {
237            let before = entries.len();
238            entries.retain(|stored| stored.project_id.as_deref() != Some(project_id));
239            removed += (before - entries.len()) as u64;
240        }
241
242        Ok(removed)
243    }
244
245    async fn delete_user(&self, app_name: &str, user_id: &str) -> Result<()> {
246        let key = MemoryKey { app_name: app_name.to_string(), user_id: user_id.to_string() };
247
248        let mut store = self.store.write().unwrap();
249        store.remove(&key);
250
251        Ok(())
252    }
253
254    async fn search(&self, req: SearchRequest) -> Result<SearchResponse> {
255        let query_words = crate::text::extract_words(&req.query);
256        let limit = req.limit.unwrap_or(10);
257
258        let key = MemoryKey { app_name: req.app_name, user_id: req.user_id };
259
260        let store = self.store.read().unwrap();
261        let sessions = match store.get(&key) {
262            Some(s) => s,
263            None => return Ok(SearchResponse { memories: Vec::new() }),
264        };
265
266        let mut memories = Vec::new();
267        for stored_entries in sessions.values() {
268            for stored in stored_entries {
269                if !Self::has_intersection(&stored.words, &query_words) {
270                    continue;
271                }
272
273                match &req.project_id {
274                    // Global search: only include global entries
275                    None => {
276                        if stored.project_id.is_none() {
277                            memories.push(stored.entry.clone());
278                        }
279                    }
280                    // Project search: include global + matching project entries
281                    Some(pid) => {
282                        if stored.project_id.is_none()
283                            || stored.project_id.as_deref() == Some(pid.as_str())
284                        {
285                            memories.push(stored.entry.clone());
286                        }
287                    }
288                }
289            }
290        }
291
292        memories.truncate(limit);
293
294        Ok(SearchResponse { memories })
295    }
296}