1use std::collections::HashMap;
26use std::ops::Range;
27use std::sync::{Arc, RwLock};
28
29use async_trait::async_trait;
30use bytes::Bytes;
31use futures::StreamExt;
32use futures::stream::BoxStream;
33use serde::{Deserialize, Serialize};
34
35use crate::error::invalid_data;
36use crate::io::{
37 FileMetadata, FileRead, FileWrite, InputFile, OutputFile, Storage, StorageConfig,
38 StorageFactory,
39};
40use crate::{Error, ErrorKind, Result};
41
42#[derive(Debug, Clone, Default, Serialize, Deserialize)]
63pub struct MemoryStorage {
64 #[serde(skip, default = "default_memory_data")]
65 data: Arc<RwLock<HashMap<String, Bytes>>>,
66}
67
68fn default_memory_data() -> Arc<RwLock<HashMap<String, Bytes>>> {
69 Arc::new(RwLock::new(HashMap::new()))
70}
71
72impl MemoryStorage {
73 pub fn new() -> Self {
75 Self {
76 data: Arc::new(RwLock::new(HashMap::new())),
77 }
78 }
79
80 pub(crate) fn normalize_path(path: &str) -> String {
88 let path = path.strip_prefix("memory://").unwrap_or(path);
90 let path = path.strip_prefix("memory:/").unwrap_or(path);
92 path.trim_start_matches('/').to_string()
94 }
95}
96
97#[async_trait]
98#[typetag::serde]
99impl Storage for MemoryStorage {
100 async fn exists(&self, path: &str) -> Result<bool> {
101 let normalized = Self::normalize_path(path);
102 let data = self.data.read().map_err(|e| {
103 Error::new(
104 ErrorKind::Unexpected,
105 format!("Failed to acquire read lock: {e}"),
106 )
107 })?;
108 Ok(data.contains_key(&normalized))
109 }
110
111 async fn metadata(&self, path: &str) -> Result<FileMetadata> {
112 let normalized = Self::normalize_path(path);
113 let data = self.data.read().map_err(|e| {
114 Error::new(
115 ErrorKind::Unexpected,
116 format!("Failed to acquire read lock: {e}"),
117 )
118 })?;
119 match data.get(&normalized) {
120 Some(bytes) => Ok(FileMetadata {
121 size: bytes.len() as u64,
122 }),
123 None => Err(invalid_data!("File not found: {path}")),
124 }
125 }
126
127 async fn read(&self, path: &str) -> Result<Bytes> {
128 let normalized = Self::normalize_path(path);
129 let data = self.data.read().map_err(|e| {
130 Error::new(
131 ErrorKind::Unexpected,
132 format!("Failed to acquire read lock: {e}"),
133 )
134 })?;
135 match data.get(&normalized) {
136 Some(bytes) => Ok(bytes.clone()),
137 None => Err(invalid_data!("File not found: {path}")),
138 }
139 }
140
141 async fn reader(&self, path: &str) -> Result<Box<dyn FileRead>> {
142 let normalized = Self::normalize_path(path);
143 let data = self.data.read().map_err(|e| {
144 Error::new(
145 ErrorKind::Unexpected,
146 format!("Failed to acquire read lock: {e}"),
147 )
148 })?;
149 match data.get(&normalized) {
150 Some(bytes) => Ok(Box::new(MemoryFileRead::new(bytes.clone()))),
151 None => Err(invalid_data!("File not found: {path}")),
152 }
153 }
154
155 async fn write(&self, path: &str, bs: Bytes) -> Result<()> {
156 let normalized = Self::normalize_path(path);
157 let mut data = self.data.write().map_err(|e| {
158 Error::new(
159 ErrorKind::Unexpected,
160 format!("Failed to acquire write lock: {e}"),
161 )
162 })?;
163 data.insert(normalized, bs);
164 Ok(())
165 }
166
167 async fn writer(&self, path: &str) -> Result<Box<dyn FileWrite>> {
168 let normalized = Self::normalize_path(path);
169 Ok(Box::new(MemoryFileWrite::new(
170 self.data.clone(),
171 normalized,
172 )))
173 }
174
175 async fn delete(&self, path: &str) -> Result<()> {
176 let normalized = Self::normalize_path(path);
177 let mut data = self.data.write().map_err(|e| {
178 Error::new(
179 ErrorKind::Unexpected,
180 format!("Failed to acquire write lock: {e}"),
181 )
182 })?;
183 data.remove(&normalized);
184 Ok(())
185 }
186
187 async fn delete_prefix(&self, path: &str) -> Result<()> {
188 let normalized = Self::normalize_path(path);
189 let prefix = if normalized.ends_with('/') {
190 normalized
191 } else {
192 format!("{normalized}/")
193 };
194
195 let mut data = self.data.write().map_err(|e| {
196 Error::new(
197 ErrorKind::Unexpected,
198 format!("Failed to acquire write lock: {e}"),
199 )
200 })?;
201
202 let keys_to_remove: Vec<String> = data
204 .keys()
205 .filter(|k| k.starts_with(&prefix))
206 .cloned()
207 .collect();
208
209 for key in keys_to_remove {
210 data.remove(&key);
211 }
212
213 Ok(())
214 }
215
216 async fn delete_stream(&self, mut paths: BoxStream<'static, String>) -> Result<()> {
217 while let Some(path) = paths.next().await {
218 self.delete(&path).await?;
219 }
220 Ok(())
221 }
222
223 fn new_input(&self, path: &str) -> Result<InputFile> {
224 Ok(InputFile::new(Arc::new(self.clone()), path.to_string()))
225 }
226
227 fn new_output(&self, path: &str) -> Result<OutputFile> {
228 Ok(OutputFile::new(Arc::new(self.clone()), path.to_string()))
229 }
230}
231
232#[derive(Clone, Debug, Default, Serialize, Deserialize)]
238pub struct MemoryStorageFactory;
239
240#[typetag::serde]
241impl StorageFactory for MemoryStorageFactory {
242 fn build(&self, _config: &StorageConfig) -> Result<Arc<dyn Storage>> {
243 Ok(Arc::new(MemoryStorage::new()))
244 }
245}
246
247#[derive(Debug)]
249pub struct MemoryFileRead {
250 data: Bytes,
251}
252
253impl MemoryFileRead {
254 pub fn new(data: Bytes) -> Self {
256 Self { data }
257 }
258}
259
260#[async_trait]
261impl FileRead for MemoryFileRead {
262 async fn read(&self, range: Range<u64>) -> Result<Bytes> {
263 let start = range.start as usize;
264 let end = range.end as usize;
265
266 if start > self.data.len() || end > self.data.len() {
267 return Err(invalid_data!(
268 "Range {}..{} is out of bounds for data of length {}",
269 start,
270 end,
271 self.data.len()
272 ));
273 }
274
275 Ok(self.data.slice(start..end))
276 }
277}
278
279#[derive(Debug)]
285pub struct MemoryFileWrite {
286 data: Arc<RwLock<HashMap<String, Bytes>>>,
287 path: String,
288 buffer: Vec<u8>,
289 closed: bool,
290}
291
292impl MemoryFileWrite {
293 pub fn new(data: Arc<RwLock<HashMap<String, Bytes>>>, path: String) -> Self {
295 Self {
296 data,
297 path,
298 buffer: Vec::new(),
299 closed: false,
300 }
301 }
302}
303
304#[async_trait]
305impl FileWrite for MemoryFileWrite {
306 async fn write(&mut self, bs: Bytes) -> Result<()> {
307 if self.closed {
308 return Err(invalid_data!("Cannot write to closed file"));
309 }
310 self.buffer.extend_from_slice(&bs);
311 Ok(())
312 }
313
314 async fn close(&mut self) -> Result<FileMetadata> {
315 if self.closed {
316 return Err(invalid_data!("File already closed"));
317 }
318
319 let mut data = self.data.write().map_err(|e| {
320 Error::new(
321 ErrorKind::Unexpected,
322 format!("Failed to acquire write lock: {e}"),
323 )
324 })?;
325
326 let size = self.buffer.len() as u64;
327 data.insert(
328 self.path.clone(),
329 Bytes::from(std::mem::take(&mut self.buffer)),
330 );
331 self.closed = true;
332 Ok(FileMetadata { size })
333 }
334}
335
336#[cfg(test)]
337mod tests {
338 use super::*;
339
340 #[test]
341 fn test_normalize_path() {
342 assert_eq!(
344 MemoryStorage::normalize_path("memory://path/to/file"),
345 "path/to/file"
346 );
347
348 assert_eq!(
350 MemoryStorage::normalize_path("memory:/path/to/file"),
351 "path/to/file"
352 );
353
354 assert_eq!(
356 MemoryStorage::normalize_path("/path/to/file"),
357 "path/to/file"
358 );
359
360 assert_eq!(
362 MemoryStorage::normalize_path("path/to/file"),
363 "path/to/file"
364 );
365
366 assert_eq!(
368 MemoryStorage::normalize_path("///path/to/file"),
369 "path/to/file"
370 );
371
372 assert_eq!(
374 MemoryStorage::normalize_path("memory:///path/to/file"),
375 "path/to/file"
376 );
377 }
378
379 #[tokio::test]
380 async fn test_memory_storage_write_read() {
381 let storage = MemoryStorage::new();
382 let path = "memory://test/file.txt";
383 let content = Bytes::from("Hello, World!");
384
385 storage.write(path, content.clone()).await.unwrap();
387
388 let read_content = storage.read(path).await.unwrap();
390 assert_eq!(read_content, content);
391 }
392
393 #[tokio::test]
394 async fn test_memory_storage_exists() {
395 let storage = MemoryStorage::new();
396 let path = "memory://test/file.txt";
397
398 assert!(!storage.exists(path).await.unwrap());
400
401 storage.write(path, Bytes::from("test")).await.unwrap();
403
404 assert!(storage.exists(path).await.unwrap());
406 }
407
408 #[tokio::test]
409 async fn test_memory_storage_metadata() {
410 let storage = MemoryStorage::new();
411 let path = "memory://test/file.txt";
412 let content = Bytes::from("Hello, World!");
413
414 storage.write(path, content.clone()).await.unwrap();
415
416 let metadata = storage.metadata(path).await.unwrap();
417 assert_eq!(metadata.size, content.len() as u64);
418 }
419
420 #[tokio::test]
421 async fn test_memory_storage_delete() {
422 let storage = MemoryStorage::new();
423 let path = "memory://test/file.txt";
424
425 storage.write(path, Bytes::from("test")).await.unwrap();
426 assert!(storage.exists(path).await.unwrap());
427
428 storage.delete(path).await.unwrap();
429 assert!(!storage.exists(path).await.unwrap());
430 }
431
432 #[tokio::test]
433 async fn test_memory_storage_delete_prefix() {
434 let storage = MemoryStorage::new();
435
436 storage
438 .write("memory://dir/file1.txt", Bytes::from("1"))
439 .await
440 .unwrap();
441 storage
442 .write("memory://dir/file2.txt", Bytes::from("2"))
443 .await
444 .unwrap();
445 storage
446 .write("memory://other/file.txt", Bytes::from("3"))
447 .await
448 .unwrap();
449
450 storage.delete_prefix("memory://dir").await.unwrap();
452
453 assert!(!storage.exists("memory://dir/file1.txt").await.unwrap());
455 assert!(!storage.exists("memory://dir/file2.txt").await.unwrap());
456
457 assert!(storage.exists("memory://other/file.txt").await.unwrap());
459 }
460
461 #[tokio::test]
462 async fn test_memory_storage_reader() {
463 let storage = MemoryStorage::new();
464 let path = "memory://test/file.txt";
465 let content = Bytes::from("Hello, World!");
466
467 storage.write(path, content.clone()).await.unwrap();
468
469 let reader = storage.reader(path).await.unwrap();
470 let read_content = reader.read(0..content.len() as u64).await.unwrap();
471 assert_eq!(read_content, content);
472
473 let partial = reader.read(0..5).await.unwrap();
475 assert_eq!(partial, Bytes::from("Hello"));
476 }
477
478 #[tokio::test]
479 async fn test_memory_storage_writer() {
480 let storage = MemoryStorage::new();
481 let path = "memory://test/file.txt";
482
483 let mut writer = storage.writer(path).await.unwrap();
484 writer.write(Bytes::from("Hello, ")).await.unwrap();
485 writer.write(Bytes::from("World!")).await.unwrap();
486 let metadata = writer.close().await.unwrap();
487
488 let content = storage.read(path).await.unwrap();
489 assert_eq!(content, Bytes::from("Hello, World!"));
490 assert_eq!(metadata.size, content.len() as u64);
491 }
492
493 #[tokio::test]
494 async fn test_memory_file_write_double_close() {
495 let storage = MemoryStorage::new();
496 let path = "memory://test/file.txt";
497
498 let mut writer = storage.writer(path).await.unwrap();
499 writer.write(Bytes::from("test")).await.unwrap();
500 writer.close().await.unwrap();
501
502 let result = writer.close().await;
504 assert!(result.is_err());
505 }
506
507 #[tokio::test]
508 async fn test_memory_file_write_after_close() {
509 let storage = MemoryStorage::new();
510 let path = "memory://test/file.txt";
511
512 let mut writer = storage.writer(path).await.unwrap();
513 assert_eq!(writer.close().await.unwrap().size, 0);
514
515 let result = writer.write(Bytes::from("test")).await;
517 assert!(result.is_err());
518 }
519
520 #[tokio::test]
521 async fn test_memory_file_read_out_of_bounds() {
522 let storage = MemoryStorage::new();
523 let path = "memory://test/file.txt";
524 let content = Bytes::from("Hello");
525
526 storage.write(path, content).await.unwrap();
527
528 let reader = storage.reader(path).await.unwrap();
529 let result = reader.read(0..100).await;
530 assert!(result.is_err());
531 }
532
533 #[test]
534 fn test_memory_storage_serialization() {
535 let storage = MemoryStorage::new();
536
537 let serialized = serde_json::to_string(&storage).unwrap();
539
540 let deserialized: MemoryStorage = serde_json::from_str(&serialized).unwrap();
542
543 assert!(deserialized.data.read().unwrap().is_empty());
545 }
546
547 #[test]
548 fn test_memory_storage_factory() {
549 let factory = MemoryStorageFactory;
550 let config = StorageConfig::new();
551 let storage = factory.build(&config).unwrap();
552
553 assert!(format!("{storage:?}").contains("MemoryStorage"));
555 }
556
557 #[test]
558 fn test_memory_storage_factory_serialization() {
559 let factory = MemoryStorageFactory;
560
561 let serialized = serde_json::to_string(&factory).unwrap();
563
564 let deserialized: MemoryStorageFactory = serde_json::from_str(&serialized).unwrap();
566
567 let config = StorageConfig::new();
569 let storage = deserialized.build(&config).unwrap();
570 assert!(format!("{storage:?}").contains("MemoryStorage"));
571 }
572
573 #[tokio::test]
574 async fn test_path_normalization_consistency() {
575 let storage = MemoryStorage::new();
576 let content = Bytes::from("test content");
577
578 storage
580 .write("memory://path/to/file", content.clone())
581 .await
582 .unwrap();
583
584 assert_eq!(
586 storage.read("memory://path/to/file").await.unwrap(),
587 content
588 );
589 assert_eq!(storage.read("memory:/path/to/file").await.unwrap(), content);
590 assert_eq!(storage.read("/path/to/file").await.unwrap(), content);
591 assert_eq!(storage.read("path/to/file").await.unwrap(), content);
592 }
593
594 #[tokio::test]
595 async fn test_memory_storage_delete_stream() {
596 use futures::stream;
597
598 let storage = MemoryStorage::new();
599
600 storage
602 .write("memory://file1.txt", Bytes::from("1"))
603 .await
604 .unwrap();
605 storage
606 .write("memory://file2.txt", Bytes::from("2"))
607 .await
608 .unwrap();
609 storage
610 .write("memory://file3.txt", Bytes::from("3"))
611 .await
612 .unwrap();
613
614 assert!(storage.exists("memory://file1.txt").await.unwrap());
616 assert!(storage.exists("memory://file2.txt").await.unwrap());
617 assert!(storage.exists("memory://file3.txt").await.unwrap());
618
619 let paths = vec![
621 "memory://file1.txt".to_string(),
622 "memory://file2.txt".to_string(),
623 ];
624 let path_stream = stream::iter(paths).boxed();
625 storage.delete_stream(path_stream).await.unwrap();
626
627 assert!(!storage.exists("memory://file1.txt").await.unwrap());
629 assert!(!storage.exists("memory://file2.txt").await.unwrap());
630
631 assert!(storage.exists("memory://file3.txt").await.unwrap());
633 }
634
635 #[tokio::test]
636 async fn test_memory_storage_delete_stream_empty() {
637 use futures::stream;
638
639 let storage = MemoryStorage::new();
640
641 let path_stream = stream::iter(Vec::<String>::new()).boxed();
643 storage.delete_stream(path_stream).await.unwrap();
644 }
645}