use std::sync::Arc;
use super::advanced::{Gpdc, KnnDependence, MixedKnnDependence, OracleCi, SymbolicCmi};
use super::bayes::{BayesFactorCi, PosteriorDependenceCi, PosteriorPredictiveCi};
use super::gsquared::{GSquared, RegressionCi};
use super::pairwise_mv::PairwiseMultivariateCi;
use super::parcorr::PartialCorrelation;
use super::parcorr_variants::{MultivariatePartialCorrelation, RobustPartialCorrelation};
use super::types::ConditionalIndependence;
use crate::error::StatsError;
pub fn ci_from_name(
name: &str,
) -> Result<Arc<dyn ConditionalIndependence + Send + Sync>, StatsError> {
let key = name.trim().to_ascii_lowercase();
let ci: Arc<dyn ConditionalIndependence + Send + Sync> = match key.as_str() {
"parcorr" | "partial_corr" | "partial_correlation" => Arc::new(PartialCorrelation::new()),
"robust_parcorr" | "robust_partial_corr" => Arc::new(RobustPartialCorrelation::new()),
"weighted_parcorr" | "weighted_partial_corr" => {
return Err(StatsError::Backend(
"weighted_parcorr requires observation weights; use \
WeightedPartialCorrelation::new(weights) or causal::resolve_ci(..., Some(weights))"
.into(),
));
}
"multivariate_parcorr" | "multivariate_partial_corr" => {
Arc::new(MultivariatePartialCorrelation::new())
}
"pairwise_multivariate" | "pairwise_mv" => Arc::new(PairwiseMultivariateCi::new()),
"gsquared" | "g_squared" => Arc::new(GSquared::new()),
"regression" => Arc::new(RegressionCi::new()),
"knn_dependence" => Arc::new(KnnDependence::new(5)),
"mixed_knn_dependence" => Arc::new(MixedKnnDependence::new(5)),
"symbolic_cmi" => Arc::new(SymbolicCmi::new()),
"gpdc" => Arc::new(Gpdc::new()),
"oracle" => Arc::new(OracleCi::new([])),
"bayes_factor" | "bayes_factor_ci" => Arc::new(BayesFactorCi::new()),
"posterior_dependence" | "posterior_dependence_ci" => {
Arc::new(PosteriorDependenceCi::new())
}
"posterior_predictive_ci" | "ppc_ci" => Arc::new(PosteriorPredictiveCi::new(199)),
_ => {
return Err(StatsError::Backend(format!("unknown CI test name: {name}")));
}
};
Ok(ci)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn resolves_parcorr_and_rejects_unknown() {
assert!(ci_from_name("parcorr").is_ok());
assert!(ci_from_name("gpdc").is_ok());
assert!(ci_from_name("nope").is_err());
assert!(ci_from_name("weighted_parcorr").is_err());
assert!(ci_from_name("knn_dependence").is_ok());
assert!(ci_from_name("mixed_knn_dependence").is_ok());
assert!(ci_from_name("cmi_knn").is_err());
assert!(ci_from_name("knn_cmi").is_err());
}
}