ruvio-client 0.1.0

RESP2 client for the Ruvio key-value server
Documentation
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;

/// How `SET` treats a key that already exists or is missing.
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum SetMode {
    /// Write the key either way.
    #[default]
    Always,
    /// Write only when the key is absent (`NX`).
    IfNotExists,
    /// Write only when the key exists (`XX`).
    IfExists,
}

/// Where to connect, and an optional password sent as `AUTH`.
#[derive(Debug, Clone)]
pub struct ClientOptions {
    /// Server host. Defaults to loopback.
    pub host: String,
    /// RESP port. Defaults to 6379.
    pub port: u16,
    /// ACL user. When set together with [`Self::password`], the client sends `AUTH user password`.
    pub username: Option<String>,
    /// When set, the client sends `AUTH` before [`Client::connect_with`] returns.
    pub password: Option<String>,
    /// How long TCP connect may take.
    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),
        }
    }
}

/// One TCP connection to Ruvio.
///
/// The client speaks RESP2 and does not send `CLIENT SETINFO` or `HELLO`.
pub struct Client {
    stream: BufReader<TcpStream>,
}

impl Client {
    /// Opens a connection and sends nothing until the first command.
    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)
    }

    /// Opens a connection. When [`ClientOptions::password`] is set, sends `AUTH` before returning.
    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")
        })))
    }

    /// Sends `PING` and returns the simple-string reply.
    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()))
    }

    /// Sends `GET` and returns the raw bulk, or `None` when the key is missing.
    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:?}"))),
        }
    }

    /// Sends `GET` and decodes the bulk as UTF-8. `None` when the key is missing.
    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(),
        }
    }

    /// Sends `SET`. Returns false when `NX` or `XX` skips the write.
    pub fn set(&mut self, key: &str, value: &str) -> Result<bool, Error> {
        self.set_with(key, value, None, SetMode::Always)
    }

    /// Sends `SET` with an optional millisecond expiry and `NX` or `XX`.
    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),
        }
    }

    /// Sends `INCR`.
    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:?}"))),
        }
    }

    /// Sends `EXPIRE`. Returns false when the key does not exist.
    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:?}"))),
        }
    }

    /// Sends `DEL` and returns how many keys were removed.
    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:?}"))),
        }
    }

    /// Sends one command. A RESP error becomes [`Error::Server`] and the connection stays open.
    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())),
        }
    }

    /// Writes every command, then reads one reply each. Error replies stay as [`RespValue::Error`].
    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);
    }
}