Skip to main content

agent_base/engine/
session_store.rs

1use 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/// Session Persistence Adapter
10///
11/// `SessionStore` is an optional persistence interface for agent sessions.
12/// Under the lightweight kernel design:
13/// - `AgentRuntime.sessions` is the authoritative live state during execution
14/// - `SessionStore` is a persistence adapter for save/load/list/delete
15/// - Does not participate in the execution control flow
16///
17/// Replace the default [`InMemorySessionStore`] with a custom implementation
18/// to persist sessions to a database, filesystem, or other storage.
19#[async_trait]
20pub trait SessionStore: Send + Sync {
21    /// Save a session snapshot to the persistence layer
22    async fn save(&self, session: &AgentSession) -> AgentResult<()>;
23
24    /// Load a session from the persistence layer
25    async fn load(&self, session_id: &SessionId) -> AgentResult<Option<AgentSession>>;
26
27    /// List all saved session IDs
28    async fn list(&self) -> AgentResult<Vec<SessionId>>;
29
30    /// Delete a specific session
31    async fn delete(&self, session_id: &SessionId) -> AgentResult<()>;
32
33    /// Optionally persist a runtime event for audit/replay.
34    /// Default implementation is a no-op — override to implement
35    /// append-log style persistence (e.g. JSONL file, database event log).
36    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// ── SqliteSessionStore ──
91
92#[cfg(feature = "sqlite-session")]
93use std::sync::Mutex as StdMutex;
94
95/// SQLite-backed session persistence.
96///
97/// Stores each [`AgentSession`] as a JSON blob in a single `sessions` table.
98/// Enable with `features = ["sqlite-session"]` in your `Cargo.toml`.
99///
100/// # Example
101///
102/// ```ignore
103/// use agent_base::engine::SqliteSessionStore;
104/// let store = SqliteSessionStore::open("sessions.db").unwrap();
105/// ```
106#[cfg(feature = "sqlite-session")]
107pub struct SqliteSessionStore {
108    db: StdMutex<rusqlite::Connection>,
109}
110
111#[cfg(feature = "sqlite-session")]
112impl SqliteSessionStore {
113    /// Open (or create) a SQLite database at `path` and initialize the schema.
114    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    /// Build from an already-open [`rusqlite::Connection`].
124    ///
125    /// The connection must outlive the store. Caller is responsible for closing it.
126    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        // Use serde_json for robust round-tripping — never depend on Display format.
145        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            // Keys are serialized with serde_json — deserialize for robust round-tripping.
198            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// ── Tests ──
225
226#[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        // Should not error
307        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(); // id is None
332        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        // Default append_event should not error
340        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}