gemini_adk_rs/live/
persistence.rs

1//! Session persistence — survive process restarts.
2//!
3//! The Gemini Live API supports session resumption via opaque handles.
4//! This module persists the SDK's client-side state (State, phase position,
5//! transcript summary) so it can be restored on reconnection.
6
7use std::collections::HashMap;
8use std::path::PathBuf;
9
10use async_trait::async_trait;
11use serde::{Deserialize, Serialize};
12use serde_json::Value;
13
14/// Serializable snapshot of the control plane state.
15#[derive(Debug, Clone, Serialize, Deserialize)]
16pub struct SessionSnapshot {
17    /// All state key-value pairs.
18    pub state: HashMap<String, Value>,
19    /// Current phase name.
20    pub phase: String,
21    /// Turn count at time of snapshot.
22    pub turn_count: u32,
23    /// Human-readable summary of recent transcript.
24    pub transcript_summary: String,
25    /// Resume handle from the Gemini server.
26    pub resume_handle: Option<String>,
27    /// ISO 8601 timestamp.
28    pub saved_at: String,
29}
30
31/// Error returned by a [`SessionPersistence`] backend.
32#[derive(Debug, thiserror::Error)]
33pub enum PersistenceError {
34    /// Filesystem or network I/O failed.
35    #[error("persistence I/O error: {0}")]
36    Io(#[from] std::io::Error),
37    /// A snapshot could not be encoded or decoded.
38    #[error("persistence serialization error: {0}")]
39    Serde(#[from] serde_json::Error),
40    /// No snapshot is stored under the given session id (for backends that
41    /// distinguish this from an empty `load`).
42    #[error("no persisted session '{0}'")]
43    NotFound(String),
44    /// A backend-specific failure (Redis, Firestore, DynamoDB, …).
45    #[error("persistence backend error: {0}")]
46    Backend(String),
47}
48
49/// Trait for persisting session state across process restarts.
50///
51/// Implementations might write to the filesystem, Redis, Firestore, etc.
52#[async_trait]
53pub trait SessionPersistence: Send + Sync {
54    /// Save a session snapshot.
55    async fn save(
56        &self,
57        session_id: &str,
58        snapshot: &SessionSnapshot,
59    ) -> Result<(), PersistenceError>;
60
61    /// Load a previously saved session snapshot.
62    async fn load(&self, session_id: &str) -> Result<Option<SessionSnapshot>, PersistenceError>;
63
64    /// Delete a saved session.
65    async fn delete(&self, session_id: &str) -> Result<(), PersistenceError>;
66}
67
68/// File-system persistence (good for development and single-server deployments).
69pub struct FsPersistence {
70    dir: PathBuf,
71}
72
73impl FsPersistence {
74    /// Create a new file-system persistence backend.
75    ///
76    /// The directory will be created if it doesn't exist.
77    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        // Write to a sibling temp file, then atomically rename over the
100        // destination. `rename(2)` is atomic on the same filesystem, so a
101        // crash mid-write (or a concurrent `load`) only ever observes the
102        // previous complete snapshot or the new complete snapshot — never a
103        // torn half-write.
104        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
129/// In-memory persistence (good for tests).
130pub struct MemoryPersistence {
131    store: std::sync::Arc<dashmap::DashMap<String, SessionSnapshot>>,
132}
133
134impl MemoryPersistence {
135    /// Create a new in-memory persistence backend.
136    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        // Cleanup
237        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        // Hammer save() while load()ing concurrently: with the tmp+rename
271        // scheme every load parses; with the old direct `fs::write` a reader
272        // could observe a truncated/partial file mid-write.
273        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        // A snapshot large enough that a write is not a single tiny syscall.
280        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}