use std::io::{Error as IoError, ErrorKind};
use crate::prelude::*;
#[derive(Debug, Clone)]
enum ListpackEntry {
String(Arc<str>),
Integer(i64),
}
impl ListpackEntry {
fn to_arc_str(&self) -> Arc<str> {
match self {
Self::String(s) => s.clone(),
Self::Integer(i) => i.to_string().into(),
}
}
}
fn parse(data: &[u8]) -> Result<Vec<ListpackEntry>> {
if data.len() < 7 {
return Ok(Vec::new());
}
let num_elements = u16::from_le_bytes([data[4], data[5]]) as usize;
let mut entries = Vec::with_capacity(num_elements);
let mut pos = 6;
while pos < data.len() && data[pos] != super::symbols::LISTPACK_END {
let (entry, consumed) = parse_entry(data, pos)?;
entries.push(entry);
pos += consumed;
}
Ok(entries)
}
pub fn parse_hash(data: &[u8]) -> Result<HashMap<Value, Value>> {
let entries = parse(data)?;
let mut map = HashMap::with_capacity(entries.len() / 2);
let mut iter = entries.into_iter();
while let Some(key) = iter.next() {
let value = iter.next().ok_or_else(|| {
IoError::new(
ErrorKind::InvalidData,
"Listpack hash has odd number of entries",
)
})?;
map.insert(
Value::String(key.to_arc_str()),
Value::String(value.to_arc_str()),
);
}
Ok(map)
}
fn parse_entry(data: &[u8], pos: usize) -> Result<(ListpackEntry, usize)> {
if pos >= data.len() {
return Err(IoError::new(ErrorKind::InvalidData, "Listpack entry out of bounds").into());
}
let first_byte = data[pos];
if first_byte <= 0x7F {
let backlen = backlen_size(1);
return Ok((ListpackEntry::Integer(i64::from(first_byte)), 1 + backlen));
}
if first_byte & 0xC0 == 0x80 {
return parse_6bit_string(data, pos, first_byte);
}
if first_byte & 0xE0 == 0xC0 {
return parse_13bit_int(data, pos, first_byte);
}
parse_special_encoding(data, pos, first_byte)
}
fn parse_6bit_string(data: &[u8], pos: usize, first_byte: u8) -> Result<(ListpackEntry, usize)> {
let str_len = (first_byte & 0x3F) as usize;
if pos + 1 + str_len > data.len() {
return Err(IoError::new(ErrorKind::InvalidData, "Listpack string too short").into());
}
let s = String::from_utf8_lossy(&data[pos + 1..pos + 1 + str_len]);
let entry_len = 1 + str_len;
let backlen = backlen_size(entry_len);
Ok((
ListpackEntry::String(s.into_owned().into()),
entry_len + backlen,
))
}
#[allow(clippy::cast_possible_wrap)]
fn parse_13bit_int(data: &[u8], pos: usize, first_byte: u8) -> Result<(ListpackEntry, usize)> {
if pos + 2 > data.len() {
return Err(IoError::new(ErrorKind::InvalidData, "Listpack int13 too short").into());
}
let val = (i16::from(first_byte & 0x1F) << 8) | i16::from(data[pos + 1]);
let val = if val & 0x1000 != 0 {
val | (0xE000_u16 as i16)
} else {
val
};
let backlen = backlen_size(2);
Ok((ListpackEntry::Integer(i64::from(val)), 2 + backlen))
}
fn parse_special_encoding(
data: &[u8],
pos: usize,
first_byte: u8,
) -> Result<(ListpackEntry, usize)> {
if first_byte & 0xF0 == 0xE0 {
return parse_12bit_string(data, pos, first_byte);
}
match first_byte {
0xF0 => parse_32bit_string(data, pos),
0xF1 => parse_int16(data, pos),
0xF2 => parse_int24(data, pos),
0xF3 => parse_int32(data, pos),
0xF4 => parse_int64(data, pos),
_ => Err(IoError::new(
ErrorKind::InvalidData,
format!("Unknown listpack encoding: {first_byte:#04x}"),
)
.into()),
}
}
fn parse_12bit_string(data: &[u8], pos: usize, first_byte: u8) -> Result<(ListpackEntry, usize)> {
if pos + 1 >= data.len() {
return Err(IoError::new(
ErrorKind::InvalidData,
"Listpack 12-bit string header short",
)
.into());
}
let str_len = ((usize::from(first_byte & 0x0F)) << 8) | (data[pos + 1] as usize);
if pos + 2 + str_len > data.len() {
return Err(
IoError::new(ErrorKind::InvalidData, "Listpack 12-bit string too short").into(),
);
}
let s = String::from_utf8_lossy(&data[pos + 2..pos + 2 + str_len]);
let entry_len = 2 + str_len;
let backlen = backlen_size(entry_len);
Ok((
ListpackEntry::String(s.into_owned().into()),
entry_len + backlen,
))
}
fn parse_32bit_string(data: &[u8], pos: usize) -> Result<(ListpackEntry, usize)> {
if pos + 5 > data.len() {
return Err(IoError::new(
ErrorKind::InvalidData,
"Listpack 32-bit string header short",
)
.into());
}
let str_len =
u32::from_le_bytes([data[pos + 1], data[pos + 2], data[pos + 3], data[pos + 4]]) as usize;
if pos + 5 + str_len > data.len() {
return Err(
IoError::new(ErrorKind::InvalidData, "Listpack 32-bit string too short").into(),
);
}
let s = String::from_utf8_lossy(&data[pos + 5..pos + 5 + str_len]);
let entry_len = 5 + str_len;
let backlen = backlen_size(entry_len);
Ok((
ListpackEntry::String(s.into_owned().into()),
entry_len + backlen,
))
}
fn parse_int16(data: &[u8], pos: usize) -> Result<(ListpackEntry, usize)> {
if pos + 3 > data.len() {
return Err(IoError::new(ErrorKind::InvalidData, "Listpack int16 too short").into());
}
let val = i16::from_le_bytes([data[pos + 1], data[pos + 2]]);
let backlen = backlen_size(3);
Ok((ListpackEntry::Integer(i64::from(val)), 3 + backlen))
}
#[allow(clippy::cast_possible_wrap)]
fn parse_int24(data: &[u8], pos: usize) -> Result<(ListpackEntry, usize)> {
if pos + 4 > data.len() {
return Err(IoError::new(ErrorKind::InvalidData, "Listpack int24 too short").into());
}
let val = i32::from_le_bytes([data[pos + 1], data[pos + 2], data[pos + 3], 0]);
let val = if val & 0x0080_0000 != 0 {
val | (0xFF00_0000_u32 as i32)
} else {
val
};
let backlen = backlen_size(4);
Ok((ListpackEntry::Integer(i64::from(val)), 4 + backlen))
}
fn parse_int32(data: &[u8], pos: usize) -> Result<(ListpackEntry, usize)> {
if pos + 5 > data.len() {
return Err(IoError::new(ErrorKind::InvalidData, "Listpack int32 too short").into());
}
let val = i32::from_le_bytes([data[pos + 1], data[pos + 2], data[pos + 3], data[pos + 4]]);
let backlen = backlen_size(5);
Ok((ListpackEntry::Integer(i64::from(val)), 5 + backlen))
}
fn parse_int64(data: &[u8], pos: usize) -> Result<(ListpackEntry, usize)> {
if pos + 9 > data.len() {
return Err(IoError::new(ErrorKind::InvalidData, "Listpack int64 too short").into());
}
let val = i64::from_le_bytes([
data[pos + 1],
data[pos + 2],
data[pos + 3],
data[pos + 4],
data[pos + 5],
data[pos + 6],
data[pos + 7],
data[pos + 8],
]);
let backlen = backlen_size(9);
Ok((ListpackEntry::Integer(val), 9 + backlen))
}
const fn backlen_size(entry_len: usize) -> usize {
if entry_len <= 127 {
1
} else if entry_len <= 16383 {
2
} else if entry_len <= 2_097_151 {
3
} else if entry_len <= 268_435_455 {
4
} else {
5
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::parsers::rdb;
fn make_listpack(entries: &[u8]) -> Vec<u8> {
let total_len = 7 + entries.len();
let mut data = Vec::with_capacity(total_len);
#[allow(clippy::cast_possible_truncation)]
data.extend_from_slice(&(total_len as u32).to_le_bytes());
data.extend_from_slice(&[0u8; 2]);
data.extend_from_slice(entries);
data.push(rdb::symbols::LISTPACK_END);
data
}
#[test]
fn test_parse_empty() {
let data = make_listpack(&[]);
let entries = parse(&data).unwrap();
assert!(entries.is_empty());
}
#[test]
fn test_parse_small_int() {
let data = make_listpack(&[42, 1]);
let entries = parse(&data).unwrap();
assert_eq!(entries.len(), 1);
assert!(matches!(&entries[0], ListpackEntry::Integer(42)));
}
#[test]
fn test_parse_6bit_string() {
let data = make_listpack(&[0x82, b'h', b'i', 3]);
let entries = parse(&data).unwrap();
assert_eq!(entries.len(), 1);
assert!(matches!(&entries[0], ListpackEntry::String(s) if s.as_ref() == "hi"));
}
#[test]
fn test_parse_hash() {
let data = make_listpack(&[
0x84, b'n', b'a', b'm', b'e', 5, 0x85, b'p', b'x', b's', b'e', b'u', 6, ]);
let map = parse_hash(&data).unwrap();
assert_eq!(map.len(), 1);
let key = Value::String("name".into());
assert!(map.contains_key(&key));
assert_eq!(map.get(&key), Some(&Value::String("pxseu".into())));
}
#[test]
fn test_parse_int16() {
let data = make_listpack(&[0xF1, 0xE8, 0x03, 3]);
let entries = parse(&data).unwrap();
assert_eq!(entries.len(), 1);
assert!(matches!(&entries[0], ListpackEntry::Integer(1000)));
}
#[test]
fn test_parse_negative_int13() {
let data = make_listpack(&[0xDF, 0xFF, 2]);
let entries = parse(&data).unwrap();
assert_eq!(entries.len(), 1);
assert!(matches!(&entries[0], ListpackEntry::Integer(-1)));
}
#[test]
fn test_parse_12bit_string() {
let mut entry = vec![0xE0, 0x0C]; entry.extend_from_slice(b"hello world!");
entry.push(14); let data = make_listpack(&entry);
let entries = parse(&data).unwrap();
assert_eq!(entries.len(), 1);
assert!(matches!(&entries[0], ListpackEntry::String(s) if s.as_ref() == "hello world!"));
}
#[test]
fn test_parse_12bit_string_longer() {
let test_str = "a".repeat(300);
let mut entry = vec![0xE1, 0x2C]; entry.extend_from_slice(test_str.as_bytes());
entry.push(0x82); entry.push(0x02);
let data = make_listpack(&entry);
let entries = parse(&data).unwrap();
assert_eq!(entries.len(), 1);
assert!(matches!(&entries[0], ListpackEntry::String(s) if s.as_ref() == test_str));
}
}