use std::sync::Arc;
use crate::ids::VariableId;
use crate::intervention::Intervention;
use crate::value::Value;
use super::AverageEffectQuery;
use super::TargetPopulation;
use super::error::QueryError;
#[derive(Clone, Copy, Debug, Eq, PartialEq, Hash)]
pub enum MediationContrast {
Total,
Direct,
Mediated,
NaturalDirect,
NaturalIndirect,
}
#[derive(Clone, Debug, Eq, PartialEq, Hash)]
pub struct MediationQuery {
pub treatment: VariableId,
pub outcome: VariableId,
pub mediators: Arc<[VariableId]>,
pub contrast: MediationContrast,
pub control: Intervention,
pub active: Intervention,
pub target_population: TargetPopulation,
}
impl MediationQuery {
#[must_use]
pub fn binary(
treatment: VariableId,
outcome: VariableId,
mediators: impl Into<Arc<[VariableId]>>,
contrast: MediationContrast,
) -> Self {
Self {
treatment,
outcome,
mediators: mediators.into(),
contrast,
control: Intervention::set(treatment, Value::f64(0.0)),
active: Intervention::set(treatment, Value::f64(1.0)),
target_population: TargetPopulation::AllObserved,
}
}
pub fn validate(&self) -> Result<(), QueryError> {
if self.treatment == self.outcome {
return Err(QueryError::TreatmentEqualsOutcome { id: self.treatment });
}
if self.mediators.is_empty() {
return Err(QueryError::EmptyMediators);
}
if self.mediators.iter().any(|&m| m == self.treatment || m == self.outcome) {
return Err(QueryError::MediatorOverlapsTreatmentOrOutcome);
}
let control_var =
self.control.primary_variable().ok_or(QueryError::AmbiguousInterventionTarget)?;
if control_var != self.treatment {
return Err(QueryError::InterventionVariableMismatch {
expected: self.treatment,
got: control_var,
});
}
let active_var =
self.active.primary_variable().ok_or(QueryError::AmbiguousInterventionTarget)?;
if active_var != self.treatment {
return Err(QueryError::InterventionVariableMismatch {
expected: self.treatment,
got: active_var,
});
}
self.target_population.validate()?;
Ok(())
}
}
#[derive(Clone, Debug, Eq, PartialEq, Hash)]
pub struct ConditionalEffectQuery {
pub inner: AverageEffectQuery,
}
impl ConditionalEffectQuery {
pub fn try_new(inner: AverageEffectQuery) -> Result<Self, QueryError> {
if inner.effect_modifiers.is_empty() {
return Err(QueryError::EmptyEffectModifiers);
}
inner.validate()?;
Ok(Self { inner })
}
pub fn validate(&self) -> Result<(), QueryError> {
if self.inner.effect_modifiers.is_empty() {
return Err(QueryError::EmptyEffectModifiers);
}
self.inner.validate()
}
}