1use std::{
2 collections::{BTreeMap, BTreeSet},
3 panic::{AssertUnwindSafe, catch_unwind},
4};
5
6use candid::Principal;
7use pocket_ic::{PocketIc, RejectResponse};
8
9use super::{PocketIcOperationError, transport};
10
11#[derive(Clone, Debug, Eq, PartialEq)]
12struct ControllerSnapshot {
13 snapshot_id: Vec<u8>,
14 sender: Option<Principal>,
15}
16
17#[derive(Clone, Debug, Eq, PartialEq)]
19pub struct ControllerSnapshots(BTreeMap<Principal, ControllerSnapshot>);
20
21#[derive(Clone, Copy, Debug, Eq, PartialEq)]
23pub struct CanisterSnapshotTarget {
24 canister_id: Principal,
25 sender: Option<Principal>,
26}
27
28#[derive(Clone, Debug, Eq, PartialEq)]
30pub struct SnapshotAttemptFailure {
31 sender: Option<Principal>,
32 response: RejectResponse,
33}
34
35#[derive(Clone, Debug, Eq, PartialEq)]
37pub struct SnapshotCleanupFailure {
38 canister_id: Principal,
39 sender: Option<Principal>,
40 response: Option<Box<RejectResponse>>,
41 panic_message: Option<String>,
42}
43
44#[non_exhaustive]
46#[derive(Clone, Copy, Debug, Eq, PartialEq)]
47pub enum SnapshotRestoreFunding {
48 Preserve,
50 TopUpTo {
52 minimum_cycles: u128,
54 },
55}
56
57#[non_exhaustive]
59#[derive(Clone, Debug, Eq, PartialEq)]
60pub enum ControllerSnapshotError {
61 DuplicateCanisterId {
63 canister_id: Principal,
65 },
66 CaptureFailed {
68 canister_id: Principal,
70 attempts: Vec<SnapshotAttemptFailure>,
72 cleanup_failures: Vec<SnapshotCleanupFailure>,
74 },
75 CapturePanicked {
77 canister_id: Principal,
79 source: PocketIcOperationError,
81 cleanup_failures: Vec<SnapshotCleanupFailure>,
83 },
84 RestoreFailed {
86 canister_id: Principal,
88 attempts: Vec<SnapshotAttemptFailure>,
90 },
91 RestorePanicked {
93 canister_id: Principal,
95 source: PocketIcOperationError,
97 },
98}
99
100enum SnapshotCaptureFailure {
101 Rejected(Vec<SnapshotAttemptFailure>),
102 Panicked(PocketIcOperationError),
103}
104
105pub trait PocketIcSnapshotExt {
107 fn capture_snapshots_with_senders<I>(
113 &self,
114 targets: I,
115 ) -> Result<ControllerSnapshots, ControllerSnapshotError>
116 where
117 I: IntoIterator<Item = CanisterSnapshotTarget>;
118
119 fn capture_controller_snapshots<I>(
125 &self,
126 controller_id: Principal,
127 canister_ids: I,
128 ) -> Result<ControllerSnapshots, ControllerSnapshotError>
129 where
130 I: IntoIterator<Item = Principal>;
131
132 fn restore_controller_snapshots(
137 &self,
138 controller_id: Principal,
139 snapshots: &ControllerSnapshots,
140 ) -> Result<(), ControllerSnapshotError>;
141
142 fn restore_controller_snapshots_with_funding(
147 &self,
148 controller_id: Principal,
149 snapshots: &ControllerSnapshots,
150 funding: SnapshotRestoreFunding,
151 ) -> Result<(), ControllerSnapshotError>;
152
153 fn restore_snapshots_with_captured_senders(
159 &self,
160 snapshots: &ControllerSnapshots,
161 ) -> Result<(), ControllerSnapshotError>;
162
163 fn restore_snapshots_with_captured_senders_and_funding(
165 &self,
166 snapshots: &ControllerSnapshots,
167 funding: SnapshotRestoreFunding,
168 ) -> Result<(), ControllerSnapshotError>;
169}
170
171impl PocketIcSnapshotExt for PocketIc {
172 fn capture_snapshots_with_senders<I>(
173 &self,
174 targets: I,
175 ) -> Result<ControllerSnapshots, ControllerSnapshotError>
176 where
177 I: IntoIterator<Item = CanisterSnapshotTarget>,
178 {
179 let targets = ordered_unique_snapshot_targets(targets)?;
180 capture_snapshot_set(
181 self,
182 targets.into_iter().map(|target| {
183 (
184 target.canister_id,
185 std::iter::once(target.sender).collect::<Vec<_>>(),
186 )
187 }),
188 )
189 }
190
191 fn capture_controller_snapshots<I>(
192 &self,
193 controller_id: Principal,
194 canister_ids: I,
195 ) -> Result<ControllerSnapshots, ControllerSnapshotError>
196 where
197 I: IntoIterator<Item = Principal>,
198 {
199 let canister_ids = ordered_unique_canister_ids(canister_ids)?;
200 capture_snapshot_set(
201 self,
202 canister_ids.into_iter().map(|canister_id| {
203 (
204 canister_id,
205 controller_sender_candidates(controller_id, canister_id).to_vec(),
206 )
207 }),
208 )
209 }
210
211 fn restore_controller_snapshots(
212 &self,
213 controller_id: Principal,
214 snapshots: &ControllerSnapshots,
215 ) -> Result<(), ControllerSnapshotError> {
216 self.restore_controller_snapshots_with_funding(
217 controller_id,
218 snapshots,
219 SnapshotRestoreFunding::Preserve,
220 )
221 }
222
223 fn restore_controller_snapshots_with_funding(
224 &self,
225 controller_id: Principal,
226 snapshots: &ControllerSnapshots,
227 funding: SnapshotRestoreFunding,
228 ) -> Result<(), ControllerSnapshotError> {
229 for (canister_id, snapshot_id, sender) in snapshots.iter() {
230 restore_controller_snapshot(
231 self,
232 canister_id,
233 snapshot_id,
234 funding,
235 [
236 sender,
237 if sender.is_some() {
238 None
239 } else {
240 Some(controller_id)
241 },
242 ],
243 )?;
244 }
245 Ok(())
246 }
247
248 fn restore_snapshots_with_captured_senders(
249 &self,
250 snapshots: &ControllerSnapshots,
251 ) -> Result<(), ControllerSnapshotError> {
252 self.restore_snapshots_with_captured_senders_and_funding(
253 snapshots,
254 SnapshotRestoreFunding::Preserve,
255 )
256 }
257
258 fn restore_snapshots_with_captured_senders_and_funding(
259 &self,
260 snapshots: &ControllerSnapshots,
261 funding: SnapshotRestoreFunding,
262 ) -> Result<(), ControllerSnapshotError> {
263 for (canister_id, snapshot_id, sender) in snapshots.iter() {
264 restore_controller_snapshot(
265 self,
266 canister_id,
267 snapshot_id,
268 funding,
269 std::iter::once(sender),
270 )?;
271 }
272 Ok(())
273 }
274}
275
276impl CanisterSnapshotTarget {
277 #[must_use]
279 pub const fn new(canister_id: Principal, sender: Option<Principal>) -> Self {
280 Self {
281 canister_id,
282 sender,
283 }
284 }
285
286 #[must_use]
288 pub const fn canister_id(self) -> Principal {
289 self.canister_id
290 }
291
292 #[must_use]
294 pub const fn sender(self) -> Option<Principal> {
295 self.sender
296 }
297}
298
299fn capture_snapshot_set<I>(
300 pocket_ic: &PocketIc,
301 targets: I,
302) -> Result<ControllerSnapshots, ControllerSnapshotError>
303where
304 I: IntoIterator<Item = (Principal, Vec<Option<Principal>>)>,
305{
306 let mut snapshots = BTreeMap::new();
307 for (canister_id, senders) in targets {
308 match try_take_snapshot(pocket_ic, canister_id, senders) {
309 Ok(snapshot) => {
310 snapshots.insert(canister_id, snapshot);
311 }
312 Err(SnapshotCaptureFailure::Rejected(attempts)) => {
313 let cleanup_failures = cleanup_captured_snapshots(pocket_ic, &snapshots);
314 return Err(ControllerSnapshotError::CaptureFailed {
315 canister_id,
316 attempts,
317 cleanup_failures,
318 });
319 }
320 Err(SnapshotCaptureFailure::Panicked(source)) => {
321 let cleanup_failures = cleanup_captured_snapshots(pocket_ic, &snapshots);
322 return Err(ControllerSnapshotError::CapturePanicked {
323 canister_id,
324 source,
325 cleanup_failures,
326 });
327 }
328 }
329 }
330 Ok(ControllerSnapshots(snapshots))
331}
332
333impl ControllerSnapshots {
334 #[must_use]
336 pub fn len(&self) -> usize {
337 self.0.len()
338 }
339
340 #[must_use]
342 pub fn is_empty(&self) -> bool {
343 self.0.is_empty()
344 }
345
346 pub fn canister_ids(&self) -> impl Iterator<Item = Principal> + '_ {
348 self.0.keys().copied()
349 }
350
351 pub(super) fn iter(&self) -> impl Iterator<Item = (Principal, &[u8], Option<Principal>)> + '_ {
352 self.0.iter().map(|(canister_id, snapshot)| {
353 (
354 *canister_id,
355 snapshot.snapshot_id.as_slice(),
356 snapshot.sender,
357 )
358 })
359 }
360}
361
362impl SnapshotAttemptFailure {
363 #[must_use]
365 pub const fn sender(&self) -> Option<Principal> {
366 self.sender
367 }
368
369 #[must_use]
371 pub const fn response(&self) -> &RejectResponse {
372 &self.response
373 }
374}
375
376impl SnapshotCleanupFailure {
377 #[must_use]
379 pub const fn canister_id(&self) -> Principal {
380 self.canister_id
381 }
382
383 #[must_use]
385 pub const fn sender(&self) -> Option<Principal> {
386 self.sender
387 }
388
389 #[must_use]
391 pub fn response(&self) -> Option<&RejectResponse> {
392 self.response.as_deref()
393 }
394
395 #[must_use]
397 pub fn panic_message(&self) -> Option<&str> {
398 self.panic_message.as_deref()
399 }
400}
401
402impl std::fmt::Display for ControllerSnapshotError {
403 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
404 match self {
405 Self::DuplicateCanisterId { canister_id } => {
406 write!(f, "duplicate canister id in snapshot set: {canister_id}")
407 }
408 Self::CaptureFailed {
409 canister_id,
410 attempts,
411 cleanup_failures,
412 } => write!(
413 f,
414 "failed to capture snapshot for {canister_id} after {} sender attempts; {} partial snapshots could not be cleaned up",
415 attempts.len(),
416 cleanup_failures.len()
417 ),
418 Self::CapturePanicked {
419 canister_id,
420 source,
421 cleanup_failures,
422 } => write!(
423 f,
424 "snapshot capture panicked for {canister_id}: {source}; {} partial snapshots could not be cleaned up",
425 cleanup_failures.len()
426 ),
427 Self::RestoreFailed {
428 canister_id,
429 attempts,
430 } => write!(
431 f,
432 "failed to restore snapshot for {canister_id} after {} sender attempts",
433 attempts.len()
434 ),
435 Self::RestorePanicked {
436 canister_id,
437 source,
438 } => write!(f, "snapshot restore panicked for {canister_id}: {source}"),
439 }
440 }
441}
442
443impl std::error::Error for ControllerSnapshotError {
444 fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
445 match self {
446 Self::CapturePanicked { source, .. } | Self::RestorePanicked { source, .. } => {
447 Some(source)
448 }
449 _ => None,
450 }
451 }
452}
453
454fn ordered_unique_canister_ids<I>(
455 canister_ids: I,
456) -> Result<Vec<Principal>, ControllerSnapshotError>
457where
458 I: IntoIterator<Item = Principal>,
459{
460 let mut unique = BTreeSet::new();
461 for canister_id in canister_ids {
462 if !unique.insert(canister_id) {
463 return Err(ControllerSnapshotError::DuplicateCanisterId { canister_id });
464 }
465 }
466 Ok(unique.into_iter().collect())
467}
468
469fn ordered_unique_snapshot_targets<I>(
470 targets: I,
471) -> Result<Vec<CanisterSnapshotTarget>, ControllerSnapshotError>
472where
473 I: IntoIterator<Item = CanisterSnapshotTarget>,
474{
475 let mut unique = BTreeMap::new();
476 for target in targets {
477 if unique.insert(target.canister_id, target).is_some() {
478 return Err(ControllerSnapshotError::DuplicateCanisterId {
479 canister_id: target.canister_id,
480 });
481 }
482 }
483 Ok(unique.into_values().collect())
484}
485
486fn try_take_snapshot(
487 pocket_ic: &PocketIc,
488 canister_id: Principal,
489 candidates: impl IntoIterator<Item = Option<Principal>>,
490) -> Result<ControllerSnapshot, SnapshotCaptureFailure> {
491 let mut attempts = Vec::new();
492
493 for sender in candidates {
494 let capture = catch_unwind(AssertUnwindSafe(|| {
495 pocket_ic.take_canister_snapshot(canister_id, sender, None)
496 }));
497 match capture {
498 Err(payload) => {
499 return Err(SnapshotCaptureFailure::Panicked(
500 PocketIcOperationError::from_panic(payload.as_ref()),
501 ));
502 }
503 Ok(snapshot) => match snapshot {
504 Ok(snapshot) => {
505 return Ok(ControllerSnapshot {
506 snapshot_id: snapshot.id,
507 sender,
508 });
509 }
510 Err(response) => attempts.push(SnapshotAttemptFailure { sender, response }),
511 },
512 }
513 }
514
515 Err(SnapshotCaptureFailure::Rejected(attempts))
516}
517
518fn cleanup_captured_snapshots(
519 pocket_ic: &PocketIc,
520 snapshots: &BTreeMap<Principal, ControllerSnapshot>,
521) -> Vec<SnapshotCleanupFailure> {
522 let mut failures = Vec::new();
523 for (canister_id, snapshot) in snapshots {
524 let cleanup = catch_unwind(AssertUnwindSafe(|| {
525 pocket_ic.delete_canister_snapshot(
526 *canister_id,
527 snapshot.sender,
528 snapshot.snapshot_id.clone(),
529 )
530 }));
531 match cleanup {
532 Ok(Ok(())) => {}
533 Ok(Err(response)) => failures.push(SnapshotCleanupFailure {
534 canister_id: *canister_id,
535 sender: snapshot.sender,
536 response: Some(Box::new(response)),
537 panic_message: None,
538 }),
539 Err(payload) => failures.push(SnapshotCleanupFailure {
540 canister_id: *canister_id,
541 sender: snapshot.sender,
542 response: None,
543 panic_message: Some(transport::panic_payload_to_string(payload.as_ref())),
544 }),
545 }
546 }
547 failures
548}
549
550fn restore_controller_snapshot(
551 pocket_ic: &PocketIc,
552 canister_id: Principal,
553 snapshot_id: &[u8],
554 funding: SnapshotRestoreFunding,
555 candidates: impl IntoIterator<Item = Option<Principal>>,
556) -> Result<(), ControllerSnapshotError> {
557 let mut attempts = Vec::new();
558
559 for sender in candidates {
560 let restore = catch_unwind(AssertUnwindSafe(|| {
561 apply_snapshot_restore_funding(pocket_ic, canister_id, funding);
562 pocket_ic.load_canister_snapshot(canister_id, sender, snapshot_id.to_vec())
563 }));
564 match restore {
565 Err(payload) => {
566 return Err(ControllerSnapshotError::RestorePanicked {
567 canister_id,
568 source: PocketIcOperationError::from_panic(payload.as_ref()),
569 });
570 }
571 Ok(Ok(())) => return Ok(()),
572 Ok(Err(response)) => attempts.push(SnapshotAttemptFailure { sender, response }),
573 }
574 }
575
576 Err(ControllerSnapshotError::RestoreFailed {
577 canister_id,
578 attempts,
579 })
580}
581
582fn apply_snapshot_restore_funding(
583 pocket_ic: &PocketIc,
584 canister_id: Principal,
585 funding: SnapshotRestoreFunding,
586) {
587 if funding == SnapshotRestoreFunding::Preserve {
588 return;
589 }
590
591 let balance = pocket_ic.cycle_balance(canister_id);
592 let top_up = snapshot_restore_top_up(balance, funding);
593 if top_up > 0 {
594 let _ = pocket_ic.add_cycles(canister_id, top_up);
595 }
596}
597
598const fn snapshot_restore_top_up(balance: u128, funding: SnapshotRestoreFunding) -> u128 {
599 match funding {
600 SnapshotRestoreFunding::Preserve => 0,
601 SnapshotRestoreFunding::TopUpTo { minimum_cycles } => {
602 minimum_cycles.saturating_sub(balance)
603 }
604 }
605}
606
607fn controller_sender_candidates(
608 controller_id: Principal,
609 canister_id: Principal,
610) -> [Option<Principal>; 2] {
611 if canister_id == controller_id {
612 [None, Some(controller_id)]
613 } else {
614 [Some(controller_id), None]
615 }
616}
617
618#[cfg(test)]
619mod tests {
620 use candid::Principal;
621
622 use super::{
623 ControllerSnapshotError, SnapshotRestoreFunding, ordered_unique_canister_ids,
624 snapshot_restore_top_up,
625 };
626
627 #[test]
628 fn duplicate_canister_ids_are_rejected_before_capture() {
629 let canister_id = Principal::from_slice(&[1]);
630 let error = ordered_unique_canister_ids([canister_id, canister_id]).unwrap_err();
631
632 assert_eq!(
633 error,
634 ControllerSnapshotError::DuplicateCanisterId { canister_id }
635 );
636 }
637
638 #[test]
639 fn canister_ids_are_sorted_deterministically() {
640 let first = Principal::from_slice(&[1]);
641 let second = Principal::from_slice(&[2]);
642
643 assert_eq!(
644 ordered_unique_canister_ids([second, first]).unwrap(),
645 vec![first, second]
646 );
647 }
648
649 #[test]
650 fn snapshot_restore_funding_is_explicit() {
651 assert_eq!(
652 snapshot_restore_top_up(10, SnapshotRestoreFunding::Preserve),
653 0
654 );
655 assert_eq!(
656 snapshot_restore_top_up(10, SnapshotRestoreFunding::TopUpTo { minimum_cycles: 25 }),
657 15
658 );
659 assert_eq!(
660 snapshot_restore_top_up(30, SnapshotRestoreFunding::TopUpTo { minimum_cycles: 25 }),
661 0
662 );
663 }
664}