///|
pub struct ObservationalAnalysis {
  dataset : CausalDataset
  propensity_model : ModelFit
  propensity_scores : Array[Double]
  ipw_estimate : Estimate
  balance : Array[BalanceMetric]
  overlap : OverlapReport
}

///|
pub fn run_observational_analysis(
  dataset : CausalDataset,
) -> ObservationalAnalysis {
  let model = fit_logistic_regression(dataset.covariates, dataset.treatment)
  let propensity = predict_propensity(model, dataset.covariates)
  let estimate = estimate_ipw_ate(
    dataset.outcome,
    dataset.treatment,
    propensity,
  )
  {
    dataset,
    propensity_model: model,
    propensity_scores: propensity,
    ipw_estimate: estimate,
    balance: balance_table(
      dataset.covariates,
      dataset.treatment,
      dataset.feature_names,
    ),
    overlap: overlap_report(propensity, dataset.treatment),
  }
}

///|
pub fn analysis_quality_score(analysis : ObservationalAnalysis) -> Double {
  let overlap_score = analysis.overlap.overlap_fraction
  let balance_score = 1.0 -
    clamp(maximum_absolute_smd(analysis.balance), 0.0, 1.0)
  let convergence_score = if analysis.propensity_model.converged {
    1.0
  } else {
    0.5
  }
  (overlap_score + balance_score + convergence_score) / 3.0
}

///|
pub fn doubly_robust_from_dataset(analysis : ObservationalAnalysis) -> Estimate {
  let treated_model = fit_linear_outcome_model(
    analysis.dataset.covariates,
    analysis.dataset.outcome,
  )
  let control_model = fit_linear_outcome_model(
    analysis.dataset.covariates,
    analysis.dataset.outcome,
  )
  let predicted_treated = predict_outcomes(
    treated_model,
    analysis.dataset.covariates,
  )
  let predicted_control = predict_outcomes(
    control_model,
    analysis.dataset.covariates,
  )
  let estimate = doubly_robust_ate(
    analysis.dataset.outcome,
    analysis.dataset.treatment,
    analysis.propensity_scores,
    predicted_treated,
    predicted_control,
  )
  Estimate::from_standard_error(
    estimate,
    analysis.ipw_estimate.standard_error,
    analysis.dataset.n(),
    analysis.ipw_estimate.effective_sample_size,
    "ATE (doubly robust)",
  )
}

///|
pub fn analysis_summary_vector(
  analysis : ObservationalAnalysis,
) -> Array[Double] {
  [
    analysis.ipw_estimate.estimate,
    analysis.ipw_estimate.standard_error,
    analysis.overlap.overlap_fraction,
    maximum_absolute_smd(analysis.balance),
    analysis.quality_score(),
  ]
}

///|
pub fn ObservationalAnalysis::quality_score(
  self : ObservationalAnalysis,
) -> Double {
  analysis_quality_score(self)
}