use core::ptr::NonNull;
use std::cell::RefCell;
use std::mem;
use luau_vm::native::{NativeCallContext, NativeCallResult};
use luau_vm::thread::StackGuard;
use crate::callback::{
Arguments, Callback, CallbackEntry, CallbackReturn, finish_callback, invalidate_callback,
push_callback,
};
use crate::error::Error;
use crate::lua::runtime::RuntimeData;
use crate::lua::{Lua, LuaRef};
use crate::value::{IntoLua, Value};
use super::Function;
struct WrappedFunction<F> {
function: F,
}
struct WrappedFunctionMut<F> {
function: F,
}
impl Lua {
pub fn create_function<'lua, F>(&'lua self, function: F) -> Result<Function<'lua>, Error>
where
F: for<'call> Fn(LuaRef<'call>, Arguments<'call>) -> Result<CallbackReturn<'call>, Error>
+ 'static,
{
self.lua_ref().create_function(function)
}
pub fn create_function_mut<'lua, F>(&'lua self, function: F) -> Result<Function<'lua>, Error>
where
F: for<'call> FnMut(
LuaRef<'call>,
Arguments<'call>,
) -> Result<CallbackReturn<'call>, Error>
+ 'static,
{
self.lua_ref().create_function_mut(function)
}
}
impl<'lua> LuaRef<'lua> {
pub fn create_function<F>(&self, function: F) -> Result<Function<'lua>, Error>
where
F: for<'call> Fn(LuaRef<'call>, Arguments<'call>) -> Result<CallbackReturn<'call>, Error>
+ 'static,
{
create(self, Box::new(CallbackFn::new(function)))
}
pub fn create_function_mut<F>(&self, function: F) -> Result<Function<'lua>, Error>
where
F: for<'call> FnMut(
LuaRef<'call>,
Arguments<'call>,
) -> Result<CallbackReturn<'call>, Error>
+ 'static,
{
create(self, Box::new(CallbackMut::new(function)))
}
}
impl Function<'_> {
pub fn wrap<F>(function: F) -> impl for<'value> IntoLua<'value>
where
F: for<'call> Fn(LuaRef<'call>, Arguments<'call>) -> Result<CallbackReturn<'call>, Error>
+ 'static,
{
WrappedFunction { function }
}
pub fn wrap_mut<F>(function: F) -> impl for<'value> IntoLua<'value>
where
F: for<'call> FnMut(
LuaRef<'call>,
Arguments<'call>,
) -> Result<CallbackReturn<'call>, Error>
+ 'static,
{
WrappedFunctionMut { function }
}
}
impl<'lua, F> IntoLua<'lua> for WrappedFunction<F>
where
F: for<'call> Fn(LuaRef<'call>, Arguments<'call>) -> Result<CallbackReturn<'call>, Error>
+ 'static,
{
fn into_lua(self, lua: crate::LuaRef<'lua>) -> Result<Value<'lua>, Error> {
create(&lua, Box::new(CallbackFn::new(self.function))).map(Value::Function)
}
}
impl<'lua, F> IntoLua<'lua> for WrappedFunctionMut<F>
where
F: for<'call> FnMut(LuaRef<'call>, Arguments<'call>) -> Result<CallbackReturn<'call>, Error>
+ 'static,
{
fn into_lua(self, lua: crate::LuaRef<'lua>) -> Result<Value<'lua>, Error> {
create(&lua, Box::new(CallbackMut::new(self.function))).map(Value::Function)
}
}
pub(crate) struct ScopedCallback<'lua> {
function: Function<'lua>,
entry: NonNull<CallbackEntry>,
}
pub(crate) struct CallbackFn<F> {
function: F,
}
impl<F> CallbackFn<F> {
pub(crate) const fn new(function: F) -> Self {
Self { function }
}
}
pub(crate) struct CallbackMut<F> {
function: RefCell<F>,
}
impl<F> CallbackMut<F> {
pub(crate) const fn new(function: F) -> Self {
Self {
function: RefCell::new(function),
}
}
}
impl<F> Callback for CallbackFn<F>
where
F: for<'call> Fn(LuaRef<'call>, Arguments<'call>) -> Result<CallbackReturn<'call>, Error>,
{
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), None)
}
}
impl<F> Callback for CallbackMut<F>
where
F: for<'call> FnMut(LuaRef<'call>, Arguments<'call>) -> Result<CallbackReturn<'call>, Error>,
{
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, None)
}
}
pub(crate) fn create<'lua>(
lua: &LuaRef<'lua>,
callback: Box<dyn Callback>,
) -> Result<Function<'lua>, Error> {
unsafe {
let thread = lua.current_thread();
let vm_thread = thread.as_vm();
let _stack = StackGuard::new(vm_thread);
push_callback(vm_thread, callback)?;
Function::from_stack(&thread, -1).map_err(|exit| Error::from_thread_exit(vm_thread, exit))
}
}
pub(crate) unsafe fn create_scoped<'lua, 'env>(
lua: &LuaRef<'lua>,
callback: Box<dyn Callback + 'env>,
) -> Result<ScopedCallback<'lua>, Error> {
unsafe {
let callback = mem::transmute::<Box<dyn Callback + 'env>, Box<dyn Callback>>(callback);
let thread = lua.current_thread();
let vm_thread = thread.as_vm();
let _stack = StackGuard::new(vm_thread);
let entry = push_callback(vm_thread, callback)?;
let function = match Function::from_stack(&thread, -1) {
Ok(function) => function,
Err(exit) => {
invalidate_callback(entry);
return Err(Error::from_thread_exit(vm_thread, exit));
}
};
Ok(ScopedCallback { function, entry })
}
}
impl<'lua> ScopedCallback<'lua> {
pub(crate) fn function(&self) -> &Function<'lua> {
&self.function
}
fn invalidate(&mut self) {
unsafe {
invalidate_callback(self.entry);
}
}
}
impl Drop for ScopedCallback<'_> {
fn drop(&mut self) {
self.invalidate();
}
}