use super::super::table::{Spec, arity_ok, lookup};
use super::super::{Args, Server, Session, resolved};
use super::{convert, nested, sha1};
use crate::frame;
use crate::proto::Limits;
use crate::reply::Out;
use crate::request::{Argv, Step};
use mlua::{Lua, MultiValue, Scope, Table, Value, ffi};
use std::cell::{Cell, RefCell};
use std::os::raw::c_int;
pub(super) struct Ctx<'a> {
server: &'a Server,
session: RefCell<&'a mut Session>,
ro: bool,
resp3: Cell<bool>,
}
impl<'a> Ctx<'a> {
pub(super) fn new(server: &'a Server, session: &'a mut Session, ro: bool) -> Ctx<'a> {
Ctx {
server,
session: RefCell::new(session),
ro,
resp3: Cell::new(false),
}
}
}
pub(super) fn table(lua: &Lua) -> mlua::Result<Table> {
let t = lua.create_table()?;
t.raw_set("LOG_DEBUG", 0)?;
t.raw_set("LOG_VERBOSE", 1)?;
t.raw_set("LOG_NOTICE", 2)?;
t.raw_set("LOG_WARNING", 3)?;
t.raw_set("REPL_NONE", 0)?;
t.raw_set("REPL_AOF", 1)?;
t.raw_set("REPL_SLAVE", 2)?;
t.raw_set("REPL_REPLICA", 2)?;
t.raw_set("REPL_ALL", 3)?;
t.raw_set("REDIS_VERSION", super::super::server::REPORTED_VERSION)?;
t.raw_set("REDIS_VERSION_NUM", version_num())?;
Ok(t)
}
fn version_num() -> i64 {
let mut parts = super::super::server::REPORTED_VERSION.split('.');
let mut at = |shift: u32| -> i64 {
let n: i64 = parts.next().and_then(|p| p.parse().ok()).unwrap_or(0);
(n & 0xff) << shift
};
at(16) | at(8) | at(0)
}
pub(super) fn statics(lua: &Lua, raw: &Table) -> mlua::Result<()> {
raw.raw_set(
"sha1hex",
lua.create_function(|_, value: Value| {
let bytes = match &value {
Value::String(s) => s.as_bytes().to_vec(),
Value::Integer(n) => n.to_string().into_bytes(),
Value::Number(d) => lua_number(*d).into_bytes(),
_ => Vec::new(),
};
Ok(String::from_utf8_lossy(&sha1::hex(&bytes)).into_owned())
})?,
)?;
raw.raw_set(
"log",
lua.create_function(|_, (level, message): (i64, String)| {
let word = match level {
0 => "debug",
1 => "verbose",
3 => "warning",
_ => "notice",
};
eprintln!("yodb script {word}: {message}");
Ok(())
})?,
)?;
raw.raw_set(
"known",
lua.create_function(|_, name: mlua::LuaString| Ok(lookup(&name.as_bytes()).is_some()))?,
)?;
raw.raw_set(
"arity_ok",
lua.create_function(|_, (name, count): (mlua::LuaString, usize)| {
Ok(lookup(&name.as_bytes()).is_some_and(|spec| arity_ok(spec, count)))
})?,
)?;
Ok(())
}
unsafe extern "C-unwind" fn forward(state: *mut mlua::lua_State) -> c_int {
unsafe {
let given = ffi::lua_gettop(state);
if ffi::lua_checkstack(state, 1) == 0 {
return 0;
}
ffi::lua_pushvalue(state, ffi::lua_upvalueindex(1));
ffi::lua_insert(state, 1);
ffi::lua_call(state, given, ffi::LUA_MULTRET);
ffi::lua_gettop(state)
}
}
pub(super) fn bridge(lua: &Lua) -> mlua::Result<mlua::Function> {
lua.create_function(|lua, body: mlua::Function| {
unsafe {
lua.exec_raw::<mlua::Function>(body, |state| {
ffi::lua_pushcclosure(state, forward, 1);
})
}
})
}
pub(super) fn lend<'scope, 'env: 'scope>(
lua: &Lua,
scope: &'scope Scope<'scope, 'env>,
ctx: &'env Ctx<'env>,
) -> mlua::Result<()> {
let raw: Table = lua.named_registry_value("yo_raw")?;
raw.raw_set(
"pcall",
scope.create_function(move |lua, args: MultiValue| command(lua, ctx, args))?,
)?;
raw.raw_set(
"setresp",
scope.create_function(move |_, three: bool| {
ctx.resp3.set(three);
Ok(())
})?,
)?;
Ok(())
}
fn command(lua: &Lua, ctx: &Ctx<'_>, args: MultiValue) -> mlua::Result<Value> {
let words = match flatten(lua, &args) {
Ok(w) => w,
Err(msg) => return failed(lua, msg),
};
const REFUSED: &[u8] = b"ERR This Redis command is not allowed from script";
let Some(spec) = lookup(&words[0]) else {
if shut(&words) {
return failed(lua, REFUSED);
}
return failed(lua, b"ERR Unknown Redis command called from script");
};
if !arity_ok(spec, words.len()) {
return failed(
lua,
b"ERR Wrong number of args calling Redis command from script",
);
}
if spec.flags.contains(&"noscript") || shut(&words) {
return failed(lua, REFUSED);
}
if ctx.ro && spec.flags.contains(&"write") {
return failed(
lua,
b"ERR Write commands are not allowed from read-only scripts.",
);
}
reply(lua, ctx, spec, &words)
}
fn shut(words: &[Vec<u8>]) -> bool {
const ALONE: [&[u8]; 39] = [
b"auth",
b"bgrewriteaof",
b"bgsave",
b"debug",
b"discard",
b"eval",
b"eval_ro",
b"evalsha",
b"evalsha_ro",
b"exec",
b"failover",
b"fcall",
b"fcall_ro",
b"hello",
b"monitor",
b"multi",
b"psubscribe",
b"psync",
b"punsubscribe",
b"quit",
b"replconf",
b"replicaof",
b"reset",
b"role",
b"save",
b"search.clusterinfo",
b"search.clusterrefresh",
b"search.clusterset",
b"shutdown",
b"slaveof",
b"ssubscribe",
b"subscribe",
b"sunsubscribe",
b"sync",
b"timeseries.clusterset",
b"timeseries.refreshcluster",
b"unsubscribe",
b"unwatch",
b"watch",
];
const CONTAINERS: [&[u8]; 9] = [
b"acl",
b"backup",
b"client",
b"config",
b"function",
b"hotkeys",
b"latency",
b"module",
b"script",
];
let name = words[0].to_ascii_lowercase();
if ALONE.contains(&name.as_slice()) {
return true;
}
if !CONTAINERS.contains(&name.as_slice()) {
return false;
}
!words
.get(1)
.is_some_and(|sub| sub.eq_ignore_ascii_case(b"help"))
}
fn reply(lua: &Lua, ctx: &Ctx<'_>, spec: &'static Spec, words: &[Vec<u8>]) -> mlua::Result<Value> {
let wire = encode(words);
let limits = Limits::default();
let mut argv = Argv::new();
match argv.decode(&wire, &limits) {
Ok(Step::Command { .. }) => {}
_ => return failed(lua, b"ERR Lua redis lib command arguments are too long"),
}
let mut scratch = Out::new(nested(ctx.resp3.get()));
{
let mut session = ctx.session.borrow_mut();
resolved(
ctx.server,
&mut session,
Some(spec),
Args::new(&argv, &wire),
&mut scratch,
);
}
match frame::decode(scratch.as_slice(), &limits) {
Ok(Some((f, _))) => convert::pull(lua, &f, ctx.resp3.get()),
_ => Ok(Value::Boolean(false)),
}
}
fn flatten(_lua: &Lua, args: &MultiValue) -> Result<Vec<Vec<u8>>, &'static [u8]> {
if args.is_empty() {
return Err(b"ERR Please specify at least one argument for this redis lib call");
}
let mut words = Vec::with_capacity(args.len());
for arg in args {
words.push(match arg {
Value::String(s) => s.as_bytes().to_vec(),
Value::Integer(n) => n.to_string().into_bytes(),
Value::Number(d) => lua_number(*d).into_bytes(),
_ => return Err(b"ERR Lua redis lib command arguments must be strings or integers"),
});
}
Ok(words)
}
fn lua_number(d: f64) -> String {
if d == d.trunc() && d.abs() < 1e15 {
return format!("{}", d as i64);
}
format!("{d:.14}")
.trim_end_matches('0')
.trim_end_matches('.')
.to_string()
}
fn encode(words: &[Vec<u8>]) -> Vec<u8> {
let mut wire = format!("*{}\r\n", words.len()).into_bytes();
for w in words {
wire.extend_from_slice(format!("${}\r\n", w.len()).as_bytes());
wire.extend_from_slice(w);
wire.extend_from_slice(b"\r\n");
}
wire
}
fn failed(lua: &Lua, msg: &[u8]) -> mlua::Result<Value> {
let t = lua.create_table()?;
t.raw_set("err", lua.create_string(msg)?)?;
t.raw_set("ignore_error_stats_update", true)?;
Ok(Value::Table(t))
}