use crate::trap::{Trap, TrapCode};
use crate::vmcontext::{VMCallerCheckedAnyfunc, VMTableDefinition};
use serde::{Deserialize, Serialize};
use std::borrow::{Borrow, BorrowMut};
use std::cell::UnsafeCell;
use std::convert::TryFrom;
use std::fmt;
use std::ptr::NonNull;
use std::sync::Mutex;
use wasmer_types::{TableType, Type as ValType};
#[derive(Debug, Clone, Hash, Serialize, Deserialize)]
pub enum TableStyle {
CallerChecksSignature,
}
pub trait Table: fmt::Debug + Send + Sync {
fn style(&self) -> &TableStyle;
fn ty(&self) -> &TableType;
fn size(&self) -> u32;
fn grow(&self, delta: u32) -> Option<u32>;
fn get(&self, index: u32) -> Option<VMCallerCheckedAnyfunc>;
fn set(&self, index: u32, func: VMCallerCheckedAnyfunc) -> Result<(), Trap>;
fn vmtable(&self) -> NonNull<VMTableDefinition>;
fn copy(
&self,
src_table: &dyn Table,
dst_index: u32,
src_index: u32,
len: u32,
) -> Result<(), Trap> {
if src_index
.checked_add(len)
.map_or(true, |n| n > src_table.size())
{
return Err(Trap::new_from_runtime(TrapCode::TableAccessOutOfBounds));
}
if dst_index.checked_add(len).map_or(true, |m| m > self.size()) {
return Err(Trap::new_from_runtime(TrapCode::TableSetterOutOfBounds));
}
let srcs = src_index..src_index + len;
let dsts = dst_index..dst_index + len;
if dst_index <= src_index {
for (s, d) in (srcs).zip(dsts) {
self.set(d, src_table.get(s).unwrap())?;
}
} else {
for (s, d) in srcs.rev().zip(dsts.rev()) {
self.set(d, src_table.get(s).unwrap())?;
}
}
Ok(())
}
}
#[derive(Debug)]
pub struct LinearTable {
vec: Mutex<Vec<VMCallerCheckedAnyfunc>>,
maximum: Option<u32>,
table: TableType,
style: TableStyle,
vm_table_definition: VMTableDefinitionOwnership,
}
#[derive(Debug)]
enum VMTableDefinitionOwnership {
VMOwned(NonNull<VMTableDefinition>),
HostOwned(Box<UnsafeCell<VMTableDefinition>>),
}
unsafe impl Send for LinearTable {}
unsafe impl Sync for LinearTable {}
impl LinearTable {
pub fn new(table: &TableType, style: &TableStyle) -> Result<Self, String> {
unsafe { Self::new_inner(table, style, None) }
}
pub unsafe fn from_definition(
table: &TableType,
style: &TableStyle,
vm_table_location: NonNull<VMTableDefinition>,
) -> Result<Self, String> {
Self::new_inner(table, style, Some(vm_table_location))
}
unsafe fn new_inner(
table: &TableType,
style: &TableStyle,
vm_table_location: Option<NonNull<VMTableDefinition>>,
) -> Result<Self, String> {
match table.ty {
ValType::FuncRef => (),
ty => return Err(format!("tables of types other than anyfunc ({})", ty)),
};
if let Some(max) = table.maximum {
if max < table.minimum {
return Err(format!(
"Table minimum ({}) is larger than maximum ({})!",
table.minimum, max
));
}
}
let table_minimum = usize::try_from(table.minimum)
.map_err(|_| "Table minimum is bigger than usize".to_string())?;
let mut vec = vec![VMCallerCheckedAnyfunc::default(); table_minimum];
let base = vec.as_mut_ptr();
match style {
TableStyle::CallerChecksSignature => Ok(Self {
vec: Mutex::new(vec),
maximum: table.maximum,
table: *table,
style: style.clone(),
vm_table_definition: if let Some(table_loc) = vm_table_location {
{
let mut ptr = table_loc;
let td = ptr.as_mut();
td.base = base as _;
td.current_elements = table_minimum as _;
}
VMTableDefinitionOwnership::VMOwned(table_loc)
} else {
VMTableDefinitionOwnership::HostOwned(Box::new(UnsafeCell::new(
VMTableDefinition {
base: base as _,
current_elements: table_minimum as _,
},
)))
},
}),
}
}
unsafe fn get_vm_table_definition(&self) -> NonNull<VMTableDefinition> {
match &self.vm_table_definition {
VMTableDefinitionOwnership::VMOwned(ptr) => *ptr,
VMTableDefinitionOwnership::HostOwned(boxed_ptr) => {
NonNull::new_unchecked(boxed_ptr.get())
}
}
}
}
impl Table for LinearTable {
fn ty(&self) -> &TableType {
&self.table
}
fn style(&self) -> &TableStyle {
&self.style
}
fn size(&self) -> u32 {
unsafe {
let td_ptr = self.get_vm_table_definition();
let td = td_ptr.as_ref();
td.current_elements
}
}
fn grow(&self, delta: u32) -> Option<u32> {
let mut vec_guard = self.vec.lock().unwrap();
let vec = vec_guard.borrow_mut();
let size = self.size();
let new_len = size.checked_add(delta)?;
if self.maximum.map_or(false, |max| new_len > max) {
return None;
}
vec.resize(
usize::try_from(new_len).unwrap(),
VMCallerCheckedAnyfunc::default(),
);
unsafe {
let mut td_ptr = self.get_vm_table_definition();
let td = td_ptr.as_mut();
td.current_elements = new_len;
td.base = vec.as_mut_ptr() as _;
}
Some(size)
}
fn get(&self, index: u32) -> Option<VMCallerCheckedAnyfunc> {
let vec_guard = self.vec.lock().unwrap();
vec_guard.borrow().get(index as usize).cloned()
}
fn set(&self, index: u32, func: VMCallerCheckedAnyfunc) -> Result<(), Trap> {
let mut vec_guard = self.vec.lock().unwrap();
let vec = vec_guard.borrow_mut();
match vec.get_mut(index as usize) {
Some(slot) => {
*slot = func;
Ok(())
}
None => Err(Trap::new_from_runtime(TrapCode::TableAccessOutOfBounds)),
}
}
fn vmtable(&self) -> NonNull<VMTableDefinition> {
let _vec_guard = self.vec.lock().unwrap();
unsafe { self.get_vm_table_definition() }
}
}