use std::io::{BufReader, Write};
use std::net::{TcpStream, ToSocketAddrs};
use std::time::Duration;
use crate::codec::{encode_strings, read_value};
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 connect_timeout: Duration,
}
impl Default for ClientOptions {
fn default() -> Self {
Self {
host: "127.0.0.1".to_owned(),
port: 6379,
username: None,
password: None,
connect_timeout: Duration::from_secs(5),
}
}
}
pub struct Client {
stream: BufReader<TcpStream>,
}
impl Client {
pub fn connect(host: impl Into<String>, port: u16) -> Result<Self, Error> {
let mut options = ClientOptions::default();
options.host = host.into();
options.port = port;
Self::connect_with(options)
}
pub fn connect_with(options: ClientOptions) -> Result<Self, Error> {
if options.host.is_empty() {
return Err(Error::Protocol("host is required".to_owned()));
}
let mut last_error = None;
for address in (options.host.as_str(), options.port).to_socket_addrs()? {
match TcpStream::connect_timeout(&address, options.connect_timeout) {
Ok(stream) => {
stream.set_nodelay(true)?;
let mut client = Client {
stream: BufReader::new(stream),
};
if let Some(password) = options.password.clone() {
let reply = if let Some(username) = options.username.clone() {
client.execute(&["AUTH", &username, &password])
} else {
client.execute(&["AUTH", &password])
};
reply?;
}
return Ok(client);
}
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 fn ping(&mut self) -> Result<String, Error> {
let reply = self.execute(&["PING"])?;
reply
.as_string()?
.ok_or_else(|| Error::Protocol("PING returned a null reply".to_owned()))
}
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 mut arguments = vec!["SET", key, value];
let millis = expiry.map(|ttl| duration_millis(ttl).to_string());
if let Some(millis) = millis.as_deref() {
arguments.push("PX");
arguments.push(millis);
}
match mode {
SetMode::Always => {}
SetMode::IfNotExists => arguments.push("NX"),
SetMode::IfExists => arguments.push("XX"),
}
match self.execute(&arguments)? {
RespValue::Null => Ok(false),
_ => Ok(true),
}
}
pub fn incr(&mut self, key: &str) -> Result<i64, Error> {
match self.execute(&["INCR", key])? {
RespValue::Integer(value) => Ok(value),
other => Err(Error::Protocol(format!("INCR returned {other:?}"))),
}
}
pub fn expire(&mut self, key: &str, ttl: Duration) -> Result<bool, Error> {
let seconds = duration_secs(ttl).to_string();
match self.execute(&["EXPIRE", key, &seconds])? {
RespValue::Integer(value) => Ok(value > 0),
other => Err(Error::Protocol(format!("EXPIRE returned {other:?}"))),
}
}
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()));
}
let mut arguments = Vec::with_capacity(keys.len() + 1);
arguments.push("DEL");
arguments.extend_from_slice(keys);
match self.execute(&arguments)? {
RespValue::Integer(value) => Ok(value),
other => Err(Error::Protocol(format!("DEL returned {other:?}"))),
}
}
pub fn execute(&mut self, arguments: &[&str]) -> Result<RespValue, Error> {
let values = self.execute_many(&[arguments])?;
match values.into_iter().next() {
Some(RespValue::Error(message)) => Err(Error::Server(message)),
Some(value) => Ok(value),
None => Err(Error::Protocol("missing reply".to_owned())),
}
}
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(),
));
}
let mut payload = Vec::new();
for command in commands {
payload.extend(encode_strings(command)?);
}
self.stream.get_mut().write_all(&payload)?;
self.stream.get_mut().flush()?;
let mut values = Vec::with_capacity(commands.len());
for _ in 0..commands.len() {
values.push(read_value(&mut self.stream)?);
}
Ok(values)
}
}
fn duration_millis(ttl: Duration) -> u128 {
let millis = ttl.as_millis();
if ttl.subsec_nanos() % 1_000_000 == 0 {
return millis;
}
millis.saturating_add(1)
}
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};
#[test]
fn connect_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("127.0.0.1", port).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 mut client = Client::connect("127.0.0.1", port).unwrap();
let mut server = accepted.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();
connecting.join().unwrap();
}
fn read_some(stream: &mut TcpStream) {
let mut buffer = [0_u8; 64];
let read = stream.read(&mut buffer).unwrap();
assert!(read > 0);
}
}