1use crate::{
2 AppendEventRequest, CreateRequest, DeleteRequest, Event, Events, GetRequest, KEY_PREFIX_TEMP,
3 ListRequest, Session, SessionService, State, state_utils,
4};
5use adk_core::Result;
6use adk_core::identity::{AdkIdentity, AppName, SessionId, UserId};
7use async_trait::async_trait;
8use chrono::{DateTime, Utc};
9use serde_json::Value;
10use std::collections::HashMap;
11use std::sync::{Arc, RwLock};
12use uuid::Uuid;
13
14type StateMap = HashMap<String, Value>;
15
16#[derive(Clone)]
17struct SessionData {
18 identity: AdkIdentity,
19 events: Vec<Event>,
20 state: StateMap,
21 updated_at: DateTime<Utc>,
22}
23
24pub struct InMemorySessionService {
28 sessions: Arc<RwLock<HashMap<AdkIdentity, SessionData>>>,
29 app_state: Arc<RwLock<HashMap<String, StateMap>>>,
30 user_state: Arc<RwLock<HashMap<String, HashMap<String, StateMap>>>>,
31}
32
33impl InMemorySessionService {
34 pub fn new() -> Self {
36 Self {
37 sessions: Arc::new(RwLock::new(HashMap::new())),
38 app_state: Arc::new(RwLock::new(HashMap::new())),
39 user_state: Arc::new(RwLock::new(HashMap::new())),
40 }
41 }
42
43 fn extract_state_deltas(delta: &HashMap<String, Value>) -> (StateMap, StateMap, StateMap) {
44 state_utils::extract_state_deltas(delta)
45 }
46
47 fn merge_states(app: &StateMap, user: &StateMap, session: &StateMap) -> StateMap {
48 state_utils::merge_states(app, user, session)
49 }
50
51 fn make_identity(app_name: &str, user_id: &str, session_id: &str) -> Result<AdkIdentity> {
53 Ok(AdkIdentity::new(
54 AppName::try_from(app_name)?,
55 UserId::try_from(user_id)?,
56 SessionId::try_from(session_id)?,
57 ))
58 }
59
60 async fn rewind_to_empty(&self, session_id: &str) -> Result<Box<dyn Session>> {
62 let mut sessions = self.sessions.write().unwrap_or_else(|e| e.into_inner());
63 let data = sessions
64 .values_mut()
65 .find(|d| d.identity.session_id.as_ref() == session_id)
66 .ok_or_else(|| adk_core::AdkError::session("session not found"))?;
67
68 data.events.clear();
69 data.state = HashMap::new();
70 data.updated_at = Utc::now();
71
72 let app_name = data.identity.app_name.as_ref().to_string();
73 let user_id = data.identity.user_id.as_ref().to_string();
74 let identity = data.identity.clone();
75 let updated_at = data.updated_at;
76 drop(sessions);
77
78 let app_state_lock = self.app_state.read().unwrap_or_else(|e| e.into_inner());
79 let app_state = app_state_lock.get(&app_name).cloned().unwrap_or_default();
80 drop(app_state_lock);
81
82 let user_state_lock = self.user_state.read().unwrap_or_else(|e| e.into_inner());
83 let user_state = user_state_lock
84 .get(&app_name)
85 .and_then(|m| m.get(&user_id))
86 .cloned()
87 .unwrap_or_default();
88 drop(user_state_lock);
89
90 let merged_state = state_utils::merge_states(&app_state, &user_state, &HashMap::new());
91
92 Ok(Box::new(InMemorySession {
93 identity,
94 state: merged_state,
95 events: Vec::new(),
96 updated_at,
97 }))
98 }
99}
100
101impl Default for InMemorySessionService {
102 fn default() -> Self {
103 Self::new()
104 }
105}
106
107#[async_trait]
108impl SessionService for InMemorySessionService {
109 async fn create(&self, req: CreateRequest) -> Result<Box<dyn Session>> {
110 let session_id_str = req.session_id.unwrap_or_else(|| Uuid::new_v4().to_string());
111
112 let identity = Self::make_identity(&req.app_name, &req.user_id, &session_id_str)?;
113
114 let (app_delta, user_delta, session_state) = Self::extract_state_deltas(&req.state);
115
116 let mut app_state_lock = self.app_state.write().unwrap_or_else(|e| e.into_inner());
117 let app_state = app_state_lock.entry(req.app_name.clone()).or_default();
118 app_state.extend(app_delta);
119 let app_state_clone = app_state.clone();
120 drop(app_state_lock);
121
122 let mut user_state_lock = self.user_state.write().unwrap_or_else(|e| e.into_inner());
123 let user_map = user_state_lock.entry(req.app_name.clone()).or_default();
124 let user_state = user_map.entry(req.user_id.clone()).or_default();
125 user_state.extend(user_delta);
126 let user_state_clone = user_state.clone();
127 drop(user_state_lock);
128
129 let merged_state = Self::merge_states(&app_state_clone, &user_state_clone, &session_state);
130
131 let data = SessionData {
132 identity: identity.clone(),
133 events: Vec::new(),
134 state: merged_state.clone(),
135 updated_at: Utc::now(),
136 };
137
138 let mut sessions = self.sessions.write().unwrap_or_else(|e| e.into_inner());
139 sessions.insert(identity.clone(), data);
140 drop(sessions);
141
142 Ok(Box::new(InMemorySession {
143 identity,
144 state: merged_state,
145 events: Vec::new(),
146 updated_at: Utc::now(),
147 }))
148 }
149
150 async fn get(&self, req: GetRequest) -> Result<Box<dyn Session>> {
151 let identity = Self::make_identity(&req.app_name, &req.user_id, &req.session_id)?;
152
153 let sessions = self.sessions.read().unwrap_or_else(|e| e.into_inner());
154 let data =
155 sessions.get(&identity).ok_or_else(|| crate::service::session_not_found(&req))?;
156
157 let app_state_lock = self.app_state.read().unwrap_or_else(|e| e.into_inner());
158 let app_state = app_state_lock.get(&req.app_name).cloned().unwrap_or_default();
159 drop(app_state_lock);
160
161 let user_state_lock = self.user_state.read().unwrap_or_else(|e| e.into_inner());
162 let user_state = user_state_lock
163 .get(&req.app_name)
164 .and_then(|m| m.get(&req.user_id))
165 .cloned()
166 .unwrap_or_default();
167 drop(user_state_lock);
168
169 let merged_state = Self::merge_states(&app_state, &user_state, &data.state);
170
171 let mut events = data.events.clone();
172 if let Some(num) = req.num_recent_events {
173 let start = events.len().saturating_sub(num);
174 events = events[start..].to_vec();
175 }
176 if let Some(after) = req.after {
177 events.retain(|e| e.timestamp >= after);
178 }
179
180 Ok(Box::new(InMemorySession {
181 identity: data.identity.clone(),
182 state: merged_state,
183 events,
184 updated_at: data.updated_at,
185 }))
186 }
187
188 async fn list(&self, req: ListRequest) -> Result<Vec<Box<dyn Session>>> {
189 let sessions = self.sessions.read().unwrap_or_else(|e| e.into_inner());
190 let offset = req.offset.unwrap_or(0);
191 let limit = req.limit.unwrap_or(usize::MAX);
192 let mut result = Vec::new();
193
194 for data in sessions.values() {
195 if data.identity.app_name.as_ref() == req.app_name
196 && data.identity.user_id.as_ref() == req.user_id
197 {
198 result.push(data.clone());
199 }
200 }
201
202 result.sort_by_key(|b| std::cmp::Reverse(b.updated_at));
204
205 let result: Vec<Box<dyn Session>> = result
206 .into_iter()
207 .skip(offset)
208 .take(limit)
209 .map(|data| {
210 Box::new(InMemorySession {
211 identity: data.identity,
212 state: data.state,
213 events: data.events,
214 updated_at: data.updated_at,
215 }) as Box<dyn Session>
216 })
217 .collect();
218
219 Ok(result)
220 }
221
222 async fn delete(&self, req: DeleteRequest) -> Result<()> {
223 let identity = Self::make_identity(&req.app_name, &req.user_id, &req.session_id)?;
224
225 let mut sessions = self.sessions.write().unwrap_or_else(|e| e.into_inner());
226 sessions.remove(&identity);
227 Ok(())
228 }
229
230 async fn delete_all_sessions(&self, app_name: &str, user_id: &str) -> Result<()> {
231 let mut sessions = self.sessions.write().unwrap_or_else(|e| e.into_inner());
232 sessions.retain(|_, data| {
233 !(data.identity.app_name.as_ref() == app_name
234 && data.identity.user_id.as_ref() == user_id)
235 });
236 Ok(())
237 }
238
239 async fn append_event(&self, session_id: &str, mut event: Event) -> Result<()> {
240 event.actions.state_delta.retain(|k, _| !k.starts_with(KEY_PREFIX_TEMP));
241
242 let (app_name, user_id, app_delta, user_delta, _session_delta) = {
243 let mut sessions = self.sessions.write().unwrap_or_else(|e| e.into_inner());
244 let data = sessions
245 .values_mut()
246 .find(|d| d.identity.session_id.as_ref() == session_id)
247 .ok_or_else(|| adk_core::AdkError::session("session not found"))?;
248
249 data.events.push(event.clone());
250 data.updated_at = event.timestamp;
251
252 let (app_delta, user_delta, session_delta) =
253 Self::extract_state_deltas(&event.actions.state_delta);
254 data.state.extend(session_delta.clone());
255
256 (
257 data.identity.app_name.as_ref().to_string(),
258 data.identity.user_id.as_ref().to_string(),
259 app_delta,
260 user_delta,
261 session_delta,
262 )
263 };
264
265 if !app_delta.is_empty() {
266 let mut app_state_lock = self.app_state.write().unwrap_or_else(|e| e.into_inner());
267 let app_state = app_state_lock.entry(app_name.clone()).or_default();
268 app_state.extend(app_delta);
269 }
270
271 if !user_delta.is_empty() {
272 let mut user_state_lock = self.user_state.write().unwrap_or_else(|e| e.into_inner());
273 let user_map = user_state_lock.entry(app_name).or_default();
274 let user_state = user_map.entry(user_id).or_default();
275 user_state.extend(user_delta);
276 }
277
278 Ok(())
279 }
280
281 async fn append_event_for_identity(&self, req: AppendEventRequest) -> Result<()> {
282 let mut event = req.event;
283 event.actions.state_delta.retain(|k, _| !k.starts_with(KEY_PREFIX_TEMP));
284
285 let identity = req.identity;
286
287 let (app_name_str, user_id_str, app_delta, user_delta) = {
288 let mut sessions = self.sessions.write().unwrap_or_else(|e| e.into_inner());
289 let data = sessions
290 .get_mut(&identity)
291 .ok_or_else(|| adk_core::AdkError::session("session not found"))?;
292
293 data.events.push(event.clone());
294 data.updated_at = event.timestamp;
295
296 let (app_delta, user_delta, session_delta) =
297 Self::extract_state_deltas(&event.actions.state_delta);
298 data.state.extend(session_delta);
299
300 (
301 identity.app_name.as_ref().to_string(),
302 identity.user_id.as_ref().to_string(),
303 app_delta,
304 user_delta,
305 )
306 };
307
308 if !app_delta.is_empty() {
309 let mut app_state_lock = self.app_state.write().unwrap_or_else(|e| e.into_inner());
310 let app_state = app_state_lock.entry(app_name_str.clone()).or_default();
311 app_state.extend(app_delta);
312 }
313
314 if !user_delta.is_empty() {
315 let mut user_state_lock = self.user_state.write().unwrap_or_else(|e| e.into_inner());
316 let user_map = user_state_lock.entry(app_name_str).or_default();
317 let user_state = user_map.entry(user_id_str).or_default();
318 user_state.extend(user_delta);
319 }
320
321 Ok(())
322 }
323
324 async fn get_for_identity(&self, identity: &AdkIdentity) -> Result<Box<dyn Session>> {
325 let sessions = self.sessions.read().unwrap_or_else(|e| e.into_inner());
326 let data = sessions
327 .get(identity)
328 .ok_or_else(|| adk_core::AdkError::session("session not found"))?;
329
330 let app_state_lock = self.app_state.read().unwrap_or_else(|e| e.into_inner());
331 let app_state = app_state_lock.get(identity.app_name.as_ref()).cloned().unwrap_or_default();
332 drop(app_state_lock);
333
334 let user_state_lock = self.user_state.read().unwrap_or_else(|e| e.into_inner());
335 let user_state = user_state_lock
336 .get(identity.app_name.as_ref())
337 .and_then(|m| m.get(identity.user_id.as_ref()))
338 .cloned()
339 .unwrap_or_default();
340 drop(user_state_lock);
341
342 let merged_state = Self::merge_states(&app_state, &user_state, &data.state);
343
344 Ok(Box::new(InMemorySession {
345 identity: data.identity.clone(),
346 state: merged_state,
347 events: data.events.clone(),
348 updated_at: data.updated_at,
349 }))
350 }
351
352 async fn delete_for_identity(&self, identity: &AdkIdentity) -> Result<()> {
353 let mut sessions = self.sessions.write().unwrap_or_else(|e| e.into_inner());
354 sessions.remove(identity);
355 Ok(())
356 }
357
358 async fn rewind(&self, session_id: &str, target_event_id: &str) -> Result<Box<dyn Session>> {
359 let mut sessions = self.sessions.write().unwrap_or_else(|e| e.into_inner());
360
361 let data = sessions
363 .values_mut()
364 .find(|d| d.identity.session_id.as_ref() == session_id)
365 .ok_or_else(|| adk_core::AdkError::session("session not found"))?;
366
367 let target_index =
369 data.events.iter().position(|e| e.id == target_event_id).ok_or_else(|| {
370 adk_core::AdkError::session(format!("target event not found: {target_event_id}"))
371 })?;
372
373 data.events.truncate(target_index + 1);
375
376 let mut rebuilt_session_state: HashMap<String, Value> = HashMap::new();
378 for event in &data.events {
379 let (_app_delta, _user_delta, session_delta) =
380 state_utils::extract_state_deltas(&event.actions.state_delta);
381 rebuilt_session_state.extend(session_delta);
382 }
383
384 let app_name = data.identity.app_name.as_ref().to_string();
386 let user_id = data.identity.user_id.as_ref().to_string();
387
388 data.state = rebuilt_session_state.clone();
390 data.updated_at = data.events.last().map(|e| e.timestamp).unwrap_or(Utc::now());
391
392 let identity = data.identity.clone();
393 let events = data.events.clone();
394 let updated_at = data.updated_at;
395 drop(sessions);
396
397 let app_state_lock = self.app_state.read().unwrap_or_else(|e| e.into_inner());
399 let app_state = app_state_lock.get(&app_name).cloned().unwrap_or_default();
400 drop(app_state_lock);
401
402 let user_state_lock = self.user_state.read().unwrap_or_else(|e| e.into_inner());
403 let user_state = user_state_lock
404 .get(&app_name)
405 .and_then(|m| m.get(&user_id))
406 .cloned()
407 .unwrap_or_default();
408 drop(user_state_lock);
409
410 let merged_state =
411 state_utils::merge_states(&app_state, &user_state, &rebuilt_session_state);
412
413 Ok(Box::new(InMemorySession { identity, state: merged_state, events, updated_at }))
414 }
415
416 async fn rewind_steps(&self, session_id: &str, steps: usize) -> Result<Box<dyn Session>> {
417 if steps == 0 {
418 let sessions = self.sessions.read().unwrap_or_else(|e| e.into_inner());
420 let data = sessions
421 .values()
422 .find(|d| d.identity.session_id.as_ref() == session_id)
423 .ok_or_else(|| adk_core::AdkError::session("session not found"))?;
424
425 let app_name = data.identity.app_name.as_ref().to_string();
426 let user_id = data.identity.user_id.as_ref().to_string();
427 let identity = data.identity.clone();
428 let events = data.events.clone();
429 let session_state = data.state.clone();
430 let updated_at = data.updated_at;
431 drop(sessions);
432
433 let app_state_lock = self.app_state.read().unwrap_or_else(|e| e.into_inner());
434 let app_state = app_state_lock.get(&app_name).cloned().unwrap_or_default();
435 drop(app_state_lock);
436
437 let user_state_lock = self.user_state.read().unwrap_or_else(|e| e.into_inner());
438 let user_state = user_state_lock
439 .get(&app_name)
440 .and_then(|m| m.get(&user_id))
441 .cloned()
442 .unwrap_or_default();
443 drop(user_state_lock);
444
445 let merged_state = state_utils::merge_states(&app_state, &user_state, &session_state);
446
447 return Ok(Box::new(InMemorySession {
448 identity,
449 state: merged_state,
450 events,
451 updated_at,
452 }));
453 }
454
455 let rewind_target = {
457 let sessions = self.sessions.read().unwrap_or_else(|e| e.into_inner());
458 let data = sessions
459 .values()
460 .find(|d| d.identity.session_id.as_ref() == session_id)
461 .ok_or_else(|| adk_core::AdkError::session("session not found"))?;
462
463 if steps > data.events.len() {
464 return Err(adk_core::AdkError::session("rewind steps exceeds event count"));
465 }
466
467 let target_index = data.events.len() - steps;
468 if target_index == 0 {
469 None
471 } else {
472 Some(data.events[target_index - 1].id.clone())
473 }
474 };
475
476 match rewind_target {
477 Some(target_event_id) => self.rewind(session_id, &target_event_id).await,
478 None => self.rewind_to_empty(session_id).await,
479 }
480 }
481}
482
483struct InMemorySession {
484 identity: AdkIdentity,
485 state: StateMap,
486 events: Vec<Event>,
487 updated_at: DateTime<Utc>,
488}
489
490impl Session for InMemorySession {
491 fn id(&self) -> &str {
492 self.identity.session_id.as_ref()
493 }
494
495 fn app_name(&self) -> &str {
496 self.identity.app_name.as_ref()
497 }
498
499 fn user_id(&self) -> &str {
500 self.identity.user_id.as_ref()
501 }
502
503 fn state(&self) -> &dyn State {
504 self
505 }
506
507 fn events(&self) -> &dyn Events {
508 self
509 }
510
511 fn last_update_time(&self) -> DateTime<Utc> {
512 self.updated_at
513 }
514}
515
516impl State for InMemorySession {
517 fn get(&self, key: &str) -> Option<Value> {
518 self.state.get(key).cloned()
519 }
520
521 fn set(&mut self, key: String, value: Value) {
522 if let Err(msg) = adk_core::validate_state_key(&key) {
523 tracing::warn!(key = %key, "rejecting invalid state key: {msg}");
524 return;
525 }
526 self.state.insert(key, value);
527 }
528
529 fn all(&self) -> HashMap<String, Value> {
530 self.state.clone()
531 }
532}
533
534impl Events for InMemorySession {
535 fn all(&self) -> Vec<Event> {
536 self.events.clone()
537 }
538
539 fn len(&self) -> usize {
540 self.events.len()
541 }
542
543 fn at(&self, index: usize) -> Option<&Event> {
544 self.events.get(index)
545 }
546}