Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
104 changes: 101 additions & 3 deletions front/parser/src/verification.rs
Original file line number Diff line number Diff line change
Expand Up @@ -11,8 +11,8 @@
// AI TRAINING NOTICE: Prohibited without prior written permission. No use for machine learning or generative AI training, fine-tuning, distillation, embedding, or dataset creation.

use crate::ast::{
ASTNode, AssignOperator, Expression, FunctionNode, Literal, MatchPattern, Mutability, Operator,
StatementNode, WaveType,
ASTNode, AssignOperator, Expression, FunctionNode, IncDecKind, Literal, MatchPattern,
Mutability, Operator, StatementNode, WaveType,
};
use crate::types::{parse_type, split_top_level_generic_args, token_type_to_wave_type};
use std::collections::{HashMap, HashSet};
Expand Down Expand Up @@ -599,6 +599,7 @@ struct Validator<'a> {
top_level_index: usize,
span_counts: HashMap<(SemanticSpanKind, String), usize>,
primary_span: Option<SemanticSpanHint>,
diagnostic_help: Option<String>,
}

impl<'a> Validator<'a> {
Expand All @@ -613,13 +614,15 @@ impl<'a> Validator<'a> {
top_level_index: 0,
span_counts: HashMap::new(),
primary_span: None,
diagnostic_help: None,
}
}

fn begin_top_level(&mut self, index: usize, hint: SemanticSpanHint) {
self.top_level_index = index;
self.span_counts.clear();
self.primary_span = Some(hint);
self.diagnostic_help = None;
}

fn mark_span(&mut self, kind: SemanticSpanKind, text: impl Into<String>) {
Expand All @@ -644,7 +647,9 @@ impl<'a> Validator<'a> {
top_level_index: self.top_level_index,
primary: self.primary_span.clone(),
note: None,
help: "fix type, mutability, scope, and control-flow errors".to_string(),
help: self.diagnostic_help.clone().unwrap_or_else(|| {
"fix type, mutability, scope, and control-flow errors".to_string()
}),
}
}

