use std::{str, sync::Arc};
use whasher::GxPapayaMap;
use super::{
lua_options::LuaOptions,
lua_runner::RespOut,
scratch_buffer_network_sender::ScratchBufferNetworkSender,
script_hash_key::ScriptHashKey,
scripting_api::ScriptingApi,
session_script_cache::{LuaScriptHandle, RunnerCreateOptions, SHA1_LEN, SessionScriptCache},
};
use crate::resp::{
cmd_strings::GENERIC_ERR_WRONG_NUM_ARGS, parser::session_parse_state::strict_i64,
};
const ERR_VALUE_NOT_INTEGER: &[u8] = b"ERR value is not an integer or out of range.";
const ERR_SCRIPT_FLUSH_OPTION: &[u8] = b"ERR SCRIPT FLUSH only support SYNC|ASYNC option";
#[derive(Default)]
pub struct StoreScriptCache {
map: GxPapayaMap<ScriptHashKey, Arc<LuaScriptHandle>>,
}
impl StoreScriptCache {
pub fn try_get(&self, key: &ScriptHashKey) -> Option<Arc<LuaScriptHandle>> {
self.map.pin().get(key).cloned()
}
pub fn try_add(&self, key: ScriptHashKey, handle: Arc<LuaScriptHandle>) -> bool {
self.map.pin().insert(key, handle).is_none()
}
pub fn try_remove(&self, key: &ScriptHashKey) -> Option<Arc<LuaScriptHandle>> {
self.map.pin().remove(key).cloned()
}
pub fn contains_key(&self, key: &ScriptHashKey) -> bool {
self.map.pin().contains_key(key)
}
pub fn keys(&self) -> Vec<ScriptHashKey> {
self.map.pin().keys().cloned().collect()
}
}
pub struct LuaSessionContext<'a> {
pub args: &'a [Vec<u8>],
pub out: &'a mut Vec<u8>,
pub session_cache: &'a mut SessionScriptCache,
pub store_cache: &'a StoreScriptCache,
pub session: &'a mut dyn ScriptingApi,
pub lua_enabled: bool,
pub txn_mode: bool,
pub redis_version: &'a str,
pub lua_options: &'a LuaOptions,
}
impl<'a> LuaSessionContext<'a> {
fn runner_options(&self) -> RunnerCreateOptions {
RunnerCreateOptions {
mem_limit_bytes: self.lua_options.get_memory_limit_bytes(),
log_mode: Some(self.lua_options.log_mode),
allowed_functions: if self.lua_options.allowed_functions.is_empty() {
None
} else {
Some(self.lua_options.allowed_functions.iter().cloned().collect())
},
txn_mode: self.txn_mode,
redis_version: self.redis_version.to_string(),
}
}
}
pub struct LuaCommands;
impl LuaCommands {
#[allow(clippy::too_many_lines)]
pub fn try_evalsha(ctx: &mut LuaSessionContext) -> bool {
if !Self::check_lua_enabled(ctx) {
return true;
}
let count = ctx.args.len();
if count < 2 {
return Self::abort_with_wrong_number_of_arguments(ctx, "EVALSHA");
}
let Some(n) = strict_i64(&ctx.args[1]) else {
return Self::abort_with_error_message(ctx, ERR_VALUE_NOT_INTEGER);
};
if !(0..=(count as i64 - 2)).contains(&n) {
return Self::abort_with_error_message(ctx, ERR_VALUE_NOT_INTEGER);
}
let mut digest = ctx.args[0].clone();
let mut converted_to_lower = false;
let mut resolved: Option<ScriptHashKey> = None;
while digest.len() == SHA1_LEN {
let Some(script_key) = ScriptHashKey::from_hex(&digest) else {
break;
};
if ctx.session_cache.try_get_runner(&script_key).is_some() {
resolved = Some(script_key);
break;
}
if let Some(global_script_handle) = ctx.store_cache.try_get(&script_key) {
let mut handle = Some(Arc::clone(&global_script_handle));
let source = global_script_handle.script_data().to_vec();
let options = ctx.runner_options();
let mut load_out = Vec::new();
let loaded = ctx.session_cache.try_load_runner(
&source,
&script_key,
&mut handle,
&options,
&mut load_out,
);
if loaded.is_none() {
ctx.out.extend_from_slice(&load_out);
_ = ctx.store_cache.try_remove(&script_key);
return true;
}
resolved = Some(script_key);
break;
}
if !converted_to_lower {
digest.make_ascii_lowercase();
converted_to_lower = true;
continue;
}
break;
}
let Some(script_key) = resolved else {
let mut resp = RespOut::session(ctx.out, 2);
resp.write_error(b"NOSCRIPT No matching script. Please use EVAL.");
return true;
};
Self::run_script_for_session(ctx, count, &script_key);
true
}
pub fn try_eval(ctx: &mut LuaSessionContext) -> bool {
if !Self::check_lua_enabled(ctx) {
return true;
}
let count = ctx.args.len();
if count < 2 {
return Self::abort_with_wrong_number_of_arguments(ctx, "EVAL");
}
let Some(n) = strict_i64(&ctx.args[1]) else {
return Self::abort_with_error_message(ctx, ERR_VALUE_NOT_INTEGER);
};
if !(0..=(count as i64 - 2)).contains(&n) {
return Self::abort_with_error_message(ctx, ERR_VALUE_NOT_INTEGER);
}
let script = ctx.args[0].clone();
let on_stack_script_key = SessionScriptCache::get_script_digest(&script);
let global_script_handle = ctx.store_cache.try_get(&on_stack_script_key);
let mut session_script_handle = global_script_handle.clone();
let options = ctx.runner_options();
let mut load_out = Vec::new();
let loaded = ctx.session_cache.try_load_runner(
&script,
&on_stack_script_key,
&mut session_script_handle,
&options,
&mut load_out,
);
let Some((_, created)) = loaded else {
ctx.out.extend_from_slice(&load_out);
return true;
};
if let Some(new_handle) = created {
_ = ctx
.store_cache
.try_add(on_stack_script_key.clone(), new_handle);
}
Self::run_script_for_session(ctx, count, &on_stack_script_key);
true
}
pub fn network_script_exists(ctx: &mut LuaSessionContext) -> bool {
if !Self::check_lua_enabled(ctx) {
return true;
}
if ctx.args.is_empty() {
return Self::abort_with_wrong_number_of_arguments(ctx, "script|exists");
}
let mut resp = RespOut::session(ctx.out, 2);
resp.write_array_len(ctx.args.len());
for sha1 in ctx.args {
let mut exists = 0;
if let Some(key) = ScriptHashKey::from_hex(sha1) {
exists = i64::from(ctx.store_cache.contains_key(&key));
}
resp.write_int64(exists);
}
true
}
pub fn network_script_flush(ctx: &mut LuaSessionContext) -> bool {
if !Self::check_lua_enabled(ctx) {
return true;
}
if ctx.args.len() > 1 {
return Self::abort_with_error_message(ctx, ERR_SCRIPT_FLUSH_OPTION);
} else if ctx.args.len() == 1 {
let arg = ctx.args[0].to_ascii_uppercase();
if arg != b"ASYNC" && arg != b"SYNC" {
return Self::abort_with_error_message(ctx, ERR_SCRIPT_FLUSH_OPTION);
}
}
for digest in ctx.store_cache.keys() {
if let Some(script_handle) = ctx.store_cache.try_remove(&digest) {
script_handle.dispose();
}
}
let mut resp = RespOut::session(ctx.out, 2);
resp.write_direct(b"+OK\r\n");
true
}
pub fn network_script_load(ctx: &mut LuaSessionContext) -> bool {
if !Self::check_lua_enabled(ctx) {
return true;
}
if ctx.args.len() != 1 {
return Self::abort_with_wrong_number_of_arguments(ctx, "script|load");
}
let source = ctx.args[0].clone();
let digest = SessionScriptCache::get_script_digest(&source);
let global_script_handle = ctx.store_cache.try_get(&digest);
let mut session_script_handle = global_script_handle.clone();
let options = ctx.runner_options();
let mut load_out = Vec::new();
let loaded = ctx.session_cache.try_load_runner(
&source,
&digest,
&mut session_script_handle,
&options,
&mut load_out,
);
let Some((_, created)) = loaded else {
ctx.out.extend_from_slice(&load_out);
return true;
};
if let Some(new_handle) = created {
_ = ctx.store_cache.try_add(digest.clone(), new_handle);
}
let mut resp = RespOut::session(ctx.out, 2);
resp.write_bulk_string(digest.as_str().as_bytes());
true
}
pub fn check_lua_enabled(ctx: &mut LuaSessionContext) -> bool {
if !ctx.lua_enabled {
let mut resp = RespOut::session(ctx.out, 2);
resp.write_error(b"ERR This instance has Lua scripting support disabled");
return false;
}
true
}
pub fn run_script_for_session(
ctx: &mut LuaSessionContext,
count: usize,
script_key: &ScriptHashKey,
) {
if ctx.session_cache.is_running(script_key) {
return;
}
ctx.session_cache.start_running_script(script_key);
Self::try_execute_script(ctx, count - 1, script_key);
ctx.session_cache.stop_running_script(script_key);
}
pub fn try_execute_script(
ctx: &mut LuaSessionContext,
count: usize,
script_key: &ScriptHashKey,
) -> bool {
let Some(runner) = ctx.session_cache.try_get_runner(script_key) else {
return false;
};
let args = ctx.args[1..(count + 1).min(ctx.args.len())].to_vec();
runner.run_for_session(&args, ctx.session, ctx.out);
let keep = !runner.needs_dispose();
if !keep {
ctx.session_cache.remove_runner(script_key);
}
keep
}
fn abort_with_wrong_number_of_arguments(ctx: &mut LuaSessionContext, command: &str) -> bool {
let text = GENERIC_ERR_WRONG_NUM_ARGS.replace("{0}", command);
Self::abort_with_error_message(ctx, text.as_bytes())
}
fn abort_with_error_message(ctx: &mut LuaSessionContext, message: &[u8]) -> bool {
let mut resp = RespOut::session(ctx.out, 2);
resp.write_error(message);
true
}
}
pub fn script_digest(source: &[u8]) -> ScriptHashKey {
SessionScriptCache::get_script_digest(source)
}
#[derive(Default)]
pub struct NoopScriptingApi;
impl ScriptingApi for NoopScriptingApi {
fn dispatch_resp(&mut self, _request: &[u8], _sender: &mut ScratchBufferNetworkSender) {}
fn get(&mut self, _key: &[u8]) -> Result<Option<Vec<u8>>, &'static str> {
Ok(None)
}
fn set(&mut self, _key: &[u8], _value: &[u8]) -> Result<(), &'static str> {
Ok(())
}
fn resp_protocol_version(&self) -> u8 {
2
}
fn update_resp_protocol_version(&mut self, _version: u8) {}
fn check_acl_permissions(&self, _command: &str) -> bool {
true
}
}