use core::cell::Cell;
use std::sync::Arc;
use luau_vm::Thread as VmThread;
use luau_vm::state::{GcInterrupt as VmGcInterrupt, GcPhase as VmGcPhase, InterruptRequest};
use luau_vm::{VmErrorResult, VmResult};
use crate::callback::{raise_callback_error, raise_callback_vm_error};
use crate::error::{Error, Result};
use crate::lua::runtime::RuntimeData;
use crate::lua::{Lua, LuaRef, StackInfo};
#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)]
#[non_exhaustive]
pub enum InterruptAction {
#[default]
Continue,
Yield,
Break,
}
#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)]
#[non_exhaustive]
pub enum InterruptMode {
#[default]
Continuous,
Requested,
}
#[must_use = "retain the handle to request interrupts"]
#[derive(Clone, Debug)]
pub struct InterruptHandle {
request: Arc<InterruptRequest>,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
#[non_exhaustive]
pub enum GcPhase {
Pause,
Propagate,
PropagateAgain,
Atomic,
Sweep,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
#[non_exhaustive]
pub enum GcInterruptStage {
BeforeStep,
AfterStep {
previous_phase: GcPhase,
},
}
#[derive(Clone, Copy)]
pub struct ExecutionInterruptContext<'callback> {
lua: LuaRef<'callback>,
deferred: &'callback Cell<bool>,
}
#[derive(Clone, Copy)]
pub struct PatternInterruptContext<'callback> {
lua: LuaRef<'callback>,
deferred: &'callback Cell<bool>,
}
#[derive(Clone, Copy)]
pub struct GcInterruptContext<'callback> {
lua: LuaRef<'callback>,
stage: GcInterruptStage,
deferred: &'callback Cell<bool>,
}
pub trait InterruptHandler: 'static {
fn mode(&self) -> InterruptMode {
InterruptMode::Continuous
}
fn execution(&self, context: ExecutionInterruptContext<'_>) -> Result<InterruptAction> {
context.defer_interrupt();
Ok(InterruptAction::Continue)
}
fn pattern(&self, context: PatternInterruptContext<'_>) -> Result<()> {
context.defer_interrupt();
Ok(())
}
fn garbage_collection(&self, context: GcInterruptContext<'_>) {
context.defer_interrupt();
}
}
type ExecutionCallback =
Box<dyn for<'callback> Fn(ExecutionInterruptContext<'callback>) -> Result<InterruptAction>>;
type PatternCallback = Box<dyn for<'callback> Fn(PatternInterruptContext<'callback>) -> Result<()>>;
type GcCallback = Box<dyn for<'callback> Fn(GcInterruptContext<'callback>)>;
#[derive(Default)]
pub struct InterruptHooks {
mode: InterruptMode,
execution: Option<ExecutionCallback>,
pattern: Option<PatternCallback>,
garbage_collection: Option<GcCallback>,
}
impl InterruptHooks {
pub const fn new() -> Self {
Self {
mode: InterruptMode::Continuous,
execution: None,
pattern: None,
garbage_collection: None,
}
}
#[must_use]
pub const fn set_mode(mut self, mode: InterruptMode) -> Self {
self.mode = mode;
self
}
#[must_use]
pub fn on_execution<F>(mut self, callback: F) -> Self
where
F: for<'callback> Fn(ExecutionInterruptContext<'callback>) -> Result<InterruptAction>
+ 'static,
{
self.execution = Some(Box::new(callback));
self
}
#[must_use]
pub fn on_pattern<F>(mut self, callback: F) -> Self
where
F: for<'callback> Fn(PatternInterruptContext<'callback>) -> Result<()> + 'static,
{
self.pattern = Some(Box::new(callback));
self
}
#[must_use]
pub fn on_garbage_collection<F>(mut self, callback: F) -> Self
where
F: for<'callback> Fn(GcInterruptContext<'callback>) + 'static,
{
self.garbage_collection = Some(Box::new(callback));
self
}
}
impl InterruptHandler for InterruptHooks {
fn mode(&self) -> InterruptMode {
self.mode
}
fn execution(&self, context: ExecutionInterruptContext<'_>) -> Result<InterruptAction> {
match &self.execution {
Some(callback) => callback(context),
None => {
context.defer_interrupt();
Ok(InterruptAction::Continue)
}
}
}
fn pattern(&self, context: PatternInterruptContext<'_>) -> Result<()> {
match &self.pattern {
Some(callback) => callback(context),
None => {
context.defer_interrupt();
Ok(())
}
}
}
fn garbage_collection(&self, context: GcInterruptContext<'_>) {
if let Some(callback) = &self.garbage_collection {
callback(context);
} else {
context.defer_interrupt();
}
}
}
impl Lua {
pub fn interrupt_handle(&self) -> InterruptHandle {
self.runtime.callbacks().interrupt_handle()
}
pub fn set_interrupt_handler(&mut self, handler: impl InterruptHandler) {
self.runtime
.callbacks_mut()
.set_interrupt_handler(Box::new(handler));
self.install_callbacks();
}
pub fn remove_interrupt_handler(&mut self) {
self.runtime.callbacks_mut().remove_interrupt_handler();
self.install_callbacks();
}
}
impl InterruptHandle {
pub(in crate::hooks) fn new() -> Self {
Self {
request: Arc::new(InterruptRequest::new()),
}
}
pub fn request(&self) {
self.request.request();
}
pub(in crate::hooks) fn clear(&self) {
self.request.clear();
}
pub(in crate::hooks) fn as_ptr(&self) -> *const InterruptRequest {
Arc::as_ptr(&self.request)
}
}
impl<'callback> ExecutionInterruptContext<'callback> {
pub(in crate::hooks) const fn new(
lua: LuaRef<'callback>,
deferred: &'callback Cell<bool>,
) -> Self {
Self { lua, deferred }
}
pub const fn lua(&self) -> LuaRef<'callback> {
self.lua
}
pub fn is_yieldable(&self) -> bool {
self.lua.is_yieldable()
}
pub fn defer_interrupt(&self) {
self.deferred.set(true);
}
}
impl<'callback> PatternInterruptContext<'callback> {
pub(in crate::hooks) const fn new(
lua: LuaRef<'callback>,
deferred: &'callback Cell<bool>,
) -> Self {
Self { lua, deferred }
}
pub const fn lua(&self) -> LuaRef<'callback> {
self.lua
}
pub fn defer_interrupt(&self) {
self.deferred.set(true);
}
}
impl<'callback> GcInterruptContext<'callback> {
pub(in crate::hooks) const fn new(
lua: LuaRef<'callback>,
stage: GcInterruptStage,
deferred: &'callback Cell<bool>,
) -> Self {
Self {
lua,
stage,
deferred,
}
}
pub const fn stage(&self) -> GcInterruptStage {
self.stage
}
pub fn inspect_stack<R>(
&self,
level: usize,
inspect: impl FnOnce(&StackInfo<'_>) -> R,
) -> Result<Option<R>> {
self.lua.inspect_stack(level, inspect)
}
pub fn defer_interrupt(&self) {
self.deferred.set(true);
}
}
pub(in crate::hooks) fn execution_interrupt(thread: &VmThread) -> VmResult {
let runtime = RuntimeData::from_thread(thread);
let result = runtime.callbacks().invoke_interrupt(|handler, delivery| {
runtime.with_thread(thread, || {
let lua = LuaRef::new(thread, runtime);
let action = match handler.execution(ExecutionInterruptContext::new(lua, delivery)) {
Ok(action) => action,
Err(error) => return raise_callback_error(thread, runtime, error),
};
match action {
InterruptAction::Continue => Ok(()),
InterruptAction::Yield if lua.is_yieldable() => unsafe {
thread.yield_current(0).map(|_| ())
},
InterruptAction::Yield => raise_callback_error(
thread,
runtime,
Error::runtime(
"interrupt handler attempted to yield from a non-yieldable boundary",
),
),
InterruptAction::Break => unsafe { thread.break_current().map(|_| ()) },
}
})
});
let Some(result) = result else {
return Ok(());
};
result
}
pub(in crate::hooks) fn pattern_interrupt(thread: &VmThread) -> VmErrorResult {
let runtime = RuntimeData::from_thread(thread);
let result = runtime.callbacks().invoke_interrupt(|handler, delivery| {
runtime.with_thread(thread, || {
let lua = LuaRef::new(thread, runtime);
handler
.pattern(PatternInterruptContext::new(lua, delivery))
.or_else(|error| raise_callback_vm_error(thread, runtime, error))
})
});
let Some(result) = result else {
return Ok(());
};
result
}
pub(in crate::hooks) fn gc_interrupt(thread: &VmThread, event: VmGcInterrupt) -> VmErrorResult {
let runtime = RuntimeData::from_thread(thread);
let stage = match event {
VmGcInterrupt::BeforeStep => GcInterruptStage::BeforeStep,
VmGcInterrupt::AfterStep { previous_phase } => GcInterruptStage::AfterStep {
previous_phase: map_gc_phase(previous_phase),
},
};
let result = runtime.callbacks().invoke_interrupt(|handler, delivery| {
runtime.with_thread(thread, || {
handler.garbage_collection(GcInterruptContext::new(
LuaRef::new(thread, runtime),
stage,
delivery,
))
})
});
let Some(()) = result else {
return Ok(());
};
Ok(())
}
const fn map_gc_phase(phase: VmGcPhase) -> GcPhase {
match phase {
VmGcPhase::Pause => GcPhase::Pause,
VmGcPhase::Propagate => GcPhase::Propagate,
VmGcPhase::PropagateAgain => GcPhase::PropagateAgain,
VmGcPhase::Atomic => GcPhase::Atomic,
VmGcPhase::Sweep => GcPhase::Sweep,
}
}