Skip to main content

agent_base/engine/runtime/
session_manager.rs

1use std::collections::HashMap;
2use std::sync::Arc;
3use std::time::Instant;
4
5use tokio::sync::{Mutex, RwLock};
6
7use crate::engine::AgentSession;
8use crate::engine::context::ContextWindowManager;
9use crate::engine::session_store::SessionStore;
10use crate::types::{
11    AgentError, AgentResult, MessageRole, SessionConfig, SessionId, SessionIdGenerator,
12};
13
14#[derive(Clone)]
15pub struct SessionManager {
16    session_id_generator: Arc<dyn SessionIdGenerator>,
17    sessions: Arc<RwLock<HashMap<SessionId, AgentSession>>>,
18    /// Separate LRU timestamp map — protected by a lightweight Mutex,
19    /// so session reads (RwLock::read) don't need exclusive access just to
20    /// update the last-active time.
21    lru_times: Arc<Mutex<HashMap<SessionId, Instant>>>,
22    session_store: Arc<dyn SessionStore>,
23    config: SessionConfig,
24}
25
26impl SessionManager {
27    pub fn new(
28        session_id_generator: Arc<dyn SessionIdGenerator>,
29        session_store: Arc<dyn SessionStore>,
30        config: SessionConfig,
31    ) -> Self {
32        Self {
33            session_id_generator,
34            sessions: Arc::new(RwLock::new(HashMap::new())),
35            lru_times: Arc::new(Mutex::new(HashMap::new())),
36            session_store,
37            config,
38        }
39    }
40
41    pub async fn create_session(&self, system_prompt: Option<&str>) -> SessionId {
42        // Eviction is best-effort — if it fails, we still create the session
43        if let Err(e) = self.evict_if_needed().await {
44            tracing::warn!(error = %e, "session eviction failed, proceeding with creation");
45        }
46
47        let id = self.session_id_generator.generate();
48        let mut session = AgentSession::new(id.clone());
49        if let Some(prompt) = system_prompt {
50            session.push_message(MessageRole::System, prompt);
51        }
52        {
53            let mut sessions = self.sessions.write().await;
54            sessions.insert(id.clone(), session);
55        }
56        {
57            let mut lru = self.lru_times.lock().await;
58            lru.insert(id.clone(), Instant::now());
59        }
60        tracing::debug!(session_id = id.id, "session created");
61        id
62    }
63
64    pub async fn restore_session(&self, session_id: &SessionId) -> Option<AgentSession> {
65        {
66            let sessions = self.sessions.read().await;
67            if sessions.contains_key(session_id) {
68                let mut lru = self.lru_times.lock().await;
69                lru.insert(session_id.clone(), Instant::now());
70                tracing::debug!(session_id = session_id.id, "session restore cache hit");
71                return sessions.get(session_id).cloned();
72            }
73        }
74        match self.session_store.load(session_id).await {
75            Ok(Some(session)) => {
76                let msg_count = session.chat_messages().len();
77                // Validate the restored message sequence — corrupt persisted data
78                // would cause LLM API errors downstream, so warn early.
79                if let Err(e) =
80                    crate::engine::session::validate_message_sequence(session.chat_messages())
81                {
82                    tracing::warn!(session_id = session_id.id, error = %e, "restored session has invalid message sequence");
83                }
84                // Evict before inserting to make room if needed
85                self.evict_if_needed().await.ok();
86                {
87                    let mut sessions = self.sessions.write().await;
88                    sessions.insert(session_id.clone(), session.clone());
89                }
90                {
91                    let mut lru = self.lru_times.lock().await;
92                    lru.insert(session_id.clone(), Instant::now());
93                }
94                tracing::debug!(
95                    session_id = session_id.id,
96                    msg_count,
97                    "session restored from store"
98                );
99                Some(session)
100            }
101            Ok(None) => {
102                tracing::debug!(session_id = session_id.id, "session not found in store");
103                None
104            }
105            Err(e) => {
106                tracing::warn!(session_id = session_id.id, error = %e, "session restore failed");
107                None
108            }
109        }
110    }
111
112    pub async fn session(&self, session_id: &SessionId) -> Option<AgentSession> {
113        let sessions = self.sessions.read().await;
114        let result = sessions.get(session_id).cloned();
115        if result.is_some() {
116            let mut lru = self.lru_times.lock().await;
117            lru.insert(session_id.clone(), Instant::now());
118        }
119        result
120    }
121
122    pub async fn session_or_err(&self, session_id: &SessionId) -> AgentResult<AgentSession> {
123        let sessions = self.sessions.read().await;
124        let result = sessions
125            .get(session_id)
126            .cloned()
127            .ok_or_else(|| AgentError::session_not_found(session_id.id));
128        if result.is_ok() {
129            let mut lru = self.lru_times.lock().await;
130            lru.insert(session_id.clone(), Instant::now());
131        }
132        result
133    }
134
135    pub async fn with_session_mut<F, R>(&self, session_id: &SessionId, f: F) -> AgentResult<R>
136    where
137        F: FnOnce(&mut AgentSession) -> R,
138    {
139        // The write lock on `sessions` is held only for the closure execution.
140        // It MUST be dropped before `enforce_session_limits` — that method may
141        // call `save_session` → `session_or_err` which re-acquires a (read) lock
142        // on the same `sessions` map. Holding the write lock across that call
143        // would deadlock.
144        let result = {
145            let mut sessions = self.sessions.write().await;
146            let session = sessions
147                .get_mut(session_id)
148                .ok_or_else(|| AgentError::session_not_found(session_id.id))?;
149            f(session)
150        };
151        // Update LRU timestamp (lightweight Mutex, not RwLock)
152        {
153            let mut lru = self.lru_times.lock().await;
154            lru.insert(session_id.clone(), Instant::now());
155        }
156        self.enforce_session_limits(session_id).await;
157        Ok(result)
158    }
159
160    pub async fn cached_approval(&self, session_id: &SessionId, action_key: &str) -> bool {
161        let sessions = self.sessions.read().await;
162        sessions
163            .get(session_id)
164            .is_some_and(|session| session.is_action_allowed(action_key))
165    }
166
167    pub async fn cache_approval(&self, session_id: &SessionId, action_key: String) {
168        let mut sessions = self.sessions.write().await;
169        if let Some(session) = sessions.get_mut(session_id) {
170            session.allow_action(action_key);
171        } else {
172            // `get_mut` returns None when the session doesn't exist yet — the
173            // approval cache is silently a no-op in that case. Surface it so a
174            // caller caching before session creation isn't left wondering why
175            // approval was re-prompted.
176            tracing::warn!(
177                session_id = session_id.id,
178                action = %action_key,
179                "cache_approval ignored: session not found"
180            );
181        }
182    }
183
184    pub async fn save_session(&self, session_id: &SessionId) -> AgentResult<()> {
185        let session = self.session_or_err(session_id).await?;
186        let msg_count = session.chat_messages().len();
187        tracing::debug!(session_id = session_id.id, msg_count, "saving session");
188        self.session_store
189            .save(&session)
190            .await
191            .map_err(|e| AgentError::internal(format!("Session persistence failed: {e}")))
192    }
193
194    pub fn session_store(&self) -> &Arc<dyn SessionStore> {
195        &self.session_store
196    }
197
198    /// Evict the least recently used session if at capacity.
199    ///
200    /// Note: there is a benign TOCTOU window between the read-lock capacity
201    /// check and the write-lock removal — another task may concurrently create
202    /// a session and also trigger eviction. The worst case is one extra eviction,
203    /// which is harmless since `restore_session` transparently reloads from store.
204    async fn evict_if_needed(&self) -> AgentResult<()> {
205        let max = match self.config.max_sessions {
206            Some(m) => m,
207            None => return Ok(()),
208        };
209
210        // Find the LRU victim (read lock on sessions, Mutex on lru_times)
211        let victim = {
212            let sessions = self.sessions.read().await;
213            if sessions.len() < max {
214                return Ok(());
215            }
216            let lru = self.lru_times.lock().await;
217            sessions
218                .keys()
219                .min_by_key(|id| lru.get(*id).copied().unwrap_or(Instant::now()))
220                .cloned()
221        };
222
223        let Some(victim_id) = victim else {
224            return Ok(());
225        };
226
227        // Persist before evicting
228        if let Err(e) = self.save_session(&victim_id).await {
229            tracing::warn!(session_id = victim_id.id, error = %e, "failed to persist session before eviction");
230        }
231
232        // Remove from both maps
233        {
234            let mut sessions = self.sessions.write().await;
235            sessions.remove(&victim_id);
236        }
237        {
238            let mut lru = self.lru_times.lock().await;
239            lru.remove(&victim_id);
240        }
241        tracing::info!(session_id = victim_id.id, "session evicted (LRU)");
242
243        Ok(())
244    }
245
246    /// Enforce turn count and message size limits after a session mutation.
247    async fn enforce_session_limits(&self, session_id: &SessionId) {
248        // Layer 2: Turn count trimming
249        if let Some(max_turns) = self.config.max_turns_per_session {
250            let needs_trim = {
251                let sessions = self.sessions.read().await;
252                sessions
253                    .get(session_id)
254                    .is_some_and(|session| session.turn_count() > max_turns)
255            };
256
257            if needs_trim {
258                // Persist full history before trimming
259                if let Err(e) = self.save_session(session_id).await {
260                    tracing::warn!(session_id = session_id.id, error = %e, "failed to persist before turn trim");
261                }
262                // Trim under write lock
263                let mut sessions = self.sessions.write().await;
264                if let Some(session) = sessions.get_mut(session_id) {
265                    let before = session.turn_count();
266                    session.trim_oldest_turns(max_turns);
267                    tracing::info!(
268                        session_id = session_id.id,
269                        before,
270                        after = session.turn_count(),
271                        max_turns,
272                        "session turns trimmed"
273                    );
274                }
275            }
276        }
277
278        // Layer 3: Oversized message safety valve
279        if let Some(max_tokens) = self.config.max_message_tokens {
280            let mut sessions = self.sessions.write().await;
281            if let Some(session) = sessions.get_mut(session_id)
282                && let Some(last) = session.chat_messages().last()
283            {
284                let tokens = ContextWindowManager::message_tokens(last);
285                if tokens > max_tokens {
286                    session.pop_last_message();
287                    tracing::warn!(
288                        session_id = session_id.id,
289                        tokens,
290                        max_tokens,
291                        "oversized message removed from session (safety valve)"
292                    );
293                }
294            }
295        }
296    }
297}
298
299#[cfg(test)]
300mod tests {
301    use super::*;
302    use crate::engine::InMemorySessionStore;
303    use crate::types::{AtomicU64SessionIdGenerator, ChatMessage};
304    use std::sync::Arc;
305
306    fn manager() -> SessionManager {
307        SessionManager::new(
308            Arc::new(AtomicU64SessionIdGenerator::default()),
309            Arc::new(InMemorySessionStore::new()),
310            SessionConfig::default(),
311        )
312    }
313
314    fn manager_with_config(config: SessionConfig) -> SessionManager {
315        SessionManager::new(
316            Arc::new(AtomicU64SessionIdGenerator::default()),
317            Arc::new(InMemorySessionStore::new()),
318            config,
319        )
320    }
321
322    struct FailingStore;
323
324    #[async_trait::async_trait]
325    impl SessionStore for FailingStore {
326        async fn save(&self, _session: &AgentSession) -> AgentResult<()> {
327            Ok(())
328        }
329        async fn load(&self, _session_id: &SessionId) -> AgentResult<Option<AgentSession>> {
330            Err(AgentError::internal("boom"))
331        }
332        async fn list(&self) -> AgentResult<Vec<SessionId>> {
333            Ok(vec![])
334        }
335        async fn delete(&self, _session_id: &SessionId) -> AgentResult<()> {
336            Ok(())
337        }
338    }
339
340    #[tokio::test]
341    async fn create_session_with_system_prompt() {
342        let m = manager();
343        let id = m.create_session(Some("be helpful")).await;
344        assert_eq!(id.id, 1);
345
346        let session = m.session(&id).await.unwrap();
347        assert_eq!(session.chat_messages().len(), 1);
348        assert!(matches!(
349            session.chat_messages()[0],
350            ChatMessage::System { .. }
351        ));
352    }
353
354    #[tokio::test]
355    async fn create_session_without_prompt_is_empty() {
356        let m = manager();
357        let id = m.create_session(None).await;
358        let session = m.session(&id).await.unwrap();
359        assert!(session.chat_messages().is_empty());
360    }
361
362    #[tokio::test]
363    async fn session_returns_none_for_unknown() {
364        let m = manager();
365        assert!(m.session(&SessionId::new(999)).await.is_none());
366    }
367
368    #[tokio::test]
369    async fn session_or_err_returns_error_for_unknown() {
370        let m = manager();
371        let err = m.session_or_err(&SessionId::new(999)).await.unwrap_err();
372        assert!(matches!(err, AgentError::SessionNotFound(_)));
373    }
374
375    #[tokio::test]
376    async fn restore_session_cache_hit() {
377        let m = manager();
378        let id = m.create_session(Some("sys")).await;
379        let restored = m.restore_session(&id).await;
380        assert!(restored.is_some());
381        assert_eq!(restored.unwrap().chat_messages().len(), 1);
382    }
383
384    #[tokio::test]
385    async fn restore_session_from_store() {
386        let store = Arc::new(InMemorySessionStore::new());
387        let mut session = AgentSession::new(SessionId::new(42));
388        session.push_message(MessageRole::User, "persisted");
389        store.save(&session).await.unwrap();
390
391        let m = SessionManager::new(
392            Arc::new(AtomicU64SessionIdGenerator::default()),
393            store.clone(),
394            SessionConfig::default(),
395        );
396        let restored = m.restore_session(&SessionId::new(42)).await;
397        assert!(restored.is_some());
398        assert_eq!(restored.unwrap().chat_messages().len(), 1);
399        // now cached in memory
400        assert!(m.session(&SessionId::new(42)).await.is_some());
401    }
402
403    #[tokio::test]
404    async fn restore_session_returns_none_when_not_found() {
405        let m = manager();
406        assert!(m.restore_session(&SessionId::new(999)).await.is_none());
407    }
408
409    #[tokio::test]
410    async fn restore_session_returns_none_on_store_error() {
411        let m = SessionManager::new(
412            Arc::new(AtomicU64SessionIdGenerator::default()),
413            Arc::new(FailingStore),
414            SessionConfig::default(),
415        );
416        assert!(m.restore_session(&SessionId::new(1)).await.is_none());
417    }
418
419    #[tokio::test]
420    async fn with_session_mut_applies_closure() {
421        let m = manager();
422        let id = m.create_session(None).await;
423        m.with_session_mut(&id, |s| {
424            s.push_message(MessageRole::User, "hello");
425            s.push_message(MessageRole::Assistant, "hi");
426        })
427        .await
428        .unwrap();
429
430        let session = m.session(&id).await.unwrap();
431        assert_eq!(session.chat_messages().len(), 2);
432    }
433
434    #[tokio::test]
435    async fn with_session_mut_errors_for_unknown() {
436        let m = manager();
437        let err = m
438            .with_session_mut(&SessionId::new(999), |_s| ())
439            .await
440            .unwrap_err();
441        assert!(matches!(err, AgentError::SessionNotFound(_)));
442    }
443
444    #[tokio::test]
445    async fn approval_cache_roundtrip() {
446        let m = manager();
447        let id = m.create_session(None).await;
448        assert!(!m.cached_approval(&id, "read_file").await);
449
450        m.cache_approval(&id, "read_file".into()).await;
451        assert!(m.cached_approval(&id, "read_file").await);
452    }
453
454    #[tokio::test]
455    async fn save_session_persists_to_store() {
456        let store = Arc::new(InMemorySessionStore::new());
457        let m = SessionManager::new(
458            Arc::new(AtomicU64SessionIdGenerator::default()),
459            store.clone(),
460            SessionConfig::default(),
461        );
462        let id = m.create_session(Some("sys")).await;
463        m.with_session_mut(&id, |s| s.push_message(MessageRole::User, "hello"))
464            .await
465            .unwrap();
466        m.save_session(&id).await.unwrap();
467
468        let loaded = store.load(&id).await.unwrap();
469        assert!(loaded.is_some());
470        assert_eq!(loaded.unwrap().chat_messages().len(), 2);
471    }
472
473    #[tokio::test]
474    async fn save_session_errors_for_unknown() {
475        let m = manager();
476        let err = m.save_session(&SessionId::new(999)).await.unwrap_err();
477        assert!(matches!(err, AgentError::SessionNotFound(_)));
478    }
479
480    #[tokio::test]
481    async fn session_store_getter_returns_store() {
482        let m = manager();
483        assert!(m.session_store().list().await.unwrap().is_empty());
484    }
485
486    #[tokio::test]
487    async fn eviction_evicts_lru_when_at_capacity() {
488        let cfg = SessionConfig {
489            max_sessions: Some(2),
490            ..Default::default()
491        };
492        let m = manager_with_config(cfg);
493        let id1 = m.create_session(Some("s1")).await;
494        let id2 = m.create_session(Some("s2")).await;
495        let id3 = m.create_session(Some("s3")).await;
496
497        assert!(m.session(&id1).await.is_none()); // LRU evicted
498        assert!(m.session(&id2).await.is_some());
499        assert!(m.session(&id3).await.is_some());
500    }
501
502    #[tokio::test]
503    async fn turn_trimming_enforced_after_mutation() {
504        let cfg = SessionConfig {
505            max_turns_per_session: Some(1),
506            ..Default::default()
507        };
508        let m = manager_with_config(cfg);
509        let id = m.create_session(None).await;
510        m.with_session_mut(&id, |s| {
511            s.push_message(MessageRole::User, "u1");
512            s.push_message(MessageRole::Assistant, "a1");
513            s.push_message(MessageRole::User, "u2");
514            s.push_message(MessageRole::Assistant, "a2");
515        })
516        .await
517        .unwrap();
518
519        let session = m.session(&id).await.unwrap();
520        assert_eq!(session.turn_count(), 1);
521    }
522
523    #[tokio::test]
524    async fn oversized_message_removed_after_mutation() {
525        let cfg = SessionConfig {
526            max_message_tokens: Some(10),
527            ..Default::default()
528        };
529        let m = manager_with_config(cfg);
530        let id = m.create_session(None).await;
531        let big = "x".repeat(200);
532        m.with_session_mut(&id, |s| s.push_message(MessageRole::User, big))
533            .await
534            .unwrap();
535
536        let session = m.session(&id).await.unwrap();
537        assert!(session.chat_messages().is_empty());
538    }
539}