use fugue::*;
use rand::{rngs::StdRng, SeedableRng};
#[test]
fn test_api_contract_distribution_interfaces() {
let mut rng = StdRng::seed_from_u64(42);
let normal = Normal::new(0.0, 1.0).expect("Valid Normal parameters");
let sample_n = normal.sample(&mut rng);
let log_prob_n = normal.log_prob(&sample_n);
assert!(sample_n.is_finite());
assert!(log_prob_n.is_finite());
let bernoulli = Bernoulli::new(0.5).expect("Valid Bernoulli parameters");
let sample_b = bernoulli.sample(&mut rng);
let log_prob_b = bernoulli.log_prob(&sample_b);
let _ = sample_b; assert!(log_prob_b.is_finite());
let poisson = Poisson::new(2.0).expect("Valid Poisson parameters");
let sample_p = poisson.sample(&mut rng);
let log_prob_p = poisson.log_prob(&sample_p);
let _ = sample_p; assert!(log_prob_p.is_finite());
let invalid_normal = Normal::new(0.0, -1.0);
assert!(invalid_normal.is_err());
match invalid_normal {
Err(error) => {
assert_eq!(error.category(), ErrorCategory::DistributionValidation);
assert_eq!(error.code(), ErrorCode::InvalidVariance);
}
Ok(_) => panic!("Expected error for negative variance"),
}
}
#[test]
fn test_api_contract_model_composition() {
let mut rng = StdRng::seed_from_u64(42);
let model1 = pure(42.0);
let model2 = sample(addr!("x"), Normal::new(0.0, 1.0).unwrap());
let model3 = observe(addr!("y"), Normal::new(0.0, 1.0).unwrap(), 0.5);
let model4 = factor(-1.0);
let handler1 = runtime::interpreters::PriorHandler {
rng: &mut rng,
trace: runtime::trace::Trace::default(),
};
let (result1, _) = runtime::handler::run(handler1, model1);
assert_eq!(result1, 42.0);
let handler2 = runtime::interpreters::PriorHandler {
rng: &mut rng,
trace: runtime::trace::Trace::default(),
};
let (result2, trace2) = runtime::handler::run(handler2, model2);
assert!(result2.is_finite());
assert!(trace2.get_f64(&addr!("x")).is_some());
let handler3 = runtime::interpreters::PriorHandler {
rng: &mut rng,
trace: runtime::trace::Trace::default(),
};
let (_, trace3) = runtime::handler::run(handler3, model3);
assert!(trace3.log_likelihood.is_finite());
let handler4 = runtime::interpreters::PriorHandler {
rng: &mut rng,
trace: runtime::trace::Trace::default(),
};
let (_, trace4) = runtime::handler::run(handler4, model4);
assert!((trace4.log_factors + 1.0).abs() < 1e-12);
}
#[test]
fn test_public_exports_accessibility() {
let _addr = addr!("test");
let _address = Address::new("test");
let _normal = Normal::new(0.0, 1.0).unwrap();
let _bernoulli = Bernoulli::new(0.5).unwrap();
let _uniform = Uniform::new(0.0, 1.0).unwrap();
let _exponential = Exponential::new(1.0).unwrap();
let _beta = Beta::new(1.0, 1.0).unwrap();
let _gamma = Gamma::new(1.0, 1.0).unwrap();
let _lognormal = LogNormal::new(0.0, 1.0).unwrap();
let _poisson = Poisson::new(1.0).unwrap();
let _binomial = Binomial::new(10, 0.5).unwrap();
let _categorical = Categorical::new(vec![0.5, 0.5]).unwrap();
let _pure_model = pure(1.0);
let _sample_model = sample(addr!("x"), Normal::new(0.0, 1.0).unwrap());
let _observe_model = observe(addr!("y"), Normal::new(0.0, 1.0).unwrap(), 0.0);
let _factor_model = factor(-0.5);
let models = vec![pure(1.0), pure(2.0)];
let _sequence_model = sequence_vec(models);
let _zip_model = zip(pure(1.0), pure(2.0));
let _traverse_model = traverse_vec(vec![1.0, 2.0], |x| pure(x * 2.0));
let _trace = runtime::trace::Trace::default();
let _error_code = ErrorCode::InvalidVariance;
let _error_category = ErrorCategory::DistributionValidation;
let _lse = log_sum_exp(&[-1.0, -2.0]);
let _log1p = log1p_exp(0.5);
let _safe = safe_ln(1.0);
}
#[test]
fn test_api_consistency_error_handling() {
let errors = vec![
Normal::new(0.0, -1.0).unwrap_err(),
Bernoulli::new(1.5).unwrap_err(),
Uniform::new(1.0, 0.0).unwrap_err(),
Exponential::new(-1.0).unwrap_err(),
];
for error in errors {
assert_eq!(error.category(), ErrorCategory::DistributionValidation);
let message = format!("{}", error);
assert!(!message.is_empty());
assert!(message.len() > 10);
let code = error.code();
assert!(matches!(
code,
ErrorCode::InvalidMean
| ErrorCode::InvalidVariance
| ErrorCode::InvalidProbability
| ErrorCode::InvalidRange
| ErrorCode::InvalidShape
| ErrorCode::InvalidRate
));
}
}
#[test]
fn test_api_consistency_naming_conventions() {
let _n1 = Normal::new(0.0, 1.0);
let _n2 = Bernoulli::new(0.5);
let _n3 = Uniform::new(0.0, 1.0);
let _n4 = Exponential::new(1.0);
let _n5 = Beta::new(1.0, 1.0);
let _n6 = Gamma::new(1.0, 1.0);
let _n7 = LogNormal::new(0.0, 1.0);
let _n8 = Poisson::new(1.0);
let _n9 = Binomial::new(10, 0.5);
let _n10 = Categorical::new(vec![0.5, 0.5]);
let _pure = pure(1.0);
let _sample = sample(addr!("x"), Normal::new(0.0, 1.0).unwrap());
let _observe = observe(addr!("y"), Normal::new(0.0, 1.0).unwrap(), 0.0);
let _factor = factor(-0.5);
let trace = runtime::trace::Trace::default();
let _f64_opt = trace.get_f64(&addr!("x"));
let _bool_opt = trace.get_bool(&addr!("x"));
let _u64_opt = trace.get_u64(&addr!("x"));
let _usize_opt = trace.get_usize(&addr!("x"));
let _f64_res = trace.get_f64_result(&addr!("x"));
let _bool_res = trace.get_bool_result(&addr!("x"));
let _u64_res = trace.get_u64_result(&addr!("x"));
let _usize_res = trace.get_usize_result(&addr!("x"));
}
#[test]
fn test_ergonomics_type_inference() {
let mut rng = StdRng::seed_from_u64(42);
let model1 = sample(addr!("x"), Normal::new(0.0, 1.0).unwrap());
let handler1 = runtime::interpreters::PriorHandler {
rng: &mut rng,
trace: runtime::trace::Trace::default(),
};
let (result1, _) = runtime::handler::run(handler1, model1);
let _: f64 = result1;
let model2 = sample(addr!("b"), Bernoulli::new(0.5).unwrap());
let handler2 = runtime::interpreters::PriorHandler {
rng: &mut rng,
trace: runtime::trace::Trace::default(),
};
let (result2, _) = runtime::handler::run(handler2, model2);
let _: bool = result2;
let model3 = sample(addr!("p"), Poisson::new(2.0).unwrap());
let handler3 = runtime::interpreters::PriorHandler {
rng: &mut rng,
trace: runtime::trace::Trace::default(),
};
let (result3, _) = runtime::handler::run(handler3, model3);
let _: u64 = result3;
let composed = sample(addr!("x"), Normal::new(0.0, 1.0).unwrap())
.map(|x| x * 2.0)
.bind(|x| pure(x + 1.0));
let handler4 = runtime::interpreters::PriorHandler {
rng: &mut rng,
trace: runtime::trace::Trace::default(),
};
let (result4, _) = runtime::handler::run(handler4, composed);
let _: f64 = result4;
}
#[test]
fn test_ergonomics_common_patterns() {
let mut rng = StdRng::seed_from_u64(42);
let bayesian_model = sample(addr!("mu"), Normal::new(0.0, 1.0).unwrap())
.bind(|mu| observe(addr!("y"), Normal::new(mu, 0.5).unwrap(), 1.2).map(move |_| mu));
let handler = runtime::interpreters::PriorHandler {
rng: &mut rng,
trace: runtime::trace::Trace::default(),
};
let (mu_sample, trace) = runtime::handler::run(handler, bayesian_model);
assert!(mu_sample.is_finite());
assert!(trace.log_likelihood.is_finite());
let multi_param = sample(addr!("a"), Normal::new(0.0, 1.0).unwrap())
.bind(|a| sample(addr!("b"), Normal::new(a, 0.5).unwrap()).map(move |b| (a, b)));
let handler2 = runtime::interpreters::PriorHandler {
rng: &mut rng,
trace: runtime::trace::Trace::default(),
};
let ((a_val, b_val), trace2) = runtime::handler::run(handler2, multi_param);
assert!(a_val.is_finite());
assert!(b_val.is_finite());
assert_eq!(a_val, trace2.get_f64(&addr!("a")).unwrap());
assert_eq!(b_val, trace2.get_f64(&addr!("b")).unwrap());
let vectorized = sequence_vec(vec![
sample(addr!("x1"), Normal::new(0.0, 1.0).unwrap()),
sample(addr!("x2"), Normal::new(1.0, 1.0).unwrap()),
sample(addr!("x3"), Normal::new(2.0, 1.0).unwrap()),
]);
let handler3 = runtime::interpreters::PriorHandler {
rng: &mut rng,
trace: runtime::trace::Trace::default(),
};
let (vec_results, trace3) = runtime::handler::run(handler3, vectorized);
assert_eq!(vec_results.len(), 3);
assert!(vec_results.iter().all(|x| x.is_finite()));
assert!(trace3.get_f64(&addr!("x1")).is_some());
assert!(trace3.get_f64(&addr!("x2")).is_some());
assert!(trace3.get_f64(&addr!("x3")).is_some());
}
#[test]
fn test_api_contract_inference_algorithms() {
let mut rng = StdRng::seed_from_u64(42);
let model_fn = || sample(addr!("theta"), Normal::new(0.0, 1.0).unwrap());
let mcmc_samples = adaptive_mcmc_chain(&mut rng, model_fn, 50, 10);
assert_eq!(mcmc_samples.len(), 50);
assert!(mcmc_samples.iter().all(|(theta, _)| theta.is_finite()));
let smc_model_fn = || sample(addr!("mu"), Normal::new(0.0, 1.0).unwrap());
let smc_config = SMCConfig {
resampling_method: ResamplingMethod::Systematic,
ess_threshold: 0.5,
rejuvenation_steps: 0,
};
let particles = adaptive_smc(&mut rng, 20, smc_model_fn, smc_config);
assert_eq!(particles.len(), 20);
assert!(particles.iter().all(|p| p.log_weight.is_finite()));
let abc_model_fn = || sample(addr!("param"), Normal::new(0.0, 1.0).unwrap());
let simulator =
|trace: &runtime::trace::Trace| -> f64 { trace.get_f64(&addr!("param")).unwrap_or(0.0) };
let abc_samples = abc_scalar_summary(
&mut rng,
abc_model_fn,
simulator,
0.0, 1.0, 20, );
assert!(abc_samples.len() <= 20);
let vi_model_fn = || sample(addr!("x"), Normal::new(0.0, 1.0).unwrap());
let mut guide = MeanFieldGuide::new();
guide.params.insert(
addr!("x"),
VariationalParam::Normal {
mu: 0.0,
log_sigma: 0.0,
},
);
let optimized_guide = optimize_meanfield_vi(
&mut rng,
vi_model_fn,
guide,
5, 10, 0.1, );
assert!(!optimized_guide.params.is_empty());
}
#[test]
fn test_compatibility_legacy_patterns() {
let mut rng = StdRng::seed_from_u64(42);
let model = sample(addr!("param"), Normal::new(0.0, 1.0).unwrap());
let handler = runtime::interpreters::PriorHandler {
rng: &mut rng,
trace: runtime::trace::Trace::default(),
};
let (result, trace) = runtime::handler::run(handler, model);
assert!(result.is_finite());
assert!(trace.get_f64(&addr!("param")).is_some());
let mut manual_trace = runtime::trace::Trace::default();
manual_trace.insert_choice(addr!("manual"), runtime::trace::ChoiceValue::F64(2.5), -1.0);
manual_trace.log_prior += -1.0;
assert_eq!(manual_trace.get_f64(&addr!("manual")), Some(2.5));
assert!((manual_trace.log_prior + 1.0).abs() < 1e-12);
let legacy_mcmc_model = || sample(addr!("theta"), Normal::new(0.0, 1.0).unwrap());
let legacy_samples = adaptive_mcmc_chain(&mut rng, legacy_mcmc_model, 20, 5);
assert_eq!(legacy_samples.len(), 20);
assert!(legacy_samples.iter().all(|(val, _)| val.is_finite()));
let legacy_composition = pure(1.0).bind(|x| pure(x + 1.0)).map(|x| x * 2.0);
let handler2 = runtime::interpreters::PriorHandler {
rng: &mut rng,
trace: runtime::trace::Trace::default(),
};
let (legacy_result, _) = runtime::handler::run(handler2, legacy_composition);
assert_eq!(legacy_result, 4.0); }
#[test]
fn test_compatibility_api_stability() {
let _normal: Normal = Normal::new(0.0, 1.0).unwrap();
let _bernoulli: Bernoulli = Bernoulli::new(0.5).unwrap();
let _uniform: Uniform = Uniform::new(0.0, 1.0).unwrap();
let _pure_model: Model<i32> = pure(42);
let _sample_model: Model<f64> = sample(addr!("x"), Normal::new(0.0, 1.0).unwrap());
let _observe_model: Model<()> = observe(addr!("y"), Normal::new(0.0, 1.0).unwrap(), 0.0);
let _factor_model: Model<()> = factor(-0.5);
let trace = runtime::trace::Trace::default();
let _f64_option: Option<f64> = trace.get_f64(&addr!("test"));
let _bool_option: Option<bool> = trace.get_bool(&addr!("test"));
let _f64_result: Result<f64, FugueError> = trace.get_f64_result(&addr!("test"));
let _bool_result: Result<bool, FugueError> = trace.get_bool_result(&addr!("test"));
let mut rng = StdRng::seed_from_u64(42);
let _prior_handler = runtime::interpreters::PriorHandler {
rng: &mut rng,
trace: runtime::trace::Trace::default(),
};
let model_fn = || sample(addr!("x"), Normal::new(0.0, 1.0).unwrap());
let _mcmc_samples: Vec<(f64, runtime::trace::Trace)> =
adaptive_mcmc_chain(&mut rng, model_fn, 10, 2);
let _error_code: ErrorCode = ErrorCode::InvalidVariance;
let _error_category: ErrorCategory = ErrorCategory::DistributionValidation;
}
#[test]
fn test_api_contract_comprehensive_validation() {
let mut rng = StdRng::seed_from_u64(42);
let distributions: Vec<Box<dyn Distribution<f64>>> = vec![
Box::new(Normal::new(0.0, 1.0).unwrap()),
Box::new(Uniform::new(0.0, 1.0).unwrap()),
Box::new(Exponential::new(1.0).unwrap()),
Box::new(Beta::new(1.0, 1.0).unwrap()),
Box::new(Gamma::new(1.0, 1.0).unwrap()),
Box::new(LogNormal::new(0.0, 1.0).unwrap()),
];
for dist in distributions {
let sample = dist.sample(&mut rng);
let log_prob = dist.log_prob(&sample);
assert!(sample.is_finite());
assert!(log_prob.is_finite());
}
let make_test_model = || sample(addr!("test"), Normal::new(0.0, 1.0).unwrap());
let prior_handler = runtime::interpreters::PriorHandler {
rng: &mut rng,
trace: runtime::trace::Trace::default(),
};
let (_, base_trace) = runtime::handler::run(prior_handler, make_test_model());
let replay_handler = runtime::interpreters::ReplayHandler {
rng: &mut rng,
base: base_trace.clone(),
trace: runtime::trace::Trace::default(),
};
let (_, replay_trace) = runtime::handler::run(replay_handler, make_test_model());
assert_eq!(
base_trace.get_f64(&addr!("test")),
replay_trace.get_f64(&addr!("test"))
);
let safe_replay_handler = runtime::interpreters::SafeReplayHandler {
rng: &mut rng,
base: base_trace.clone(),
trace: runtime::trace::Trace::default(),
warn_on_mismatch: false,
};
let (_, safe_replay_trace) = runtime::handler::run(safe_replay_handler, make_test_model());
assert_eq!(
base_trace.get_f64(&addr!("test")),
safe_replay_trace.get_f64(&addr!("test"))
);
let score_handler = runtime::interpreters::ScoreGivenTrace {
base: base_trace.clone(),
trace: runtime::trace::Trace::default(),
};
let (_, score_trace) = runtime::handler::run(score_handler, make_test_model());
assert_eq!(
base_trace.get_f64(&addr!("test")),
score_trace.get_f64(&addr!("test"))
);
let safe_score_handler = runtime::interpreters::SafeScoreGivenTrace {
base: base_trace.clone(),
trace: runtime::trace::Trace::default(),
warn_on_error: false,
};
let (_, safe_score_trace) = runtime::handler::run(safe_score_handler, make_test_model());
assert_eq!(
base_trace.get_f64(&addr!("test")),
safe_score_trace.get_f64(&addr!("test"))
);
let inference_model = || {
sample(addr!("param"), Normal::new(0.0, 1.0).unwrap()).bind(|param| {
observe(addr!("obs"), Normal::new(param, 0.5).unwrap(), 0.5).map(move |_| param)
})
};
let mcmc_samples = adaptive_mcmc_chain(&mut rng, inference_model, 20, 5);
assert_eq!(mcmc_samples.len(), 20);
let smc_config = SMCConfig::default();
let smc_particles = adaptive_smc(&mut rng, 15, inference_model, smc_config);
assert_eq!(smc_particles.len(), 15);
let abc_model = || sample(addr!("param"), Normal::new(0.0, 1.0).unwrap());
let simulator = |trace: &runtime::trace::Trace| trace.get_f64(&addr!("param")).unwrap_or(0.0);
let abc_samples = abc_scalar_summary(&mut rng, abc_model, simulator, 0.5, 1.0, 10);
assert!(abc_samples.len() <= 10);
let vi_model = || sample(addr!("param"), Normal::new(0.0, 1.0).unwrap());
let mut guide = MeanFieldGuide::new();
guide.params.insert(
addr!("param"),
VariationalParam::Normal {
mu: 0.0,
log_sigma: 0.0,
},
);
let optimized_guide = optimize_meanfield_vi(&mut rng, vi_model, guide, 5, 5, 0.1);
assert!(!optimized_guide.params.is_empty());
}