use std::{
collections::VecDeque,
io::{self, Read, Write},
};
use crate::{
inputs::ObjectInput,
objects::types::object_output::{ObjectOutput, ObjectOutputFlags},
types::GarnetObjectType,
};
#[derive(
Debug, Clone, Copy, PartialEq, Eq, num_enum::TryFromPrimitive, num_enum::IntoPrimitive,
)]
#[repr(u8)]
pub enum ListOperation {
Lpop = 0,
Lpush = 1,
Lpushx = 2,
Rpop = 3,
Rpush = 4,
Rpushx = 5,
Llen = 6,
Ltrim = 7,
Lrange = 8,
Lindex = 9,
Linsert = 10,
Lrem = 11,
Rpoplpush = 12,
Lmove = 13,
Lset = 14,
Brpop = 15,
Blpop = 16,
Lpos = 17,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[repr(u8)]
pub enum OperationDirection {
Left = 0,
Right = 1,
Unknown = 2,
}
#[derive(Debug, Clone, Default)]
pub struct ListObject {
pub list: VecDeque<Vec<u8>>,
pub heap_memory_size: i64,
}
impl ListObject {
pub fn new() -> Self {
Self::default()
}
pub fn deserialize<R: Read>(reader: &mut R) -> io::Result<Self> {
let mut obj = Self::new();
let mut len_buf = [0_u8; 4];
reader.read_exact(&mut len_buf)?;
let count = i32::from_le_bytes(len_buf);
for _ in 0..count {
reader.read_exact(&mut len_buf)?;
let mut item = vec![0_u8; i32::from_le_bytes(len_buf) as usize];
reader.read_exact(&mut item)?;
obj.list.push_back(item.clone());
obj.update_size(&item, true);
}
Ok(obj)
}
pub fn serialize<W: Write>(&self, writer: &mut W) -> io::Result<()> {
writer.write_all(&(self.list.len() as i32).to_le_bytes())?;
for item in &self.list {
writer.write_all(&(item.len() as i32).to_le_bytes())?;
writer.write_all(item)?;
}
Ok(())
}
pub fn from_items(items: Vec<Vec<u8>>) -> Self {
let mut obj = Self::new();
for item in items {
obj.list.push_back(item.clone());
obj.update_size(&item, true);
}
obj
}
pub fn to_items(&self) -> Vec<Vec<u8>> {
self.list.iter().cloned().collect()
}
pub fn operate(
&mut self,
input: &ObjectInput,
output: &mut ObjectOutput,
resp_protocol_version: u8,
) -> bool {
if input.header.data[0] != GarnetObjectType::List as u8 {
output.output_flags |= ObjectOutputFlags::WRONG_TYPE;
output.payload.clear();
return true;
}
let Some(op) = list_op_from_header(input) else {
output.write_error(b"ERR unsupported operation");
return true;
};
match op {
ListOperation::Lpush | ListOperation::Lpushx => self.list_push(input, output, true),
ListOperation::Rpush | ListOperation::Rpushx => self.list_push(input, output, false),
ListOperation::Lpop => self.list_pop(input, output, resp_protocol_version, true),
ListOperation::Rpop => self.list_pop(input, output, resp_protocol_version, false),
ListOperation::Llen => self.list_length(output),
ListOperation::Ltrim => self.list_trim(input, output),
ListOperation::Lrange => self.list_range(input, output, resp_protocol_version),
ListOperation::Lindex => self.list_index(input, output, resp_protocol_version),
ListOperation::Linsert => self.list_insert(input, output),
ListOperation::Lrem => self.list_remove(input, output),
ListOperation::Lset => self.list_set(input, output, resp_protocol_version),
ListOperation::Lpos => self.list_position(input, output, resp_protocol_version),
ListOperation::Rpoplpush
| ListOperation::Lmove
| ListOperation::Brpop
| ListOperation::Blpop => {
output.write_error(b"ERR unsupported operation");
}
}
if self.list.is_empty() {
output.output_flags |= ObjectOutputFlags::REMOVE_KEY;
}
true
}
pub fn update_size(&mut self, item: &[u8], add: bool) {
let memory_size = (item.len().div_ceil(8) * 8 + 16 + 16) as i64;
if add {
self.heap_memory_size += memory_size;
} else {
self.heap_memory_size -= memory_size;
}
}
#[inline]
pub fn nodes(&self) -> impl Iterator<Item = usize> {
0..self.list.len()
}
}
#[inline]
pub fn list_op_from_header(input: &ObjectInput) -> Option<ListOperation> {
ListOperation::try_from(input.header.sub_id()).ok()
}
#[cfg(test)]
mod tests {
use super::*;
fn obj_with_items(items: &[&str]) -> ListObject {
let mut obj = ListObject::new();
for item in items {
obj.list.push_back(item.as_bytes().to_vec());
obj.update_size(item.as_bytes(), true);
}
obj
}
#[test]
fn serde_round_trip() {
let obj = obj_with_items(&["a", "b", "c"]);
let mut bytes = Vec::new();
obj.serialize(&mut bytes).unwrap();
let restored = ListObject::deserialize(&mut io::Cursor::new(&bytes)).unwrap();
assert_eq!(restored.to_items(), obj.to_items());
assert_eq!(restored.list.len(), 3);
}
#[test]
fn items_round_trip_bitcode_path() {
let obj = obj_with_items(&["x", "y"]);
let restored = ListObject::from_items(obj.to_items());
assert_eq!(restored.to_items(), vec![b"x".to_vec(), b"y".to_vec()]);
}
#[test]
fn nodes_walk_all_positions() {
let obj = obj_with_items(&["a", "b", "c"]);
assert_eq!(obj.nodes().collect::<Vec<_>>(), vec![0, 1, 2]);
}
}