luau-vm 0.732.0

Pure-Rust Luau virtual machine, garbage collector, and standard libraries
Documentation
use core::mem::MaybeUninit;
use core::slice;

use luau_common::ByteSlice;

use crate::native::{NativeCallContext, NativeCallResult, NativeFunction};
use crate::string::MAX_STRING_SIZE;
use crate::thread::{LUA_BUFFER_SIZE, LuaStringBuilder, LuaStringBuilderStorage, Thread};

mod format;
mod pack;
mod pattern;

use format::string_format;
use pack::{string_pack, string_pack_size, string_unpack};
use pattern::{string_find, string_gmatch, string_gsub, string_match};

const L_ESC: u8 = b'%';

static STRING_LIB: [NativeFunction; 17] = [
    NativeFunction {
        name: "byte",
        function: string_byte,
    },
    NativeFunction {
        name: "char",
        function: string_char,
    },
    NativeFunction {
        name: "find",
        function: string_find,
    },
    NativeFunction {
        name: "format",
        function: string_format,
    },
    NativeFunction {
        name: "gmatch",
        function: string_gmatch,
    },
    NativeFunction {
        name: "gsub",
        function: string_gsub,
    },
    NativeFunction {
        name: "len",
        function: string_len,
    },
    NativeFunction {
        name: "lower",
        function: string_lower,
    },
    NativeFunction {
        name: "match",
        function: string_match,
    },
    NativeFunction {
        name: "rep",
        function: string_rep,
    },
    NativeFunction {
        name: "reverse",
        function: string_reverse,
    },
    NativeFunction {
        name: "sub",
        function: string_sub,
    },
    NativeFunction {
        name: "upper",
        function: string_upper,
    },
    NativeFunction {
        name: "split",
        function: string_split,
    },
    NativeFunction {
        name: "pack",
        function: string_pack,
    },
    NativeFunction {
        name: "packsize",
        function: string_pack_size,
    },
    NativeFunction {
        name: "unpack",
        function: string_unpack,
    },
];

/// `posrelat`
fn pos_relat(pos: i32, len: usize) -> i32 {
    if pos < 0 {
        pos + len as i32 + 1
    } else {
        pos.max(0)
    }
}

/// `uchar`
fn uchar(byte: u8) -> u8 {
    byte
}

/// `digit`
fn digit(byte: u8) -> bool {
    byte.is_ascii_digit()
}

/// `str_len`
fn string_len(ctx: NativeCallContext) -> NativeCallResult {
    ctx.push_integer(unsafe { ctx.arg(1).string()? }.len() as i32)?;
    Ok(1)
}

/// `str_sub`
fn string_sub(ctx: NativeCallContext) -> NativeCallResult {
    let thread = ctx.raw_thread();
    unsafe {
        let bytes = thread.check_string(1)?;
        let mut start = pos_relat(thread.check_integer(2)?, bytes.len());
        let mut end = pos_relat(thread.opt_integer(3, -1)?, bytes.len());

        if start < 1 {
            start = 1;
        }
        if end > bytes.len() as i32 {
            end = bytes.len() as i32;
        }

        if start <= end {
            thread.push_string(&bytes[(start - 1) as usize..end as usize])?;
        } else {
            thread.push_string("")?;
        }
    }

    Ok(1)
}

/// `str_reverse`
fn string_reverse(ctx: NativeCallContext) -> NativeCallResult {
    let thread = ctx.raw_thread();
    unsafe {
        let bytes = thread.check_string(1)?;
        let len = bytes.len();
        let mut buffer_storage = LuaStringBuilderStorage::uninit();
        let mut buffer = LuaStringBuilder::new(thread, &mut buffer_storage);
        let out = buffer.reserve(len)?;
        let out = slice::from_raw_parts_mut(out.as_ptr(), len);

        for (dst, src) in out.iter_mut().zip(bytes.iter().rev()) {
            *dst = *src;
        }

        buffer.finish_with_reserved(len)?;
    }
    Ok(1)
}

/// `str_lower`
fn string_lower(ctx: NativeCallContext) -> NativeCallResult {
    let thread = ctx.raw_thread();
    unsafe {
        let bytes = thread.check_string(1)?;
        let len = bytes.len();
        let mut buffer_storage = LuaStringBuilderStorage::uninit();
        let mut buffer = LuaStringBuilder::new(thread, &mut buffer_storage);
        let out = buffer.reserve(len)?;
        let out = slice::from_raw_parts_mut(out.as_ptr(), len);

        for (dst, src) in out.iter_mut().zip(bytes.iter()) {
            *dst = src.to_ascii_lowercase();
        }

        buffer.finish_with_reserved(len)?;
    }
    Ok(1)
}

/// `str_upper`
fn string_upper(ctx: NativeCallContext) -> NativeCallResult {
    let thread = ctx.raw_thread();
    unsafe {
        let bytes = thread.check_string(1)?;
        let len = bytes.len();
        let mut buffer_storage = LuaStringBuilderStorage::uninit();
        let mut buffer = LuaStringBuilder::new(thread, &mut buffer_storage);
        let out = buffer.reserve(len)?;
        let out = slice::from_raw_parts_mut(out.as_ptr(), len);

        for (dst, src) in out.iter_mut().zip(bytes.iter()) {
            *dst = src.to_ascii_uppercase();
        }

        buffer.finish_with_reserved(len)?;
    }
    Ok(1)
}

