luau 0.732.0

Safe lifetime-bound Rust embedding API for the Luau runtime
use std::fmt;
use std::io::{self, SeekFrom};
use std::ptr;

use luau_vm::Thread as VmThread;
use luau_vm::VmErrorResult;
use luau_vm::thread::StackGuard;

use crate::error::Error;
use crate::lua::{Lua, LuaRef};
use crate::thread::Thread;
use crate::value::{FromLua, IntoLua, LuaType, Value, ValueRef};

/// A handle to a mutable Luau buffer.
pub struct Buffer<'lua> {
    reference: ValueRef<'lua>,
    data: *mut u8,
    len: usize,
}

impl Lua {
    /// Creates a buffer containing a copy of `data`.
    pub fn create_buffer(&self, data: impl AsRef<[u8]>) -> Result<Buffer<'_>, Error> {
        self.lua_ref().create_buffer(data)
    }

    /// Creates a zero-initialized buffer of `size` bytes.
    pub fn create_buffer_with_capacity(&self, size: usize) -> Result<Buffer<'_>, Error> {
        self.lua_ref().create_buffer_with_capacity(size)
    }
}

impl<'lua> LuaRef<'lua> {
    /// Creates a buffer containing a copy of `data`.
    pub fn create_buffer(&self, data: impl AsRef<[u8]>) -> Result<Buffer<'lua>, Error> {
        self.current_thread().create_buffer(data)
    }

    /// Creates a zero-initialized buffer of `size` bytes.
    pub fn create_buffer_with_capacity(&self, size: usize) -> Result<Buffer<'lua>, Error> {
        self.current_thread().create_buffer_with_capacity(size)
    }
}

impl<'lua> Thread<'lua> {
    pub(crate) fn create_buffer(&self, data: impl AsRef<[u8]>) -> Result<Buffer<'lua>, Error> {
        unsafe {
            let data = data.as_ref();
            let thread = self.as_vm();
            let _stack = StackGuard::new(thread);
            let buffer = thread
                .new_buffer(data.len())
                .map_err(|exit| Error::from_thread_exit(thread, exit))?;
            if !data.is_empty() {
                core::ptr::copy_nonoverlapping(data.as_ptr(), buffer, data.len());
            }
            Buffer::from_stack(self, -1).map_err(|exit| Error::from_thread_exit(thread, exit))
        }
    }

    pub(crate) fn create_buffer_with_capacity(&self, size: usize) -> Result<Buffer<'lua>, Error> {
        unsafe {
            let thread = self.as_vm();
            let _stack = StackGuard::new(thread);
            thread
                .new_buffer(size)
                .map_err(|exit| Error::from_thread_exit(thread, exit))?;
            Buffer::from_stack(self, -1).map_err(|exit| Error::from_thread_exit(thread, exit))
        }
    }
}

impl<'lua> Buffer<'lua> {
    pub(crate) unsafe fn from_stack(thread: &Thread<'lua>, index: i32) -> VmErrorResult<Self> {
        let stack_thread = thread.as_vm();
        debug_assert_ne!(unsafe { stack_thread.is_buffer(index) }, 0);
        let (data, len) = unsafe {
            stack_thread
                .to_buffer(index)
                .expect("buffer stack slot should contain a buffer")
        };
        Ok(Self {
            reference: ValueRef::from_stack(thread, index)?,
            data,
            len,
        })
    }

    /// Creates another rooted handle to this buffer.
    pub fn try_clone(&self) -> Result<Self, Error> {
        Ok(Self {
            reference: self.reference.try_clone()?,
            data: self.data,
            len: self.len,
        })
    }

    /// Copies the buffer contents into a vector.
    pub fn to_vec(&self) -> Vec<u8> {
        if self.len == 0 {
            return Vec::new();
        }

        // `reference` roots the fixed-size buffer while `self` is alive.
        unsafe { core::slice::from_raw_parts(self.data.cast_const(), self.len).to_vec() }
    }

    /// Returns the buffer length in bytes.
    pub fn len(&self) -> usize {
        self.len
    }

    /// Returns whether the buffer has a length of zero.
    pub fn is_empty(&self) -> bool {
        self.len() == 0
    }

    #[track_caller]
    /// Reads `N` bytes beginning at `offset`.
    ///
    /// # Panics
    ///
    /// Panics if the requested range is outside the buffer.
    pub fn read_bytes<const N: usize>(&self, offset: usize) -> [u8; N] {
        assert_buffer_range(offset, N, self.len);

        let mut bytes = [0u8; N];
        if N == 0 {
            return bytes;
        }

        unsafe {
            ptr::copy_nonoverlapping(self.data.add(offset), bytes.as_mut_ptr(), N);
        }
        bytes
    }

    #[track_caller]
    /// Writes `bytes` beginning at `offset`.
    ///
    /// # Panics
    ///
    /// Panics if the requested range is outside the buffer.
    pub fn write_bytes(&self, offset: usize, bytes: &[u8]) {
        assert_buffer_range(offset, bytes.len(), self.len);
        if bytes.is_empty() {
            return;
        }

        unsafe {
            ptr::copy(bytes.as_ptr(), self.data.add(offset), bytes.len());
        }
    }

