holos_tda/intervention/
model.rs1use std::fmt;
2
3use crate::{
4 EdgeKey, Error, ExplainedDiagram, IntervalGroupId, ProgramTraceArtifact,
5 ProgramTraceDecodeLimits,
6};
7
8#[derive(Debug, Clone, Copy, PartialEq, Eq)]
10pub struct InterventionBudget {
11 pub max_candidates: usize,
13}
14
15impl InterventionBudget {
16 pub fn new(max_candidates: usize) -> Self {
18 Self { max_candidates }
19 }
20}
21
22impl Default for InterventionBudget {
23 fn default() -> Self {
24 Self { max_candidates: 1 }
25 }
26}
27
28#[derive(Debug, Clone, Copy, PartialEq, Eq)]
30#[non_exhaustive]
31pub enum InterventionStatus {
32 Optimal,
34 BoundedGap,
36 BudgetLimited,
38}
39
40#[derive(Debug, Clone, Copy, PartialEq)]
42pub struct EdgeWeightEdit {
43 pub edge: EdgeKey,
45 pub before: f64,
47 pub after: f64,
49}
50
51#[derive(Debug, Clone)]
53pub struct H1Intervention {
54 pub target: IntervalGroupId,
56 pub target_scale: f64,
58 pub status: InterventionStatus,
60 pub lower_bound: f64,
62 pub upper_bound: Option<f64>,
64 pub edits: Vec<EdgeWeightEdit>,
66 pub result: Option<ExplainedDiagram>,
68 pub artifact: Option<InterventionArtifact>,
70}
71#[derive(Debug, Clone, PartialEq, Eq)]
73pub struct InterventionError {
74 message: String,
75}
76
77impl InterventionError {
78 pub(super) fn new(message: impl Into<String>) -> Self {
79 Self {
80 message: message.into(),
81 }
82 }
83
84 pub fn message(&self) -> &str {
86 &self.message
87 }
88}
89
90impl fmt::Display for InterventionError {
91 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
92 write!(f, "intervention artifact: {}", self.message)
93 }
94}
95
96impl std::error::Error for InterventionError {}
97
98impl From<InterventionError> for Error {
99 fn from(error: InterventionError) -> Self {
100 Self::InvalidInput(error.to_string())
101 }
102}
103
104#[derive(Debug, Clone, Copy, PartialEq, Eq)]
106#[non_exhaustive]
107pub struct InterventionDecodeLimits {
108 pub max_bytes: usize,
110 pub max_edits: usize,
112 pub max_trace_bytes: usize,
114 pub trace: ProgramTraceDecodeLimits,
116}
117
118impl Default for InterventionDecodeLimits {
119 fn default() -> Self {
120 Self {
121 max_bytes: 1 << 30,
122 max_edits: 100_000_000,
123 max_trace_bytes: 1 << 30,
124 trace: ProgramTraceDecodeLimits::default(),
125 }
126 }
127}
128
129#[derive(Debug, Clone)]
131pub struct InterventionArtifact {
132 pub(super) target: IntervalGroupId,
133 pub(super) target_scale: f64,
134 pub(super) status: InterventionStatus,
135 pub(super) lower_bound: f64,
136 pub(super) upper_bound: f64,
137 pub(super) edits: Vec<EdgeWeightEdit>,
138 pub(super) trace: ProgramTraceArtifact,
139}
140
141impl InterventionArtifact {
142 pub fn target(&self) -> IntervalGroupId {
144 self.target
145 }
146
147 pub fn target_scale(&self) -> f64 {
149 self.target_scale
150 }
151
152 pub fn status(&self) -> InterventionStatus {
154 self.status
155 }
156
157 pub fn lower_bound(&self) -> f64 {
159 self.lower_bound
160 }
161
162 pub fn upper_bound(&self) -> f64 {
164 self.upper_bound
165 }
166
167 pub fn edits(&self) -> &[EdgeWeightEdit] {
169 &self.edits
170 }
171
172 pub fn trace(&self) -> &ProgramTraceArtifact {
174 &self.trace
175 }
176}
177
178#[derive(Debug, Clone)]
180pub struct VerifiedIntervention {
181 pub status: InterventionStatus,
183 pub target: IntervalGroupId,
185 pub target_scale: f64,
187 pub lower_bound: f64,
189 pub upper_bound: f64,
191 pub edits: Vec<EdgeWeightEdit>,
193 pub result: ExplainedDiagram,
195}