use std::str;
use gxhash::HashMap;
use mlua::{Lua, MultiValue, Value, Variadic};
use super::i_lua_allocator::ILuaAllocator;
pub type StackValue = Value;
#[derive(Default)]
pub struct LuaInterp {
pub stack: Vec<StackValue>,
pub refs: HashMap<i32, mlua::RegistryKey>,
pub next_ref_id: i32,
pub deadline_monotonic_millis: Option<i64>,
}
pub struct LuaStateWrapper {
lua: Lua,
allocator: Option<Box<dyn ILuaAllocator>>,
}
impl Default for LuaStateWrapper {
fn default() -> Self {
Self::new()
}
}
impl LuaStateWrapper {
pub fn new() -> Self {
let lua = Lua::new();
lua.set_app_data(LuaInterp::default());
Self {
lua,
allocator: None,
}
}
pub fn view(lua: &Lua) -> Self {
if lua.app_data_ref::<LuaInterp>().is_none() {
lua.set_app_data(LuaInterp::default());
}
Self {
lua: lua.clone(),
allocator: None,
}
}
pub fn lua(&self) -> &Lua {
&self.lua
}
pub fn interp(&self) -> mlua::AppDataRef<'_, LuaInterp> {
self
.lua
.app_data_ref::<LuaInterp>()
.expect("app data 已在构造时装载")
}
pub fn interp_mut(&self) -> mlua::AppDataRefMut<'_, LuaInterp> {
self
.lua
.app_data_mut::<LuaInterp>()
.expect("app data 已在构造时装载")
}
pub fn expect_lua_stack_empty(&self) -> bool {
self.interp().stack.is_empty()
}
pub fn try_ensure_minimum_stack_capacity(&mut self, min_capacity: usize) -> bool {
self.interp_mut().stack.reserve(min_capacity);
true
}
pub fn call_from_lua_entered(&mut self, nargs: usize) -> Result<(), mlua::Error> {
self.known_call_from_lua_entered(nargs)
}
pub fn known_call_from_lua_entered(&mut self, nargs: usize) -> Result<(), mlua::Error> {
let total = nargs + 1;
let args: Vec<StackValue> = {
let mut interp = self.interp_mut();
args_from_stack(&mut interp.stack, total)
};
let Some((first, rest)) = args.split_first() else {
return Err(mlua::Error::RuntimeError(
"attempt to call a non-function object".into(),
));
};
let function = as_function(first)?;
let results: MultiValue = function.call(Variadic::from_iter(rest.iter().cloned()))?;
let mut interp = self.interp_mut();
interp.stack.extend(results.into_vec());
Ok(())
}
pub fn type_name(&self, idx: i32) -> Option<&'static str> {
self.peek(idx).as_ref().map(value_type_name)
}
pub fn try_push_buffer(&mut self, buffer: &[u8]) -> bool {
let Ok(s) = self.lua.create_string(buffer) else {
return false;
};
self.interp_mut().stack.push(Value::String(s));
true
}
pub fn push_nil(&mut self) {
self.interp_mut().stack.push(Value::Nil);
}
pub fn push_number(&mut self, number: f64) {
self.interp_mut().stack.push(Value::Number(number));
}
pub fn push_integer(&mut self, integer: i64) {
self.interp_mut().stack.push(Value::Integer(integer));
}
pub fn push_boolean(&mut self, boolean: bool) {
self.interp_mut().stack.push(Value::Boolean(boolean));
}
pub fn pop(&mut self, count: usize) {
let mut interp = self.interp_mut();
let keep = interp.stack.len().saturating_sub(count);
interp.stack.truncate(keep);
}
pub fn remove(&mut self, idx: i32) {
let Some(abs) = self.abs_index(idx) else {
return;
};
self.interp_mut().stack.remove(abs);
}
pub fn pcall(&mut self, nargs: usize) -> Result<(), mlua::Error> {
self.pcall_n(nargs, usize::MAX).map(|_| ())
}
pub fn pcall_n(&mut self, nargs: usize, nresults: usize) -> Result<usize, mlua::Error> {
let total = nargs + 1;
let args: Vec<StackValue> = {
let mut interp = self.interp_mut();
args_from_stack(&mut interp.stack, total)
};
let Some((first, rest)) = args.split_first() else {
return Err(mlua::Error::RuntimeError(
"attempt to call a non-function object".into(),
));
};
let Ok(function) = as_function(first) else {
return Err(mlua::Error::RuntimeError(
"attempt to call a non-function object".into(),
));
};
let called: Result<MultiValue, mlua::Error> =
function.call(Variadic::from_iter(rest.iter().cloned()));
let mut interp = self.interp_mut();
match called {
Ok(results) => {
let mut values = results.into_vec();
if nresults != usize::MAX {
values.truncate(nresults);
values.resize(nresults, Value::Nil);
}
let count = values.len();
interp.stack.extend(values);
Ok(count)
}
Err(error) => {
let message = error_message(&error);
let Ok(s) = self.lua.create_string(message) else {
return Err(error);
};
interp.stack.push(Value::String(s));
Err(error)
}
}
}
pub fn raw_set_integer(&mut self, table_idx: i32, key: i64, value: StackValue) -> bool {
let Some(Value::Table(table)) = self.peek(table_idx) else {
return false;
};
table.raw_set(key, value).is_ok()
}
pub fn raw_set(&mut self, table_idx: i32) -> bool {
let Some(Value::Table(table)) = self.peek(table_idx) else {
return false;
};
let Some(value) = self.interp_mut().stack.pop() else {
return false;
};
let Some(key) = self.interp_mut().stack.pop() else {
return false;
};
table.raw_set(key, value).is_ok()
}
pub fn raw_get_integer(&mut self, table_idx: i32, key: i64) -> bool {
let Some(Value::Table(table)) = self.peek(table_idx) else {
return false;
};
match table.raw_get(key) {
Ok(value) => {
self.interp_mut().stack.push(value);
true
}
Err(_) => false,
}
}
pub fn raw_get_top(&mut self, table_idx: i32) -> bool {
let Some(Value::Table(table)) = self.peek(table_idx) else {
return false;
};
let Some(key) = self.interp_mut().stack.pop() else {
return false;
};
match table.raw_get(key) {
Ok(value) => {
self.interp_mut().stack.push(value);
true
}
Err(_) => false,
}
}
pub fn try_ref(&mut self) -> Option<i32> {
let value = self.interp_mut().stack.pop()?;
let key = self.lua.create_registry_value(value).ok()?;
let id = {
let mut interp = self.interp_mut();
interp.next_ref_id += 1;
interp.next_ref_id
};
self.interp_mut().refs.insert(id, key);
Some(id)
}
pub fn unref(&mut self, ref_id: i32) {
if let Some(key) = self.interp_mut().refs.remove(&ref_id) {
self.lua.remove_registry_value(key).ok();
}
}
pub fn ref_value(&self, ref_id: i32) -> Option<StackValue> {
let interp = self.interp();
let key = interp.refs.get(&ref_id)?;
self.lua.registry_value::<Value>(key).ok()
}
pub fn push_ref(&mut self, ref_id: i32) -> bool {
match self.ref_value(ref_id) {
Some(value) => {
self.interp_mut().stack.push(value);
true
}
None => false,
}
}
pub fn try_create_table(&mut self, narr: usize, nrec: usize) -> bool {
match self.lua.create_table_with_capacity(narr, nrec) {
Ok(table) => {
self.interp_mut().stack.push(Value::Table(table));
true
}
Err(_) => false,
}
}
pub fn get_global(&mut self, name: &[u8]) -> bool {
let Ok(name) = str::from_utf8(name) else {
return false;
};
match self.lua.globals().get(name) {
Ok(value) => {
self.interp_mut().stack.push(value);
true
}
Err(_) => false,
}
}
pub fn try_set_global(&mut self, name: &[u8]) -> bool {
let (Some(value), Ok(name)) = (self.interp_mut().stack.pop(), str::from_utf8(name)) else {
return false;
};
self.lua.globals().set(name, value).is_ok()
}
pub fn register_function<F, A, R>(&mut self, name: &[u8], function: F) -> bool
where
F: Fn(&Lua, A) -> Result<R, mlua::Error> + 'static,
A: mlua::FromLuaMulti,
R: mlua::IntoLuaMulti,
{
let Ok(name) = str::from_utf8(name) else {
return false;
};
let Ok(f) = self.lua.create_function(function) else {
return false;
};
self.lua.globals().set(name, f).is_ok()
}
pub fn load_buffer(&mut self, buffer: &[u8], chunk_name: &str) -> Result<(), mlua::Error> {
let Ok(source) = str::from_utf8(buffer) else {
return Err(mlua::Error::RuntimeError("non-utf8 chunk".into()));
};
let chunk = self.lua.load(source).set_name(chunk_name);
let function = chunk.into_function()?;
self.interp_mut().stack.push(Value::Function(function));
Ok(())
}
pub fn load_string(&mut self, source: &str) -> Result<(), mlua::Error> {
let function = self.lua.load(source).into_function()?;
self.interp_mut().stack.push(Value::Function(function));
Ok(())
}
pub fn try_number_to_string(&mut self) -> bool {
self.try_number_to_string_at(-1)
}
pub fn try_number_to_string_at(&mut self, idx: i32) -> bool {
let number = match self.peek(idx) {
Some(Value::Number(n)) => n,
Some(Value::Integer(i)) => i as f64,
_ => return false,
};
let Ok(s) = self.lua.create_string(format_number_text(number)) else {
return false;
};
let Some(abs) = self.abs_index(idx) else {
return false;
};
let mut interp = self.interp_mut();
interp.stack[abs] = Value::String(s);
true
}
pub fn known_string_to_buffer(&self, idx: i32) -> Option<Vec<u8>> {
match self.peek(idx)? {
Value::String(s) => Some(s.as_bytes().to_vec()),
_ => None,
}
}
pub fn check_number(&self, idx: i32) -> Option<f64> {
match self.peek(idx)? {
Value::Number(n) => Some(n),
Value::Integer(i) => Some(i as f64),
Value::String(s) => str::from_utf8(&s.as_bytes()).ok()?.parse().ok(),
_ => None,
}
}
pub fn to_boolean(&self, idx: i32) -> bool {
!matches!(
self.peek(idx),
None | Some(Value::Nil) | Some(Value::Boolean(false))
)
}
pub fn raw_len(&self, idx: i32) -> i64 {
match self.peek(idx) {
Some(Value::Table(t)) => t.raw_len() as i64,
Some(Value::String(s)) => s.as_bytes().len() as i64,
_ => 0,
}
}
pub fn push_c_function(&mut self, function: mlua::Function) {
self.interp_mut().stack.push(Value::Function(function));
}
pub fn push_constant_string(&mut self, constant: &[u8]) -> bool {
self.try_push_buffer(constant)
}
pub fn lua_next(&mut self) -> bool {
let key = self.interp_mut().stack.pop();
let Some(key) = key else {
return false;
};
let Some(Value::Table(table)) = self.peek(-1) else {
self.interp_mut().stack.push(key);
return false;
};
let mut next_pair: Option<(StackValue, StackValue)> = None;
if matches!(key, Value::Nil) {
next_pair = table
.clone()
.pairs::<StackValue, StackValue>()
.find_map(Result::ok);
} else {
let mut passed = false;
for pair in table.clone().pairs::<StackValue, StackValue>() {
let Ok((k, v)) = pair else { continue };
if passed {
next_pair = Some((k, v));
break;
}
passed = k == key;
}
}
match next_pair {
Some((next_key, value)) => {
let mut interp = self.interp_mut();
interp.stack.push(next_key);
interp.stack.push(value);
true
}
None => false,
}
}
pub fn push_value(&mut self, idx: i32) {
if let Some(value) = self.peek(idx) {
self.interp_mut().stack.push(value);
}
}
pub fn rotate(&mut self, idx: i32, n: i32) {
let start = self.abs_index(idx);
let Some(start) = start else { return };
let mut interp = self.interp_mut();
if n == 0 || interp.stack.len() <= start {
return;
}
let len = interp.stack.len() - start;
let n = ((n % len as i32) + len as i32) as usize % len;
if n == 0 {
return;
}
interp.stack[start..].rotate_right(n);
}
pub fn try_set_hook(&mut self, deadline_monotonic_millis: Option<i64>) {
self.interp_mut().deadline_monotonic_millis = deadline_monotonic_millis;
self.lua.set_interrupt(move |lua: &Lua| {
let expired = lua
.app_data_ref::<LuaInterp>()
.and_then(|interp| interp.deadline_monotonic_millis)
.is_some_and(|deadline| now_monotonic_millis() >= deadline);
if expired {
Err(mlua::Error::RuntimeError(
"ERR Lua script exceeded configured timeout".into(),
))
} else {
Ok(mlua::VmState::Continue)
}
});
}
pub fn deadline(&self) -> Option<i64> {
self.interp().deadline_monotonic_millis
}
pub fn assert_lua_stack_index_in_bounds(&self, idx: i32) -> bool {
self.abs_index(idx).is_some()
}
pub fn assert_lua_stack_expected(&self, idx: i32, expected: &str) -> bool {
self.type_name(idx) == Some(expected)
}
pub fn assert_lua_stack_not_full(&self) -> bool {
self.interp().stack.len() < i32::MAX as usize
}
pub fn assert_lua_stack_not_empty(&self) -> bool {
!self.interp().stack.is_empty()
}
pub fn lua_at_panic(&mut self) -> i32 {
0
}
pub fn lua_allocate_bytes(&self) -> usize {
self.lua.used_memory()
}
pub fn clear_stack(&mut self) {
self.interp_mut().stack.clear();
}
pub fn update_stack_top(&mut self, new_top: usize) {
self.interp_mut().stack.resize(new_top, Value::Nil);
}
pub fn get_top(&self) -> usize {
self.interp().stack.len()
}
pub fn enter_infallible_allocation_region(&mut self) {
if let Some(allocator) = &mut self.allocator {
allocator.enter_infallible_allocation_region();
}
}
pub fn try_exit_infallible_allocation_region(&mut self) -> bool {
self
.allocator
.as_mut()
.is_none_or(|a| a.try_exit_infallible_allocation_region())
}
pub fn set_allocator(&mut self, allocator: Box<dyn ILuaAllocator>) {
self.allocator = Some(allocator);
}
fn abs_index(&self, idx: i32) -> Option<usize> {
if idx > 0 {
usize::try_from(idx).ok().map(|i| i - 1)
} else {
usize::try_from(-idx)
.ok()
.and_then(|i| self.interp().stack.len().checked_sub(i))
}
}
fn peek(&self, idx: i32) -> Option<StackValue> {
let i = self.abs_index(idx)?;
self.interp().stack.get(i).cloned()
}
}
fn args_from_stack(stack: &mut Vec<StackValue>, total: usize) -> Vec<StackValue> {
stack.split_off(stack.len().saturating_sub(total))
}
pub fn error_message(error: &mlua::Error) -> String {
match error {
mlua::Error::RuntimeError(msg) => msg.clone(),
other => {
let text = other.to_string();
text
.strip_prefix("runtime error: ")
.map_or_else(|| text.clone(), str::to_owned)
}
}
}
pub fn now_monotonic_millis() -> i64 {
coarsetime::Clock::now_since_epoch().as_millis() as i64
}
fn format_number_text(number: f64) -> String {
if number == number.trunc() && number.abs() < 1e15 {
format!("{}", number as i64)
} else {
format!("{number}")
}
}
fn as_function(value: &StackValue) -> Result<mlua::Function, mlua::Error> {
match value {
Value::Function(function) => Ok(function.clone()),
other => Err(mlua::Error::RuntimeError(format!(
"attempt to call a {} value",
value_type_name(other)
))),
}
}
pub fn value_type_name(value: &StackValue) -> &'static str {
match value {
Value::Nil => "nil",
Value::Boolean(_) => "boolean",
Value::Integer(_) | Value::Number(_) => "number",
Value::String(_) => "string",
Value::Table(_) => "table",
Value::Function(_) => "function",
Value::LightUserData(_) | Value::UserData(_) => "userdata",
Value::Thread(_) => "thread",
_ => "userdata",
}
}
#[cfg(test)]
mod tests {
use super::{LuaStateWrapper, Value};
#[test]
fn push_pop_and_types() {
let mut state = LuaStateWrapper::new();
state.push_integer(7);
state.push_number(3.5);
state.push_boolean(true);
state.push_nil();
state.try_push_buffer(b"hello");
assert_eq!(state.get_top(), 5);
assert_eq!(state.type_name(-1), Some("string"));
assert_eq!(state.type_name(-2), Some("nil"));
assert_eq!(state.check_number(-5), Some(7.0));
assert!(state.to_boolean(-3));
assert!(!state.to_boolean(-2));
assert_eq!(state.raw_len(-1), 5);
state.pop(5);
assert!(state.expect_lua_stack_empty());
}
#[test]
fn load_and_pcall() {
let mut state = LuaStateWrapper::new();
state.load_string("return 1, 'x'").unwrap();
state.pcall(0).unwrap();
assert_eq!(state.get_top(), 2);
assert_eq!(state.check_number(-2), Some(1.0));
assert_eq!(state.type_name(-1), Some("string"));
state.clear_stack();
state.load_string("error('boom')").unwrap();
assert!(state.pcall(0).is_err());
assert_eq!(state.get_top(), 1);
assert!(String::from_utf8_lossy(&state.known_string_to_buffer(-1).unwrap()).contains("boom"));
}
#[test]
fn table_ops_and_globals() {
let mut state = LuaStateWrapper::new();
assert!(state.try_create_table(4, 0));
state.push_integer(42);
assert!(state.raw_set_integer(-2, 1, Value::Integer(42)));
assert!(state.raw_get_integer(-2, 1));
assert_eq!(state.check_number(-1), Some(42.0));
state.pop(1);
assert_eq!(state.raw_len(-2), 1);
state.clear_stack();
state.try_push_buffer(b"answer");
assert!(state.try_set_global(b"ANSWER"));
assert!(state.get_global(b"ANSWER"));
assert_eq!(state.known_string_to_buffer(-1).unwrap(), b"answer");
}
#[test]
fn rotate_and_stack_top() {
let mut state = LuaStateWrapper::new();
for i in 1..=3 {
state.push_integer(i);
}
state.rotate(1, 1);
assert_eq!(state.check_number(1), Some(3.0));
assert_eq!(state.check_number(2), Some(1.0));
state.update_stack_top(5);
assert_eq!(state.get_top(), 5);
state.update_stack_top(1);
assert_eq!(state.get_top(), 1);
}
#[test]
fn view_shares_interp_and_refs() {
let mut state = LuaStateWrapper::new();
state.push_integer(1);
assert!(state.try_ref().is_some());
let mut view = LuaStateWrapper::view(state.lua());
assert_eq!(view.get_top(), 0);
assert!(view.push_ref(1));
assert_eq!(view.check_number(-1), Some(1.0));
view.push_integer(2);
assert_eq!(state.get_top(), 2);
}
#[test]
fn pcall_n_pads_and_truncates() {
let mut state = LuaStateWrapper::new();
state.load_string("return 1, 2").unwrap();
state.pcall_n(0, 3).unwrap();
assert_eq!(state.get_top(), 3);
assert!(state.ref_value(0).is_none());
state.clear_stack();
state.load_string("return 1, 2").unwrap();
state.pcall_n(0, 1).unwrap();
assert_eq!(state.get_top(), 1);
assert_eq!(state.check_number(-1), Some(1.0));
}
#[test]
fn next_and_raw_get_top_semantics() {
let mut state = LuaStateWrapper::new();
assert!(state.try_create_table(0, 2));
state.push_constant_string(b"k");
state.push_integer(7);
assert!(state.raw_set(-3));
assert_eq!(state.get_top(), 1);
state.push_nil();
let mut seen = 0;
while state.lua_next() {
seen += 1;
state.pop(1);
}
assert_eq!(seen, 1);
assert_eq!(state.get_top(), 1);
state.push_constant_string(b"k");
assert!(state.raw_get_top(-2));
assert_eq!(state.get_top(), 2);
assert_eq!(state.check_number(-1), Some(7.0));
}
}