    /// Creates a cursor that reads and writes this buffer in place.
    pub fn cursor(self) -> impl io::Read + io::Write + io::Seek {
        BufferCursor {
            buffer: self,
            offset: 0,
        }
    }

    pub(crate) fn push_to(&self, target: impl AsRef<VmThread>) -> Result<(), Error> {
        self.reference.push_to(target)
    }

    pub(crate) fn thread(&self) -> Thread<'lua> {
        self.reference.thread()
    }

    pub(crate) fn pointer(&self) -> *const () {
        self.reference.pointer()
    }
}

impl LuaType for Buffer<'_> {
    fn push_type_key(thread: impl AsRef<VmThread>) -> Result<(), Error> {
        let thread = thread.as_ref();
        unsafe {
            thread
                .new_buffer(0)
                .map(drop)
                .map_err(|exit| Error::from_thread_exit(thread, exit))
        }
    }
}

impl<'lua, 'buffer> IntoLua<'lua> for Buffer<'buffer>
where
    'buffer: 'lua,
{
    fn into_lua(self, _: crate::LuaRef<'lua>) -> Result<Value<'lua>, Error> {
        Ok(Value::Buffer(self))
    }

    unsafe fn push_into_stack(self, thread: &Thread<'lua>) -> Result<(), Error> {
        self.push_to(thread)
    }
}

impl<'lua, 'buffer> IntoLua<'lua> for &Buffer<'buffer>
where
    'buffer: 'lua,
{
    fn into_lua(self, _: crate::LuaRef<'lua>) -> Result<Value<'lua>, Error> {
        Ok(Value::Buffer(self.try_clone()?))
    }

    unsafe fn push_into_stack(self, thread: &Thread<'lua>) -> Result<(), Error> {
        self.push_to(thread)
    }
}

impl<'lua> FromLua<'lua> for Buffer<'lua> {
    fn from_lua(value: Value<'lua>, _: crate::LuaRef<'lua>) -> Result<Self, Error> {
        match value {
            Value::Buffer(buffer) => Ok(buffer),
            value => Err(Error::from_lua_conversion(
                value.type_name(),
                "buffer",
                None,
            )),
        }
    }
}

struct BufferCursor<'lua> {
    buffer: Buffer<'lua>,
    offset: usize,
}

impl io::Read for BufferCursor<'_> {
    fn read(&mut self, output: &mut [u8]) -> io::Result<usize> {
        if output.is_empty() || self.offset == self.buffer.len {
            return Ok(0);
        }

        let count = output.len().min(self.buffer.len - self.offset);
        unsafe {
            ptr::copy_nonoverlapping(
                self.buffer.data.add(self.offset),
                output.as_mut_ptr(),
                count,
            );
        }
        self.offset += count;
        Ok(count)
    }
}

impl io::Write for BufferCursor<'_> {
    fn write(&mut self, input: &[u8]) -> io::Result<usize> {
        if input.is_empty() || self.offset == self.buffer.len {
            return Ok(0);
        }

        let count = input.len().min(self.buffer.len - self.offset);
        unsafe {
            ptr::copy(input.as_ptr(), self.buffer.data.add(self.offset), count);
        }
        self.offset += count;
        Ok(count)
    }

    fn flush(&mut self) -> io::Result<()> {
        Ok(())
    }
}

impl io::Seek for BufferCursor<'_> {
    fn seek(&mut self, pos: SeekFrom) -> io::Result<u64> {
        let len = self.buffer.len();
        let offset = match pos {
            SeekFrom::Start(offset) => i128::from(offset),
            SeekFrom::End(offset) => len as i128 + i128::from(offset),
            SeekFrom::Current(offset) => self.offset as i128 + i128::from(offset),
        };

        if offset < 0 {
            return Err(io::Error::new(
                io::ErrorKind::InvalidInput,
                "invalid seek to a negative position",
            ));
        }

        let Ok(offset) = usize::try_from(offset) else {
            return Err(io::Error::new(
                io::ErrorKind::InvalidInput,
                "invalid seek to a position beyond the end of the buffer",
            ));
        };
        if offset > len {
            return Err(io::Error::new(
                io::ErrorKind::InvalidInput,
                "invalid seek to a position beyond the end of the buffer",
            ));
        }

        self.offset = offset;
        Ok(offset as u64)
    }
}

#[track_caller]
fn assert_buffer_range(offset: usize, count: usize, len: usize) {
    let Some(end) = offset.checked_add(count) else {
        panic!("buffer access out of bounds");
    };

    if end > len {
        panic!("buffer access out of bounds");
    }
}

impl PartialEq for Buffer<'_> {
    fn eq(&self, other: &Self) -> bool {
        self.reference == other.reference
    }
}

impl Eq for Buffer<'_> {}

impl fmt::Debug for Buffer<'_> {
    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
        formatter
            .debug_tuple("Buffer")
            .field(&self.reference)
            .finish()
    }
}