use luau_common::{ByteSlice, flags};
use crate::call::{ErrorRuntime, ProtectedCall, ThreadStack};
use crate::debug::DebugRuntime;
use crate::function::FunctionRuntime;
use crate::gc::GcBarrier;
use crate::handle::RawHandle;
use crate::native::NativeCallContext;
use crate::state::ThreadState;
use crate::state::{
CallInfoCursor, LUA_CALLINFO_HANDLE, THREAD_STATUS_BREAK, THREAD_STATUS_ERR_ERR,
THREAD_STATUS_ERR_MEM, THREAD_STATUS_ERR_RUN, THREAD_STATUS_OK,
THREAD_STATUS_SCHEDULED_REENTRY, THREAD_STATUS_YIELD,
};
use crate::state::{LUA_OK, ProtectedErrorAction};
use crate::string::StringRuntime;
use crate::thread::Thread;
use crate::value::TValueCursor;
use crate::vm::{PreCallResult, VmCallFrame, VmExecution};
use crate::{VmControl, VmError, VmExit, VmResult};
use super::protected::{ErrorFunctionContext, call_error_function};
impl Thread {
pub(super) unsafe fn restore_stack_limit(&self) -> VmResult {
unsafe {
debug_assert!(
self.stack_last().offset_from(self.restore_stack(0))
== self.as_ptr().as_ref().unwrap_unchecked().stack_size as isize
- crate::state::EXTRA_STACK as isize
);
if self.as_ptr().as_ref().unwrap_unchecked().size_ci as usize
> crate::thread::LUAI_MAX_CALLS
{
let in_use =
self.current_call_info_cursor()
.offset_from(self.base_call_info_cursor()) as usize;
if in_use + 1 < crate::thread::LUAI_MAX_CALLS {
self.realloc_ci(crate::thread::LUAI_MAX_CALLS as i32)?;
}
}
}
Ok(())
}
unsafe fn resume_continue(&self) -> VmResult {
loop {
unsafe {
let status = self.as_ptr().as_ref().unwrap_unchecked().status;
if !matches!(
status,
x if x == THREAD_STATUS_OK || x == THREAD_STATUS_SCHEDULED_REENTRY
) {
break;
}
let current_call_info = self.current_call_info();
if current_call_info == self.base_call_info() {
break;
}
debug_assert_eq!(
self.as_ptr()
.as_ref()
.unwrap_unchecked()
.base_native_call_depth,
self.as_ptr().as_ref().unwrap_unchecked().native_call_depth
);
self.as_ptr().as_mut().unwrap_unchecked().status = THREAD_STATUS_OK;
let closure = current_call_info.function_closure();
if closure.is_native() {
current_call_info.as_ptr().as_mut().unwrap_unchecked().flags &=
!LUA_CALLINFO_HANDLE;
let continuation = closure.native_data().continuation.unwrap_unchecked();
let result_count = match continuation(NativeCallContext::new(self), LUA_OK) {
Ok(count) => count,
Err(VmExit::Control(VmControl::Yield))
if self.as_ptr().as_ref().unwrap_unchecked().status
== THREAD_STATUS_SCHEDULED_REENTRY =>
{
continue;
}
Err(exit) => return Err(exit),
};
let first_result = self.stack_top().sub(result_count);
self.pos_call(first_result);
} else {
if flags::LuauYieldIter2.get()
&& current_call_info.as_ptr().as_ref().unwrap_unchecked().flags
& crate::state::LUA_CALLINFO_OP_YIELD
!= 0
{
self.finish_op()?;
}
match self.execute() {
Ok(()) => {}
Err(VmExit::Control(VmControl::Yield))
if self.as_ptr().as_ref().unwrap_unchecked().status
== THREAD_STATUS_SCHEDULED_REENTRY =>
{
continue;
}
Err(exit) => return Err(exit),
}
}
}
}
Ok(())
}
pub(crate) unsafe fn resume_find_handler(&self) -> Option<CallInfoCursor> {
unsafe {
let mut call_info_cursor = self.current_call_info_cursor();
let base_call_info_cursor = self.base_call_info_cursor();
while call_info_cursor != base_call_info_cursor {
let call_info = call_info_cursor.call_info_unchecked();
let flags = call_info.as_ptr().as_ref().unwrap_unchecked().flags;
if flags & crate::state::LUA_CALLINFO_HANDLE != 0 {
return Some(call_info_cursor);
}
call_info_cursor = call_info_cursor.sub(1);
}
None
}
}
unsafe fn resume_start_error(&self, message: &[u8], argument_count: i32) -> VmResult {
unsafe {
let top = self.stack_top();
let error_slot = top.sub(argument_count as usize);
self.set_stack_top(error_slot);
let message = self
.intern_string(message.as_bstr())
.expect("resume start error message is fixed");
error_slot.value_unchecked().set_string_value(message);
self.set_stack_top(error_slot.add(1));
}
Err(VmError::Runtime.into())
}
pub(crate) unsafe fn resume_start(
&self,
from: Option<&Thread>,
argument_count: i32,
) -> VmResult {
debug_assert!(argument_count >= 0);
unsafe {
debug_assert!(self.stack_top().offset_from(self.stack_base()) as i32 >= argument_count);
if self.as_ptr().as_ref().unwrap_unchecked().status != THREAD_STATUS_YIELD
&& self.as_ptr().as_ref().unwrap_unchecked().status != THREAD_STATUS_BREAK
&& (self.as_ptr().as_ref().unwrap_unchecked().status != THREAD_STATUS_OK
|| self.current_call_info() != self.base_call_info())
{
return self
.resume_start_error(b"cannot resume non-suspended coroutine", argument_count);
}
self.as_ptr().as_mut().unwrap_unchecked().native_call_depth =
from.map_or(0, |thread| {
thread
.as_ptr()
.as_ref()
.unwrap_unchecked()
.native_call_depth
});
if self.as_ptr().as_ref().unwrap_unchecked().native_call_depth
>= crate::thread::LUAI_MAX_NATIVE_CALLS
{
return self.resume_start_error(b"C stack overflow", argument_count);
}
self.increment_native_call_depth();
self.as_ptr()
.as_mut()
.unwrap_unchecked()
.base_native_call_depth =
self.as_ptr().as_ref().unwrap_unchecked().native_call_depth;
self.as_ptr().as_mut().unwrap_unchecked().is_active = true;
self.thread_barrier();
}
Ok(())
}
pub(crate) unsafe fn resume_finish(
&self,
mut result: VmResult,
old_native_call_depth: u16,
) -> VmResult {
while let Err(VmExit::Error(error)) = result {
unsafe {
let Some(handler) = self.resume_find_handler() else {
break;
};
if self.is_yieldable() != 0
&& let Some(callback) = self.global().protected_error_callback()
&& callback(self) == ProtectedErrorAction::Break
{
self.as_ptr().as_mut().unwrap_unchecked().status = THREAD_STATUS_BREAK;
result = Err(VmExit::Control(VmControl::Break));
break;
}
if flags::LuauXpcallFixMessageYieldPath.get() {
self.as_ptr()
.as_mut()
.unwrap_unchecked()
.base_native_call_depth = old_native_call_depth;
} else {
self.as_ptr().as_mut().unwrap_unchecked().native_call_depth =
old_native_call_depth;
self.as_ptr()
.as_mut()
.unwrap_unchecked()
.base_native_call_depth =
self.as_ptr().as_ref().unwrap_unchecked().native_call_depth;
}
self.as_ptr().as_mut().unwrap_unchecked().status = error.status() as u8;
let mut handler = handler;
result = self.raw_run_protected(resume_handle, &mut handler);
}
}
unsafe {
self.as_ptr().as_mut().unwrap_unchecked().native_call_depth = old_native_call_depth - 1;
self.as_ptr()
.as_mut()
.unwrap_unchecked()
.base_native_call_depth =
self.as_ptr().as_ref().unwrap_unchecked().native_call_depth;
self.as_ptr().as_mut().unwrap_unchecked().is_active = false;
match result {
Ok(()) => {
self.expand_stack_limit(self.stack_top());
Ok(())
}
Err(VmExit::Error(error)) => {
self.as_ptr().as_mut().unwrap_unchecked().status = error.status() as u8;
let top = self.stack_top();
self.set_error_object(error.status(), top);
self.current_call_info().set_top(top);
Err(VmExit::Error(error))
}
err => err,
}
}
}
}
pub(crate) unsafe fn resume(thread: &Thread, first_argument: &mut TValueCursor) -> VmResult {
let mut first_argument = *first_argument;
unsafe {
if thread.as_ptr().as_ref().unwrap_unchecked().status == THREAD_STATUS_OK {
let stack_base = thread.stack_base();
debug_assert!(thread.current_call_info() == thread.base_call_info());
debug_assert!(first_argument >= stack_base);
if first_argument == stack_base {
return crate::run_error!(thread, "cannot resume dead coroutine")
.map_err(Into::into);
}
let pre_call_result =
thread.pre_call(first_argument.sub(1), crate::thread::LUA_MULTRET);
if thread.as_ptr().as_ref().unwrap_unchecked().status == THREAD_STATUS_SCHEDULED_REENTRY
{
debug_assert!(matches!(
pre_call_result,
Err(VmExit::Control(VmControl::Yield))
));
first_argument = thread.stack_base();
} else {
let pre_call_result = pre_call_result?;
if pre_call_result != PreCallResult::Lua {
return Ok(());
}
let call_info = thread.current_call_info();
call_info.as_ptr().as_mut().unwrap_unchecked().flags |=
crate::state::LUA_CALLINFO_RETURN;
}
}
if thread.as_ptr().as_ref().unwrap_unchecked().status != THREAD_STATUS_OK {
let stack_base = thread.stack_base();
debug_assert!(first_argument >= stack_base);
debug_assert!(
thread.as_ptr().as_ref().unwrap_unchecked().status == THREAD_STATUS_YIELD
|| thread.as_ptr().as_ref().unwrap_unchecked().status == THREAD_STATUS_BREAK
|| thread.as_ptr().as_ref().unwrap_unchecked().status
== THREAD_STATUS_SCHEDULED_REENTRY
);
thread.as_ptr().as_mut().unwrap_unchecked().status = THREAD_STATUS_OK;
let closure = thread.current_call_info().function_closure();
if closure.is_native() {
if closure.native_data().continuation.is_none() {
thread.pos_call(first_argument);
} else {
thread.set_stack_base(thread.current_call_info().base());
}
} else {
thread.set_stack_base(thread.current_call_info().base());
}
}
thread.resume_continue()?;
}
Ok(())
}
pub(crate) unsafe fn resume_handle(
thread: &Thread,
call_info_cursor: &mut CallInfoCursor,
) -> VmResult {
let mut call_info_cursor = *call_info_cursor;
let mut call_info = unsafe { call_info_cursor.call_info_unchecked() };
debug_assert!(
unsafe { call_info.as_ptr().as_ref().unwrap_unchecked().flags } & LUA_CALLINFO_HANDLE != 0
);
unsafe {
let closure = call_info.function_closure();
debug_assert!(closure.is_native());
debug_assert!(closure.native_data().continuation.is_some());
debug_assert_ne!(
thread.as_ptr().as_ref().unwrap_unchecked().status,
THREAD_STATUS_OK
);
call_info.as_ptr().as_mut().unwrap_unchecked().flags &= !LUA_CALLINFO_HANDLE;
let status_byte = thread.as_ptr().as_ref().unwrap_unchecked().status;
thread.as_ptr().as_mut().unwrap_unchecked().status = THREAD_STATUS_OK;
if status_byte != THREAD_STATUS_ERR_RUN {
let top = thread.stack_top();
thread.set_error_object(status_byte as i32, top);
}
let mut status = status_byte as i32;
if call_info.errfunc() != 0 {
let saved_ci = thread.save_ci(call_info_cursor);
let errfunc = call_info.errfunc();
let mut error_context = ErrorFunctionContext {
error_function: call_info
.base()
.add((errfunc - 1) as usize)
.value_unchecked(),
};
let error_result = thread.raw_run_protected(call_error_function, &mut error_context);
if !flags::LuauXpcallFixMessageYieldPath.get() {
thread
.as_ptr()
.as_mut()
.unwrap_unchecked()
.native_call_depth = thread
.as_ptr()
.as_ref()
.unwrap_unchecked()
.base_native_call_depth;
}
status = match error_result {
Ok(()) => THREAD_STATUS_ERR_RUN as i32,
Err(VmExit::Error(VmError::Memory)) if status == THREAD_STATUS_ERR_MEM as i32 => {
THREAD_STATUS_ERR_MEM as i32
}
Err(VmExit::Error(_)) => THREAD_STATUS_ERR_ERR as i32,
Err(exit) => return Err(exit),
};
thread.set_error_object(status, thread.stack_top().sub(1));
call_info_cursor = thread.restore_ci(saved_ci);
call_info = call_info_cursor.call_info_unchecked();
call_info.set_errfunc(0);
}
if flags::LuauXpcallFixMessageYieldPath.get() {
thread
.as_ptr()
.as_mut()
.unwrap_unchecked()
.native_call_depth = thread
.as_ptr()
.as_ref()
.unwrap_unchecked()
.base_native_call_depth;
}
thread.set_current_call_info(call_info_cursor);
thread.close(call_info.base().value_unchecked());
thread.set_stack_base(call_info.base());
call_info.set_top(thread.stack_top());
thread.restore_stack_limit()?;
let continuation = closure.native_data().continuation.unwrap_unchecked();
let result_count = match continuation(NativeCallContext::new(thread), status) {
Ok(count) => count,
Err(VmExit::Control(VmControl::Yield))
if thread.as_ptr().as_ref().unwrap_unchecked().status
== THREAD_STATUS_SCHEDULED_REENTRY =>
{
thread.resume_continue()?;
return Ok(());
}
Err(exit) => return Err(exit),
};
let first_result = thread.stack_top().sub(result_count);
thread.pos_call(first_result);
thread.resume_continue()?;
}
Ok(())
}