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}