use antecedent_core::{AverageEffectQuery, ExecutionContext};
use antecedent_data::TabularData;
use antecedent_graph::Dag;
use crate::error::CausalError;
use crate::result::CausalAnalysisResult;
use super::builder::{CausalAnalysisBuilder, RefuteSuite};
use super::execute::CausalAnalysis;
use super::latency::LatencyMode;
#[derive(Clone, Debug)]
pub struct BatchAnalysis {
data: TabularData,
graph: Dag,
bootstrap_replicates: u32,
refute: RefuteSuite,
latency_mode: Option<LatencyMode>,
identifier: Option<String>,
estimator: Option<String>,
}
impl BatchAnalysis {
#[must_use]
pub fn new(data: TabularData, graph: Dag) -> Self {
Self {
data,
graph,
bootstrap_replicates: 50,
refute: RefuteSuite::PlaceboAndRcc,
latency_mode: None,
identifier: None,
estimator: None,
}
}
#[must_use]
pub fn bootstrap_replicates(mut self, n: u32) -> Self {
self.bootstrap_replicates = n;
self
}
#[must_use]
pub fn refute(mut self, suite: RefuteSuite) -> Self {
self.refute = suite;
self
}
#[must_use]
pub fn latency_mode(mut self, mode: LatencyMode) -> Self {
self.latency_mode = Some(mode);
self
}
#[must_use]
pub fn identifier(mut self, id: impl Into<String>) -> Self {
self.identifier = Some(id.into());
self
}
#[must_use]
pub fn estimator(mut self, id: impl Into<String>) -> Self {
self.estimator = Some(id.into());
self
}
pub fn estimate_many(
&self,
queries: &[AverageEffectQuery],
ctx: &ExecutionContext,
) -> Result<Vec<CausalAnalysisResult>, CausalError> {
if queries.is_empty() {
return Err(CausalError::Compile {
message: "batch estimate_many requires at least one query".into(),
});
}
let mut out = Vec::with_capacity(queries.len());
for q in queries {
let mut builder = CausalAnalysisBuilder::new()
.data(self.data.clone())
.graph(self.graph.clone())
.query(q.clone())
.refute(self.refute)
.bootstrap_replicates(self.bootstrap_replicates);
if let Some(mode) = self.latency_mode {
builder = builder.latency_mode(mode);
}
if let Some(id) = self.identifier.as_deref() {
builder = builder.identifier(id);
}
if let Some(est) = self.estimator.as_deref() {
builder = builder.estimator(est);
}
let analysis: CausalAnalysis = builder.build()?;
out.push(analysis.run(ctx)?);
}
Ok(out)
}
}