use cranelift::codegen::ir::FuncRef;
use cranelift::prelude::*;
use shape_vm::type_tracking::NativeKind;
pub const V2_HEAP_HEADER_SIZE: u32 = 8;
pub const V2_HEADER_REFCOUNT_OFFSET: u32 = 0;
pub const V2_HEADER_KIND_OFFSET: u32 = 4;
pub const V2_HEADER_FLAGS_OFFSET: u32 = 6;
pub fn cranelift_type_for_slot(kind: NativeKind) -> types::Type {
match kind {
NativeKind::Float64 | NativeKind::NullableFloat64 => types::F64,
NativeKind::Int64 | NativeKind::NullableInt64 | NativeKind::UInt64 | NativeKind::NullableUInt64 => {
types::I64
}
NativeKind::Int32 | NativeKind::NullableInt32 | NativeKind::UInt32 | NativeKind::NullableUInt32 => {
types::I32
}
NativeKind::Int16 | NativeKind::NullableInt16 | NativeKind::UInt16 | NativeKind::NullableUInt16 => {
types::I16
}
NativeKind::Int8
| NativeKind::NullableInt8
| NativeKind::UInt8
| NativeKind::NullableUInt8
| NativeKind::Bool => types::I8,
NativeKind::IntSize
| NativeKind::NullableIntSize
| NativeKind::UIntSize
| NativeKind::NullableUIntSize => types::I64,
NativeKind::Float32 => types::F32,
NativeKind::Char => types::I32,
NativeKind::StringV2 | NativeKind::DecimalV2 => types::I64,
NativeKind::String | NativeKind::Ptr(_) => types::I64,
NativeKind::Null => types::I8,
}
}
pub fn slot_byte_width(kind: NativeKind) -> u32 {
match kind {
NativeKind::Float64 | NativeKind::NullableFloat64 => 8,
NativeKind::Int64 | NativeKind::NullableInt64 | NativeKind::UInt64 | NativeKind::NullableUInt64 => {
8
}
NativeKind::Int32 | NativeKind::NullableInt32 | NativeKind::UInt32 | NativeKind::NullableUInt32 => {
4
}
NativeKind::Int16 | NativeKind::NullableInt16 | NativeKind::UInt16 | NativeKind::NullableUInt16 => {
2
}
NativeKind::Int8
| NativeKind::NullableInt8
| NativeKind::UInt8
| NativeKind::NullableUInt8
| NativeKind::Bool => 1,
NativeKind::IntSize
| NativeKind::NullableIntSize
| NativeKind::UIntSize
| NativeKind::NullableUIntSize => 8,
NativeKind::Float32 | NativeKind::Char => 4,
NativeKind::StringV2 | NativeKind::DecimalV2 => 8,
NativeKind::String | NativeKind::Ptr(_) => 8,
NativeKind::Null => 1,
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct FieldLayout {
pub name: String,
pub offset: u32,
pub kind: NativeKind,
}
pub fn compute_struct_layout(fields: &[(String, NativeKind)]) -> (Vec<FieldLayout>, u32) {
let mut layouts = Vec::with_capacity(fields.len());
let mut cursor = V2_HEAP_HEADER_SIZE;
for (name, kind) in fields {
let width = slot_byte_width(*kind);
let align = width;
let misalign = cursor % align;
if misalign != 0 {
cursor += align - misalign;
}
layouts.push(FieldLayout {
name: name.clone(),
offset: cursor,
kind: *kind,
});
cursor += width;
}
let misalign = cursor % 8;
if misalign != 0 {
cursor += 8 - misalign;
}
(layouts, cursor)
}
pub struct MirToIR<'a, 'b: 'a> {
pub builder: &'a mut FunctionBuilder<'b>,
}
impl<'a, 'b: 'a> MirToIR<'a, 'b> {
pub fn new(builder: &'a mut FunctionBuilder<'b>) -> Self {
Self { builder }
}
pub fn v2_field_get(
&mut self,
struct_ptr: Value,
field_offset: u32,
field_type: NativeKind,
) -> Value {
let cl_type = cranelift_type_for_slot(field_type);
self.builder
.ins()
.load(cl_type, MemFlags::trusted(), struct_ptr, field_offset as i32)
}
pub fn v2_field_set(
&mut self,
struct_ptr: Value,
field_offset: u32,
val: Value,
_field_type: NativeKind,
) {
self.builder
.ins()
.store(MemFlags::trusted(), val, struct_ptr, field_offset as i32);
}
pub fn v2_struct_alloc(&mut self, total_size: u32, alloc_fn: FuncRef) -> Value {
let size_val = self.builder.ins().iconst(types::I64, total_size as i64);
let inst = self.builder.ins().call(alloc_fn, &[size_val]);
self.builder.inst_results(inst)[0]
}
pub fn v2_write_refcount(&mut self, struct_ptr: Value, initial: u32) {
let rc = self.builder.ins().iconst(types::I32, initial as i64);
self.builder.ins().store(
MemFlags::trusted(),
rc,
struct_ptr,
V2_HEADER_REFCOUNT_OFFSET as i32,
);
}
pub fn v2_write_kind(&mut self, struct_ptr: Value, kind: u16) {
let k = self.builder.ins().iconst(types::I16, kind as i64);
self.builder.ins().store(
MemFlags::trusted(),
k,
struct_ptr,
V2_HEADER_KIND_OFFSET as i32,
);
}
pub fn v2_write_flags(&mut self, struct_ptr: Value, flags: u8) {
let f = self.builder.ins().iconst(types::I8, flags as i64);
self.builder.ins().store(
MemFlags::trusted(),
f,
struct_ptr,
V2_HEADER_FLAGS_OFFSET as i32,
);
}
}
#[cfg(test)]
mod tests {
use super::*;
use cranelift_jit::{JITBuilder, JITModule};
use cranelift_module::Module;
#[test]
fn v2_field_layout_point_all_f64() {
let fields = vec![
("x".to_string(), NativeKind::Float64),
("y".to_string(), NativeKind::Float64),
];
let (layout, total) = compute_struct_layout(&fields);
assert_eq!(layout.len(), 2);
assert_eq!(layout[0].name, "x");
assert_eq!(layout[0].offset, 8); assert_eq!(layout[0].kind, NativeKind::Float64);
assert_eq!(layout[1].name, "y");
assert_eq!(layout[1].offset, 16); assert_eq!(layout[1].kind, NativeKind::Float64);
assert_eq!(total, 24); }
#[test]
fn v2_field_layout_mixed_types() {
let fields = vec![
("flag".to_string(), NativeKind::Bool),
("count".to_string(), NativeKind::Int32),
("value".to_string(), NativeKind::Float64),
];
let (layout, total) = compute_struct_layout(&fields);
assert_eq!(layout.len(), 3);
assert_eq!(layout[0].name, "flag");
assert_eq!(layout[0].offset, 8);
assert_eq!(layout[0].kind, NativeKind::Bool);
assert_eq!(layout[1].name, "count");
assert_eq!(layout[1].offset, 12);
assert_eq!(layout[1].kind, NativeKind::Int32);
assert_eq!(layout[2].name, "value");
assert_eq!(layout[2].offset, 16);
assert_eq!(layout[2].kind, NativeKind::Float64);
assert_eq!(total, 24);
}
#[test]
fn v2_field_layout_i16_alignment() {
let fields = vec![
("a".to_string(), NativeKind::Int16),
("b".to_string(), NativeKind::Int16),
("c".to_string(), NativeKind::Int64),
];
let (layout, total) = compute_struct_layout(&fields);
assert_eq!(layout[0].offset, 8);
assert_eq!(layout[1].offset, 10);
assert_eq!(layout[2].offset, 16);
assert_eq!(total, 24);
}
#[test]
fn v2_field_layout_single_bool() {
let fields = vec![("active".to_string(), NativeKind::Bool)];
let (layout, total) = compute_struct_layout(&fields);
assert_eq!(layout[0].offset, 8);
assert_eq!(total, 16);
}
#[test]
fn v2_field_layout_empty_struct() {
let fields: Vec<(String, NativeKind)> = vec![];
let (layout, total) = compute_struct_layout(&fields);
assert_eq!(layout.len(), 0);
assert_eq!(total, 8); }
#[test]
fn v2_field_layout_padding_between_fields() {
let fields = vec![
("a".to_string(), NativeKind::Bool),
("b".to_string(), NativeKind::Int64),
];
let (layout, total) = compute_struct_layout(&fields);
assert_eq!(layout[0].offset, 8);
assert_eq!(layout[1].offset, 16); assert_eq!(total, 24); }
#[test]
fn v2_field_layout_all_i32() {
let fields = vec![
("x".to_string(), NativeKind::Int32),
("y".to_string(), NativeKind::Int32),
("z".to_string(), NativeKind::Int32),
];
let (layout, total) = compute_struct_layout(&fields);
assert_eq!(layout[0].offset, 8);
assert_eq!(layout[1].offset, 12);
assert_eq!(layout[2].offset, 16);
assert_eq!(total, 24);
}
#[test]
fn v2_field_slot_to_cranelift_type() {
assert_eq!(cranelift_type_for_slot(NativeKind::Float64), types::F64);
assert_eq!(cranelift_type_for_slot(NativeKind::NullableFloat64), types::F64);
assert_eq!(cranelift_type_for_slot(NativeKind::Int64), types::I64);
assert_eq!(cranelift_type_for_slot(NativeKind::Int32), types::I32);
assert_eq!(cranelift_type_for_slot(NativeKind::Int16), types::I16);
assert_eq!(cranelift_type_for_slot(NativeKind::Bool), types::I8);
assert_eq!(cranelift_type_for_slot(NativeKind::Int8), types::I8);
assert_eq!(
cranelift_type_for_slot(NativeKind::Ptr(shape_value::heap_value::HeapKind::TypedArray)),
types::I64
);
assert_eq!(cranelift_type_for_slot(NativeKind::String), types::I64);
}
#[test]
fn v2_field_slot_byte_widths() {
assert_eq!(slot_byte_width(NativeKind::Float64), 8);
assert_eq!(slot_byte_width(NativeKind::Int64), 8);
assert_eq!(slot_byte_width(NativeKind::Int32), 4);
assert_eq!(slot_byte_width(NativeKind::Int16), 2);
assert_eq!(slot_byte_width(NativeKind::Bool), 1);
assert_eq!(slot_byte_width(NativeKind::Int8), 1);
assert_eq!(slot_byte_width(NativeKind::String), 8);
}
fn make_jit_env() -> (
JITModule,
cranelift::codegen::Context,
FunctionBuilderContext,
) {
let mut flag_builder = settings::builder();
flag_builder.set("opt_level", "speed").unwrap();
flag_builder.set("is_pic", "false").unwrap();
let isa_builder = cranelift_native::builder().unwrap();
let isa = isa_builder
.finish(settings::Flags::new(flag_builder))
.unwrap();
let builder = JITBuilder::with_isa(isa, cranelift_module::default_libcall_names());
let module = JITModule::new(builder);
let ctx = cranelift::codegen::Context::new();
let fb_ctx = FunctionBuilderContext::new();
(module, ctx, fb_ctx)
}
unsafe fn alloc_test_struct(total_bytes: usize) -> *mut u8 {
let layout = std::alloc::Layout::from_size_align(total_bytes, 8).unwrap();
let ptr = unsafe { std::alloc::alloc_zeroed(layout) };
assert!(!ptr.is_null());
ptr
}
unsafe fn free_test_struct(ptr: *mut u8, total_bytes: usize) {
let layout = std::alloc::Layout::from_size_align(total_bytes, 8).unwrap();
unsafe { std::alloc::dealloc(ptr, layout) };
}
#[test]
fn v2_field_get_f64_codegen_and_execute() {
let (mut module, mut ctx, mut fb_ctx) = make_jit_env();
let ptr_type = module.target_config().pointer_type();
let mut sig = module.make_signature();
sig.params.push(AbiParam::new(ptr_type));
sig.returns.push(AbiParam::new(types::F64));
let func_id = module
.declare_function("test_get_f64", cranelift_module::Linkage::Local, &sig)
.unwrap();
ctx.func.signature = sig;
{
let mut builder = FunctionBuilder::new(&mut ctx.func, &mut fb_ctx);
let entry = builder.create_block();
builder.append_block_params_for_function_params(entry);
builder.switch_to_block(entry);
builder.seal_block(entry);
let struct_ptr = builder.block_params(entry)[0];
let result = {
let mut mir = MirToIR::new(&mut builder);
mir.v2_field_get(struct_ptr, 8, NativeKind::Float64)
};
builder.ins().return_(&[result]);
builder.finalize();
}
module.define_function(func_id, &mut ctx).unwrap();
module.clear_context(&mut ctx);
module.finalize_definitions().unwrap();
let code_ptr = module.get_finalized_function(func_id);
unsafe {
let struct_mem = alloc_test_struct(24);
let field_ptr = struct_mem.add(8) as *mut f64;
*field_ptr = 42.5;
let func: unsafe fn(u64) -> f64 = std::mem::transmute(code_ptr);
let result = func(struct_mem as u64);
assert_eq!(result, 42.5);
free_test_struct(struct_mem, 24);
}
}
#[test]
fn v2_field_set_f64_codegen_and_execute() {
let (mut module, mut ctx, mut fb_ctx) = make_jit_env();
let ptr_type = module.target_config().pointer_type();
let mut sig = module.make_signature();
sig.params.push(AbiParam::new(ptr_type));
sig.params.push(AbiParam::new(types::F64));
let func_id = module
.declare_function("test_set_f64", cranelift_module::Linkage::Local, &sig)
.unwrap();
ctx.func.signature = sig;
{
let mut builder = FunctionBuilder::new(&mut ctx.func, &mut fb_ctx);
let entry = builder.create_block();
builder.append_block_params_for_function_params(entry);
builder.switch_to_block(entry);
builder.seal_block(entry);
let struct_ptr = builder.block_params(entry)[0];
let val = builder.block_params(entry)[1];
{
let mut mir = MirToIR::new(&mut builder);
mir.v2_field_set(struct_ptr, 16, val, NativeKind::Float64);
}
builder.ins().return_(&[]);
builder.finalize();
}
module.define_function(func_id, &mut ctx).unwrap();
module.clear_context(&mut ctx);
module.finalize_definitions().unwrap();
let code_ptr = module.get_finalized_function(func_id);
unsafe {
let struct_mem = alloc_test_struct(24);
let func: unsafe fn(u64, f64) = std::mem::transmute(code_ptr);
func(struct_mem as u64, 99.75);
let stored = *(struct_mem.add(16) as *const f64);
assert_eq!(stored, 99.75);
free_test_struct(struct_mem, 24);
}
}
#[test]
fn v2_field_get_i32_codegen_and_execute() {
let (mut module, mut ctx, mut fb_ctx) = make_jit_env();
let ptr_type = module.target_config().pointer_type();
let mut sig = module.make_signature();
sig.params.push(AbiParam::new(ptr_type));
sig.returns.push(AbiParam::new(types::I32));
let func_id = module
.declare_function("test_get_i32", cranelift_module::Linkage::Local, &sig)
.unwrap();
ctx.func.signature = sig;
{
let mut builder = FunctionBuilder::new(&mut ctx.func, &mut fb_ctx);
let entry = builder.create_block();
builder.append_block_params_for_function_params(entry);
builder.switch_to_block(entry);
builder.seal_block(entry);
let struct_ptr = builder.block_params(entry)[0];
let result = {
let mut mir = MirToIR::new(&mut builder);
mir.v2_field_get(struct_ptr, 12, NativeKind::Int32)
};
builder.ins().return_(&[result]);
builder.finalize();
}
module.define_function(func_id, &mut ctx).unwrap();
module.clear_context(&mut ctx);
module.finalize_definitions().unwrap();
let code_ptr = module.get_finalized_function(func_id);
unsafe {
let struct_mem = alloc_test_struct(16);
let field_ptr = struct_mem.add(12) as *mut i32;
*field_ptr = 0x7FFF_FFFE;
let func: unsafe fn(u64) -> i32 = std::mem::transmute(code_ptr);
let result = func(struct_mem as u64);
assert_eq!(result, 0x7FFF_FFFE);
free_test_struct(struct_mem, 16);
}
}
#[test]
fn v2_field_get_bool_codegen_and_execute() {
let (mut module, mut ctx, mut fb_ctx) = make_jit_env();
let ptr_type = module.target_config().pointer_type();
let mut sig = module.make_signature();
sig.params.push(AbiParam::new(ptr_type));
sig.returns.push(AbiParam::new(types::I8));
let func_id = module
.declare_function("test_get_bool", cranelift_module::Linkage::Local, &sig)
.unwrap();
ctx.func.signature = sig;
{
let mut builder = FunctionBuilder::new(&mut ctx.func, &mut fb_ctx);
let entry = builder.create_block();
builder.append_block_params_for_function_params(entry);
builder.switch_to_block(entry);
builder.seal_block(entry);
let struct_ptr = builder.block_params(entry)[0];
let result = {
let mut mir = MirToIR::new(&mut builder);
mir.v2_field_get(struct_ptr, 8, NativeKind::Bool)
};
builder.ins().return_(&[result]);
builder.finalize();
}
module.define_function(func_id, &mut ctx).unwrap();
module.clear_context(&mut ctx);
module.finalize_definitions().unwrap();
let code_ptr = module.get_finalized_function(func_id);
unsafe {
let struct_mem = alloc_test_struct(16);
*struct_mem.add(8) = 1u8;
let func: unsafe fn(u64) -> i8 = std::mem::transmute(code_ptr);
assert_eq!(func(struct_mem as u64), 1);
*struct_mem.add(8) = 0u8;
assert_eq!(func(struct_mem as u64), 0);
free_test_struct(struct_mem, 16);
}
}
#[test]
fn v2_field_point_two_fields_codegen_and_execute() {
let (mut module, mut ctx, mut fb_ctx) = make_jit_env();
let ptr_type = module.target_config().pointer_type();
let mut sig = module.make_signature();
sig.params.push(AbiParam::new(ptr_type));
sig.returns.push(AbiParam::new(types::F64));
let func_id = module
.declare_function("test_point_sum", cranelift_module::Linkage::Local, &sig)
.unwrap();
let fields = vec![
("x".to_string(), NativeKind::Float64),
("y".to_string(), NativeKind::Float64),
];
let (layout, total) = compute_struct_layout(&fields);
assert_eq!(layout[0].offset, 8);
assert_eq!(layout[1].offset, 16);
ctx.func.signature = sig;
{
let mut builder = FunctionBuilder::new(&mut ctx.func, &mut fb_ctx);
let entry = builder.create_block();
builder.append_block_params_for_function_params(entry);
builder.switch_to_block(entry);
builder.seal_block(entry);
let struct_ptr = builder.block_params(entry)[0];
let (x, y) = {
let mut mir = MirToIR::new(&mut builder);
let x = mir.v2_field_get(struct_ptr, layout[0].offset, layout[0].kind);
let y = mir.v2_field_get(struct_ptr, layout[1].offset, layout[1].kind);
(x, y)
};
let sum = builder.ins().fadd(x, y);
builder.ins().return_(&[sum]);
builder.finalize();
}
module.define_function(func_id, &mut ctx).unwrap();
module.clear_context(&mut ctx);
module.finalize_definitions().unwrap();
let code_ptr = module.get_finalized_function(func_id);
unsafe {
let struct_mem = alloc_test_struct(total as usize);
*(struct_mem.add(8) as *mut f64) = 3.0;
*(struct_mem.add(16) as *mut f64) = 4.0;
let func: unsafe fn(u64) -> f64 = std::mem::transmute(code_ptr);
assert_eq!(func(struct_mem as u64), 7.0);
free_test_struct(struct_mem, total as usize);
}
}
#[test]
fn v2_field_mixed_struct_codegen_and_execute() {
let (mut module, mut ctx, mut fb_ctx) = make_jit_env();
let ptr_type = module.target_config().pointer_type();
let mut sig = module.make_signature();
sig.params.push(AbiParam::new(ptr_type));
sig.returns.push(AbiParam::new(types::F64));
let func_id = module
.declare_function(
"test_mixed_value",
cranelift_module::Linkage::Local,
&sig,
)
.unwrap();
let fields = vec![
("flag".to_string(), NativeKind::Bool),
("count".to_string(), NativeKind::Int32),
("value".to_string(), NativeKind::Float64),
];
let (layout, total) = compute_struct_layout(&fields);
ctx.func.signature = sig;
{
let mut builder = FunctionBuilder::new(&mut ctx.func, &mut fb_ctx);
let entry = builder.create_block();
builder.append_block_params_for_function_params(entry);
builder.switch_to_block(entry);
builder.seal_block(entry);
let struct_ptr = builder.block_params(entry)[0];
let value = {
let mut mir = MirToIR::new(&mut builder);
mir.v2_field_get(struct_ptr, layout[2].offset, layout[2].kind)
};
builder.ins().return_(&[value]);
builder.finalize();
}
module.define_function(func_id, &mut ctx).unwrap();
module.clear_context(&mut ctx);
module.finalize_definitions().unwrap();
let code_ptr = module.get_finalized_function(func_id);
unsafe {
let struct_mem = alloc_test_struct(total as usize);
*struct_mem.add(layout[0].offset as usize) = 1u8;
*(struct_mem.add(layout[1].offset as usize) as *mut i32) = 42;
*(struct_mem.add(layout[2].offset as usize) as *mut f64) = 3.14;
let func: unsafe fn(u64) -> f64 = std::mem::transmute(code_ptr);
assert_eq!(func(struct_mem as u64), 3.14);
free_test_struct(struct_mem, total as usize);
}
}
#[test]
fn v2_field_write_then_read_roundtrip() {
let (mut module, mut ctx, mut fb_ctx) = make_jit_env();
let ptr_type = module.target_config().pointer_type();
let mut writer_sig = module.make_signature();
writer_sig.params.push(AbiParam::new(ptr_type));
writer_sig.params.push(AbiParam::new(types::F64));
let writer_id = module
.declare_function(
"test_roundtrip_w",
cranelift_module::Linkage::Local,
&writer_sig,
)
.unwrap();
ctx.func.signature = writer_sig;
{
let mut builder = FunctionBuilder::new(&mut ctx.func, &mut fb_ctx);
let entry = builder.create_block();
builder.append_block_params_for_function_params(entry);
builder.switch_to_block(entry);
builder.seal_block(entry);
let sp = builder.block_params(entry)[0];
let val = builder.block_params(entry)[1];
{
let mut mir = MirToIR::new(&mut builder);
mir.v2_field_set(sp, 8, val, NativeKind::Float64);
}
builder.ins().return_(&[]);
builder.finalize();
}
module.define_function(writer_id, &mut ctx).unwrap();
module.clear_context(&mut ctx);
let mut reader_sig = module.make_signature();
reader_sig.params.push(AbiParam::new(ptr_type));
reader_sig.returns.push(AbiParam::new(types::F64));
let reader_id = module
.declare_function(
"test_roundtrip_r",
cranelift_module::Linkage::Local,
&reader_sig,
)
.unwrap();
ctx.func.signature = reader_sig;
{
let mut builder = FunctionBuilder::new(&mut ctx.func, &mut fb_ctx);
let entry = builder.create_block();
builder.append_block_params_for_function_params(entry);
builder.switch_to_block(entry);
builder.seal_block(entry);
let sp = builder.block_params(entry)[0];
let result = {
let mut mir = MirToIR::new(&mut builder);
mir.v2_field_get(sp, 8, NativeKind::Float64)
};
builder.ins().return_(&[result]);
builder.finalize();
}
module.define_function(reader_id, &mut ctx).unwrap();
module.clear_context(&mut ctx);
module.finalize_definitions().unwrap();
let writer_ptr = module.get_finalized_function(writer_id);
let reader_ptr = module.get_finalized_function(reader_id);
unsafe {
let struct_mem = alloc_test_struct(16);
let writer: unsafe fn(u64, f64) = std::mem::transmute(writer_ptr);
let reader: unsafe fn(u64) -> f64 = std::mem::transmute(reader_ptr);
writer(struct_mem as u64, -123.456);
let got = reader(struct_mem as u64);
assert_eq!(got, -123.456);
writer(struct_mem as u64, f64::INFINITY);
assert_eq!(reader(struct_mem as u64), f64::INFINITY);
free_test_struct(struct_mem, 16);
}
}
}