use crate::model::{
artifacts::ArtifactChecksumRecord,
consistency::ApplicationFenceState,
ic_request::IcManagementMethodRecord,
restore_safety::{
RestoreFenceEvidence, RestoreSafetyEvidence, RestoreSafetyLaneRecord,
RestoreSafetyObservation, RestoreSafetyRequest, RestoreSafetyRequirementRecord,
TargetRestoreEvidence,
},
};
use ic_management_canister_types::CanisterStatusType;
use thiserror::Error;
#[derive(Clone, Debug)]
pub struct RestoreSafetyView<'a> {
request: ArtifactChecksumRecord,
requirement: ArtifactChecksumRecord,
observation: &'a RestoreSafetyObservation,
}
impl RestoreSafetyView<'_> {
#[must_use]
pub const fn request(&self) -> &ArtifactChecksumRecord {
&self.request
}
#[must_use]
pub const fn requirement(&self) -> &ArtifactChecksumRecord {
&self.requirement
}
#[must_use]
pub fn targets(&self) -> &[TargetRestoreEvidence] {
self.observation.targets()
}
#[must_use]
pub fn fence(&self) -> Option<&RestoreFenceEvidence> {
match self.observation.safety() {
RestoreSafetyEvidence::ApplicationFenced(fence) => Some(fence),
RestoreSafetyEvidence::NoIrreversibleEffects(_)
| RestoreSafetyEvidence::Unresolved(_) => None,
}
}
#[must_use]
pub const fn evidence(&self) -> &ArtifactChecksumRecord {
self.observation.evidence()
}
#[must_use]
pub const fn remote_observations(&self) -> u32 {
self.observation.remote_observations()
}
}
pub fn validate<'a>(
request: &RestoreSafetyRequest<'_>,
observation: &'a RestoreSafetyObservation,
) -> Result<RestoreSafetyView<'a>, RestoreSafetyError> {
let digest = request.digest();
if *observation.request() != digest {
return Err(RestoreSafetyError::RequestMismatch);
}
let binding = request.binding();
for (field, expected, actual) in [
(
"network",
binding.network(),
observation.context().network(),
),
("caller", binding.caller(), observation.context().caller()),
(
"release",
binding.release(),
observation.context().release(),
),
] {
if actual != expected {
return Err(RestoreSafetyError::ContextMismatch(field));
}
}
if observation.inventory() != request.inventory() {
return Err(RestoreSafetyError::InventoryMismatch);
}
if !observation
.targets()
.iter()
.map(|target| &target.target)
.eq(request.selected_targets())
{
return Err(RestoreSafetyError::SelectionMismatch);
}
let requirement = request.requirement();
if observation.source_plan_intent().hash() != requirement.source_plan_intent() {
return Err(RestoreSafetyError::SourcePlanMismatch);
}
if observation.source_artifacts() != requirement.source_artifacts() {
return Err(RestoreSafetyError::SourceArtifactsMismatch);
}
if observation.remote_observations() > request.max_remote_observations() {
return Err(RestoreSafetyError::ObservationLimitExceeded {
limit: request.max_remote_observations(),
reported: observation.remote_observations(),
});
}
let before_start = request.wire().method() == IcManagementMethodRecord::StartCanister;
let lifecycle_matches = observation.targets().iter().all(|target| {
if before_start && target.target != binding.target() {
target.state != CanisterStatusType::Stopping
} else {
target.state == CanisterStatusType::Stopped
}
});
if !lifecycle_matches {
return Err(RestoreSafetyError::TargetNotStopped);
}
if before_start
&& observation
.targets()
.iter()
.any(|target| target.restored_acceptance.is_none())
{
return Err(RestoreSafetyError::RestoredAcceptanceRequired);
}
validate_safety(requirement, observation.safety(), before_start)?;
Ok(RestoreSafetyView {
request: digest,
requirement: requirement.digest(),
observation,
})
}
fn validate_safety(
requirement: &RestoreSafetyRequirementRecord,
evidence: &RestoreSafetyEvidence,
before_start: bool,
) -> Result<(), RestoreSafetyError> {
match (requirement.safety(), evidence) {
(_, RestoreSafetyEvidence::Unresolved(_)) => {
return Err(RestoreSafetyError::UnresolvedExternalObligations);
}
(
RestoreSafetyLaneRecord::NoIrreversibleEffects,
RestoreSafetyEvidence::NoIrreversibleEffects(_),
) => {}
(
RestoreSafetyLaneRecord::ApplicationFenced,
RestoreSafetyEvidence::ApplicationFenced(fence),
) => {
let original = requirement
.expected_fence()
.ok_or(RestoreSafetyError::FenceMismatch)?;
if fence.binding.identity != original.identity {
return Err(RestoreSafetyError::FenceMismatch);
}
if fence.binding.membership_revision != original.membership_revision {
return Err(RestoreSafetyError::MembershipRevisionMismatch);
}
if fence.binding.external_obligations_revision != original.external_obligations_revision
{
return Err(RestoreSafetyError::ExternalObligationsRevisionMismatch);
}
if fence.state != ApplicationFenceState::Active {
return Err(RestoreSafetyError::FenceNotActive);
}
if before_start && fence.controlled_execution.is_none() {
return Err(RestoreSafetyError::ControlledExecutionRequired);
}
}
_ => return Err(RestoreSafetyError::SafetyLaneMismatch),
}
Ok(())
}
#[derive(Debug, Eq, Error, PartialEq)]
pub enum RestoreSafetyError {
#[error("restore safety request mismatch")]
RequestMismatch,
#[error("restore safety observed {0} mismatch")]
ContextMismatch(&'static str),
#[error("restore safety inventory mismatch")]
InventoryMismatch,
#[error("restore safety selection mismatch")]
SelectionMismatch,
#[error("restore safety source plan mismatch")]
SourcePlanMismatch,
#[error("restore safety source artifacts mismatch")]
SourceArtifactsMismatch,
#[error("restore safety evidence lane mismatch")]
SafetyLaneMismatch,
#[error("restore safety external obligations unresolved")]
UnresolvedExternalObligations,
#[error("restore safety target is not stopped")]
TargetNotStopped,
#[error("restore safety restored-state acceptance required")]
RestoredAcceptanceRequired,
#[error("restore safety fence identity mismatch")]
FenceMismatch,
#[error("restore safety membership revision mismatch")]
MembershipRevisionMismatch,
#[error("restore safety external-obligations revision mismatch")]
ExternalObligationsRevisionMismatch,
#[error("restore safety fence not active")]
FenceNotActive,
#[error("restore safety controlled execution evidence required")]
ControlledExecutionRequired,
#[error("restore safety reports {reported} observations above ceiling {limit}")]
ObservationLimitExceeded {
limit: u32,
reported: u32,
},
}
#[cfg(test)]
mod tests;