1use std::panic::{AssertUnwindSafe, catch_unwind};
2use std::time::Duration;
3
4use candid::Principal;
5use pocket_ic::{ErrorCode, PocketIc, RejectResponse};
6
7use super::{
8 CanisterDiagnosticsRequest, CanisterInstallError, CanisterInstallPhase, PocketIcDiagnosticsExt,
9 PocketIcOperationError,
10};
11
12#[non_exhaustive]
14pub struct InstallSpec {
15 pub wasm: Vec<u8>,
17 pub init_bytes: Vec<u8>,
19 pub cycles: u128,
21 pub install_sender: Option<Principal>,
23 pub label: Option<String>,
25}
26
27impl InstallSpec {
28 #[must_use]
30 pub const fn new(wasm: Vec<u8>, init_bytes: Vec<u8>, cycles: u128) -> Self {
31 Self {
32 wasm,
33 init_bytes,
34 cycles,
35 install_sender: None,
36 label: None,
37 }
38 }
39
40 #[must_use]
42 pub const fn install_sender(mut self, sender: Principal) -> Self {
43 self.install_sender = Some(sender);
44 self
45 }
46
47 #[must_use]
49 pub fn label(mut self, label: impl Into<String>) -> Self {
50 self.label = Some(label.into());
51 self
52 }
53}
54
55#[derive(Clone, Copy, Debug, Eq, PartialEq)]
57pub struct RetryPolicy {
58 max_attempts: usize,
59 cooldown: Duration,
60}
61
62#[non_exhaustive]
64#[derive(Clone, Copy, Debug, Eq, PartialEq)]
65pub enum RetryPolicyError {
66 ZeroMaxAttempts,
68}
69
70impl std::fmt::Display for RetryPolicyError {
71 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
72 match self {
73 Self::ZeroMaxAttempts => {
74 formatter.write_str("retry policy requires at least one attempt")
75 }
76 }
77 }
78}
79
80impl std::error::Error for RetryPolicyError {}
81
82impl RetryPolicy {
83 pub const fn try_new(
85 max_attempts: usize,
86 cooldown: Duration,
87 ) -> Result<Self, RetryPolicyError> {
88 if max_attempts == 0 {
89 Err(RetryPolicyError::ZeroMaxAttempts)
90 } else {
91 Ok(Self {
92 max_attempts,
93 cooldown,
94 })
95 }
96 }
97
98 #[must_use]
100 pub const fn max_attempts(self) -> usize {
101 self.max_attempts
102 }
103
104 #[must_use]
106 pub const fn cooldown(self) -> Duration {
107 self.cooldown
108 }
109}
110
111pub trait CanisterInstallExt {
116 #[must_use]
118 fn create_and_install(&self, spec: InstallSpec) -> Principal;
119
120 fn try_create_and_install(&self, spec: InstallSpec) -> Result<Principal, CanisterInstallError>;
125
126 #[must_use]
128 fn create_and_install_many<I>(&self, specs: I) -> Vec<Principal>
129 where
130 I: IntoIterator<Item = InstallSpec>;
131
132 fn try_create_and_install_many<I>(
134 &self,
135 specs: I,
136 ) -> Result<Vec<Principal>, CanisterInstallError>
137 where
138 I: IntoIterator<Item = InstallSpec>;
139
140 fn wait_out_install_code_rate_limit(&self, cooldown: Duration);
142
143 fn retry_install_code<T, F>(&self, policy: RetryPolicy, op: F) -> Result<T, RejectResponse>
149 where
150 F: FnMut() -> Result<T, RejectResponse>;
151}
152
153impl CanisterInstallExt for PocketIc {
154 fn create_and_install(&self, spec: InstallSpec) -> Principal {
156 self.try_create_and_install(spec)
157 .unwrap_or_else(|err| panic!("{err}"))
158 }
159
160 fn try_create_and_install(&self, spec: InstallSpec) -> Result<Principal, CanisterInstallError> {
162 try_create_funded_and_install(self, spec)
163 }
164
165 fn create_and_install_many<I>(&self, specs: I) -> Vec<Principal>
172 where
173 I: IntoIterator<Item = InstallSpec>,
174 {
175 self.try_create_and_install_many(specs)
176 .unwrap_or_else(|err| panic!("{err}"))
177 }
178
179 fn try_create_and_install_many<I>(
186 &self,
187 specs: I,
188 ) -> Result<Vec<Principal>, CanisterInstallError>
189 where
190 I: IntoIterator<Item = InstallSpec>,
191 {
192 specs
193 .into_iter()
194 .map(|spec| self.try_create_and_install(spec))
195 .collect()
196 }
197
198 fn wait_out_install_code_rate_limit(&self, cooldown: Duration) {
200 self.advance_time(cooldown);
201 self.tick();
202 self.tick();
203 }
204
205 fn retry_install_code<T, F>(&self, policy: RetryPolicy, op: F) -> Result<T, RejectResponse>
206 where
207 F: FnMut() -> Result<T, RejectResponse>,
208 {
209 retry_install_code_with(policy, op, || {
210 self.wait_out_install_code_rate_limit(policy.cooldown());
211 })
212 }
213}
214
215fn try_create_funded_and_install(
217 pocket_ic: &PocketIc,
218 spec: InstallSpec,
219) -> Result<Principal, CanisterInstallError> {
220 let canister_id = try_install_step(CanisterInstallPhase::CreateCanister, None, &spec, || {
221 pocket_ic.create_canister()
222 })?;
223 let diagnostic_sender = spec.install_sender.unwrap_or_else(Principal::anonymous);
224 if spec.cycles > 0 {
225 try_install_step(
226 CanisterInstallPhase::AddCycles,
227 Some(canister_id),
228 &spec,
229 || pocket_ic.add_cycles(canister_id, spec.cycles),
230 )?;
231 }
232
233 let install = catch_unwind(AssertUnwindSafe(|| {
235 pocket_ic.install_canister(canister_id, spec.wasm, spec.init_bytes, spec.install_sender);
236 }));
237 if let Err(payload) = install {
238 let source = PocketIcOperationError::from_panic(payload.as_ref());
239 let context = if let Some(label) = &spec.label {
240 format!("install_canister trapped ({label})")
241 } else {
242 "install_canister trapped".to_string()
243 };
244 let _ = catch_unwind(AssertUnwindSafe(|| {
247 let report = pocket_ic.collect_canister_diagnostics(CanisterDiagnosticsRequest::new(
248 canister_id,
249 diagnostic_sender,
250 diagnostic_sender,
251 ));
252 eprintln!("{context}: {report}");
253 }));
254
255 return Err(CanisterInstallError::new(
256 CanisterInstallPhase::InstallCode,
257 Some(canister_id),
258 spec.label,
259 source,
260 ));
261 }
262
263 Ok(canister_id)
264}
265
266fn try_install_step<T>(
267 phase: CanisterInstallPhase,
268 canister_id: Option<Principal>,
269 spec: &InstallSpec,
270 operation: impl FnOnce() -> T,
271) -> Result<T, CanisterInstallError> {
272 catch_unwind(AssertUnwindSafe(operation)).map_err(|payload| {
273 CanisterInstallError::new(
274 phase,
275 canister_id,
276 spec.label.clone(),
277 PocketIcOperationError::from_panic(payload.as_ref()),
278 )
279 })
280}
281
282fn retry_install_code_with<T, F, W>(
283 policy: RetryPolicy,
284 mut op: F,
285 mut wait_out_cooldown: W,
286) -> Result<T, RejectResponse>
287where
288 F: FnMut() -> Result<T, RejectResponse>,
289 W: FnMut(),
290{
291 let mut retries = 1..policy.max_attempts();
292 loop {
293 match op() {
294 Err(err)
295 if err.error_code == ErrorCode::CanisterInstallCodeRateLimited
296 && retries.next().is_some() =>
297 {
298 wait_out_cooldown();
299 }
300 result => return result,
301 }
302 }
303}
304
305#[cfg(test)]
306mod tests {
307 use std::{
308 cell::{Cell, RefCell},
309 time::Duration,
310 };
311
312 use pocket_ic::{ErrorCode, RejectCode, RejectResponse};
313
314 use super::{RetryPolicy, RetryPolicyError, retry_install_code_with};
315
316 fn rejection(error_code: ErrorCode, message: &str) -> RejectResponse {
317 RejectResponse {
318 reject_code: RejectCode::SysTransient,
319 reject_message: message.to_string(),
320 error_code,
321 certified: false,
322 }
323 }
324
325 #[test]
326 fn retry_policy_counts_the_first_attempt() {
327 for max_attempts in [1, 3] {
328 let attempts = Cell::new(0);
329 let waits = Cell::new(0);
330 let result = retry_install_code_with(
331 RetryPolicy::try_new(max_attempts, Duration::from_secs(1))
332 .expect("valid retry policy"),
333 || {
334 attempts.set(attempts.get() + 1);
335 Err::<(), _>(rejection(
336 ErrorCode::CanisterInstallCodeRateLimited,
337 &format!("rate limit on attempt {}", attempts.get()),
338 ))
339 },
340 || waits.set(waits.get() + 1),
341 );
342
343 assert_eq!(
344 result,
345 Err(rejection(
346 ErrorCode::CanisterInstallCodeRateLimited,
347 &format!("rate limit on attempt {max_attempts}"),
348 ))
349 );
350 assert_eq!(attempts.get(), max_attempts);
351 assert_eq!(waits.get(), max_attempts - 1);
352 }
353 }
354
355 #[test]
356 fn retry_policy_waits_only_between_retryable_attempts() {
357 for success_attempt in [1, 3] {
358 let events = RefCell::new(Vec::new());
359 let attempts = Cell::new(0);
360 let result = retry_install_code_with(
361 RetryPolicy::try_new(3, Duration::from_secs(1)).unwrap(),
362 || {
363 events.borrow_mut().push("attempt");
364 attempts.set(attempts.get() + 1);
365 if attempts.get() == success_attempt {
366 Ok(42)
367 } else {
368 Err(rejection(
369 ErrorCode::CanisterInstallCodeRateLimited,
370 "rate limit",
371 ))
372 }
373 },
374 || events.borrow_mut().push("wait"),
375 );
376 assert_eq!(result, Ok(42));
377 assert_eq!(attempts.get(), success_attempt);
378 assert_eq!(
379 events.into_inner(),
380 if success_attempt == 1 {
381 vec!["attempt"]
382 } else {
383 vec!["attempt", "wait", "attempt", "wait", "attempt"]
384 }
385 );
386 }
387 }
388
389 #[test]
390 fn retry_policy_stops_on_non_rate_limit_failure() {
391 let attempts = Cell::new(0);
392 let not_retryable = rejection(ErrorCode::CanisterRejectedMessage, "not retryable");
393 let result = retry_install_code_with(
394 RetryPolicy::try_new(3, Duration::from_secs(1)).expect("valid retry policy"),
395 || {
396 attempts.set(attempts.get() + 1);
397 Err::<(), _>(not_retryable.clone())
398 },
399 || panic!("non-rate-limit failure must not wait"),
400 );
401
402 assert_eq!(result, Err(not_retryable));
403 assert_eq!(attempts.get(), 1);
404 }
405
406 #[test]
407 fn retry_policy_rejects_zero_attempts() {
408 assert_eq!(
409 RetryPolicy::try_new(0, Duration::from_secs(1)),
410 Err(RetryPolicyError::ZeroMaxAttempts)
411 );
412 }
413}