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())
);
}
}