use std::cell::RefCell;
use std::marker::PhantomData;
use luau_vm::Thread as VmThread;
use luau_vm::native::{NativeCallContext, NativeCallResult};
use crate::callback::{Arguments, Callback, CallbackReturn, finish_callback, push_callback};
use crate::error::Error;
use crate::lua::LuaRef;
use crate::lua::runtime::RuntimeData;
use crate::userdata::AnyUserdata;
pub(super) struct UserdataMethodCallback<T, M> {
pub(super) method: M,
pub(super) expected_results: Option<usize>,
pub(super) _marker: PhantomData<fn(T)>,
}
pub(super) struct UserdataMethodMutCallback<T, M> {
pub(super) method: RefCell<M>,
pub(super) expected_results: Option<usize>,
pub(super) _marker: PhantomData<fn(T)>,
}
pub(super) struct UserdataFunctionCallback<F> {
pub(super) function: F,
pub(super) expected_results: Option<usize>,
}
pub(super) struct UserdataFunctionMutCallback<F> {
pub(super) function: RefCell<F>,
pub(super) expected_results: Option<usize>,
}
impl<T, M> Callback for UserdataMethodCallback<T, M>
where
T: 'static,
M: for<'call> Fn(LuaRef<'call>, &T, Arguments<'call>) -> Result<CallbackReturn<'call>, Error>
+ 'static,
{
fn call<'call>(
&self,
context: NativeCallContext<'call>,
runtime: &'call RuntimeData,
) -> NativeCallResult {
let thread = context.raw_thread();
let lua = LuaRef::new(thread, runtime);
let safe_thread = lua.current_thread();
let arguments = Arguments::new(lua, 1, context.arg_count().saturating_sub(1));
let result = (|| {
let userdata = AnyUserdata::from_stack(&safe_thread, 1)?;
let userdata = userdata.borrow::<T>()?;
(self.method)(lua, &userdata, arguments)
})();
finish_callback(thread, runtime, result, self.expected_results)
}
}
impl<T, M> Callback for UserdataMethodMutCallback<T, M>
where
T: 'static,
M: for<'call> FnMut(
LuaRef<'call>,
&mut T,
Arguments<'call>,
) -> Result<CallbackReturn<'call>, Error>
+ 'static,
{
fn call<'call>(
&self,
context: NativeCallContext<'call>,
runtime: &'call RuntimeData,
) -> NativeCallResult {
let thread = context.raw_thread();
let lua = LuaRef::new(thread, runtime);
let safe_thread = lua.current_thread();
let arguments = Arguments::new(lua, 1, context.arg_count().saturating_sub(1));
let result = (|| {
let userdata = AnyUserdata::from_stack(&safe_thread, 1)?;
let mut userdata = userdata.borrow_mut::<T>()?;
let mut method = self
.method
.try_borrow_mut()
.map_err(|_| Error::RecursiveMutCallback)?;
method(lua, &mut userdata, arguments)
})();
finish_callback(thread, runtime, result, self.expected_results)
}
}
impl<F> Callback for UserdataFunctionCallback<F>
where
F: for<'call> Fn(LuaRef<'call>, Arguments<'call>) -> Result<CallbackReturn<'call>, Error>
+ 'static,
{
fn call<'call>(
&self,
context: NativeCallContext<'call>,
runtime: &'call RuntimeData,
) -> NativeCallResult {
let thread = context.raw_thread();
let lua = LuaRef::new(thread, runtime);
let arguments = Arguments::new(lua, 0, context.arg_count());
finish_callback(
thread,
runtime,
(self.function)(lua, arguments),
self.expected_results,
)
}
}
impl<F> Callback for UserdataFunctionMutCallback<F>
where
F: for<'call> FnMut(LuaRef<'call>, Arguments<'call>) -> Result<CallbackReturn<'call>, Error>
+ 'static,
{
fn call<'call>(
&self,
context: NativeCallContext<'call>,
runtime: &'call RuntimeData,
) -> NativeCallResult {
let thread = context.raw_thread();
let lua = LuaRef::new(thread, runtime);
let arguments = Arguments::new(lua, 0, context.arg_count());
let result = match self.function.try_borrow_mut() {
Ok(mut function) => {
runtime.invalidate_managed_safe_env();
function(lua, arguments)
}
Err(_) => Err(Error::RecursiveMutCallback),
};
finish_callback(thread, runtime, result, self.expected_results)
}
}
pub(super) fn push_userdata_callback(
thread: &VmThread,
callback: Box<dyn Callback>,
) -> Result<(), Error> {
push_callback(thread, callback).map(|_| ())
}