1use std::{
2 collections::BTreeMap,
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 capture_snapshot_set(
180 self,
181 targets
182 .into_iter()
183 .map(|target| (target.canister_id, std::iter::once(target.sender))),
184 )
185 }
186
187 fn capture_controller_snapshots<I>(
188 &self,
189 controller_id: Principal,
190 canister_ids: I,
191 ) -> Result<ControllerSnapshots, ControllerSnapshotError>
192 where
193 I: IntoIterator<Item = Principal>,
194 {
195 capture_snapshot_set(
196 self,
197 canister_ids.into_iter().map(|canister_id| {
198 (
199 canister_id,
200 controller_sender_candidates(controller_id, canister_id),
201 )
202 }),
203 )
204 }
205
206 fn restore_controller_snapshots(
207 &self,
208 controller_id: Principal,
209 snapshots: &ControllerSnapshots,
210 ) -> Result<(), ControllerSnapshotError> {
211 self.restore_controller_snapshots_with_funding(
212 controller_id,
213 snapshots,
214 SnapshotRestoreFunding::Preserve,
215 )
216 }
217
218 fn restore_controller_snapshots_with_funding(
219 &self,
220 controller_id: Principal,
221 snapshots: &ControllerSnapshots,
222 funding: SnapshotRestoreFunding,
223 ) -> Result<(), ControllerSnapshotError> {
224 for (canister_id, snapshot_id, sender) in snapshots.iter() {
225 restore_controller_snapshot(
226 self,
227 canister_id,
228 snapshot_id,
229 funding,
230 [
231 sender,
232 if sender.is_some() {
233 None
234 } else {
235 Some(controller_id)
236 },
237 ],
238 )?;
239 }
240 Ok(())
241 }
242
243 fn restore_snapshots_with_captured_senders(
244 &self,
245 snapshots: &ControllerSnapshots,
246 ) -> Result<(), ControllerSnapshotError> {
247 self.restore_snapshots_with_captured_senders_and_funding(
248 snapshots,
249 SnapshotRestoreFunding::Preserve,
250 )
251 }
252
253 fn restore_snapshots_with_captured_senders_and_funding(
254 &self,
255 snapshots: &ControllerSnapshots,
256 funding: SnapshotRestoreFunding,
257 ) -> Result<(), ControllerSnapshotError> {
258 for (canister_id, snapshot_id, sender) in snapshots.iter() {
259 restore_controller_snapshot(
260 self,
261 canister_id,
262 snapshot_id,
263 funding,
264 std::iter::once(sender),
265 )?;
266 }
267 Ok(())
268 }
269}
270
271impl CanisterSnapshotTarget {
272 #[must_use]
274 pub const fn new(canister_id: Principal, sender: Option<Principal>) -> Self {
275 Self {
276 canister_id,
277 sender,
278 }
279 }
280
281 #[must_use]
283 pub const fn canister_id(self) -> Principal {
284 self.canister_id
285 }
286
287 #[must_use]
289 pub const fn sender(self) -> Option<Principal> {
290 self.sender
291 }
292}
293
294fn capture_snapshot_set<I, S>(
295 pocket_ic: &PocketIc,
296 targets: I,
297) -> Result<ControllerSnapshots, ControllerSnapshotError>
298where
299 I: IntoIterator<Item = (Principal, S)>,
300 S: IntoIterator<Item = Option<Principal>>,
301{
302 let mut ordered_targets = BTreeMap::new();
305 for (canister_id, senders) in targets {
306 if ordered_targets.insert(canister_id, senders).is_some() {
307 return Err(ControllerSnapshotError::DuplicateCanisterId { canister_id });
308 }
309 }
310 let mut snapshots = BTreeMap::new();
311 for (canister_id, senders) in ordered_targets {
312 match try_take_snapshot(pocket_ic, canister_id, senders) {
313 Ok(snapshot) => {
314 snapshots.insert(canister_id, snapshot);
315 }
316 Err(SnapshotCaptureFailure::Rejected(attempts)) => {
317 let cleanup_failures = cleanup_captured_snapshots(pocket_ic, snapshots);
318 return Err(ControllerSnapshotError::CaptureFailed {
319 canister_id,
320 attempts,
321 cleanup_failures,
322 });
323 }
324 Err(SnapshotCaptureFailure::Panicked(source)) => {
325 let cleanup_failures = cleanup_captured_snapshots(pocket_ic, snapshots);
326 return Err(ControllerSnapshotError::CapturePanicked {
327 canister_id,
328 source,
329 cleanup_failures,
330 });
331 }
332 }
333 }
334 Ok(ControllerSnapshots(snapshots))
335}
336
337impl ControllerSnapshots {
338 #[must_use]
340 pub fn len(&self) -> usize {
341 self.0.len()
342 }
343
344 #[must_use]
346 pub fn is_empty(&self) -> bool {
347 self.0.is_empty()
348 }
349
350 pub fn canister_ids(&self) -> impl Iterator<Item = Principal> + '_ {
352 self.0.keys().copied()
353 }
354
355 pub(super) fn iter(&self) -> impl Iterator<Item = (Principal, &[u8], Option<Principal>)> + '_ {
356 self.0.iter().map(|(canister_id, snapshot)| {
357 (
358 *canister_id,
359 snapshot.snapshot_id.as_slice(),
360 snapshot.sender,
361 )
362 })
363 }
364}
365
366impl SnapshotAttemptFailure {
367 #[must_use]
369 pub const fn sender(&self) -> Option<Principal> {
370 self.sender
371 }
372
373 #[must_use]
375 pub const fn response(&self) -> &RejectResponse {
376 &self.response
377 }
378}
379
380impl SnapshotCleanupFailure {
381 #[must_use]
383 pub const fn canister_id(&self) -> Principal {
384 self.canister_id
385 }
386
387 #[must_use]
389 pub const fn sender(&self) -> Option<Principal> {
390 self.sender
391 }
392
393 #[must_use]
395 pub fn response(&self) -> Option<&RejectResponse> {
396 self.response.as_deref()
397 }
398
399 #[must_use]
401 pub fn panic_message(&self) -> Option<&str> {
402 self.panic_message.as_deref()
403 }
404}
405
406impl std::fmt::Display for ControllerSnapshotError {
407 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
408 match self {
409 Self::DuplicateCanisterId { canister_id } => {
410 write!(f, "duplicate canister id in snapshot set: {canister_id}")
411 }
412 Self::CaptureFailed {
413 canister_id,
414 attempts,
415 cleanup_failures,
416 } => write!(
417 f,
418 "failed to capture snapshot for {canister_id} after {} sender attempts; {} partial snapshots could not be cleaned up",
419 attempts.len(),
420 cleanup_failures.len()
421 ),
422 Self::CapturePanicked {
423 canister_id,
424 source,
425 cleanup_failures,
426 } => write!(
427 f,
428 "snapshot capture panicked for {canister_id}: {source}; {} partial snapshots could not be cleaned up",
429 cleanup_failures.len()
430 ),
431 Self::RestoreFailed {
432 canister_id,
433 attempts,
434 } => write!(
435 f,
436 "failed to restore snapshot for {canister_id} after {} sender attempts",
437 attempts.len()
438 ),
439 Self::RestorePanicked {
440 canister_id,
441 source,
442 } => write!(f, "snapshot restore panicked for {canister_id}: {source}"),
443 }
444 }
445}
446
447impl std::error::Error for ControllerSnapshotError {
448 fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
449 match self {
450 Self::CapturePanicked { source, .. } | Self::RestorePanicked { source, .. } => {
451 Some(source)
452 }
453 _ => None,
454 }
455 }
456}
457
458fn try_take_snapshot(
459 pocket_ic: &PocketIc,
460 canister_id: Principal,
461 candidates: impl IntoIterator<Item = Option<Principal>>,
462) -> Result<ControllerSnapshot, SnapshotCaptureFailure> {
463 let mut attempts = Vec::new();
464
465 for sender in candidates {
466 let capture = catch_unwind(AssertUnwindSafe(|| {
467 pocket_ic.take_canister_snapshot(canister_id, sender, None)
468 }));
469 match capture {
470 Err(payload) => {
471 return Err(SnapshotCaptureFailure::Panicked(
472 PocketIcOperationError::from_panic(payload.as_ref()),
473 ));
474 }
475 Ok(snapshot) => match snapshot {
476 Ok(snapshot) => {
477 return Ok(ControllerSnapshot {
478 snapshot_id: snapshot.id,
479 sender,
480 });
481 }
482 Err(response) => attempts.push(SnapshotAttemptFailure { sender, response }),
483 },
484 }
485 }
486
487 Err(SnapshotCaptureFailure::Rejected(attempts))
488}
489
490fn cleanup_captured_snapshots(
491 pocket_ic: &PocketIc,
492 snapshots: BTreeMap<Principal, ControllerSnapshot>,
493) -> Vec<SnapshotCleanupFailure> {
494 let mut failures = Vec::new();
495 for (canister_id, snapshot) in snapshots {
496 let cleanup = catch_unwind(AssertUnwindSafe(|| {
497 pocket_ic.delete_canister_snapshot(canister_id, snapshot.sender, snapshot.snapshot_id)
498 }));
499 match cleanup {
500 Ok(Ok(())) => {}
501 Ok(Err(response)) => failures.push(SnapshotCleanupFailure {
502 canister_id,
503 sender: snapshot.sender,
504 response: Some(Box::new(response)),
505 panic_message: None,
506 }),
507 Err(payload) => failures.push(SnapshotCleanupFailure {
508 canister_id,
509 sender: snapshot.sender,
510 response: None,
511 panic_message: Some(transport::panic_payload_to_string(payload.as_ref())),
512 }),
513 }
514 }
515 failures
516}
517
518fn restore_controller_snapshot(
519 pocket_ic: &PocketIc,
520 canister_id: Principal,
521 snapshot_id: &[u8],
522 funding: SnapshotRestoreFunding,
523 candidates: impl IntoIterator<Item = Option<Principal>>,
524) -> Result<(), ControllerSnapshotError> {
525 let mut attempts = Vec::new();
526
527 for sender in candidates {
528 let restore = catch_unwind(AssertUnwindSafe(|| {
529 apply_snapshot_restore_funding(pocket_ic, canister_id, funding);
530 pocket_ic.load_canister_snapshot(canister_id, sender, snapshot_id.to_vec())
531 }));
532 match restore {
533 Err(payload) => {
534 return Err(ControllerSnapshotError::RestorePanicked {
535 canister_id,
536 source: PocketIcOperationError::from_panic(payload.as_ref()),
537 });
538 }
539 Ok(Ok(())) => return Ok(()),
540 Ok(Err(response)) => attempts.push(SnapshotAttemptFailure { sender, response }),
541 }
542 }
543
544 Err(ControllerSnapshotError::RestoreFailed {
545 canister_id,
546 attempts,
547 })
548}
549
550fn apply_snapshot_restore_funding(
551 pocket_ic: &PocketIc,
552 canister_id: Principal,
553 funding: SnapshotRestoreFunding,
554) {
555 if funding == SnapshotRestoreFunding::Preserve {
556 return;
557 }
558
559 let balance = pocket_ic.cycle_balance(canister_id);
560 let top_up = snapshot_restore_top_up(balance, funding);
561 if top_up > 0 {
562 let _ = pocket_ic.add_cycles(canister_id, top_up);
563 }
564}
565
566const fn snapshot_restore_top_up(balance: u128, funding: SnapshotRestoreFunding) -> u128 {
567 match funding {
568 SnapshotRestoreFunding::Preserve => 0,
569 SnapshotRestoreFunding::TopUpTo { minimum_cycles } => {
570 minimum_cycles.saturating_sub(balance)
571 }
572 }
573}
574
575fn controller_sender_candidates(
576 controller_id: Principal,
577 canister_id: Principal,
578) -> [Option<Principal>; 2] {
579 if canister_id == controller_id {
580 [None, Some(controller_id)]
581 } else {
582 [Some(controller_id), None]
583 }
584}
585
586#[cfg(test)]
587mod tests {
588 use super::{SnapshotRestoreFunding, snapshot_restore_top_up};
589
590 #[test]
591 fn snapshot_restore_funding_is_explicit() {
592 assert_eq!(
593 snapshot_restore_top_up(10, SnapshotRestoreFunding::Preserve),
594 0
595 );
596 assert_eq!(
597 snapshot_restore_top_up(10, SnapshotRestoreFunding::TopUpTo { minimum_cycles: 25 }),
598 15
599 );
600 assert_eq!(
601 snapshot_restore_top_up(30, SnapshotRestoreFunding::TopUpTo { minimum_cycles: 25 }),
602 0
603 );
604 }
605}