Skip to main content

ai_agents_storage/
storage.rs

1use std::collections::BTreeSet;
2use std::path::{Component, Path, PathBuf};
3
4use async_trait::async_trait;
5use serde::{Deserialize, Serialize};
6use sha2::{Digest, Sha256};
7
8#[cfg(test)]
9use ai_agents_core::traits::storage::StorageCapability;
10use ai_agents_core::{AgentError, AgentSnapshot, AgentStorage, Result};
11
12const HASHED_PREFIX: &str = "v2-";
13const ENCODED_PREFIX: &str = "v1-";
14const ENCODED_EXTENSION: &str = "snapshot";
15const ENVELOPE_VERSION: u8 = 1;
16
17#[derive(Serialize, Deserialize)]
18struct SnapshotEnvelope {
19    version: u8,
20    session_id: String,
21    snapshot: AgentSnapshot,
22}
23
24pub struct FileStorage {
25    base_path: PathBuf,
26}
27
28impl FileStorage {
29    pub fn new(base_path: impl AsRef<Path>) -> Self {
30        Self {
31            base_path: base_path.as_ref().to_path_buf(),
32        }
33    }
34
35    fn session_path(&self, session_id: &str) -> PathBuf {
36        self.base_path.join(format!(
37            "{}.{}",
38            hashed_session_id(session_id),
39            ENCODED_EXTENSION
40        ))
41    }
42
43    fn encoded_session_path(&self, session_id: &str) -> PathBuf {
44        self.base_path.join(format!(
45            "{}.{}",
46            encode_session_id(session_id),
47            ENCODED_EXTENSION
48        ))
49    }
50
51    fn legacy_session_path(&self, session_id: &str) -> Option<PathBuf> {
52        if !is_safe_legacy_session_id(session_id) {
53            return None;
54        }
55        Some(self.base_path.join(format!("{session_id}.json")))
56    }
57}
58
59fn hashed_session_id(session_id: &str) -> String {
60    let digest = Sha256::digest(session_id.as_bytes());
61    format!("{HASHED_PREFIX}{digest:x}")
62}
63
64fn encode_session_id(session_id: &str) -> String {
65    let mut encoded = String::with_capacity(ENCODED_PREFIX.len() + session_id.len() * 2);
66    encoded.push_str(ENCODED_PREFIX);
67    for byte in session_id.as_bytes() {
68        use std::fmt::Write;
69        write!(&mut encoded, "{byte:02x}").expect("writing to a String cannot fail");
70    }
71    encoded
72}
73
74fn decode_session_id(encoded: &str) -> Option<String> {
75    let encoded = encoded.strip_prefix(ENCODED_PREFIX)?;
76    let chunks = encoded.as_bytes().chunks_exact(2);
77    if !chunks.remainder().is_empty() {
78        return None;
79    }
80
81    let bytes = chunks
82        .map(|chunk| {
83            let high = decode_hex_digit(chunk[0])?;
84            let low = decode_hex_digit(chunk[1])?;
85            Some((high << 4) | low)
86        })
87        .collect::<Option<Vec<_>>>()?;
88    String::from_utf8(bytes).ok()
89}
90
91fn decode_hex_digit(byte: u8) -> Option<u8> {
92    match byte {
93        b'0'..=b'9' => Some(byte - b'0'),
94        b'a'..=b'f' => Some(byte - b'a' + 10),
95        b'A'..=b'F' => Some(byte - b'A' + 10),
96        _ => None,
97    }
98}
99
100fn is_safe_legacy_session_id(session_id: &str) -> bool {
101    if session_id.is_empty() || session_id.contains('/') || session_id.contains('\\') {
102        return false;
103    }
104
105    let filename = format!("{session_id}.json");
106    let mut components = Path::new(&filename).components();
107    matches!(
108        (components.next(), components.next()),
109        (Some(Component::Normal(_)), None)
110    )
111}
112
113#[async_trait]
114impl AgentStorage for FileStorage {
115    async fn save(&self, session_id: &str, snapshot: &AgentSnapshot) -> Result<()> {
116        tokio::fs::create_dir_all(&self.base_path).await?;
117        let envelope = SnapshotEnvelope {
118            version: ENVELOPE_VERSION,
119            session_id: session_id.to_string(),
120            snapshot: snapshot.clone(),
121        };
122        tokio::fs::write(
123            self.session_path(session_id),
124            serde_json::to_vec_pretty(&envelope)?,
125        )
126        .await?;
127
128        let encoded_path = self.encoded_session_path(session_id);
129        if encoded_path.exists() {
130            tokio::fs::remove_file(encoded_path).await?;
131        }
132        if let Some(legacy_path) = self.legacy_session_path(session_id)
133            && legacy_path.exists()
134        {
135            tokio::fs::remove_file(legacy_path).await?;
136        }
137        Ok(())
138    }
139
140    async fn load(&self, session_id: &str) -> Result<Option<AgentSnapshot>> {
141        let path = self.session_path(session_id);
142        if path.exists() {
143            let envelope: SnapshotEnvelope = serde_json::from_slice(&tokio::fs::read(path).await?)?;
144            if envelope.version != ENVELOPE_VERSION || envelope.session_id != session_id {
145                return Err(AgentError::Persistence(
146                    "File storage envelope does not match the requested session".into(),
147                ));
148            }
149            return Ok(Some(envelope.snapshot));
150        }
151
152        let encoded_path = self.encoded_session_path(session_id);
153        let fallback = if encoded_path.exists() {
154            Some(encoded_path)
155        } else {
156            self.legacy_session_path(session_id)
157                .filter(|legacy_path| legacy_path.exists())
158        };
159        let Some(path) = fallback else {
160            return Ok(None);
161        };
162        let snapshot = serde_json::from_slice(&tokio::fs::read(path).await?)?;
163        Ok(Some(snapshot))
164    }
165
166    async fn delete(&self, session_id: &str) -> Result<()> {
167        for path in [
168            self.session_path(session_id),
169            self.encoded_session_path(session_id),
170        ] {
171            if path.exists() {
172                tokio::fs::remove_file(path).await?;
173            }
174        }
175        if let Some(legacy_path) = self.legacy_session_path(session_id)
176            && legacy_path.exists()
177        {
178            tokio::fs::remove_file(legacy_path).await?;
179        }
180        Ok(())
181    }
182
183    async fn list_sessions(&self) -> Result<Vec<String>> {
184        let mut sessions = BTreeSet::new();
185        if !self.base_path.exists() {
186            return Ok(Vec::new());
187        }
188
189        let mut entries = tokio::fs::read_dir(&self.base_path).await?;
190        while let Some(entry) = entries.next_entry().await? {
191            let path = entry.path();
192            let Some(extension) = path.extension().and_then(|value| value.to_str()) else {
193                continue;
194            };
195            let Some(stem) = path.file_stem().and_then(|value| value.to_str()) else {
196                continue;
197            };
198
199            if extension == ENCODED_EXTENSION && stem.starts_with(HASHED_PREFIX) {
200                let envelope: SnapshotEnvelope =
201                    serde_json::from_slice(&tokio::fs::read(&path).await?)?;
202                if envelope.version != ENVELOPE_VERSION
203                    || hashed_session_id(&envelope.session_id) != stem
204                {
205                    return Err(AgentError::Persistence(
206                        "File storage envelope does not match its filename".into(),
207                    ));
208                }
209                sessions.insert(envelope.session_id);
210            } else if extension == ENCODED_EXTENSION {
211                if let Some(session_id) = decode_session_id(stem) {
212                    sessions.insert(session_id);
213                }
214            } else if extension == "json" && is_safe_legacy_session_id(stem) {
215                sessions.insert(stem.to_string());
216            }
217        }
218        Ok(sessions.into_iter().collect())
219    }
220}
221
222#[cfg(test)]
223mod tests {
224    use super::*;
225    use tempfile::TempDir;
226
227    #[test]
228    fn reports_snapshot_capability() {
229        let storage = FileStorage::new("unused");
230
231        assert!(storage.supports(StorageCapability::Snapshot));
232        assert!(!storage.supports(StorageCapability::SessionMetadata));
233    }
234
235    #[tokio::test]
236    async fn test_save_and_load() {
237        let temp_dir = TempDir::new().unwrap();
238        let storage = FileStorage::new(temp_dir.path());
239
240        let snapshot = AgentSnapshot::new("test-agent".into());
241        storage.save("session-1", &snapshot).await.unwrap();
242
243        let loaded = storage.load("session-1").await.unwrap();
244        assert!(loaded.is_some());
245        assert_eq!(loaded.unwrap().agent_id, "test-agent");
246    }
247
248    #[tokio::test]
249    async fn test_load_nonexistent() {
250        let temp_dir = TempDir::new().unwrap();
251        let storage = FileStorage::new(temp_dir.path());
252
253        let loaded = storage.load("nonexistent").await.unwrap();
254        assert!(loaded.is_none());
255    }
256
257    #[tokio::test]
258    async fn test_delete() {
259        let temp_dir = TempDir::new().unwrap();
260        let storage = FileStorage::new(temp_dir.path());
261
262        let snapshot = AgentSnapshot::new("test-agent".into());
263        storage.save("session-1", &snapshot).await.unwrap();
264        assert!(storage.load("session-1").await.unwrap().is_some());
265
266        storage.delete("session-1").await.unwrap();
267        assert!(storage.load("session-1").await.unwrap().is_none());
268    }
269
270    #[tokio::test]
271    async fn test_list_sessions() {
272        let temp_dir = TempDir::new().unwrap();
273        let storage = FileStorage::new(temp_dir.path());
274
275        storage
276            .save("session-1", &AgentSnapshot::new("agent".into()))
277            .await
278            .unwrap();
279        storage
280            .save("session-2", &AgentSnapshot::new("agent".into()))
281            .await
282            .unwrap();
283
284        let sessions = storage.list_sessions().await.unwrap();
285        assert_eq!(
286            sessions,
287            vec!["session-1".to_string(), "session-2".to_string()]
288        );
289    }
290
291    #[tokio::test]
292    async fn arbitrary_session_ids_use_flat_round_trip_paths() {
293        let temp_dir = TempDir::new().unwrap();
294        let storage = FileStorage::new(temp_dir.path());
295        let session_ids = [
296            "../escape",
297            "nested/session",
298            "nested\\session",
299            ".",
300            "..",
301            "unicode-雪",
302            "",
303        ];
304
305        for session_id in session_ids {
306            let path = storage.session_path(session_id);
307            assert_eq!(path.parent(), Some(temp_dir.path()));
308            assert_eq!(
309                path.extension().and_then(|value| value.to_str()),
310                Some("snapshot")
311            );
312            storage
313                .save(session_id, &AgentSnapshot::new(session_id.to_string()))
314                .await
315                .unwrap();
316        }
317
318        let sessions = storage.list_sessions().await.unwrap();
319        assert_eq!(sessions.len(), session_ids.len());
320        for session_id in session_ids {
321            assert!(sessions.contains(&session_id.to_string()));
322            assert_eq!(
323                storage.load(session_id).await.unwrap().unwrap().agent_id,
324                session_id
325            );
326        }
327
328        let mut entries = tokio::fs::read_dir(temp_dir.path()).await.unwrap();
329        while let Some(entry) = entries.next_entry().await.unwrap() {
330            assert!(entry.file_type().await.unwrap().is_file());
331        }
332    }
333
334    #[tokio::test]
335    async fn fixed_length_filename_supports_long_session_ids() {
336        let temp_dir = TempDir::new().unwrap();
337        let storage = FileStorage::new(temp_dir.path());
338        let session_id = "segment/".repeat(1024);
339        let snapshot = AgentSnapshot::new("agent".into());
340
341        storage.save(&session_id, &snapshot).await.unwrap();
342
343        let path = storage.session_path(&session_id);
344        assert!(path.file_name().unwrap().len() < 100);
345        assert_eq!(
346            storage.list_sessions().await.unwrap(),
347            vec![session_id.clone()]
348        );
349        assert_eq!(
350            storage.load(&session_id).await.unwrap().unwrap().agent_id,
351            "agent"
352        );
353    }
354
355    #[tokio::test]
356    async fn envelope_mismatch_fails_closed() {
357        let temp_dir = TempDir::new().unwrap();
358        let storage = FileStorage::new(temp_dir.path());
359        let requested = "requested";
360        let envelope = SnapshotEnvelope {
361            version: ENVELOPE_VERSION,
362            session_id: "different".into(),
363            snapshot: AgentSnapshot::new("agent".into()),
364        };
365        tokio::fs::write(
366            storage.session_path(requested),
367            serde_json::to_vec(&envelope).unwrap(),
368        )
369        .await
370        .unwrap();
371
372        assert!(matches!(
373            storage.load(requested).await,
374            Err(AgentError::Persistence(message)) if message.contains("does not match")
375        ));
376    }
377
378    #[tokio::test]
379    async fn reversible_v1_snapshot_is_read_and_migrated_on_save() {
380        let temp_dir = TempDir::new().unwrap();
381        let storage = FileStorage::new(temp_dir.path());
382        let session_id = "legacy/encoded";
383        let old_path = storage.encoded_session_path(session_id);
384        tokio::fs::write(
385            &old_path,
386            serde_json::to_vec(&AgentSnapshot::new("old".into())).unwrap(),
387        )
388        .await
389        .unwrap();
390
391        assert_eq!(
392            storage.load(session_id).await.unwrap().unwrap().agent_id,
393            "old"
394        );
395        storage
396            .save(session_id, &AgentSnapshot::new("new".into()))
397            .await
398            .unwrap();
399        assert!(!old_path.exists());
400        assert_eq!(
401            storage.load(session_id).await.unwrap().unwrap().agent_id,
402            "new"
403        );
404    }
405
406    #[tokio::test]
407    async fn legacy_fallback_accepts_only_safe_flat_ids() {
408        let temp_dir = TempDir::new().unwrap();
409        let storage_path = temp_dir.path().join("storage");
410        tokio::fs::create_dir_all(&storage_path).await.unwrap();
411        let storage = FileStorage::new(&storage_path);
412        let snapshot = AgentSnapshot::new("legacy-agent".into());
413        tokio::fs::write(
414            storage_path.join("legacy.json"),
415            serde_json::to_string(&snapshot).unwrap(),
416        )
417        .await
418        .unwrap();
419        tokio::fs::write(
420            temp_dir.path().join("escape.json"),
421            serde_json::to_string(&snapshot).unwrap(),
422        )
423        .await
424        .unwrap();
425
426        assert_eq!(
427            storage.load("legacy").await.unwrap().unwrap().agent_id,
428            "legacy-agent"
429        );
430        assert_eq!(
431            storage.list_sessions().await.unwrap(),
432            vec!["legacy".to_string()]
433        );
434        assert!(storage.load("../escape").await.unwrap().is_none());
435
436        storage.delete("legacy").await.unwrap();
437        assert!(!storage_path.join("legacy.json").exists());
438        assert!(temp_dir.path().join("escape.json").exists());
439    }
440}