use std::error::Error;
use std::fmt;
use bamts_bytecode::{Program as BytecodeProgram, Verified};
use cranelift_codegen::Context;
use cranelift_codegen::ir::{ExternalName, Function, UserExternalName};
use cranelift_codegen::isa;
use cranelift_codegen::settings::{self, Flags};
use cranelift_module::{DataDescription, FuncId, Linkage, Module, default_libcall_names};
use cranelift_object::{ObjectBuilder, ObjectModule};
use crate::{
HELPER_NAMESPACE, Helper, LowerError, LoweredProgram, ProgramLowerError, function_symbol,
lower_program,
};
const HELPER_COUNT: u32 = 32;
const AOT_MAGIC: u64 = u64::from_le_bytes(*b"BMTSAOT1");
const AOT_ABI_VERSION: u32 = 3;
const UNIT_DESCRIPTOR_BYTES: usize = 16;
const PROGRAM_DESCRIPTOR_BYTES: usize = 56;
const BYTECODE_SYMBOL: &str = "bamts_bytecode_blob";
const UNITS_SYMBOL: &str = "bamts_unit_descriptors";
pub const PROGRAM_DESCRIPTOR_SYMBOL: &str = "bamts_program_descriptor";
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct AotObject {
pub bytes: Vec<u8>,
pub target: String,
pub descriptor_symbol: &'static str,
pub entry_module: u32,
pub entry_function: u32,
pub entry_symbol: String,
pub required_helpers: Vec<&'static str>,
}
#[derive(Debug)]
pub enum AotError {
TargetLookup(String),
TargetBuild(String),
TargetEndianness(String),
Lower(ProgramLowerError),
InvalidLoweredModule(String),
Module(String),
Emit(String),
}
impl fmt::Display for AotError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::TargetLookup(message) => write!(f, "unsupported AOT target: {message}"),
Self::TargetBuild(message) => write!(f, "could not configure AOT target: {message}"),
Self::TargetEndianness(target) => {
write!(f, "AOT target has no known byte order: {target}")
}
Self::Lower(error) => write!(f, "AOT lowering failed: {error}"),
Self::InvalidLoweredModule(message) => {
write!(f, "invalid lowered module for AOT emission: {message}")
}
Self::Module(message) => write!(f, "AOT object definition failed: {message}"),
Self::Emit(message) => write!(f, "AOT object serialization failed: {message}"),
}
}
}
impl Error for AotError {
fn source(&self) -> Option<&(dyn Error + 'static)> {
match self {
Self::Lower(error) => Some(error),
_ => None,
}
}
}
impl From<ProgramLowerError> for AotError {
fn from(error: ProgramLowerError) -> Self {
Self::Lower(error)
}
}
fn require_64_bit_pointer_width(bits: u8) -> Result<(), LowerError> {
if bits != 64 {
Err(LowerError::UnsupportedPointerWidth { bits })
} else {
Ok(())
}
}
pub fn compile_aot(
bytecode: &BytecodeProgram<Verified>,
target: &str,
) -> Result<AotObject, AotError> {
let flags = Flags::new(settings::builder());
let isa_builder =
isa::lookup_by_name(target).map_err(|error| AotError::TargetLookup(error.to_string()))?;
let isa = isa_builder
.finish(flags)
.map_err(|error| AotError::TargetBuild(error.to_string()))?;
require_64_bit_pointer_width(isa.frontend_config().pointer_bits()).map_err(|kind| {
AotError::Lower(ProgramLowerError {
module: bamts_bytecode::ModuleId::new(0),
kind,
})
})?;
let target_endianness = isa
.triple()
.endianness()
.map_err(|()| AotError::TargetEndianness(isa.triple().to_string()))?;
let little_endianness = isa::lookup_by_name("x86_64")
.expect("the all-native-arch build includes x86-64")
.triple()
.endianness()
.expect("x86-64 has a defined byte order");
let little_endian = target_endianness == little_endianness;
let lowered = lower_program(bytecode, isa.frontend_config())?;
let normalized_target = isa.triple().to_string();
let call_conv = isa.frontend_config().default_call_conv;
let builder = ObjectBuilder::new(isa, "bamts", default_libcall_names())
.map_err(|error| AotError::Module(error.to_string()))?;
let mut object = ObjectModule::new(builder);
let function_ids = declare_functions(&mut object, &lowered)?;
let helper_ids = declare_helpers(&mut object, call_conv)?;
define_functions(&mut object, &lowered, &function_ids, &helper_ids)?;
define_program_data(
&mut object,
&lowered,
&function_ids,
bytecode.encode(),
little_endian,
)?;
let required_helpers = (0..HELPER_COUNT)
.filter_map(Helper::from_external_index)
.filter(|helper| {
lowered.modules.iter().any(|module| {
module
.functions
.iter()
.any(|function| function.helpers.contains(helper))
})
})
.map(Helper::symbol)
.collect();
let entry_module = lowered.entry_module.get();
let entry_function = lowered.entry_function.get();
let bytes = object
.finish()
.emit()
.map_err(|error| AotError::Emit(error.to_string()))?;
Ok(AotObject {
bytes,
target: normalized_target,
descriptor_symbol: PROGRAM_DESCRIPTOR_SYMBOL,
entry_module,
entry_function,
entry_symbol: function_symbol(entry_module, entry_function),
required_helpers,
})
}
struct DeclaredUnit {
module_id: u32,
function_id: u32,
function: FuncId,
}
fn declare_functions(
object: &mut ObjectModule,
lowered: &LoweredProgram,
) -> Result<Vec<DeclaredUnit>, AotError> {
let function_count = lowered
.modules
.iter()
.map(|module| module.functions.len())
.sum();
let mut units = Vec::with_capacity(function_count);
for (module_index, module) in lowered.modules.iter().enumerate() {
if module.id.get() as usize != module_index {
return Err(AotError::InvalidLoweredModule(format!(
"module {} appears at index {module_index}",
module.id.get()
)));
}
for (function_index, function) in module.functions.iter().enumerate() {
if function.id.get() as usize != function_index {
return Err(AotError::InvalidLoweredModule(format!(
"module {} function {} appears at local index {function_index}",
module.id.get(),
function.id.get()
)));
}
units.push(DeclaredUnit {
module_id: module.id.get(),
function_id: function.id.get(),
function: object
.declare_function(&function.symbol, Linkage::Export, &function.signature)
.map_err(|error| AotError::Module(error.to_string()))?,
});
}
}
Ok(units)
}
fn declare_helpers(
object: &mut ObjectModule,
call_conv: cranelift_codegen::isa::CallConv,
) -> Result<Vec<FuncId>, AotError> {
(0..HELPER_COUNT)
.map(|index| {
let helper = Helper::from_external_index(index).ok_or_else(|| {
AotError::InvalidLoweredModule(format!("missing helper ABI index {index}"))
})?;
object
.declare_function(
helper.symbol(),
Linkage::Import,
&helper.signature(call_conv),
)
.map_err(|error| AotError::Module(error.to_string()))
})
.collect()
}
fn define_functions(
object: &mut ObjectModule,
lowered: &LoweredProgram,
units: &[DeclaredUnit],
helper_ids: &[FuncId],
) -> Result<(), AotError> {
let mut unit_index = 0;
for module in &lowered.modules {
for lowered_function in &module.functions {
let mut function = lowered_function.clif.clone();
remap_helper_names(&mut function, helper_ids)?;
let mut context = Context::for_function(function);
object
.define_function(units[unit_index].function, &mut context)
.map_err(|error| AotError::Module(error.to_string()))?;
unit_index += 1;
}
}
Ok(())
}
fn remap_helper_names(function: &mut Function, helper_ids: &[FuncId]) -> Result<(), AotError> {
let external_functions: Vec<_> = function.dfg.ext_funcs.keys().collect();
for function_ref in external_functions {
let ExternalName::User(name_ref) = function.dfg.ext_funcs[function_ref].name else {
continue;
};
let name = &function.params.user_named_funcs()[name_ref];
if name.namespace != HELPER_NAMESPACE {
continue;
}
let helper_id = helper_ids.get(name.index as usize).ok_or_else(|| {
AotError::InvalidLoweredModule(format!(
"function references unknown helper index {}",
name.index
))
})?;
let replacement = function.declare_imported_user_function(UserExternalName {
namespace: 0,
index: helper_id.as_u32(),
});
function.dfg.ext_funcs[function_ref].name = ExternalName::user(replacement);
}
Ok(())
}
fn define_program_data(
object: &mut ObjectModule,
lowered: &LoweredProgram,
function_ids: &[DeclaredUnit],
bytecode: Vec<u8>,
little_endian: bool,
) -> Result<(), AotError> {
let unit_bytes = function_ids
.len()
.checked_mul(UNIT_DESCRIPTOR_BYTES)
.ok_or_else(|| AotError::InvalidLoweredModule("unit table size overflow".to_string()))?;
if unit_bytes > u32::MAX as usize {
return Err(AotError::InvalidLoweredModule(
"unit table exceeds the relocation offset range".to_string(),
));
}
let bytecode_id = object
.declare_data(BYTECODE_SYMBOL, Linkage::Local, false, false)
.map_err(|error| AotError::Module(error.to_string()))?;
let mut bytecode_data = DataDescription::new();
bytecode_data.define(bytecode.into_boxed_slice());
bytecode_data.set_align(1);
object
.define_data(bytecode_id, &bytecode_data)
.map_err(|error| AotError::Module(error.to_string()))?;
let units_id = object
.declare_data(UNITS_SYMBOL, Linkage::Local, false, false)
.map_err(|error| AotError::Module(error.to_string()))?;
let mut unit_contents = vec![0; unit_bytes];
for (index, unit) in function_ids.iter().enumerate() {
let offset = index * UNIT_DESCRIPTOR_BYTES;
write_u32(&mut unit_contents, offset, unit.function_id, little_endian);
write_u32(
&mut unit_contents,
offset + 4,
unit.module_id,
little_endian,
);
}
let mut units_data = DataDescription::new();
units_data.define(unit_contents.into_boxed_slice());
units_data.set_align(8);
for (index, unit) in function_ids.iter().enumerate() {
let function_ref = object.declare_func_in_data(unit.function, &mut units_data);
units_data.write_function_addr((index * UNIT_DESCRIPTOR_BYTES + 8) as u32, function_ref);
}
object
.define_data(units_id, &units_data)
.map_err(|error| AotError::Module(error.to_string()))?;
let descriptor_id = object
.declare_data(PROGRAM_DESCRIPTOR_SYMBOL, Linkage::Export, false, false)
.map_err(|error| AotError::Module(error.to_string()))?;
let mut descriptor = vec![0; PROGRAM_DESCRIPTOR_BYTES];
write_u64(&mut descriptor, 0, AOT_MAGIC, little_endian);
write_u32(&mut descriptor, 8, AOT_ABI_VERSION, little_endian);
write_u64(
&mut descriptor,
24,
bytecode_data.init.size() as u64,
little_endian,
);
write_u64(
&mut descriptor,
40,
function_ids.len() as u64,
little_endian,
);
write_u32(
&mut descriptor,
48,
lowered.entry_function.get(),
little_endian,
);
write_u32(
&mut descriptor,
52,
lowered.entry_module.get(),
little_endian,
);
let mut descriptor_data = DataDescription::new();
descriptor_data.define(descriptor.into_boxed_slice());
descriptor_data.set_align(8);
descriptor_data.set_used(true);
let bytecode_ref = object.declare_data_in_data(bytecode_id, &mut descriptor_data);
descriptor_data.write_data_addr(16, bytecode_ref, 0);
let units_ref = object.declare_data_in_data(units_id, &mut descriptor_data);
descriptor_data.write_data_addr(32, units_ref, 0);
object
.define_data(descriptor_id, &descriptor_data)
.map_err(|error| AotError::Module(error.to_string()))
}
fn write_u32(bytes: &mut [u8], offset: usize, value: u32, little_endian: bool) {
let encoded = if little_endian {
value.to_le_bytes()
} else {
value.to_be_bytes()
};
bytes[offset..offset + encoded.len()].copy_from_slice(&encoded);
}
fn write_u64(bytes: &mut [u8], offset: usize, value: u64, little_endian: bool) {
let encoded = if little_endian {
value.to_le_bytes()
} else {
value.to_be_bytes()
};
bytes[offset..offset + encoded.len()].copy_from_slice(&encoded);
}
#[cfg(test)]
mod tests {
use super::*;
use bamts_bytecode::{
Constant, ConstantId, EcmaString, Function as BytecodeFunction, FunctionFlags, FunctionId,
Instruction, Module, ModuleId, Program, ProgramDecodeLimits, ProgramModule, Register,
decode_verified_program,
};
use cranelift_object::object::{
Object, ObjectSection, ObjectSymbol, RelocationTarget, SymbolIndex,
};
fn function(code: Vec<Instruction>, register_count: u32) -> BytecodeFunction {
BytecodeFunction::new(
None,
0,
0,
register_count,
FunctionFlags::default(),
code,
Vec::new(),
)
}
fn module(name: &str, value: i32, loads_constant: bool) -> ProgramModule<Verified> {
let code = if loads_constant {
vec![
Instruction::LoadConst {
dst: Register::new(0),
constant: ConstantId::new(1),
},
Instruction::Return {
value: Register::new(0),
},
]
} else {
vec![Instruction::Halt]
};
ProgramModule {
name: ConstantId::new(0),
code: Module::new(
vec![
Constant::String(EcmaString::from_utf8(name)),
Constant::Int32(value),
],
vec![function(code, u32::from(loads_constant))],
FunctionId::new(0),
)
.verify()
.expect("test module verifies"),
edges: Vec::new(),
bindings: Vec::new(),
exports: Vec::new(),
}
}
fn test_program() -> Program<Verified> {
Program::link(
vec![module("dependency", 7, true), module("entry", 42, false)],
ModuleId::new(1),
)
.expect("test program verifies")
}
fn target() -> &'static str {
if cfg!(target_arch = "x86_64") {
"x86_64-unknown-linux-gnu"
} else if cfg!(target_arch = "aarch64") {
"aarch64-unknown-linux-gnu"
} else {
panic!("AOT object test needs a supported 64-bit host architecture")
}
}
fn symbol_bytes<'a>(
file: &'a cranelift_object::object::File<'a>,
symbol_name: &str,
) -> (&'a [u8], SymbolIndex) {
let symbol = file
.symbols()
.find(|symbol| symbol.name() == Ok(symbol_name))
.unwrap_or_else(|| panic!("missing symbol {symbol_name}"));
let section_index = symbol.section_index().expect("defined symbol section");
let section = file
.section_by_index(section_index)
.expect("symbol section");
let section_data = section.data().expect("section data");
let start = usize::try_from(symbol.address() - section.address()).expect("symbol offset");
let size = usize::try_from(symbol.size()).expect("symbol size");
(§ion_data[start..start + size], symbol.index())
}
fn relocation_targets(
file: &cranelift_object::object::File<'_>,
owner: SymbolIndex,
) -> Vec<String> {
let owner = file.symbol_by_index(owner).expect("owner symbol");
let section = file
.section_by_index(owner.section_index().expect("owner section"))
.expect("owner section data");
let start = owner.address() - section.address();
let end = start + owner.size();
section
.relocations()
.filter(|(offset, _)| (start..end).contains(offset))
.filter_map(|(_, relocation)| match relocation.target() {
RelocationTarget::Symbol(index) => file
.symbol_by_index(index)
.ok()
.and_then(|symbol| symbol.name().ok())
.map(str::to_owned),
_ => None,
})
.collect()
}
#[test]
fn emits_two_module_tuple_units_and_canonical_program() {
let program = test_program();
let canonical = program.encode();
let emitted = compile_aot(&program, target()).expect("AOT object emits");
let file = cranelift_object::object::File::parse(&*emitted.bytes).expect("object parses");
let (descriptor, descriptor_index) = symbol_bytes(&file, PROGRAM_DESCRIPTOR_SYMBOL);
assert_eq!(descriptor.len(), PROGRAM_DESCRIPTOR_BYTES);
assert_eq!(&descriptor[0..8], b"BMTSAOT1");
assert_eq!(u32::from_le_bytes(descriptor[8..12].try_into().unwrap()), 3);
assert_eq!(
u64::from_le_bytes(descriptor[24..32].try_into().unwrap()),
canonical.len() as u64
);
assert_eq!(
u64::from_le_bytes(descriptor[40..48].try_into().unwrap()),
2
);
assert_eq!(
u32::from_le_bytes(descriptor[48..52].try_into().unwrap()),
0
);
assert_eq!(
u32::from_le_bytes(descriptor[52..56].try_into().unwrap()),
1
);
assert_eq!(
relocation_targets(&file, descriptor_index),
[BYTECODE_SYMBOL.to_string(), UNITS_SYMBOL.to_string()]
);
let (embedded, _) = symbol_bytes(&file, BYTECODE_SYMBOL);
assert_eq!(embedded, canonical);
let decoded = decode_verified_program(embedded, &ProgramDecodeLimits::default())
.expect("embedded canonical program decodes");
assert_eq!(decoded, program);
assert_eq!(decoded.encode(), embedded);
let (units, units_index) = symbol_bytes(&file, UNITS_SYMBOL);
assert_eq!(units.len(), 2 * UNIT_DESCRIPTOR_BYTES);
assert_eq!(u32::from_le_bytes(units[0..4].try_into().unwrap()), 0);
assert_eq!(u32::from_le_bytes(units[4..8].try_into().unwrap()), 0);
assert_eq!(u32::from_le_bytes(units[16..20].try_into().unwrap()), 0);
assert_eq!(u32::from_le_bytes(units[20..24].try_into().unwrap()), 1);
assert_eq!(
relocation_targets(&file, units_index),
[function_symbol(0, 0), function_symbol(1, 0)]
);
assert!(file.symbols().any(|symbol| {
symbol.name() == Ok(emitted.entry_symbol.as_str()) && !symbol.is_undefined()
}));
assert_eq!(
emitted.required_helpers,
[Helper::LoadConstant.symbol(), Helper::ConsumeFuel.symbol()]
);
assert!(file.symbols().any(|symbol| {
symbol.name() == Ok(Helper::ConsumeFuel.symbol()) && symbol.is_undefined()
}));
assert_eq!((emitted.entry_module, emitted.entry_function), (1, 0));
assert_eq!(emitted.entry_symbol, function_symbol(1, 0));
}
#[test]
fn object_emission_is_deterministic() {
let program = test_program();
let first = compile_aot(&program, target()).expect("first AOT object emits");
let second = compile_aot(&program, target()).expect("second AOT object emits");
assert_eq!(first, second);
}
#[test]
fn require_64_bit_pointer_width_rejects_32() {
assert!(matches!(
require_64_bit_pointer_width(32),
Err(LowerError::UnsupportedPointerWidth { bits: 32 })
));
}
#[test]
fn require_64_bit_pointer_width_accepts_64() {
assert!(require_64_bit_pointer_width(64).is_ok());
}
#[test]
fn compile_aot_rejects_i686_without_panic() {
let error = compile_aot(&test_program(), "i686-unknown-linux-gnu")
.expect_err("i686 AOT target is rejected");
assert!(matches!(error, AotError::TargetLookup(_)));
}
}