1use std::path::PathBuf;
7
8use async_trait::async_trait;
9use tokio::fs;
10
11use super::{Artifact, ArtifactError, ArtifactMetadata, ArtifactService, now_secs};
12
13pub struct FileArtifactService {
19 root_dir: PathBuf,
20}
21
22impl FileArtifactService {
23 pub fn new(root_dir: impl Into<PathBuf>) -> Result<Self, ArtifactError> {
27 let root = root_dir.into();
28 std::fs::create_dir_all(&root)
30 .map_err(|e| ArtifactError::Storage(format!("Failed to create root dir: {e}")))?;
31 Ok(Self { root_dir: root })
32 }
33
34 fn artifact_dir(&self, session_id: &str, name: &str) -> PathBuf {
35 let safe_session = sanitize_path_component(session_id);
37 let safe_name = sanitize_path_component(name);
38 self.root_dir.join(&safe_session).join(&safe_name)
39 }
40
41 fn version_dir(&self, session_id: &str, name: &str, version: u32) -> PathBuf {
42 self.artifact_dir(session_id, name)
43 .join(format!("v{version}"))
44 }
45
46 async fn next_version(&self, session_id: &str, name: &str) -> u32 {
48 let dir = self.artifact_dir(session_id, name);
49 if !dir.exists() {
50 return 1;
51 }
52 let mut max_version = 0u32;
53 if let Ok(mut entries) = fs::read_dir(&dir).await {
54 while let Ok(Some(entry)) = entries.next_entry().await {
55 if let Some(name) = entry.file_name().to_str()
56 && let Some(v) = name.strip_prefix('v')
57 && let Ok(version) = v.parse::<u32>()
58 {
59 max_version = max_version.max(version);
60 }
61 }
62 }
63 max_version + 1
64 }
65}
66
67fn sanitize_path_component(s: &str) -> String {
69 s.replace(['/', '\\', '.'], "_")
70}
71
72#[async_trait]
73impl ArtifactService for FileArtifactService {
74 async fn save(
75 &self,
76 session_id: &str,
77 artifact: Artifact,
78 ) -> Result<ArtifactMetadata, ArtifactError> {
79 let version = self.next_version(session_id, &artifact.metadata.name).await;
80 let ver_dir = self.version_dir(session_id, &artifact.metadata.name, version);
81 fs::create_dir_all(&ver_dir)
82 .await
83 .map_err(|e| ArtifactError::Storage(e.to_string()))?;
84
85 fs::write(ver_dir.join("data"), &artifact.data)
87 .await
88 .map_err(|e| ArtifactError::Storage(e.to_string()))?;
89
90 let mut metadata = artifact.metadata;
92 metadata.version = version;
93 metadata.updated_at = now_secs();
94 if version == 1 {
95 metadata.created_at = metadata.updated_at;
96 }
97
98 let metadata_json = serde_json::to_string_pretty(&metadata)
99 .map_err(|e| ArtifactError::Storage(e.to_string()))?;
100 fs::write(ver_dir.join("metadata.json"), metadata_json)
101 .await
102 .map_err(|e| ArtifactError::Storage(e.to_string()))?;
103
104 Ok(metadata)
105 }
106
107 async fn load(&self, session_id: &str, name: &str) -> Result<Option<Artifact>, ArtifactError> {
108 let latest = self.next_version(session_id, name).await;
109 if latest == 1 {
110 return Ok(None);
111 }
112 self.load_version(session_id, name, latest - 1).await
113 }
114
115 async fn load_version(
116 &self,
117 session_id: &str,
118 name: &str,
119 version: u32,
120 ) -> Result<Option<Artifact>, ArtifactError> {
121 let ver_dir = self.version_dir(session_id, name, version);
122 if !ver_dir.exists() {
123 return Ok(None);
124 }
125
126 let data = fs::read(ver_dir.join("data"))
127 .await
128 .map_err(|e| ArtifactError::Storage(e.to_string()))?;
129 let metadata_str = fs::read_to_string(ver_dir.join("metadata.json"))
130 .await
131 .map_err(|e| ArtifactError::Storage(e.to_string()))?;
132 let metadata: ArtifactMetadata = serde_json::from_str(&metadata_str)
133 .map_err(|e| ArtifactError::Storage(e.to_string()))?;
134
135 Ok(Some(Artifact { metadata, data }))
136 }
137
138 async fn list(&self, session_id: &str) -> Result<Vec<ArtifactMetadata>, ArtifactError> {
139 let session_dir = self.root_dir.join(sanitize_path_component(session_id));
140 if !session_dir.exists() {
141 return Ok(vec![]);
142 }
143
144 let mut result = vec![];
145 let mut entries = fs::read_dir(&session_dir)
146 .await
147 .map_err(|e| ArtifactError::Storage(e.to_string()))?;
148
149 while let Some(entry) = entries
150 .next_entry()
151 .await
152 .map_err(|e| ArtifactError::Storage(e.to_string()))?
153 {
154 if entry.file_type().await.map(|t| t.is_dir()).unwrap_or(false) {
155 let name = entry.file_name().to_string_lossy().to_string();
156 if let Ok(Some(artifact)) = self.load(session_id, &name).await {
158 result.push(artifact.metadata);
159 }
160 }
161 }
162 Ok(result)
163 }
164
165 async fn delete(&self, session_id: &str, name: &str) -> Result<(), ArtifactError> {
166 let dir = self.artifact_dir(session_id, name);
167 if dir.exists() {
168 fs::remove_dir_all(&dir)
169 .await
170 .map_err(|e| ArtifactError::Storage(e.to_string()))?;
171 }
172 Ok(())
173 }
174}
175
176#[cfg(test)]
177mod tests {
178 use super::*;
179 use std::sync::atomic::{AtomicU32, Ordering};
180
181 fn test_dir() -> PathBuf {
183 static COUNTER: AtomicU32 = AtomicU32::new(0);
184 let id = COUNTER.fetch_add(1, Ordering::SeqCst);
185 let dir = std::env::temp_dir()
186 .join("gemini_adk_rs_file_artifact_tests")
187 .join(format!("test_{}_{}", std::process::id(), id));
188 let _ = std::fs::remove_dir_all(&dir);
190 dir
191 }
192
193 #[tokio::test]
194 async fn save_and_load_round_trip() {
195 let dir = test_dir();
196 let svc = FileArtifactService::new(&dir).unwrap();
197
198 let artifact = Artifact::text("notes", "Hello, world!");
199 let meta = svc.save("session1", artifact).await.unwrap();
200 assert_eq!(meta.name, "notes");
201 assert_eq!(meta.version, 1);
202
203 let loaded = svc.load("session1", "notes").await.unwrap().unwrap();
204 assert_eq!(std::str::from_utf8(&loaded.data).unwrap(), "Hello, world!");
205 assert_eq!(loaded.metadata.version, 1);
206 assert_eq!(loaded.metadata.mime_type, "text/plain");
207
208 std::fs::remove_dir_all(&dir).ok();
209 }
210
211 #[tokio::test]
212 async fn versioning_increments_and_load_gets_latest() {
213 let dir = test_dir();
214 let svc = FileArtifactService::new(&dir).unwrap();
215
216 let m1 = svc
217 .save("s1", Artifact::text("doc", "version 1"))
218 .await
219 .unwrap();
220 assert_eq!(m1.version, 1);
221
222 let m2 = svc
223 .save("s1", Artifact::text("doc", "version 2"))
224 .await
225 .unwrap();
226 assert_eq!(m2.version, 2);
227
228 let m3 = svc
229 .save("s1", Artifact::text("doc", "version 3"))
230 .await
231 .unwrap();
232 assert_eq!(m3.version, 3);
233
234 let latest = svc.load("s1", "doc").await.unwrap().unwrap();
236 assert_eq!(latest.metadata.version, 3);
237 assert_eq!(std::str::from_utf8(&latest.data).unwrap(), "version 3");
238
239 std::fs::remove_dir_all(&dir).ok();
240 }
241
242 #[tokio::test]
243 async fn load_specific_version() {
244 let dir = test_dir();
245 let svc = FileArtifactService::new(&dir).unwrap();
246
247 svc.save("s1", Artifact::text("doc", "v1 data"))
248 .await
249 .unwrap();
250 svc.save("s1", Artifact::text("doc", "v2 data"))
251 .await
252 .unwrap();
253 svc.save("s1", Artifact::text("doc", "v3 data"))
254 .await
255 .unwrap();
256
257 let v1 = svc.load_version("s1", "doc", 1).await.unwrap().unwrap();
258 assert_eq!(std::str::from_utf8(&v1.data).unwrap(), "v1 data");
259 assert_eq!(v1.metadata.version, 1);
260
261 let v2 = svc.load_version("s1", "doc", 2).await.unwrap().unwrap();
262 assert_eq!(std::str::from_utf8(&v2.data).unwrap(), "v2 data");
263
264 let v3 = svc.load_version("s1", "doc", 3).await.unwrap().unwrap();
265 assert_eq!(std::str::from_utf8(&v3.data).unwrap(), "v3 data");
266
267 let v99 = svc.load_version("s1", "doc", 99).await.unwrap();
269 assert!(v99.is_none());
270
271 std::fs::remove_dir_all(&dir).ok();
272 }
273
274 #[tokio::test]
275 async fn list_artifacts() {
276 let dir = test_dir();
277 let svc = FileArtifactService::new(&dir).unwrap();
278
279 svc.save("s1", Artifact::text("alpha", "data"))
280 .await
281 .unwrap();
282 svc.save("s1", Artifact::text("beta", "data"))
283 .await
284 .unwrap();
285 svc.save("s2", Artifact::text("gamma", "data"))
286 .await
287 .unwrap();
288
289 let list = svc.list("s1").await.unwrap();
290 assert_eq!(list.len(), 2);
291 let names: Vec<&str> = list.iter().map(|m| m.name.as_str()).collect();
292 assert!(names.contains(&"alpha"));
293 assert!(names.contains(&"beta"));
294
295 let list2 = svc.list("s2").await.unwrap();
297 assert_eq!(list2.len(), 1);
298 assert_eq!(list2[0].name, "gamma");
299
300 std::fs::remove_dir_all(&dir).ok();
301 }
302
303 #[tokio::test]
304 async fn delete_artifact() {
305 let dir = test_dir();
306 let svc = FileArtifactService::new(&dir).unwrap();
307
308 svc.save("s1", Artifact::text("notes", "data"))
309 .await
310 .unwrap();
311 svc.save("s1", Artifact::text("notes", "v2")).await.unwrap();
312
313 svc.delete("s1", "notes").await.unwrap();
314
315 let result = svc.load("s1", "notes").await.unwrap();
316 assert!(result.is_none());
317
318 svc.delete("s1", "notes").await.unwrap();
320
321 std::fs::remove_dir_all(&dir).ok();
322 }
323
324 #[tokio::test]
325 async fn load_nonexistent_returns_none() {
326 let dir = test_dir();
327 let svc = FileArtifactService::new(&dir).unwrap();
328
329 let result = svc.load("no_session", "no_artifact").await.unwrap();
330 assert!(result.is_none());
331
332 std::fs::remove_dir_all(&dir).ok();
333 }
334
335 #[tokio::test]
336 async fn path_traversal_prevention() {
337 let dir = test_dir();
338 let svc = FileArtifactService::new(&dir).unwrap();
339
340 let artifact = Artifact::text("../../../etc/passwd", "malicious");
342 let meta = svc.save("../../hack", artifact).await.unwrap();
343 assert_eq!(meta.version, 1);
344
345 let sanitized_session = sanitize_path_component("../../hack");
347 let sanitized_name = sanitize_path_component("../../../etc/passwd");
348 assert!(!sanitized_session.contains('/'));
349 assert!(!sanitized_session.contains('\\'));
350 assert!(!sanitized_session.contains('.'));
351 assert!(!sanitized_name.contains('/'));
352 assert!(!sanitized_name.contains('\\'));
353 assert!(!sanitized_name.contains('.'));
354
355 let loaded = svc.load("../../hack", "../../../etc/passwd").await.unwrap();
357 assert!(loaded.is_some());
358 assert_eq!(
359 std::str::from_utf8(&loaded.unwrap().data).unwrap(),
360 "malicious"
361 );
362
363 assert!(dir.exists());
365
366 std::fs::remove_dir_all(&dir).ok();
367 }
368
369 #[test]
370 fn sanitize_removes_dangerous_chars() {
371 assert_eq!(sanitize_path_component("normal"), "normal");
372 assert_eq!(sanitize_path_component(".."), "__");
373 assert_eq!(sanitize_path_component("a/b"), "a_b");
374 assert_eq!(sanitize_path_component("a\\b"), "a_b");
375 assert_eq!(sanitize_path_component("../../etc"), "______etc");
376 assert_eq!(sanitize_path_component("file.txt"), "file_txt");
377 }
378}