use std::net::{IpAddr, SocketAddr};
use std::sync::{Arc, Mutex};
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,
pub reconnect: bool,
pub max_reconnect_attempts: u32,
pub reconnect_base_delay: Duration,
pub reconnect_max_delay: Duration,
}
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,
reconnect: true,
max_reconnect_attempts: 8,
reconnect_base_delay: Duration::from_millis(100),
reconnect_max_delay: Duration::from_secs(2),
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ConnectionNotice {
pub host: String,
pub port: u16,
pub subscriber: bool,
}
pub struct Client {
transport: Transport,
database: u32,
options: ClientOptions,
dedicated_subscriber: bool,
pubsub: Option<Box<Client>>,
on_lost: Arc<Mutex<Option<Arc<dyn Fn(ConnectionNotice) + Send + Sync>>>>,
on_restored: Arc<Mutex<Option<Arc<dyn Fn(ConnectionNotice) + Send + Sync>>>>,
reconnect_delay: Option<Arc<dyn Fn(u32) -> Option<Duration> + Send + Sync>>,
subscriptions: Vec<Vec<String>>,
read_timeout: Option<Option<Duration>>,
}
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(mut router) => {
let on_lost = Arc::new(Mutex::new(None));
let on_restored = Arc::new(Mutex::new(None));
router.share_hooks(Arc::clone(&on_lost), Arc::clone(&on_restored));
Ok(Self {
transport: Transport::Cluster(router),
database: 0,
options,
dedicated_subscriber: false,
pubsub: None,
read_timeout: None,
on_lost,
on_restored,
reconnect_delay: None,
subscriptions: Vec::new(),
})
}
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> {
self.open_pubsub(&["SUBSCRIBE"])
}
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> {
self.read_timeout = Some(timeout);
match &mut self.transport {
Transport::Single(connection) => connection.set_read_timeout(timeout)?,
Transport::Cluster(router) => router.set_read_timeout(timeout)?,
}
if let Some(pubsub) = &mut self.pubsub {
pubsub.set_read_timeout(timeout)?;
}
Ok(())
}
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> {
if let Some(pubsub) = &mut self.pubsub {
return pubsub.read_message();
}
self.ensure_open()?;
let reply = match &mut self.transport {
Transport::Single(connection) => connection.read()?,
Transport::Cluster(router) => router.read_message()?,
};
throw_error(reply)
}
pub fn on_connection_lost<F>(&mut self, hook: F)
where
F: Fn(ConnectionNotice) + Send + Sync + 'static,
{
*self.on_lost.lock().expect("connection hook") = Some(Arc::new(hook));
}
pub fn on_connection_restored<F>(&mut self, hook: F)
where
F: Fn(ConnectionNotice) + Send + Sync + 'static,
{
*self.on_restored.lock().expect("connection hook") = Some(Arc::new(hook));
}
pub fn set_reconnect_delay<F>(&mut self, delay: F)
where
F: Fn(u32) -> Option<Duration> + Send + Sync + 'static,
{
self.reconnect_delay = Some(Arc::new(delay));
if let Transport::Cluster(router) = &mut self.transport {
router.set_reconnect_delay(Arc::clone(self.reconnect_delay.as_ref().unwrap()));
}
}
pub fn reconnect(&mut self) -> Result<(), Error> {
if let Transport::Cluster(router) = &mut self.transport {
return router.reconnect_broken();
}
if !self.transport_is_broken() {
return Ok(());
}
self.reopen()
}
pub fn execute(&mut self, arguments: &[&str]) -> Result<RespValue, Error> {
self.ensure_open()?;
if is_subscription_command(arguments) && !self.dedicated_subscriber {
return self.ensure_pubsub(arguments)?.execute(arguments);
}
let reply = match &mut self.transport {
Transport::Single(connection) => connection.send(arguments)?,
Transport::Cluster(router) => router.execute(arguments)?,
};
let reply = throw_error(reply)?;
self.remember_subscription(arguments);
Ok(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();
if is_subscription_command(&borrowed) && !self.dedicated_subscriber {
return self.ensure_pubsub(&borrowed)?.run_replies(owned, replies);
}
self.ensure_open()?;
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())?;
}
self.remember_subscription(&borrowed);
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)
}
fn ensure_pubsub(&mut self, arguments: &[&str]) -> Result<&mut Client, Error> {
if self.pubsub.is_none() {
self.pubsub = Some(Box::new(self.open_pubsub(arguments)?));
}
Ok(self.pubsub.as_mut().unwrap())
}
fn open_pubsub(&self, arguments: &[&str]) -> Result<Client, Error> {
let mut options = self.options.clone();
options.discover_cluster = false;
if let Transport::Cluster(router) = &self.transport {
let (host, port) = router.subscription_endpoint(arguments);
options.host = host;
options.port = port;
}
let mut client = Client::connect_with(options)?;
client.dedicated_subscriber = true;
client.on_lost = Arc::clone(&self.on_lost);
client.on_restored = Arc::clone(&self.on_restored);
client.reconnect_delay = self.reconnect_delay.clone();
if let Some(timeout) = self.read_timeout {
client.set_read_timeout(timeout)?;
}
Ok(client)
}
fn transport_is_broken(&self) -> bool {
match &self.transport {
Transport::Single(connection) => connection.is_broken(),
Transport::Cluster(_) => false,
}
}
fn ensure_open(&mut self) -> Result<(), Error> {
let broken = self.transport_is_broken();
if !broken {
return Ok(());
}
if !self.options.reconnect {
return Err(Error::Io(std::io::Error::new(
std::io::ErrorKind::BrokenPipe,
"connection is closed",
)));
}
self.reopen()
}
fn reopen(&mut self) -> Result<(), Error> {
let host = self.options.host.clone();
let port = self.options.port;
self.notify_lost(host.clone(), port);
let mut attempt = 0_u32;
let mut opened = loop {
attempt += 1;
match crate::cluster::Connection::open(&host, port, &self.options) {
Ok(connection) => break connection,
Err(error) => match self.next_delay(attempt) {
Some(delay) => std::thread::sleep(delay),
None => return Err(error),
},
}
};
if self.database != 0 {
let index = self.database.to_string();
throw_error(opened.send(&["SELECT", &index])?)?;
}
for command in &self.subscriptions {
let borrowed: Vec<&str> = command.iter().map(String::as_str).collect();
throw_error(opened.send(&borrowed)?)?;
}
self.transport = Transport::Single(opened);
self.notify_restored(host, port);
Ok(())
}
fn next_delay(&self, attempt: u32) -> Option<Duration> {
if !self.options.reconnect {
return None;
}
if let Some(delay) = &self.reconnect_delay {
return delay(attempt);
}
if attempt >= self.options.max_reconnect_attempts {
return None;
}
let multiplier = 2_u32.saturating_pow(attempt.saturating_sub(1));
let delay = self.options.reconnect_base_delay.saturating_mul(multiplier);
Some(delay.min(self.options.reconnect_max_delay))
}
fn notify_lost(&self, host: String, port: u16) {
let notice = ConnectionNotice {
host,
port,
subscriber: self.dedicated_subscriber,
};
if let Some(hook) = self.on_lost.lock().expect("connection hook").as_ref() {
hook(notice);
}
}
fn notify_restored(&self, host: String, port: u16) {
let notice = ConnectionNotice {
host,
port,
subscriber: self.dedicated_subscriber,
};
if let Some(hook) = self.on_restored.lock().expect("connection hook").as_ref() {
hook(notice);
}
}
fn remember_subscription(&mut self, arguments: &[&str]) {
if !self.dedicated_subscriber || arguments.is_empty() {
return;
}
let command = arguments[0].to_ascii_uppercase();
let (subscribe, dropping) = match command.as_str() {
"SUBSCRIBE" => ("SUBSCRIBE", false),
"UNSUBSCRIBE" => ("SUBSCRIBE", true),
"PSUBSCRIBE" => ("PSUBSCRIBE", false),
"PUNSUBSCRIBE" => ("PSUBSCRIBE", true),
"SSUBSCRIBE" => ("SSUBSCRIBE", false),
"SUNSUBSCRIBE" => ("SSUBSCRIBE", true),
_ => return,
};
if dropping && arguments.len() <= 1 {
self.subscriptions
.retain(|item| item.first().map(String::as_str) != Some(subscribe));
return;
}
if dropping {
for channel in &arguments[1..] {
self.subscriptions.retain(|item| {
item.first().map(String::as_str) != Some(subscribe)
|| item.get(1).map(String::as_str) != Some(*channel)
});
}
return;
}
for channel in &arguments[1..] {
let entry = vec![subscribe.to_owned(), (*channel).to_owned()];
if !self.subscriptions.iter().any(|item| item == &entry) {
self.subscriptions.push(entry);
}
}
}
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,
dedicated_subscriber: false,
pubsub: None,
read_timeout: None,
on_lost: Arc::new(Mutex::new(None)),
on_restored: Arc::new(Mutex::new(None)),
reconnect_delay: None,
subscriptions: Vec::new(),
})
}
fn is_subscription_command(arguments: &[&str]) -> bool {
arguments.first().is_some_and(|name| {
matches!(
name.to_ascii_uppercase().as_str(),
"SUBSCRIBE"
| "UNSUBSCRIBE"
| "PSUBSCRIBE"
| "PUNSUBSCRIBE"
| "SSUBSCRIBE"
| "SUNSUBSCRIBE"
)
})
}
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 bound(stream: TcpStream) -> TcpStream {
stream
.set_read_timeout(Some(Duration::from_secs(5)))
.unwrap();
stream
.set_write_timeout(Some(Duration::from_secs(5)))
.unwrap();
stream
}
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 || bound(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 || bound(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 || bound(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 || bound(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 || bound(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 subscribe_opens_a_second_socket_and_leaves_ping_on_the_first() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let port = listener.local_addr().unwrap().port();
let connecting = thread::spawn(move || Client::connect("127.0.0.1", port).unwrap());
let mut command = bound(listener.accept().unwrap().0);
reply_standalone_cluster(&mut command);
let mut client = connecting.join().unwrap();
let working = thread::spawn(move || {
let confirms = client.subscribe(&["news"]).unwrap();
let pong = client.ping().unwrap();
(confirms.len(), pong)
});
let mut subscriber = bound(listener.accept().unwrap().0);
let expected = b"*2\r\n$9\r\nSUBSCRIBE\r\n$4\r\nnews\r\n";
let mut got = vec![0_u8; expected.len()];
subscriber.read_exact(&mut got).unwrap();
assert_eq!(got, expected);
subscriber
.write_all(b"*3\r\n$9\r\nsubscribe\r\n$4\r\nnews\r\n:1\r\n")
.unwrap();
let mut ping = [0_u8; 14];
command.read_exact(&mut ping).unwrap();
assert_eq!(&ping, b"*1\r\n$4\r\nPING\r\n");
command.write_all(b"+PONG\r\n").unwrap();
let (confirms, pong) = working.join().unwrap();
assert_eq!(confirms, 1);
assert_eq!(pong, "PONG");
}
#[test]
fn explicit_subscriber_uses_one_extra_socket() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let port = listener.local_addr().unwrap().port();
let connecting = thread::spawn(move || Client::connect("127.0.0.1", port).unwrap());
let mut command = bound(listener.accept().unwrap().0);
reply_standalone_cluster(&mut command);
let client = connecting.join().unwrap();
let working = thread::spawn(move || {
let mut subscriber = client.subscriber().unwrap();
subscriber.subscribe(&["news"]).unwrap();
});
let mut subscriber = bound(listener.accept().unwrap().0);
let expected = b"*2\r\n$9\r\nSUBSCRIBE\r\n$4\r\nnews\r\n";
let mut got = vec![0_u8; expected.len()];
subscriber.read_exact(&mut got).unwrap();
assert_eq!(got, expected);
subscriber
.write_all(b"*3\r\n$9\r\nsubscribe\r\n$4\r\nnews\r\n:1\r\n")
.unwrap();
working.join().unwrap();
listener.set_nonblocking(true).unwrap();
assert!(listener.accept().is_err());
}
#[test]
fn the_command_after_a_dropped_socket_reconnects() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let port = listener.local_addr().unwrap().port();
let connecting = thread::spawn(move || {
Client::connect_with(ClientOptions {
host: "127.0.0.1".to_owned(),
port,
discover_cluster: false,
reconnect_base_delay: std::time::Duration::ZERO,
reconnect_max_delay: std::time::Duration::ZERO,
..ClientOptions::default()
})
.unwrap()
});
let accepted = bound(listener.accept().unwrap().0);
drop(accepted);
let mut client = connecting.join().unwrap();
let failed = thread::spawn(move || {
assert!(client.ping().is_err());
client.ping().unwrap()
});
let mut again = bound(listener.accept().unwrap().0);
let mut got = [0_u8; 14];
again.read_exact(&mut got).unwrap();
assert_eq!(&got, b"*1\r\n$4\r\nPING\r\n");
again.write_all(b"+PONG\r\n").unwrap();
assert_eq!(failed.join().unwrap(), "PONG");
}
#[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);
}
}