luau 0.732.0

Safe lifetime-bound Rust embedding API for the Luau runtime
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 {
    /// Creates a Luau function from an immutable Rust callback.
    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)
    }

    /// Creates a Luau function from a mutable Rust callback.
    ///
    /// Recursively invoking the callback returns
    /// [`Error::RecursiveMutCallback`].
    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> {
    /// Creates a Luau function from an immutable Rust callback.
    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)))
    }

    /// Creates a Luau function from a mutable Rust callback.
    ///
    /// Recursively invoking the callback returns
    /// [`Error::RecursiveMutCallback`].
    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<'_> {
    /// Wraps an immutable Rust callback as an [`IntoLua`] value.
    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 }
    }

    /// Wraps a mutable Rust callback as an [`IntoLua`] value.
    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))
    }
}

/// # Safety
///
/// The returned scoped callback must be dropped before `callback`'s captures can expire.
pub(crate) unsafe fn create_scoped<'lua, 'env>(
    lua: &LuaRef<'lua>,
    callback: Box<dyn Callback + 'env>,
) -> Result<ScopedCallback<'lua>, Error> {
    unsafe {
        // The scoped handle clears the entry before the erased callback lifetime can expire.
        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();
    }
}