Skip to main content

adk_session/
inmemory.rs

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
24/// In-memory session service for testing and lightweight deployments.
25///
26/// All data is stored in process memory and lost on restart.
27pub 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    /// Creates a new empty in-memory session service.
35    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    /// Builds an [`AdkIdentity`] from raw string fields.
52    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    /// Rewind a session to before all events (remove all events and reset state).
61    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        // Sort by updated_at descending for consistency with other backends
203        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        // Find the session by session_id
362        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        // Find the target event index
368        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        // Truncate events after the target (keep 0..=target_index)
374        data.events.truncate(target_index + 1);
375
376        // Rebuild session state from remaining events' state deltas
377        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        // Get app and user state (these are not rewound — they are separate)
385        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        // Update the stored session state with rebuilt session-level state
389        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        // Merge with app and user state for the returned session
398        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            // Return the session unchanged
419            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        // Read the event count and determine target
456        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                // Rewinding all events
470                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}