Skip to main content

turbo_tasks_backend/backend/
storage.rs

1use std::{
2    cell::Cell,
3    fmt::{Display, Formatter},
4    hash::{BuildHasher, Hash},
5    ops::{Deref, DerefMut},
6    sync::{
7        Arc,
8        atomic::{AtomicBool, AtomicU64, Ordering},
9    },
10};
11
12use crossbeam_utils::CachePadded;
13use hashbrown::hash_table;
14use thread_local::ThreadLocal;
15use tracing::span::Id;
16use turbo_bincode::TurboBincodeBuffer;
17use turbo_tasks::{FxDashMap, TaskId, backend::CachedTaskTypeArc, event::Event, parallel};
18
19use crate::{
20    backend::storage_schema::{
21        DropPartialOutcome, KeyEvictability, TaskStorage, UnevictableReason, ValueEvictability,
22    },
23    backing_storage::SnapshotItem,
24    database::key_value_database::KeySpace,
25    utils::{
26        dash_map_drop_contents::drop_contents,
27        dash_map_entry::{TryLockAndRemove, try_lock_and_remove},
28        dash_map_multi::{RefMut, get_disjoint_mut},
29    },
30};
31
32#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
33pub enum TaskDataCategory {
34    Meta,
35    Data,
36    All,
37}
38impl PartialOrd for TaskDataCategory {
39    /// `All` is greater than both `Meta` and `Data`; `Meta` and `Data` are unordered.
40    fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
41        use std::cmp::Ordering::*;
42
43        use TaskDataCategory::All;
44        match (self, other) {
45            _ if self == other => Some(Equal),
46            (All, _) => Some(Greater),
47            (_, All) => Some(Less),
48            _ => None,
49        }
50    }
51}
52
53/// Counts of tasks evicted at each level.
54#[derive(Debug, Default)]
55pub struct EvictionCounts {
56    pub key_evictions: usize,
57    pub full: usize,
58    pub data_and_meta: usize,
59    pub data_only: usize,
60    pub meta_only: usize,
61    /// Per-reason counts of tasks we considered but could not evict, indexed by
62    /// `UnevictableReason::index()`.
63    pub unevictable_reasons: [usize; UnevictableReason::COUNT],
64}
65
66impl std::ops::AddAssign for EvictionCounts {
67    fn add_assign(&mut self, rhs: Self) {
68        self.key_evictions += rhs.key_evictions;
69        self.full += rhs.full;
70        self.data_and_meta += rhs.data_and_meta;
71        self.data_only += rhs.data_only;
72        self.meta_only += rhs.meta_only;
73        for i in 0..UnevictableReason::COUNT {
74            self.unevictable_reasons[i] += rhs.unevictable_reasons[i];
75        }
76    }
77}
78
79impl Display for EvictionCounts {
80    /// Compact `field=value,...` form used as a single tracing span field so that
81    /// adding a new counter or `UnevictableReason` variant doesn't require updating
82    /// the span field list.
83    fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
84        let skipped: usize = self.unevictable_reasons.iter().sum();
85        write!(
86            f,
87            "task_cache_evictions={},full={},data_and_meta={},data_only={},meta_only={},skipped={}",
88            self.key_evictions,
89            self.full,
90            self.data_and_meta,
91            self.data_only,
92            self.meta_only,
93            skipped,
94        )?;
95        for reason in UnevictableReason::ALL {
96            write!(
97                f,
98                ",{}={}",
99                reason.span_name(),
100                self.unevictable_reasons[reason.index()],
101            )?;
102        }
103        Ok(())
104    }
105}
106
107impl TaskDataCategory {
108    pub fn includes_data(self) -> bool {
109        matches!(self, TaskDataCategory::Data | TaskDataCategory::All)
110    }
111
112    pub fn includes_meta(self) -> bool {
113        matches!(self, TaskDataCategory::Meta | TaskDataCategory::All)
114    }
115}
116
117#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
118pub enum SpecificTaskDataCategory {
119    Meta,
120    Data,
121}
122
123impl From<SpecificTaskDataCategory> for TaskDataCategory {
124    fn from(category: SpecificTaskDataCategory) -> Self {
125        match category {
126            SpecificTaskDataCategory::Meta => TaskDataCategory::Meta,
127            SpecificTaskDataCategory::Data => TaskDataCategory::Data,
128        }
129    }
130}
131
132impl SpecificTaskDataCategory {
133    /// Returns the KeySpace for storing data of this category
134    pub fn key_space(self) -> KeySpace {
135        match self {
136            SpecificTaskDataCategory::Meta => KeySpace::TaskMeta,
137            SpecificTaskDataCategory::Data => KeySpace::TaskData,
138        }
139    }
140}
141
142/// Records exactly what a `track_modification` call changed, so that
143/// [`StorageWriteGuard::undo_track_modification`] can reverse it precisely when the mutation it
144/// guarded turns out to be a no-op.  This allows us to track modifications 'optimistically' and
145/// undo it if the modification turned out to be a no op.  Useful when dealing with datastructures
146/// like `AutoSet` that can efficiently say whether or not they were modified.
147#[must_use = "a no-op mutation must undo its TrackOutcome; dropping it leaks an over-track"]
148pub enum TrackOutcome {
149    /// Nothing was tracked: either the category was already modified, or (in snapshot mode) it was
150    /// already modified-during-snapshot. Undo is a no-op.
151    NoChange,
152    /// Non-snapshot path: `modified(category)` was set. `bumped` is true if this call also
153    /// incremented the per-shard modified counter (i.e. the task had no prior modifications).
154    Tracked {
155        category: SpecificTaskDataCategory,
156        bumped: bool,
157    },
158    /// Snapshot path: `modified_during_snapshot(category)` was set. `inserted_snapshot` is true if
159    /// this call also inserted the task's entry into the `snapshots` map (the pre-mutation copy or
160    /// a `None` marker).
161    TrackedDuringSnapshot {
162        category: SpecificTaskDataCategory,
163        inserted_snapshot: bool,
164    },
165}
166
167pub struct Storage {
168    snapshot_mode: AtomicBool,
169    /// Per-shard counts of tasks with modified flags set. Incremented when a task
170    /// transitions from unmodified to modified (outside snapshot mode). Reset to zero when
171    /// snapshot mode begins, and re-incremented in `end_snapshot` for tasks that still have
172    /// modifications (promoted from `modified_during_snapshot`). Used to skip unmodified shards
173    /// in `take_snapshot`, avoiding unnecessary iteration and enabling early returns
174    ///
175    /// Indexed by `map.determine_shard(map.hash_usize(&key))` and guaranteed by construction so
176    /// that  `shard_modified_counts.len()==map.shards().len()`
177    ///
178    /// Should only be modified while holding the corresponding dashmap shard lock.
179    shard_modified_counts: Box<[CachePadded<AtomicU64>]>,
180    /// Stores snapshots of task state for tasks accessed during snapshot mode.
181    /// - `Some(snapshot)`: Task was modified before snapshot mode and accessed again during it.
182    ///   Contains a copy of the pre-snapshot state that needs to be persisted.
183    /// - `None`: Task was first modified during snapshot mode (not part of current snapshot). Will
184    ///   be marked as modified at the beginning of the next snapshot cycle.
185    ///
186    /// Lock Ordering: `snapshots` locks are acquired **after** `map` locks (see the comment on
187    /// `map` below). Holding a `snapshots` shard write lock and then trying to take a `map` shard
188    /// write lock is forbidden — it would deadlock against `track_modification_internal` /
189    /// `SnapshotShardIter::next`, which take map first.
190    ///
191    /// Shard Invariant: `snapshots` is constructed with the same `shard_amount`, the same key
192    /// type (`TaskId`), and the same stateless hasher (`FxBuildHasher`) as `map`. Therefore shard
193    /// index `N` in `snapshots` corresponds exactly to shard index `N` in `map`: any `TaskId`
194    /// present in `snapshots.shards()[N]` (if present in `map` at all) is in `map.shards()[N]`.
195    /// Code that walks both maps in parallel (e.g. `end_snapshot`) relies on this to lock pairs
196    /// of shards by index instead of going through the top-level `DashMap` accessors.
197    snapshots: FxDashMap<TaskId, Option<Box<TaskStorage>>>,
198    /// The main storage map
199    ///
200    /// Lock Ordering: Task creation acquires a `task_cache` lock and then inserts into this map.
201    /// Because both datastructures are sharded on different keys, the locks are not 'strictly'
202    /// ordered but we should treat them as such
203    /// Acquiring locks in the opposite order should be defensive
204    ///
205    /// Lock Ordering vs. `snapshots`: `map` locks are acquired **before** `snapshots` locks.
206    /// `track_modification_internal` and `SnapshotShardIter::next` both hold a `map` shard write
207    /// lock (via `StorageWriteGuard` / `map.get_mut`) and then take a `snapshots` shard lock.
208    /// `end_snapshot` must lock in the same order — see the shard-zipping pattern there.
209    map: FxDashMap<TaskId, Box<TaskStorage>>,
210    /// A shared event notified whenever any task finishes restoring (successfully or not).
211    ///
212    /// Threads waiting for another thread's in-progress restore subscribe to this event,
213    /// then re-check the specific task's `restoring`/`restored` bits after waking.
214    pub(crate) restored: Event,
215    /// Maps `CachedTaskType` → `TaskId` for deduplication of persistent task creation.
216    /// This is backed by the TaskCache table in the database.
217    ///
218    /// LockOrdering: See the comments on [map].
219    pub task_cache: FxDashMap<CachedTaskTypeArc, TaskId>,
220}
221
222impl Storage {
223    pub fn new(shard_amount: usize, small_preallocation: bool) -> Self {
224        let map_capacity: usize = if small_preallocation {
225            1024
226        } else {
227            1024 * 1024
228        };
229
230        let map = FxDashMap::with_capacity_and_hasher_and_shard_amount(
231            map_capacity,
232            Default::default(),
233            shard_amount,
234        );
235        let shard_modified_counts = (0..shard_amount)
236            .map(|_| CachePadded::new(AtomicU64::new(0)))
237            .collect::<Vec<_>>()
238            .into_boxed_slice();
239        Self {
240            snapshot_mode: AtomicBool::new(false),
241            shard_modified_counts,
242            snapshots: FxDashMap::with_capacity_and_hasher_and_shard_amount(
243                // We expect very few updates to this map since it will only happen when updates
244                // race with snapshots.  This never happens in a build and only rarely happens in
245                // dev sessions
246                0,
247                Default::default(),
248                shard_amount,
249            ),
250            map,
251            restored: Event::new(|| || "Storage::restored".to_string()),
252            task_cache: FxDashMap::default(),
253        }
254    }
255
256    /// Returns the shard index for the given key in the `map` DashMap.
257    fn shard_index(&self, key: &TaskId) -> usize {
258        let hash = self.map.hash_usize(key);
259        self.map.determine_shard(hash)
260    }
261
262    /// Promote `modified_during_snapshot` → `modified` flags on a task, and increment the
263    /// per-shard modified count if the task was not already marked as modified.
264    ///
265    /// This is used after persisting a snapshot: _during_snapshot flags represent changes
266    /// that occurred concurrently and were not included in the persisted snapshot, so they
267    /// must be carried forward as `modified` for the next snapshot cycle.
268    fn promote_during_snapshot_flags(&self, task: &mut TaskStorage, shard_idx: usize) {
269        let already_modified = task.flags.any_modified();
270        let mut promoted = false;
271        if task.flags.meta_modified_during_snapshot() {
272            task.flags.set_meta_modified_during_snapshot(false);
273            task.flags.set_meta_modified(true);
274            promoted = true;
275        }
276        if task.flags.data_modified_during_snapshot() {
277            task.flags.set_data_modified_during_snapshot(false);
278            task.flags.set_data_modified(true);
279            promoted = true;
280        }
281        if !already_modified && promoted {
282            self.shard_modified_counts[shard_idx].fetch_add(1, Ordering::Relaxed);
283        }
284    }
285
286    /// Mark a newly allocated task as restored (skip DB queries) and new (include in persistence
287    /// snapshots). Optionally sets the `persistent_task_type` eagerly so it's available for
288    /// persistence snapshots without needing to propagate it through `connect_child`.
289    pub fn initialize_new_task(&self, task_id: TaskId, task_type: Option<CachedTaskTypeArc>) {
290        let mut task = self.access_mut(task_id);
291        task.flags.set_restored(TaskDataCategory::All);
292        task.flags.set_new_task(true);
293        task.gc_pin_for_construction();
294        if let Some(task_type) = task_type {
295            task.set_persistent_task_type(task_type);
296            if !task_id.is_transient() {
297                // Unconditional track: a new task's type is always a real persistable change.
298                let _ =
299                    task.track_modification(SpecificTaskDataCategory::Data, "persistent_task_type");
300            }
301        }
302    }
303
304    /// Processes every modified item (resp. a snapshot of it) with the given function and returns
305    /// the results. Ends snapshot mode when the returned `SnapshotGuard` (held by each shard) is
306    /// dropped.
307    ///
308    /// `process` is called while holding a read lock on the task storage, so it can access
309    /// the TaskStorage directly without cloning.
310    ///
311    /// Both callbacks receive a mutable scratch buffer that can be reused across iterations
312    /// to avoid repeated allocations.
313    ///
314    /// The returned shards implement `IntoIterator`. Empty shards (no modified or snapshot
315    /// entries) are filtered out, but shards may still yield no items if all entries produce
316    /// empty `SnapshotItem`s (this is rare and only happens under error conditions).
317    ///
318    /// When `drain_entries` is true (shutdown only), the scan drains the map: unmodified entries
319    /// are erased and freed immediately, and the modified entries are moved out into the
320    /// returned shard iterators, which free each task's memory as it is serialized rather than
321    /// after the whole batch is written.
322    pub fn take_snapshot<
323        'l,
324        P: for<'a> Fn(TaskId, &'a TaskStorage, &mut TurboBincodeBuffer) -> SnapshotItem + Sync,
325    >(
326        &'l self,
327        guard: SnapshotGuard<'l>,
328        process: &'l P,
329        drain_entries: bool,
330    ) -> Vec<SnapshotShard<'l, P>> {
331        let guard = Arc::new(guard);
332
333        let shards: Vec<_> = self.map.shards().iter().enumerate().collect();
334
335        // The number of shards is much larger than the number of threads, so the effect of the
336        // locks held is negligible.
337        parallel::map_collect::<_, _, Vec<_>>(&shards, |&(shard_idx, shard)| {
338            // Check how many modifications there are in this shard, because we have entered
339            // snapshot_mode, there are no racing writes
340            // So we can safely clear it out now that we are processing the modifications
341            let modified_count = self.shard_modified_counts[shard_idx].swap(0, Ordering::Relaxed);
342
343            if modified_count == 0 && !drain_entries {
344                // Nothing to persist in this shard and we're keeping the map, so skip the scan.
345                // TODO: when not draining but eviction is enabled we should run that logic here as
346                // well
347                return None;
348            }
349
350            // Scan the shard once, building the work this shard's iterator will perform. The two
351            // modes carry different data so that `next` has no per-item `drain` branch:
352            // - keep mode collects the modified `TaskId`s and looks them up again while iterating.
353            // - drain mode erases the unmodified entries here and then moves the remaining
354            //   (modified-only) table out of the map, so the iterator owns and drains it directly.
355            let work = {
356                let mut shard_guard = shard.write();
357                if drain_entries {
358                    shard_guard.retain(|(key, task)| {
359                        let modified_task = task.flags.any_modified();
360                        if modified_task {
361                            debug_assert!(
362                                !key.is_transient(),
363                                "found a modified transient task: {key:?}"
364                            );
365                        }
366                        // Unmodified entries are not part of the snapshot. Remove and free them
367                        // now so the table we move out below holds only modified entries.
368                        modified_task
369                    });
370                    if shard_guard.is_empty() {
371                        // The shard held only unmodified entries, which we've now erased and freed.
372                        // No iterator is created for an empty shard.
373                        return None;
374                    }
375                    // Move the modified-only table out of the map. Iterating it frees each task box
376                    // as it is serialized, and the shard's table allocation is released here.
377                    ShardWork::Drain(std::mem::take(&mut *shard_guard).into_iter())
378                } else {
379                    let mut modified = Vec::with_capacity(modified_count as usize);
380                    for (key, task) in shard_guard.iter() {
381                        // Only check modified flags — transient tasks never have modified flags set
382                        // (track_modification guards against it), so this naturally excludes them.
383                        // new_task always comes with modified flags (set_persistent_task_type calls
384                        // track_modification), so any_modified() is sufficient.
385                        if task.flags.any_modified() {
386                            debug_assert!(
387                                !key.is_transient(),
388                                "found a modified transient task: {key:?}"
389                            );
390                            modified.push(*key);
391                        }
392                    }
393                    // modified_count > 0 (we returned early otherwise), so this is never empty.
394                    debug_assert!(!modified.is_empty());
395                    ShardWork::Keep(modified)
396                }
397            };
398
399            Some(SnapshotShard {
400                shard_idx,
401                work,
402                storage: self,
403                process,
404                _guard: guard.clone(),
405            })
406        })
407        .into_iter()
408        .flatten()
409        .collect()
410    }
411
412    /// Enter snapshot mode and return a guard that will call `end_snapshot` on drop.
413    ///
414    /// Returns whether any shard has modifications. Per-shard counts are reset
415    /// in `take_snapshot` as each shard is processed, not here — resetting eagerly
416    /// would lose track of modifications for shards that haven't been persisted yet.
417    ///
418    /// Safety invariant: `start_snapshot` and `end_snapshot` are always called
419    /// sequentially within a single `snapshot_and_persist` invocation (the sole
420    /// caller). There is no concurrent snapshot lifecycle, so they cannot race.
421    pub fn start_snapshot(&self) -> (SnapshotGuard<'_>, bool) {
422        // Enter snapshot mode first so concurrent track_modification calls switch
423        // to the _during_snapshot path and stop incrementing shard_modified_counts.
424        self.snapshot_mode.store(true, Ordering::Release);
425        // Check if any shard has modifications. Don't reset counts here —
426        // take_snapshot resets per-shard counts as it processes each shard,
427        // which avoids losing track of modifications for shards that haven't
428        // been persisted yet.
429        let has_modifications = self
430            .shard_modified_counts
431            .iter()
432            .any(|c| c.load(Ordering::Relaxed) > 0);
433        (SnapshotGuard::new(self), has_modifications)
434    }
435
436    /// End snapshot mode.
437    ///
438    /// Modified/new flags on tasks are cleared incrementally during snapshot iteration
439    /// (in `take_snapshot` for direct_snapshots, and in `SnapshotShardIter::next` for
440    /// modified tasks), so no full-map scan is needed here.
441    ///
442    /// This method only needs to:
443    /// 1. Leave snapshot mode so new modifications go to the modified flags directly.
444    /// 2. Promote `modified_during_snapshot` → `modified` for tasks that were accessed during
445    ///    snapshot mode (tracked in the small `snapshots` map).
446    fn end_snapshot(&self) {
447        // Leave snapshot mode first. After this, concurrent track_modification calls
448        // will set modified flags directly instead of going through the snapshots map.
449        self.snapshot_mode.store(false, Ordering::Release);
450
451        // Promote modified_during_snapshot → modified for tasks that had snapshots.
452        // The snapshots map should be small (only tasks concurrently accessed during snapshot
453        // mode). Increment the per-shard modified counts for promoted tasks.
454
455        // Lock Ordering: we must acquire `map` shards BEFORE `snapshots` shards, matching the
456        // order used by `track_modification_internal` and `SnapshotShardIter::next`. The
457        // previous implementation drained `snapshots` first and then called `self.map.get_mut`,
458        // which is the opposite order — a concurrent `track_modification` (holding map[N], about
459        // to insert into snapshots[N]) could deadlock against it through the
460        // `snapshot_mode = false` race window.
461        //
462        // Shard pairing: `map` and `snapshots` are constructed with the same `shard_amount`,
463        // same `TaskId` keys, and the same stateless `FxBuildHasher`. Therefore shard `N` in
464        // `snapshots` pairs with shard `N` in `map`: every key drained from `snapshots[N]` (if
465        // it still exists in `map`) lives in `map[N]`. We zip them and lock each pair in order.
466        let map_shards = self.map.shards();
467        let snapshot_shards = self.snapshots.shards();
468        debug_assert_eq!(
469            map_shards.len(),
470            snapshot_shards.len(),
471            "map and snapshots must share shard count for zipped locking; see Shard Invariant on \
472             `snapshots` field"
473        );
474
475        let shard_indices: Vec<usize> = (0..map_shards.len()).collect();
476        parallel::for_each(&shard_indices, |&shard_idx| {
477            let map_shard = &map_shards[shard_idx];
478            let snap_shard = &snapshot_shards[shard_idx];
479
480            // Acquire in documented order: map first, snapshots second.
481            let mut map_guard = map_shard.write();
482            let mut snap_guard = snap_shard.write();
483
484            for (key, _) in snap_guard.drain() {
485                // The key is in this shard's `map` (or absent entirely), by the shard
486                // invariant above. Resolve directly in the held map guard rather than going
487                // through `self.map.get_mut`, which would attempt to re-acquire this shard's
488                // write lock and would also obscure the pairing.
489                let hash = self.map.hasher().hash_one(key);
490                if let Some((_, task)) = map_guard.find_mut(hash, |(k, _)| *k == key) {
491                    self.promote_during_snapshot_flags(task, shard_idx);
492                }
493            }
494            // If we are saving a non-trivial amount of memory just clear it out.
495            if snap_guard.capacity() > 1024 {
496                snap_guard.shrink_to(0, |_entry| {
497                    unreachable!("nothing is hashed when resizing an empty shard to zero");
498                });
499            }
500
501            drop(snap_guard);
502            drop(map_guard);
503        });
504    }
505
506    /// Returns true if actively snapshotting (modifications should go to snapshots map).
507    /// Returns false if inactive (modifications go to modified list).
508    fn snapshot_mode(&self) -> bool {
509        self.snapshot_mode.load(Ordering::Acquire)
510    }
511
512    pub fn access_mut(&self, key: TaskId) -> StorageWriteGuard<'_> {
513        let inner = match self.map.entry(key) {
514            dashmap::mapref::entry::Entry::Occupied(e) => e.into_ref(),
515            dashmap::mapref::entry::Entry::Vacant(e) => e.insert(Box::new(TaskStorage::new())),
516        };
517        StorageWriteGuard {
518            storage: self,
519            inner: inner.into(),
520        }
521    }
522
523    /// Like [`Self::access_mut`], but keeps the map entry so the caller can still remove it.
524    pub fn access_entry_mut(&self, key: TaskId) -> TaskEntryGuard<'_> {
525        let entry = match self.map.entry(key) {
526            dashmap::mapref::entry::Entry::Occupied(e) => e,
527            dashmap::mapref::entry::Entry::Vacant(e) => {
528                e.insert_entry(Box::new(TaskStorage::new()))
529            }
530        };
531        TaskEntryGuard {
532            storage: self,
533            entry,
534        }
535    }
536
537    /// Read-only access to an already resident task. Returns `None` if the task isnt in memory
538    /// resident. The closure runs while a shard read lock is held, so it must be cheap and must
539    /// not re-enter the map.
540    pub fn with_task<R>(&self, key: TaskId, f: impl FnOnce(&TaskStorage) -> R) -> Option<R> {
541        let task = self.map.get(&key)?;
542        Some(f(task.value()))
543    }
544
545    /// The number of tasks resident in the map.
546    #[doc(hidden)]
547    pub fn resident_task_count_for_testing(&self) -> usize {
548        self.map.len()
549    }
550
551    /// The number of shards in the resident map. GC seeds one `ScanShard` job per index; the slice
552    /// returned by `map.shards()` is fixed for the map's lifetime, so an index is a stable handle
553    /// to one shard.
554    pub fn shard_count(&self) -> usize {
555        self.map.shards().len()
556    }
557
558    /// Scans a **single** shard by index, invoking `on_candidate` for each resident task whose
559    /// storage passes [`TaskStorage::gc_collectible`].
560    pub fn gc_scan_shard(&self, index: usize, mut on_candidate: impl FnMut(TaskId)) {
561        let shard = self.map.shards()[index].read();
562        for (task_id, task) in shard.iter() {
563            if task.gc_collectible() {
564                on_candidate(*task_id);
565            }
566        }
567    }
568
569    /// Return the set of all known live roots.
570    pub fn gc_scan_roots(&self) -> impl Iterator<Item = TaskId> {
571        let per_shard: Vec<Vec<TaskId>> =
572            parallel::map_collect(&(0..self.shard_count()).collect::<Vec<_>>(), |&index| {
573                let mut roots = Vec::new();
574                let shard = self.map.shards()[index].read();
575                for (task_id, task) in shard.iter() {
576                    if !task_id.is_transient() && task.gc_is_root() {
577                        roots.push(*task_id);
578                    }
579                }
580                roots
581            });
582
583        per_shard.into_iter().flatten()
584    }
585
586    pub fn access_pair_mut(
587        &self,
588        key1: TaskId,
589        key2: TaskId,
590    ) -> (StorageWriteGuard<'_>, StorageWriteGuard<'_>) {
591        let (a, b) = get_disjoint_mut(&self.map, key1, key2, || Box::new(TaskStorage::new()));
592        (
593            StorageWriteGuard {
594                storage: self,
595                inner: a,
596            },
597            StorageWriteGuard {
598                storage: self,
599                inner: b,
600            },
601        )
602    }
603
604    pub fn drop_contents(&self) {
605        drop_contents(&self.map);
606        drop_contents(&self.snapshots);
607    }
608
609    /// Drop the `task_cache` map, freeing its memory.
610    pub(crate) fn drop_task_cache(&self) {
611        drop_contents(&self.task_cache);
612    }
613
614    /// Evict tasks from in-memory storage after a successful snapshot.
615    ///
616    /// Iterates all tasks and applies the eviction level returned by
617    /// `TaskStorage::evictability()`:
618    /// - `Full`: remove from map entirely
619    /// - `DataAndMeta`: drop both data and meta fields, keep task in map
620    /// - `DataOnly`: drop data fields only
621    /// - `MetaOnly`: drop meta fields only
622    /// - `No`: skip
623    ///
624    /// Must be called when NOT in snapshot mode (i.e., after `end_snapshot()`).
625    pub fn evict_after_snapshot(&self, parent_span: Option<Id>) -> EvictionCounts {
626        let span = tracing::trace_span!(
627            parent: parent_span,
628            "evict_after_snapshot",
629            total_task_cache_keys = self.task_cache.len(),
630            total_map_keys = self.map.len(),
631            counts = tracing::field::Empty,
632        )
633        .entered();
634        debug_assert!(
635            !self.snapshot_mode(),
636            "evict_after_snapshot must not be called during snapshot mode"
637        );
638
639        let counts: Vec<EvictionCounts> = parallel::map_collect(self.map.shards(), |shard| {
640            let mut shard = shard.write();
641            let mut evicted = EvictionCounts::default();
642            // task_cache removals that we couldn't perform inline because the target shard
643            // was contended. We defer them until after the map shard lock is released to
644            // avoid a lock cycle with get_or_create_persistent_task, which takes task_cache
645            // before map. Allocated lazily on first conflict.
646            let mut deferred_task_cache_removals: Vec<CachedTaskTypeArc> = Vec::new();
647            // Remove a task type from `task_cache`, deferring on contention. Shared by the
648            // GC-deleted path below and the ordinary key eviction.
649            let remove_from_task_cache =
650                |evicted: &mut EvictionCounts,
651                 deferred: &mut Vec<CachedTaskTypeArc>,
652                 task_type: &CachedTaskTypeArc| {
653                    match try_lock_and_remove(&self.task_cache, task_type.as_ref()) {
654                        TryLockAndRemove::Removed => {
655                            evicted.key_evictions += 1;
656                        }
657                        TryLockAndRemove::NotFound => {
658                            // Generally this should be rare, it more or less implies something
659                            // else is concurrently holding the Arc
660                        }
661                        TryLockAndRemove::WouldBlock => {
662                            // Contention, to avoid a deadlock just defer
663                            deferred.push(task_type.clone());
664                        }
665                    }
666                };
667            shard.retain(|(task_id, task)| {
668                // Transient tasks can not be evicted at all, unless they are fully
669                // delete by the GC.
670                if task_id.is_transient() && !task.flags.deleted() {
671                    evicted.unevictable_reasons[UnevictableReason::Transient.index()] += 1;
672                    return true;
673                }
674                // All GC'd tasks were tombstoned during the snapshot (or are not persisted) so we
675                // can drop them fully now.
676                if task.flags.deleted() {
677                    if let Some(task_type) = task.get_persistent_task_type() {
678                        remove_from_task_cache(
679                            &mut evicted,
680                            &mut deferred_task_cache_removals,
681                            task_type,
682                        );
683                    }
684                    evicted.full += 1;
685                    return false;
686                }
687                let (key_evictability, value_evictability) = task.evictability();
688                match key_evictability {
689                    KeyEvictability::Evictable => {
690                        // The task type is persisted to backing storage (new_task = false),
691                        // so task_cache is a pure perf cache. Remove it now; it will be
692                        // re-populated by task_by_type() on the next cache miss.
693                        let task_type = task.get_persistent_task_type().unwrap();
694                        // Only try to acquire the lock, if we cannot just remove at the end
695                        // Because `get_or_create_task` acquires 'task_cache' then `storage.map` and
696                        // we do the opposite we need to be defensive here.  Attempting here is just
697                        // an optimization to avoid pushing into `deferred_task_cache_removals`
698                        remove_from_task_cache(
699                            &mut evicted,
700                            &mut deferred_task_cache_removals,
701                            task_type,
702                        );
703                    }
704                    KeyEvictability::AlreadyEvicted | KeyEvictability::Unevictable => {}
705                }
706                match value_evictability {
707                    ValueEvictability::Evictable { meta, data } => {
708                        match task.drop_partial(data, meta) {
709                            DropPartialOutcome::Empty => {
710                                evicted.full += 1;
711                                return false;
712                            }
713                            DropPartialOutcome::HasResidue => {
714                                if data && meta {
715                                    evicted.data_and_meta += 1;
716                                } else if data {
717                                    evicted.data_only += 1;
718                                } else {
719                                    debug_assert!(meta);
720                                    evicted.meta_only += 1;
721                                }
722                            }
723                        }
724                    }
725                    ValueEvictability::Unevictable(reason) => {
726                        evicted.unevictable_reasons[reason.index()] += 1;
727                    }
728                }
729                true
730            });
731            // Shrink the shard if it's less than half full, to reclaim slack capacity
732            // after bulk evictions. We already hold the write lock, so this is free
733            // from a locking perspective. TaskId hashing is cheap (it's just an integer).
734            let len = shard.len();
735            if shard.capacity() > len * 2 {
736                shard.shrink_to(len, |(k, _v)| self.map.hasher().hash_one(k));
737            }
738            // Release the map shard lock before draining deferred removals so that a thread
739            // holding a task_cache shard lock and waiting on this map shard can make progress.
740            drop(shard);
741            for task_type in deferred_task_cache_removals {
742                if self.task_cache.remove(task_type.as_ref()).is_some() {
743                    evicted.key_evictions += 1;
744                }
745            }
746            evicted
747        });
748
749        let mut totals = EvictionCounts::default();
750        for evicted in counts {
751            totals += evicted;
752        }
753        // Shrink task_cache only when we evicted more entries than remain — i.e. the map
754        // is less than half full. Rehashing each surviving CachedTaskType isn't free, so
755        // we gate it on meaningful slack. Within that, walk shards in parallel and shrink
756        // each one independently if it is itself less than half full.
757        if totals.key_evictions > self.task_cache.len() {
758            parallel::for_each(self.task_cache.shards(), |shard| {
759                let mut shard = shard.write();
760                let len = shard.len();
761                if shard.capacity() > len * 2 {
762                    shard.shrink_to(len, |(k, _v)| self.task_cache.hasher().hash_one(k));
763                }
764            });
765        }
766        span.record("counts", tracing::field::display(&totals));
767
768        totals
769    }
770}
771
772/// A write guard that still owns its map entry, so the task can be removed under the lock that is
773/// already held.
774///
775/// Use [`Storage::access_entry_mut`] to obtain one. Convert it with [`Self::into_write_guard`] once
776/// removal is no longer a possibility, or call [`Self::discard`] to drop the entry outright.
777pub struct TaskEntryGuard<'a> {
778    storage: &'a Storage,
779    entry: dashmap::mapref::entry::OccupiedEntry<'a, TaskId, Box<TaskStorage>>,
780}
781
782impl<'a> TaskEntryGuard<'a> {
783    /// Removes this task's entry.
784    pub fn discard(self) {
785        self.entry.remove();
786    }
787
788    /// Gives up the ability to remove the entry, yielding an ordinary write guard.
789    pub fn into_write_guard(self) -> StorageWriteGuard<'a> {
790        StorageWriteGuard {
791            storage: self.storage,
792            inner: self.entry.into_ref().into(),
793        }
794    }
795}
796
797impl Deref for TaskEntryGuard<'_> {
798    type Target = TaskStorage;
799    fn deref(&self) -> &Self::Target {
800        self.entry.get()
801    }
802}
803
804impl DerefMut for TaskEntryGuard<'_> {
805    fn deref_mut(&mut self) -> &mut Self::Target {
806        self.entry.get_mut()
807    }
808}
809
810pub struct StorageWriteGuard<'a> {
811    storage: &'a Storage,
812    inner: RefMut<'a, TaskId, Box<TaskStorage>>,
813}
814
815impl StorageWriteGuard<'_> {
816    /// Tracks mutation of this task.
817    #[inline(always)]
818    pub fn track_modification(
819        &mut self,
820        category: SpecificTaskDataCategory,
821        #[allow(unused_variables)] name: &str,
822    ) -> TrackOutcome {
823        debug_assert!(
824            !self.inner.key().is_transient(),
825            "transient task_ids should never be enqueued to be persisted"
826        );
827        self.track_modification_internal(
828            category,
829            #[cfg(feature = "trace_task_modification")]
830            name,
831        )
832    }
833
834    fn track_modification_internal(
835        &mut self,
836        category: SpecificTaskDataCategory,
837        #[cfg(feature = "trace_task_modification")] name: &str,
838    ) -> TrackOutcome {
839        // Transient tasks are never persisted, so tracking modifications is meaningless.
840        // All callers (TaskGuard, initialize_new_task) already
841        // guard against this, but we enforce it here as defense-in-depth.
842        debug_assert!(
843            !self.inner.key().is_transient(),
844            "track_modification called on transient task {:?}",
845            self.inner.key()
846        );
847        let flags = &self.inner.flags;
848        if flags.is_modified_during_snapshot(category) {
849            // We can early return since `end_snapshot` is responsible for reconciling.
850            return TrackOutcome::NoChange;
851        }
852        #[cfg(feature = "trace_task_modification")]
853        let _span = (!modified).then(|| tracing::trace_span!("mark_modified", name).entered());
854        match (self.storage.snapshot_mode(), flags.is_modified(category)) {
855            (false, false) => {
856                // Not in snapshot mode and item is unmodified
857                let bumped = !flags.any_modified();
858                if bumped {
859                    let shard_idx = self.storage.shard_index(self.inner.key());
860                    self.storage.shard_modified_counts[shard_idx].fetch_add(1, Ordering::Relaxed);
861                }
862                self.inner.flags.set_modified(category, true);
863                TrackOutcome::Tracked { category, bumped }
864            }
865            (false, true) => {
866                // Not in snapshot mode and item is already modified
867                // Do nothing
868                TrackOutcome::NoChange
869            }
870            (true, false) => {
871                // In snapshot mode and item is unmodified (so it's not part of the snapshot)
872                // Mark it so it gets re-added as Modified after this snapshot completes.
873                // Insert a None entry into snapshots so end_snapshot discovers this task
874                // and promotes its _during_snapshot flags.
875                let inserted_snapshot = !flags.any_modified_during_snapshot();
876                if inserted_snapshot {
877                    self.storage.snapshots.insert(*self.inner.key(), None);
878                }
879                self.inner
880                    .flags
881                    .set_modified_during_snapshot(category, true);
882                TrackOutcome::TrackedDuringSnapshot {
883                    category,
884                    inserted_snapshot,
885                }
886            }
887            (true, true) => {
888                // In snapshot mode and item is modified (so it's part of the snapshot)
889                // We need to store the original version that is part of the snapshot
890                let inserted_snapshot = !flags.any_modified_during_snapshot();
891                if inserted_snapshot {
892                    // Snapshot all non-transient fields, carrying the modified bits into
893                    // the copy so the iterator knows which categories to persist.
894                    let mut snapshot = self.inner.clone_snapshot();
895                    snapshot.flags.set_data_modified(flags.data_modified());
896                    snapshot.flags.set_meta_modified(flags.meta_modified());
897                    snapshot.flags.set_new_task(flags.new_task());
898                    self.storage
899                        .snapshots
900                        .insert(*self.inner.key(), Some(Box::new(snapshot)));
901                }
902                self.inner
903                    .flags
904                    .set_modified_during_snapshot(category, true);
905                TrackOutcome::TrackedDuringSnapshot {
906                    category,
907                    inserted_snapshot,
908                }
909            }
910        }
911    }
912
913    /// Reverse a [`TrackOutcome`] produced by [`Self::track_modification`] when the mutation it
914    /// guarded changed nothing persistable.
915    ///
916    /// # Correctness
917    ///
918    /// The `outcome` MUST be applied to the **same `StorageWriteGuard`** that produced it, with the
919    /// map shard write lock held continuously in between — i.e. `track_modification`, the mutation,
920    /// and `undo_track_modification` all run within one guard's lifetime. The guard holds its shard
921    /// write lock for its whole lifetime, so this guarantees no other thread observed the tracked
922    /// state, and that `bumped` / `inserted_snapshot` still describe reality (the counter and
923    /// `snapshots` entry are only mutated under that lock). Because those flags record whether
924    /// *this* call created the state, undo never clears a flag, counter, or snapshot entry that a
925    /// prior modification owns.
926    pub fn undo_track_modification(&mut self, outcome: TrackOutcome) {
927        match outcome {
928            TrackOutcome::NoChange => {}
929            TrackOutcome::Tracked { category, bumped } => {
930                self.inner.flags.set_modified(category, false);
931                if bumped {
932                    let shard_idx = self.storage.shard_index(self.inner.key());
933                    self.storage.shard_modified_counts[shard_idx].fetch_sub(1, Ordering::Relaxed);
934                }
935            }
936            TrackOutcome::TrackedDuringSnapshot {
937                category,
938                inserted_snapshot,
939            } => {
940                self.inner
941                    .flags
942                    .set_modified_during_snapshot(category, false);
943                if inserted_snapshot {
944                    self.storage.snapshots.remove(self.inner.key());
945                }
946            }
947        }
948    }
949
950    /// Clears all modified/new flags for a GC-collected task that was **never persisted**
951    /// (`new_task`).
952    pub fn discard_modifications_for_gc_new_task(&mut self) {
953        debug_assert!(
954            !self.storage.snapshot_mode(),
955            "discard_modifications_for_gc_new_task must run before the snapshot starts"
956        );
957        debug_assert!(
958            self.inner.flags.new_task(),
959            "only a never-persisted (new_task) collected task may be discarded this way"
960        );
961        if self.inner.flags.any_modified() {
962            let shard_idx = self.storage.shard_index(self.inner.key());
963            self.storage.shard_modified_counts[shard_idx].fetch_sub(1, Ordering::Relaxed);
964        }
965        self.inner.flags.set_meta_modified(false);
966        self.inner.flags.set_data_modified(false);
967        self.inner.flags.set_new_task(false);
968    }
969}
970
971impl Deref for StorageWriteGuard<'_> {
972    type Target = TaskStorage;
973
974    fn deref(&self) -> &Self::Target {
975        &self.inner
976    }
977}
978
979impl DerefMut for StorageWriteGuard<'_> {
980    fn deref_mut(&mut self) -> &mut Self::Target {
981        &mut self.inner
982    }
983}
984
985/// How big of a buffer to allocate initially. Based on metrics from a large
986/// application this should cover about 98% of values with no resizes.
987const SCRATCH_BUFFER_INITIAL_SIZE: usize = 4096;
988
989/// State machine for a per-thread scratch buffer slot.
990///
991/// Transitions:
992/// - `Uninit` → `Taken` (first take)
993/// - `Available` → `Taken` (subsequent takes)
994/// - `Taken` → `Available` (return)
995///
996/// Any other transition is a bug (e.g. double-take or double-return).
997#[derive(Default)]
998enum ScratchBufferSlot {
999    /// No buffer has been allocated on this thread yet.
1000    #[default]
1001    Uninit,
1002    /// The buffer is currently checked out.
1003    Taken,
1004    /// The buffer is available for reuse.
1005    Available(TurboBincodeBuffer),
1006}
1007
1008pub struct SnapshotGuard<'l> {
1009    storage: &'l Storage,
1010    /// Per-thread scratch buffers for encoding task data. Buffers are taken
1011    /// by `SnapshotShardIter` on creation and returned on drop, allowing reuse
1012    /// across multiple shards processed by the same thread. When the guard is
1013    /// dropped (after all iterators are done), the `ThreadLocal` drops too,
1014    /// freeing all buffers.
1015    scratch_buffers: ThreadLocal<Cell<ScratchBufferSlot>>,
1016}
1017
1018impl<'l> SnapshotGuard<'l> {
1019    fn new(storage: &'l Storage) -> Self {
1020        Self {
1021            storage,
1022            scratch_buffers: ThreadLocal::new(),
1023        }
1024    }
1025
1026    fn take_scratch_buffer(&self) -> TurboBincodeBuffer {
1027        let cell = self.scratch_buffers.get_or_default();
1028        match cell.take() {
1029            ScratchBufferSlot::Available(buf) => {
1030                cell.set(ScratchBufferSlot::Taken);
1031                buf
1032            }
1033            ScratchBufferSlot::Uninit => {
1034                cell.set(ScratchBufferSlot::Taken);
1035                TurboBincodeBuffer::with_capacity(SCRATCH_BUFFER_INITIAL_SIZE)
1036            }
1037            ScratchBufferSlot::Taken => {
1038                panic!("scratch buffer taken twice without being returned");
1039            }
1040        }
1041    }
1042
1043    fn return_scratch_buffer(&self, buffer: TurboBincodeBuffer) {
1044        let cell = self.scratch_buffers.get_or_default();
1045        match cell.take() {
1046            ScratchBufferSlot::Taken => cell.set(ScratchBufferSlot::Available(buffer)),
1047            ScratchBufferSlot::Available(_) => {
1048                panic!("scratch buffer returned without being taken (already available)");
1049            }
1050            ScratchBufferSlot::Uninit => {
1051                panic!("scratch buffer returned without being taken (uninit)");
1052            }
1053        }
1054    }
1055}
1056
1057impl Drop for SnapshotGuard<'_> {
1058    fn drop(&mut self) {
1059        self.storage.end_snapshot();
1060    }
1061}
1062
1063/// The work a single shard's iterator performs, with the snapshot mode encoded in the data rather
1064/// than a runtime flag re-checked per item. Built by `take_snapshot`'s scan.
1065enum ShardWork {
1066    /// Normal snapshot: look each task up in the map while iterating, serialize it, then clear and
1067    /// promote its modified flags so it stays dirty for the next snapshot cycle.
1068    Keep(Vec<TaskId>),
1069    /// Shutdown drain: the scan already erased the unmodified entries and moved the remaining
1070    /// (modified-only) shard table out of the map. The iterator owns that table and drains it
1071    /// directly, freeing each task box as it is serialized. No second map lookup, no flag
1072    /// bookkeeping (the whole map is discarded right after this snapshot).
1073    Drain(hash_table::IntoIter<(TaskId, Box<TaskStorage>)>),
1074}
1075
1076pub struct SnapshotShard<'l, P> {
1077    shard_idx: usize,
1078    work: ShardWork,
1079    storage: &'l Storage,
1080    process: &'l P,
1081    /// Held for its `Drop` impl — ensures snapshot mode ends when all shards are done.
1082    _guard: Arc<SnapshotGuard<'l>>,
1083}
1084
1085impl<'l, P> IntoIterator for SnapshotShard<'l, P>
1086where
1087    P: Fn(TaskId, &TaskStorage, &mut TurboBincodeBuffer) -> SnapshotItem + Sync,
1088{
1089    type Item = SnapshotItem;
1090    type IntoIter = SnapshotShardIter<'l, P>;
1091
1092    fn into_iter(self) -> Self::IntoIter {
1093        let buffer = self._guard.take_scratch_buffer();
1094        SnapshotShardIter {
1095            shard: self,
1096            buffer,
1097        }
1098    }
1099}
1100
1101/// Iterator over a single shard's snapshot items. Holds a thread-local scratch
1102/// buffer for the duration of iteration and returns it on drop.
1103pub struct SnapshotShardIter<'l, P> {
1104    shard: SnapshotShard<'l, P>,
1105    buffer: TurboBincodeBuffer,
1106}
1107
1108impl<'l, P> Iterator for SnapshotShardIter<'l, P>
1109where
1110    P: Fn(TaskId, &TaskStorage, &mut TurboBincodeBuffer) -> SnapshotItem + Sync,
1111{
1112    type Item = SnapshotItem;
1113
1114    fn next(&mut self) -> Option<Self::Item> {
1115        let process = self.shard.process;
1116        let snapshots = &self.shard.storage.snapshots;
1117        let buffer = &mut self.buffer;
1118        let mut serialize_task = |task_id: TaskId, inner: &TaskStorage| {
1119            // If the task was re-modified during snapshot, the snapshots map may
1120            // hold a pre-modification copy we must serialize instead of the live
1121            // data. Remove the entry so end_snapshot doesn't double-promote it;
1122            // we promote manually below.
1123            if inner.flags.any_modified_during_snapshot() {
1124                match snapshots.remove(&task_id) {
1125                    Some((_, Some(snapshot))) => process(task_id, &snapshot, buffer),
1126                    Some((_, None)) | None => process(task_id, inner, buffer),
1127                }
1128            } else {
1129                process(task_id, inner, buffer)
1130            }
1131        };
1132
1133        match &mut self.shard.work {
1134            ShardWork::Keep(modified) => {
1135                let task_id = modified.pop()?;
1136                let mut inner = self.shard.storage.map.get_mut(&task_id).unwrap();
1137                let item = serialize_task(task_id, &inner);
1138                // Clear the modified flags that were captured into the snapshot copy,
1139                // then promote modified_during_snapshot → modified so the task stays
1140                // dirty for the next snapshot cycle.
1141                inner.flags.set_data_modified(false);
1142                inner.flags.set_meta_modified(false);
1143                inner.flags.set_new_task(false);
1144                self.shard
1145                    .storage
1146                    .promote_during_snapshot_flags(&mut inner, self.shard.shard_idx);
1147                Some(item)
1148            }
1149            ShardWork::Drain(entries) => {
1150                // Shutdown only: the scan already moved this shard's modified entries out of the
1151                // map, so we own each `Box<TaskStorage>` here. Serialize from a borrow of the owned
1152                // box and let it drop at the end of this branch — freeing the task's memory as it
1153                // is persisted rather than after the whole batch is written. We skip the flag
1154                // bookkeeping the normal path does, since the entire map is discarded right after
1155                // this snapshot.
1156                let (task_id, inner) = entries.next()?;
1157                Some(serialize_task(task_id, &inner))
1158                // we don't need to update any bits because everything is getting dropped.
1159            }
1160        }
1161    }
1162}
1163
1164impl<P> Drop for SnapshotShardIter<'_, P> {
1165    fn drop(&mut self) {
1166        self.shard
1167            ._guard
1168            .return_scratch_buffer(std::mem::take(&mut self.buffer));
1169    }
1170}
1171
1172#[cfg(test)]
1173mod tests {
1174    use turbo_bincode::TurboBincodeBuffer;
1175    use turbo_tasks::TaskId;
1176
1177    use super::{SpecificTaskDataCategory, Storage, TrackOutcome};
1178    use crate::backing_storage::SnapshotItem;
1179
1180    fn non_transient_task(id: u32) -> TaskId {
1181        // TRANSIENT_TASK_BIT is 0x2000_0000; any id without that bit is non-transient.
1182        TaskId::new(id).expect("id must be non-zero")
1183    }
1184
1185    #[test]
1186    fn new_task_is_pinned_during_construction() {
1187        let storage = Storage::new(2, true);
1188        let task_id = non_transient_task(1);
1189
1190        storage.initialize_new_task(task_id, None);
1191
1192        let task = storage.access_mut(task_id);
1193        assert_eq!(task.gc_transient_ref_count(), 1);
1194        assert!(!task.gc_collectible());
1195    }
1196
1197    /// A process fn that returns a non-empty SnapshotItem so the iterator doesn't
1198    /// silently skip items via the "encoding failed" error path.
1199    fn dummy_process(
1200        task_id: TaskId,
1201        _: &super::TaskStorage,
1202        _: &mut TurboBincodeBuffer,
1203    ) -> SnapshotItem {
1204        SnapshotItem::Put {
1205            task_id,
1206            meta: Some(TurboBincodeBuffer::default()),
1207            data: None,
1208            task_type_hash: None,
1209        }
1210    }
1211
1212    /// Regression test: a task modified before a snapshot and then modified *again* during
1213    /// snapshot iteration must serialize the pre-snapshot state and carry the during-snapshot
1214    /// modification forward to the next cycle.
1215    ///
1216    /// Sequence of events:
1217    /// 1. Task is modified (data_modified = true) → added to shard_modified_counts.
1218    /// 2. `start_snapshot` puts us in snapshot mode.
1219    /// 3. `take_snapshot` scans the shard: task has `any_modified()=true` → goes into the
1220    ///    `modified` list.
1221    /// 4. **Between scan and iteration**: `track_modification` is called on the same category. This
1222    ///    is the `(true, true)` branch: already modified AND in snapshot mode. A snapshot copy of
1223    ///    the pre-second-modification state is stored in `snapshots` as `Some(copy)`, and
1224    ///    `data_modified_during_snapshot` is set.
1225    /// 5. `SnapshotShardIter::next` processes the task from the `modified` list, detects
1226    ///    `any_modified_during_snapshot()=true`, finds the `Some(copy)` in `snapshots`, encodes the
1227    ///    pre-snapshot copy, clears the live modified flags, removes the snapshots entry, and
1228    ///    promotes `data_modified_during_snapshot → data_modified` for the next cycle.
1229    // `end_snapshot` uses `parallel::for_each` which calls `block_in_place` internally,
1230    // requiring a multi-threaded Tokio runtime.
1231    #[tokio::test(flavor = "multi_thread")]
1232    async fn modify_during_snapshot_clears_live_modified_flags() {
1233        let storage = Storage::new(2, true);
1234        let task_id = non_transient_task(1);
1235
1236        // Step 1: modify the task outside snapshot mode (data_modified = true).
1237        {
1238            let mut guard = storage.access_mut(task_id);
1239            let _ = guard.track_modification(SpecificTaskDataCategory::Data, "test");
1240        }
1241
1242        // Step 2: enter snapshot mode.
1243        let (snapshot_guard, has_modifications) = storage.start_snapshot();
1244        assert!(has_modifications);
1245
1246        // Step 3: `take_snapshot` scans the shard. At this point the task has
1247        // `any_modified()=true` and `any_modified_during_snapshot()=false`, so it
1248        // goes into the `modified` list inside the returned `SnapshotShard`.
1249        let shards = storage.take_snapshot(snapshot_guard, &dummy_process, false);
1250
1251        // Step 4: now that the scan is done but before we consume the iterator,
1252        // modify the task again. We're still in snapshot mode, the task is already
1253        // modified → `(true, true)` branch: creates a snapshot copy (carrying the
1254        // modified bits) and sets `data_modified_during_snapshot=true`.
1255        {
1256            let mut guard = storage.access_mut(task_id);
1257            let _ = guard.track_modification(SpecificTaskDataCategory::Data, "test");
1258            // We should have set a snapshot bit
1259            assert!(guard.flags.data_modified_during_snapshot())
1260        }
1261
1262        // Step 5: consume the iterator. The iterator encodes from the pre-snapshot copy,
1263        // clears the live modified flags, removes the snapshots entry, and promotes
1264        // `data_modified_during_snapshot → data_modified` for the next cycle.
1265        let items: Vec<_> = shards
1266            .into_iter()
1267            .flat_map(|shard| shard.into_iter())
1268            .collect();
1269
1270        // The pre-snapshot snapshot copy should have been encoded and returned.
1271        assert_eq!(items.len(), 1);
1272        assert_eq!(items[0].task_id(), task_id);
1273
1274        {
1275            let guard = storage.access_mut(task_id);
1276            // The iterator should have promoted modified_during_snapshot → modified.
1277            assert!(guard.flags.data_modified());
1278        }
1279
1280        // The during-snapshot modification must be reflected in shard_modified_counts so
1281        // the next snapshot cycle picks it up. Verify by starting another snapshot.
1282        let (_guard2, has_modifications) = storage.start_snapshot();
1283        assert!(
1284            has_modifications,
1285            "shard_modified_counts must be non-zero after promoting modified_during_snapshot"
1286        );
1287    }
1288
1289    /// Regression test for the `(true, false)` during-snapshot case: a task modified in one
1290    /// category before a snapshot, then modified in a *different* category during snapshot
1291    /// iteration, must not panic and must carry both modifications forward correctly.
1292    ///
1293    /// Sequence of events:
1294    /// 1. Task meta is modified (meta_modified = true).
1295    /// 2. `start_snapshot` puts us in snapshot mode.
1296    /// 3. `take_snapshot` scans the shard: task goes into the `modified` list.
1297    /// 4. Task data is modified during snapshot → `(true, false)` branch: data was not previously
1298    ///    modified, so `snapshots` gets a `None` entry and `data_modified_during_snapshot` is set.
1299    /// 5. `SnapshotShardIter::next` processes the task: finds `any_modified_during_snapshot()`,
1300    ///    sees `None` in snapshots, encodes from live data (correct — live data for the
1301    ///    unmodified-before-snapshot category is still the pre-snapshot state), clears pre-snapshot
1302    ///    flags, and promotes `data_modified_during_snapshot → data_modified`.
1303    #[tokio::test(flavor = "multi_thread")]
1304    async fn modify_different_category_during_snapshot() {
1305        let storage = Storage::new(2, true);
1306        let task_id = non_transient_task(1);
1307
1308        // Step 1: modify meta only, outside snapshot mode.
1309        {
1310            let mut guard = storage.access_mut(task_id);
1311            let _ = guard.track_modification(SpecificTaskDataCategory::Meta, "test");
1312            assert!(guard.flags.meta_modified());
1313            assert!(!guard.flags.data_modified());
1314        }
1315
1316        // Step 2: enter snapshot mode.
1317        let (snapshot_guard, has_modifications) = storage.start_snapshot();
1318        assert!(has_modifications);
1319
1320        // Step 3: take_snapshot — task goes into modified list (meta_modified = true).
1321        let shards = storage.take_snapshot(snapshot_guard, &dummy_process, false);
1322
1323        // Step 4: modify data during snapshot. The `(true, false)` branch fires:
1324        // data was not previously modified, so snapshots gets a None entry.
1325        {
1326            let mut guard = storage.access_mut(task_id);
1327            let _ = guard.track_modification(SpecificTaskDataCategory::Data, "test");
1328            assert!(guard.flags.data_modified_during_snapshot());
1329            assert!(!guard.flags.meta_modified_during_snapshot());
1330        }
1331
1332        // Step 5: consume the iterator — must not panic.
1333        let items: Vec<_> = shards
1334            .into_iter()
1335            .flat_map(|shard| shard.into_iter())
1336            .collect();
1337
1338        assert_eq!(items.len(), 1);
1339        assert_eq!(items[0].task_id(), task_id);
1340
1341        {
1342            let guard = storage.access_mut(task_id);
1343            // meta_modified was cleared by the iterator (it was the pre-snapshot flag).
1344            assert!(!guard.flags.meta_modified());
1345            // data_modified_during_snapshot was promoted to data_modified.
1346            assert!(guard.flags.data_modified());
1347            assert!(!guard.flags.data_modified_during_snapshot());
1348        }
1349
1350        // Next snapshot cycle must pick up the promoted data_modified.
1351        let (_guard2, has_modifications) = storage.start_snapshot();
1352        assert!(
1353            has_modifications,
1354            "shard_modified_counts must be non-zero after promoting data_modified_during_snapshot"
1355        );
1356    }
1357
1358    /// With `drain_entries = true` (shutdown path), the modified entries are moved out of the map
1359    /// (during the scan) and serialized by the iterator, freeing each task's memory as it is
1360    /// persisted rather than retaining it until the whole snapshot is written. Either way the
1361    /// entry must be gone from the map by the time the snapshot is consumed.
1362    #[tokio::test(flavor = "multi_thread")]
1363    async fn drain_entries_removes_entry_from_map() {
1364        let storage = Storage::new(2, true);
1365        let task_id = non_transient_task(1);
1366
1367        // Modify the task outside snapshot mode so it lands in the modified list.
1368        {
1369            let mut guard = storage.access_mut(task_id);
1370            let _ = guard.track_modification(SpecificTaskDataCategory::Data, "test");
1371        }
1372        assert!(storage.map.get(&task_id).is_some());
1373
1374        let (snapshot_guard, has_modifications) = storage.start_snapshot();
1375        assert!(has_modifications);
1376
1377        // Take the snapshot in drain mode.
1378        let shards = storage.take_snapshot(snapshot_guard, &dummy_process, true);
1379
1380        // Consume the iterator: the task is serialized and then removed from the map.
1381        let items: Vec<_> = shards
1382            .into_iter()
1383            .flat_map(|shard| shard.into_iter())
1384            .collect();
1385
1386        assert_eq!(items.len(), 1);
1387        assert_eq!(items[0].task_id(), task_id);
1388
1389        // The entry must be gone from the map now that it has been persisted.
1390        assert!(
1391            storage.map.get(&task_id).is_none(),
1392            "task entry should be removed from the map after being persisted in drain mode"
1393        );
1394    }
1395
1396    /// In drain mode, fully consuming the iterators should release each drained shard's table
1397    /// allocation entirely (reset-to-empty in `SnapshotShardIter::drop`), not just shrink it.
1398    #[tokio::test(flavor = "multi_thread")]
1399    async fn drain_entries_releases_drained_shards() {
1400        // dashmap requires at least 2 shards.
1401        let storage = Storage::new(2, true);
1402
1403        // Insert and modify enough tasks to grow the shards' tables beyond their minimum.
1404        let task_ids: Vec<_> = (1..=256).map(non_transient_task).collect();
1405        for &task_id in &task_ids {
1406            let mut guard = storage.access_mut(task_id);
1407            let _ = guard.track_modification(SpecificTaskDataCategory::Data, "test");
1408        }
1409        let grown_capacity: usize = storage
1410            .map
1411            .shards()
1412            .iter()
1413            .map(|s| s.read().capacity())
1414            .sum();
1415        assert!(grown_capacity >= task_ids.len());
1416
1417        let (snapshot_guard, has_modifications) = storage.start_snapshot();
1418        assert!(has_modifications);
1419
1420        let shards = storage.take_snapshot(snapshot_guard, &dummy_process, true);
1421        let items: Vec<_> = shards
1422            .into_iter()
1423            .flat_map(|shard| shard.into_iter())
1424            .collect();
1425        assert_eq!(items.len(), task_ids.len());
1426
1427        // Every shard is now empty and its table allocation has been released (capacity 0),
1428        // since the reset swaps in the allocation-free default table.
1429        for shard in storage.map.shards() {
1430            let shard = shard.read();
1431            assert_eq!(shard.len(), 0);
1432            assert_eq!(
1433                shard.capacity(),
1434                0,
1435                "drained shard should have released its table allocation"
1436            );
1437        }
1438    }
1439
1440    /// In drain mode, `take_snapshot`'s scan removes *both* kinds of entry from the map: unmodified
1441    /// entries are erased and freed (never serialized), and the remaining modified-only table is
1442    /// moved out into the shard iterators (to be serialized, then freed as each is consumed). So
1443    /// the map is already empty when `take_snapshot` returns, and only the modified task is
1444    /// yielded.
1445    #[tokio::test(flavor = "multi_thread")]
1446    async fn drain_entries_removes_unmodified_during_take_snapshot() {
1447        let storage = Storage::new(2, true);
1448        let modified_id = non_transient_task(1);
1449        let unmodified_id = non_transient_task(2);
1450
1451        // One modified task (gets serialized) and one unmodified task (e.g. restored from disk but
1452        // never dirtied) that just occupies memory and must not be serialized.
1453        {
1454            let mut guard = storage.access_mut(modified_id);
1455            let _ = guard.track_modification(SpecificTaskDataCategory::Data, "test");
1456        }
1457        // `access_mut` inserts an entry; leaving it without track_modification keeps it unmodified.
1458        let _ = storage.access_mut(unmodified_id);
1459        assert!(storage.map.get(&unmodified_id).is_some());
1460
1461        let (snapshot_guard, has_modifications) = storage.start_snapshot();
1462        assert!(has_modifications);
1463
1464        let shards = storage.take_snapshot(snapshot_guard, &dummy_process, true);
1465
1466        // The scan moved the modified table out and freed the unmodified entry, so both ids are
1467        // already absent from the map before any iterator is consumed.
1468        assert!(
1469            storage.map.get(&unmodified_id).is_none(),
1470            "unmodified entry should be removed during take_snapshot in drain mode"
1471        );
1472        assert!(
1473            storage.map.get(&modified_id).is_none(),
1474            "modified entry should be moved out of the map during take_snapshot in drain mode"
1475        );
1476
1477        // Consuming the iterators yields only the modified task (the unmodified one was never part
1478        // of the snapshot).
1479        let items: Vec<_> = shards
1480            .into_iter()
1481            .flat_map(|shard| shard.into_iter())
1482            .collect();
1483        assert_eq!(items.len(), 1);
1484        assert_eq!(items[0].task_id(), modified_id);
1485    }
1486
1487    #[tokio::test(flavor = "multi_thread")]
1488    async fn undo_non_snapshot_reverses_flag_and_counter() {
1489        let storage = Storage::new(2, true);
1490        let task_id = non_transient_task(1);
1491
1492        {
1493            let mut guard = storage.access_mut(task_id);
1494            let outcome = guard.track_modification(SpecificTaskDataCategory::Data, "test");
1495            assert!(guard.flags.data_modified());
1496            guard.undo_track_modification(outcome);
1497            assert!(!guard.flags.data_modified());
1498            assert!(!guard.flags.any_modified());
1499        }
1500
1501        // Counter is back to zero: the next snapshot sees no modifications.
1502        let (_guard, has_modifications) = storage.start_snapshot();
1503        assert!(
1504            !has_modifications,
1505            "undo must decrement the shard counter so no modifications remain"
1506        );
1507    }
1508
1509    /// A second track on an already-modified category returns `NoChange`; undoing it is a no-op and
1510    /// must NOT clear the real modification recorded by the first track.
1511    #[tokio::test(flavor = "multi_thread")]
1512    async fn undo_nochange_preserves_prior_modification() {
1513        let storage = Storage::new(2, true);
1514        let task_id = non_transient_task(1);
1515
1516        let mut guard = storage.access_mut(task_id);
1517        // First track is the real modification.
1518        let _first = guard.track_modification(SpecificTaskDataCategory::Data, "test");
1519        // Second track on the same category changes nothing.
1520        let second = guard.track_modification(SpecificTaskDataCategory::Data, "test");
1521        assert!(matches!(second, TrackOutcome::NoChange));
1522        // Undoing the no-op must leave the prior modification intact.
1523        guard.undo_track_modification(second);
1524        assert!(
1525            guard.flags.data_modified(),
1526            "undoing a NoChange outcome must not clear a real prior modification"
1527        );
1528    }
1529
1530    /// Undo only reverses the category it tracked: tracking Data then Meta, undoing only the Meta
1531    /// outcome must leave Data modified and the shard counter still non-zero.
1532    #[tokio::test(flavor = "multi_thread")]
1533    async fn undo_only_reverses_its_own_category() {
1534        let storage = Storage::new(2, true);
1535        let task_id = non_transient_task(1);
1536
1537        {
1538            let mut guard = storage.access_mut(task_id);
1539            let _data = guard.track_modification(SpecificTaskDataCategory::Data, "test");
1540            let meta = guard.track_modification(SpecificTaskDataCategory::Meta, "test");
1541            assert!(guard.flags.meta_modified());
1542            guard.undo_track_modification(meta);
1543            assert!(!guard.flags.meta_modified());
1544            assert!(guard.flags.data_modified());
1545        }
1546
1547        // Data is still modified, so the counter is still non-zero.
1548        let (_guard, has_modifications) = storage.start_snapshot();
1549        assert!(has_modifications);
1550    }
1551
1552    /// During-snapshot `(true, false)` arm: a task unmodified-before-snapshot, tracked during
1553    /// snapshot, inserts a `None` marker into `snapshots` and sets the `_during_snapshot` bit.
1554    /// Undo must remove the marker and clear the bit.
1555    #[tokio::test(flavor = "multi_thread")]
1556    async fn undo_during_snapshot_true_false_removes_marker() {
1557        let storage = Storage::new(2, true);
1558        let task_id = non_transient_task(1);
1559        // Insert the task (unmodified) so it exists in the map.
1560        let _ = storage.access_mut(task_id);
1561
1562        let (_snapshot_guard, _) = storage.start_snapshot();
1563
1564        let mut guard = storage.access_mut(task_id);
1565        let outcome = guard.track_modification(SpecificTaskDataCategory::Data, "test");
1566        assert!(matches!(
1567            outcome,
1568            TrackOutcome::TrackedDuringSnapshot {
1569                inserted_snapshot: true,
1570                ..
1571            }
1572        ));
1573        assert!(guard.flags.data_modified_during_snapshot());
1574        assert!(storage.snapshots.get(&task_id).is_some());
1575
1576        guard.undo_track_modification(outcome);
1577        assert!(!guard.flags.data_modified_during_snapshot());
1578        assert!(
1579            storage.snapshots.get(&task_id).is_none(),
1580            "undo must remove the snapshots marker it inserted"
1581        );
1582    }
1583
1584    /// During-snapshot `(true, true)` arm: a task modified-before-snapshot, tracked again during
1585    /// snapshot, stores a pre-mutation copy in `snapshots`. Undo must remove that copy and clear
1586    /// the `_during_snapshot` bit, while leaving the pre-existing `modified` flag intact (it
1587    /// belongs to the snapshot, not to this call).
1588    #[tokio::test(flavor = "multi_thread")]
1589    async fn undo_during_snapshot_true_true_removes_copy_preserves_modified() {
1590        let storage = Storage::new(2, true);
1591        let task_id = non_transient_task(1);
1592
1593        // Modify before snapshot so the category is part of the snapshot.
1594        {
1595            let mut guard = storage.access_mut(task_id);
1596            let _ = guard.track_modification(SpecificTaskDataCategory::Data, "test");
1597        }
1598
1599        let (_snapshot_guard, _) = storage.start_snapshot();
1600
1601        let mut guard = storage.access_mut(task_id);
1602        let outcome = guard.track_modification(SpecificTaskDataCategory::Data, "test");
1603        assert!(matches!(
1604            outcome,
1605            TrackOutcome::TrackedDuringSnapshot {
1606                inserted_snapshot: true,
1607                ..
1608            }
1609        ));
1610        assert!(matches!(
1611            storage.snapshots.get(&task_id).as_deref(),
1612            Some(Some(_))
1613        ));
1614
1615        guard.undo_track_modification(outcome);
1616        assert!(!guard.flags.data_modified_during_snapshot());
1617        assert!(
1618            guard.flags.data_modified(),
1619            "the pre-snapshot modification belongs to the snapshot and must survive undo"
1620        );
1621        assert!(
1622            storage.snapshots.get(&task_id).is_none(),
1623            "undo must remove the pre-mutation copy it inserted"
1624        );
1625    }
1626}