use std::sync::Arc;
use crate::{AssumptionSet, Diagnostic, IdentificationStatus, ResponseFunctional};
#[derive(Clone, Copy, Debug, Eq, PartialEq, Hash)]
pub enum SupportStatus {
Supported,
WeakOverlap,
Extrapolative,
OutsideEmpiricalSupport,
}
#[derive(Clone, Debug, PartialEq)]
pub struct SupportRegion {
pub minima: Arc<[f64]>,
pub maxima: Arc<[f64]>,
}
#[derive(Clone, Debug, PartialEq)]
pub struct SupportDiagnostic {
pub id: Arc<str>,
pub values: Arc<[f64]>,
pub detail: Arc<str>,
}
#[derive(Clone, Debug, PartialEq)]
pub struct SupportReport {
pub status: SupportStatus,
pub query_region: SupportRegion,
pub diagnostics: Vec<SupportDiagnostic>,
pub warnings: Vec<Diagnostic>,
}
#[derive(Clone, Debug, PartialEq)]
pub struct IdentifiedSet<T> {
pub lower: T,
pub upper: T,
}
impl IdentifiedSet<f64> {
pub fn try_new(lower: f64, upper: f64) -> Result<Self, &'static str> {
if !lower.is_finite() || !upper.is_finite() || lower > upper {
return Err("identified interval requires finite lower <= upper");
}
Ok(Self { lower, upper })
}
#[must_use]
pub fn intersect(&self, other: &Self) -> Option<Self> {
let lower = self.lower.max(other.lower);
let upper = self.upper.min(other.upper);
(lower <= upper).then_some(Self { lower, upper })
}
}
#[derive(Clone, Debug, PartialEq)]
pub struct ResponseEnvelope {
pub grid: Arc<[f64]>,
pub dimension: usize,
pub lower: Arc<[f64]>,
pub upper: Arc<[f64]>,
}
#[derive(Clone, Debug, PartialEq)]
pub enum ResponseValue {
Scalar(f64),
Surface {
grid: Arc<[f64]>,
dimension: usize,
mean: Arc<[f64]>,
},
Vector(Arc<[f64]>),
Jacobian {
outcomes: usize,
treatments: usize,
values: Arc<[f64]>,
},
Envelope(ResponseEnvelope),
}
#[derive(Clone, Debug, PartialEq)]
pub enum ResponseUncertainty {
None,
Scalar {
standard_error: f64,
level: f64,
lower: f64,
upper: f64,
},
PointwiseBand {
level: f64,
lower: Arc<[f64]>,
upper: Arc<[f64]>,
},
SimultaneousBand {
level: f64,
lower: Arc<[f64]>,
upper: Arc<[f64]>,
replicates: u32,
},
IdentifiedEnvelopeBand {
level: f64,
lower_outer: Arc<[f64]>,
upper_outer: Arc<[f64]>,
},
Posterior {
artifact_id: Arc<str>,
},
}
#[derive(Clone, Debug, PartialEq)]
pub enum ResponseIdentification {
PointIdentified(ResponseValue),
PartiallyIdentified(ResponseValue),
GraphDependent(Vec<(u64, ResponseValue)>),
Unidentified {
certificate: Arc<str>,
},
}
#[derive(Clone, Debug, PartialEq)]
pub struct CausalResponse {
pub estimand: ResponseFunctional,
pub identification_status: IdentificationStatus,
pub estimate: ResponseIdentification,
pub uncertainty: ResponseUncertainty,
pub support: SupportReport,
pub assumptions: AssumptionSet,
pub provenance_id: Arc<str>,
}
#[cfg(test)]
mod tests {
use super::IdentifiedSet;
#[test]
fn identified_set_intersection_never_widens() {
let a = IdentifiedSet::try_new(-1.0, 3.0).unwrap();
let b = IdentifiedSet::try_new(0.0, 2.0).unwrap();
assert_eq!(a.intersect(&b), Some(b));
assert!(a.intersect(&IdentifiedSet::try_new(4.0, 5.0).unwrap()).is_none());
}
}