use core::mem::size_of;
use core::ptr::{self, NonNull};
use luau_common::{ByteSlice, flags};
use crate::gc::GcObject;
use crate::handle::RawHandle;
use crate::handle::sealed::Sealed;
use crate::memory::{LuaPage, MemoryRuntime};
use crate::metamethod::{MetamethodRuntime, TmEvent};
use crate::string::{LUA_MIN_STRING_TABLE_SIZE, StringRuntime};
use crate::table::TableRuntime;
use crate::thread::{LUA_MIN_STACK, Thread};
use crate::types::LUA_TTHREAD;
use crate::value::{RawTValue, TValue, TValueCursor};
use crate::{VmErrorResult, VmResult};
use super::{
BASIC_CI_SIZE, EXTRA_STACK, GlobalState, INITIAL_STACK_SIZE, LUA_ERRERRMSG, LUA_MEMERRMSG,
RawCallInfo, RawLuaState, THREAD_STATUS_OK, ThreadState,
};
#[allow(
clippy::missing_safety_doc,
reason = "all methods share the capability-level safety contract"
)]
pub trait ThreadLifecycle: Sealed {
unsafe fn allocation_size(&self) -> usize;
unsafe fn new_thread_internal(&self) -> VmErrorResult<Thread>;
unsafe fn free_thread(&self, thread: &Thread, page: LuaPage);
}
impl ThreadLifecycle for Thread {
unsafe fn allocation_size(&self) -> usize {
unsafe {
size_of::<RawLuaState>()
+ size_of::<RawTValue>()
* self.as_ptr().as_ref().unwrap_unchecked().stack_size as usize
+ size_of::<RawCallInfo>()
* self.as_ptr().as_ref().unwrap_unchecked().size_ci as usize
}
}
unsafe fn new_thread_internal(&self) -> VmErrorResult<Thread> {
unsafe {
let global = self.global();
let memcat = self.as_ptr().as_ref().unwrap_unchecked().active_memcat;
let thread = self.new_gco::<Thread>(size_of::<RawLuaState>(), memcat)?;
GcObject::from(&thread).init_header(self, LUA_TTHREAD as u8);
thread.preinit_state(global);
thread.as_ptr().as_mut().unwrap_unchecked().active_memcat = memcat;
self.init_stack(&thread)?;
let source = self.as_ptr().as_ref().unwrap_unchecked();
thread.set_globals(self.globals());
thread.as_ptr().as_mut().unwrap_unchecked().single_step = source.single_step;
debug_assert!(GcObject::from(&thread).is_white());
Ok(thread)
}
}
unsafe fn free_thread(&self, thread: &Thread, page: LuaPage) {
unsafe {
if let Some(user_thread) = self.global().user_thread_callback() {
user_thread(None, thread);
}
self.free_stack(thread);
self.free_gco(
thread.into(),
size_of::<RawLuaState>(),
thread.as_ptr().as_ref().unwrap_unchecked().memcat,
page,
);
}
}
}
impl Thread {
pub(crate) unsafe fn init_stack(&self, thread: &Thread) -> VmErrorResult {
unsafe {
let memcat = thread.as_ptr().as_ref().unwrap_unchecked().active_memcat;
let base_ci = self.new_array::<RawCallInfo>(BASIC_CI_SIZE, memcat)?;
{
let raw = thread.as_ptr().as_mut().unwrap_unchecked();
raw.base_ci = base_ci;
raw.size_ci = BASIC_CI_SIZE as i32;
raw.end_ci = base_ci.add(BASIC_CI_SIZE - 1);
}
let stack = self.new_array::<RawTValue>(INITIAL_STACK_SIZE, memcat)?;
{
let raw = thread.as_ptr().as_mut().unwrap_unchecked();
raw.stack = stack;
raw.stack_size = INITIAL_STACK_SIZE as i32;
raw.stack_last = stack.add(INITIAL_STACK_SIZE - EXTRA_STACK);
}
for index in 0..INITIAL_STACK_SIZE {
TValue::from_raw(NonNull::new_unchecked(stack.add(index))).set_nil();
}
thread.set_current_call_info(thread.base_call_info_cursor());
let function = TValueCursor::from_ptr(stack);
function.value_unchecked().set_nil();
let base = function.add(1);
thread
.base_call_info()
.init_call(function, base.add(LUA_MIN_STACK), 0, None);
thread.set_stack_top(base);
thread.set_stack_base(base);
}
Ok(())
}
pub(crate) unsafe fn free_stack(&self, thread: &Thread) {
unsafe {
let memcat = thread.as_ptr().as_ref().unwrap_unchecked().active_memcat;
self.free_array(
thread.as_ptr().as_ref().unwrap_unchecked().base_ci,
thread.as_ptr().as_ref().unwrap_unchecked().size_ci as usize,
memcat,
);
self.free_array(
thread.as_ptr().as_ref().unwrap_unchecked().stack,
thread.as_ptr().as_ref().unwrap_unchecked().stack_size as usize,
memcat,
);
}
}
pub(crate) unsafe fn preinit_state(&self, global: GlobalState) {
let raw = unsafe { self.as_ptr().as_mut().unwrap_unchecked() };
raw.global = global.as_ptr();
raw.stack = ptr::null_mut();
raw.stack_size = 0;
raw.gt = ptr::null_mut();
raw.open_upval = ptr::null_mut();
raw.size_ci = 0;
raw.native_call_depth = 0;
raw.base_native_call_depth = 0;
raw.status = THREAD_STATUS_OK;
raw.base_ci = ptr::null_mut();
raw.ci = ptr::null_mut();
raw.name_call = ptr::null_mut();
raw.cached_slot = 0;
raw.single_step = false;
raw.is_active = false;
raw.active_memcat = 0;
raw.userdata = ptr::null_mut();
raw.top = ptr::null_mut();
raw.base = ptr::null_mut();
raw.stack_last = ptr::null_mut();
raw.end_ci = ptr::null_mut();
raw.gc_list = ptr::null_mut();
}
pub(crate) unsafe fn open_main_state(&self) -> VmErrorResult {
unsafe {
self.init_stack(self)?;
let globals = self.new_table_internal(0, 2)?;
let registry = self.new_table_internal(0, 2)?;
self.set_globals(globals);
self.global().registry().set_table_value(registry);
self.resize(LUA_MIN_STRING_TABLE_SIZE as i32)?;
self.init()?;
if flags::LuauGcTraceUdata.get() {
let weak_registry = self.new_table_internal(0, 0)?;
let metatable = self.new_table_internal(0, 1)?;
let mode = self.intern_string(b"v".as_bstr())?;
let mode_slot = self
.set_str(metatable, self.global().tm_name(TmEvent::Mode as usize))?
.node_unchecked()
.value_unchecked();
mode_slot.set_string_value(mode);
weak_registry.set_metatable(Some(metatable));
self.global().weak_registry().set_table_value(weak_registry);
}
self.intern_string(LUA_MEMERRMSG.as_bstr())?.fix();
self.intern_string(LUA_ERRERRMSG.as_bstr())?.fix();
let global = self.global();
let global_ref = global.as_ptr().as_mut().unwrap_unchecked();
global_ref.gc_threshold = 4 * global_ref.total_bytes;
}
Ok(())
}
}
pub(crate) unsafe fn open_main_state(thread: &Thread, _: &mut ()) -> VmResult {
unsafe { thread.open_main_state() }?;
Ok(())
}