use std::{cell::Cell, mem, sync::Arc};
use gxhash::{HashMap, HashSet};
use super::{lua_options::LuaLoggingMode, lua_runner::LuaRunner, script_hash_key::ScriptHashKey};
pub const SHA1_LEN: usize = 40;
pub struct LuaScriptHandle {
disposed: Cell<bool>,
script_data: Vec<u8>,
}
impl LuaScriptHandle {
pub fn new(script_data: Vec<u8>) -> Arc<Self> {
Arc::new(Self {
disposed: Cell::new(false),
script_data,
})
}
pub fn is_disposed(&self) -> bool {
self.disposed.get()
}
pub fn script_data(&self) -> &[u8] {
&self.script_data
}
pub fn dispose(&self) {
self.disposed.set(true);
}
}
#[derive(Clone, Default)]
pub struct RunnerCreateOptions {
pub mem_limit_bytes: Option<usize>,
pub log_mode: Option<LuaLoggingMode>,
pub allowed_functions: Option<HashSet<String>>,
pub txn_mode: bool,
pub redis_version: String,
}
pub struct ScriptCacheEntry {
pub runner: LuaRunner,
pub handle: Arc<LuaScriptHandle>,
}
#[derive(Default)]
pub struct SessionScriptCache {
script_cache: HashMap<ScriptHashKey, ScriptCacheEntry>,
scripts: HashMap<ScriptHashKey, Vec<u8>>,
running: HashMap<ScriptHashKey, u32>,
user_handle: Option<u64>,
timeout_requested: bool,
}
impl SessionScriptCache {
pub fn set_user_handle(&mut self, user_handle: Option<u64>) {
self.user_handle = user_handle;
}
pub fn user_handle(&self) -> Option<u64> {
self.user_handle
}
pub fn start_running_script(&mut self, hash: &ScriptHashKey) {
*self.running.entry(hash.clone()).or_insert(0) += 1;
}
pub fn stop_running_script(&mut self, hash: &ScriptHashKey) {
if let Some(count) = self.running.get_mut(hash) {
*count = count.saturating_sub(1);
if *count == 0 {
self.running.remove(hash);
}
}
}
pub fn is_running(&self, hash: &ScriptHashKey) -> bool {
self.running.contains_key(hash)
}
pub fn request_timeout(&mut self) {
self.timeout_requested = true;
}
pub fn take_timeout_requested(&mut self) -> bool {
mem::take(&mut self.timeout_requested)
}
pub fn try_get_from_digest(&self, hash: &ScriptHashKey) -> Option<&Vec<u8>> {
self.scripts.get(hash)
}
pub fn try_get_runner(&mut self, digest: &ScriptHashKey) -> Option<&mut LuaRunner> {
if self
.script_cache
.get(digest)
.is_some_and(|entry| entry.handle.is_disposed())
{
self.script_cache.remove(digest);
return None;
}
self
.script_cache
.get_mut(digest)
.map(|entry| &mut entry.runner)
}
pub fn try_load(&mut self, hash: &ScriptHashKey, script: &[u8]) -> bool {
self
.scripts
.entry(hash.clone())
.or_insert_with(|| script.to_vec());
true
}
pub fn try_load_runner(
&mut self,
source: &[u8],
digest: &ScriptHashKey,
global_handle: &mut Option<Arc<LuaScriptHandle>>,
options: &RunnerCreateOptions,
out: &mut Vec<u8>,
) -> Option<(&mut LuaRunner, Option<Arc<LuaScriptHandle>>)> {
if let Some(entry) = self.script_cache.get_mut(digest) {
if !entry.handle.is_disposed() {
let promoted = Arc::clone(&entry.handle);
let created = global_handle.take().map_or_else(
|| Some(promoted),
|existing| {
*global_handle = Some(existing);
None
},
);
*global_handle = Some(Arc::clone(&entry.handle));
return Some((&mut entry.runner, created));
}
self.script_cache.remove(digest);
}
let compiled_source = LuaRunner_LoaderProxy::compile_source(source);
let mut runner = LuaRunner::new(
options.log_mode.unwrap_or_default(),
options.mem_limit_bytes,
options.allowed_functions.clone().unwrap_or_default(),
compiled_source.clone(),
options.txn_mode,
&options.redis_version,
)
.ok()?;
if !runner.compile_for_session(out) {
return None;
}
let (handle, created) = match global_handle.take() {
Some(existing) => {
*global_handle = Some(Arc::clone(&existing));
(existing, None)
}
None => {
let fresh = LuaScriptHandle::new(compiled_source);
*global_handle = Some(Arc::clone(&fresh));
(Arc::clone(&fresh), Some(fresh))
}
};
self.script_cache.insert(
digest.clone(),
ScriptCacheEntry {
runner,
handle: Arc::clone(&handle),
},
);
let entry = self.script_cache.get_mut(digest)?;
Some((&mut entry.runner, created))
}
pub fn remove_runner(&mut self, key: &ScriptHashKey) {
self.script_cache.remove(key);
}
pub fn clear(&mut self) {
self.script_cache.clear();
self.scripts.clear();
self.running.clear();
}
pub fn try_swap_database_sessions(&mut self, _old_db: i32, _new_db: i32) -> bool {
true
}
pub fn get_script_digest(script: &[u8]) -> ScriptHashKey {
use sha1_smol::Sha1;
let mut hasher = Sha1::new();
hasher.update(script);
let digest = hasher.digest().bytes();
ScriptHashKey::new(&digest)
}
pub fn len(&self) -> usize {
self.script_cache.len().max(self.scripts.len())
}
pub fn is_empty(&self) -> bool {
self.script_cache.is_empty() && self.scripts.is_empty()
}
}
struct LuaRunner_LoaderProxy;
impl LuaRunner_LoaderProxy {
fn compile_source(source: &[u8]) -> Vec<u8> {
super::lua_runner__loader::LuaRunner_Loader::compile_source(source)
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use super::{LuaScriptHandle, RunnerCreateOptions, SessionScriptCache};
use crate::lua::script_hash_key::ScriptHashKey;
#[test]
fn load_get_and_digest() {
let mut cache = SessionScriptCache::default();
let script = b"return 1";
let hash = SessionScriptCache::get_script_digest(script);
assert!(cache.try_load(&hash, script));
assert_eq!(cache.try_get_from_digest(&hash).unwrap(), &script.to_vec());
assert_eq!(cache.len(), 1);
let hash2 = SessionScriptCache::get_script_digest(script);
assert!(hash.equals(&hash2));
}
#[test]
fn running_lifecycle_and_timeout() {
let mut cache = SessionScriptCache::default();
let hash = SessionScriptCache::get_script_digest(b"return redis.call('PING')");
assert!(cache.try_load(&hash, b"return redis.call('PING')"));
cache.start_running_script(&hash);
assert!(cache.is_running(&hash));
cache.request_timeout();
assert!(cache.take_timeout_requested());
cache.stop_running_script(&hash);
assert!(!cache.is_running(&hash));
cache.set_user_handle(Some(42));
assert_eq!(cache.user_handle(), Some(42));
assert!(cache.try_swap_database_sessions(0, 1));
assert_eq!(cache.len(), 1);
}
#[test]
fn try_load_runner_compiles_and_reuses() {
let mut cache = SessionScriptCache::default();
let source = b"return 'ok'";
let digest = SessionScriptCache::get_script_digest(source);
let mut out = Vec::new();
let mut global = None;
let options = RunnerCreateOptions::default();
{
let (runner, created) = cache
.try_load_runner(source, &digest, &mut global, &options, &mut out)
.expect("首次装载成功");
assert!(created.is_some());
assert!(runner.source().starts_with(b"return 'ok'"));
}
let mut global = None;
let mut out = Vec::new();
let (_, created) = cache
.try_load_runner(source, &digest, &mut global, &options, &mut out)
.expect("二次装载命中");
assert!(created.is_some(), "应上升会话句柄供全局缓存登记");
assert!(Arc::ptr_eq(
created.as_ref().unwrap(),
global.as_ref().unwrap()
));
assert_eq!(cache.len(), 1);
let mut global = created;
let mut out = Vec::new();
let (_, created) = cache
.try_load_runner(source, &digest, &mut global, &options, &mut out)
.expect("三次装载命中");
assert!(created.is_none());
assert_eq!(cache.len(), 1);
}
#[test]
fn try_load_runner_writes_error_on_bad_source() {
let mut cache = SessionScriptCache::default();
let source = b"return ]]";
let digest = SessionScriptCache::get_script_digest(source);
let mut out = Vec::new();
let mut global = None;
let options = RunnerCreateOptions::default();
assert!(
cache
.try_load_runner(source, &digest, &mut global, &options, &mut out)
.is_none()
);
assert!(out.starts_with(b"-"));
}
#[test]
fn handle_dispose_invalidates_session_entry() {
let mut cache = SessionScriptCache::default();
let digest = SessionScriptCache::get_script_digest(b"return 7");
let mut out = Vec::new();
let mut global = None;
let options = RunnerCreateOptions::default();
let (_, created) = cache
.try_load_runner(b"return 7", &digest, &mut global, &options, &mut out)
.expect("装载成功");
let handle = created.expect("新建句柄");
assert!(!handle.is_disposed());
handle.dispose();
assert!(cache.try_get_runner(&digest).is_none());
}
#[test]
fn script_hex_key_roundtrip() {
let digest = SessionScriptCache::get_script_digest(b"return 7");
let from_hex = ScriptHashKey::from_hex(digest.as_str().as_bytes()).unwrap();
assert!(digest.equals(&from_hex));
assert!(ScriptHashKey::from_hex(b"zz").is_none());
}
#[test]
fn lua_script_handle_lifecycle() {
let handle = LuaScriptHandle::new(b"return 1".to_vec());
assert_eq!(handle.script_data(), b"return 1");
assert!(!handle.is_disposed());
handle.dispose();
assert!(handle.is_disposed());
}
}