1use std::collections::HashMap;
8use std::path::PathBuf;
9
10use async_trait::async_trait;
11use serde::{Deserialize, Serialize};
12use serde_json::Value;
13
14#[derive(Debug, Clone, Serialize, Deserialize)]
16pub struct SessionSnapshot {
17 pub state: HashMap<String, Value>,
19 pub phase: String,
21 pub turn_count: u32,
23 pub transcript_summary: String,
25 pub resume_handle: Option<String>,
27 pub saved_at: String,
29}
30
31#[derive(Debug, thiserror::Error)]
33pub enum PersistenceError {
34 #[error("persistence I/O error: {0}")]
36 Io(#[from] std::io::Error),
37 #[error("persistence serialization error: {0}")]
39 Serde(#[from] serde_json::Error),
40 #[error("no persisted session '{0}'")]
43 NotFound(String),
44 #[error("persistence backend error: {0}")]
46 Backend(String),
47}
48
49#[async_trait]
53pub trait SessionPersistence: Send + Sync {
54 async fn save(
56 &self,
57 session_id: &str,
58 snapshot: &SessionSnapshot,
59 ) -> Result<(), PersistenceError>;
60
61 async fn load(&self, session_id: &str) -> Result<Option<SessionSnapshot>, PersistenceError>;
63
64 async fn delete(&self, session_id: &str) -> Result<(), PersistenceError>;
66}
67
68pub struct FsPersistence {
70 dir: PathBuf,
71}
72
73impl FsPersistence {
74 pub fn new(dir: impl Into<PathBuf>) -> Self {
78 Self { dir: dir.into() }
79 }
80
81 fn path(&self, session_id: &str) -> PathBuf {
82 self.dir.join(format!("{session_id}.json"))
83 }
84
85 fn tmp_path(&self, session_id: &str) -> PathBuf {
86 self.dir.join(format!("{session_id}.json.tmp"))
87 }
88}
89
90#[async_trait]
91impl SessionPersistence for FsPersistence {
92 async fn save(
93 &self,
94 session_id: &str,
95 snapshot: &SessionSnapshot,
96 ) -> Result<(), PersistenceError> {
97 tokio::fs::create_dir_all(&self.dir).await?;
98 let json = serde_json::to_string_pretty(snapshot)?;
99 let tmp = self.tmp_path(session_id);
105 tokio::fs::write(&tmp, json).await?;
106 tokio::fs::rename(&tmp, self.path(session_id)).await?;
107 Ok(())
108 }
109
110 async fn load(&self, session_id: &str) -> Result<Option<SessionSnapshot>, PersistenceError> {
111 let path = self.path(session_id);
112 match tokio::fs::read_to_string(&path).await {
113 Ok(json) => Ok(Some(serde_json::from_str(&json)?)),
114 Err(e) if e.kind() == std::io::ErrorKind::NotFound => Ok(None),
115 Err(e) => Err(e.into()),
116 }
117 }
118
119 async fn delete(&self, session_id: &str) -> Result<(), PersistenceError> {
120 let path = self.path(session_id);
121 match tokio::fs::remove_file(&path).await {
122 Ok(()) => Ok(()),
123 Err(e) if e.kind() == std::io::ErrorKind::NotFound => Ok(()),
124 Err(e) => Err(e.into()),
125 }
126 }
127}
128
129pub struct MemoryPersistence {
131 store: std::sync::Arc<dashmap::DashMap<String, SessionSnapshot>>,
132}
133
134impl MemoryPersistence {
135 pub fn new() -> Self {
137 Self {
138 store: std::sync::Arc::new(dashmap::DashMap::new()),
139 }
140 }
141}
142
143impl Default for MemoryPersistence {
144 fn default() -> Self {
145 Self::new()
146 }
147}
148
149#[async_trait]
150impl SessionPersistence for MemoryPersistence {
151 async fn save(
152 &self,
153 session_id: &str,
154 snapshot: &SessionSnapshot,
155 ) -> Result<(), PersistenceError> {
156 self.store.insert(session_id.to_string(), snapshot.clone());
157 Ok(())
158 }
159
160 async fn load(&self, session_id: &str) -> Result<Option<SessionSnapshot>, PersistenceError> {
161 Ok(self.store.get(session_id).map(|v| v.value().clone()))
162 }
163
164 async fn delete(&self, session_id: &str) -> Result<(), PersistenceError> {
165 self.store.remove(session_id);
166 Ok(())
167 }
168}
169
170#[cfg(test)]
171mod tests {
172 use super::*;
173
174 #[tokio::test]
175 async fn memory_persistence_round_trip() {
176 let p = MemoryPersistence::new();
177 let snapshot = SessionSnapshot {
178 state: [("name".into(), Value::String("Alice".into()))]
179 .into_iter()
180 .collect(),
181 phase: "greeting".into(),
182 turn_count: 5,
183 transcript_summary: "User: Hello\nAssistant: Hi!".into(),
184 resume_handle: Some("handle-123".into()),
185 saved_at: "2026-03-07T00:00:00Z".into(),
186 };
187
188 p.save("session-1", &snapshot).await.unwrap();
189
190 let loaded = p.load("session-1").await.unwrap().unwrap();
191 assert_eq!(loaded.phase, "greeting");
192 assert_eq!(loaded.turn_count, 5);
193 assert_eq!(loaded.resume_handle, Some("handle-123".into()));
194 }
195
196 #[tokio::test]
197 async fn memory_persistence_load_missing() {
198 let p = MemoryPersistence::new();
199 assert!(p.load("nonexistent").await.unwrap().is_none());
200 }
201
202 #[tokio::test]
203 async fn memory_persistence_delete() {
204 let p = MemoryPersistence::new();
205 let snapshot = SessionSnapshot {
206 state: HashMap::new(),
207 phase: "test".into(),
208 turn_count: 0,
209 transcript_summary: String::new(),
210 resume_handle: None,
211 saved_at: "2026-03-07T00:00:00Z".into(),
212 };
213
214 p.save("session-1", &snapshot).await.unwrap();
215 p.delete("session-1").await.unwrap();
216 assert!(p.load("session-1").await.unwrap().is_none());
217 }
218
219 #[tokio::test]
220 async fn fs_persistence_round_trip() {
221 let dir = std::env::temp_dir().join("gemini_rs_test_persistence");
222 let p = FsPersistence::new(&dir);
223 let snapshot = SessionSnapshot {
224 state: [("key".into(), Value::from(42))].into_iter().collect(),
225 phase: "main".into(),
226 turn_count: 3,
227 transcript_summary: "test".into(),
228 resume_handle: None,
229 saved_at: "2026-03-07T00:00:00Z".into(),
230 };
231
232 p.save("test-session", &snapshot).await.unwrap();
233 let loaded = p.load("test-session").await.unwrap().unwrap();
234 assert_eq!(loaded.phase, "main");
235
236 p.delete("test-session").await.unwrap();
238 let _ = tokio::fs::remove_dir_all(&dir).await;
239 }
240
241 #[tokio::test]
242 async fn fs_persistence_save_is_atomic_and_leaves_no_tmp_file() {
243 let dir = std::env::temp_dir().join(format!(
244 "gemini_rs_test_persistence_atomic_{}",
245 uuid::Uuid::new_v4()
246 ));
247 let p = FsPersistence::new(&dir);
248 let snapshot = SessionSnapshot {
249 state: HashMap::new(),
250 phase: "main".into(),
251 turn_count: 1,
252 transcript_summary: "x".into(),
253 resume_handle: None,
254 saved_at: "now".into(),
255 };
256
257 p.save("atomic-session", &snapshot).await.unwrap();
258
259 assert!(
260 !p.tmp_path("atomic-session").exists(),
261 "tmp file must be renamed away after save"
262 );
263 assert!(p.path("atomic-session").exists());
264
265 let _ = tokio::fs::remove_dir_all(&dir).await;
266 }
267
268 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
269 async fn fs_persistence_concurrent_loads_never_observe_torn_snapshot() {
270 let dir = std::env::temp_dir().join(format!(
274 "gemini_rs_test_persistence_torn_{}",
275 uuid::Uuid::new_v4()
276 ));
277 let p = std::sync::Arc::new(FsPersistence::new(&dir));
278
279 let big = "x".repeat(256 * 1024);
281 let snapshot = SessionSnapshot {
282 state: [("blob".to_string(), Value::String(big))]
283 .into_iter()
284 .collect(),
285 phase: "main".into(),
286 turn_count: 0,
287 transcript_summary: String::new(),
288 resume_handle: None,
289 saved_at: "now".into(),
290 };
291 p.save("torn", &snapshot).await.unwrap();
292
293 let writer = {
294 let p = p.clone();
295 let snapshot = snapshot.clone();
296 tokio::spawn(async move {
297 for _ in 0..50 {
298 p.save("torn", &snapshot).await.unwrap();
299 }
300 })
301 };
302
303 for _ in 0..200 {
304 let loaded = p
305 .load("torn")
306 .await
307 .expect("load must never observe a torn snapshot");
308 assert!(loaded.is_some(), "snapshot must always be present");
309 }
310
311 writer.await.unwrap();
312 let _ = tokio::fs::remove_dir_all(&dir).await;
313 }
314}