use std::io::{Read, Write};
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, Mutex};
static START_GATE: Mutex<()> = Mutex::new(());
fn free_port() -> u16 {
let l = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
l.local_addr().unwrap().port()
}
fn req(parts: &[&[u8]]) -> Vec<u8> {
let mut v = format!("*{}\r\n", parts.len()).into_bytes();
for p in parts {
v.extend_from_slice(format!("${}\r\n", p.len()).as_bytes());
v.extend_from_slice(p);
v.extend_from_slice(b"\r\n");
}
v
}
fn read_reply(s: &mut std::net::TcpStream, expected: &[u8]) {
let mut buf = vec![0u8; expected.len()];
s.read_exact(&mut buf).unwrap();
assert_eq!(
&buf,
expected,
"expected {:?}",
String::from_utf8_lossy(expected)
);
}
struct Server {
port: u16,
dir: std::path::PathBuf,
stop: Arc<AtomicBool>,
handle: Option<std::thread::JoinHandle<()>>,
}
impl Server {
fn start(nshards: usize) -> Server {
let _gate = START_GATE.lock().unwrap_or_else(std::sync::PoisonError::into_inner);
let port = free_port();
let dir = std::env::temp_dir().join(format!(
"kevy-sharded-{}",
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos()
));
std::fs::create_dir_all(&dir).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::new([127, 0, 0, 1], port, nshards, kevy::KevyCommands)
.with_data_dir(dir_thread);
rt.run(stop_thread).unwrap();
});
let mut ready = false;
for _ in 0..200 {
if std::net::TcpStream::connect(("127.0.0.1", port)).is_ok() {
ready = true;
break;
}
std::thread::sleep(std::time::Duration::from_millis(5));
}
assert!(ready, "runtime did not come up");
Server {
port,
dir,
stop,
handle: Some(handle),
}
}
fn connect(&self) -> std::net::TcpStream {
std::net::TcpStream::connect(("127.0.0.1", self.port)).unwrap()
}
}
impl Drop for Server {
fn drop(&mut self) {
self.stop.store(true, Ordering::Relaxed);
if let Some(h) = self.handle.take() {
let _ = h.join();
}
let _ = std::fs::remove_dir_all(&self.dir);
}
}
#[test]
fn keyspace_is_shared_across_cores() {
let srv = Server::start(4);
let mut writer = srv.connect();
for i in 0..200u32 {
let key = format!("k{i}");
let val = format!("v{i}");
writer
.write_all(&req(&[b"SET", key.as_bytes(), val.as_bytes()]))
.unwrap();
read_reply(&mut writer, b"+OK\r\n");
}
let mut reader = srv.connect();
for i in 0..200u32 {
let key = format!("k{i}");
let want = format!("v{i}");
reader.write_all(&req(&[b"GET", key.as_bytes()])).unwrap();
let expected = format!("${}\r\n{}\r\n", want.len(), want);
read_reply(&mut reader, expected.as_bytes());
}
}
#[test]
fn pipelined_order_is_preserved() {
let srv = Server::start(4);
let mut c = srv.connect();
let mut batch = Vec::new();
let mut expected = Vec::new();
for i in 0..50u32 {
let key = format!("ord{i}");
batch.extend_from_slice(&req(&[b"SET", key.as_bytes(), format!("{i}").as_bytes()]));
expected.extend_from_slice(b"+OK\r\n");
batch.extend_from_slice(&req(&[b"GET", key.as_bytes()]));
let v = format!("{i}");
expected.extend_from_slice(format!("${}\r\n{}\r\n", v.len(), v).as_bytes());
}
c.write_all(&batch).unwrap();
let mut got = vec![0u8; expected.len()];
c.read_exact(&mut got).unwrap();
assert_eq!(got, expected, "pipelined replies out of order");
}
#[test]
fn fanout_dbsize_del_flush() {
let srv = Server::start(4);
let mut c = srv.connect();
for i in 0..30u32 {
c.write_all(&req(&[b"SET", format!("f{i}").as_bytes(), b"x"]))
.unwrap();
read_reply(&mut c, b"+OK\r\n");
}
c.write_all(&req(&[b"DBSIZE"])).unwrap();
read_reply(&mut c, b":30\r\n");
c.write_all(&req(&[b"DEL", b"f0", b"f1", b"f2", b"nope"]))
.unwrap();
read_reply(&mut c, b":3\r\n");
c.write_all(&req(&[b"DBSIZE"])).unwrap();
read_reply(&mut c, b":27\r\n");
c.write_all(&req(&[b"FLUSHALL"])).unwrap();
read_reply(&mut c, b"+OK\r\n");
c.write_all(&req(&[b"DBSIZE"])).unwrap();
read_reply(&mut c, b":0\r\n");
}
#[test]
fn hash_type_across_cores() {
let srv = Server::start(4);
let mut w = srv.connect();
w.write_all(&req(&[
b"HSET", b"user:1", b"name", b"alice", b"age", b"30",
]))
.unwrap();
read_reply(&mut w, b":2\r\n");
w.write_all(&req(&[b"HSET", b"user:1", b"age", b"31"]))
.unwrap(); read_reply(&mut w, b":0\r\n");
let mut r = srv.connect();
r.write_all(&req(&[b"HGET", b"user:1", b"name"])).unwrap();
read_reply(&mut r, b"$5\r\nalice\r\n");
r.write_all(&req(&[b"HLEN", b"user:1"])).unwrap();
read_reply(&mut r, b":2\r\n");
r.write_all(&req(&[b"HINCRBY", b"user:1", b"age", b"1"]))
.unwrap();
read_reply(&mut r, b":32\r\n");
r.write_all(&req(&[b"TYPE", b"user:1"])).unwrap();
read_reply(&mut r, b"+hash\r\n");
r.write_all(&req(&[b"GET", b"user:1"])).unwrap();
let mut buf = [0u8; 64];
let n = r.read(&mut buf).unwrap();
assert!(
buf[..n].starts_with(b"-WRONGTYPE"),
"got {:?}",
String::from_utf8_lossy(&buf[..n])
);
}
#[test]
fn list_type_across_cores() {
let srv = Server::start(4);
let mut w = srv.connect();
w.write_all(&req(&[b"RPUSH", b"q", b"a", b"b", b"c"]))
.unwrap();
read_reply(&mut w, b":3\r\n");
w.write_all(&req(&[b"LPUSH", b"q", b"z"])).unwrap();
read_reply(&mut w, b":4\r\n");
let mut r = srv.connect();
r.write_all(&req(&[b"LRANGE", b"q", b"0", b"-1"])).unwrap();
read_reply(
&mut r,
b"*4\r\n$1\r\nz\r\n$1\r\na\r\n$1\r\nb\r\n$1\r\nc\r\n",
);
r.write_all(&req(&[b"LPOP", b"q"])).unwrap();
read_reply(&mut r, b"$1\r\nz\r\n");
r.write_all(&req(&[b"LLEN", b"q"])).unwrap();
read_reply(&mut r, b":3\r\n");
}
fn read_len(s: &mut std::net::TcpStream, prefix: u8) -> i64 {
let mut b = [0u8; 1];
s.read_exact(&mut b).unwrap();
assert_eq!(b[0], prefix, "unexpected RESP prefix");
let mut num = Vec::new();
loop {
let mut c = [0u8; 1];
s.read_exact(&mut c).unwrap();
if c[0] == b'\r' {
s.read_exact(&mut [0u8; 1]).unwrap(); break;
}
num.push(c[0]);
}
String::from_utf8(num).unwrap().parse().unwrap()
}
fn read_array_sorted(s: &mut std::net::TcpStream) -> Vec<Vec<u8>> {
let n = read_len(s, b'*');
let mut items = Vec::new();
for _ in 0..n {
let len = read_len(s, b'$') as usize;
let mut buf = vec![0u8; len];
s.read_exact(&mut buf).unwrap();
s.read_exact(&mut [0u8; 2]).unwrap(); items.push(buf);
}
items.sort();
items
}
#[test]
fn cross_shard_multikey() {
let srv = Server::start(4);
let mut c = srv.connect();
c.write_all(&req(&[b"MSET", b"a", b"1", b"b", b"2", b"c", b"3"]))
.unwrap();
read_reply(&mut c, b"+OK\r\n");
c.write_all(&req(&[b"MGET", b"a", b"missing", b"c"]))
.unwrap();
read_reply(&mut c, b"*3\r\n$1\r\n1\r\n$-1\r\n$1\r\n3\r\n");
c.write_all(&req(&[b"SADD", b"s1", b"x", b"y", b"z"]))
.unwrap();
read_reply(&mut c, b":3\r\n");
c.write_all(&req(&[b"SADD", b"s2", b"y", b"z", b"w"]))
.unwrap();
read_reply(&mut c, b":3\r\n");
c.write_all(&req(&[b"SINTER", b"s1", b"s2"])).unwrap();
assert_eq!(
read_array_sorted(&mut c),
vec![b"y".to_vec(), b"z".to_vec()]
);
c.write_all(&req(&[b"SUNION", b"s1", b"s2"])).unwrap();
assert_eq!(
read_array_sorted(&mut c),
vec![b"w".to_vec(), b"x".to_vec(), b"y".to_vec(), b"z".to_vec()]
);
c.write_all(&req(&[b"SDIFF", b"s1", b"s2"])).unwrap();
assert_eq!(read_array_sorted(&mut c), vec![b"x".to_vec()]);
}
#[test]
fn keys_scan_randomkey_across_cores() {
let srv = Server::start(4);
let mut c = srv.connect();
for i in 0..6u32 {
c.write_all(&req(&[b"SET", format!("u:{i}").as_bytes(), b"x"]))
.unwrap();
read_reply(&mut c, b"+OK\r\n");
}
c.write_all(&req(&[b"SET", b"other", b"y"])).unwrap();
read_reply(&mut c, b"+OK\r\n");
c.write_all(&req(&[b"KEYS", b"u:*"])).unwrap();
assert_eq!(read_array_sorted(&mut c).len(), 6);
c.write_all(&req(&[b"KEYS", b"*"])).unwrap();
assert_eq!(read_array_sorted(&mut c).len(), 7);
c.write_all(&req(&[b"SCAN", b"0", b"MATCH", b"u:*"]))
.unwrap();
assert_eq!(read_len(&mut c, b'*'), 2);
let curlen = read_len(&mut c, b'$') as usize;
let mut cur = vec![0u8; curlen];
c.read_exact(&mut cur).unwrap();
c.read_exact(&mut [0u8; 2]).unwrap();
assert_eq!(cur, b"0");
assert_eq!(read_array_sorted(&mut c).len(), 6);
c.write_all(&req(&[b"RANDOMKEY"])).unwrap();
let l = read_len(&mut c, b'$');
assert!(l > 0);
c.read_exact(&mut vec![0u8; l as usize]).unwrap();
c.read_exact(&mut [0u8; 2]).unwrap();
}
#[test]
fn pubsub_across_cores() {
let srv = Server::start(4);
let mut sub = srv.connect();
sub.write_all(&req(&[b"SUBSCRIBE", b"news"])).unwrap();
read_reply(&mut sub, b"*3\r\n$9\r\nsubscribe\r\n$4\r\nnews\r\n:1\r\n");
let mut publisher = srv.connect();
publisher
.write_all(&req(&[b"PUBLISH", b"news", b"hello"]))
.unwrap();
read_reply(&mut publisher, b":1\r\n");
read_reply(
&mut sub,
b"*3\r\n$7\r\nmessage\r\n$4\r\nnews\r\n$5\r\nhello\r\n",
);
publisher
.write_all(&req(&[b"PUBLISH", b"empty", b"x"]))
.unwrap();
read_reply(&mut publisher, b":0\r\n");
}
#[test]
fn subscriber_disconnect_unregisters_subs() {
let srv = Server::start(4);
let mut sub = srv.connect();
sub.write_all(&req(&[b"SUBSCRIBE", b"chA", b"chB"])).unwrap();
read_reply(&mut sub, b"*3\r\n$9\r\nsubscribe\r\n$3\r\nchA\r\n:1\r\n");
read_reply(&mut sub, b"*3\r\n$9\r\nsubscribe\r\n$3\r\nchB\r\n:2\r\n");
let mut publisher = srv.connect();
publisher.write_all(&req(&[b"PUBLISH", b"chA", b"x"])).unwrap();
read_reply(&mut publisher, b":1\r\n");
read_reply(&mut sub, b"*3\r\n$7\r\nmessage\r\n$3\r\nchA\r\n$1\r\nx\r\n");
drop(sub);
std::thread::sleep(std::time::Duration::from_millis(50));
publisher.write_all(&req(&[b"PUBLISH", b"chA", b"y"])).unwrap();
read_reply(&mut publisher, b":0\r\n");
publisher.write_all(&req(&[b"PUBLISH", b"chB", b"z"])).unwrap();
read_reply(&mut publisher, b":0\r\n");
}
#[test]
fn transactions() {
let srv = Server::start(4);
let mut c = srv.connect();
c.write_all(&req(&[b"MULTI"])).unwrap();
read_reply(&mut c, b"+OK\r\n");
c.write_all(&req(&[b"SET", b"tx:a", b"1"])).unwrap();
read_reply(&mut c, b"+QUEUED\r\n");
c.write_all(&req(&[b"INCR", b"tx:a"])).unwrap();
read_reply(&mut c, b"+QUEUED\r\n");
c.write_all(&req(&[b"GET", b"tx:a"])).unwrap();
read_reply(&mut c, b"+QUEUED\r\n");
c.write_all(&req(&[b"EXEC"])).unwrap();
read_reply(&mut c, b"*3\r\n+OK\r\n:2\r\n$1\r\n2\r\n");
c.write_all(&req(&[b"GET", b"tx:a"])).unwrap();
read_reply(&mut c, b"$1\r\n2\r\n");
c.write_all(&req(&[b"MULTI"])).unwrap();
read_reply(&mut c, b"+OK\r\n");
c.write_all(&req(&[b"SET", b"tx:b", b"x"])).unwrap();
read_reply(&mut c, b"+QUEUED\r\n");
c.write_all(&req(&[b"DISCARD"])).unwrap();
read_reply(&mut c, b"+OK\r\n");
c.write_all(&req(&[b"GET", b"tx:b"])).unwrap();
read_reply(&mut c, b"$-1\r\n");
c.write_all(&req(&[b"EXEC"])).unwrap();
let mut buf = [0u8; 64];
let n = c.read(&mut buf).unwrap();
assert!(buf[..n].starts_with(b"-ERR EXEC without MULTI"));
}
#[test]
fn single_shard_still_works() {
let srv = Server::start(1);
let mut c = srv.connect();
c.write_all(&req(&[b"SET", b"a", b"1"])).unwrap();
read_reply(&mut c, b"+OK\r\n");
c.write_all(&req(&[b"INCR", b"a"])).unwrap();
read_reply(&mut c, b":2\r\n");
c.write_all(&req(&[b"GET", b"a"])).unwrap();
read_reply(&mut c, b"$1\r\n2\r\n");
}
#[test]
fn pipelined_cross_shard_no_deadlock() {
let srv = Server::start(4);
let mut c = srv.connect();
c.set_read_timeout(Some(std::time::Duration::from_secs(30)))
.unwrap();
let n = 10_000usize;
let mut buf = Vec::new();
for i in 0..n {
let key = format!("pp:{i}"); buf.extend_from_slice(&req(&[b"INCR", key.as_bytes()]));
}
c.write_all(&buf).unwrap();
let mut expected = Vec::with_capacity(n * 4);
for _ in 0..n {
expected.extend_from_slice(b":1\r\n");
}
let mut got = vec![0u8; expected.len()];
c.read_exact(&mut got).unwrap();
assert_eq!(got, expected);
}
fn read_info(s: &mut std::net::TcpStream, section: &str) -> String {
s.set_read_timeout(Some(std::time::Duration::from_secs(2)))
.unwrap();
s.write_all(&req(&[b"INFO", section.as_bytes()])).unwrap();
let mut buf = Vec::new();
let mut tmp = [0u8; 8192];
loop {
if let Some(hdr_end) = buf.windows(2).position(|w| w == b"\r\n")
&& buf.first() == Some(&b'$')
{
let len: usize = std::str::from_utf8(&buf[1..hdr_end])
.unwrap()
.parse()
.unwrap();
if buf.len() >= hdr_end + 2 + len {
return String::from_utf8_lossy(&buf[hdr_end + 2..hdr_end + 2 + len]).into_owned();
}
}
let n = s.read(&mut tmp).unwrap();
if n == 0 {
break;
}
buf.extend_from_slice(&tmp[..n]);
}
String::from_utf8_lossy(&buf).into_owned()
}
fn info_field(body: &str, field: &str) -> u64 {
body.lines()
.find_map(|l| l.strip_prefix(&format!("{field}:")))
.unwrap_or_else(|| panic!("field {field} not in INFO:\n{body}"))
.trim()
.parse()
.unwrap_or_else(|_| panic!("field {field} not a u64 in INFO:\n{body}"))
}
#[test]
fn info_aggregates_across_shards() {
const N: u32 = 400; const M: u32 = 120; let srv = Server::start(4);
let mut c = srv.connect();
for i in 0..N {
c.write_all(&req(&[b"SET", format!("k{i}").as_bytes(), b"value-payload"]))
.unwrap();
read_reply(&mut c, b"+OK\r\n");
}
for i in 0..M {
c.write_all(&req(&[
b"SET",
format!("t{i}").as_bytes(),
b"v",
b"EX",
b"600",
]))
.unwrap();
read_reply(&mut c, b"+OK\r\n");
}
let total = u64::from(N + M);
let db0_field = |body: &str, k: &str| -> Option<u64> {
body.lines()
.find_map(|l| l.strip_prefix("db0:"))
.and_then(|db0| db0.split(',').find_map(|p| p.strip_prefix(&format!("{k}="))))
.and_then(|v| v.parse().ok())
};
let deadline = std::time::Instant::now() + std::time::Duration::from_secs(15);
let ks = loop {
let ks = read_info(&mut c, "keyspace");
if db0_field(&ks, "keys") == Some(total) {
break ks;
}
assert!(
std::time::Instant::now() < deadline,
"INFO keyspace never summed to {total} keys across shards:\n{ks}"
);
std::thread::sleep(std::time::Duration::from_millis(50));
};
assert_eq!(db0_field(&ks, "keys"), Some(total), "keyspace keys not summed");
assert_eq!(
db0_field(&ks, "expires"),
Some(u64::from(M)),
"expire-set count wrong"
);
let mem = read_info(&mut c, "memory");
let used = info_field(&mem, "used_memory");
assert!(
used >= total * 48,
"used_memory {used} looks like a single shard, not the {total}-key sum"
);
let stats = read_info(&mut c, "stats");
assert!(
info_field(&stats, "total_commands_processed") >= total,
"commands_processed not summed across shards"
);
assert!(
info_field(&stats, "total_connections_received") >= 1,
"connections_received not counted"
);
assert_eq!(info_field(&stats, "expired_keys"), 0, "nothing should expire");
}