agent_base/engine/
session_store.rs1use std::collections::HashMap;
2
3use async_trait::async_trait;
4use tokio::sync::Mutex;
5
6use super::AgentSession;
7use crate::types::{AgentError, AgentResult, RuntimeEvent, SessionId};
8
9#[async_trait]
20pub trait SessionStore: Send + Sync {
21 async fn save(&self, session: &AgentSession) -> AgentResult<()>;
23
24 async fn load(&self, session_id: &SessionId) -> AgentResult<Option<AgentSession>>;
26
27 async fn list(&self) -> AgentResult<Vec<SessionId>>;
29
30 async fn delete(&self, session_id: &SessionId) -> AgentResult<()>;
32
33 async fn append_event(
37 &self,
38 _session_id: &SessionId,
39 _event: &RuntimeEvent,
40 ) -> AgentResult<()> {
41 Ok(())
42 }
43}
44
45pub struct InMemorySessionStore {
46 sessions: Mutex<HashMap<SessionId, AgentSession>>,
47}
48
49impl InMemorySessionStore {
50 pub fn new() -> Self {
51 Self {
52 sessions: Mutex::new(HashMap::new()),
53 }
54 }
55}
56
57impl Default for InMemorySessionStore {
58 fn default() -> Self {
59 Self::new()
60 }
61}
62
63#[async_trait]
64impl SessionStore for InMemorySessionStore {
65 async fn save(&self, session: &AgentSession) -> AgentResult<()> {
66 let session_id = session
67 .id()
68 .ok_or_else(|| AgentError::internal("session has no id"))?;
69 self.sessions
70 .lock()
71 .await
72 .insert(session_id, session.clone());
73 Ok(())
74 }
75
76 async fn load(&self, session_id: &SessionId) -> AgentResult<Option<AgentSession>> {
77 Ok(self.sessions.lock().await.get(session_id).cloned())
78 }
79
80 async fn list(&self) -> AgentResult<Vec<SessionId>> {
81 Ok(self.sessions.lock().await.keys().cloned().collect())
82 }
83
84 async fn delete(&self, session_id: &SessionId) -> AgentResult<()> {
85 self.sessions.lock().await.remove(session_id);
86 Ok(())
87 }
88}
89
90#[cfg(feature = "sqlite-session")]
93use std::sync::Mutex as StdMutex;
94
95#[cfg(feature = "sqlite-session")]
107pub struct SqliteSessionStore {
108 db: StdMutex<rusqlite::Connection>,
109}
110
111#[cfg(feature = "sqlite-session")]
112impl SqliteSessionStore {
113 pub fn open(path: impl AsRef<std::path::Path>) -> AgentResult<Self> {
115 let conn = rusqlite::Connection::open(path)
116 .map_err(|e| AgentError::internal(format!("sqlite open: {e}")))?;
117 Self::init_tables(&conn)?;
118 Ok(Self {
119 db: StdMutex::new(conn),
120 })
121 }
122
123 pub fn from_connection(conn: rusqlite::Connection) -> AgentResult<Self> {
127 Self::init_tables(&conn)?;
128 Ok(Self {
129 db: StdMutex::new(conn),
130 })
131 }
132
133 fn init_tables(conn: &rusqlite::Connection) -> AgentResult<()> {
134 conn.execute_batch(
135 "CREATE TABLE IF NOT EXISTS sessions (
136 id TEXT PRIMARY KEY,
137 data TEXT NOT NULL
138 );",
139 )
140 .map_err(|e| AgentError::internal(format!("sqlite init: {e}")))
141 }
142
143 fn session_key(id: &SessionId) -> String {
144 serde_json::to_string(id).unwrap_or_else(|_| id.to_string())
146 }
147}
148
149#[cfg(feature = "sqlite-session")]
150#[async_trait]
151impl SessionStore for SqliteSessionStore {
152 async fn save(&self, session: &AgentSession) -> AgentResult<()> {
153 let session_id = session
154 .id()
155 .ok_or_else(|| AgentError::internal("session has no id"))?;
156 let key = Self::session_key(&session_id);
157 let data = serde_json::to_string(session).map_err(|e| AgentError::json(e.to_string()))?;
158 let db = self.db.lock().unwrap_or_else(|e| e.into_inner());
159 db.execute(
160 "INSERT OR REPLACE INTO sessions (id, data) VALUES (?1, ?2)",
161 rusqlite::params![key, data],
162 )
163 .map_err(|e| AgentError::internal(format!("sqlite save: {e}")))?;
164 Ok(())
165 }
166
167 async fn load(&self, session_id: &SessionId) -> AgentResult<Option<AgentSession>> {
168 let key = Self::session_key(session_id);
169 let db = self.db.lock().unwrap_or_else(|e| e.into_inner());
170 let mut stmt = db
171 .prepare("SELECT data FROM sessions WHERE id = ?1")
172 .map_err(|e| AgentError::internal(format!("sqlite prepare: {e}")))?;
173 let result: Result<String, rusqlite::Error> =
174 stmt.query_row(rusqlite::params![key], |row| row.get(0));
175 match result {
176 Ok(data) => {
177 let session: AgentSession = serde_json::from_str(&data)
178 .map_err(|e| AgentError::json(format!("deserialize session: {e}")))?;
179 Ok(Some(session))
180 }
181 Err(rusqlite::Error::QueryReturnedNoRows) => Ok(None),
182 Err(e) => Err(AgentError::internal(format!("sqlite load: {e}"))),
183 }
184 }
185
186 async fn list(&self) -> AgentResult<Vec<SessionId>> {
187 let db = self.db.lock().unwrap_or_else(|e| e.into_inner());
188 let mut stmt = db
189 .prepare("SELECT id FROM sessions ORDER BY id")
190 .map_err(|e| AgentError::internal(format!("sqlite prepare: {e}")))?;
191 let rows = stmt
192 .query_map([], |row| row.get::<_, String>(0))
193 .map_err(|e| AgentError::internal(format!("sqlite list: {e}")))?;
194 let mut ids = Vec::new();
195 for row in rows {
196 let id_str = row.map_err(|e| AgentError::internal(format!("sqlite list row: {e}")))?;
197 match serde_json::from_str::<SessionId>(&id_str) {
199 Ok(sid) => ids.push(sid),
200 Err(e) => {
201 tracing::warn!(key = id_str, error = %e, "failed to deserialize session key, skipping");
202 }
203 }
204 }
205 Ok(ids)
206 }
207
208 async fn delete(&self, session_id: &SessionId) -> AgentResult<()> {
209 let key = Self::session_key(session_id);
210 let db = self.db.lock().unwrap_or_else(|e| e.into_inner());
211 db.execute("DELETE FROM sessions WHERE id = ?1", rusqlite::params![key])
212 .map_err(|e| AgentError::internal(format!("sqlite delete: {e}")))?;
213 Ok(())
214 }
215}
216
217#[cfg(feature = "sqlite-session")]
218impl std::fmt::Debug for SqliteSessionStore {
219 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
220 f.debug_struct("SqliteSessionStore").finish_non_exhaustive()
221 }
222}
223
224#[cfg(test)]
227#[cfg(feature = "sqlite-session")]
228mod sqlite_tests {
229 use super::*;
230
231 fn make_store() -> SqliteSessionStore {
232 let conn = rusqlite::Connection::open_in_memory().expect("open in-memory sqlite");
233 SqliteSessionStore::from_connection(conn).expect("init tables")
234 }
235
236 fn make_session(id: u64) -> AgentSession {
237 let mut s = AgentSession::new(SessionId::new(id));
238 s.push_message(crate::types::MessageRole::User, "hello");
239 s.push_message(crate::types::MessageRole::Assistant, "hi");
240 s
241 }
242
243 #[tokio::test]
244 async fn save_and_load() {
245 let store = make_store();
246 let session = make_session(1);
247 store.save(&session).await.unwrap();
248
249 let loaded = store.load(&SessionId::new(1)).await.unwrap();
250 assert!(loaded.is_some());
251 let loaded = loaded.unwrap();
252 assert_eq!(loaded.chat_messages().len(), 2);
253 }
254
255 #[tokio::test]
256 async fn load_nonexistent_returns_none() {
257 let store = make_store();
258 let result = store.load(&SessionId::new(999)).await.unwrap();
259 assert!(result.is_none());
260 }
261
262 #[tokio::test]
263 async fn save_overwrites() {
264 let store = make_store();
265 let mut session = make_session(1);
266 store.save(&session).await.unwrap();
267
268 session.push_message(crate::types::MessageRole::User, "another");
269 store.save(&session).await.unwrap();
270
271 let loaded = store.load(&SessionId::new(1)).await.unwrap().unwrap();
272 assert_eq!(loaded.chat_messages().len(), 3);
273 }
274
275 #[tokio::test]
276 async fn list_sessions() {
277 let store = make_store();
278 store.save(&make_session(1)).await.unwrap();
279 store.save(&make_session(2)).await.unwrap();
280 store.save(&make_session(3)).await.unwrap();
281
282 let ids = store.list().await.unwrap();
283 assert_eq!(ids.len(), 3);
284 }
285
286 #[tokio::test]
287 async fn list_empty() {
288 let store = make_store();
289 let ids = store.list().await.unwrap();
290 assert!(ids.is_empty());
291 }
292
293 #[tokio::test]
294 async fn delete_session() {
295 let store = make_store();
296 store.save(&make_session(1)).await.unwrap();
297 assert!(store.load(&SessionId::new(1)).await.unwrap().is_some());
298
299 store.delete(&SessionId::new(1)).await.unwrap();
300 assert!(store.load(&SessionId::new(1)).await.unwrap().is_none());
301 }
302
303 #[tokio::test]
304 async fn delete_nonexistent_is_noop() {
305 let store = make_store();
306 store.delete(&SessionId::new(999)).await.unwrap();
308 }
309
310 #[tokio::test]
311 async fn save_session_with_external_id() {
312 let store = make_store();
313 let mut s = AgentSession::new(SessionId::with_external_id(42, "my-ext-id"));
314 s.push_message(crate::types::MessageRole::User, "test");
315 store.save(&s).await.unwrap();
316
317 let loaded = store
318 .load(&SessionId::with_external_id(42, "my-ext-id"))
319 .await
320 .unwrap();
321 assert!(loaded.is_some());
322
323 let ids = store.list().await.unwrap();
324 assert_eq!(ids.len(), 1);
325 assert_eq!(ids[0].to_string(), "42(my-ext-id)");
326 }
327
328 #[tokio::test]
329 async fn save_without_id_errors() {
330 let store = make_store();
331 let session = AgentSession::default(); let err = store.save(&session).await.unwrap_err();
333 assert!(err.to_string().contains("no id"));
334 }
335
336 #[tokio::test]
337 async fn append_event_default_noop() {
338 let store = make_store();
339 store
341 .append_event(
342 &SessionId::new(1),
343 &RuntimeEvent::UserEvent {
344 session_id: SessionId::new(1),
345 event: crate::types::UserEvent::Progress {
346 text: "test".into(),
347 },
348 agent_id: None,
349 trace_id: None,
350 },
351 )
352 .await
353 .unwrap();
354 }
355
356 #[tokio::test]
357 async fn roundtrip_preserves_fields() {
358 let store = make_store();
359 let mut session = AgentSession::new(SessionId::new(7));
360 session.push_message(crate::types::MessageRole::System, "system prompt");
361 session.push_message(crate::types::MessageRole::User, "question");
362 session.push_message(crate::types::MessageRole::Assistant, "answer");
363 session.allow_action("read_file");
364 session.total_tool_calls = 3;
365
366 store.save(&session).await.unwrap();
367 let loaded = store.load(&SessionId::new(7)).await.unwrap().unwrap();
368
369 assert_eq!(loaded.chat_messages().len(), 3);
370 assert!(loaded.is_action_allowed("read_file"));
371 assert_eq!(loaded.total_tool_calls, 3);
372 }
373}