use std::fmt;
use crate::{
EdgeKey, Error, ExplainedDiagram, IntervalGroupId, ProgramTraceArtifact,
ProgramTraceDecodeLimits,
};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct InterventionBudget {
pub max_candidates: usize,
}
impl InterventionBudget {
pub fn new(max_candidates: usize) -> Self {
Self { max_candidates }
}
}
impl Default for InterventionBudget {
fn default() -> Self {
Self { max_candidates: 1 }
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum InterventionStatus {
Optimal,
BoundedGap,
BudgetLimited,
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct EdgeWeightEdit {
pub edge: EdgeKey,
pub before: f64,
pub after: f64,
}
#[derive(Debug, Clone)]
pub struct H1Intervention {
pub target: IntervalGroupId,
pub target_scale: f64,
pub status: InterventionStatus,
pub lower_bound: f64,
pub upper_bound: Option<f64>,
pub edits: Vec<EdgeWeightEdit>,
pub result: Option<ExplainedDiagram>,
pub artifact: Option<InterventionArtifact>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct InterventionError {
message: String,
}
impl InterventionError {
pub(super) fn new(message: impl Into<String>) -> Self {
Self {
message: message.into(),
}
}
pub fn message(&self) -> &str {
&self.message
}
}
impl fmt::Display for InterventionError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "intervention artifact: {}", self.message)
}
}
impl std::error::Error for InterventionError {}
impl From<InterventionError> for Error {
fn from(error: InterventionError) -> Self {
Self::InvalidInput(error.to_string())
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub struct InterventionDecodeLimits {
pub max_bytes: usize,
pub max_edits: usize,
pub max_trace_bytes: usize,
pub trace: ProgramTraceDecodeLimits,
}
impl Default for InterventionDecodeLimits {
fn default() -> Self {
Self {
max_bytes: 1 << 30,
max_edits: 100_000_000,
max_trace_bytes: 1 << 30,
trace: ProgramTraceDecodeLimits::default(),
}
}
}
#[derive(Debug, Clone)]
pub struct InterventionArtifact {
pub(super) target: IntervalGroupId,
pub(super) target_scale: f64,
pub(super) status: InterventionStatus,
pub(super) lower_bound: f64,
pub(super) upper_bound: f64,
pub(super) edits: Vec<EdgeWeightEdit>,
pub(super) trace: ProgramTraceArtifact,
}
impl InterventionArtifact {
pub fn target(&self) -> IntervalGroupId {
self.target
}
pub fn target_scale(&self) -> f64 {
self.target_scale
}
pub fn status(&self) -> InterventionStatus {
self.status
}
pub fn lower_bound(&self) -> f64 {
self.lower_bound
}
pub fn upper_bound(&self) -> f64 {
self.upper_bound
}
pub fn edits(&self) -> &[EdgeWeightEdit] {
&self.edits
}
pub fn trace(&self) -> &ProgramTraceArtifact {
&self.trace
}
}
#[derive(Debug, Clone)]
pub struct VerifiedIntervention {
pub status: InterventionStatus,
pub target: IntervalGroupId,
pub target_scale: f64,
pub lower_bound: f64,
pub upper_bound: f64,
pub edits: Vec<EdgeWeightEdit>,
pub result: ExplainedDiagram,
}