use serde::{Deserialize, Serialize};
use super::ast::Comprehension;
use super::cardinality::CardinalityClass;
use super::metadata::{IndexFn, Metadata};
use super::source::Source;
use super::strategy::{StrategyName, ZipMode};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[derive(Default)]
pub enum Mode {
#[default]
Permissive,
Strict,
}
#[derive(Debug, Clone)]
pub struct ValidationReport {
pub warnings: Vec<ValidationWarning>,
}
#[derive(Debug, Clone, PartialEq)]
pub enum ValidationError {
V1DuplicateName { combinator: &'static str, name: String },
V2ShapeMismatch { expected: Vec<String>, actual: Vec<String> },
V3UnresolvedNames {
predicate: String,
coords: Vec<String>,
unresolved: Vec<String>,
},
V4InputShape {
strategy: StrategyName,
reason: String,
},
V6UnboundedDiscrete {
operator: &'static str,
cardinality: CardinalityClass,
},
V7ZipCardinality {
mode: ZipMode,
reason: String,
},
V8ContinuousRequirement { reason: String },
V9UnionClassMismatch { reason: String },
}
#[derive(Debug, Clone, PartialEq)]
pub enum ValidationWarning {
DegenerateGeometric { strategy: StrategyName },
LhsDegenerate,
TriviallyTrueFilter,
TriviallyFalseFilter,
SingletonCombinator { combinator: &'static str },
}
pub fn validate(c: &Comprehension, mode: Mode) -> Result<ValidationReport, ValidationError> {
let mut report = ValidationReport { warnings: Vec::new() };
visit(c, &mut report)?;
if mode == Mode::Strict
&& let Some(_warning) = report.warnings.first()
{
return Err(ValidationError::V8ContinuousRequirement {
reason: format!(
"strict mode: warning promoted: {:?}",
report.warnings.first().unwrap()
),
});
}
Ok(report)
}
fn visit(c: &Comprehension, report: &mut ValidationReport) -> Result<(), ValidationError> {
for child in c.children() {
visit(child, report)?;
}
match c {
Comprehension::Clause { source, .. } => visit_clause(source, report),
Comprehension::Cartesian { children } => visit_cartesian(children, report),
Comprehension::Zip { children, mode } => visit_zip(children, *mode, report),
Comprehension::Union { children } => visit_union(children, report),
Comprehension::Filter { child, predicate } => visit_filter(child, predicate, report),
Comprehension::Order { child, strategy, truncation } => {
visit_order(child, *strategy, *truncation, report)
}
}
}
fn visit_clause(source: &Source, report: &mut ValidationReport) -> Result<(), ValidationError> {
if let Source::ContinuousInterval { interval, measure } = source
&& !measure.is_integrable(std::slice::from_ref(interval))
{
let _ = report; return Err(ValidationError::V8ContinuousRequirement {
reason: format!(
"continuous source has non-integrable measure: \
interval [{}, {}] + {:?}",
interval.lo, interval.hi, measure
),
});
}
Ok(())
}
fn visit_cartesian(
children: &[Comprehension],
report: &mut ValidationReport,
) -> Result<(), ValidationError> {
check_disjoint_names("cartesian", children)?;
if children.len() == 1 {
report.warnings.push(ValidationWarning::SingletonCombinator {
combinator: "cartesian",
});
}
Ok(())
}
fn visit_zip(
children: &[Comprehension],
mode: ZipMode,
report: &mut ValidationReport,
) -> Result<(), ValidationError> {
check_disjoint_names("zip", children)?;
for child in children {
if contains_continuous_source(child) {
return Err(ValidationError::V7ZipCardinality {
mode,
reason: "zip children must all be discrete; \
a continuous source was found"
.to_string(),
});
}
}
if children.len() == 1 {
report.warnings.push(ValidationWarning::SingletonCombinator {
combinator: "zip",
});
}
if matches!(mode, ZipMode::Strict | ZipMode::Truncate) {
for child in children {
if let Some(card) = direct_source_cardinality(child)
&& matches!(card, CardinalityClass::Unbounded)
{
return Err(ValidationError::V6UnboundedDiscrete {
operator: "zip",
cardinality: card,
});
}
}
}
Ok(())
}
fn visit_union(
children: &[Comprehension],
report: &mut ValidationReport,
) -> Result<(), ValidationError> {
for child in children {
if contains_continuous_source(child) {
return Err(ValidationError::V9UnionClassMismatch {
reason: "union children must all be discrete; \
a continuous source was found"
.to_string(),
});
}
}
if let Some(first) = children.first() {
let expected = first.coordinate_names();
for sibling in &children[1..] {
let actual = sibling.coordinate_names();
if actual != expected {
return Err(ValidationError::V2ShapeMismatch {
expected,
actual,
});
}
}
}
if children.len() == 1 {
report.warnings.push(ValidationWarning::SingletonCombinator {
combinator: "union",
});
}
Ok(())
}
fn visit_filter(
child: &Comprehension,
predicate: &str,
report: &mut ValidationReport,
) -> Result<(), ValidationError> {
let coords = child.coordinate_names();
let referenced = extract_interpolated_names(predicate);
let unresolved: Vec<String> = referenced
.into_iter()
.filter(|n| !coords.contains(n))
.collect();
let _ = unresolved;
let trimmed = predicate.trim();
if trimmed.eq_ignore_ascii_case("true") {
report.warnings.push(ValidationWarning::TriviallyTrueFilter);
} else if trimmed.eq_ignore_ascii_case("false") {
report.warnings.push(ValidationWarning::TriviallyFalseFilter);
}
let _ = child;
Ok(())
}
fn visit_order(
child: &Comprehension,
strategy: StrategyName,
truncation: Option<u64>,
report: &mut ValidationReport,
) -> Result<(), ValidationError> {
let metadata_target = match child {
Comprehension::Filter { child: inner, .. } => inner.as_ref(),
other => other,
};
if !matches!(strategy, StrategyName::Lex)
&& matches!(metadata_target, Comprehension::Filter { .. })
{
return Err(ValidationError::V4InputShape {
strategy,
reason: "non-Lex strategy applied to nested filter; \
fold filters first (spec F1 / R6)"
.to_string(),
});
}
let target_metadata = metadata_target.metadata();
check_strategy_input_shape(strategy, &target_metadata, report)?;
if !matches!(strategy, StrategyName::Lex)
&& matches!(target_metadata.cardinality, CardinalityClass::Unbounded)
{
return Err(ValidationError::V6UnboundedDiscrete {
operator: "order",
cardinality: target_metadata.cardinality.clone(),
});
}
let is_continuous = matches!(
target_metadata.cardinality,
CardinalityClass::Continuous { .. }
| CardinalityClass::ContinuousAtMost { .. }
| CardinalityClass::Hybrid(_)
);
if is_continuous {
if truncation.is_none() {
return Err(ValidationError::V8ContinuousRequirement {
reason: "continuous comprehension requires order(_, \
sampling-strategy, Some(n)) with finite \
truncation"
.to_string(),
});
}
if matches!(strategy, StrategyName::Lex) {
return Err(ValidationError::V8ContinuousRequirement {
reason: "Lex does not sample continuous inputs; use \
Halton / Sobol / Lhs / Shuffle / Extrema"
.to_string(),
});
}
}
Ok(())
}
fn check_strategy_input_shape(
strategy: StrategyName,
metadata: &Metadata,
report: &mut ValidationReport,
) -> Result<(), ValidationError> {
if matches!(strategy, StrategyName::Lex) {
return Ok(());
}
let idx = match &metadata.index_addressable {
Some(i) => i,
None => {
return Err(ValidationError::V4InputShape {
strategy,
reason: "input has no closed-form index function \
(raw filter output, dependent cartesian, or \
nested non-Lex order)"
.to_string(),
});
}
};
let has_continuous = idx.has_continuous_axis();
if has_continuous {
match strategy {
StrategyName::Shuffle
| StrategyName::Halton
| StrategyName::Sobol
| StrategyName::Lhs => {}
StrategyName::Extrema => {}
StrategyName::ReverseLex
| StrategyName::Shells
| StrategyName::Diagonal
| StrategyName::Antidiagonal => {
return Err(ValidationError::V4InputShape {
strategy,
reason: format!(
"{} does not accept continuous input",
strategy.as_str()
),
});
}
StrategyName::Lex => unreachable!("Lex handled above"),
}
if matches!(strategy, StrategyName::Lhs | StrategyName::Extrema) {
let dim = continuous_dim(idx);
if dim < 2 {
if matches!(strategy, StrategyName::Lhs) {
report.warnings.push(ValidationWarning::LhsDegenerate);
} else {
report.warnings.push(ValidationWarning::DegenerateGeometric { strategy });
}
}
}
return Ok(());
}
if strategy.is_lattice_geometric() {
match idx {
IndexFn::Lattice { axis_sizes } => {
if axis_sizes.len() < 2 {
report.warnings.push(ValidationWarning::DegenerateGeometric { strategy });
}
}
IndexFn::Concatenation { .. } => {
return Err(ValidationError::V4InputShape {
strategy,
reason: format!(
"{} requires a cartesian input; got union",
strategy.as_str()
),
});
}
IndexFn::Lockstep { .. } | IndexFn::Modular { .. } => {
report.warnings.push(ValidationWarning::DegenerateGeometric { strategy });
}
IndexFn::Continuous { .. } | IndexFn::Hybrid { .. } => unreachable!(),
}
return Ok(());
}
if matches!(strategy, StrategyName::Lhs) {
match idx {
IndexFn::Lattice { axis_sizes } if axis_sizes.len() < 2 => {
report.warnings.push(ValidationWarning::LhsDegenerate);
}
IndexFn::Lockstep { .. } | IndexFn::Modular { .. } => {
report.warnings.push(ValidationWarning::LhsDegenerate);
}
_ => {}
}
}
Ok(())
}
fn continuous_dim(idx: &IndexFn) -> usize {
match idx {
IndexFn::Continuous { intervals, .. } => intervals.len(),
IndexFn::Hybrid {
discrete_axes,
continuous_axes,
..
} => discrete_axes.len() + continuous_axes.len(),
_ => 0,
}
}
fn check_disjoint_names(
combinator: &'static str,
children: &[Comprehension],
) -> Result<(), ValidationError> {
let mut seen: Vec<String> = Vec::new();
for child in children {
for name in child.coordinate_names() {
if seen.contains(&name) {
return Err(ValidationError::V1DuplicateName {
combinator,
name,
});
}
seen.push(name);
}
}
Ok(())
}
fn contains_continuous_source(c: &Comprehension) -> bool {
match c {
Comprehension::Clause { source, .. } => source.is_continuous(),
Comprehension::Cartesian { children } | Comprehension::Zip { children, .. } | Comprehension::Union { children } => {
children.iter().any(contains_continuous_source)
}
Comprehension::Filter { child, .. } | Comprehension::Order { child, .. } => {
contains_continuous_source(child)
}
}
}
fn direct_source_cardinality(c: &Comprehension) -> Option<CardinalityClass> {
match c {
Comprehension::Clause { source, .. } => Some(source.cardinality()),
_ => None,
}
}
fn extract_interpolated_names(predicate: &str) -> Vec<String> {
let mut out = Vec::new();
let bytes = predicate.as_bytes();
let mut i = 0;
while i < bytes.len() {
if bytes[i] == b'{'
&& let Some(close) = predicate[i + 1..].find('}') {
let name = predicate[i + 1..i + 1 + close].trim();
if !name.is_empty() && name.chars().all(|c| c.is_alphanumeric() || c == '_') {
out.push(name.to_string());
}
i += close + 2;
continue;
}
i += 1;
}
out
}
#[cfg(test)]
mod tests {
use super::*;
use crate::iteration::comprehension::source::{LiteralValue, Source};
use crate::iteration::comprehension::cardinality::{Interval, ProductMeasure};
fn clause(name: &str, vs: &[i64]) -> Comprehension {
Comprehension::clause(
name,
Source::Literal {
values: vs.iter().map(|n| LiteralValue::Int(*n)).collect(),
},
)
}
fn continuous_clause(name: &str) -> Comprehension {
Comprehension::clause(
name,
Source::ContinuousInterval {
interval: Interval::closed(0.0, 1.0),
measure: ProductMeasure::Uniform,
},
)
}
#[test]
fn v1_rejects_duplicate_names_in_cartesian() {
let bad = Comprehension::cartesian(vec![clause("k", &[1]), clause("k", &[2])]);
let result = validate(&bad, Mode::Permissive);
assert!(matches!(
result,
Err(ValidationError::V1DuplicateName { combinator: "cartesian", .. })
));
}
#[test]
fn v1_accepts_disjoint_names() {
let ok = Comprehension::cartesian(vec![clause("k", &[1]), clause("limit", &[10])]);
assert!(validate(&ok, Mode::Permissive).is_ok());
}
#[test]
fn v2_rejects_union_shape_mismatch() {
let bad = Comprehension::union(vec![
Comprehension::cartesian(vec![clause("k", &[1]), clause("limit", &[10])]),
Comprehension::cartesian(vec![clause("limit", &[100]), clause("k", &[100])]),
]);
let result = validate(&bad, Mode::Permissive);
assert!(matches!(result, Err(ValidationError::V2ShapeMismatch { .. })));
}
#[test]
fn v2_accepts_matching_union_shape() {
let ok = Comprehension::union(vec![
Comprehension::cartesian(vec![clause("k", &[1]), clause("limit", &[10])]),
Comprehension::cartesian(vec![clause("k", &[100]), clause("limit", &[100])]),
]);
assert!(validate(&ok, Mode::Permissive).is_ok());
}
#[test]
fn v4_rejects_lattice_geometric_over_union() {
let bad = Comprehension::order(
Comprehension::union(vec![
clause("k", &[1, 2, 3]),
clause("k", &[10, 20, 30]),
]),
StrategyName::Extrema,
Some(2),
);
assert!(matches!(
validate(&bad, Mode::Permissive),
Err(ValidationError::V4InputShape { strategy: StrategyName::Extrema, .. })
));
}
#[test]
fn v4_lattice_geometric_over_1axis_warns_not_errors() {
let degenerate = Comprehension::order(
clause("k", &[1, 2, 3]),
StrategyName::Extrema,
Some(2),
);
let report = validate(°enerate, Mode::Permissive).unwrap();
assert!(report.warnings.iter().any(|w| matches!(
w,
ValidationWarning::DegenerateGeometric { strategy: StrategyName::Extrema }
)));
}
#[test]
fn v4_strict_mode_promotes_warning() {
let degenerate = Comprehension::order(
clause("k", &[1, 2, 3]),
StrategyName::Extrema,
Some(2),
);
assert!(validate(°enerate, Mode::Strict).is_err());
}
#[test]
fn v7_rejects_continuous_in_zip() {
let bad = Comprehension::zip(
vec![continuous_clause("alpha"), continuous_clause("beta")],
ZipMode::Strict,
);
assert!(matches!(
validate(&bad, Mode::Permissive),
Err(ValidationError::V7ZipCardinality { .. })
));
}
#[test]
fn v8_rejects_continuous_without_sampling() {
let bad = continuous_clause("theta");
assert!(validate(&bad, Mode::Permissive).is_ok());
let bad_lex = Comprehension::order(continuous_clause("theta"), StrategyName::Lex, None);
assert!(matches!(
validate(&bad_lex, Mode::Permissive),
Err(ValidationError::V8ContinuousRequirement { .. })
));
}
#[test]
fn v8_accepts_continuous_with_sampling() {
let ok = Comprehension::order(
Comprehension::cartesian(vec![continuous_clause("alpha"), continuous_clause("beta")]),
StrategyName::Halton,
Some(100),
);
assert!(validate(&ok, Mode::Permissive).is_ok());
}
#[test]
fn v8_rejects_unbounded_uniform_at_source() {
let bad = Comprehension::clause(
"x",
Source::ContinuousInterval {
interval: Interval { lo: 0.0, hi: f64::INFINITY, lo_open: false, hi_open: true },
measure: ProductMeasure::Uniform,
},
);
assert!(matches!(
validate(&bad, Mode::Permissive),
Err(ValidationError::V8ContinuousRequirement { .. })
));
}
#[test]
fn v9_rejects_continuous_in_union() {
let bad = Comprehension::union(vec![
Comprehension::cartesian(vec![continuous_clause("k"), continuous_clause("limit")]),
Comprehension::cartesian(vec![continuous_clause("k"), continuous_clause("limit")]),
]);
assert!(matches!(
validate(&bad, Mode::Permissive),
Err(ValidationError::V9UnionClassMismatch { .. })
));
}
#[test]
fn singleton_combinator_warns() {
let degenerate = Comprehension::cartesian(vec![clause("k", &[1, 2])]);
let report = validate(°enerate, Mode::Permissive).unwrap();
assert!(report.warnings.iter().any(|w| matches!(
w,
ValidationWarning::SingletonCombinator { combinator: "cartesian" }
)));
}
#[test]
fn trivially_true_filter_warns() {
let degenerate = Comprehension::filter(clause("k", &[1, 2]), "true");
let report = validate(°enerate, Mode::Permissive).unwrap();
assert!(report.warnings.iter().any(|w| matches!(w, ValidationWarning::TriviallyTrueFilter)));
}
#[test]
fn name_extraction_handles_simple_predicates() {
assert_eq!(extract_interpolated_names("{k} > 0"), vec!["k"]);
assert_eq!(
extract_interpolated_names("{k} * {limit} <= 1000"),
vec!["k", "limit"]
);
assert_eq!(extract_interpolated_names("no refs here"), Vec::<String>::new());
}
}