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 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 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 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 None => {
276 if stored.project_id.is_none() {
277 memories.push(stored.entry.clone());
278 }
279 }
280 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}