use crate::cmd::wrong_args;
use crate::state::{Ctx, RuntimeState, ShardCtx};
use kevy_lua_host::LuaHost;
use kevy_resp::{Argv, ArgvView, encode_error};
use kevy_store::Store;
use std::sync::Arc;
fn make_lua_host(state: Arc<RuntimeState>, my_shard: usize) -> LuaHost<Store> {
let cfg = state.config();
let mut host = LuaHost::<Store>::new(move |store, argv, read_only| {
lua_redis_call(&state, my_shard, store, argv, read_only)
});
apply_lua_config(&mut host, &cfg);
host
}
fn read_only_violation(argv: &[&[u8]], read_only: bool) -> Option<Vec<u8>> {
if !read_only {
return None;
}
let cmd = argv.first()?;
let mut buf = [0u8; 32];
if crate::cmd::is_write_verb(crate::cmd::upper_verb(cmd, &mut buf)) {
return Some(b"-READONLY can't write against a read-only script\r\n".to_vec());
}
None
}
fn cross_shard_violation(
state: &Arc<RuntimeState>,
my_shard: usize,
argv: &[&[u8]],
) -> Option<Vec<u8>> {
let cfg = state.config();
let nshards = cfg.server.threads;
if nshards > 1
&& let Some(target_key) = argv.get(1)
{
let target_shard = kevy_rt::shard_of_key(target_key, nshards, cfg.cluster.enabled);
if target_shard != my_shard {
return Some(b"-CROSSSLOT Lua redis.call target key is on a different shard than the EVAL. Use {hashtag} to colocate keys, or run kevy --threads 1.\r\n".to_vec());
}
}
None
}
fn lua_redis_call(
state: &Arc<RuntimeState>,
my_shard: usize,
store: &mut Store,
argv: &[&[u8]],
read_only: bool,
) -> Vec<u8> {
if let Some(err) = read_only_violation(argv, read_only) {
return err;
}
if let Some(err) = cross_shard_violation(state, my_shard, argv) {
return err;
}
let mut a = Argv::default();
for slice in argv {
a.push(slice);
}
let mut out = Vec::new();
let shard = ShardCtx::default();
shard.set_shard_id(my_shard);
let ctx = Ctx { state, shard: &shard };
crate::dispatch::dispatch_into(&ctx, store, &a, &mut out);
kevy_rt::propagation::discard_override();
bridge_lua_wake_keys(argv, &out);
out
}
fn bridge_lua_wake_keys(argv: &[&[u8]], out: &[u8]) {
if !out.is_empty()
&& out[0] != b'-'
&& let Some(verb) = argv.first()
{
let mut buf = [0u8; 32];
let upper = crate::cmd::upper_verb(verb, &mut buf);
if crate::cmd_block::wake_idx_for_verb(upper).is_some()
&& let Some(key) = argv.get(1)
{
kevy_rt::push_lua_wake_key(key);
}
}
}
fn apply_lua_config(host: &mut LuaHost<Store>, cfg: &kevy_config::Config) {
if cfg.lua.time_limit_ms > 0 {
let budget = (cfg.lua.time_limit_ms as i64).saturating_mul(40_000);
host.set_instr_budget(budget);
} else {
host.set_instr_budget(0); }
if !cfg.lua.allow_dialects.is_empty() {
let versions: Vec<kevy_lua::LuaVersion> = cfg
.lua
.allow_dialects
.iter()
.filter_map(|s| match s.as_str() {
"5.1" | "51" => Some(kevy_lua::LuaVersion::Lua51),
"5.2" | "52" => Some(kevy_lua::LuaVersion::Lua52),
"5.3" | "53" => Some(kevy_lua::LuaVersion::Lua53),
"5.4" | "54" => Some(kevy_lua::LuaVersion::Lua54),
"5.5" | "55" => Some(kevy_lua::LuaVersion::Lua55),
_ => None,
})
.collect();
if !versions.is_empty() {
host.set_allowed_dialects(&versions);
}
}
}
fn with_host<R>(ctx: &Ctx<'_>, f: impl FnOnce(&mut LuaHost<Store>) -> R) -> Option<R> {
kevy_lua_host::with_thread_host(
|| make_lua_host(Arc::clone(ctx.state), ctx.shard.shard_id()),
f,
)
}
fn emit_reentry_err(out: &mut Vec<u8>) {
encode_error(out, "ERR EVAL inside EVAL is not supported in v1.27");
}
pub(crate) fn dispatch_lua<A: ArgvView + ?Sized>(
ctx: &Ctx<'_>,
cmd: &[u8],
store: &mut Store,
args: &A,
out: &mut Vec<u8>,
) -> bool {
match cmd {
b"EVAL" => {
cmd_eval(ctx, store, args, out, false);
true
}
b"EVAL_RO" => {
cmd_eval(ctx, store, args, out, true);
true
}
b"EVALSHA" => {
cmd_evalsha(ctx, store, args, out, false);
true
}
b"EVALSHA_RO" => {
cmd_evalsha(ctx, store, args, out, true);
true
}
b"SCRIPT" => {
cmd_script(ctx, args, out);
true
}
_ => false,
}
}
fn cmd_eval<A: ArgvView + ?Sized>(
ctx: &Ctx<'_>,
store: &mut Store,
args: &A,
out: &mut Vec<u8>,
read_only: bool,
) {
if args.len() < 3 {
wrong_args(out, if read_only { "eval_ro" } else { "eval" });
return;
}
let script: &[u8] = args.get(1).unwrap_or(b"");
let Some((keys, argv)) = parse_eval_keys_argv(ctx, args, out) else {
return;
};
let sha = kevy_lua::sha1::sha1(script);
scripts(ctx).insert(sha, script.to_vec());
let reply = with_host(ctx, |h| {
if read_only {
h.eval_ro(store, script, &keys, &argv)
} else {
h.eval(store, script, &keys, &argv)
}
});
match reply {
Some(bytes) => out.extend_from_slice(&bytes),
None => emit_reentry_err(out),
}
}
fn scripts<'a>(
ctx: &'a Ctx<'_>,
) -> std::sync::MutexGuard<'a, std::collections::HashMap<[u8; 20], Vec<u8>>> {
ctx.state.catalogs.scripts.lock().unwrap_or_else(std::sync::PoisonError::into_inner)
}
fn script_source(ctx: &Ctx<'_>, sha: &[u8; 20]) -> Option<Vec<u8>> {
scripts(ctx).get(sha).cloned()
}
fn cmd_evalsha<A: ArgvView + ?Sized>(
ctx: &Ctx<'_>,
store: &mut Store,
args: &A,
out: &mut Vec<u8>,
read_only: bool,
) {
if args.len() < 3 {
wrong_args(out, if read_only { "evalsha_ro" } else { "evalsha" });
return;
}
let sha_hex: &[u8] = args.get(1).unwrap_or(b"");
let sha = match kevy_lua::sha1::parse_hex(sha_hex) {
Some(s) => s,
None => {
encode_error(out, "NOSCRIPT No matching script. Please use EVAL.");
return;
}
};
let Some((keys, argv)) = parse_eval_keys_argv(ctx, args, out) else {
return;
};
let Some(source) = script_source(ctx, &sha) else {
encode_error(out, "NOSCRIPT No matching script. Please use EVAL.");
return;
};
let reply = with_host(ctx, |h| {
if read_only {
h.eval_ro(store, &source, &keys, &argv)
} else {
h.eval(store, &source, &keys, &argv)
}
});
match reply {
Some(bytes) => out.extend_from_slice(&bytes),
None => emit_reentry_err(out),
}
}
fn cmd_script<A: ArgvView + ?Sized>(ctx: &Ctx<'_>, args: &A, out: &mut Vec<u8>) {
if args.len() < 2 {
wrong_args(out, "script");
return;
}
let sub_upper: Vec<u8> =
args.get(1).unwrap_or(b"").iter().map(|b| b.to_ascii_uppercase()).collect();
match sub_upper.as_slice() {
b"LOAD" => script_load(ctx, args, out),
b"EXISTS" => script_exists(ctx, args, out),
b"FLUSH" => script_flush(ctx, args, out),
_ => encode_error(out, "ERR SCRIPT subcommand must be one of LOAD, EXISTS, FLUSH"),
}
}
fn script_load<A: ArgvView + ?Sized>(ctx: &Ctx<'_>, args: &A, out: &mut Vec<u8>) {
if args.len() != 3 {
wrong_args(out, "script|load");
return;
}
let source = args.get(2).unwrap_or(b"");
let sha = kevy_lua::sha1::sha1(source);
scripts(ctx).insert(sha, source.to_vec());
let hex = kevy_lua::sha1::hex(&sha);
out.push(b'$');
out.extend_from_slice(b"40\r\n");
out.extend_from_slice(&hex);
out.extend_from_slice(b"\r\n");
}
fn script_exists<A: ArgvView + ?Sized>(ctx: &Ctx<'_>, args: &A, out: &mut Vec<u8>) {
if args.len() < 3 {
wrong_args(out, "script|exists");
return;
}
let cache =
ctx.state.catalogs.scripts.lock().unwrap_or_else(std::sync::PoisonError::into_inner);
let count = args.len() - 2;
out.extend_from_slice(format!("*{count}\r\n").as_bytes());
for i in 2..args.len() {
let hit = kevy_lua::sha1::parse_hex(args.get(i).unwrap_or(b""))
.is_some_and(|sha| cache.contains_key(&sha));
out.extend_from_slice(if hit { b":1\r\n" } else { b":0\r\n" });
}
}
fn script_flush<A: ArgvView + ?Sized>(ctx: &Ctx<'_>, args: &A, out: &mut Vec<u8>) {
if args.len() == 3 {
let mode = args.get(2).unwrap_or(b"");
if !mode.eq_ignore_ascii_case(b"SYNC") && !mode.eq_ignore_ascii_case(b"ASYNC") {
encode_error(out, "ERR SCRIPT FLUSH mode must be SYNC or ASYNC");
return;
}
} else if args.len() != 2 {
wrong_args(out, "script|flush");
return;
}
ctx.state.catalogs.scripts.lock().unwrap_or_else(std::sync::PoisonError::into_inner).clear();
out.extend_from_slice(b"+OK\r\n");
}
type KeysArgv<'a> = (Vec<&'a [u8]>, Vec<&'a [u8]>);
fn parse_eval_keys_argv<'a, A: ArgvView + ?Sized>(
ctx: &Ctx<'_>,
args: &'a A,
out: &mut Vec<u8>,
) -> Option<KeysArgv<'a>> {
let numkeys: usize = match parse_uint(args.get(2).unwrap_or(b"")) {
Some(n) => n,
None => {
encode_error(out, "ERR value is not an integer or out of range");
return None;
}
};
let total_after_numkeys = args.len().saturating_sub(3);
if numkeys > total_after_numkeys {
encode_error(out, "ERR Number of keys can't be greater than number of args");
return None;
}
let keys: Vec<&[u8]> = (0..numkeys).map(|i| args.get(3 + i).unwrap_or(b"")).collect();
let argv: Vec<&[u8]> =
((3 + numkeys)..args.len()).map(|i| args.get(i).unwrap_or(b"")).collect();
if let Some(crossslot) = cross_slot_check(ctx, &keys) {
out.extend_from_slice(&crossslot);
return None;
}
Some((keys, argv))
}
fn parse_uint(bytes: &[u8]) -> Option<usize> {
let s = std::str::from_utf8(bytes).ok()?;
let n: i64 = s.parse().ok()?;
if n < 0 { None } else { Some(n as usize) }
}
fn cross_slot_check(ctx: &Ctx<'_>, keys: &[&[u8]]) -> Option<Vec<u8>> {
if keys.len() < 2 {
return None;
}
let cfg = ctx.state.config();
if !cfg.cluster.enabled {
return None;
}
let first = kevy_hash::key_hash_slot(keys[0]);
for k in &keys[1..] {
if kevy_hash::key_hash_slot(k) != first {
let mut out = Vec::new();
encode_error(&mut out, "CROSSSLOT Keys in request don't hash to the same slot");
return Some(out);
}
}
None
}