#![forbid(unsafe_code)]
#![warn(missing_docs)]
use luna_core::runtime::value::Value;
use luna_core::vm::exec::Vm;
use std::cell::Cell;
use std::rc::Rc;
mod dispatch;
mod host;
mod marshal;
mod resp;
mod shebang;
pub mod sha1;
mod cmsgpack;
mod cjson;
pub use luna_core::version::LuaVersion;
pub(crate) use dispatch::{DispatchHandle, DispatchSlot, DISPATCH_KEY};
const N_DIALECTS: usize = 6;
const DEFAULT_INSTR_BUDGET: i64 = 200_000_000;
fn dialect_slot(v: LuaVersion) -> usize {
match v {
LuaVersion::Lua51 => 0,
LuaVersion::Lua52 => 1,
LuaVersion::Lua53 => 2,
LuaVersion::Lua54 => 3,
LuaVersion::MacroLua => 4,
LuaVersion::Lua55 => 5,
}
}
pub type Reply = Vec<u8>;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum FlushMode {
Sync,
Async,
}
pub type ScriptSha1 = [u8; 20];
pub struct Bridge {
vms: [Option<Vm>; N_DIALECTS],
dispatch: DispatchHandle,
read_only: Rc<Cell<bool>>,
instr_budget: i64,
allow: [bool; N_DIALECTS],
script_cache: std::collections::HashMap<ScriptSha1, Vec<u8>>,
}
impl Bridge {
pub fn new<F>(dispatch: F) -> Self
where
F: Fn(&[&[u8]], bool) -> Vec<u8> + 'static,
{
Self {
vms: [const { None }; N_DIALECTS],
dispatch: Rc::new(dispatch),
read_only: Rc::new(Cell::new(false)),
allow: [true; N_DIALECTS],
script_cache: std::collections::HashMap::new(),
instr_budget: DEFAULT_INSTR_BUDGET,
}
}
pub fn set_instr_budget(&mut self, n: i64) {
self.instr_budget = n;
}
#[must_use]
pub fn with_no_dispatch() -> Self {
Self::new(|_argv: &[&[u8]], _ro: bool| {
b"-ERR redis.call: no host dispatch wired\r\n".to_vec()
})
}
pub fn set_allowed_dialects(&mut self, versions: &[LuaVersion]) {
if versions.is_empty() {
self.allow = [true; N_DIALECTS];
return;
}
self.allow = [false; N_DIALECTS];
for v in versions {
self.allow[dialect_slot(*v)] = true;
}
}
pub fn eval(&mut self, script: &[u8], keys: &[&[u8]], args: &[&[u8]]) -> Reply {
let (sh, body) = match shebang::parse(script) {
Ok(t) => t,
Err(e) => return resp::err(format!("{e}").as_bytes()),
};
if !self.allow[dialect_slot(sh.version)] {
return resp::err(
format!(
"dialect {} disabled by [lua] allow_dialects",
version_tag(sh.version)
)
.as_bytes(),
);
}
let src = match std::str::from_utf8(body) {
Ok(s) => s,
Err(_) => return resp::err(b"script body is not valid UTF-8"),
};
let digest = sha1::sha1(script);
self.script_cache.entry(digest).or_insert_with(|| script.to_vec());
let vm = self.vm_for(sh.version);
host::bind_keys_argv(vm, keys, args);
match vm.eval(src) {
Ok(results) => {
let first = results.first().copied().unwrap_or(Value::Nil);
marshal::value(vm, first)
}
Err(e) => resp::err(format_lua_error(&e).as_bytes()),
}
}
pub fn eval_ro(&mut self, script: &[u8], keys: &[&[u8]], args: &[&[u8]]) -> Reply {
self.read_only.set(true);
let r = self.eval(script, keys, args);
self.read_only.set(false);
r
}
pub fn evalsha_ro(&mut self, sha1: ScriptSha1, keys: &[&[u8]], args: &[&[u8]]) -> Reply {
self.read_only.set(true);
let r = self.evalsha(sha1, keys, args);
self.read_only.set(false);
r
}
pub fn evalsha(&mut self, sha1: ScriptSha1, keys: &[&[u8]], args: &[&[u8]]) -> Reply {
let Some(script) = self.script_cache.get(&sha1).cloned() else {
return resp::err(b"NOSCRIPT No matching script. Please use EVAL.");
};
self.eval(&script, keys, args)
}
pub fn script_load(&mut self, script: &[u8]) -> ScriptSha1 {
let digest = sha1::sha1(script);
self.script_cache.insert(digest, script.to_vec());
digest
}
#[must_use]
pub fn script_exists(&self, sha1s: &[ScriptSha1]) -> Vec<bool> {
sha1s.iter().map(|s| self.script_cache.contains_key(s)).collect()
}
pub fn script_flush(&mut self, _mode: FlushMode) {
for slot in &mut self.vms {
*slot = None;
}
self.script_cache.clear();
}
#[cfg(test)]
fn vm_count(&self) -> usize {
self.vms.iter().filter(|s| s.is_some()).count()
}
fn vm_for(&mut self, version: LuaVersion) -> &mut Vm {
let slot = &mut self.vms[dialect_slot(version)];
if slot.is_none() {
let mut builder = Vm::sandbox(version)
.open_base()
.open_math()
.open_string()
.open_table();
if self.instr_budget > 0 {
builder = builder.with_instr_budget(self.instr_budget);
}
let mut vm = builder.build();
host::install_redis_table(&mut vm);
cmsgpack::install_cmsgpack(&mut vm);
cjson::install_cjson(&mut vm);
let _ = vm.set_userdata(
DISPATCH_KEY,
DispatchSlot {
f: Rc::clone(&self.dispatch),
read_only: Rc::clone(&self.read_only),
},
);
*slot = Some(vm);
}
slot.as_mut().expect("just-inserted Vm")
}
}
impl Default for Bridge {
fn default() -> Self {
Self::with_no_dispatch()
}
}
fn format_lua_error(e: &luna_core::vm::error::LuaError) -> String {
format!("{e}")
}
fn version_tag(v: LuaVersion) -> &'static str {
match v {
LuaVersion::Lua51 => "5.1",
LuaVersion::Lua52 => "5.2",
LuaVersion::Lua53 => "5.3",
LuaVersion::Lua54 => "5.4",
LuaVersion::MacroLua => "macro",
LuaVersion::Lua55 => "5.5",
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn eval_reuses_vm_across_calls() {
let mut b = Bridge::with_no_dispatch();
assert_eq!(b.eval(b"return 1", &[], &[]), b":1\r\n");
assert_eq!(b.eval(b"return 2", &[], &[]), b":2\r\n");
assert_eq!(b.vm_count(), 1);
}
#[test]
fn script_flush_drops_vm_pool() {
let mut b = Bridge::with_no_dispatch();
let _ = b.eval(b"return 1", &[], &[]);
assert_eq!(b.vm_count(), 1);
b.script_flush(FlushMode::Sync);
assert_eq!(b.vm_count(), 0);
}
}