use std::net::{IpAddr, SocketAddr};
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use crate::cluster::{ClusterRouter, Connection, Discovery, throw_error};
use crate::error::Error;
use crate::value::RespValue;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum SetMode {
#[default]
Always,
IfNotExists,
IfExists,
}
#[derive(Debug, Clone)]
pub struct ClientOptions {
pub host: String,
pub port: u16,
pub username: Option<String>,
pub password: Option<String>,
pub database: u32,
pub connect_timeout: Duration,
pub discover_cluster: bool,
}
impl Default for ClientOptions {
fn default() -> Self {
Self {
host: "127.0.0.1".to_owned(),
port: 6379,
username: None,
password: None,
database: 0,
connect_timeout: Duration::from_secs(5),
discover_cluster: true,
}
}
}
pub struct Client {
transport: Transport,
database: u32,
options: ClientOptions,
}
enum Transport {
Single(Connection),
Cluster(Box<ClusterRouter>),
}
impl Client {
pub fn connect(host: impl Into<String>, port: u16) -> Result<Self, Error> {
Self::connect_database(host, port, 0)
}
pub fn connect_database(
host: impl Into<String>,
port: u16,
database: u32,
) -> Result<Self, Error> {
Self::connect_with(ClientOptions {
host: host.into(),
port,
database,
..ClientOptions::default()
})
}
pub fn connect_ip(ip: IpAddr, port: u16, database: u32) -> Result<Self, Error> {
Self::connect_addr(SocketAddr::new(ip, port), database)
}
pub fn connect_addr(address: SocketAddr, database: u32) -> Result<Self, Error> {
Self::connect_database(address.ip().to_string(), address.port(), database)
}
pub fn connect_with(options: ClientOptions) -> Result<Self, Error> {
if options.host.is_empty() {
return Err(Error::Protocol("host is required".to_owned()));
}
let seed = Connection::open(&options.host, options.port, &options)?;
if !options.discover_cluster {
return finish_standalone(seed, options);
}
match ClusterRouter::discover(seed, options.clone())? {
Discovery::Cluster(_) if options.database != 0 => Err(Error::Protocol(
"SELECT is not supported in cluster mode".to_owned(),
)),
Discovery::Cluster(router) => Ok(Self {
transport: Transport::Cluster(router),
database: 0,
options,
}),
Discovery::Standalone(seed) => finish_standalone(seed, options),
}
}
pub fn is_cluster(&self) -> bool {
matches!(self.transport, Transport::Cluster(_))
}
pub fn shard_count(&self) -> usize {
match &self.transport {
Transport::Single(_) => 1,
Transport::Cluster(router) => router.node_count(),
}
}
pub fn ping(&mut self) -> Result<String, Error> {
read_ping(self.run(["PING"])?)
}
pub fn ping_message(&mut self, message: &str) -> Result<String, Error> {
read_ping(self.run(["PING", message])?)
}
pub fn database(&self) -> u32 {
self.database
}
pub fn subscriber(&self) -> Result<Self, Error> {
let mut options = self.options.clone();
options.discover_cluster = false;
Self::connect_with(options)
}
pub fn select(&mut self, database: u32) -> Result<(), Error> {
self.run(["SELECT".to_owned(), database.to_string()])?;
self.database = database;
Ok(())
}
pub fn into_database(mut self, database: u32) -> Result<Self, Error> {
self.select(database)?;
Ok(self)
}
pub fn set_read_timeout(&mut self, timeout: Option<Duration>) -> Result<(), Error> {
match &mut self.transport {
Transport::Single(connection) => connection.set_read_timeout(timeout),
Transport::Cluster(router) => router.set_read_timeout(timeout),
}
}
pub fn get(&mut self, key: &str) -> Result<Option<Vec<u8>>, Error> {
match self.execute(&["GET", key])? {
RespValue::Null => Ok(None),
RespValue::Bulk(bytes) => Ok(Some(bytes)),
other => Err(Error::Protocol(format!("GET returned {other:?}"))),
}
}
pub fn get_string(&mut self, key: &str) -> Result<Option<String>, Error> {
match self.get(key)? {
None => Ok(None),
Some(bytes) => RespValue::Bulk(bytes).as_string(),
}
}
pub fn set(&mut self, key: &str, value: &str) -> Result<bool, Error> {
self.set_with(key, value, None, SetMode::Always)
}
pub fn set_with(
&mut self,
key: &str,
value: &str,
expiry: Option<Duration>,
mode: SetMode,
) -> Result<bool, Error> {
let expiry = expiry.map(|ttl| ["PX".to_owned(), duration_millis(ttl).to_string()]);
let arguments = set_arguments(key, value, expiry.as_ref().map(|pair| &pair[..]), mode);
Ok(!matches!(self.run(arguments)?, RespValue::Null))
}
pub fn set_expires_at(
&mut self,
key: &str,
value: &str,
expires_at: SystemTime,
mode: SetMode,
) -> Result<bool, Error> {
let expiry = ["PXAT".to_owned(), unix_millis(expires_at)?.to_string()];
let arguments = set_arguments(key, value, Some(&expiry), mode);
Ok(!matches!(self.run(arguments)?, RespValue::Null))
}
pub fn set_keep_ttl(&mut self, key: &str, value: &str, mode: SetMode) -> Result<bool, Error> {
let expiry = ["KEEPTTL".to_owned()];
let arguments = set_arguments(key, value, Some(&expiry), mode);
Ok(!matches!(self.run(arguments)?, RespValue::Null))
}
pub fn set_and_get(
&mut self,
key: &str,
value: &str,
expiry: Option<Duration>,
mode: SetMode,
) -> Result<Option<String>, Error> {
let expiry = expiry.map(|ttl| ["PX".to_owned(), duration_millis(ttl).to_string()]);
let mut arguments = set_arguments(key, value, expiry.as_ref().map(|pair| &pair[..]), mode);
arguments.push("GET".to_owned());
self.bulk_or_null(arguments)
}
pub fn incr(&mut self, key: &str) -> Result<i64, Error> {
self.integer(["INCR", key])
}
pub fn expire(&mut self, key: &str, ttl: Duration) -> Result<bool, Error> {
if ttl.subsec_nanos() == 0 {
return self.flag([
"EXPIRE".to_owned(),
key.to_owned(),
duration_secs(ttl).to_string(),
]);
}
self.flag([
"PEXPIRE".to_owned(),
key.to_owned(),
duration_millis(ttl).to_string(),
])
}
pub fn del_key(&mut self, key: &str) -> Result<i64, Error> {
self.del(&[key])
}
pub fn del(&mut self, keys: &[&str]) -> Result<i64, Error> {
if keys.is_empty() {
return Err(Error::Protocol("DEL needs at least one key".to_owned()));
}
self.integer(join("DEL", keys))
}
pub fn read_message(&mut self) -> Result<RespValue, Error> {
let reply = match &mut self.transport {
Transport::Single(connection) => connection.read()?,
Transport::Cluster(router) => router.read_message()?,
};
throw_error(reply)
}
pub fn execute(&mut self, arguments: &[&str]) -> Result<RespValue, Error> {
let reply = match &mut self.transport {
Transport::Single(connection) => connection.send(arguments)?,
Transport::Cluster(router) => router.execute(arguments)?,
};
throw_error(reply)
}
pub fn execute_many(&mut self, commands: &[&[&str]]) -> Result<Vec<RespValue>, Error> {
if commands.is_empty() {
return Err(Error::Protocol(
"a pipeline needs at least one command".to_owned(),
));
}
match &mut self.transport {
Transport::Single(connection) => connection.send_many(commands),
Transport::Cluster(router) => router.execute_many(commands),
}
}
pub(crate) fn run_replies<I, S>(
&mut self,
arguments: I,
replies: usize,
) -> Result<Vec<RespValue>, Error>
where
I: IntoIterator<Item = S>,
S: AsRef<str>,
{
let owned: Vec<S> = arguments.into_iter().collect();
let borrowed: Vec<&str> = owned.iter().map(AsRef::as_ref).collect();
let values = match &mut self.transport {
Transport::Single(connection) => connection.send_replies(&borrowed, replies)?,
Transport::Cluster(router) => router.run_replies(&borrowed, replies)?,
};
if values.len() < replies {
return Err(Error::Protocol("missing subscribe confirmation".to_owned()));
}
for value in &values {
throw_error(value.clone())?;
}
Ok(values)
}
pub(crate) fn run<I, S>(&mut self, arguments: I) -> Result<RespValue, Error>
where
I: IntoIterator<Item = S>,
S: AsRef<str>,
{
let owned: Vec<S> = arguments.into_iter().collect();
let borrowed: Vec<&str> = owned.iter().map(AsRef::as_ref).collect();
self.execute(&borrowed)
}
pub(crate) fn integer<I, S>(&mut self, arguments: I) -> Result<i64, Error>
where
I: IntoIterator<Item = S>,
S: AsRef<str>,
{
let (command, reply) = self.run_named(arguments)?;
match reply {
RespValue::Integer(value) => Ok(value),
other => Err(unexpected(&command, &other)),
}
}
pub(crate) fn flag<I, S>(&mut self, arguments: I) -> Result<bool, Error>
where
I: IntoIterator<Item = S>,
S: AsRef<str>,
{
Ok(self.integer(arguments)? > 0)
}
pub(crate) fn ok<I, S>(&mut self, arguments: I) -> Result<(), Error>
where
I: IntoIterator<Item = S>,
S: AsRef<str>,
{
self.run(arguments)?;
Ok(())
}
pub(crate) fn text<I, S>(&mut self, arguments: I) -> Result<String, Error>
where
I: IntoIterator<Item = S>,
S: AsRef<str>,
{
let (command, reply) = self.run_named(arguments)?;
reply
.as_string()?
.ok_or_else(|| Error::Protocol(format!("{command} returned a null reply")))
}
pub(crate) fn bulk_or_null<I, S>(&mut self, arguments: I) -> Result<Option<String>, Error>
where
I: IntoIterator<Item = S>,
S: AsRef<str>,
{
self.run(arguments)?.as_string()
}
pub(crate) fn strings<I, S>(&mut self, arguments: I) -> Result<Vec<String>, Error>
where
I: IntoIterator<Item = S>,
S: AsRef<str>,
{
let (command, reply) = self.run_named(arguments)?;
match reply {
RespValue::Null => Ok(Vec::new()),
RespValue::Array(items) => items
.iter()
.map(|item| {
item.as_string()?
.ok_or_else(|| Error::Protocol(format!("{command} returned a null bulk")))
})
.collect(),
other => Err(unexpected(&command, &other)),
}
}
pub(crate) fn optional_strings<I, S>(
&mut self,
arguments: I,
) -> Result<Vec<Option<String>>, Error>
where
I: IntoIterator<Item = S>,
S: AsRef<str>,
{
let (command, reply) = self.run_named(arguments)?;
match reply {
RespValue::Null => Ok(Vec::new()),
RespValue::Array(items) => items.iter().map(RespValue::as_string).collect(),
other => Err(unexpected(&command, &other)),
}
}
pub(crate) fn run_named<I, S>(&mut self, arguments: I) -> Result<(String, RespValue), Error>
where
I: IntoIterator<Item = S>,
S: AsRef<str>,
{
let owned: Vec<S> = arguments.into_iter().collect();
let command = owned
.first()
.map(|first| first.as_ref().to_owned())
.unwrap_or_default();
let borrowed: Vec<&str> = owned.iter().map(AsRef::as_ref).collect();
Ok((command, self.execute(&borrowed)?))
}
}
fn read_ping(reply: RespValue) -> Result<String, Error> {
if let Some(items) = reply.as_array() {
if let Some(first) = items.first() {
if first
.as_string()?
.is_some_and(|kind| kind.eq_ignore_ascii_case("pong"))
{
if items.len() > 1 {
return Ok(items[1].as_string()?.unwrap_or_default());
}
return Ok("PONG".to_owned());
}
}
}
reply
.as_string()?
.ok_or_else(|| Error::Protocol("PING returned a null reply".to_owned()))
}
fn finish_standalone(mut seed: Connection, options: ClientOptions) -> Result<Client, Error> {
if options.database != 0 {
let index = options.database.to_string();
throw_error(seed.send(&["SELECT", &index])?)?;
}
Ok(Client {
transport: Transport::Single(seed),
database: options.database,
options,
})
}
pub(crate) fn unexpected(command: &str, reply: &RespValue) -> Error {
Error::Protocol(format!("{command} returned {reply:?}"))
}
pub(crate) fn join(command: &str, arguments: &[&str]) -> Vec<String> {
let mut values = Vec::with_capacity(arguments.len() + 1);
values.push(command.to_owned());
values.extend(arguments.iter().map(|argument| (*argument).to_owned()));
values
}
fn set_arguments(key: &str, value: &str, expiry: Option<&[String]>, mode: SetMode) -> Vec<String> {
let mut arguments = vec!["SET".to_owned(), key.to_owned(), value.to_owned()];
if let Some(expiry) = expiry {
arguments.extend_from_slice(expiry);
}
match mode {
SetMode::Always => {}
SetMode::IfNotExists => arguments.push("NX".to_owned()),
SetMode::IfExists => arguments.push("XX".to_owned()),
}
arguments
}
pub(crate) fn unix_millis(time: SystemTime) -> Result<u128, Error> {
time.duration_since(UNIX_EPOCH)
.map(|elapsed| elapsed.as_millis())
.map_err(|_| Error::Protocol("time is before the Unix epoch".to_owned()))
}
pub(crate) fn duration_millis(ttl: Duration) -> u128 {
let millis = ttl.as_millis();
if ttl.subsec_nanos().is_multiple_of(1_000_000) {
return millis;
}
millis.saturating_add(1)
}
pub(crate) fn duration_secs(ttl: Duration) -> u64 {
let seconds = ttl.as_secs();
if ttl.subsec_nanos() == 0 {
return seconds;
}
seconds.saturating_add(1)
}
#[cfg(test)]
mod tests {
use std::io::{Read, Write};
use std::net::{TcpListener, TcpStream};
use std::thread;
use std::time::Duration;
use super::{Client, ClientOptions, RespValue, read_ping};
fn reply_standalone_cluster(server: &mut TcpStream) {
let expected = b"*2\r\n$7\r\nCLUSTER\r\n$5\r\nSLOTS\r\n";
let mut got = vec![0_u8; expected.len()];
server.read_exact(&mut got).unwrap();
assert_eq!(got, expected);
server
.write_all(b"-ERR cluster mode is disabled\r\n")
.unwrap();
}
#[test]
fn connect_asks_for_cluster_slots_then_is_quiet() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let port = listener.local_addr().unwrap().port();
let accepted = thread::spawn(move || listener.accept().unwrap().0);
let connecting = thread::spawn(move || Client::connect("127.0.0.1", port).unwrap());
let mut server = accepted.join().unwrap();
reply_standalone_cluster(&mut server);
let mut client = connecting.join().unwrap();
assert!(!client.is_cluster());
thread::sleep(Duration::from_millis(50));
server.set_nonblocking(true).unwrap();
let mut peeked = [0_u8; 1];
assert!(server.peek(&mut peeked).is_err());
server.set_nonblocking(false).unwrap();
let ping = thread::spawn(move || client.ping().unwrap());
let mut got = [0_u8; 14];
server.read_exact(&mut got).unwrap();
assert_eq!(&got, b"*1\r\n$4\r\nPING\r\n");
server.write_all(b"+PONG\r\n").unwrap();
assert_eq!(ping.join().unwrap(), "PONG");
}
#[test]
fn discover_cluster_off_sends_nothing_until_the_first_command() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let port = listener.local_addr().unwrap().port();
let accepted = thread::spawn(move || listener.accept().unwrap().0);
let mut client = Client::connect_with(ClientOptions {
host: "127.0.0.1".to_owned(),
port,
discover_cluster: false,
..ClientOptions::default()
})
.unwrap();
let mut server = accepted.join().unwrap();
thread::sleep(Duration::from_millis(50));
server.set_nonblocking(true).unwrap();
let mut peeked = [0_u8; 1];
assert!(server.peek(&mut peeked).is_err());
server.set_nonblocking(false).unwrap();
let ping = thread::spawn(move || client.ping().unwrap());
let mut got = [0_u8; 14];
server.read_exact(&mut got).unwrap();
assert_eq!(&got, b"*1\r\n$4\r\nPING\r\n");
server.write_all(b"+PONG\r\n").unwrap();
assert_eq!(ping.join().unwrap(), "PONG");
}
#[test]
fn a_server_error_leaves_the_next_command_usable() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let port = listener.local_addr().unwrap().port();
let accepted = thread::spawn(move || listener.accept().unwrap().0);
let connecting = thread::spawn(move || Client::connect("127.0.0.1", port).unwrap());
let mut server = accepted.join().unwrap();
reply_standalone_cluster(&mut server);
let mut client = connecting.join().unwrap();
let error = thread::spawn(move || {
let error = client.get_string("missing").unwrap_err();
let pong = client.ping().unwrap();
(error.to_string(), pong)
});
read_some(&mut server);
server.write_all(b"-ERR no such key\r\n").unwrap();
read_some(&mut server);
server.write_all(b"+PONG\r\n").unwrap();
let (message, pong) = error.join().unwrap();
assert_eq!(message, "ERR no such key");
assert_eq!(pong, "PONG");
}
#[test]
fn password_is_sent_as_auth() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let port = listener.local_addr().unwrap().port();
let accepted = thread::spawn(move || listener.accept().unwrap().0);
let connecting = thread::spawn(move || {
Client::connect_with(ClientOptions {
host: "127.0.0.1".to_owned(),
port,
password: Some("secret".to_owned()),
..ClientOptions::default()
})
.unwrap()
});
let mut server = accepted.join().unwrap();
let expected = b"*2\r\n$4\r\nAUTH\r\n$6\r\nsecret\r\n";
let mut got = vec![0_u8; expected.len()];
server.read_exact(&mut got).unwrap();
assert_eq!(got, expected);
server.write_all(b"+OK\r\n").unwrap();
reply_standalone_cluster(&mut server);
connecting.join().unwrap();
}
#[test]
fn database_is_selected_after_auth() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let port = listener.local_addr().unwrap().port();
let accepted = thread::spawn(move || listener.accept().unwrap().0);
let connecting = thread::spawn(move || {
Client::connect_with(ClientOptions {
host: "127.0.0.1".to_owned(),
port,
password: Some("secret".to_owned()),
database: 3,
..ClientOptions::default()
})
.unwrap()
});
let mut server = accepted.join().unwrap();
let auth = b"*2\r\n$4\r\nAUTH\r\n$6\r\nsecret\r\n";
let select = b"*2\r\n$6\r\nSELECT\r\n$1\r\n3\r\n";
let mut got = vec![0_u8; auth.len()];
server.read_exact(&mut got).unwrap();
assert_eq!(got, auth);
server.write_all(b"+OK\r\n").unwrap();
reply_standalone_cluster(&mut server);
got = vec![0_u8; select.len()];
server.read_exact(&mut got).unwrap();
assert_eq!(got, select);
server.write_all(b"+OK\r\n").unwrap();
assert_eq!(connecting.join().unwrap().database(), 3);
}
#[test]
fn ping_reads_a_pubsub_array() {
let reply = RespValue::Array(vec![
RespValue::Bulk(b"pong".to_vec()),
RespValue::Bulk(b"hello".to_vec()),
]);
assert_eq!(read_ping(reply).unwrap(), "hello");
}
fn read_some(stream: &mut TcpStream) {
let mut buffer = [0_u8; 64];
let read = stream.read(&mut buffer).unwrap();
assert!(read > 0);
}
}