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 is_install_code_rate_limited(response: &RejectResponse) -> bool {
267 response.error_code == ErrorCode::CanisterInstallCodeRateLimited
268}
269
270fn try_install_step<T>(
271 phase: CanisterInstallPhase,
272 canister_id: Option<Principal>,
273 spec: &InstallSpec,
274 operation: impl FnOnce() -> T,
275) -> Result<T, CanisterInstallError> {
276 catch_unwind(AssertUnwindSafe(operation)).map_err(|payload| {
277 CanisterInstallError::new(
278 phase,
279 canister_id,
280 spec.label.clone(),
281 PocketIcOperationError::from_panic(payload.as_ref()),
282 )
283 })
284}
285
286fn retry_install_code_with<T, F, W>(
287 policy: RetryPolicy,
288 mut op: F,
289 mut wait_out_cooldown: W,
290) -> Result<T, RejectResponse>
291where
292 F: FnMut() -> Result<T, RejectResponse>,
293 W: FnMut(),
294{
295 for attempt in 1..=policy.max_attempts() {
296 match op() {
297 Ok(value) => return Ok(value),
298 Err(err) if is_install_code_rate_limited(&err) && attempt < policy.max_attempts() => {
299 wait_out_cooldown();
300 }
301 Err(err) => return Err(err),
302 }
303 }
304
305 unreachable!("RetryPolicy guarantees at least one attempt")
306}
307
308#[cfg(test)]
309mod tests {
310 use std::{cell::Cell, time::Duration};
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 let attempts = Cell::new(0);
328 let waits = Cell::new(0);
329 let rate_limited = rejection(
330 ErrorCode::CanisterInstallCodeRateLimited,
331 "install-code rate limit",
332 );
333 let result = retry_install_code_with(
334 RetryPolicy::try_new(3, Duration::from_secs(1)).expect("valid retry policy"),
335 || {
336 attempts.set(attempts.get() + 1);
337 Err::<(), _>(rate_limited.clone())
338 },
339 || waits.set(waits.get() + 1),
340 );
341
342 assert_eq!(result, Err(rate_limited));
343 assert_eq!(attempts.get(), 3);
344 assert_eq!(waits.get(), 2);
345 }
346
347 #[test]
348 fn retry_policy_stops_on_non_rate_limit_failure() {
349 let attempts = Cell::new(0);
350 let not_retryable = rejection(ErrorCode::CanisterRejectedMessage, "not retryable");
351 let result = retry_install_code_with(
352 RetryPolicy::try_new(3, Duration::from_secs(1)).expect("valid retry policy"),
353 || {
354 attempts.set(attempts.get() + 1);
355 Err::<(), _>(not_retryable.clone())
356 },
357 || panic!("non-rate-limit failure must not wait"),
358 );
359
360 assert_eq!(result, Err(not_retryable));
361 assert_eq!(attempts.get(), 1);
362 }
363
364 #[test]
365 fn retry_policy_rejects_zero_attempts() {
366 assert_eq!(
367 RetryPolicy::try_new(0, Duration::from_secs(1)),
368 Err(RetryPolicyError::ZeroMaxAttempts)
369 );
370 }
371}