ruvio-client 0.1.0

RESP2 client for the Ruvio key-value server
Documentation
use std::io::{BufRead, Write};

use crate::error::Error;
use crate::value::RespValue;

const MAX_LINE_BYTES: usize = 1024 * 1024;
const MAX_BULK_BYTES: usize = 64 * 1024 * 1024;
const MAX_ARRAY_LENGTH: usize = 1024 * 1024;

pub(crate) fn encode(arguments: &[&[u8]]) -> Result<Vec<u8>, Error> {
    if arguments.is_empty() {
        return Err(Error::Protocol(
            "a command needs at least one argument".to_owned(),
        ));
    }

    let mut payload = Vec::new();

    write!(payload, "*{}\r\n", arguments.len())?;

    for argument in arguments {
        write!(payload, "${}\r\n", argument.len())?;
        payload.write_all(argument)?;
        payload.write_all(b"\r\n")?;
    }

    Ok(payload)
}

pub(crate) fn encode_strings(arguments: &[&str]) -> Result<Vec<u8>, Error> {
    let encoded: Vec<&[u8]> = arguments
        .iter()
        .map(|argument| argument.as_bytes())
        .collect();

    encode(&encoded)
}

pub(crate) fn read_value(reader: &mut impl BufRead) -> Result<RespValue, Error> {
    let mut prefix = [0_u8; 1];

    reader.read_exact(&mut prefix)?;

    match prefix[0] {
        b'+' => Ok(RespValue::Simple(read_text(reader)?)),
        b'-' => Ok(RespValue::Error(read_text(reader)?)),
        b':' => Ok(RespValue::Integer(parse_integer(
            &read_line(reader)?,
            "integer",
        )?)),
        b'$' => read_bulk(reader),
        b'*' => read_array(reader),
        other => Err(Error::Protocol(format!(
            "unexpected RESP prefix 0x{other:02x}"
        ))),
    }
}

fn read_bulk(reader: &mut impl BufRead) -> Result<RespValue, Error> {
    let length = parse_integer(&read_line(reader)?, "bulk length")?;

    if length == -1 {
        return Ok(RespValue::Null);
    }

    if length < 0 || length as usize > MAX_BULK_BYTES {
        return Err(Error::Protocol("bulk length is out of range".to_owned()));
    }

    let mut payload = vec![0_u8; length as usize];

    reader.read_exact(&mut payload)?;

    let mut trailer = [0_u8; 2];

    reader.read_exact(&mut trailer)?;

    if trailer != *b"\r\n" {
        return Err(Error::Protocol(
            "bulk payload was not terminated".to_owned(),
        ));
    }

    Ok(RespValue::Bulk(payload))
}

fn read_array(reader: &mut impl BufRead) -> Result<RespValue, Error> {
    let length = parse_integer(&read_line(reader)?, "array length")?;

    if length == -1 {
        return Ok(RespValue::Null);
    }

    if length < 0 || length as usize > MAX_ARRAY_LENGTH {
        return Err(Error::Protocol("array length is out of range".to_owned()));
    }

    let mut items = Vec::with_capacity(length as usize);

    for _ in 0..length {
        items.push(read_value(reader)?);
    }

    Ok(RespValue::Array(items))
}

fn read_text(reader: &mut impl BufRead) -> Result<String, Error> {
    let line = read_line(reader)?;

    String::from_utf8(line).map_err(|_| Error::Protocol("RESP line is not UTF-8".to_owned()))
}

fn read_line(reader: &mut impl BufRead) -> Result<Vec<u8>, Error> {
    let mut line = Vec::new();

    let read = reader.read_until(b'\n', &mut line)?;

    if read == 0 || line.last() != Some(&b'\n') {
        return Err(Error::Protocol("connection closed".to_owned()));
    }

    line.pop();

    if line.last() == Some(&b'\r') {
        line.pop();
    } else {
        return Err(Error::Protocol("expected CR before LF".to_owned()));
    }

    if line.len() > MAX_LINE_BYTES {
        return Err(Error::Protocol("RESP line is too long".to_owned()));
    }

    Ok(line)
}

fn parse_integer(bytes: &[u8], label: &str) -> Result<i64, Error> {
    let text =
        std::str::from_utf8(bytes).map_err(|_| Error::Protocol(format!("bad RESP {label}")))?;

    text.parse()
        .map_err(|_| Error::Protocol(format!("bad RESP {label}")))
}

#[cfg(test)]
mod tests {
    use std::io::Cursor;

    use super::{encode_strings, read_value};

    use crate::value::RespValue;

    #[test]
    fn encode_writes_a_resp2_array_of_bulk_strings() {
        let payload = encode_strings(&["GET", "key"]).expect("encode");

        assert_eq!(payload, b"*2\r\n$3\r\nGET\r\n$3\r\nkey\r\n");
    }

    #[test]
    fn reader_decodes_the_reply_kinds_ruvio_sends() {
        let payload = b"+OK\r\n-ERR nope\r\n:7\r\n$4\r\nping\r\n$-1\r\n*2\r\n$1\r\na\r\n:2\r\n";
        let mut reader = Cursor::new(&payload[..]);

        assert_eq!(
            read_value(&mut reader).unwrap(),
            RespValue::Simple("OK".to_owned())
        );
        assert_eq!(
            read_value(&mut reader).unwrap(),
            RespValue::Error("ERR nope".to_owned())
        );
        assert_eq!(read_value(&mut reader).unwrap(), RespValue::Integer(7));
        assert_eq!(
            read_value(&mut reader).unwrap(),
            RespValue::Bulk(b"ping".to_vec())
        );
        assert_eq!(read_value(&mut reader).unwrap(), RespValue::Null);

        assert_eq!(
            read_value(&mut reader).unwrap(),
            RespValue::Array(vec![RespValue::Bulk(b"a".to_vec()), RespValue::Integer(2)])
        );
    }

    #[test]
    fn reader_keeps_crlf_inside_a_bulk_string() {
        let mut framed = b"$4\r\n".to_vec();

        framed.extend_from_slice(b"a\r\nb");
        framed.extend_from_slice(b"\r\n");

        let mut reader = Cursor::new(framed);

        assert_eq!(
            read_value(&mut reader).unwrap(),
            RespValue::Bulk(b"a\r\nb".to_vec())
        );
    }
}