use std::io::{Error as IoError, ErrorKind};
use crate::prelude::*;
#[derive(Debug, Clone)]
enum ZiplistEntry {
String(Arc<str>),
Integer(i64),
}
impl ZiplistEntry {
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<ZiplistEntry>> {
if data.len() < 11 {
return Ok(Vec::new());
}
let num_entries = u16::from_le_bytes([data[8], data[9]]) as usize;
let mut entries = Vec::with_capacity(num_entries);
let mut pos = 10;
while pos < data.len() && data[pos] != super::symbols::ZIPLIST_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,
"Ziplist 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<(ZiplistEntry, usize)> {
if pos >= data.len() {
return Err(IoError::new(ErrorKind::InvalidData, "Ziplist entry out of bounds").into());
}
let (prevlen_size, _prevlen) = parse_prevlen(data, pos)?;
let encoding_pos = pos + prevlen_size;
if encoding_pos >= data.len() {
return Err(IoError::new(ErrorKind::InvalidData, "Ziplist encoding out of bounds").into());
}
let encoding = data[encoding_pos];
let (entry, data_size) = parse_encoding(data, encoding_pos, encoding)?;
Ok((entry, prevlen_size + data_size))
}
fn parse_prevlen(data: &[u8], pos: usize) -> Result<(usize, u32)> {
if pos >= data.len() {
return Err(IoError::new(ErrorKind::InvalidData, "Prevlen out of bounds").into());
}
let first = data[pos];
if first < super::symbols::ZIPLIST_BIGLEN {
Ok((1, u32::from(first)))
} else {
if pos + 5 > data.len() {
return Err(
IoError::new(ErrorKind::InvalidData, "Prevlen 5-byte out of bounds").into(),
);
}
let prevlen =
u32::from_le_bytes([data[pos + 1], data[pos + 2], data[pos + 3], data[pos + 4]]);
Ok((5, prevlen))
}
}
fn parse_encoding(data: &[u8], pos: usize, encoding: u8) -> Result<(ZiplistEntry, usize)> {
let top_two = encoding & 0xC0;
match top_two {
super::symbols::ZIPLIST_STR_6BIT => {
let len = (encoding & 0x3F) as usize;
if pos + 1 + len > data.len() {
return Err(
IoError::new(ErrorKind::InvalidData, "Ziplist string too short").into(),
);
}
let s = String::from_utf8_lossy(&data[pos + 1..pos + 1 + len]);
Ok((ZiplistEntry::String(s.into_owned().into()), 1 + len))
}
super::symbols::ZIPLIST_STR_14BIT => {
if pos + 2 > data.len() {
return Err(IoError::new(
ErrorKind::InvalidData,
"Ziplist 14-bit length too short",
)
.into());
}
let len = (((encoding & 0x3F) as usize) << 8) | (data[pos + 1] as usize);
if pos + 2 + len > data.len() {
return Err(
IoError::new(ErrorKind::InvalidData, "Ziplist string too short").into(),
);
}
let s = String::from_utf8_lossy(&data[pos + 2..pos + 2 + len]);
Ok((ZiplistEntry::String(s.into_owned().into()), 2 + len))
}
super::symbols::ZIPLIST_STR_32BIT => {
if pos + 5 > data.len() {
return Err(IoError::new(
ErrorKind::InvalidData,
"Ziplist 32-bit length too short",
)
.into());
}
let len =
u32::from_le_bytes([data[pos + 1], data[pos + 2], data[pos + 3], data[pos + 4]])
as usize;
if pos + 5 + len > data.len() {
return Err(
IoError::new(ErrorKind::InvalidData, "Ziplist string too short").into(),
);
}
let s = String::from_utf8_lossy(&data[pos + 5..pos + 5 + len]);
Ok((ZiplistEntry::String(s.into_owned().into()), 5 + len))
}
_ => {
parse_integer_encoding(data, pos, encoding)
}
}
}
#[allow(clippy::cast_possible_wrap)]
fn parse_integer_encoding(data: &[u8], pos: usize, encoding: u8) -> Result<(ZiplistEntry, usize)> {
match encoding {
super::symbols::ZIPLIST_INT_16 => {
if pos + 3 > data.len() {
return Err(IoError::new(ErrorKind::InvalidData, "Ziplist int16 too short").into());
}
let val = i16::from_le_bytes([data[pos + 1], data[pos + 2]]);
Ok((ZiplistEntry::Integer(i64::from(val)), 3))
}
super::symbols::ZIPLIST_INT_32 => {
if pos + 5 > data.len() {
return Err(IoError::new(ErrorKind::InvalidData, "Ziplist int32 too short").into());
}
let val =
i32::from_le_bytes([data[pos + 1], data[pos + 2], data[pos + 3], data[pos + 4]]);
Ok((ZiplistEntry::Integer(i64::from(val)), 5))
}
super::symbols::ZIPLIST_INT_64 => {
if pos + 9 > data.len() {
return Err(IoError::new(ErrorKind::InvalidData, "Ziplist 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],
]);
Ok((ZiplistEntry::Integer(val), 9))
}
super::symbols::ZIPLIST_INT_24 => {
if pos + 4 > data.len() {
return Err(IoError::new(ErrorKind::InvalidData, "Ziplist 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
};
Ok((ZiplistEntry::Integer(i64::from(val)), 4))
}
super::symbols::ZIPLIST_INT_8 => {
if pos + 2 > data.len() {
return Err(IoError::new(ErrorKind::InvalidData, "Ziplist int8 too short").into());
}
let val = data[pos + 1] as i8;
Ok((ZiplistEntry::Integer(i64::from(val)), 2))
}
_ if (0xF1..=0xFD).contains(&encoding) => {
let val = i64::from((encoding & 0x0F) - 1);
Ok((ZiplistEntry::Integer(val), 1))
}
_ => Err(IoError::new(
ErrorKind::InvalidData,
format!("Unknown ziplist encoding: {encoding:#04x}"),
)
.into()),
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::parsers::rdb;
fn make_ziplist(entries: &[u8]) -> Vec<u8> {
let total_len = 11 + 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; 4]);
data.extend_from_slice(&[0u8; 2]);
data.extend_from_slice(entries);
data.push(rdb::symbols::ZIPLIST_END);
data
}
#[test]
fn test_parse_empty() {
let data = make_ziplist(&[]);
let entries = parse(&data).unwrap();
assert!(entries.is_empty());
}
#[test]
fn test_parse_6bit_string() {
let entries_data = [0x00, 0x05, b'h', b'e', b'l', b'l', b'o'];
let data = make_ziplist(&entries_data);
let entries = parse(&data).unwrap();
assert_eq!(entries.len(), 1);
assert!(matches!(&entries[0], ZiplistEntry::String(s) if s.as_ref() == "hello"));
}
#[test]
fn test_parse_small_int() {
let entries_data = [0x00, 0xF2];
let data = make_ziplist(&entries_data);
let entries = parse(&data).unwrap();
assert_eq!(entries.len(), 1);
assert!(matches!(&entries[0], ZiplistEntry::Integer(1)));
}
#[test]
fn test_parse_int8() {
let entries_data = [0x00, rdb::symbols::ZIPLIST_INT_8, 42];
let data = make_ziplist(&entries_data);
let entries = parse(&data).unwrap();
assert_eq!(entries.len(), 1);
assert!(matches!(&entries[0], ZiplistEntry::Integer(42)));
}
#[test]
fn test_parse_int16() {
let entries_data = [0x00, rdb::symbols::ZIPLIST_INT_16, 0xE8, 0x03];
let data = make_ziplist(&entries_data);
let entries = parse(&data).unwrap();
assert_eq!(entries.len(), 1);
assert!(matches!(&entries[0], ZiplistEntry::Integer(1000)));
}
#[test]
fn test_parse_hash() {
let entries_data = [
0x00, 0x03, b'k', b'e', b'y', 0x05, 0x05, b'v', b'a', b'l', b'u', b'e', ];
let data = make_ziplist(&entries_data);
let map = parse_hash(&data).unwrap();
assert_eq!(map.len(), 1);
let key = Value::String("key".into());
assert!(map.contains_key(&key));
assert_eq!(map.get(&key), Some(&Value::String("value".into())));
}
}