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
9 changes: 8 additions & 1 deletion .github/workflows/rust.yml
Original file line number Diff line number Diff line change
Expand Up @@ -861,7 +861,7 @@ jobs:
env:
WAVE_RUNTIME_ARTIFACT_DIR: ${{ runner.temp }}
run: >-
cargo test --locked --no-fail-fast --test runtime_regressions --test std_io_regressions --test native_providers --target aarch64-pc-windows-msvc
cargo test --locked --no-fail-fast --test runtime_regressions --test std_io_regressions --test native_providers --test stabilization_17 --target aarch64-pc-windows-msvc
--no-default-features --features llvm-target-aarch64 --jobs 2

- name: Save failed native runtime compiler and fixture
Expand Down Expand Up @@ -1014,6 +1014,13 @@ jobs:
cargo test --locked --no-default-features --features llvm-target-wasm
--test release_stabilization webassembly_ --jobs 2

- name: Run checked numeric conversion and shift regressions
env:
WAVE_RUN_WASM_RUNTIME_TESTS: "1"
run: >-
cargo test --locked --no-default-features --features llvm-target-wasm
--test stabilization_17 webassembly_ --jobs 2

- name: Build and run browser and WASI modules
run: |
set -euo pipefail
Expand Down
12 changes: 12 additions & 0 deletions Cargo.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

2 changes: 2 additions & 0 deletions Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@ path = "src/lib.rs"
utils = { path = "utils" }
lexer = { path = "front/lexer" }
parser = { path = "front/parser" }
hir = { path = "front/hir" }
error = { path = "front/error" }
llvm = { path = "./llvm", default-features = false }

