1use async_trait::async_trait;
9use serde::de::{self, DeserializeSeed, SeqAccess, Visitor};
10use serde::{Deserialize, Deserializer, Serialize};
11use std::collections::HashMap;
12use std::fmt;
13use std::path::{Path, PathBuf};
14use std::sync::Arc;
15use std::sync::atomic::{AtomicU64, Ordering};
16use tokio::fs::{self, File, OpenOptions};
17use tokio::io::{AsyncReadExt, AsyncWriteExt};
18use tokio::sync::Mutex;
19
20use crate::backends::{StorageKey, StorageMeta, StorageValue};
21use crate::storage::EntryMap;
22use crate::{CacheEntry, CacheError, Result, StorageBackend};
23
24const SNAPSHOT_FILE_NAME: &str = "cache.json";
25const TEMP_FILE_PREFIX: &str = ".cache.json.tmp-";
26const SNAPSHOT_FORMAT_VERSION: u32 = 1;
27const MAX_SNAPSHOT_KEYS: usize = 100_000;
28const MAX_SNAPSHOT_ENTRIES: usize = 1_000_000;
29static TEMP_FILE_SEQUENCE: AtomicU64 = AtomicU64::new(0);
30
31pub const DEFAULT_MAX_SNAPSHOT_BYTES: u64 = 64 * 1024 * 1024;
33
34type PhantomTypes<K, V, M> = std::marker::PhantomData<(K, V, M)>;
36type SnapshotEntriesRef<'a, K, V, M> = Vec<(&'a K, &'a Vec<CacheEntry<K, V, M>>)>;
37
38struct BoundedBuffer {
39 bytes: Vec<u8>,
40 max_bytes: usize,
41 attempted_bytes: u64,
42 exceeded: bool,
43}
44
45impl BoundedBuffer {
46 fn new(max_bytes: usize) -> Self {
47 Self {
48 bytes: Vec::new(),
49 max_bytes,
50 attempted_bytes: 0,
51 exceeded: false,
52 }
53 }
54}
55
56impl std::io::Write for BoundedBuffer {
57 fn write(&mut self, buffer: &[u8]) -> std::io::Result<usize> {
58 let attempted = self.bytes.len().saturating_add(buffer.len());
59 self.attempted_bytes = u64::try_from(attempted).unwrap_or(u64::MAX);
60 if attempted > self.max_bytes {
61 self.exceeded = true;
62 return Err(std::io::Error::other("snapshot byte limit exceeded"));
63 }
64 self.bytes
65 .try_reserve(buffer.len())
66 .map_err(|error| std::io::Error::other(error.to_string()))?;
67 self.bytes.extend_from_slice(buffer);
68 Ok(buffer.len())
69 }
70
71 fn flush(&mut self) -> std::io::Result<()> {
72 Ok(())
73 }
74}
75
76#[derive(Serialize)]
77struct SnapshotRef<'a, K, V, M>
78where
79 K: Clone + std::hash::Hash + Eq,
80 V: Clone,
81 M: Clone,
82{
83 version: u32,
84 entries: SnapshotEntriesRef<'a, K, V, M>,
85}
86
87#[derive(Deserialize)]
88#[serde(bound(deserialize = "
89 K: Deserialize<'de> + Clone + std::hash::Hash + Eq,
90 V: Deserialize<'de> + Clone,
91 M: Deserialize<'de> + Clone
92"))]
93struct Snapshot<K, V, M>
94where
95 K: Clone + std::hash::Hash + Eq,
96 V: Clone,
97 M: Clone,
98{
99 version: u32,
100 #[serde(deserialize_with = "deserialize_snapshot_entries")]
101 entries: EntryMap<K, V, M>,
102}
103
104struct RejectAdditionalElement(&'static str);
105
106impl<'de> DeserializeSeed<'de> for RejectAdditionalElement {
107 type Value = ();
108
109 fn deserialize<D>(self, _deserializer: D) -> std::result::Result<Self::Value, D::Error>
110 where
111 D: Deserializer<'de>,
112 {
113 Err(de::Error::custom(self.0))
114 }
115}
116
117struct HistorySeed<K, V, M> {
118 max_entries: usize,
119 marker: PhantomTypes<K, V, M>,
120}
121
122impl<K, V, M> HistorySeed<K, V, M> {
123 fn new(max_entries: usize) -> Self {
124 Self {
125 max_entries,
126 marker: std::marker::PhantomData,
127 }
128 }
129}
130
131impl<'de, K, V, M> DeserializeSeed<'de> for HistorySeed<K, V, M>
132where
133 K: Deserialize<'de> + Clone + std::hash::Hash + Eq,
134 V: Deserialize<'de> + Clone,
135 M: Deserialize<'de> + Clone,
136{
137 type Value = Vec<CacheEntry<K, V, M>>;
138
139 fn deserialize<D>(self, deserializer: D) -> std::result::Result<Self::Value, D::Error>
140 where
141 D: Deserializer<'de>,
142 {
143 deserializer.deserialize_seq(HistoryVisitor {
144 max_entries: self.max_entries,
145 marker: self.marker,
146 })
147 }
148}
149
150struct HistoryVisitor<K, V, M> {
151 max_entries: usize,
152 marker: PhantomTypes<K, V, M>,
153}
154
155impl<'de, K, V, M> Visitor<'de> for HistoryVisitor<K, V, M>
156where
157 K: Deserialize<'de> + Clone + std::hash::Hash + Eq,
158 V: Deserialize<'de> + Clone,
159 M: Deserialize<'de> + Clone,
160{
161 type Value = Vec<CacheEntry<K, V, M>>;
162
163 fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
164 formatter.write_str("a cache entry history within the configured snapshot limit")
165 }
166
167 fn visit_seq<A>(self, mut sequence: A) -> std::result::Result<Self::Value, A::Error>
168 where
169 A: SeqAccess<'de>,
170 {
171 let capacity = sequence.size_hint().unwrap_or(0).min(self.max_entries);
172 let mut entries = Vec::new();
173 entries.try_reserve(capacity).map_err(de::Error::custom)?;
174
175 loop {
176 if entries.len() == self.max_entries {
177 let additional = sequence.next_element_seed(RejectAdditionalElement(
178 "snapshot contains more than the maximum number of entries",
179 ))?;
180 debug_assert!(additional.is_none());
181 break;
182 }
183 match sequence.next_element()? {
184 Some(entry) => entries.push(entry),
185 None => break,
186 }
187 }
188
189 Ok(entries)
190 }
191}
192
193struct SnapshotPairSeed<K, V, M> {
194 remaining_entries: usize,
195 marker: PhantomTypes<K, V, M>,
196}
197
198impl<K, V, M> SnapshotPairSeed<K, V, M> {
199 fn new(remaining_entries: usize) -> Self {
200 Self {
201 remaining_entries,
202 marker: std::marker::PhantomData,
203 }
204 }
205}
206
207impl<'de, K, V, M> DeserializeSeed<'de> for SnapshotPairSeed<K, V, M>
208where
209 K: Deserialize<'de> + Clone + std::hash::Hash + Eq,
210 V: Deserialize<'de> + Clone,
211 M: Deserialize<'de> + Clone,
212{
213 type Value = (K, Vec<CacheEntry<K, V, M>>);
214
215 fn deserialize<D>(self, deserializer: D) -> std::result::Result<Self::Value, D::Error>
216 where
217 D: Deserializer<'de>,
218 {
219 deserializer.deserialize_tuple(
220 2,
221 SnapshotPairVisitor {
222 remaining_entries: self.remaining_entries,
223 marker: self.marker,
224 },
225 )
226 }
227}
228
229struct SnapshotPairVisitor<K, V, M> {
230 remaining_entries: usize,
231 marker: PhantomTypes<K, V, M>,
232}
233
234impl<'de, K, V, M> Visitor<'de> for SnapshotPairVisitor<K, V, M>
235where
236 K: Deserialize<'de> + Clone + std::hash::Hash + Eq,
237 V: Deserialize<'de> + Clone,
238 M: Deserialize<'de> + Clone,
239{
240 type Value = (K, Vec<CacheEntry<K, V, M>>);
241
242 fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
243 formatter.write_str("a two-element [key, entry_history] pair")
244 }
245
246 fn visit_seq<A>(self, mut sequence: A) -> std::result::Result<Self::Value, A::Error>
247 where
248 A: SeqAccess<'de>,
249 {
250 let key = sequence
251 .next_element()?
252 .ok_or_else(|| de::Error::custom("snapshot entry pair is missing its key"))?;
253 let history = sequence
254 .next_element_seed(HistorySeed::new(self.remaining_entries))?
255 .ok_or_else(|| de::Error::custom("snapshot entry pair is missing its history"))?;
256 let additional = sequence.next_element_seed(RejectAdditionalElement(
257 "snapshot entry pair must contain exactly two elements",
258 ))?;
259 debug_assert!(additional.is_none());
260 Ok((key, history))
261 }
262}
263
264struct SnapshotEntriesVisitor<K, V, M> {
265 max_keys: usize,
266 max_entries: usize,
267 marker: PhantomTypes<K, V, M>,
268}
269
270impl<'de, K, V, M> Visitor<'de> for SnapshotEntriesVisitor<K, V, M>
271where
272 K: Deserialize<'de> + Clone + std::hash::Hash + Eq,
273 V: Deserialize<'de> + Clone,
274 M: Deserialize<'de> + Clone,
275{
276 type Value = EntryMap<K, V, M>;
277
278 fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
279 formatter.write_str("a bounded array of cache entry pairs")
280 }
281
282 fn visit_seq<A>(self, mut sequence: A) -> std::result::Result<Self::Value, A::Error>
283 where
284 A: SeqAccess<'de>,
285 {
286 let capacity = sequence.size_hint().unwrap_or(0).min(self.max_keys);
287 let mut entries = HashMap::new();
288 entries.try_reserve(capacity).map_err(de::Error::custom)?;
289 let mut total_entries = 0usize;
290
291 loop {
292 if entries.len() == self.max_keys {
293 let additional = sequence.next_element_seed(RejectAdditionalElement(
294 "snapshot contains more than the maximum number of keys",
295 ))?;
296 debug_assert!(additional.is_none());
297 break;
298 }
299
300 let remaining_entries = self.max_entries - total_entries;
301 let Some((key, key_entries)) =
302 sequence.next_element_seed(SnapshotPairSeed::new(remaining_entries))?
303 else {
304 break;
305 };
306 total_entries = total_entries
307 .checked_add(key_entries.len())
308 .ok_or_else(|| de::Error::custom("snapshot entry count overflowed usize"))?;
309 if key_entries.iter().any(|entry| entry.key != key) {
310 return Err(de::Error::custom(
311 "snapshot contains an entry whose embedded key does not match",
312 ));
313 }
314 if entries.insert(key, key_entries).is_some() {
315 return Err(de::Error::custom("snapshot contains a duplicate key"));
316 }
317 }
318
319 Ok(entries)
320 }
321}
322
323fn deserialize_snapshot_entries<'de, D, K, V, M>(
324 deserializer: D,
325) -> std::result::Result<EntryMap<K, V, M>, D::Error>
326where
327 D: Deserializer<'de>,
328 K: Deserialize<'de> + Clone + std::hash::Hash + Eq,
329 V: Deserialize<'de> + Clone,
330 M: Deserialize<'de> + Clone,
331{
332 deserializer.deserialize_seq(SnapshotEntriesVisitor {
333 max_keys: MAX_SNAPSHOT_KEYS,
334 max_entries: MAX_SNAPSHOT_ENTRIES,
335 marker: std::marker::PhantomData,
336 })
337}
338
339#[allow(clippy::type_complexity)]
344pub struct FilesystemBackend<K, V, M = ()>
345where
346 K: StorageKey,
347 V: StorageValue,
348 M: StorageMeta,
349{
350 base_path: PathBuf,
351 max_snapshot_bytes: u64,
352 io_lock: Arc<Mutex<()>>,
353 _phantom: PhantomTypes<K, V, M>,
354}
355
356impl<K, V, M> FilesystemBackend<K, V, M>
357where
358 K: StorageKey,
359 V: StorageValue,
360 M: StorageMeta,
361{
362 pub async fn new<P: AsRef<Path>>(base_path: P) -> Result<Self> {
364 let base_path = base_path.as_ref().to_path_buf();
365 fs::create_dir_all(&base_path).await?;
366 let metadata = fs::symlink_metadata(&base_path).await?;
367 if metadata.file_type().is_symlink() || !metadata.is_dir() {
368 return Err(CacheError::InvalidConfiguration(
369 "filesystem cache path must be a real directory, not a symlink".to_string(),
370 ));
371 }
372
373 Ok(Self {
374 base_path,
375 max_snapshot_bytes: DEFAULT_MAX_SNAPSHOT_BYTES,
376 io_lock: Arc::new(Mutex::new(())),
377 _phantom: std::marker::PhantomData,
378 })
379 }
380
381 pub fn with_max_snapshot_bytes(mut self, max_snapshot_bytes: u64) -> Self {
385 self.max_snapshot_bytes = max_snapshot_bytes;
386 self
387 }
388
389 fn snapshot_path(&self) -> PathBuf {
390 self.base_path.join(SNAPSHOT_FILE_NAME)
391 }
392
393 fn temporary_path(&self) -> PathBuf {
394 let sequence = TEMP_FILE_SEQUENCE.fetch_add(1, Ordering::Relaxed);
395 self.base_path.join(format!(
396 "{TEMP_FILE_PREFIX}{}-{sequence}",
397 std::process::id()
398 ))
399 }
400
401 async fn has_legacy_layout(&self) -> Result<bool> {
402 let mut directory = fs::read_dir(&self.base_path).await?;
403 while let Some(entry) = directory.next_entry().await? {
404 let file_type = entry.file_type().await?;
405 let path = entry.path();
406 let is_legacy_extension = matches!(
407 path.extension().and_then(|extension| extension.to_str()),
408 Some("json" | "bin")
409 );
410 let is_snapshot =
411 path.file_name().and_then(|name| name.to_str()) == Some(SNAPSHOT_FILE_NAME);
412 if file_type.is_symlink() && is_legacy_extension {
413 return Err(CacheError::StorageBackend(
414 "cache directory contains a symlink with a persistence-file extension"
415 .to_string(),
416 ));
417 }
418 if file_type.is_file() && is_legacy_extension && !is_snapshot {
419 return Ok(true);
420 }
421 }
422 Ok(false)
423 }
424
425 async fn reject_legacy_layout(&self) -> Result<()> {
426 if self.has_legacy_layout().await? {
427 Err(CacheError::UnsupportedPersistenceFormat(
428 "legacy per-key cache files were found; clear or migrate the cache directory"
429 .to_string(),
430 ))
431 } else {
432 Ok(())
433 }
434 }
435
436 async fn read_snapshot_bytes(&self) -> Result<Option<Vec<u8>>> {
437 self.reject_legacy_layout().await?;
438 let path = self.snapshot_path();
439 let metadata = match fs::symlink_metadata(&path).await {
440 Ok(metadata) => metadata,
441 Err(error) if error.kind() == std::io::ErrorKind::NotFound => {
442 return Ok(None);
443 }
444 Err(error) => return Err(error.into()),
445 };
446 if metadata.file_type().is_symlink() || !metadata.is_file() {
447 return Err(CacheError::StorageBackend(
448 "cache snapshot must be a regular file".to_string(),
449 ));
450 }
451 if metadata.len() > self.max_snapshot_bytes {
452 return Err(CacheError::SnapshotTooLarge {
453 actual_bytes: metadata.len(),
454 max_bytes: self.max_snapshot_bytes,
455 });
456 }
457
458 let mut bytes = Vec::new();
459 let read_limit = self.max_snapshot_bytes.saturating_add(1);
460 File::open(&path)
461 .await?
462 .take(read_limit)
463 .read_to_end(&mut bytes)
464 .await?;
465 let actual_bytes = u64::try_from(bytes.len()).unwrap_or(u64::MAX);
466 if actual_bytes > self.max_snapshot_bytes {
467 return Err(CacheError::SnapshotTooLarge {
468 actual_bytes,
469 max_bytes: self.max_snapshot_bytes,
470 });
471 }
472 Ok(Some(bytes))
473 }
474
475 async fn load_unlocked(&self) -> Result<EntryMap<K, V, M>> {
476 let Some(bytes) = self.read_snapshot_bytes().await? else {
477 return Ok(HashMap::new());
478 };
479 let snapshot: Snapshot<K, V, M> = serde_json::from_slice(&bytes)
480 .map_err(|error| CacheError::Deserialization(error.to_string()))?;
481 if snapshot.version != SNAPSHOT_FORMAT_VERSION {
482 return Err(CacheError::UnsupportedPersistenceFormat(format!(
483 "snapshot version {} is not supported (expected {SNAPSHOT_FORMAT_VERSION})",
484 snapshot.version
485 )));
486 }
487 Ok(snapshot.entries)
488 }
489
490 async fn replace_snapshot(&self, bytes: &[u8]) -> Result<()> {
491 let actual_bytes = u64::try_from(bytes.len()).unwrap_or(u64::MAX);
492 if actual_bytes > self.max_snapshot_bytes {
493 return Err(CacheError::SnapshotTooLarge {
494 actual_bytes,
495 max_bytes: self.max_snapshot_bytes,
496 });
497 }
498
499 let temporary_path = self.temporary_path();
500 let write_result = async {
501 let mut options = OpenOptions::new();
502 options.write(true).create_new(true);
503 #[cfg(unix)]
504 {
505 options.mode(0o600);
506 }
507 let mut file = options.open(&temporary_path).await?;
508 file.write_all(bytes).await?;
509 file.flush().await?;
510 file.sync_all().await?;
511 drop(file);
512
513 #[cfg(not(windows))]
514 fs::rename(&temporary_path, self.snapshot_path()).await?;
515
516 #[cfg(windows)]
520 {
521 let snapshot_path = self.snapshot_path();
522 match fs::rename(&temporary_path, &snapshot_path).await {
523 Ok(()) => {}
524 Err(error)
525 if matches!(
526 error.kind(),
527 std::io::ErrorKind::AlreadyExists
528 | std::io::ErrorKind::PermissionDenied
529 ) =>
530 {
531 fs::remove_file(&snapshot_path).await?;
532 fs::rename(&temporary_path, &snapshot_path).await?;
533 }
534 Err(error) => return Err(error),
535 }
536 }
537
538 #[cfg(unix)]
539 File::open(&self.base_path).await?.sync_all().await?;
540 Ok::<(), std::io::Error>(())
541 }
542 .await;
543
544 if write_result.is_err() {
545 let _ = fs::remove_file(&temporary_path).await;
546 }
547 write_result.map_err(Into::into)
548 }
549
550 async fn save_unlocked(&self, entries: &EntryMap<K, V, M>) -> Result<()> {
551 self.reject_legacy_layout().await?;
552 if entries.len() > MAX_SNAPSHOT_KEYS {
553 return Err(CacheError::CapacityExceeded {
554 message: format!(
555 "snapshot contains {} keys; limit is {MAX_SNAPSHOT_KEYS}",
556 entries.len()
557 ),
558 });
559 }
560 let total_entries = entries.values().try_fold(0usize, |total, key_entries| {
561 total
562 .checked_add(key_entries.len())
563 .ok_or_else(|| CacheError::CapacityExceeded {
564 message: "snapshot entry count overflowed usize".to_string(),
565 })
566 })?;
567 if total_entries > MAX_SNAPSHOT_ENTRIES {
568 return Err(CacheError::CapacityExceeded {
569 message: format!("snapshot contains more than {MAX_SNAPSHOT_ENTRIES} entries"),
570 });
571 }
572 for (key, key_entries) in entries {
573 if key_entries.iter().any(|entry| &entry.key != key) {
574 return Err(CacheError::Serialization(
575 "cannot persist an entry whose embedded key does not match".to_string(),
576 ));
577 }
578 }
579 let snapshot = SnapshotRef {
580 version: SNAPSHOT_FORMAT_VERSION,
581 entries: entries.iter().collect(),
582 };
583 let max_bytes = usize::try_from(self.max_snapshot_bytes).unwrap_or(usize::MAX);
584 let mut writer = BoundedBuffer::new(max_bytes);
585 if let Err(error) = serde_json::to_writer(&mut writer, &snapshot) {
586 if writer.exceeded {
587 return Err(CacheError::SnapshotTooLarge {
588 actual_bytes: writer.attempted_bytes,
589 max_bytes: self.max_snapshot_bytes,
590 });
591 }
592 return Err(CacheError::Serialization(error.to_string()));
593 }
594 self.replace_snapshot(&writer.bytes).await
595 }
596}
597
598#[async_trait]
599impl<K, V, M> StorageBackend for FilesystemBackend<K, V, M>
600where
601 K: StorageKey,
602 V: StorageValue,
603 M: StorageMeta,
604{
605 type Value = V;
606 type Key = K;
607 type Metadata = M;
608
609 async fn save(&self, entries: &EntryMap<K, V, M>) -> Result<()> {
610 let _guard = self.io_lock.lock().await;
611 self.save_unlocked(entries).await
612 }
613
614 async fn load(&self) -> Result<EntryMap<K, V, M>> {
615 let _guard = self.io_lock.lock().await;
616 self.load_unlocked().await
617 }
618
619 async fn remove(&self, key: &K) -> Result<()> {
620 let _guard = self.io_lock.lock().await;
621 let mut entries = self.load_unlocked().await?;
622 if entries.remove(key).is_some() {
623 self.save_unlocked(&entries).await?;
624 }
625 Ok(())
626 }
627
628 async fn clear(&self) -> Result<()> {
629 let _guard = self.io_lock.lock().await;
630 self.reject_legacy_layout().await?;
631 let snapshot_path = self.snapshot_path();
632 let metadata = match fs::symlink_metadata(&snapshot_path).await {
633 Ok(metadata) => metadata,
634 Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(()),
635 Err(error) => return Err(error.into()),
636 };
637 if metadata.file_type().is_symlink() || !metadata.is_file() {
638 return Err(CacheError::StorageBackend(
639 "cache snapshot must be a regular file".to_string(),
640 ));
641 }
642 fs::remove_file(snapshot_path).await?;
643 #[cfg(unix)]
644 File::open(&self.base_path).await?.sync_all().await?;
645 Ok(())
646 }
647
648 async fn contains(&self, key: &K) -> Result<bool> {
649 let _guard = self.io_lock.lock().await;
650 Ok(self.load_unlocked().await?.contains_key(key))
651 }
652
653 async fn size_bytes(&self) -> Result<u64> {
654 let _guard = self.io_lock.lock().await;
655 self.reject_legacy_layout().await?;
656 match fs::symlink_metadata(self.snapshot_path()).await {
657 Ok(metadata) if metadata.is_file() && !metadata.file_type().is_symlink() => {
658 Ok(metadata.len())
659 }
660 Ok(_) => Err(CacheError::StorageBackend(
661 "cache snapshot must be a regular file".to_string(),
662 )),
663 Err(error) if error.kind() == std::io::ErrorKind::NotFound => Ok(0),
664 Err(error) => Err(error.into()),
665 }
666 }
667
668 async fn compact(&self) -> Result<()> {
669 let _guard = self.io_lock.lock().await;
670 let entries = self.load_unlocked().await?;
671 self.save_unlocked(&entries).await
672 }
673}
674
675#[cfg(test)]
676mod tests {
677 use super::*;
678 use tempfile::TempDir;
679
680 async fn new_backend() -> (TempDir, FilesystemBackend<String, String>) {
681 let temp_dir = TempDir::new().unwrap();
682 let backend = FilesystemBackend::new(temp_dir.path()).await.unwrap();
683 (temp_dir, backend)
684 }
685
686 fn entries(values: &[(&str, &str)]) -> EntryMap<String, String, ()> {
687 values
688 .iter()
689 .map(|(key, value)| {
690 (
691 (*key).to_string(),
692 vec![CacheEntry::new((*key).to_string(), (*value).to_string())],
693 )
694 })
695 .collect()
696 }
697
698 #[tokio::test]
699 async fn persists_and_atomically_replaces_snapshot() {
700 let (temp_dir, backend) = new_backend().await;
701 backend.save(&entries(&[("one", "1")])).await.unwrap();
702 backend.save(&entries(&[("two", "2")])).await.unwrap();
703
704 let loaded = backend.load().await.unwrap();
705 assert_eq!(loaded.len(), 1);
706 assert_eq!(loaded["two"][0].value, "2");
707 assert!(!loaded.contains_key("one"));
708
709 let file_names: Vec<_> = std::fs::read_dir(temp_dir.path())
710 .unwrap()
711 .map(|entry| entry.unwrap().file_name())
712 .collect();
713 assert_eq!(
714 file_names,
715 vec![std::ffi::OsString::from(SNAPSHOT_FILE_NAME)]
716 );
717 }
718
719 #[tokio::test]
720 async fn formerly_colliding_and_traversal_keys_round_trip() {
721 let (_temp_dir, backend) = new_backend().await;
722 let values = entries(&[("a/b", "slash"), ("a\\b", "backslash"), ("../x", "dot")]);
723 backend.save(&values).await.unwrap();
724 let loaded = backend.load().await.unwrap();
725 assert_eq!(loaded.len(), 3);
726 assert_eq!(loaded["a/b"][0].value, "slash");
727 assert_eq!(loaded["a\\b"][0].value, "backslash");
728 assert_eq!(loaded["../x"][0].value, "dot");
729 }
730
731 #[tokio::test]
732 async fn corrupted_snapshot_is_an_error() {
733 let (temp_dir, backend) = new_backend().await;
734 fs::write(temp_dir.path().join(SNAPSHOT_FILE_NAME), b"not json")
735 .await
736 .unwrap();
737 assert!(matches!(
738 backend.load().await,
739 Err(CacheError::Deserialization(_))
740 ));
741 }
742
743 #[tokio::test]
744 async fn oversized_snapshot_is_rejected_before_deserialization() {
745 let (temp_dir, backend) = new_backend().await;
746 fs::write(temp_dir.path().join(SNAPSHOT_FILE_NAME), b"123456789")
747 .await
748 .unwrap();
749 let backend = backend.with_max_snapshot_bytes(8);
750 assert!(matches!(
751 backend.load().await,
752 Err(CacheError::SnapshotTooLarge {
753 actual_bytes: 9,
754 max_bytes: 8
755 })
756 ));
757 }
758
759 #[tokio::test]
760 async fn oversized_save_does_not_replace_the_previous_snapshot() {
761 let (temp_dir, backend) = new_backend().await;
762 backend
763 .save(&entries(&[("stable", "value")]))
764 .await
765 .unwrap();
766
767 let constrained: FilesystemBackend<String, String> =
768 FilesystemBackend::new(temp_dir.path())
769 .await
770 .unwrap()
771 .with_max_snapshot_bytes(32);
772 let result = constrained
773 .save(&entries(&[("large", &"x".repeat(1_024))]))
774 .await;
775 assert!(matches!(result, Err(CacheError::SnapshotTooLarge { .. })));
776
777 let loaded = backend.load().await.unwrap();
778 assert!(loaded.contains_key("stable"));
779 assert!(!loaded.contains_key("large"));
780 }
781
782 #[test]
783 fn snapshot_deserialization_enforces_limits_while_streaming() {
784 let too_many_keys = serde_json::to_string(&vec![
785 (
786 "one".to_string(),
787 vec![CacheEntry::<String, String>::new(
788 "one".to_string(),
789 "1".to_string(),
790 )],
791 ),
792 ("two".to_string(), Vec::new()),
793 ])
794 .unwrap();
795 let mut deserializer = serde_json::Deserializer::from_str(&too_many_keys);
796 let key_result =
797 deserializer.deserialize_seq(SnapshotEntriesVisitor::<String, String, ()> {
798 max_keys: 1,
799 max_entries: 10,
800 marker: std::marker::PhantomData,
801 });
802 assert!(key_result.is_err());
803
804 let too_many_entries = serde_json::to_string(&vec![(
805 "one".to_string(),
806 vec![
807 CacheEntry::<String, String>::new("one".to_string(), "1".to_string()),
808 CacheEntry::<String, String>::new("one".to_string(), "2".to_string()),
809 ],
810 )])
811 .unwrap();
812 let mut deserializer = serde_json::Deserializer::from_str(&too_many_entries);
813 let entry_result =
814 deserializer.deserialize_seq(SnapshotEntriesVisitor::<String, String, ()> {
815 max_keys: 10,
816 max_entries: 1,
817 marker: std::marker::PhantomData,
818 });
819 assert!(entry_result.is_err());
820 }
821
822 #[tokio::test]
823 async fn legacy_layout_is_reported_explicitly() {
824 let (temp_dir, backend) = new_backend().await;
825 fs::write(temp_dir.path().join("metadata.json"), b"{}")
826 .await
827 .unwrap();
828 assert!(matches!(
829 backend.load().await,
830 Err(CacheError::UnsupportedPersistenceFormat(_))
831 ));
832 }
833
834 #[tokio::test]
835 async fn legacy_bincode_layout_is_reported_explicitly() {
836 let (temp_dir, backend) = new_backend().await;
837 fs::write(temp_dir.path().join("entry.bin"), b"legacy")
838 .await
839 .unwrap();
840 assert!(matches!(
841 backend.load().await,
842 Err(CacheError::UnsupportedPersistenceFormat(_))
843 ));
844 }
845
846 #[tokio::test]
847 async fn clear_rejects_legacy_files_before_removing_the_snapshot() {
848 let (temp_dir, backend) = new_backend().await;
849 backend
850 .save(&entries(&[("stable", "value")]))
851 .await
852 .unwrap();
853 let legacy_path = temp_dir.path().join("legacy.json");
854 fs::write(&legacy_path, b"{}").await.unwrap();
855
856 assert!(matches!(
857 backend.clear().await,
858 Err(CacheError::UnsupportedPersistenceFormat(_))
859 ));
860 assert!(
861 fs::try_exists(temp_dir.path().join(SNAPSHOT_FILE_NAME))
862 .await
863 .unwrap()
864 );
865 assert!(fs::try_exists(legacy_path).await.unwrap());
866 }
867
868 #[tokio::test]
869 async fn unknown_snapshot_version_is_rejected() {
870 let (temp_dir, backend) = new_backend().await;
871 fs::write(
872 temp_dir.path().join(SNAPSHOT_FILE_NAME),
873 br#"{"version":99,"entries":[]}"#,
874 )
875 .await
876 .unwrap();
877 assert!(matches!(
878 backend.load().await,
879 Err(CacheError::UnsupportedPersistenceFormat(_))
880 ));
881 }
882
883 #[tokio::test]
884 async fn filesystem_backend_size_tracks_snapshot() {
885 let (_temp_dir, backend) = new_backend().await;
886 assert_eq!(backend.size_bytes().await.unwrap(), 0);
887 backend.save(&entries(&[("key", "value")])).await.unwrap();
888 assert!(backend.size_bytes().await.unwrap() > 0);
889 backend.clear().await.unwrap();
890 assert_eq!(backend.size_bytes().await.unwrap(), 0);
891 }
892
893 #[cfg(unix)]
894 #[tokio::test]
895 async fn rejects_symlink_cache_directory() {
896 use std::os::unix::fs::symlink;
897
898 let target = TempDir::new().unwrap();
899 let parent = TempDir::new().unwrap();
900 let link = parent.path().join("cache-link");
901 symlink(target.path(), &link).unwrap();
902 let result = FilesystemBackend::<String, String>::new(&link).await;
903 assert!(matches!(result, Err(CacheError::InvalidConfiguration(_))));
904 }
905}