use crate::codegen::FunctionTranslator;
use crate::compiler::JIT;
use crate::error::{JITError, JITErrorType};
use crate::prelude::SSARepr;
use cranelift_codegen::ir::{types, InstBuilder, MemFlags, TrapCode, Value};
use cranelift_jit::JITModule;
use cranelift_module::{DataDescription, DataId, Linkage, Module, ModuleError};
use edlc_core::prelude::mir_type::MirTypeId;
use edlc_core::prelude::{MirError, MirPhase};
use std::sync::OnceLock;
use std::{mem, ptr, slice};
#[macro_export]
macro_rules! jit_intrinsic_panic(
($val:expr) => (
log::error!("JIT panic: {}", $val);
$crate::compiler::panic_handle::RawPanicHandle::panic_global($val).unwrap();
return unsafe { std::mem::MaybeUninit::uninit().assume_init() };
);
);
pub use jit_intrinsic_panic;
use crate::trap;
static mut GLOBAL_JIT_PANIC_HANDLE: OnceLock<RawPanicHandle> = OnceLock::new();
pub struct RawPanicHandle {
location: *const u8,
size: usize,
#[allow(dead_code)]
stack_trace_size: usize,
target_ptr_size: usize,
}
pub struct PanicHandle {
location: DataId,
size: usize,
stack_trace_size: usize,
}
impl PanicHandle {
pub fn new(
data_context: &mut DataDescription,
module: &mut JITModule,
size: usize,
stack_trace_size: usize,
) -> Result<Self, ModuleError> {
assert!(size <= 256, "stack trace message in panic handler must be <= 256 bytes in size");
let header_size = module.target_config().pointer_bytes() as usize * 2;
data_context.define_zeroinit(header_size + size * stack_trace_size);
let id = module
.declare_data("__panic_handle", Linkage::Export, true, false)?;
module.define_data(id, data_context)?;
module.finalize_definitions()?;
data_context.clear();
Ok(Self {
location: id,
size,
stack_trace_size,
})
}
pub fn set_global(&self, module: &JITModule) -> Result<(), JITError> {
let (ptr, _) = module
.get_finalized_data(self.location);
unsafe {
#[allow(static_mut_refs)]
GLOBAL_JIT_PANIC_HANDLE.set(RawPanicHandle {
location: ptr,
size: self.size,
stack_trace_size: self.stack_trace_size,
target_ptr_size: module.target_config().pointer_bytes() as usize,
}).map_err(|_| JITError {
ty: JITErrorType::RuntimeError("tried to redefine global panic handle".to_string())
})?;
}
Ok(())
}
}
impl<Runtime> JIT<Runtime> {
pub unsafe fn has_panicked(&self) -> Result<bool, MirError<JIT<Runtime>>> {
let (ptr, _) = self.module
.get_finalized_data(self.panic_handle.location);
let has_panicked = ptr::read::<u8>(ptr);
if has_panicked != 0 {
return Ok(true)
}
Ok(false)
}
pub unsafe fn unwind_panic(&self) -> Result<(), JITError> {
let (ptr, _) = self.module
.get_finalized_data(self.panic_handle.location);
let has_panicked = ptr::read::<u8>(ptr);
if has_panicked == 0 {
return Ok(());
}
let mut raw_ptr = ptr as usize + self.module.target_config().pointer_bytes() as usize;
let stack_trace_size = ptr::read::<usize>(raw_ptr as *const u8 as *const usize);
raw_ptr += self.module.target_config().pointer_bytes() as usize;
let mut stack_trace = Vec::new();
for depth in 0..stack_trace_size {
let len = ptr::read::<u8>(raw_ptr as *const u8) as usize;
raw_ptr += 1;
let slice = slice::from_raw_parts(raw_ptr as *const u8, len);
stack_trace.push(format!("{}: {}\n", depth, std::str::from_utf8_unchecked(slice)));
raw_ptr += self.panic_handle.size - 1;
}
self.reset_panic_handler()?;
Err(JITError {
ty: JITErrorType::RuntimePanic(stack_trace),
})
}
pub unsafe fn reset_panic_handler(&self) -> Result<(), JITError> {
let (ptr, _) = self.module
.get_finalized_data(self.panic_handle.location);
ptr::write_bytes(ptr as *mut u8, 0, mem::size_of::<usize>() * 2);
Ok(())
}
pub unsafe fn panic(&self, msg: &str) -> Result<(), JITError> {
let (ptr, _) = self.module
.get_finalized_data(self.panic_handle.location);
let mut raw_ptr = ptr as usize;
ptr::write::<u8>(raw_ptr as *mut u8, 0xff);
raw_ptr += self.module.target_config().pointer_bytes() as usize;
ptr::write::<usize>(raw_ptr as *mut usize, 1);
raw_ptr += self.module.target_config().pointer_bytes() as usize;
let bytes = msg.as_bytes();
let len = usize::min(bytes.len(), self.panic_handle.size - 1);
ptr::write::<u8>(raw_ptr as *mut u8, len as u8);
for (idx, &b) in bytes.iter().enumerate() {
raw_ptr += 1;
if idx >= len {
break;
}
ptr::write::<u8>(raw_ptr as *mut u8, b);
}
Ok(())
}
}
impl RawPanicHandle {
pub fn has_global_panicked() -> Result<bool, JITError> {
unsafe {
#[allow(static_mut_refs)]
if let Some(global) = GLOBAL_JIT_PANIC_HANDLE.get() {
global.has_panicked()
} else {
Ok(false)
}
}
}
pub fn unwind_global() -> Result<(), JITError> {
unsafe {
#[allow(static_mut_refs)]
if let Some(global) = GLOBAL_JIT_PANIC_HANDLE.get() {
global.unwind_panic()
} else {
Err(JITError {
ty: JITErrorType::RuntimeError(
"Tried to unwind global panic on empty panic handler".to_string())
})
}
}
}
pub fn unwind_global_no_reset() -> Result<(), JITError> {
unsafe {
#[allow(static_mut_refs)]
if let Some(global) = GLOBAL_JIT_PANIC_HANDLE.get() {
global.unwind_panic_no_reset()
} else {
Err(JITError {
ty: JITErrorType::RuntimeError(
"Tried to unwind global panic on empty panic handler".to_string())
})
}
}
}
pub fn reset_global() -> Result<(), JITError> {
unsafe {
#[allow(static_mut_refs)]
if let Some(global) = GLOBAL_JIT_PANIC_HANDLE.get() {
global.reset()
} else {
Err(JITError {
ty: JITErrorType::RuntimeError(
"Tried to reset global panic on empty panic handler".to_string())
})
}
}
}
pub fn panic_global(msg: &str) -> Result<(), JITError> {
unsafe {
#[allow(static_mut_refs)]
if let Some(global) = GLOBAL_JIT_PANIC_HANDLE.get() {
global.panic(msg)
} else {
Err(JITError {
ty: JITErrorType::RuntimeError(
"Tried to invoke global panic on empty panic handler".to_string())
})
}
}
}
pub fn has_panicked(&self) -> Result<bool, JITError> {
unsafe {
let has_panicked = ptr::read::<u8>(self.location);
if has_panicked != 0 {
return Ok(true);
}
}
Ok(false)
}
pub fn unwind_panic(&self) -> Result<(), JITError> {
unsafe {
let res = self.unwind_panic_no_reset();
if res.is_err() {
self.reset()?;
}
res
}
}
pub unsafe fn unwind_panic_no_reset(&self) -> Result<(), JITError> {
let has_panicked = ptr::read::<u8>(self.location);
if has_panicked == 0 {
return Ok(());
}
let mut raw_ptr = self.location as usize + self.target_ptr_size;
let stack_trace_size = ptr::read::<usize>(raw_ptr as *const u8 as *const usize);
raw_ptr += self.target_ptr_size;
let mut stack_trace = Vec::new();
for depth in 0..stack_trace_size {
let len = ptr::read::<u8>(raw_ptr as *const u8) as usize;
raw_ptr += 1;
let slice = slice::from_raw_parts(raw_ptr as *const u8, len);
stack_trace.push(format!("{}: {}\n", depth, std::str::from_utf8_unchecked(slice)));
raw_ptr += self.size - 1;
}
Err(JITError {
ty: JITErrorType::RuntimePanic(stack_trace),
})
}
pub fn reset(&self) -> Result<(), JITError> {
unsafe {
ptr::write_bytes(self.location as *mut u8, 0, self.target_ptr_size * 2);
}
Ok(())
}
pub fn panic(&self, msg: &str) -> Result<(), JITError> {
unsafe {
let mut raw_ptr = self.location as usize;
ptr::write::<u8>(raw_ptr as *mut u8, 0xff);
raw_ptr += self.target_ptr_size;
ptr::write::<usize>(raw_ptr as *mut usize, 1);
raw_ptr += self.target_ptr_size;
let bytes = msg.as_bytes();
let len = usize::min(bytes.len(), self.size - 1);
ptr::write::<u8>(raw_ptr as *mut u8, len as u8);
for (idx, &b) in bytes.iter().enumerate() {
raw_ptr += 1;
if idx >= len {
break;
}
ptr::write::<u8>(raw_ptr as *mut u8, b);
}
}
Ok(())
}
}
impl<'jit, Runtime> FunctionTranslator<'jit, Runtime> {
fn store_stack_trace_entry(&mut self, msg: &str, ptr: Value) -> Result<(), MirError<JIT<Runtime>>> {
let bytes = msg.as_bytes();
let len = usize::min(bytes.len(), self.panic_handle.size - 1);
let val = self.builder.ins().iconst(types::I8, len as i64);
let mut off = self.module.target_config().pointer_bytes() as i32 * 2;
self.builder
.ins()
.store(MemFlags::new(), val, ptr, off);
for (idx, &b) in bytes.iter().enumerate() {
off += 1;
if idx >= len {
break;
}
let val = self.builder
.ins()
.iconst(types::I8, b as i64);
self.builder
.ins()
.store(MemFlags::new(), val, ptr, off);
}
Ok(())
}
pub fn panic(&mut self, msg: &str) -> Result<(), MirError<JIT<Runtime>>> {
let data = self.module.declare_data_in_func(
self.panic_handle.location,
self.builder.func
);
let ptr = self.builder.ins().symbol_value(
self.module.target_config().pointer_type(),
data,
);
let val = self.builder.ins().iconst(types::I8, 0xff);
self.builder.ins().store(MemFlags::new(), val, ptr, 0);
let ptr_ty = self.module.target_config().pointer_type();
let val = self.builder.ins().iconst(ptr_ty, 1);
self.builder.ins().store(MemFlags::new(), val, ptr, ptr_ty.bytes() as i32);
self.store_stack_trace_entry(msg, ptr)?;
self.builder.ins().trap(TrapCode::unwrap_user(trap::EXPLICIT_PANIC));
Ok(())
}
}