use std::io::{Read, Write};
use std::net::TcpStream;
use std::sync::{Arc, Mutex, OnceLock};
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::Duration;
fn gate() -> std::sync::MutexGuard<'static, ()> {
static G: OnceLock<Mutex<()>> = OnceLock::new();
G.get_or_init(|| Mutex::new(()))
.lock()
.unwrap_or_else(|p| p.into_inner())
}
fn free_port() -> u16 {
let l = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
let p = l.local_addr().unwrap().port();
drop(l);
p
}
struct Server {
port: u16,
dir: std::path::PathBuf,
stop: Arc<AtomicBool>,
handle: Option<std::thread::JoinHandle<()>>,
}
impl Server {
fn start(nshards: usize) -> Server {
let port = free_port();
let dir = std::env::temp_dir().join(format!(
"kevy-lua-multishard-{}",
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos()
));
std::fs::create_dir_all(&dir).unwrap();
let mut cfg = kevy_config::Config::default();
cfg.server.port = port;
cfg.server.threads = nshards;
let state = Arc::new(
kevy::RuntimeState::new(Arc::new(cfg), std::path::PathBuf::new(), nshards).unwrap(),
);
let stop = Arc::new(AtomicBool::new(false));
let stop_thread = stop.clone();
let dir_thread = dir.clone();
let handle = std::thread::spawn(move || {
let rt = kevy_rt::Runtime::builder(kevy::KevyCommands::with_state(state)).bind([127, 0, 0, 1], port).shards(nshards)
.with_data_dir(dir_thread);
rt.run(stop_thread).unwrap();
});
for _ in 0..400 {
if TcpStream::connect(("127.0.0.1", port)).is_ok() {
return Server {
port,
dir,
stop,
handle: Some(handle),
};
}
std::thread::sleep(Duration::from_millis(5));
}
panic!("kevy server didn't bind {port}");
}
fn req(&self, parts: &[&[u8]]) -> Vec<u8> {
let mut s = TcpStream::connect(("127.0.0.1", self.port)).unwrap();
s.set_read_timeout(Some(Duration::from_secs(5))).unwrap();
let mut buf = Vec::new();
buf.extend_from_slice(format!("*{}\r\n", parts.len()).as_bytes());
for p in parts {
buf.extend_from_slice(format!("${}\r\n", p.len()).as_bytes());
buf.extend_from_slice(p);
buf.extend_from_slice(b"\r\n");
}
s.write_all(&buf).unwrap();
let mut reply = Vec::new();
let mut chunk = [0u8; 4096];
for _ in 0..8 {
match s.read(&mut chunk) {
Ok(0) => break,
Ok(n) => {
reply.extend_from_slice(&chunk[..n]);
if looks_complete(&reply) {
break;
}
}
Err(_) => break,
}
}
reply
}
}
fn looks_complete(reply: &[u8]) -> bool {
reply.ends_with(b"\r\n")
}
impl Drop for Server {
fn drop(&mut self) {
self.stop.store(true, Ordering::Relaxed);
let _ = TcpStream::connect(("127.0.0.1", self.port));
if let Some(h) = self.handle.take() {
let _ = h.join();
}
let _ = std::fs::remove_dir_all(&self.dir);
}
}
#[test]
fn eval_writes_visible_across_shards() {
let _g = gate();
let s = Server::start(4);
for i in 0..50 {
let key = format!("multi:shard:{i}");
let val = format!("v-{i}");
let r_set = s.req(&[b"SET", key.as_bytes(), val.as_bytes()]);
assert_eq!(r_set, b"+OK\r\n", "SET failed for {key}");
let script = b"return redis.call('GET', KEYS[1])";
let r_eval = s.req(&[b"EVAL", script, b"1", key.as_bytes()]);
let want = format!("${}\r\n{}\r\n", val.len(), val);
assert_eq!(
r_eval,
want.as_bytes(),
"EVAL GET wrong reply for {key}: {:?}",
String::from_utf8_lossy(&r_eval)
);
}
}
#[test]
fn redlock_canonical_works_across_shards() {
let _g = gate();
let s = Server::start(4);
let unlock_script = b"if redis.call('GET', KEYS[1]) == ARGV[1] then\n\
return redis.call('DEL', KEYS[1])\n\
else\n\
return 0\n\
end";
for i in 0..30 {
let key = format!("lock:order:{i}");
let token = format!("tok-{i}");
assert_eq!(
s.req(&[b"SET", key.as_bytes(), token.as_bytes()]),
b"+OK\r\n"
);
let r = s.req(&[b"EVAL", unlock_script, b"1", key.as_bytes(), token.as_bytes()]);
assert_eq!(r, b":1\r\n", "redlock unlock returned wrong reply for {key}");
assert_eq!(s.req(&[b"GET", key.as_bytes()]), b"$-1\r\n", "lock leaked");
}
}
#[test]
fn script_load_then_evalsha_across_shards() {
let _g = gate();
let s = Server::start(4);
let r_load = s.req(&[b"SCRIPT", b"LOAD", b"return 'cached-' .. KEYS[1]"]);
assert!(r_load.starts_with(b"$40\r\n"), "LOAD got {:?}", String::from_utf8_lossy(&r_load));
let sha = r_load[5..45].to_vec();
for i in 0..30 {
let key = format!("k-{i}");
let r = s.req(&[b"EVALSHA", &sha, b"1", key.as_bytes()]);
let want = format!("$8\r\ncached-{}\r\n", &key[2..3]);
let want_payload = format!("cached-{key}");
let want_bulk = format!("${}\r\n{}\r\n", want_payload.len(), want_payload);
assert_eq!(
r,
want_bulk.as_bytes(),
"EVALSHA missing for {key}: {:?} (sanity want={want})",
String::from_utf8_lossy(&r),
);
}
}
#[test]
fn script_flush_clears_global_cache() {
let _g = gate();
let s = Server::start(4);
let r_load = s.req(&[b"SCRIPT", b"LOAD", b"return 42"]);
let sha = r_load[5..45].to_vec();
let r_exists = s.req(&[b"SCRIPT", b"EXISTS", &sha]);
assert_eq!(r_exists, b"*1\r\n:1\r\n");
assert_eq!(s.req(&[b"SCRIPT", b"FLUSH"]), b"+OK\r\n");
let r_after = s.req(&[b"SCRIPT", b"EXISTS", &sha]);
assert_eq!(r_after, b"*1\r\n:0\r\n");
}
#[test]
fn route_eval_to_key1_shard() {
use kevy_resp::Argv;
use kevy_rt::Commands;
let mut a = Argv::default();
a.push(b"EVAL");
a.push(b"return 1");
a.push(b"1");
a.push(b"mykey");
let r = kevy::KevyCommands::new().route(&a);
let s = format!("{r:?}");
assert!(s.contains("Single(3)"), "EVAL with numkeys=1 must Route::Single(3) (KEYS[1] at args[3]), got: {s}");
}
#[test]
fn route_eval_numkeys0_local() {
use kevy_resp::Argv;
use kevy_rt::Commands;
let mut a = Argv::default();
a.push(b"EVAL");
a.push(b"return 1");
a.push(b"0");
let r = kevy::KevyCommands::new().route(&a);
let s = format!("{r:?}");
assert!(s.contains("Local"), "EVAL with numkeys=0 must Route::Local, got: {s}");
}
#[test]
fn route_script_local() {
use kevy_resp::Argv;
use kevy_rt::Commands;
let mut a = Argv::default();
a.push(b"SCRIPT");
a.push(b"LOAD");
a.push(b"return 1");
let r = kevy::KevyCommands::new().route(&a);
let s = format!("{r:?}");
assert!(s.contains("Local"), "SCRIPT must Route::Local (global cache), got: {s}");
}