Expand Down Expand Up @@ -904,6 +909,7 @@ impl<'a> Validator<'a> {

if let Some(blocks) = else_if_blocks {
for (condition, block) in blocks.iter() {
self.mark_span(SemanticSpanKind::Keyword, "if");
self.validate_condition(condition, "else-if condition")?;
any_branch_falls_through |= self.validate_scoped_block(block)?;
}
Expand Down Expand Up @@ -1063,6 +1069,16 @@ impl<'a> Validator<'a> {
}

fn validate_condition(&mut self, expression: &Expression, context: &str) -> Result<(), String> {
if let Some(mutation) = condition_mutation(expression) {
self.diagnostic_help = Some(mutation.help().to_string());
return Err(format!(
"{} `{}` is not allowed in {}",
mutation.description(),
mutation.symbol(),
context
));
}

let ty = self.validate_expr(expression)?;
if self.is_truthy_expression(&ty) {
return Ok(());
Expand Down Expand Up @@ -2145,6 +2161,88 @@ fn is_codegen_supported_index(expression: &Expression) -> bool {
}
}

#[derive(Clone, Copy)]
enum ConditionMutation {
Assignment(&'static str),
CompoundAssignment(&'static str),
IncrementOrDecrement(&'static str),
}

impl ConditionMutation {
fn symbol(self) -> &'static str {
match self {
Self::Assignment(symbol)
| Self::CompoundAssignment(symbol)
| Self::IncrementOrDecrement(symbol) => symbol,
}
}

fn description(self) -> &'static str {
match self {
Self::Assignment(_) => "assignment",
Self::CompoundAssignment(_) => "compound assignment",
Self::IncrementOrDecrement(_) => "increment or decrement",
}
}

fn help(self) -> &'static str {
match self {
Self::Assignment(_) => {
"use `==` for comparison, or move the assignment before the condition"
}
Self::CompoundAssignment(_) | Self::IncrementOrDecrement(_) => {
"move the mutation before the condition"
}
}
}
}

fn condition_mutation(expression: &Expression) -> Option<ConditionMutation> {
match expression {
Expression::Assignment { .. } => Some(ConditionMutation::Assignment("=")),
Expression::AssignOperation { operator, .. } => {
let symbol = assign_operator_source_symbol(operator);
if matches!(operator, AssignOperator::Assign) {
Some(ConditionMutation::Assignment(symbol))
} else {
Some(ConditionMutation::CompoundAssignment(symbol))
}
}
Expression::IncDec { kind, .. } => {
Some(ConditionMutation::IncrementOrDecrement(match kind {
IncDecKind::PreInc | IncDecKind::PostInc => "++",
IncDecKind::PreDec | IncDecKind::PostDec => "--",
}))
}
Expression::StructLiteral { fields, .. } => fields
.iter()
.find_map(|(_, value)| condition_mutation(value)),
Expression::FunctionCall { args, .. } => args.iter().find_map(condition_mutation),
Expression::MethodCall { object, args, .. } => {
condition_mutation(object).or_else(|| args.iter().find_map(condition_mutation))
}
Expression::Deref(inner)
| Expression::AddressOf(inner)
| Expression::Grouped(inner)
| Expression::Unary { expr: inner, .. }
| Expression::Cast { expr: inner, .. }
| Expression::FieldAccess { object: inner, .. } => condition_mutation(inner),
Expression::BinaryExpression { left, right, .. }
| Expression::IndexAccess {
target: left,
index: right,
} => condition_mutation(left).or_else(|| condition_mutation(right)),
Expression::ArrayLiteral(values) => values.iter().find_map(condition_mutation),
Expression::AsmBlock {
inputs, outputs, ..
} => inputs
.iter()
.chain(outputs.iter())
.find_map(|(_, value)| condition_mutation(value)),
Expression::Null | Expression::Literal(_) | Expression::Variable(_) => None,
}
}

fn expression_is_true(expression: &Expression) -> bool {
matches!(expression, Expression::Literal(Literal::Bool(true)))
}
Expand Down
108 changes: 108 additions & 0 deletions tests/codegen_regressions.rs
Original file line number Diff line number Diff line change
Expand Up @@ -441,6 +441,114 @@ fn semantic_validation_rejects_backend_only_type_failures_early() {
}
}

#[test]
fn semantic_validation_rejects_mutation_in_conditions() {
let dir = temp_case_dir("condition-mutation");
let cases = [
(
"indexed_assignment_in_if.wave",
"fun main() { var command: array<char, 2> = ['x', 'y']; if (command[0] = 'h') {} }\n",
"assignment `=` is not allowed in if condition",
"use `==` for comparison, or move the assignment before the condition",
),
(
"assignment_in_else_if.wave",
"fun main() { var value: i32 = 0; if (false) {} else if (value = 1) {} }\n",
"assignment `=` is not allowed in else-if condition",
"use `==` for comparison, or move the assignment before the condition",
),
(
"compound_assignment_in_while.wave",
"fun main() { var value: i32 = 0; while (value += 1) {} }\n",
"compound assignment `+=` is not allowed in while condition",
"move the mutation before the condition",
),
(
"assignment_in_for_condition.wave",
"fun main() { for (var value: i32 = 0; value = 1; value += 1) {} }\n",
"assignment `=` is not allowed in for condition",
"use `==` for comparison, or move the assignment before the condition",
),
(
"nested_assignment_in_if.wave",
"fun main() { var value: i32 = 0; if ((value = 1) == 1) {} }\n",
"assignment `=` is not allowed in if condition",
"use `==` for comparison, or move the assignment before the condition",
),
(
"increment_in_if.wave",
"fun main() { var value: i32 = 0; if (value++) {} }\n",
"increment or decrement `++` is not allowed in if condition",
"move the mutation before the condition",
),
];

for (file_name, source, expected, help) in cases {
let source = write_wave(&dir, file_name, source);
for mode in ["check", "build"] {
let output = if mode == "check" {
run_wavec_raw([OsStr::new("check"), source.as_os_str()])
} else {
run_wavec_raw([
OsStr::new("build"),
source.as_os_str(),
OsStr::new("--emit=obj"),
OsStr::new("--out-dir"),
dir.as_os_str(),
])
};
assert!(
!output.status.success(),
"{} unexpectedly accepted mutation in a condition in {} mode",
file_name,
mode
);
let stderr = String::from_utf8_lossy(&output.stderr);
assert!(stderr.contains("error[E3001]"), "{}: {}", file_name, stderr);
assert!(stderr.contains(expected), "{}: {}", file_name, stderr);
assert!(stderr.contains(help), "{}: {}", file_name, stderr);
assert!(
!stderr.contains("E9001") && !stderr.contains("compiler internal error"),
"{} ({}) leaked a backend failure: {}",
file_name,
mode,
stderr
);
}
}
}

#[test]
fn semantic_validation_allows_mutation_outside_conditions() {
let dir = temp_case_dir("condition-mutation-valid");
let source = write_wave(
&dir,
"valid.wave",
r#"
fun main() {
var value: i32 = 0;
value = 1;
if (value == 1) {}

while (value < 2) {
value += 1;
}

for (var index: i32 = 0; index < 2; index += 1) {}
}
"#,
);

run_wavec([OsStr::new("check"), source.as_os_str()]);
run_wavec([
OsStr::new("build"),
source.as_os_str(),
OsStr::new("--emit=obj"),
OsStr::new("--out-dir"),
dir.as_os_str(),
]);
}

#[test]
fn second_semantic_audit_rejects_unsafe_programs_before_codegen() {
let dir = temp_case_dir("semantic-audit-two");
Expand Down
Loading