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 label = spec.label.clone();
235 let install = catch_unwind(AssertUnwindSafe(|| {
236 pocket_ic.install_canister(canister_id, spec.wasm, spec.init_bytes, spec.install_sender);
237 }));
238 if let Err(payload) = install {
239 let source = PocketIcOperationError::from_panic(payload.as_ref());
240 let context = if let Some(label) = &spec.label {
241 format!("install_canister trapped ({label})")
242 } else {
243 "install_canister trapped".to_string()
244 };
245 let _ = catch_unwind(AssertUnwindSafe(|| {
248 let report = pocket_ic.collect_canister_diagnostics(CanisterDiagnosticsRequest::new(
249 canister_id,
250 diagnostic_sender,
251 diagnostic_sender,
252 ));
253 eprintln!("{context}: {report}");
254 }));
255
256 return Err(CanisterInstallError::new(
257 CanisterInstallPhase::InstallCode,
258 Some(canister_id),
259 label,
260 source,
261 ));
262 }
263
264 Ok(canister_id)
265}
266
267fn is_install_code_rate_limited(response: &RejectResponse) -> bool {
268 response.error_code == ErrorCode::CanisterInstallCodeRateLimited
269}
270
271fn try_install_step<T>(
272 phase: CanisterInstallPhase,
273 canister_id: Option<Principal>,
274 spec: &InstallSpec,
275 operation: impl FnOnce() -> T,
276) -> Result<T, CanisterInstallError> {
277 catch_unwind(AssertUnwindSafe(operation)).map_err(|payload| {
278 CanisterInstallError::new(
279 phase,
280 canister_id,
281 spec.label.clone(),
282 PocketIcOperationError::from_panic(payload.as_ref()),
283 )
284 })
285}
286
287fn retry_install_code_with<T, F, W>(
288 policy: RetryPolicy,
289 mut op: F,
290 mut wait_out_cooldown: W,
291) -> Result<T, RejectResponse>
292where
293 F: FnMut() -> Result<T, RejectResponse>,
294 W: FnMut(),
295{
296 for attempt in 1..=policy.max_attempts() {
297 match op() {
298 Ok(value) => return Ok(value),
299 Err(err) if is_install_code_rate_limited(&err) && attempt < policy.max_attempts() => {
300 wait_out_cooldown();
301 }
302 Err(err) => return Err(err),
303 }
304 }
305
306 unreachable!("RetryPolicy guarantees at least one attempt")
307}
308
309#[cfg(test)]
310mod tests {
311 use std::{cell::Cell, time::Duration};
312
313 use pocket_ic::{ErrorCode, RejectCode, RejectResponse};
314
315 use super::{RetryPolicy, RetryPolicyError, retry_install_code_with};
316
317 fn rejection(error_code: ErrorCode, message: &str) -> RejectResponse {
318 RejectResponse {
319 reject_code: RejectCode::SysTransient,
320 reject_message: message.to_string(),
321 error_code,
322 certified: false,
323 }
324 }
325
326 #[test]
327 fn retry_policy_counts_the_first_attempt() {
328 let attempts = Cell::new(0);
329 let waits = Cell::new(0);
330 let rate_limited = rejection(
331 ErrorCode::CanisterInstallCodeRateLimited,
332 "install-code rate limit",
333 );
334 let result = retry_install_code_with(
335 RetryPolicy::try_new(3, Duration::from_secs(1)).expect("valid retry policy"),
336 || {
337 attempts.set(attempts.get() + 1);
338 Err::<(), _>(rate_limited.clone())
339 },
340 || waits.set(waits.get() + 1),
341 );
342
343 assert_eq!(result, Err(rate_limited));
344 assert_eq!(attempts.get(), 3);
345 assert_eq!(waits.get(), 2);
346 }
347
348 #[test]
349 fn retry_policy_stops_on_non_rate_limit_failure() {
350 let attempts = Cell::new(0);
351 let not_retryable = rejection(ErrorCode::CanisterRejectedMessage, "not retryable");
352 let result = retry_install_code_with(
353 RetryPolicy::try_new(3, Duration::from_secs(1)).expect("valid retry policy"),
354 || {
355 attempts.set(attempts.get() + 1);
356 Err::<(), _>(not_retryable.clone())
357 },
358 || panic!("non-rate-limit failure must not wait"),
359 );
360
361 assert_eq!(result, Err(not_retryable));
362 assert_eq!(attempts.get(), 1);
363 }
364
365 #[test]
366 fn retry_policy_rejects_zero_attempts() {
367 assert_eq!(
368 RetryPolicy::try_new(0, Duration::from_secs(1)),
369 Err(RetryPolicyError::ZeroMaxAttempts)
370 );
371 }
372}