/// `str_rep`
fn string_rep(ctx: NativeCallContext) -> NativeCallResult {
    let thread = ctx.raw_thread();
    unsafe {
        let bytes = thread.check_string(1)?;
        let n = thread.check_integer(2)?;

        if n <= 0 {
            thread.push_string("")?;
            return Ok(1);
        }

        if bytes.len() > MAX_STRING_SIZE / n as usize {
            return crate::error!(thread, "resulting string too large").map_err(Into::into);
        }

        let total = bytes.len() * n as usize;
        if total <= LUA_BUFFER_SIZE {
            let mut buffer =
                MaybeUninit::<[MaybeUninit<u8>; LUA_BUFFER_SIZE]>::uninit().assume_init();
            let out = buffer.as_mut_ptr().cast::<u8>();
            core::ptr::copy_nonoverlapping(bytes.as_ptr(), out, bytes.len());

            let mut written = bytes.len();
            let mut left = total - bytes.len();
            let mut step = bytes.len();

            while step < left {
                core::ptr::copy_nonoverlapping(out, out.add(written), step);
                written += step;
                left -= step;
                step <<= 1;
            }

            core::ptr::copy_nonoverlapping(out, out.add(written), left);
            thread.push_string(slice::from_raw_parts(out, total))?;
            return Ok(1);
        }

        let mut buffer_storage = LuaStringBuilderStorage::uninit();
        let mut buffer = LuaStringBuilder::new(thread, &mut buffer_storage);
        let out = buffer.reserve(total)?;
        let start = out.as_ptr();

        core::ptr::copy_nonoverlapping(bytes.as_ptr(), start, bytes.len());

        let mut written = bytes.len();
        let mut left = total - bytes.len();
        let mut step = bytes.len();

        while step < left {
            core::ptr::copy_nonoverlapping(start, start.add(written), step);
            written += step;
            left -= step;
            step <<= 1;
        }

        core::ptr::copy_nonoverlapping(start, start.add(written), left);
        buffer.finish_with_reserved(total)?;
        Ok(1)
    }
}

/// `str_byte`
fn string_byte(ctx: NativeCallContext) -> NativeCallResult {
    let thread = ctx.raw_thread();
    unsafe {
        let bytes = thread.check_string(1)?;
        let mut start = pos_relat(thread.opt_integer(2, 1)?, bytes.len());
        let mut end = pos_relat(thread.opt_integer(3, start)?, bytes.len());

        if start <= 0 {
            start = 1;
        }
        if end as usize > bytes.len() {
            end = bytes.len() as i32;
        }
        if start > end {
            return Ok(0);
        }

        let count = end - start + 1;
        if start + count <= end {
            return crate::error!(thread, "string slice too long").map_err(Into::into);
        }

        thread.lua_check_stack(count, Some("string slice too long"))?;
        for index in 0..count as usize {
            thread.push_integer(uchar(bytes[start as usize + index - 1]) as i32)?;
        }
        Ok(count as usize)
    }
}

/// `str_char`
fn string_char(ctx: NativeCallContext) -> NativeCallResult {
    let thread = ctx.raw_thread();
    unsafe {
        let count = thread.get_top();
        let mut buffer_storage = LuaStringBuilderStorage::uninit();
        let mut buffer = LuaStringBuilder::new(thread, &mut buffer_storage);
        let out = buffer.reserve(count as usize)?;
        let out = slice::from_raw_parts_mut(out.as_ptr(), count as usize);

        for index in 1..=count {
            let ch = thread.check_integer(index)?;
            if uchar(ch as u8) as i32 != ch {
                return thread
                    .lua_arg_error(index, "invalid value")
                    .map_err(Into::into);
            }
            out[index as usize - 1] = ch as u8;
        }

        buffer.finish_with_reserved(count as usize)?;
        Ok(1)
    }
}

/// `str_split`
fn string_split(ctx: NativeCallContext) -> NativeCallResult {
    let thread = ctx.raw_thread();
    unsafe {
        let haystack = thread.check_string(1)?;
        let needle = thread.opt_string(2)?.unwrap_or(b",".as_bstr());

        let mut begin = 0usize;
        let end = haystack.len();
        let mut span_start = begin;
        let mut matches = 0i32;

        thread.create_table(0, 0)?;

        if needle.is_empty() {
            begin += 1;
        }

        let mut iter = begin;
        while iter <= end.saturating_sub(needle.len()) {
            if &haystack[iter..iter + needle.len()] == needle {
                matches += 1;
                thread.push_string(&haystack[span_start..iter])?;
                thread.raw_seti(-2, matches)?;

                span_start = iter + needle.len();
                if !needle.is_empty() {
                    iter += needle.len() - 1;
                }
            }

            iter += 1;
        }

        if !needle.is_empty() {
            thread.push_string(&haystack[span_start..])?;
            thread.raw_seti(-2, matches + 1)?;
        }

        Ok(1)
    }
}

/// `createmetatable`
unsafe fn create_metatable(thread: &Thread) -> NativeCallResult {
    unsafe {
        thread.create_table(0, 1)?;
        thread.push_string("")?;
        thread.push_value(-2)?;
        thread.set_metatable(-2)?;
        thread.pop(1);
        thread.push_value(-2)?;
        thread.set_field(-2, "__index")?;
        thread.pop(1);
    }
    Ok(1)
}

impl Thread {
    /// `luaopen_string`
    pub unsafe fn open_string(&self) -> NativeCallResult {
        unsafe { self.register(Some(super::LUA_STRLIB_NAME), &STRING_LIB[..])? };
        unsafe { create_metatable(self)? };
        Ok(1)
    }
}