ruvio-client 0.2.9

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

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 append_strings(payload: &mut Vec<u8>, arguments: &[&str]) -> Result<(), Error> {
    if arguments.is_empty() {
        return Err(Error::Protocol(
            "a command needs at least one argument".to_owned(),
        ));
    }

    payload.push(b'*');
    push_decimal(payload, arguments.len());
    payload.extend_from_slice(b"\r\n");

    for argument in arguments {
        append_bulk(payload, argument.as_bytes());
    }

    Ok(())
}

fn append_bulk(payload: &mut Vec<u8>, argument: &[u8]) {
    payload.push(b'$');
    push_decimal(payload, argument.len());
    payload.extend_from_slice(b"\r\n");
    payload.extend_from_slice(argument);
    payload.extend_from_slice(b"\r\n");
}

fn push_decimal(payload: &mut Vec<u8>, value: usize) {
    let mut digits = [0_u8; 20];
    let mut count = 0;
    let mut remaining = value;

    loop {
        digits[count] = b'0' + (remaining % 10) as u8;
        count += 1;
        remaining /= 10;

        if remaining == 0 {
            break;
        }
    }

    while count > 0 {
        count -= 1;
        payload.push(digits[count]);
    }
}

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::{append_strings, read_value};

    use crate::value::RespValue;

    #[test]
    fn encode_writes_a_resp2_array_of_bulk_strings() {
        let mut payload = Vec::new();

        append_strings(&mut payload, &["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())
        );
    }
}