#![forbid(unsafe_code)]
#![warn(missing_docs)]
use std::collections::HashMap;
use std::io;
use kevy_resp::Reply;
use kevy_resp_client::RespClient;
pub struct ReadWriteClient {
primary: RespClient,
replicas: Vec<RespClient>,
rr_counter: usize,
scope_writers: HashMap<String, RespClient>,
scope_key_targets: HashMap<Vec<u8>, String>,
}
const SCOPE_KEY_CACHE_CAP: usize = 4096;
const QUIESCE_RETRY_BUDGET: usize = 7;
const QUIESCE_RETRY_MIN_MS: u64 = 5;
const QUIESCE_RETRY_MAX_MS: u64 = 80;
impl ReadWriteClient {
pub fn connect(primary: (&str, u16), replicas: &[(&str, u16)]) -> io::Result<Self> {
let primary_conn = RespClient::connect(primary.0, primary.1)?;
let mut replica_conns = Vec::with_capacity(replicas.len());
for (host, port) in replicas {
replica_conns.push(RespClient::connect(host, *port)?);
}
Ok(Self {
primary: primary_conn,
replicas: replica_conns,
rr_counter: 0,
scope_writers: HashMap::new(),
scope_key_targets: HashMap::new(),
})
}
pub fn replica_count(&self) -> usize {
self.replicas.len()
}
pub fn request_write(&mut self, args: &[Vec<u8>]) -> io::Result<Reply> {
if let Some(key) = args.get(1)
&& let Some(addr) = self.scope_key_targets.get(key.as_slice()).cloned()
{
return self.request_via_writer(&addr, args);
}
self.request_write_with_quiesce_retry(args)
}
fn request_write_with_quiesce_retry(&mut self, args: &[Vec<u8>]) -> io::Result<Reply> {
let mut backoff = std::time::Duration::from_millis(QUIESCE_RETRY_MIN_MS);
for _ in 0..QUIESCE_RETRY_BUDGET {
let reply = self.primary.request(args)?;
if let Some(target_addr) = parse_misdirected(&reply) {
if let Some(key) = args.get(1) {
self.remember_key_target(key, &target_addr);
}
return self.request_via_writer(&target_addr, args);
}
if parse_quiesced(&reply).is_some() {
std::thread::sleep(backoff);
backoff = (backoff * 2).min(std::time::Duration::from_millis(QUIESCE_RETRY_MAX_MS));
continue;
}
return Ok(reply);
}
self.primary.request(args)
}
fn request_via_writer(&mut self, addr: &str, args: &[Vec<u8>]) -> io::Result<Reply> {
if !self.scope_writers.contains_key(addr) {
let (host, port) = split_host_port(addr).ok_or_else(|| {
io::Error::new(
io::ErrorKind::InvalidData,
format!("server returned MISDIRECTED with malformed target {addr:?}"),
)
})?;
let conn = RespClient::connect(host, port)?;
self.scope_writers.insert(addr.to_string(), conn);
}
let conn = self
.scope_writers
.get_mut(addr)
.expect("just inserted above");
conn.request(args)
}
fn remember_key_target(&mut self, key: &[u8], addr: &str) {
if self.scope_key_targets.len() >= SCOPE_KEY_CACHE_CAP {
self.scope_key_targets.clear();
}
self.scope_key_targets.insert(key.to_vec(), addr.to_string());
}
pub fn request_read(&mut self, args: &[Vec<u8>], consistent: bool) -> io::Result<Reply> {
if consistent || self.replicas.is_empty() {
return self.primary.request(args);
}
let idx = self.rr_counter % self.replicas.len();
self.rr_counter = self.rr_counter.wrapping_add(1);
self.replicas[idx].request(args)
}
pub fn request(&mut self, args: &[Vec<u8>]) -> io::Result<Reply> {
let Some(verb) = args.first() else {
return self.primary.request(args);
};
if is_write_verb(verb) {
self.request_write(args)
} else {
self.request_read(args, false)
}
}
}
pub fn is_write_verb(verb: &[u8]) -> bool {
let mut buf = [0u8; 32];
let upper = ascii_upper(verb, &mut buf);
matches!(
upper,
b"SET" | b"SETNX" | b"SETEX" | b"PSETEX" | b"MSET" | b"MSETNX"
| b"APPEND" | b"INCR" | b"INCRBY" | b"INCRBYFLOAT"
| b"DECR" | b"DECRBY" | b"GETSET" | b"GETDEL"
| b"SETRANGE"
| b"DEL" | b"UNLINK" | b"EXPIRE" | b"EXPIREAT" | b"PEXPIRE" | b"PEXPIREAT"
| b"PERSIST" | b"RENAME" | b"RENAMENX" | b"TYPE" | b"COPY" | b"OBJECT"
| b"HSET" | b"HSETNX" | b"HMSET" | b"HDEL" | b"HINCRBY" | b"HINCRBYFLOAT"
| b"LPUSH" | b"RPUSH" | b"LPUSHX" | b"RPUSHX" | b"LPOP" | b"RPOP"
| b"LREM" | b"LTRIM" | b"LSET" | b"LINSERT" | b"RPOPLPUSH" | b"LMOVE"
| b"BLPOP" | b"BRPOP" | b"BLMOVE"
| b"SADD" | b"SREM" | b"SPOP" | b"SMOVE" | b"SINTERSTORE" | b"SUNIONSTORE" | b"SDIFFSTORE"
| b"ZADD" | b"ZREM" | b"ZINCRBY" | b"ZPOPMIN" | b"ZPOPMAX"
| b"ZREMRANGEBYRANK" | b"ZREMRANGEBYSCORE" | b"ZREMRANGEBYLEX"
| b"XADD" | b"XDEL" | b"XTRIM" | b"XGROUP" | b"XACK" | b"XCLAIM" | b"XAUTOCLAIM"
| b"FLUSHDB" | b"FLUSHALL" | b"CONFIG" | b"SAVE" | b"BGSAVE" | b"BGREWRITEAOF"
| b"REPLICAOF" | b"SLAVEOF"
| b"PUBLISH" | b"SPUBLISH"
| b"MULTI" | b"EXEC" | b"DISCARD" | b"WATCH" | b"UNWATCH"
)
}
fn ascii_upper<'a>(s: &[u8], buf: &'a mut [u8; 32]) -> &'a [u8] {
let n = s.len().min(32);
for i in 0..n {
buf[i] = s[i].to_ascii_uppercase();
}
&buf[..n]
}
fn parse_misdirected(reply: &Reply) -> Option<String> {
let Reply::Error(bytes) = reply else { return None };
const PREFIX: &[u8] = b"MISDIRECTED writer is ";
if !bytes.starts_with(PREFIX) {
return None;
}
let addr = std::str::from_utf8(&bytes[PREFIX.len()..]).ok()?;
let addr = addr.trim_end_matches(['\r', '\n']);
if addr.is_empty() {
return None;
}
Some(addr.to_string())
}
fn parse_quiesced(reply: &Reply) -> Option<String> {
let Reply::Error(bytes) = reply else { return None };
const PREFIX: &[u8] = b"QUIESCED migrating to ";
if !bytes.starts_with(PREFIX) {
return None;
}
let addr = std::str::from_utf8(&bytes[PREFIX.len()..]).ok()?;
let addr = addr.trim_end_matches(['\r', '\n']);
if addr.is_empty() {
return None;
}
Some(addr.to_string())
}
fn split_host_port(addr: &str) -> Option<(&str, u16)> {
let colon = addr.rfind(':')?;
let host = &addr[..colon];
if host.is_empty() {
return None;
}
let port: u16 = addr[colon + 1..].parse().ok()?;
Some((host, port))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn writes_classified_correctly() {
for verb in [&b"SET"[..], b"DEL", b"LPUSH", b"HSET", b"ZADD", b"XADD", b"FLUSHDB", b"REPLICAOF"] {
assert!(is_write_verb(verb), "{:?} should be write", std::str::from_utf8(verb));
}
}
#[test]
fn reads_classified_correctly() {
for verb in [&b"GET"[..], b"HGET", b"LRANGE", b"SMEMBERS", b"ZSCORE", b"XRANGE", b"PING", b"INFO"] {
assert!(!is_write_verb(verb), "{:?} should be read", std::str::from_utf8(verb));
}
}
#[test]
fn classification_is_case_insensitive() {
assert!(is_write_verb(b"set"));
assert!(is_write_verb(b"Set"));
assert!(is_write_verb(b"SET"));
assert!(!is_write_verb(b"get"));
assert!(!is_write_verb(b"Get"));
}
#[test]
fn long_verb_doesnt_panic_on_classification() {
assert!(!is_write_verb(&[b'X'; 64]));
}
#[test]
fn parse_misdirected_basic() {
let r = Reply::Error(b"MISDIRECTED writer is 10.0.0.1:6004".to_vec());
assert_eq!(parse_misdirected(&r).as_deref(), Some("10.0.0.1:6004"));
}
#[test]
fn parse_misdirected_strips_trailing_crlf() {
let r = Reply::Error(b"MISDIRECTED writer is 10.0.0.1:6004\r\n".to_vec());
assert_eq!(parse_misdirected(&r).as_deref(), Some("10.0.0.1:6004"));
}
#[test]
fn parse_misdirected_rejects_unrelated_error() {
let r = Reply::Error(b"ERR something else".to_vec());
assert!(parse_misdirected(&r).is_none());
let r = Reply::Simple(b"OK".to_vec());
assert!(parse_misdirected(&r).is_none());
}
#[test]
fn split_host_port_dotted_v4_and_dns() {
assert_eq!(split_host_port("10.0.0.1:6004"), Some(("10.0.0.1", 6004)));
assert_eq!(split_host_port("db.local:6105"), Some(("db.local", 6105)));
}
#[test]
fn parse_quiesced_basic() {
let r = Reply::Error(b"QUIESCED migrating to 10.0.0.1:6004".to_vec());
assert_eq!(parse_quiesced(&r).as_deref(), Some("10.0.0.1:6004"));
}
#[test]
fn parse_quiesced_strips_trailing_crlf() {
let r = Reply::Error(b"QUIESCED migrating to 10.0.0.1:6004\r\n".to_vec());
assert_eq!(parse_quiesced(&r).as_deref(), Some("10.0.0.1:6004"));
}
#[test]
fn parse_quiesced_rejects_unrelated_error() {
let r = Reply::Error(b"MISDIRECTED writer is 10.0.0.1:6004".to_vec());
assert!(parse_quiesced(&r).is_none());
let r = Reply::Simple(b"OK".to_vec());
assert!(parse_quiesced(&r).is_none());
}
#[test]
fn split_host_port_rejects_bad_inputs() {
assert!(split_host_port("nohost:").is_none());
assert!(split_host_port(":6004").is_none());
assert!(split_host_port("no-colon").is_none());
assert!(split_host_port("host:99999").is_none()); }
}