use rand::Rng;
use rand_distr::{Distribution, LogNormal, Normal};
use serde::Deserialize;
#[derive(Deserialize, Clone, Debug)]
#[serde(untagged)]
pub enum RadiusSpec {
Fixed(f64),
Distribution(RadiusDistribution),
}
#[derive(Deserialize, Clone, Debug)]
#[serde(tag = "distribution", rename_all = "lowercase")]
pub enum RadiusDistribution {
Uniform {
min: f64,
max: f64,
},
Gaussian {
mean: f64,
std: f64,
},
Lognormal {
mean: f64,
std: f64,
},
Discrete {
values: Vec<f64>,
weights: Vec<f64>,
},
}
impl RadiusSpec {
pub fn try_sample(&self, rng: &mut impl Rng) -> Result<f64, String> {
match self {
RadiusSpec::Fixed(r) => Ok(*r),
RadiusSpec::Distribution(d) => d.try_sample(rng),
}
}
pub fn sample(&self, rng: &mut impl Rng) -> f64 {
self.try_sample(rng)
.expect("radius distribution parameters must be validated before sampling")
}
pub fn try_max_radius(&self) -> Result<f64, String> {
match self {
RadiusSpec::Fixed(r) => {
validate_positive_finite_radius(*r, "fixed radius")?;
Ok(*r)
}
RadiusSpec::Distribution(d) => d.try_max_radius(),
}
}
pub fn max_radius(&self) -> f64 {
self.try_max_radius()
.expect("radius distribution parameters must be validated before max_radius")
}
}
fn validate_positive_finite_radius(value: f64, context: &str) -> Result<(), String> {
if !value.is_finite() || value <= 0.0 {
return Err(format!("{context} must be finite and > 0, got {value}"));
}
Ok(())
}
impl RadiusDistribution {
fn try_max_radius(&self) -> Result<f64, String> {
match self {
RadiusDistribution::Uniform { min, max } => {
validate_positive_finite_radius(*min, "uniform radius min")?;
validate_positive_finite_radius(*max, "uniform radius max")?;
if min >= max {
return Err(format!(
"uniform radius requires min < max, got min={} max={}",
min, max
));
}
Ok(*max)
}
RadiusDistribution::Gaussian { mean, std } => {
validate_positive_finite_radius(*mean, "Gaussian radius mean")?;
if !std.is_finite() || *std <= 0.0 {
return Err(format!(
"Gaussian radius std must be finite and > 0, got {std}"
));
}
let max = mean + 4.0 * std;
validate_positive_finite_radius(max, "Gaussian radius max bound")?;
Ok(max)
}
RadiusDistribution::Lognormal { mean, std } => {
validate_positive_finite_radius(*mean, "lognormal radius mean")?;
if !std.is_finite() || *std <= 0.0 {
return Err(format!(
"lognormal radius std must be finite and > 0, got {std}"
));
}
let max = mean + 4.0 * std;
validate_positive_finite_radius(max, "lognormal radius max bound")?;
Ok(max)
}
RadiusDistribution::Discrete { values, weights } => {
if values.is_empty() {
return Err("discrete radius requires at least one value".to_string());
}
if values.len() != weights.len() {
return Err(format!(
"discrete radius requires values/weights length match, got {} values and {} weights",
values.len(),
weights.len()
));
}
for (i, value) in values.iter().enumerate() {
validate_positive_finite_radius(
*value,
&format!("discrete radius value[{i}]"),
)?;
}
for (i, weight) in weights.iter().enumerate() {
if !weight.is_finite() || *weight < 0.0 {
return Err(format!(
"discrete radius weight[{i}] must be finite and >= 0, got {weight}"
));
}
}
let total: f64 = weights.iter().sum();
if total <= 0.0 {
return Err(format!(
"discrete radius requires positive total weight, got {}",
total
));
}
Ok(values.iter().cloned().fold(f64::NEG_INFINITY, f64::max))
}
}
}
}
impl RadiusDistribution {
fn try_sample(&self, rng: &mut impl Rng) -> Result<f64, String> {
self.try_max_radius()?;
match self {
RadiusDistribution::Uniform { min, max } => Ok(rng.random_range(*min..*max)),
RadiusDistribution::Gaussian { mean, std } => {
let normal = Normal::new(*mean, *std).map_err(|e| {
format!("invalid Gaussian radius parameters (mean={mean}, std={std}): {e}")
})?;
Ok(normal.sample(rng).max(1e-15)) }
RadiusDistribution::Lognormal { mean, std } => {
let sigma_sq = (1.0 + (std / mean).powi(2)).ln();
let mu = mean.ln() - sigma_sq / 2.0;
let sigma = sigma_sq.sqrt();
let ln = LogNormal::new(mu, sigma).map_err(|e| {
format!("invalid lognormal radius parameters (mean={mean}, std={std}): {e}")
})?;
Ok(ln.sample(rng))
}
RadiusDistribution::Discrete { values, weights } => {
let total: f64 = weights.iter().sum();
let r: f64 = rng.random_range(0.0..total);
let mut cumulative = 0.0;
for (i, w) in weights.iter().enumerate() {
cumulative += w;
if r < cumulative {
return Ok(values[i]);
}
}
Ok(*values
.last()
.expect("discrete distribution was validated as non-empty"))
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use soil_core::toml;
#[test]
fn radius_spec_fixed_deserialization() {
let toml_str = "radius = 0.001";
#[derive(Deserialize)]
struct Wrapper {
radius: RadiusSpec,
}
let w: Wrapper = toml::from_str(toml_str).unwrap();
match w.radius {
RadiusSpec::Fixed(r) => assert!((r - 0.001).abs() < 1e-15),
_ => panic!("Expected Fixed variant"),
}
}
#[test]
fn radius_spec_uniform_deserialization() {
let toml_str = r#"radius = { distribution = "uniform", min = 0.0008, max = 0.0012 }"#;
#[derive(Deserialize)]
struct Wrapper {
radius: RadiusSpec,
}
let w: Wrapper = toml::from_str(toml_str).unwrap();
match &w.radius {
RadiusSpec::Distribution(RadiusDistribution::Uniform { min, max }) => {
assert!((min - 0.0008).abs() < 1e-15);
assert!((max - 0.0012).abs() < 1e-15);
}
other => panic!("Expected Uniform, got {:?}", other),
}
}
#[test]
fn radius_spec_gaussian_deserialization() {
let toml_str = r#"radius = { distribution = "gaussian", mean = 0.001, std = 0.0001 }"#;
#[derive(Deserialize)]
struct Wrapper {
radius: RadiusSpec,
}
let w: Wrapper = toml::from_str(toml_str).unwrap();
match &w.radius {
RadiusSpec::Distribution(RadiusDistribution::Gaussian { mean, std }) => {
assert!((mean - 0.001).abs() < 1e-15);
assert!((std - 0.0001).abs() < 1e-15);
}
other => panic!("Expected Gaussian, got {:?}", other),
}
}
#[test]
fn radius_spec_lognormal_deserialization() {
let toml_str = r#"radius = { distribution = "lognormal", mean = 0.001, std = 0.0001 }"#;
#[derive(Deserialize)]
struct Wrapper {
radius: RadiusSpec,
}
let w: Wrapper = toml::from_str(toml_str).unwrap();
match &w.radius {
RadiusSpec::Distribution(RadiusDistribution::Lognormal { mean, std }) => {
assert!((mean - 0.001).abs() < 1e-15);
assert!((std - 0.0001).abs() < 1e-15);
}
other => panic!("Expected Lognormal, got {:?}", other),
}
}
#[test]
fn radius_spec_discrete_deserialization() {
let toml_str = r#"radius = { distribution = "discrete", values = [0.001, 0.0015], weights = [0.7, 0.3] }"#;
#[derive(Deserialize)]
struct Wrapper {
radius: RadiusSpec,
}
let w: Wrapper = toml::from_str(toml_str).unwrap();
match &w.radius {
RadiusSpec::Distribution(RadiusDistribution::Discrete { values, weights }) => {
assert_eq!(values.len(), 2);
assert_eq!(weights.len(), 2);
assert!((values[0] - 0.001).abs() < 1e-15);
assert!((weights[0] - 0.7).abs() < 1e-15);
}
other => panic!("Expected Discrete, got {:?}", other),
}
}
#[test]
fn radius_spec_sampling_fixed() {
let spec = RadiusSpec::Fixed(0.005);
let mut rng = rand::rng();
for _ in 0..10 {
assert!((spec.sample(&mut rng) - 0.005).abs() < 1e-15);
}
}
#[test]
fn radius_spec_sampling_uniform() {
let spec = RadiusSpec::Distribution(RadiusDistribution::Uniform {
min: 0.001,
max: 0.002,
});
let mut rng = rand::rng();
for _ in 0..100 {
let r = spec.sample(&mut rng);
assert!(r >= 0.001 && r < 0.002, "uniform sample {} out of range", r);
}
}
#[test]
fn radius_spec_sampling_gaussian() {
let spec = RadiusSpec::Distribution(RadiusDistribution::Gaussian {
mean: 0.01,
std: 0.001,
});
let mut rng = rand::rng();
let samples: Vec<f64> = (0..1000).map(|_| spec.sample(&mut rng)).collect();
let mean: f64 = samples.iter().sum::<f64>() / samples.len() as f64;
assert!(
(mean - 0.01).abs() < 0.001,
"gaussian mean should be ~0.01, got {}",
mean
);
}
#[test]
fn radius_spec_sampling_lognormal() {
let spec = RadiusSpec::Distribution(RadiusDistribution::Lognormal {
mean: 0.01,
std: 0.001,
});
let mut rng = rand::rng();
let samples: Vec<f64> = (0..5000).map(|_| spec.sample(&mut rng)).collect();
let mean: f64 = samples.iter().sum::<f64>() / samples.len() as f64;
assert!(
(mean - 0.01).abs() < 0.002,
"lognormal mean should be ~0.01, got {}",
mean
);
assert!(
samples.iter().all(|&r| r > 0.0),
"lognormal samples should all be positive"
);
}
#[test]
fn radius_spec_sampling_discrete() {
let spec = RadiusSpec::Distribution(RadiusDistribution::Discrete {
values: vec![0.001, 0.002],
weights: vec![0.7, 0.3],
});
let mut rng = rand::rng();
let mut count_small = 0;
let n = 10000;
for _ in 0..n {
let r = spec.sample(&mut rng);
assert!(
(r - 0.001).abs() < 1e-15 || (r - 0.002).abs() < 1e-15,
"discrete sample should be one of the values"
);
if (r - 0.001).abs() < 1e-15 {
count_small += 1;
}
}
let ratio = count_small as f64 / n as f64;
assert!(
(ratio - 0.7).abs() < 0.05,
"discrete ratio should be ~0.7, got {}",
ratio
);
}
#[test]
fn malformed_radius_distribution_reports_error() {
let mut rng = rand::rng();
let spec = RadiusSpec::Distribution(RadiusDistribution::Lognormal {
mean: 0.0,
std: 0.1,
});
let err = spec
.try_sample(&mut rng)
.expect_err("invalid lognormal config should not panic");
assert!(err.contains("lognormal radius mean must be finite and > 0"));
}
#[test]
fn bad_gaussian_radius_parameters_report_error() {
let mut rng = rand::rng();
let spec = RadiusSpec::Distribution(RadiusDistribution::Gaussian {
mean: 0.001,
std: 0.0,
});
let err = spec
.try_sample(&mut rng)
.expect_err("zero Gaussian std should not panic");
assert!(err.contains("Gaussian radius std must be finite and > 0"));
}
#[test]
fn bad_lognormal_radius_parameters_report_error() {
let mut rng = rand::rng();
let spec = RadiusSpec::Distribution(RadiusDistribution::Lognormal {
mean: 0.001,
std: -0.1,
});
let err = spec
.try_sample(&mut rng)
.expect_err("negative lognormal std should not panic");
assert!(err.contains("lognormal radius std must be finite and > 0"));
}
#[test]
fn bad_discrete_radius_parameters_report_error() {
let mut rng = rand::rng();
let spec = RadiusSpec::Distribution(RadiusDistribution::Discrete {
values: vec![0.001, 0.002],
weights: vec![1.0],
});
let err = spec
.try_sample(&mut rng)
.expect_err("mismatched discrete values/weights should not panic");
assert!(err.contains("values/weights length match"));
}
#[test]
fn zero_weight_discrete_radius_reports_error() {
let spec = RadiusSpec::Distribution(RadiusDistribution::Discrete {
values: vec![0.001, 0.002],
weights: vec![0.0, 0.0],
});
let err = spec
.try_max_radius()
.expect_err("zero total discrete weight should not panic");
assert!(err.contains("positive total weight"));
}
#[test]
fn empty_discrete_radius_distribution_reports_error() {
let mut rng = rand::rng();
let spec = RadiusSpec::Distribution(RadiusDistribution::Discrete {
values: vec![],
weights: vec![],
});
let err = spec
.try_sample(&mut rng)
.expect_err("empty discrete config should not panic");
assert!(err.contains("discrete radius requires at least one value"));
}
#[test]
fn empty_discrete_radius_distribution_reports_max_radius_error() {
let spec = RadiusSpec::Distribution(RadiusDistribution::Discrete {
values: vec![],
weights: vec![],
});
let err = spec
.try_max_radius()
.expect_err("empty discrete config should fail before region construction");
assert!(err.contains("discrete radius requires at least one value"));
}
#[test]
fn radius_spec_max_radius() {
assert!((RadiusSpec::Fixed(0.005).max_radius() - 0.005).abs() < 1e-15);
assert!(
(RadiusSpec::Distribution(RadiusDistribution::Uniform {
min: 0.001,
max: 0.003
})
.max_radius()
- 0.003)
.abs()
< 1e-15
);
assert!(
(RadiusSpec::Distribution(RadiusDistribution::Discrete {
values: vec![0.001, 0.005, 0.002],
weights: vec![1.0, 1.0, 1.0],
})
.max_radius()
- 0.005)
.abs()
< 1e-15
);
}
}