use std::sync::Arc;
use std::time::Instant;
use antecedent_core::{CausalQuery, CausalSchema, ExecutionContext};
use antecedent_data::{TableView, TabularData};
use antecedent_estimate::EstimationWorkspace;
use crate::error::CausalError;
use crate::planner::{CompiledAnalysis, GraphInput, PhysicalExecutionPlan};
use crate::result::CausalAnalysisResult;
use crate::strategy_table::DEFAULT_ESTIMATOR;
use super::builder::{DataInput, RefuteSuite};
use super::execute::CausalAnalysis;
use super::helpers::{project_for_ate_estimate, run_refuters};
use super::stage::{STAGE_VALIDATE, StageClock};
#[derive(Clone, Debug)]
pub struct PreparedAnalysis {
analysis: CausalAnalysis,
plan: PhysicalExecutionPlan,
schema: CausalSchema,
}
impl PreparedAnalysis {
#[must_use]
pub fn schema(&self) -> &CausalSchema {
&self.schema
}
#[must_use]
pub fn plan(&self) -> &PhysicalExecutionPlan {
&self.plan
}
pub fn estimate(
&self,
data: &TabularData,
ctx: &ExecutionContext,
) -> Result<CausalAnalysisResult, CausalError> {
self.ensure_schema_compatible(data)?;
let mut analysis = self.analysis.clone();
analysis.data = DataInput::Tabular(data.clone());
analysis.execute(&CompiledAnalysis::Ready(self.plan.clone()), ctx)
}
pub fn refresh(
&mut self,
data: TabularData,
ctx: &ExecutionContext,
) -> Result<CausalAnalysisResult, CausalError> {
self.ensure_schema_compatible(&data)?;
self.analysis.data = DataInput::Tabular(data);
self.analysis.execute(&CompiledAnalysis::Ready(self.plan.clone()), ctx)
}
pub fn refute(
&self,
prior: &CausalAnalysisResult,
data: &TabularData,
suite: RefuteSuite,
ctx: &ExecutionContext,
) -> Result<CausalAnalysisResult, CausalError> {
self.ensure_schema_compatible(data)?;
let CausalQuery::AverageEffect(query) = &self.analysis.query else {
return Err(CausalError::Unsupported {
message: "PreparedAnalysis::refute requires AverageEffect",
});
};
if prior.treatment != query.treatment || prior.outcome != query.outcome {
return Err(CausalError::Compile {
message: "refute prior result treatment/outcome does not match prepared query"
.into(),
});
}
let estimator = self.analysis.estimator.as_ref().map_or(DEFAULT_ESTIMATOR, |e| e.as_str());
let (data_est, query_est, estimand_est) =
project_for_ate_estimate(data, query, &prior.estimand)?;
let mut clock = StageClock::new();
clock.begin(ctx, STAGE_VALIDATE, 0.8)?;
if ctx.cancellation.is_cancelled() {
return Err(CausalError::Cancelled { stage: STAGE_VALIDATE });
}
let mut workspace = EstimationWorkspace::default();
let started = Instant::now();
let reports = run_refuters(
&data_est,
&estimand_est,
&query_est,
&prior.estimate,
&mut workspace,
None,
ctx,
suite,
estimator,
&self.analysis.custom_validators,
None,
)?;
clock.finish(STAGE_VALIDATE);
let validate_ns = u64::try_from(started.elapsed().as_nanos()).unwrap_or(u64::MAX);
let mut out = prior.clone();
out.refutations = reports;
out.performance.stage_timings_ns.push((Arc::from(STAGE_VALIDATE), validate_ns));
out.performance.wall_time_ns =
Some(out.performance.wall_time_ns.unwrap_or(0).saturating_add(validate_ns));
let suite_label: Arc<str> = match suite {
RefuteSuite::None => Arc::from("none"),
RefuteSuite::Cheap => Arc::from("overlap+evalue"),
RefuteSuite::PlaceboAndRcc => Arc::from("placebo+rcc"),
RefuteSuite::Full => Arc::from("validation.full"),
};
out.diagnostics.push(antecedent_core::Diagnostic::new(
"exec.refute.second_click",
antecedent_core::DiagnosticKind::Execution,
antecedent_core::DiagnosticSeverity::Info,
format!("second-click refute suite={suite_label}"),
));
let _ = clock.wall_time_ns();
Ok(out)
}
fn ensure_schema_compatible(&self, data: &TabularData) -> Result<(), CausalError> {
if data.schema() != &self.schema {
return Err(CausalError::Compile {
message: "prepared analysis refresh requires the same schema \
(variable names, types, and order) as prepare-time data"
.into(),
});
}
Ok(())
}
}
impl CausalAnalysis {
pub fn prepare(&self, ctx: &ExecutionContext) -> Result<PreparedAnalysis, CausalError> {
ensure_prepared_supported(self)?;
let compiled = self.compile(ctx)?;
let CompiledAnalysis::Ready(plan) = compiled else {
return Err(CausalError::Compile {
message: "prepare requires a Ready plan; complete graph review first \
(discovery / incomplete CPDAG/PAG are not session-refreshable)"
.into(),
});
};
let schema = match &self.data {
DataInput::Tabular(data) => data.schema().clone(),
_ => {
return Err(CausalError::Unsupported {
message: "PreparedAnalysis requires tabular data",
});
}
};
Ok(PreparedAnalysis { analysis: self.clone(), plan, schema })
}
}
fn ensure_prepared_supported(analysis: &CausalAnalysis) -> Result<(), CausalError> {
let DataInput::Tabular(_) = &analysis.data else {
return Err(CausalError::Unsupported {
message: "PreparedAnalysis requires tabular data and AverageEffect",
});
};
if !matches!(analysis.query, CausalQuery::AverageEffect(_)) {
return Err(CausalError::Unsupported {
message: "PreparedAnalysis currently supports AverageEffect only",
});
}
if !is_supplied_static_graph(&analysis.graph) {
return Err(CausalError::Unsupported {
message: "PreparedAnalysis requires a supplied static Dag/Cpdag/Pag/Admg \
(discovery graphs stay on one-shot analyze / review)",
});
}
Ok(())
}
fn is_supplied_static_graph(graph: &GraphInput) -> bool {
matches!(
graph,
GraphInput::Static(_) | GraphInput::Cpdag(_) | GraphInput::Pag(_) | GraphInput::Admg(_)
)
}
#[cfg(test)]
mod tests {
use super::is_supplied_static_graph;
use crate::planner::GraphInput;
use antecedent_graph::{Admg, Cpdag, Dag, Pag};
#[test]
fn supplied_static_graphs_only() {
assert!(is_supplied_static_graph(&GraphInput::Static(Dag::with_variables(1))));
assert!(is_supplied_static_graph(&GraphInput::Cpdag(Cpdag::with_variables(1))));
assert!(is_supplied_static_graph(&GraphInput::Pag(Pag::with_variables(1))));
assert!(is_supplied_static_graph(&GraphInput::Admg(Admg::with_variables(1))));
assert!(!is_supplied_static_graph(&GraphInput::DiscoverPc {
alpha: 0.05,
max_cond_size: 3,
fdr: None,
accept_discovered: true,
}));
}
}