Expand All @@ -28,6 +29,7 @@ llvm-target-x86 = ["llvm/llvm-target-x86"]
members = [
"front/lexer",
"front/parser",
"front/hir",
"llvm",
"front/error",
"utils",
Expand Down
4 changes: 4 additions & 0 deletions MAINTAINERS
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,10 @@ F: front/lexer/
M: luna@lunastev.org
F: front/parser/

[Compiler Frontend - HIR]
M: luna@lunastev.org
F: front/hir/

[Error System]
M: luna@lunastev.org
F: front/error/
Expand Down
10 changes: 10 additions & 0 deletions front/hir/Cargo.toml
Original file line number Diff line number Diff line change
@@ -0,0 +1,10 @@
[package]
name = "hir"
version = "0.1.0"
edition = "2021"

[dependencies]
parser = { path = "../parser" }
lexer = { path = "../lexer" }
error = { path = "../error" }
utils = { path = "../../utils" }
Original file line number Diff line number Diff line change
Expand Up @@ -2,8 +2,8 @@
//!
//! Source expressions are evaluated in order into frame slots. A pending await
//! returns from the resume function; re-entry selects only its saved state.
use crate::ast::*;
use crate::hir::{HirExpressionType, TypedProgram};
use crate::{HirExpressionType, TypedProgram};
use parser::ast::*;
use std::collections::HashMap;

#[derive(Debug, Clone)]
Expand Down
178 changes: 174 additions & 4 deletions front/parser/src/hir/conversions.rs → front/hir/src/conversions.rs
Original file line number Diff line number Diff line change
Expand Up @@ -14,8 +14,8 @@
//! conversions, including equal-width signedness changes and explicit identity
//! casts. ABI storage/transport conversion is a separate backend boundary.
use super::{ExpressionId, HirExpressionType, TypedProgram};
use crate::ast::{AssignOperator, Expression, Literal, Operator, WaveType};
use error::SourceSpan;
use parser::ast::{AssignOperator, Expression, Literal, Operator, WaveType};
use std::collections::HashSet;

#[derive(Clone, Copy, Debug, PartialEq, Eq)]
Expand All @@ -30,8 +30,11 @@ pub enum ConversionKind {
SignExtend,
ZeroExtend,
Truncate,
IntegerToBool,
FloatToBool,
SignedToFloat,
UnsignedToFloat,
// Checked, truncation toward zero; invalid values trap.
FloatToSigned,
FloatToUnsigned,
FloatExtend,
Expand All @@ -58,6 +61,8 @@ pub struct NumericExpressionInfo {
/// Required destination at this occurrence, absent for an unconstrained value.
pub context_type: Option<WaveType>,
pub conversions: Vec<ConversionInfo>,
/// Original count type; checked before conversion to the LHS width.
pub shift_count_type: Option<WaveType>,
}
#[derive(Clone, Debug)]
pub struct ConversionError {
Expand Down Expand Up @@ -95,6 +100,15 @@ pub fn conversion_kind(source: &WaveType, target: &WaveType) -> Option<Conversio
if source == target {
return Some(Identity);
}
if target == &WaveType::Bool {
return if integer_width(source).is_some() {
Some(IntegerToBool)
} else if matches!(source, WaveType::Float(_)) {
Some(FloatToBool)
} else {
None
};
}
if let (Some(a), Some(b)) = (integer_width(source), integer_width(target)) {
return Some(if a == b {
ReinterpretInteger
Expand Down Expand Up @@ -127,6 +141,47 @@ pub fn conversion_kind(source: &WaveType, target: &WaveType) -> Option<Conversio
}
}

/// Source-float bounds equivalent to truncation toward zero followed by the
/// destination integer range check. Shared by constant checking and lowering.
pub struct FloatIntegerRange {
pub lower: f64,
pub lower_inclusive: bool,
pub upper: f64,
}
impl FloatIntegerRange {
pub fn contains(&self, value: f64) -> bool {
value.is_finite()
&& value < self.upper
&& if self.lower_inclusive {
value >= self.lower
} else {
value > self.lower
}
}
}
pub fn float_integer_range(source: &WaveType, target: &WaveType) -> FloatIntegerRange {
let WaveType::Float(source_bits) = source else {
panic!("non-float conversion source")
};
let round = |v: f64| {
if *source_bits == 32 {
(v as f32) as f64
} else {
v
}
};
let signed = !unsigned(target);
let bits = integer_width(target).expect("integer conversion destination");
let upper = round(2f64.powi(i32::from(bits - u16::from(signed))));
let minimum = if signed { -upper } else { 0.0 };
let before_minimum = round(minimum - 1.0);
FloatIntegerRange {
lower: before_minimum,
lower_inclusive: before_minimum == minimum,
upper,
}
}

impl NumericExpressionInfo {
pub fn verify(&self) -> Result<(), String> {
let mut previous = &self.evaluation_type;
Expand Down Expand Up @@ -272,6 +327,7 @@ impl Planner<'_> {
result_type: source,
context_type: demand.clone().filter(scalar),
conversions: Vec::new(),
shift_count_type: None,
};
match expr {
Expression::Cast {
Expand All @@ -282,7 +338,9 @@ impl Planner<'_> {
Some(target_type.clone())
} else {
literal(inner).then(|| {
if pointer(target_type) {
if target_type == &WaveType::Bool {
WaveType::Int(1024)
} else if pointer(target_type) {
WaveType::Int(64)
} else {
target_type.clone()
Expand Down Expand Up @@ -351,6 +409,16 @@ impl Planner<'_> {
Some(if numeric(&rt) { WaveType::Int(64) } else { rt }),
None,
);
} else if matches!(operator, Operator::ShiftLeft | Operator::ShiftRight) {
let left_type = self.plan(left, None, left_hint.cloned())?;
// An untyped count must not wrap at the default i32 width
// before the checked-shift range test.
let count_hint = literal(right).then_some(WaveType::Int(1024));
let right_type = self.plan(right, None, count_hint)?;
fact.evaluation_type = left_type.clone();
fact.result_type = left_type.clone();
fact.computation_type = Some(left_type);
fact.shift_count_type = Some(right_type);
} else if numeric(&lt) && numeric(&rt) {
if matches!(operator, Operator::LogicalAnd | Operator::LogicalOr) {
self.plan(left, None, None);
Expand Down Expand Up @@ -432,6 +500,17 @@ pub(super) fn verify_expression(
fact: &NumericExpressionInfo,
) -> Result<(), String> {
fact.verify()?;
if fact.shift_count_type.is_some()
!= matches!(
expr,
Expression::BinaryExpression {
operator: Operator::ShiftLeft | Operator::ShiftRight,
..
}
)
{
return Err("missing or extraneous shift count contract".into());
}
let child = |expr: &Expression| {
program
.numeric_expression_of(expr)
Expand Down Expand Up @@ -487,6 +566,20 @@ pub(super) fn verify_expression(
{
return Err("invalid logical computation".into());
}
} else if matches!(operator, Operator::ShiftLeft | Operator::ShiftRight) {
if &left.result_type != computation
|| fact.evaluation_type != *computation
|| fact.shift_count_type.as_ref() != Some(&right.result_type)
|| !matches!(
right.result_type,
WaveType::Int(_)
| WaveType::Uint(_)
| WaveType::Byte
| WaveType::Char
)
{
return Err("invalid checked shift contract".into());
}
} else {
if &left.result_type != computation || &right.result_type != computation {
return Err("operand conversion does not reach computation type".into());
Expand Down Expand Up @@ -541,12 +634,12 @@ pub(super) fn verify_expression(
#[cfg(test)]
mod tests {
use super::*;
use crate::ast::ASTNode;
use parser::ast::ASTNode;
fn program(source: &str) -> TypedProgram {
let tokens = lexer::Lexer::new_with_file(source, "conversions.wave")
.tokenize()
.unwrap();
TypedProgram::lower(crate::parse_syntax_with_spans(&tokens).unwrap()).unwrap()
TypedProgram::lower(parser::parse_syntax_with_spans(&tokens).unwrap()).unwrap()
}
fn initializer(p: &TypedProgram, index: usize) -> &Expression {
let ASTNode::Function(f) = &p.syntax()[0] else {
Expand All @@ -557,6 +650,83 @@ mod tests {
};
v.initial_value.as_ref().unwrap()
}
#[test]
fn checked_ranges_match_exact_integer_boundaries() {
use utils::const_int::ConstInt;
for source_bits in [32, 64] {
for bits in [8, 16, 32, 64, 128, 256, 512, 1024] {
for signed in [false, true] {
let target = if signed {
WaveType::Int(bits)
} else {
WaveType::Uint(bits)
};
let range = float_integer_range(&WaveType::Float(source_bits), &target);
let edge = 2f64.powi(i32::from(bits - u16::from(signed)));
let mut values = vec![
0.0,
-0.0,
-0.5,
-1.0,
f64::NAN,
f64::INFINITY,
f64::NEG_INFINITY,
edge,
-edge,
edge - 1.0,
-edge - 1.0,
];
if source_bits == 32 {
let raw = (edge as f32).to_bits();
for delta in [raw.saturating_sub(1), raw, raw.saturating_add(1)] {
let n = f32::from_bits(delta) as f64;
values.extend([n, -n]);
}
} else {
for delta in [
edge.to_bits().saturating_sub(1),
edge.to_bits(),
edge.to_bits().saturating_add(1),
] {
let n = f64::from_bits(delta);
values.extend([n, -n]);
}
}
for value in values {
let value = if source_bits == 32 {
(value as f32) as f64
} else {
value
};
let exact = ConstInt::from_f64(value).is_some_and(|n| n.fits(bits, signed));
assert_eq!(
range.contains(value),
exact,
"f{source_bits} to {target:?}: {value}"
);
}
}
}
}
}

#[test]
fn shifts_preserve_the_original_count_and_bool_casts_compare_with_zero() {
let p =
program("fun check(x: u8, n: u64) { var y: u8 = x << n; var b: bool = n as bool; }");
p.verify_conversions().unwrap();
let shift = p.numeric_expression_of(initializer(&p, 0)).unwrap();
assert_eq!(shift.shift_count_type, Some(WaveType::Uint(64)));
assert_eq!(shift.computation_type, Some(WaveType::Uint(8)));
assert_eq!(
p.numeric_expression_of(initializer(&p, 1))
.unwrap()
.conversions[0]
.kind,
ConversionKind::IntegerToBool
);
}

#[test]
fn ordered_casts_preserve_truncation_signedness_modes_and_spans() {
let source = "fun check(x: i32) { var y: i64 = x as u8 as i32; var z: u32 = x as u32; var identity: i32 = x as i32; }";
Expand Down
Loading
Loading