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};
pub struct Buffer<'lua> {
reference: ValueRef<'lua>,
data: *mut u8,
len: usize,
}
impl Lua {
pub fn create_buffer(&self, data: impl AsRef<[u8]>) -> Result<Buffer<'_>, Error> {
self.lua_ref().create_buffer(data)
}
pub fn create_buffer_with_capacity(&self, size: usize) -> Result<Buffer<'_>, Error> {
self.lua_ref().create_buffer_with_capacity(size)
}
}
impl<'lua> LuaRef<'lua> {
pub fn create_buffer(&self, data: impl AsRef<[u8]>) -> Result<Buffer<'lua>, Error> {
self.current_thread().create_buffer(data)
}
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,
})
}
pub fn try_clone(&self) -> Result<Self, Error> {
Ok(Self {
reference: self.reference.try_clone()?,
data: self.data,
len: self.len,
})
}
pub fn to_vec(&self) -> Vec<u8> {
if self.len == 0 {
return Vec::new();
}
unsafe { core::slice::from_raw_parts(self.data.cast_const(), self.len).to_vec() }
}
pub fn len(&self) -> usize {
self.len
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
#[track_caller]
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]
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());
}
}
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()
}
}