luau 0.732.0

Safe lifetime-bound Rust embedding API for the Luau runtime
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(|_| ())
}