use std::panic::{AssertUnwindSafe, catch_unwind};
use std::time::Duration;
use candid::Principal;
use pocket_ic::{ErrorCode, PocketIc, RejectResponse};
use super::{
CanisterDiagnosticsRequest, CanisterInstallError, CanisterInstallPhase, PocketIcDiagnosticsExt,
PocketIcOperationError,
};
#[non_exhaustive]
pub struct InstallSpec {
pub wasm: Vec<u8>,
pub init_bytes: Vec<u8>,
pub cycles: u128,
pub install_sender: Option<Principal>,
pub label: Option<String>,
}
impl InstallSpec {
#[must_use]
pub const fn new(wasm: Vec<u8>, init_bytes: Vec<u8>, cycles: u128) -> Self {
Self {
wasm,
init_bytes,
cycles,
install_sender: None,
label: None,
}
}
#[must_use]
pub const fn install_sender(mut self, sender: Principal) -> Self {
self.install_sender = Some(sender);
self
}
#[must_use]
pub fn label(mut self, label: impl Into<String>) -> Self {
self.label = Some(label.into());
self
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub struct RetryPolicy {
max_attempts: usize,
cooldown: Duration,
}
#[non_exhaustive]
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum RetryPolicyError {
ZeroMaxAttempts,
}
impl std::fmt::Display for RetryPolicyError {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::ZeroMaxAttempts => {
formatter.write_str("retry policy requires at least one attempt")
}
}
}
}
impl std::error::Error for RetryPolicyError {}
impl RetryPolicy {
pub const fn try_new(
max_attempts: usize,
cooldown: Duration,
) -> Result<Self, RetryPolicyError> {
if max_attempts == 0 {
Err(RetryPolicyError::ZeroMaxAttempts)
} else {
Ok(Self {
max_attempts,
cooldown,
})
}
}
#[must_use]
pub const fn max_attempts(self) -> usize {
self.max_attempts
}
#[must_use]
pub const fn cooldown(self) -> Duration {
self.cooldown
}
}
pub trait CanisterInstallExt {
#[must_use]
fn create_and_install(&self, spec: InstallSpec) -> Principal;
fn try_create_and_install(&self, spec: InstallSpec) -> Result<Principal, CanisterInstallError>;
#[must_use]
fn create_and_install_many<I>(&self, specs: I) -> Vec<Principal>
where
I: IntoIterator<Item = InstallSpec>;
fn try_create_and_install_many<I>(
&self,
specs: I,
) -> Result<Vec<Principal>, CanisterInstallError>
where
I: IntoIterator<Item = InstallSpec>;
fn wait_out_install_code_rate_limit(&self, cooldown: Duration);
fn retry_install_code<T, F>(&self, policy: RetryPolicy, op: F) -> Result<T, RejectResponse>
where
F: FnMut() -> Result<T, RejectResponse>;
}
impl CanisterInstallExt for PocketIc {
fn create_and_install(&self, spec: InstallSpec) -> Principal {
self.try_create_and_install(spec)
.unwrap_or_else(|err| panic!("{err}"))
}
fn try_create_and_install(&self, spec: InstallSpec) -> Result<Principal, CanisterInstallError> {
try_create_funded_and_install(self, spec)
}
fn create_and_install_many<I>(&self, specs: I) -> Vec<Principal>
where
I: IntoIterator<Item = InstallSpec>,
{
self.try_create_and_install_many(specs)
.unwrap_or_else(|err| panic!("{err}"))
}
fn try_create_and_install_many<I>(
&self,
specs: I,
) -> Result<Vec<Principal>, CanisterInstallError>
where
I: IntoIterator<Item = InstallSpec>,
{
specs
.into_iter()
.map(|spec| self.try_create_and_install(spec))
.collect()
}
fn wait_out_install_code_rate_limit(&self, cooldown: Duration) {
self.advance_time(cooldown);
self.tick();
self.tick();
}
fn retry_install_code<T, F>(&self, policy: RetryPolicy, op: F) -> Result<T, RejectResponse>
where
F: FnMut() -> Result<T, RejectResponse>,
{
retry_install_code_with(policy, op, || {
self.wait_out_install_code_rate_limit(policy.cooldown());
})
}
}
fn try_create_funded_and_install(
pocket_ic: &PocketIc,
spec: InstallSpec,
) -> Result<Principal, CanisterInstallError> {
let canister_id = try_install_step(CanisterInstallPhase::CreateCanister, None, &spec, || {
pocket_ic.create_canister()
})?;
let diagnostic_sender = spec.install_sender.unwrap_or_else(Principal::anonymous);
if spec.cycles > 0 {
try_install_step(
CanisterInstallPhase::AddCycles,
Some(canister_id),
&spec,
|| pocket_ic.add_cycles(canister_id, spec.cycles),
)?;
}
let install = catch_unwind(AssertUnwindSafe(|| {
pocket_ic.install_canister(canister_id, spec.wasm, spec.init_bytes, spec.install_sender);
}));
if let Err(payload) = install {
let source = PocketIcOperationError::from_panic(payload.as_ref());
let context = if let Some(label) = &spec.label {
format!("install_canister trapped ({label})")
} else {
"install_canister trapped".to_string()
};
let _ = catch_unwind(AssertUnwindSafe(|| {
let report = pocket_ic.collect_canister_diagnostics(CanisterDiagnosticsRequest::new(
canister_id,
diagnostic_sender,
diagnostic_sender,
));
eprintln!("{context}: {report}");
}));
return Err(CanisterInstallError::new(
CanisterInstallPhase::InstallCode,
Some(canister_id),
spec.label,
source,
));
}
Ok(canister_id)
}
fn try_install_step<T>(
phase: CanisterInstallPhase,
canister_id: Option<Principal>,
spec: &InstallSpec,
operation: impl FnOnce() -> T,
) -> Result<T, CanisterInstallError> {
catch_unwind(AssertUnwindSafe(operation)).map_err(|payload| {
CanisterInstallError::new(
phase,
canister_id,
spec.label.clone(),
PocketIcOperationError::from_panic(payload.as_ref()),
)
})
}
fn retry_install_code_with<T, F, W>(
policy: RetryPolicy,
mut op: F,
mut wait_out_cooldown: W,
) -> Result<T, RejectResponse>
where
F: FnMut() -> Result<T, RejectResponse>,
W: FnMut(),
{
let mut retries = 1..policy.max_attempts();
loop {
match op() {
Err(err)
if err.error_code == ErrorCode::CanisterInstallCodeRateLimited
&& retries.next().is_some() =>
{
wait_out_cooldown();
}
result => return result,
}
}
}
#[cfg(test)]
mod tests {
use std::{
cell::{Cell, RefCell},
time::Duration,
};
use pocket_ic::{ErrorCode, RejectCode, RejectResponse};
use super::{RetryPolicy, RetryPolicyError, retry_install_code_with};
fn rejection(error_code: ErrorCode, message: &str) -> RejectResponse {
RejectResponse {
reject_code: RejectCode::SysTransient,
reject_message: message.to_string(),
error_code,
certified: false,
}
}
#[test]
fn retry_policy_counts_the_first_attempt() {
for max_attempts in [1, 3] {
let attempts = Cell::new(0);
let waits = Cell::new(0);
let result = retry_install_code_with(
RetryPolicy::try_new(max_attempts, Duration::from_secs(1))
.expect("valid retry policy"),
|| {
attempts.set(attempts.get() + 1);
Err::<(), _>(rejection(
ErrorCode::CanisterInstallCodeRateLimited,
&format!("rate limit on attempt {}", attempts.get()),
))
},
|| waits.set(waits.get() + 1),
);
assert_eq!(
result,
Err(rejection(
ErrorCode::CanisterInstallCodeRateLimited,
&format!("rate limit on attempt {max_attempts}"),
))
);
assert_eq!(attempts.get(), max_attempts);
assert_eq!(waits.get(), max_attempts - 1);
}
}
#[test]
fn retry_policy_waits_only_between_retryable_attempts() {
for success_attempt in [1, 3] {
let events = RefCell::new(Vec::new());
let attempts = Cell::new(0);
let result = retry_install_code_with(
RetryPolicy::try_new(3, Duration::from_secs(1)).unwrap(),
|| {
events.borrow_mut().push("attempt");
attempts.set(attempts.get() + 1);
if attempts.get() == success_attempt {
Ok(42)
} else {
Err(rejection(
ErrorCode::CanisterInstallCodeRateLimited,
"rate limit",
))
}
},
|| events.borrow_mut().push("wait"),
);
assert_eq!(result, Ok(42));
assert_eq!(attempts.get(), success_attempt);
assert_eq!(
events.into_inner(),
if success_attempt == 1 {
vec!["attempt"]
} else {
vec!["attempt", "wait", "attempt", "wait", "attempt"]
}
);
}
}
#[test]
fn retry_policy_stops_on_non_rate_limit_failure() {
let attempts = Cell::new(0);
let not_retryable = rejection(ErrorCode::CanisterRejectedMessage, "not retryable");
let result = retry_install_code_with(
RetryPolicy::try_new(3, Duration::from_secs(1)).expect("valid retry policy"),
|| {
attempts.set(attempts.get() + 1);
Err::<(), _>(not_retryable.clone())
},
|| panic!("non-rate-limit failure must not wait"),
);
assert_eq!(result, Err(not_retryable));
assert_eq!(attempts.get(), 1);
}
#[test]
fn retry_policy_rejects_zero_attempts() {
assert_eq!(
RetryPolicy::try_new(0, Duration::from_secs(1)),
Err(RetryPolicyError::ZeroMaxAttempts)
);
}
}