use alloc::{string::String, vec::Vec}; use crate::{ Diagnostic, DiagnosticKind, Span, ast::CompareOp, runtime::{Constraint, ConstraintEntry, LiteralValue, PrimitiveType}, }; pub fn normalize_constraints( constraints: Vec, span: Span, ) -> crate::Result> { let mut primitive: Option<(PrimitiveType, Span)> = None; let mut lower: Option<(Bound, Span)> = None; let mut upper: Option<(Bound, Span)> = None; let mut rest = Vec::new(); for entry in constraints { match entry.constraint { Constraint::Type(next) => match primitive { Some((current, current_span)) if current != next => { return Err(Diagnostic::new( DiagnosticKind::Conflict, span, "primitive type constraints conflict", ) .with_label(current_span, "first primitive constraint") .with_label(entry.span, "conflicting primitive constraint")); } Some(_) => {} None => primitive = Some((next, entry.span)), }, Constraint::Compare(op, value) => { let number = Number::from_literal(&value).ok_or_else(|| { Diagnostic::new( DiagnosticKind::TypeMismatch, entry.span, "comparison constraints require numeric literals", ) })?; match op { CompareOp::Gt => { merge_lower(&mut lower, (Bound::new(number, false), entry.span)) } CompareOp::Gte => { merge_lower(&mut lower, (Bound::new(number, true), entry.span)) } CompareOp::Lt => { merge_upper(&mut upper, (Bound::new(number, false), entry.span)) } CompareOp::Lte => { merge_upper(&mut upper, (Bound::new(number, true), entry.span)) } CompareOp::Eq => { merge_lower(&mut lower, (Bound::new(number, true), entry.span)); merge_upper(&mut upper, (Bound::new(number, true), entry.span)); } } } Constraint::Regex(pattern) => rest.push(ConstraintEntry { constraint: Constraint::Regex(pattern), span: entry.span, }), Constraint::BuiltinPredicate(name) => rest.push(ConstraintEntry { constraint: Constraint::BuiltinPredicate(name), span: entry.span, }), } } if matches!( primitive, Some(( PrimitiveType::String | PrimitiveType::Bool | PrimitiveType::Array, _ )) ) && (lower.is_some() || upper.is_some()) { let mut diagnostic = Diagnostic::new( DiagnosticKind::Conflict, span, "numeric comparison constraints conflict with non-numeric primitive type", ); if let Some((_, primitive_span)) = primitive { diagnostic = diagnostic.with_label(primitive_span, "non-numeric primitive constraint"); } if let Some((_, lower_span)) = lower { diagnostic = diagnostic.with_label(lower_span, "numeric comparison constraint"); } if let Some((_, upper_span)) = upper { diagnostic = diagnostic.with_label(upper_span, "numeric comparison constraint"); } return Err(diagnostic); } if matches!(primitive, Some((PrimitiveType::Int, _))) && lower .iter() .chain(upper.iter()) .any(|(bound, _)| !matches!(bound.number, Number::Int(_))) { let mut diagnostic = Diagnostic::new( DiagnosticKind::Conflict, span, "Int comparison constraints must use integer literals", ); if let Some((_, primitive_span)) = primitive { diagnostic = diagnostic.with_label(primitive_span, "Int primitive constraint"); } if let Some((bound, bound_span)) = lower { if !matches!(bound.number, Number::Int(_)) { diagnostic = diagnostic.with_label(bound_span, "non-integer comparison constraint"); } } if let Some((bound, bound_span)) = upper { if !matches!(bound.number, Number::Int(_)) { diagnostic = diagnostic.with_label(bound_span, "non-integer comparison constraint"); } } return Err(diagnostic); } ensure_bounds_non_empty(primitive, lower, upper, span)?; let mut normalized = Vec::new(); if let Some((primitive, primitive_span)) = primitive { normalized.push(ConstraintEntry { constraint: Constraint::Type(primitive), span: primitive_span, }); } if let Some((lower, lower_span)) = lower { normalized.push(ConstraintEntry { constraint: Constraint::Compare( if lower.inclusive { CompareOp::Gte } else { CompareOp::Gt }, lower.number.into_literal(), ), span: lower_span, }); } if let Some((upper, upper_span)) = upper { normalized.push(ConstraintEntry { constraint: Constraint::Compare( if upper.inclusive { CompareOp::Lte } else { CompareOp::Lt }, upper.number.into_literal(), ), span: upper_span, }); } normalized.extend(rest); Ok(normalized) } #[derive(Debug, Clone, Copy, PartialEq)] struct Bound { number: Number, inclusive: bool, } impl Bound { fn new(number: Number, inclusive: bool) -> Self { Self { number, inclusive } } } #[derive(Debug, Clone, Copy, PartialEq)] enum Number { Int(i64), Float(f64), } impl Number { fn from_literal(value: &LiteralValue) -> Option { match value { LiteralValue::Int(value) => Some(Self::Int(*value)), LiteralValue::Float(value) => Some(Self::Float(*value)), LiteralValue::String(_) | LiteralValue::Bool(_) => None, } } fn into_literal(self) -> LiteralValue { match self { Self::Int(value) => LiteralValue::Int(value), Self::Float(value) => LiteralValue::Float(value), } } fn as_f64(self) -> f64 { match self { Self::Int(value) => value as f64, Self::Float(value) => value, } } } fn merge_lower(current: &mut Option<(Bound, Span)>, next: (Bound, Span)) { match current { None => *current = Some(next), Some((current_bound, _)) if is_stricter_lower(next.0, *current_bound) => { *current = Some(next) } Some(_) => {} } } fn merge_upper(current: &mut Option<(Bound, Span)>, next: (Bound, Span)) { match current { None => *current = Some(next), Some((current_bound, _)) if is_stricter_upper(next.0, *current_bound) => { *current = Some(next) } Some(_) => {} } } fn is_stricter_lower(next: Bound, current: Bound) -> bool { let next_value = next.number.as_f64(); let current_value = current.number.as_f64(); next_value > current_value || (next_value == current_value && !next.inclusive && current.inclusive) } fn is_stricter_upper(next: Bound, current: Bound) -> bool { let next_value = next.number.as_f64(); let current_value = current.number.as_f64(); next_value < current_value || (next_value == current_value && !next.inclusive && current.inclusive) } fn ensure_bounds_non_empty( primitive: Option<(PrimitiveType, Span)>, lower: Option<(Bound, Span)>, upper: Option<(Bound, Span)>, span: Span, ) -> crate::Result<()> { if matches!(primitive, Some((PrimitiveType::Int, _))) { let min = lower .map(|(bound, _)| int_lower_bound(bound)) .unwrap_or(i128::from(i64::MIN)); let max = upper .map(|(bound, _)| int_upper_bound(bound)) .unwrap_or(i128::from(i64::MAX)); if min > max { return Err(empty_numeric_bounds(span, lower, upper)); } return Ok(()); } if let (Some((lower_bound, _)), Some((upper_bound, _))) = (lower, upper) { let lower_value = lower_bound.number.as_f64(); let upper_value = upper_bound.number.as_f64(); if lower_value > upper_value { return Err(empty_numeric_bounds(span, lower, upper)); } if lower_value == upper_value && !(lower_bound.inclusive && upper_bound.inclusive) { return Err(empty_numeric_bounds(span, lower, upper)); } } Ok(()) } fn int_lower_bound(bound: Bound) -> i128 { let Number::Int(value) = bound.number else { unreachable!() }; if bound.inclusive { i128::from(value) } else { i128::from(value) + 1 } } fn int_upper_bound(bound: Bound) -> i128 { let Number::Int(value) = bound.number else { unreachable!() }; if bound.inclusive { i128::from(value) } else { i128::from(value) - 1 } } fn empty_numeric_bounds( span: Span, lower: Option<(Bound, Span)>, upper: Option<(Bound, Span)>, ) -> Diagnostic { let mut diagnostic = Diagnostic::new( DiagnosticKind::Conflict, span, String::from("numeric comparison constraints have an empty intersection"), ); if let Some((_, lower_span)) = lower { diagnostic = diagnostic.with_label(lower_span, "lower bound constraint"); } if let Some((_, upper_span)) = upper { diagnostic = diagnostic.with_label(upper_span, "upper bound constraint"); } diagnostic } #[cfg(test)] mod tests { use super::*; use crate::runtime::LiteralValue; fn entry(constraint: Constraint) -> ConstraintEntry { ConstraintEntry { constraint, span: Span::default(), } } #[test] fn detects_primitive_conflict() { assert!( normalize_constraints( alloc::vec![ entry(Constraint::Type(PrimitiveType::Int)), entry(Constraint::Type(PrimitiveType::String)), ], Span::default(), ) .is_err() ); } #[test] fn detects_empty_int_range() { assert!( normalize_constraints( alloc::vec![ entry(Constraint::Type(PrimitiveType::Int)), entry(Constraint::Compare(CompareOp::Gt, LiteralValue::Int(10))), entry(Constraint::Compare(CompareOp::Lt, LiteralValue::Int(5))), ], Span::default(), ) .is_err() ); } #[test] fn detects_non_integer_int_bound() { assert!( normalize_constraints( alloc::vec![ entry(Constraint::Type(PrimitiveType::Int)), entry(Constraint::Compare(CompareOp::Gt, LiteralValue::Float(1.5))), ], Span::default(), ) .is_err() ); } }