ai_agents_storage/
storage.rs1use 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}