mod method;
pub use method::{IcManagementMethodRecord, IcRequestEffect};
use crate::model::{artifacts::ArtifactChecksumRecord, attempt_journal::OperationBindingRecord};
use candid::Principal;
use ic_management_canister_types::{
CanisterIdRecord, LoadCanisterSnapshotArgs, TakeCanisterSnapshotArgs,
};
use serde::{Deserialize, Deserializer, Serialize, de};
use std::fmt;
use thiserror::Error;
pub const MAX_IC_SNAPSHOT_ID_BYTES: usize = 256;
pub const MAX_IC_ARGUMENT_BYTES: usize = 4096;
pub const MAX_IC_REQUEST_RECORD_BYTES: u64 = 8192;
#[derive(Clone, Debug)]
pub struct IcManagementRequest {
pub method: IcManagementMethodRecord,
pub target: String,
pub snapshot_id: Option<Vec<u8>>,
}
#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
#[serde(try_from = "RequestFields")]
pub struct IcManagementRequestRecord {
version: u16,
method: IcManagementMethodRecord,
target: String,
snapshot_id: Option<Vec<u8>>,
#[serde(skip)]
target_bytes: Vec<u8>,
#[serde(skip)]
arguments: Vec<u8>,
}
#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct RequestFields {
version: u16,
method: IcManagementMethodRecord,
target: String,
#[serde(deserialize_with = "required_snapshot")]
snapshot_id: Option<Vec<u8>>,
}
impl TryFrom<RequestFields> for IcManagementRequestRecord {
type Error = IcRequestError;
fn try_from(fields: RequestFields) -> Result<Self, Self::Error> {
if fields.version != 1 {
return Err(IcRequestError::UnsupportedVersion(fields.version));
}
Self::new(IcManagementRequest {
method: fields.method,
target: fields.target,
snapshot_id: fields.snapshot_id,
})
}
}
impl IcManagementRequestRecord {
pub fn new(request: IcManagementRequest) -> Result<Self, IcRequestError> {
let target = crate::model::principal::canonical_text(&request.target)
.ok_or(IcRequestError::InvalidTarget)?;
let principal = Principal::from_text(&target).map_err(|_| IcRequestError::InvalidTarget)?;
match (request.method, request.snapshot_id.as_deref()) {
(IcManagementMethodRecord::LoadCanisterSnapshot, None) => {
return Err(IcRequestError::SnapshotRequired);
}
(IcManagementMethodRecord::LoadCanisterSnapshot, Some(bytes)) => {
if bytes.is_empty() || bytes.len() > MAX_IC_SNAPSHOT_ID_BYTES {
return Err(IcRequestError::InvalidSnapshotId);
}
}
(_, Some(_)) => return Err(IcRequestError::UnexpectedSnapshot),
(_, None) => {}
}
let arguments = match request.method {
IcManagementMethodRecord::TakeCanisterSnapshot => {
candid::encode_one(TakeCanisterSnapshotArgs {
canister_id: principal,
replace_snapshot: None,
uninstall_code: Some(false),
sender_canister_version: None,
})
}
IcManagementMethodRecord::LoadCanisterSnapshot => {
candid::encode_one(LoadCanisterSnapshotArgs {
canister_id: principal,
snapshot_id: request
.snapshot_id
.clone()
.ok_or(IcRequestError::SnapshotRequired)?,
sender_canister_version: None,
})
}
IcManagementMethodRecord::CanisterStatus
| IcManagementMethodRecord::ListCanisterSnapshots
| IcManagementMethodRecord::StartCanister
| IcManagementMethodRecord::StopCanister => candid::encode_one(CanisterIdRecord {
canister_id: principal,
}),
}
.map_err(|error| IcRequestError::Encoding(error.to_string()))?;
if arguments.len() > MAX_IC_ARGUMENT_BYTES {
return Err(IcRequestError::ArgumentsTooLarge);
}
Ok(Self {
version: 1,
method: request.method,
target,
snapshot_id: request.snapshot_id,
target_bytes: principal.as_slice().to_vec(),
arguments,
})
}
#[must_use]
pub const fn method(&self) -> IcManagementMethodRecord {
self.method
}
#[must_use]
pub fn target(&self) -> &str {
&self.target
}
#[must_use]
pub fn snapshot_id(&self) -> Option<&[u8]> {
self.snapshot_id.as_deref()
}
#[must_use]
pub fn arguments(&self) -> &[u8] {
&self.arguments
}
#[must_use]
pub const fn receiver(&self) -> &'static str {
"aaaaa-aa"
}
#[must_use]
pub fn digest(&self) -> ArtifactChecksumRecord {
management_request_digest(&self.target_bytes, self.method.name(), &self.arguments)
}
pub fn validate_mutation_binding(
&self,
binding: &OperationBindingRecord,
) -> Result<(), IcRequestError> {
self.require_effect(IcRequestEffect::Mutation)?;
self.validate_identity(binding, binding.request())
}
pub fn validate_observation_binding(
&self,
binding: &OperationBindingRecord,
request: &ArtifactChecksumRecord,
) -> Result<(), IcRequestError> {
self.require_effect(IcRequestEffect::Observation)?;
self.validate_identity(binding, request.hash())
}
fn require_effect(&self, expected: IcRequestEffect) -> Result<(), IcRequestError> {
if self.method.effect() != expected {
return Err(IcRequestError::EffectMismatch { expected });
}
Ok(())
}
fn validate_identity(
&self,
binding: &OperationBindingRecord,
expected: &str,
) -> Result<(), IcRequestError> {
if self.target != binding.target() {
return Err(IcRequestError::TargetMismatch);
}
if self.digest().hash() != expected {
return Err(IcRequestError::DigestMismatch);
}
Ok(())
}
}
pub(super) fn management_request_digest(
target_bytes: &[u8],
method: &str,
arguments: &[u8],
) -> ArtifactChecksumRecord {
let mut bytes = b"ic-backup/ic-management-request/v1\0".to_vec();
bytes.push(0); bytes.push(target_bytes.len().to_le_bytes()[0]); bytes.extend_from_slice(target_bytes);
bytes.push(1); bytes.push(method.len().to_le_bytes()[0]); bytes.extend_from_slice(method.as_bytes());
append_argument_length(&mut bytes, arguments.len());
bytes.extend_from_slice(arguments);
ArtifactChecksumRecord::from_bytes(&bytes)
}
#[expect(
clippy::cast_possible_truncation,
reason = "admitted argument byte length is at most 4096"
)]
fn append_argument_length(bytes: &mut Vec<u8>, length: usize) {
bytes.extend_from_slice(&(length as u32).to_be_bytes());
}
fn required_snapshot<'de, D: Deserializer<'de>>(
deserializer: D,
) -> Result<Option<Vec<u8>>, D::Error> {
Ok(Option::<SnapshotBytes>::deserialize(deserializer)?.map(|bytes| bytes.0))
}
struct SnapshotBytes(Vec<u8>);
impl<'de> Deserialize<'de> for SnapshotBytes {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
struct BytesVisitor;
impl<'de> de::Visitor<'de> for BytesVisitor {
type Value = SnapshotBytes;
fn expecting(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("at most 256 exact snapshot bytes")
}
fn visit_seq<A: de::SeqAccess<'de>>(
self,
mut sequence: A,
) -> Result<Self::Value, A::Error> {
let mut bytes = Vec::new();
while bytes.len() < MAX_IC_SNAPSHOT_ID_BYTES {
match sequence.next_element::<u8>()? {
Some(byte) => bytes.push(byte),
None => return Ok(SnapshotBytes(bytes)),
}
}
if sequence.next_element::<de::IgnoredAny>()?.is_some() {
return Err(de::Error::custom(IcRequestError::InvalidSnapshotId));
}
Ok(SnapshotBytes(bytes))
}
}
deserializer.deserialize_seq(BytesVisitor)
}
}
#[derive(Debug, Error, Eq, PartialEq)]
pub enum IcRequestError {
#[error("unsupported IC request version {0}")]
UnsupportedVersion(u16),
#[error("invalid IC request target principal")]
InvalidTarget,
#[error("load snapshot request requires snapshot_id")]
SnapshotRequired,
#[error("snapshot_id is not admitted for this IC request method")]
UnexpectedSnapshot,
#[error("snapshot_id must contain 1..={MAX_IC_SNAPSHOT_ID_BYTES} raw bytes")]
InvalidSnapshotId,
#[error("IC request Candid encoding failed: {0}")]
Encoding(String),
#[error("IC request arguments exceed {MAX_IC_ARGUMENT_BYTES} bytes")]
ArgumentsTooLarge,
#[error("IC request must have effect class {expected:?}")]
EffectMismatch {
expected: IcRequestEffect,
},
#[error("IC request target differs from original binding")]
TargetMismatch,
#[error("IC request digest differs from original binding")]
DigestMismatch,
}
#[cfg(test)]
mod tests;