use std::collections::HashMap;
use std::io::{BufReader, Write};
use std::sync::{Arc, Mutex};
use std::net::{TcpStream, ToSocketAddrs};
use std::time::Duration;
use crate::client::ClientOptions;
use crate::codec::read_value;
use crate::error::Error;
use crate::value::RespValue;
pub const SLOT_COUNT: u16 = 16_384;
const MAX_REDIRECTS: usize = 3;
const KEYLESS: &[&str] = &[
"ACL",
"AUTH",
"BGSAVE",
"CHANGES",
"CLIENT",
"CLUSTER",
"COMMAND",
"CONFIG",
"DBSIZE",
"DISCARD",
"ECHO",
"EXEC",
"FLUSHALL",
"FLUSHDB",
"FUNCTION",
"HELLO",
"INFO",
"KEYS",
"LASTSAVE",
"MULTI",
"PING",
"PSUBSCRIBE",
"PUBLISH",
"PUNSUBSCRIBE",
"QUIT",
"RANDOMKEY",
"READONLY",
"READWRITE",
"RESET",
"SAVE",
"SCAN",
"SCRIPT",
"SELECT",
"SLOWLOG",
"SUBSCRIBE",
"SWAPDB",
"TIME",
"UNSUBSCRIBE",
"UNWATCH",
"WAIT",
];
const EVERY_ARGUMENT: &[&str] = &[
"DEL",
"EXISTS",
"MGET",
"PFCOUNT",
"PFMERGE",
"SDIFF",
"SDIFFSTORE",
"SINTER",
"SINTERSTORE",
"SSUBSCRIBE",
"SUNION",
"SUNIONSTORE",
"SUNSUBSCRIBE",
"TOUCH",
"UNLINK",
"WATCH",
];
const TWO_KEYS: &[&str] = &[
"BLMOVE",
"BRPOPLPUSH",
"COPY",
"LMOVE",
"RENAME",
"RENAMENX",
"RPOPLPUSH",
"SMOVE",
];
const SPLIT: &[&str] = &["DEL", "EXISTS", "MGET", "MSET", "TOUCH", "UNLINK"];
const TRANSACTION: &[&str] = &["DISCARD", "EXEC", "MULTI", "UNWATCH", "WATCH"];
pub fn hash_slot(key: &str) -> u16 {
crc16(hash_tag(key.as_bytes())) % SLOT_COUNT
}
fn hash_tag(key: &[u8]) -> &[u8] {
let Some(open) = find_byte(key, b'{') else {
return key;
};
let tagged = &key[open + 1..];
let Some(close) = find_byte(tagged, b'}') else {
return key;
};
if close == 0 { key } else { &tagged[..close] }
}
const CRC_TABLE: [u16; 256] = crc_table();
const fn crc_table() -> [u16; 256] {
let mut table = [0_u16; 256];
let mut value = 0;
while value < 256 {
let mut crc = (value as u16) << 8;
let mut bit = 0;
while bit < 8 {
crc = if crc & 0x8000 != 0 {
(crc << 1) ^ 0x1021
} else {
crc << 1
};
bit += 1;
}
table[value] = crc;
value += 1;
}
table
}
fn crc16(bytes: &[u8]) -> u16 {
let mut crc = 0_u16;
for byte in bytes {
let index = usize::from((crc >> 8) ^ u16::from(*byte));
crc = (crc << 8) ^ CRC_TABLE[index];
}
crc
}
fn find_byte(haystack: &[u8], needle: u8) -> Option<usize> {
#[cfg(target_arch = "x86_64")]
{
return find_byte_sse2(haystack, needle);
}
#[cfg(target_arch = "aarch64")]
{
return find_byte_neon(haystack, needle);
}
#[cfg(not(any(target_arch = "x86_64", target_arch = "aarch64")))]
{
haystack.iter().position(|byte| *byte == needle)
}
}
#[cfg(target_arch = "x86_64")]
fn find_byte_sse2(haystack: &[u8], needle: u8) -> Option<usize> {
use std::arch::x86_64::{
__m128i, _mm_cmpeq_epi8, _mm_loadu_si128, _mm_movemask_epi8, _mm_set1_epi8,
};
let width = 16;
let mut offset = 0;
let pattern = unsafe { _mm_set1_epi8(needle as i8) };
while offset + width <= haystack.len() {
let block = unsafe { _mm_loadu_si128(haystack[offset..].as_ptr() as *const __m128i) };
let mask = unsafe { _mm_movemask_epi8(_mm_cmpeq_epi8(block, pattern)) };
if mask != 0 {
return Some(offset + mask.trailing_zeros() as usize);
}
offset += width;
}
haystack[offset..]
.iter()
.position(|byte| *byte == needle)
.map(|index| offset + index)
}
#[cfg(target_arch = "aarch64")]
fn find_byte_neon(haystack: &[u8], needle: u8) -> Option<usize> {
use std::arch::aarch64::{vceqq_u8, vdupq_n_u8, vld1q_u8, vst1q_u8};
let width = 16;
let mut offset = 0;
let pattern = unsafe { vdupq_n_u8(needle) };
while offset + width <= haystack.len() {
let block = unsafe { vld1q_u8(haystack[offset..].as_ptr()) };
let equal = unsafe { vceqq_u8(block, pattern) };
let mut lanes = [0_u8; 16];
unsafe { vst1q_u8(lanes.as_mut_ptr(), equal) };
if let Some(lane) = lanes.iter().position(|lane| *lane != 0) {
return Some(offset + lane);
}
offset += width;
}
haystack[offset..]
.iter()
.position(|byte| *byte == needle)
.map(|index| offset + index)
}
fn command_name(arguments: &[&str]) -> String {
arguments
.first()
.map(|name| name.to_ascii_uppercase())
.unwrap_or_default()
}
fn listed(list: &[&str], name: &str) -> bool {
list.contains(&name)
}
pub(crate) fn command_keys<'a>(name: &str, arguments: &[&'a str]) -> Vec<&'a str> {
if arguments.len() < 2 || listed(KEYLESS, name) {
return Vec::new();
}
if listed(EVERY_ARGUMENT, name) {
return arguments[1..].to_vec();
}
if listed(TWO_KEYS, name) {
return arguments[1..].iter().take(2).copied().collect();
}
match name {
"MSET" | "MSETNX" => arguments[1..].iter().step_by(2).copied().collect(),
"BLPOP" | "BRPOP" | "BZPOPMIN" | "BZPOPMAX" => {
arguments[1..arguments.len().saturating_sub(1)].to_vec()
}
"EVAL" | "EVALSHA" | "EVAL_RO" | "EVALSHA_RO" | "FCALL" | "FCALL_RO" => {
counted(arguments, 2)
}
"ZUNION" | "ZINTER" | "ZDIFF" => counted(arguments, 1),
"ZUNIONSTORE" | "ZINTERSTORE" | "ZDIFFSTORE" => {
let mut keys = vec![arguments[1]];
keys.extend(counted(arguments, 2));
keys
}
"XREAD" | "XREADGROUP" => streams(arguments),
"OBJECT" | "MEMORY" if arguments.len() > 2 => vec![arguments[2]],
_ => vec![arguments[1]],
}
}
fn counted<'a>(arguments: &[&'a str], count_index: usize) -> Vec<&'a str> {
let Some(count) = arguments
.get(count_index)
.and_then(|value| value.parse().ok())
else {
return Vec::new();
};
arguments
.get(count_index + 1..)
.unwrap_or(&[])
.iter()
.take(count)
.copied()
.collect()
}
fn streams<'a>(arguments: &[&'a str]) -> Vec<&'a str> {
let Some(index) = arguments
.iter()
.position(|argument| argument.eq_ignore_ascii_case("STREAMS"))
else {
return Vec::new();
};
let rest = arguments.len().saturating_sub(index + 1);
arguments[index + 1..index + 1 + rest / 2].to_vec()
}
pub(crate) struct Connection {
stream: BufReader<TcpStream>,
scratch: Vec<u8>,
host: String,
port: u16,
broken: bool,
}
impl Connection {
pub(crate) fn open(host: &str, port: u16, options: &ClientOptions) -> Result<Self, Error> {
let mut last_error = None;
for address in (host, port).to_socket_addrs()? {
match TcpStream::connect_timeout(&address, options.connect_timeout) {
Ok(stream) => {
stream.set_nodelay(true)?;
let mut connection = Self {
stream: BufReader::new(stream),
scratch: Vec::with_capacity(256),
host: host.to_owned(),
port,
broken: false,
};
if let Some(password) = options.password.clone() {
let reply = if let Some(username) = options.username.clone() {
connection.send(&["AUTH", &username, &password])?
} else {
connection.send(&["AUTH", &password])?
};
throw_error(reply)?;
}
return Ok(connection);
}
Err(error) => last_error = Some(error),
}
}
Err(Error::Io(last_error.unwrap_or_else(|| {
std::io::Error::new(std::io::ErrorKind::NotFound, "no addresses for host")
})))
}
pub(crate) fn send(&mut self, arguments: &[&str]) -> Result<RespValue, Error> {
let mut values = self.send_many(&[arguments])?;
values
.pop()
.ok_or_else(|| Error::Protocol("missing reply".to_owned()))
}
pub(crate) fn send_owned(&mut self, arguments: &[String]) -> Result<RespValue, Error> {
let borrowed: Vec<&str> = arguments.iter().map(String::as_str).collect();
self.send(&borrowed)
}
pub(crate) fn send_replies(
&mut self,
arguments: &[&str],
replies: usize,
) -> Result<Vec<RespValue>, Error> {
self.write_commands(&[arguments])?;
let mut values = Vec::with_capacity(replies);
for _ in 0..replies {
values.push(self.read()?);
}
Ok(values)
}
pub(crate) fn send_many(&mut self, commands: &[&[&str]]) -> Result<Vec<RespValue>, Error> {
self.write_commands(commands)?;
let mut values = Vec::with_capacity(commands.len());
for _ in 0..commands.len() {
values.push(self.read()?);
}
Ok(values)
}
pub(crate) fn send_many_owned(
&mut self,
commands: &[Vec<String>],
) -> Result<Vec<RespValue>, Error> {
let borrowed: Vec<Vec<&str>> = commands
.iter()
.map(|command| command.iter().map(String::as_str).collect())
.collect();
let refs: Vec<&[&str]> = borrowed.iter().map(Vec::as_slice).collect();
self.send_many(&refs)
}
pub(crate) fn is_broken(&self) -> bool {
self.broken
}
fn write_commands(&mut self, commands: &[&[&str]]) -> Result<(), Error> {
if self.broken {
return Err(closed_socket());
}
let mut scratch = std::mem::take(&mut self.scratch);
scratch.clear();
for command in commands {
if let Err(error) = crate::codec::append_strings(&mut scratch, command) {
self.scratch = scratch;
return Err(error);
}
}
let write = self.stream.get_mut().write_all(&scratch);
self.scratch = scratch;
if let Err(error) = write {
self.broken = true;
return Err(Error::Io(error));
}
if let Err(error) = self.stream.get_mut().flush() {
self.broken = true;
return Err(Error::Io(error));
}
Ok(())
}
pub(crate) fn read(&mut self) -> Result<RespValue, Error> {
if self.broken {
return Err(closed_socket());
}
match read_value(&mut self.stream) {
Err(Error::Io(error)) if idle_wait(&error) && self.stream.buffer().is_empty() => {
Err(Error::Io(error))
}
Err(Error::Io(error)) => {
self.broken = true;
Err(Error::Io(error))
}
other => other,
}
}
pub(crate) fn set_read_timeout(&mut self, timeout: Option<Duration>) -> Result<(), Error> {
self.stream.get_ref().set_read_timeout(timeout)?;
Ok(())
}
}
fn idle_wait(error: &std::io::Error) -> bool {
matches!(
error.kind(),
std::io::ErrorKind::TimedOut | std::io::ErrorKind::WouldBlock
)
}
pub(crate) fn throw_error(reply: RespValue) -> Result<RespValue, Error> {
match reply {
RespValue::Error(message) => Err(Error::Server(message)),
value => Ok(value),
}
}
pub(crate) enum Discovery {
Standalone(Connection),
Cluster(Box<ClusterRouter>),
}
#[derive(Clone)]
struct Topology {
owners: Vec<Option<usize>>,
nodes: Vec<(String, u16)>,
}
impl Topology {
fn parse(reply: &RespValue) -> Option<Self> {
let RespValue::Array(ranges) = reply else {
return None;
};
if ranges.is_empty() {
return None;
}
let mut owners = vec![None; SLOT_COUNT as usize];
let mut nodes = Vec::new();
let mut ordered = ranges.clone();
ordered.sort_by_key(|range| match range {
RespValue::Array(fields) => fields.first().and_then(RespValue::as_integer).unwrap_or(0),
_ => 0,
});
for range in &ordered {
let RespValue::Array(fields) = range else {
return None;
};
if fields.len() < 3 {
return None;
}
let RespValue::Array(primary) = &fields[2] else {
return None;
};
if primary.len() < 2 {
return None;
}
let host = primary[0].as_string().ok().flatten()?;
let port = u16::try_from(primary[1].as_integer()?).ok()?;
let node = match nodes
.iter()
.position(|(existing, existing_port)| existing == &host && *existing_port == port)
{
Some(index) => index,
None => {
nodes.push((host, port));
nodes.len() - 1
}
};
let first = fields[0].as_integer()?.clamp(0, i64::from(SLOT_COUNT) - 1) as usize;
let last = fields[1].as_integer()?.clamp(0, i64::from(SLOT_COUNT) - 1) as usize;
owners[first..=last].fill(Some(node));
}
Some(Self { owners, nodes })
}
fn owner(&self, slot: u16) -> Option<&(String, u16)> {
self.owners
.get(usize::from(slot))
.and_then(|node| node.and_then(|index| self.nodes.get(index)))
}
}
pub(crate) struct ClusterRouter {
options: ClientOptions,
seed: String,
connections: HashMap<String, Connection>,
topology: Topology,
queued: Vec<Vec<String>>,
reconnect_delay: Option<Arc<dyn Fn(u32) -> Option<Duration> + Send + Sync>>,
on_lost: Arc<Mutex<Option<Arc<dyn Fn(crate::client::ConnectionNotice) + Send + Sync>>>>,
on_restored: Arc<Mutex<Option<Arc<dyn Fn(crate::client::ConnectionNotice) + Send + Sync>>>>,
pinned: Option<String>,
multi_pending: bool,
in_multi: bool,
subscriber: Option<String>,
}
impl ClusterRouter {
pub(crate) fn discover(
mut seed: Connection,
options: ClientOptions,
) -> Result<Discovery, Error> {
let reply = seed.send(&["CLUSTER", "SLOTS"])?;
let Some(topology) = Topology::parse(&reply) else {
return Ok(Discovery::Standalone(seed));
};
let key = endpoint(seed.host.as_str(), seed.port);
Ok(Discovery::Cluster(Box::new(Self {
options,
seed: key.clone(),
connections: HashMap::from([(key, seed)]),
topology,
queued: Vec::new(),
reconnect_delay: None,
on_lost: Arc::new(Mutex::new(None)),
on_restored: Arc::new(Mutex::new(None)),
pinned: None,
multi_pending: false,
in_multi: false,
subscriber: None,
})))
}
pub(crate) fn node_count(&self) -> usize {
self.topology.nodes.len()
}
pub(crate) fn set_reconnect_delay(
&mut self,
delay: Arc<dyn Fn(u32) -> Option<Duration> + Send + Sync>,
) {
self.reconnect_delay = Some(delay);
}
pub(crate) fn share_hooks(
&mut self,
on_lost: Arc<Mutex<Option<Arc<dyn Fn(crate::client::ConnectionNotice) + Send + Sync>>>>,
on_restored: Arc<Mutex<Option<Arc<dyn Fn(crate::client::ConnectionNotice) + Send + Sync>>>>,
) {
self.on_lost = on_lost;
self.on_restored = on_restored;
}
pub(crate) fn reconnect_broken(&mut self) -> Result<(), Error> {
let nodes = self.topology.nodes.clone();
for (host, port) in nodes {
self.open_endpoint(&host, port, true)?;
}
Ok(())
}
pub(crate) fn subscription_endpoint(&self, arguments: &[&str]) -> (String, u16) {
let name = command_name(arguments);
let slot = if matches!(name.as_str(), "SSUBSCRIBE" | "SUNSUBSCRIBE") && arguments.len() > 1
{
hash_slot(arguments[1])
} else {
0
};
if let Some((host, port)) = self.topology.owner(slot).cloned() {
return (host, port);
}
parse_endpoint(&self.seed)
}
pub(crate) fn execute(&mut self, arguments: &[&str]) -> Result<RespValue, Error> {
let name = command_name(arguments);
if self.pinned.is_some() || self.multi_pending {
return self.transaction(&name, arguments);
}
match name.as_str() {
"MULTI" => {
self.multi_pending = true;
return Ok(RespValue::Simple("OK".to_owned()));
}
"WATCH" => return self.watch(arguments),
"DBSIZE" | "FLUSHDB" | "FLUSHALL" => return self.every_node(&name, arguments),
"SCRIPT" | "FUNCTION" if spreads_to_every_node(arguments) => {
return self.every_node(&name, arguments);
}
"SCAN" => return self.scan(arguments),
"PUBLISH" => {
let key = self.for_slot(0)?;
return self.send_routed(&key, arguments);
}
_ => {}
}
let keys = command_keys(&name, arguments);
if matches!(name.as_str(), "RENAME" | "RENAMENX") && keys.len() == 2 && spans_slots(&keys) {
return self.rename_across(name == "RENAMENX", keys[0], keys[1]);
}
if listed(SPLIT, &name) && spans_slots(&keys) {
return self.split(&name, arguments, &keys);
}
let target = self.target(&keys)?;
self.send_routed(&target, arguments)
}
pub(crate) fn execute_many(&mut self, commands: &[&[&str]]) -> Result<Vec<RespValue>, Error> {
if self.pinned.is_some()
|| self.multi_pending
|| commands
.iter()
.any(|command| listed(TRANSACTION, &command_name(command)))
{
let key = self.first_keyed_target(commands)?;
let owned: Vec<Vec<String>> = commands
.iter()
.map(|command| command.iter().map(|item| (*item).to_owned()).collect())
.collect();
return self.connection(&key)?.send_many_owned(&owned);
}
let mut groups: HashMap<String, Vec<usize>> = HashMap::new();
for (index, command) in commands.iter().enumerate() {
let name = command_name(command);
let keys = command_keys(&name, command);
let target = self.target(&keys)?;
groups.entry(target).or_default().push(index);
}
let mut replies = vec![RespValue::Null; commands.len()];
for (target, positions) in groups {
let batch: Vec<Vec<String>> = positions
.iter()
.map(|index| {
commands[*index]
.iter()
.map(|item| (*item).to_owned())
.collect()
})
.collect();
let values = self.connection(&target)?.send_many_owned(&batch)?;
for (offset, value) in values.into_iter().enumerate() {
replies[positions[offset]] = value;
}
}
for (index, reply) in replies.iter_mut().enumerate() {
if let Some((host, port)) = moved(reply) {
self.refresh()?;
let key = self.ensure(&host, port)?;
*reply = self.send_routed(&key, commands[index])?;
}
}
Ok(replies)
}
pub(crate) fn run_replies(
&mut self,
arguments: &[&str],
replies: usize,
) -> Result<Vec<RespValue>, Error> {
let name = command_name(arguments);
let slot = if matches!(name.as_str(), "SSUBSCRIBE" | "SUNSUBSCRIBE") && arguments.len() > 1
{
hash_slot(arguments[1])
} else {
0
};
let key = self.for_slot(slot)?;
self.subscriber = Some(key.clone());
let values = self.connection(&key)?.send_replies(arguments, replies)?;
if values.len() < replies {
return Err(Error::Protocol("missing subscribe confirmation".to_owned()));
}
Ok(values)
}
pub(crate) fn read_message(&mut self) -> Result<RespValue, Error> {
let key = self.subscriber.clone().unwrap_or_else(|| self.seed.clone());
self.connection(&key)?.read()
}
pub(crate) fn set_read_timeout(&mut self, timeout: Option<Duration>) -> Result<(), Error> {
for connection in self.connections.values_mut() {
connection.set_read_timeout(timeout)?;
}
Ok(())
}
fn transaction(&mut self, name: &str, arguments: &[&str]) -> Result<RespValue, Error> {
if let Some(pinned) = self.pinned.clone() {
let reply = self.connection(&pinned)?.send(arguments)?;
if name == "EXEC" || name == "DISCARD" || (name == "UNWATCH" && !self.in_multi) {
self.pinned = None;
self.in_multi = false;
} else if name == "MULTI" && !matches!(reply, RespValue::Error(_)) {
self.in_multi = true;
}
return Ok(reply);
}
match name {
"MULTI" => {
return Ok(RespValue::Error(
"ERR MULTI calls can not be nested".to_owned(),
));
}
"WATCH" => {
return Ok(RespValue::Error(
"ERR WATCH inside MULTI is not allowed".to_owned(),
));
}
"DISCARD" => {
self.reset_transaction();
return Ok(RespValue::Simple("OK".to_owned()));
}
"EXEC" => {
let key = self.seed.clone();
let mut commands = vec![vec!["MULTI".to_owned()]];
commands.extend(self.queued.iter().cloned());
commands.push(arguments.iter().map(|item| (*item).to_owned()).collect());
let replies = self.connection(&key)?.send_many_owned(&commands)?;
self.reset_transaction();
return Ok(replies.into_iter().next_back().unwrap_or(RespValue::Null));
}
_ => {}
}
let keys = command_keys(name, arguments);
if keys.is_empty() {
self.queued
.push(arguments.iter().map(|item| (*item).to_owned()).collect());
return Ok(RespValue::Simple("QUEUED".to_owned()));
}
let owner = self.for_slot(hash_slot(keys[0]))?;
let mut commands = vec![vec!["MULTI".to_owned()]];
commands.extend(self.queued.iter().cloned());
commands.push(arguments.iter().map(|item| (*item).to_owned()).collect());
let replies = self.connection(&owner)?.send_many_owned(&commands)?;
self.queued.clear();
self.multi_pending = false;
if matches!(replies.first(), Some(RespValue::Error(_))) {
return Ok(replies.into_iter().next().unwrap());
}
self.pinned = Some(owner);
self.in_multi = true;
Ok(replies.into_iter().next_back().unwrap_or(RespValue::Null))
}
fn reset_transaction(&mut self) {
self.queued.clear();
self.multi_pending = false;
self.in_multi = false;
self.pinned = None;
}
fn watch(&mut self, arguments: &[&str]) -> Result<RespValue, Error> {
let keys = command_keys("WATCH", arguments);
let owner = self.target(&keys)?;
let reply = self.connection(&owner)?.send(arguments)?;
if !matches!(reply, RespValue::Error(_)) {
self.pinned = Some(owner);
}
Ok(reply)
}
fn every_node(&mut self, name: &str, arguments: &[&str]) -> Result<RespValue, Error> {
let nodes: Vec<(String, u16)> = self.topology.nodes.clone();
let mut replies = Vec::new();
for (host, port) in nodes {
let key = self.ensure(&host, port)?;
let reply = self.connection(&key)?.send(arguments)?;
if let RespValue::Error(_) = reply {
return Ok(reply);
}
replies.push(reply);
}
if name == "DBSIZE" {
return Ok(RespValue::Integer(
replies.iter().filter_map(RespValue::as_integer).sum(),
));
}
Ok(replies.into_iter().next().unwrap_or(RespValue::Null))
}
fn scan(&mut self, arguments: &[&str]) -> Result<RespValue, Error> {
let nodes = self.topology.nodes.len() as u64;
let cursor = arguments
.get(1)
.and_then(|value| value.parse::<u64>().ok())
.ok_or_else(|| Error::Protocol("ERR invalid cursor".to_owned()))?;
let node = (cursor % nodes) as usize;
let mut forwarded: Vec<String> = arguments.iter().map(|item| (*item).to_owned()).collect();
forwarded[1] = (cursor / nodes).to_string();
let (host, port) = self.topology.nodes[node].clone();
let key = self.ensure(&host, port)?;
let reply = self.connection(&key)?.send_owned(&forwarded)?;
let RespValue::Array(items) = &reply else {
return Ok(reply);
};
if items.len() < 2 {
return Ok(reply);
}
let server_next = items[0]
.as_string()?
.unwrap_or_else(|| "0".to_owned())
.parse::<u64>()
.unwrap_or(0);
let client_next = if server_next != 0 {
server_next
.checked_mul(nodes)
.and_then(|value| value.checked_add(node as u64))
.ok_or_else(|| Error::Protocol("scan cursor overflow".to_owned()))?
} else if (node as u64) + 1 < nodes {
node as u64 + 1
} else {
0
};
Ok(RespValue::Array(vec![
RespValue::Bulk(client_next.to_string().into_bytes()),
items[1].clone(),
]))
}
fn split(&mut self, name: &str, arguments: &[&str], keys: &[&str]) -> Result<RespValue, Error> {
let pairs = name == "MSET";
let mut by_slot: HashMap<u16, Vec<usize>> = HashMap::new();
for (index, key) in keys.iter().enumerate() {
by_slot.entry(hash_slot(key)).or_default().push(index);
}
let mut parts = Vec::new();
for (slot, positions) in &by_slot {
let mut command = vec![arguments[0].to_owned()];
for position in positions {
if pairs {
command.push(arguments[1 + position * 2].to_owned());
command.push(arguments[2 + position * 2].to_owned());
} else {
command.push(keys[*position].to_owned());
}
}
parts.push((*slot, positions.clone(), command));
}
let mut replies = Vec::new();
for (slot, _, command) in &parts {
let key = self.for_slot(*slot)?;
let reply = self.connection(&key)?.send_owned(command)?;
if let Some((host, port)) = moved(&reply) {
self.refresh()?;
let redirected = self.ensure(&host, port)?;
replies.push(self.send_routed(&redirected, &borrowed(command))?);
} else if matches!(reply, RespValue::Error(_)) {
return Ok(reply);
} else {
replies.push(reply);
}
}
match name {
"MSET" => Ok(RespValue::Simple("OK".to_owned())),
"MGET" => {
let mut values = vec![RespValue::Null; keys.len()];
for (index, (_, positions, _)) in parts.iter().enumerate() {
let Some(items) = replies[index].as_array() else {
continue;
};
for (offset, item) in items.iter().enumerate() {
values[positions[offset]] = item.clone();
}
}
Ok(RespValue::Array(values))
}
_ => Ok(RespValue::Integer(
replies.iter().filter_map(RespValue::as_integer).sum(),
)),
}
}
fn rename_across(
&mut self,
exclusive: bool,
source: &str,
destination: &str,
) -> Result<RespValue, Error> {
if exclusive {
let exists = self.call(&["EXISTS", destination])?;
if exists.as_integer() == Some(1) {
return Ok(RespValue::Integer(0));
}
}
let kind = self
.call(&["TYPE", source])?
.as_string()?
.unwrap_or_default();
if kind == "none" {
return Ok(RespValue::Error("ERR no such key".to_owned()));
}
let ttl = self.call(&["PTTL", source])?.as_integer().unwrap_or(-1);
let moved = self.read_for_rename(&kind, source)?;
self.call(&["DEL", destination])?;
self.write_renamed(&kind, destination, &moved)?;
if ttl > 0 {
let millis = ttl.to_string();
self.call(&["PEXPIRE", destination, &millis])?;
}
self.call(&["DEL", source])?;
Ok(if exclusive {
RespValue::Integer(1)
} else {
RespValue::Simple("OK".to_owned())
})
}
fn read_for_rename(&mut self, kind: &str, source: &str) -> Result<MovedKey, Error> {
match kind {
"string" => Ok(MovedKey::String(required_text(
&self.call(&["GET", source])?,
)?)),
"list" => Ok(MovedKey::List(self.read_range("LRANGE", source, &[])?)),
"set" => Ok(MovedKey::Set(texts(&self.call(&["SMEMBERS", source])?)?)),
"hash" => Ok(MovedKey::Pairs(self.read_hash(source)?)),
"zset" => Ok(MovedKey::Pairs(self.read_zset(source)?)),
"stream" => Ok(MovedKey::Stream(self.read_stream(source)?)),
other => Err(Error::Protocol(format!(
"cross-slot rename does not move a {other} value"
))),
}
}
fn write_renamed(
&mut self,
kind: &str,
destination: &str,
moved: &MovedKey,
) -> Result<(), Error> {
match (kind, moved) {
("string", MovedKey::String(value)) => {
self.call(&["SET", destination, value])?;
}
("list", MovedKey::List(values)) => self.push_many("RPUSH", destination, values)?,
("set", MovedKey::Set(values)) => self.push_many("SADD", destination, values)?,
("hash", MovedKey::Pairs(pairs)) => {
self.push_pairs("HSET", destination, pairs, false)?
}
("zset", MovedKey::Pairs(pairs)) => {
self.push_pairs("ZADD", destination, pairs, true)?
}
("stream", MovedKey::Stream(entries)) => {
for (id, fields) in entries {
let mut arguments = vec!["XADD".to_owned(), destination.to_owned(), id.clone()];
for (field, value) in fields {
arguments.push(field.clone());
arguments.push(value.clone());
}
if arguments.len() > 3 {
self.call_owned(&arguments)?;
}
}
}
_ => {
return Err(Error::Protocol(
"cross-slot rename lost the value shape".to_owned(),
));
}
}
Ok(())
}
fn read_range(
&mut self,
command: &str,
source: &str,
extra: &[&str],
) -> Result<Vec<String>, Error> {
let mut start = 0_i64;
let mut values = Vec::new();
loop {
let from = start.to_string();
let to = (start + 199).to_string();
let mut arguments = vec![command, source, &from, &to];
arguments.extend(extra);
let page = texts(&self.call(&arguments)?)?;
let count = page.len() as i64;
values.extend(page);
if count < 200 {
break;
}
start += 200;
}
Ok(values)
}
fn read_hash(&mut self, source: &str) -> Result<Vec<(String, String)>, Error> {
let mut cursor = "0".to_owned();
let mut pairs = Vec::new();
loop {
let reply = self.call(&["HSCAN", source, &cursor, "COUNT", "64"])?;
let items = reply
.as_array()
.ok_or_else(|| Error::Protocol("HSCAN did not return a page".to_owned()))?;
let next = items
.first()
.and_then(|item| item.as_string().ok().flatten())
.unwrap_or_else(|| "0".to_owned());
let flat = items.get(1).map(texts).transpose()?.unwrap_or_default();
pairs.extend(pair_up(flat));
if next == "0" {
break;
}
cursor = next;
}
Ok(pairs)
}
fn read_zset(&mut self, source: &str) -> Result<Vec<(String, String)>, Error> {
let flat = self.read_range("ZRANGE", source, &["WITHSCORES"])?;
Ok(pair_up(flat))
}
fn read_stream(&mut self, source: &str) -> Result<Vec<(String, Vec<(String, String)>)>, Error> {
let mut start = "-".to_owned();
let mut entries = Vec::new();
loop {
let reply = self.call(&["XRANGE", source, &start, "+", "COUNT", "64"])?;
let page = reply
.as_array()
.map(|items| items.to_vec())
.unwrap_or_default();
let mut fresh = Vec::new();
for entry in page {
let Some(parts) = entry.as_array() else {
continue;
};
let id = parts
.first()
.and_then(|item| item.as_string().ok().flatten())
.unwrap_or_default();
if id.is_empty() || id == start {
continue;
}
let fields = parts.get(1).map(texts).transpose()?.unwrap_or_default();
fresh.push((id, pair_up(fields)));
}
if fresh.is_empty() {
break;
}
start = fresh.last().expect("fresh entry").0.clone();
let done = fresh.len() < 63;
entries.extend(fresh);
if done {
break;
}
}
Ok(entries)
}
fn push_many(&mut self, command: &str, key: &str, values: &[String]) -> Result<(), Error> {
for chunk in values.chunks(32) {
let mut arguments = vec![command.to_owned(), key.to_owned()];
arguments.extend(chunk.iter().cloned());
self.call_owned(&arguments)?;
}
Ok(())
}
fn push_pairs(
&mut self,
command: &str,
key: &str,
pairs: &[(String, String)],
score_first: bool,
) -> Result<(), Error> {
for chunk in pairs.chunks(32) {
let mut arguments = vec![command.to_owned(), key.to_owned()];
for (left, right) in chunk {
if score_first {
arguments.push(right.clone());
arguments.push(left.clone());
} else {
arguments.push(left.clone());
arguments.push(right.clone());
}
}
self.call_owned(&arguments)?;
}
Ok(())
}
fn call(&mut self, arguments: &[&str]) -> Result<RespValue, Error> {
let name = command_name(arguments);
let keys = command_keys(&name, arguments);
let target = self.target(&keys)?;
let reply = self.send_routed(&target, arguments)?;
match reply {
RespValue::Error(message) => Err(Error::Server(message)),
value => Ok(value),
}
}
fn call_owned(&mut self, arguments: &[String]) -> Result<RespValue, Error> {
let borrowed: Vec<&str> = arguments.iter().map(String::as_str).collect();
self.call(&borrowed)
}
fn send_routed(&mut self, key: &str, arguments: &[&str]) -> Result<RespValue, Error> {
let mut current = key.to_owned();
for redirect in 0..=MAX_REDIRECTS {
let reply = self.connection(¤t)?.send(arguments)?;
if redirect == MAX_REDIRECTS {
return Ok(reply);
}
let Some((host, port)) = moved(&reply) else {
return Ok(reply);
};
self.refresh()?;
current = self.ensure(&host, port)?;
}
unreachable!("redirect loop ends")
}
fn refresh(&mut self) -> Result<(), Error> {
let keys: Vec<String> = self.connections.keys().cloned().collect();
for key in keys {
let reply = self.connection(&key)?.send(&["CLUSTER", "SLOTS"])?;
if let Some(topology) = Topology::parse(&reply) {
self.topology = topology;
return Ok(());
}
}
Ok(())
}
fn target(&mut self, keys: &[&str]) -> Result<String, Error> {
if let Some(key) = keys.first() {
return self.for_slot(hash_slot(key));
}
Ok(self.seed.clone())
}
fn first_keyed_target(&mut self, commands: &[&[&str]]) -> Result<String, Error> {
for command in commands {
let keys = command_keys(&command_name(command), command);
if let Some(key) = keys.first() {
return self.for_slot(hash_slot(key));
}
}
Ok(self.seed.clone())
}
fn for_slot(&mut self, slot: u16) -> Result<String, Error> {
if let Some((host, port)) = self.topology.owner(slot).cloned() {
return self.ensure(&host, port);
}
Ok(self.seed.clone())
}
fn ensure(&mut self, host: &str, port: u16) -> Result<String, Error> {
self.open_endpoint(host, port, false)
}
fn open_endpoint(&mut self, host: &str, port: u16, force: bool) -> Result<String, Error> {
let key = endpoint(host, port);
let broken = self
.connections
.get(&key)
.is_some_and(Connection::is_broken);
if broken && !force && !self.options.reconnect {
return Err(closed_socket());
}
if broken {
self.notify_lost(host, port);
self.connections.remove(&key);
}
if !self.connections.contains_key(&key) {
let delay = self.reconnect_delay.clone();
let opened = Connection::open_with_retry(host, port, &self.options, delay.as_ref())?;
self.connections.insert(key.clone(), opened);
if broken {
self.notify_restored(host, port);
}
}
Ok(key)
}
fn notify_lost(&self, host: &str, port: u16) {
let notice = crate::client::ConnectionNotice {
host: host.to_owned(),
port,
subscriber: false,
};
if let Some(hook) = self.on_lost.lock().expect("connection hook").as_ref() {
hook(notice);
}
}
fn notify_restored(&self, host: &str, port: u16) {
let notice = crate::client::ConnectionNotice {
host: host.to_owned(),
port,
subscriber: false,
};
if let Some(hook) = self.on_restored.lock().expect("connection hook").as_ref() {
hook(notice);
}
}
fn connection(&mut self, key: &str) -> Result<&mut Connection, Error> {
self.connections
.get_mut(key)
.ok_or_else(|| Error::Protocol(format!("no connection for {key}")))
}
}
fn borrowed(command: &[String]) -> Vec<&str> {
command.iter().map(String::as_str).collect()
}
fn spreads_to_every_node(arguments: &[&str]) -> bool {
arguments.get(1).is_some_and(|subcommand| {
matches!(
subcommand.to_ascii_uppercase().as_str(),
"LOAD" | "FLUSH" | "DELETE"
)
})
}
enum MovedKey {
String(String),
List(Vec<String>),
Set(Vec<String>),
Pairs(Vec<(String, String)>),
Stream(Vec<(String, Vec<(String, String)>)>),
}
fn required_text(value: &RespValue) -> Result<String, Error> {
value
.as_string()?
.ok_or_else(|| Error::Protocol("expected a bulk string".to_owned()))
}
fn texts(value: &RespValue) -> Result<Vec<String>, Error> {
let Some(items) = value.as_array() else {
return Ok(Vec::new());
};
items.iter().map(required_text).collect()
}
fn pair_up(flat: Vec<String>) -> Vec<(String, String)> {
let mut pairs = Vec::new();
let mut items = flat.into_iter();
while let Some(left) = items.next() {
pairs.push((left, items.next().unwrap_or_default()));
}
pairs
}
fn spans_slots(keys: &[&str]) -> bool {
let Some(first) = keys.first() else {
return false;
};
let slot = hash_slot(first);
keys.iter().any(|key| hash_slot(key) != slot)
}
fn moved(reply: &RespValue) -> Option<(String, u16)> {
let RespValue::Error(text) = reply else {
return None;
};
let rest = text.strip_prefix("MOVED ")?;
let address = rest.split_whitespace().nth(1)?;
let (host, port) = address.rsplit_once(':')?;
Some((host.to_owned(), port.parse().ok()?))
}
fn closed_socket() -> Error {
Error::Io(std::io::Error::new(
std::io::ErrorKind::BrokenPipe,
"connection is closed",
))
}
impl Connection {
pub(crate) fn open_with_retry(
host: &str,
port: u16,
options: &crate::client::ClientOptions,
custom: Option<&Arc<dyn Fn(u32) -> Option<Duration> + Send + Sync>>,
) -> Result<Self, Error> {
let mut attempt = 0_u32;
loop {
attempt += 1;
match Self::open(host, port, options) {
Ok(connection) => return Ok(connection),
Err(error) => match reconnect_delay(options, custom, attempt) {
Some(delay) => std::thread::sleep(delay),
None => return Err(error),
},
}
}
}
}
fn reconnect_delay(
options: &crate::client::ClientOptions,
custom: Option<&Arc<dyn Fn(u32) -> Option<Duration> + Send + Sync>>,
attempt: u32,
) -> Option<Duration> {
if !options.reconnect {
return None;
}
if let Some(delay) = custom {
return delay(attempt);
}
if attempt >= options.max_reconnect_attempts {
return None;
}
let multiplier = 2_u32.saturating_pow(attempt.saturating_sub(1));
let delay = options.reconnect_base_delay.saturating_mul(multiplier);
Some(delay.min(options.reconnect_max_delay))
}
fn endpoint(host: &str, port: u16) -> String {
format!("{host}:{port}")
}
fn parse_endpoint(seed: &str) -> (String, u16) {
seed.rsplit_once(':')
.and_then(|(host, port)| port.parse().ok().map(|port| (host.to_owned(), port)))
.unwrap_or_else(|| (seed.to_owned(), 6379))
}
#[cfg(test)]
mod tests {
use super::{command_keys, hash_slot};
#[test]
fn hash_slot_matches_redis_and_honors_a_hash_tag() {
assert_eq!(hash_slot("123456789"), 12_739);
assert_eq!(hash_slot("somekey"), 11_058);
assert_eq!(hash_slot("{user1000}.following"), 3_443);
assert_eq!(hash_slot("user:{42}:name"), hash_slot("cart:{42}"));
}
#[test]
fn command_keys_follow_the_redis_positions() {
assert_eq!(command_keys("MGET", &["MGET", "a", "b"]), ["a", "b"]);
assert_eq!(
command_keys("MSET", &["MSET", "a", "1", "c", "2"]),
["a", "c"]
);
assert_eq!(
command_keys("EVAL", &["EVAL", "return 1", "2", "k1", "k2", "arg"]),
["k1", "k2"]
);
assert!(command_keys("PING", &["PING"]).is_empty());
assert!(command_keys("CHANGES", &["CHANGES", "START"]).is_empty());
assert_eq!(command_keys("LEASE", &["LEASE", "lock", "5000"]), ["lock"]);
assert_eq!(
command_keys("LIMIT", &["LIMIT", "login", "5", "60000"]),
["login"]
);
assert_eq!(command_keys("ONCE", &["ONCE", "pay", "60000"]), ["pay"]);
assert_eq!(
command_keys("IDEM", &["IDEM", "pay", "BEGIN", "fp", "owner", "60000"]),
["pay"]
);
assert_eq!(
command_keys("XADD", &["XADD", "s", "DELAY", "5", "*", "f", "v"]),
["s"]
);
assert!(command_keys("PUBLISH", &["PUBLISH", "chan", "hi"]).is_empty());
assert_eq!(
command_keys("SPUBLISH", &["SPUBLISH", "chan", "hi"]),
["chan"]
);
}
}