use crate::error::{ErrorKind, Ordinal, Result};
use crate::Context as OuterContext;
use std::mem;
use wain_ast::source::Source;
use wain_ast::*;
#[derive(Copy, Clone)]
enum Type {
Known(ValType),
Unknown,
}
impl Type {
fn i32() -> Type {
Type::Known(ValType::I32)
}
fn i64() -> Type {
Type::Known(ValType::I64)
}
fn f32() -> Type {
Type::Known(ValType::F32)
}
fn f64() -> Type {
Type::Known(ValType::F64)
}
}
struct CtrlFrame {
idx: usize,
offset: usize,
}
struct FuncBodyContext<'outer, 'module: 'outer, 'source: 'module, S: Source> {
current_op: &'static str,
current_offset: usize,
outer: &'outer OuterContext<'module, 'source, S>,
op_stack: Vec<Type>,
current_frame: CtrlFrame,
label_stack: Vec<Option<ValType>>,
params: &'outer [ValType],
locals: &'outer [ValType],
ret_ty: Option<ValType>,
unreachable: bool,
}
impl<'outer, 'm, 's, S: Source> FuncBodyContext<'outer, 'm, 's, S> {
fn error<T>(&self, kind: ErrorKind) -> Result<T, S> {
self.outer.error(kind, self.current_op, self.current_offset)
}
fn ensure_ctrl_frame_not_empty(&self) -> Result<(), S> {
if self.op_stack.len() > self.current_frame.idx {
return Ok(());
}
if self.unreachable {
return Ok(());
}
self.error(ErrorKind::CtrlFrameEmpty {
op: self.current_op,
frame_start: self.current_frame.offset,
idx_in_op_stack: self.current_frame.idx,
})
}
fn ensure_op_stack_top(&self, expected: Type) -> Result<Type, S> {
self.ensure_ctrl_frame_not_empty()?;
if self.op_stack.len() == self.current_frame.idx {
assert!(self.unreachable);
return Ok(Type::Unknown);
}
let actual = self.op_stack[self.op_stack.len() - 1];
if let (Type::Known(expected), Type::Known(actual)) = (expected, actual) {
if actual != expected {
return self.error(ErrorKind::TypeMismatch { expected, actual });
}
}
Ok(actual)
}
fn pop_op_stack(&mut self, expected: Type) -> Result<Type, S> {
let ty = self.ensure_op_stack_top(expected)?;
self.op_stack.pop();
Ok(ty)
}
fn push_control_frame(&mut self, offset: usize) -> CtrlFrame {
let idx = self.op_stack.len();
let new = CtrlFrame { idx, offset };
mem::replace(&mut self.current_frame, new)
}
fn pop_control_frame(&mut self, prev: CtrlFrame) {
assert!(self.current_frame.idx <= self.op_stack.len());
self.current_frame = prev;
}
fn pop_label_stack(&mut self) -> Result<(), S> {
if let Some(ty) = self.label_stack.pop() {
if let Some(ty) = ty {
self.ensure_op_stack_top(Type::Known(ty))?;
}
Ok(())
} else {
self.error(ErrorKind::LabelStackEmpty {
op: self.current_op,
})
}
}
fn validate_label_idx(&self, idx: u32) -> Result<Option<ValType>, S> {
let len = self.label_stack.len();
if (idx as usize) >= len {
return self.error(ErrorKind::IndexOutOfBounds {
idx,
upper: len,
what: "label",
});
}
let ty = self.label_stack[len - 1 - (idx as usize)];
if let Some(ty) = ty {
self.ensure_op_stack_top(Type::Known(ty))?;
}
Ok(ty)
}
fn validate_local_idx(&self, idx: u32) -> Result<ValType, S> {
let uidx = idx as usize;
if let Some(ty) = self.params.get(uidx) {
return Ok(*ty);
}
if let Some(ty) = self.locals.get(uidx - self.params.len()) {
Ok(*ty)
} else {
self.error(ErrorKind::IndexOutOfBounds {
idx,
upper: self.locals.len(),
what: "local variable",
})
.map_err(|e| {
e.update_msg(format!(
"access to {} local variable at {}",
Ordinal(uidx - self.params.len()),
self.current_op
))
})
}
}
fn validate_memarg(&self, mem: &Mem, bits: u8) -> Result<(), S> {
self.outer
.memory_from_idx(0, self.current_op, self.current_offset)?;
if let Some(align) = mem.align {
if align > bits / 8 {
return self.error(ErrorKind::TooLargeAlign { align, bits });
}
}
Ok(())
}
fn validate_load(&mut self, mem: &Mem, bits: u8, ty: ValType) -> Result<(), S> {
self.validate_memarg(mem, bits)?;
self.pop_op_stack(Type::i32())?; self.op_stack.push(Type::Known(ty));
Ok(())
}
fn validate_store(&mut self, mem: &Mem, bits: u8, ty: ValType) -> Result<(), S> {
self.validate_memarg(mem, bits)?;
self.pop_op_stack(Type::Known(ty))?; self.pop_op_stack(Type::i32())?; Ok(())
}
fn validate_convert(&mut self, from: ValType, to: ValType) -> Result<(), S> {
self.pop_op_stack(Type::Known(from))?;
self.op_stack.push(Type::Known(to));
Ok(())
}
}
pub(crate) fn validate_func_body<'outer, 'm, 's, S: Source>(
body: &'outer [Instruction],
func_ty: &'outer FuncType,
locals: &'outer [ValType],
outer: &'outer OuterContext<'m, 's, S>,
start: usize,
) -> Result<(), S> {
let ret_ty = func_ty.results.get(0).copied();
let mut ctx = FuncBodyContext {
current_op: "",
current_offset: start,
outer,
op_stack: vec![],
label_stack: vec![],
current_frame: CtrlFrame {
idx: 0,
offset: start,
},
params: &func_ty.params,
locals,
ret_ty,
unreachable: false,
};
body.validate(&mut ctx)?;
if let Some(ty) = ret_ty {
ctx.current_op = "function return type";
ctx.current_offset = start;
ctx.ensure_op_stack_top(Type::Known(ty))?;
}
Ok(())
}
trait ValidateInsnSeq<'outer, 'm, 's, S: Source> {
fn validate(&self, ctx: &mut FuncBodyContext<'outer, 'm, 's, S>) -> Result<(), S>;
}
impl<'s, 'm, 'outer, S: Source, V: ValidateInsnSeq<'outer, 'm, 's, S>>
ValidateInsnSeq<'outer, 'm, 's, S> for [V]
{
fn validate(&self, ctx: &mut FuncBodyContext<'outer, 'm, 's, S>) -> Result<(), S> {
self.iter()
.map(|insn| insn.validate(ctx))
.collect::<Result<_, _>>()?;
ctx.unreachable = false; Ok(())
}
}
impl<'outer, 'm, 's, S: Source> ValidateInsnSeq<'outer, 'm, 's, S> for Instruction {
fn validate(&self, ctx: &mut FuncBodyContext<'outer, 'm, 's, S>) -> Result<(), S> {
ctx.current_op = self.kind.name();
ctx.current_offset = self.start;
let start = self.start;
use InsnKind::*;
match &self.kind {
Block { ty, body } => {
let saved = ctx.push_control_frame(start);
ctx.label_stack.push(*ty);
body.validate(ctx)?;
ctx.pop_label_stack()?;
ctx.pop_control_frame(saved);
if let Some(ty) = ty {
ctx.ensure_op_stack_top(Type::Known(*ty))?;
}
}
Loop { ty, body } => {
let saved = ctx.push_control_frame(start);
ctx.label_stack.push(*ty);
body.validate(ctx)?;
ctx.pop_label_stack()?;
ctx.pop_control_frame(saved);
if let Some(ty) = ty {
ctx.ensure_op_stack_top(Type::Known(*ty))?;
}
}
If {
ty,
then_body,
else_body,
} => {
ctx.pop_op_stack(Type::i32())?;
ctx.label_stack.push(*ty);
let saved = ctx.push_control_frame(start);
then_body.validate(ctx)?;
if let Some(ty) = ty {
ctx.ensure_op_stack_top(Type::Known(*ty))?;
}
ctx.pop_control_frame(saved);
let saved = ctx.push_control_frame(start);
else_body.validate(ctx)?;
if let Some(ty) = ty {
ctx.ensure_op_stack_top(Type::Known(*ty))?;
}
ctx.pop_control_frame(saved);
ctx.pop_label_stack()?;
if let Some(ty) = ty {
ctx.ensure_op_stack_top(Type::Known(*ty))?;
}
}
Unreachable => ctx.unreachable = true,
Nop => {}
Br(labelidx) => {
ctx.validate_label_idx(*labelidx)?;
ctx.unreachable = true;
}
BrIf(labelidx) => {
ctx.pop_op_stack(Type::i32())?;
ctx.validate_label_idx(*labelidx)?;
}
BrTable {
labels,
default_label,
} => {
let expected = ctx.validate_label_idx(*default_label)?;
for (i, idx) in labels.iter().enumerate() {
let ty = ctx.validate_label_idx(*idx)?;
if let (Some(l), Some(r)) = (&expected, &ty) {
if l != r {
return ctx
.error(ErrorKind::TypeMismatch {
expected: *l,
actual: *r,
})
.map_err(|e| {
e.update_msg(format!(
"{} label {} at {}",
Ordinal(i),
idx,
ctx.current_op
))
});
}
}
}
ctx.unreachable = true;
}
Return => {
if let Some(ty) = ctx.ret_ty {
ctx.ensure_op_stack_top(Type::Known(ty))?;
}
ctx.unreachable = true;
}
Call(funcidx) => {
let func = ctx.outer.func_from_idx(*funcidx, ctx.current_op, start)?;
let fty = &ctx.outer.module.types[func.idx as usize];
for (i, ty) in fty.params.iter().enumerate().rev() {
ctx.pop_op_stack(Type::Known(*ty))
.map_err(|e| e.update_msg(format!("{} parameter at call", Ordinal(i))))?;
}
for ty in fty.results.iter() {
ctx.op_stack.push(Type::Known(*ty));
}
}
CallIndirect(typeidx) => {
ctx.outer.table_from_idx(0, ctx.current_op, start)?;
ctx.pop_op_stack(Type::i32())?;
let fty = ctx.outer.type_from_idx(*typeidx, ctx.current_op, start)?;
for (i, ty) in fty.params.iter().enumerate().rev() {
ctx.pop_op_stack(Type::Known(*ty)).map_err(|e| {
e.update_msg(format!("{} parameter at call.indirect", Ordinal(i)))
})?;
}
for ty in fty.results.iter() {
ctx.op_stack.push(Type::Known(*ty));
}
}
Drop => {
ctx.pop_op_stack(Type::Unknown)?;
}
Select => {
ctx.pop_op_stack(Type::Unknown)?;
ctx.pop_op_stack(Type::Unknown)?;
ctx.pop_op_stack(Type::i32())?;
ctx.op_stack.push(Type::Unknown);
}
LocalGet(localidx) => {
let ty = ctx.validate_local_idx(*localidx)?;
ctx.op_stack.push(Type::Known(ty));
}
LocalSet(localidx) => {
let ty = Type::Known(ctx.validate_local_idx(*localidx)?);
ctx.pop_op_stack(ty)?;
}
LocalTee(localidx) => {
let ty = Type::Known(ctx.validate_local_idx(*localidx)?);
ctx.ensure_op_stack_top(ty)?;
}
GlobalGet(globalidx) => {
let global = ctx
.outer
.global_from_idx(*globalidx, ctx.current_op, start)?;
ctx.op_stack.push(Type::Known(global.ty));
}
GlobalSet(globalidx) => {
let global = ctx
.outer
.global_from_idx(*globalidx, ctx.current_op, start)?;
let ty = Type::Known(global.ty);
if !global.mutable {
return ctx.error(ErrorKind::SetImmutableGlobal {
ty: global.ty,
idx: *globalidx,
});
}
ctx.pop_op_stack(ty)?;
}
I32Load(mem) => ctx.validate_load(mem, 32, ValType::I32)?,
I64Load(mem) => ctx.validate_load(mem, 64, ValType::I64)?,
F32Load(mem) => ctx.validate_load(mem, 32, ValType::F32)?,
F64Load(mem) => ctx.validate_load(mem, 64, ValType::F64)?,
I32Load8S(mem) => ctx.validate_load(mem, 8, ValType::I32)?,
I32Load8U(mem) => ctx.validate_load(mem, 8, ValType::I32)?,
I32Load16S(mem) => ctx.validate_load(mem, 16, ValType::I32)?,
I32Load16U(mem) => ctx.validate_load(mem, 16, ValType::I32)?,
I64Load8S(mem) => ctx.validate_load(mem, 8, ValType::I64)?,
I64Load8U(mem) => ctx.validate_load(mem, 8, ValType::I64)?,
I64Load16S(mem) => ctx.validate_load(mem, 16, ValType::I64)?,
I64Load16U(mem) => ctx.validate_load(mem, 16, ValType::I64)?,
I64Load32S(mem) => ctx.validate_load(mem, 32, ValType::I64)?,
I64Load32U(mem) => ctx.validate_load(mem, 32, ValType::I64)?,
I32Store(mem) => ctx.validate_store(mem, 32, ValType::I32)?,
I64Store(mem) => ctx.validate_store(mem, 64, ValType::I64)?,
F32Store(mem) => ctx.validate_store(mem, 32, ValType::F32)?,
F64Store(mem) => ctx.validate_store(mem, 64, ValType::F64)?,
I32Store8(mem) => ctx.validate_store(mem, 8, ValType::I32)?,
I32Store16(mem) => ctx.validate_store(mem, 16, ValType::I32)?,
I64Store8(mem) => ctx.validate_store(mem, 8, ValType::I64)?,
I64Store16(mem) => ctx.validate_store(mem, 16, ValType::I64)?,
I64Store32(mem) => ctx.validate_store(mem, 32, ValType::I64)?,
MemorySize => {
if ctx.outer.module.memories.is_empty() {
return ctx.error(ErrorKind::MemoryIsNotDefined);
}
ctx.op_stack.push(Type::i32());
}
MemoryGrow => {
if ctx.outer.module.memories.is_empty() {
return ctx.error(ErrorKind::MemoryIsNotDefined);
}
ctx.ensure_op_stack_top(Type::i32())?;
}
I32Const(_) => {
ctx.op_stack.push(Type::i32());
}
I64Const(_) => {
ctx.op_stack.push(Type::i64());
}
F32Const(_) => {
ctx.op_stack.push(Type::f32());
}
F64Const(_) => {
ctx.op_stack.push(Type::f64());
}
I32Clz | I32Ctz | I32Popcnt => {
ctx.ensure_op_stack_top(Type::i32())?;
}
I64Clz | I64Ctz | I64Popcnt => {
ctx.ensure_op_stack_top(Type::i64())?;
}
F32Abs | F32Neg | F32Ceil | F32Floor | F32Trunc | F32Nearest | F32Sqrt => {
ctx.ensure_op_stack_top(Type::f32())?;
}
F64Abs | F64Neg | F64Ceil | F64Floor | F64Trunc | F64Nearest | F64Sqrt => {
ctx.ensure_op_stack_top(Type::f64())?;
}
I32Add | I32Sub | I32Mul | I32DivS | I32DivU | I32RemS | I32RemU | I32And | I32Or
| I32Xor | I32Shl | I32ShrS | I32ShrU | I32Rotl | I32Rotr => {
ctx.pop_op_stack(Type::i32())?;
ctx.ensure_op_stack_top(Type::i32())?;
}
I64Add | I64Sub | I64Mul | I64DivS | I64DivU | I64RemS | I64RemU | I64And | I64Or
| I64Xor | I64Shl | I64ShrS | I64ShrU | I64Rotl | I64Rotr => {
ctx.pop_op_stack(Type::i64())?;
ctx.ensure_op_stack_top(Type::i64())?;
}
F32Add | F32Sub | F32Mul | F32Div | F32Min | F32Max | F32Copysign => {
ctx.pop_op_stack(Type::f32())?;
ctx.ensure_op_stack_top(Type::f32())?;
}
F64Add | F64Sub | F64Mul | F64Div | F64Min | F64Max | F64Copysign => {
ctx.pop_op_stack(Type::f64())?;
ctx.ensure_op_stack_top(Type::f64())?;
}
I32Eqz => {
ctx.ensure_op_stack_top(Type::i32())?;
}
I64Eqz => {
ctx.pop_op_stack(Type::i64())?;
ctx.op_stack.push(Type::i32());
}
I32Eq | I32Ne | I32LtS | I32LtU | I32GtS | I32GtU | I32LeS | I32LeU | I32GeS
| I32GeU => {
ctx.pop_op_stack(Type::i32())?;
ctx.ensure_op_stack_top(Type::i32())?;
}
I64Eq | I64Ne | I64LtS | I64LtU | I64GtS | I64GtU | I64LeS | I64LeU | I64GeS
| I64GeU => {
ctx.pop_op_stack(Type::i64())?;
ctx.pop_op_stack(Type::i64())?;
ctx.op_stack.push(Type::i32());
}
F32Eq | F32Ne | F32Lt | F32Gt | F32Le | F32Ge => {
ctx.pop_op_stack(Type::f32())?;
ctx.pop_op_stack(Type::f32())?;
ctx.op_stack.push(Type::i32());
}
F64Eq | F64Ne | F64Lt | F64Gt | F64Le | F64Ge => {
ctx.pop_op_stack(Type::f64())?;
ctx.pop_op_stack(Type::f64())?;
ctx.op_stack.push(Type::i32());
}
I32WrapI64 => ctx.validate_convert(ValType::I64, ValType::I32)?,
I32TruncF32S => ctx.validate_convert(ValType::F32, ValType::I32)?,
I32TruncF32U => ctx.validate_convert(ValType::F32, ValType::I32)?,
I32TruncF64S => ctx.validate_convert(ValType::F64, ValType::I32)?,
I32TruncF64U => ctx.validate_convert(ValType::F64, ValType::I32)?,
I64ExtendI32S => ctx.validate_convert(ValType::I32, ValType::I64)?,
I64ExtendI32U => ctx.validate_convert(ValType::I32, ValType::I64)?,
I64TruncF32S => ctx.validate_convert(ValType::F32, ValType::I64)?,
I64TruncF32U => ctx.validate_convert(ValType::F32, ValType::I64)?,
I64TruncF64S => ctx.validate_convert(ValType::F64, ValType::I64)?,
I64TruncF64U => ctx.validate_convert(ValType::F64, ValType::I64)?,
F32ConvertI32S => ctx.validate_convert(ValType::I32, ValType::F32)?,
F32ConvertI32U => ctx.validate_convert(ValType::I32, ValType::F32)?,
F32ConvertI64S => ctx.validate_convert(ValType::I64, ValType::F32)?,
F32ConvertI64U => ctx.validate_convert(ValType::I64, ValType::F32)?,
F32DemoteF64 => ctx.validate_convert(ValType::F64, ValType::F32)?,
F64ConvertI32S => ctx.validate_convert(ValType::I32, ValType::F64)?,
F64ConvertI32U => ctx.validate_convert(ValType::I32, ValType::F64)?,
F64ConvertI64S => ctx.validate_convert(ValType::I64, ValType::F64)?,
F64ConvertI64U => ctx.validate_convert(ValType::I64, ValType::F64)?,
F64PromoteF32 => ctx.validate_convert(ValType::F32, ValType::F64)?,
I32ReinterpretF32 => ctx.validate_convert(ValType::F32, ValType::I32)?,
I64ReinterpretF64 => ctx.validate_convert(ValType::F64, ValType::I64)?,
F32ReinterpretI32 => ctx.validate_convert(ValType::I32, ValType::F32)?,
F64ReinterpretI64 => ctx.validate_convert(ValType::I64, ValType::F64)?,
}
Ok(())
}
}
pub(crate) fn validate_constant<'m, 's, S: Source>(
insns: &[Instruction],
ctx: &OuterContext<'m, 's, S>,
expr_ty: ValType,
when: &'static str,
start: usize,
) -> Result<(), S> {
let mut last_ty = None;
for insn in insns {
let name = insn.kind.name();
use InsnKind::*;
match &insn.kind {
GlobalGet(globalidx) => {
if let Some(global) = ctx.module.globals.get(*globalidx as usize) {
last_ty = Some(global.ty);
} else {
return ctx
.error(
ErrorKind::IndexOutOfBounds {
idx: *globalidx,
upper: ctx.module.globals.len(),
what: "global variable read",
},
"",
insn.start,
)
.map_err(|e| {
e.update_msg(format!("constant expression in {} at {}", name, when))
});
}
}
I32Const(_) => last_ty = Some(ValType::I32),
I64Const(_) => last_ty = Some(ValType::I64),
F32Const(_) => last_ty = Some(ValType::F32),
F64Const(_) => last_ty = Some(ValType::F64),
_ => {
return ctx
.error(ErrorKind::NotConstantInstruction(name), "", insn.start)
.map_err(|e| e.update_msg(format!("constant expression at {}", when)));
}
}
}
if let Some(ty) = last_ty {
if ty != expr_ty {
ctx.error(
ErrorKind::TypeMismatch {
expected: expr_ty,
actual: ty,
},
"",
start,
)
.map_err(|e| e.update_msg(format!("type of constant expression at {}", when)))
} else {
Ok(())
}
} else {
ctx.error(ErrorKind::NoInstructionForConstant, when, start)
}
}