agent_base/engine/runtime/
session_manager.rs1use 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 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 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 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 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 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 {
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 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 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 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 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 {
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 async fn enforce_session_limits(&self, session_id: &SessionId) {
248 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 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 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 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 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()); 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}