use std::io::{self, Read, Write};
use gxhash::HashSet;
use crate::{
inputs::ObjectInput,
objects::{
hash::hash_object::{glob_match, scan_operate_shared},
types::object_output::{ObjectOutput, ObjectOutputFlags},
},
resp::cmd_strings::RESP_ERR_GENERIC_UNSUPPORTED_OPERATION as RESP_ERR_UNSUPPORTED_OPERATION,
types::GarnetObjectType,
};
#[derive(
Debug, Clone, Copy, PartialEq, Eq, num_enum::TryFromPrimitive, num_enum::IntoPrimitive,
)]
#[repr(u8)]
pub enum SetOperation {
Sadd = 0,
Srem = 1,
Spop = 2,
Smembers = 3,
Scard = 4,
Sscan = 5,
Smove = 6,
Srandmember = 7,
Sismember = 8,
Smismember = 9,
Sunion = 10,
Sunionstore = 11,
Sdiff = 12,
Sdiffstore = 13,
Sinter = 14,
Sinterstore = 15,
}
#[derive(Debug, Clone, Default)]
pub struct SetObject {
pub set: HashSet<Vec<u8>>,
pub heap_memory_size: i64,
}
impl SetObject {
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.set.insert(item.clone());
obj.update_size(&item, true);
}
Ok(obj)
}
pub fn serialize<W: Write>(&self, writer: &mut W) -> io::Result<()> {
writer.write_all(&(self.set.len() as i32).to_le_bytes())?;
for item in &self.set {
writer.write_all(&(item.len() as i32).to_le_bytes())?;
writer.write_all(item)?;
}
Ok(())
}
pub fn from_members(members: Vec<Vec<u8>>) -> Self {
let mut obj = Self::new();
for member in members {
if obj.set.insert(member.clone()) {
obj.update_size(&member, true);
}
}
obj
}
pub fn to_members(&self) -> Vec<Vec<u8>> {
self.set.iter().cloned().collect()
}
pub fn operate(
&mut self,
input: &ObjectInput,
output: &mut ObjectOutput,
resp_protocol_version: u8,
) -> bool {
if input.header.data[0] != GarnetObjectType::Set as u8 {
output.output_flags |= ObjectOutputFlags::WRONG_TYPE;
output.payload.clear();
return true;
}
let Some(op) = set_op_from_header(input) else {
output.write_error(RESP_ERR_UNSUPPORTED_OPERATION.as_bytes());
return true;
};
match op {
SetOperation::Sadd => self.set_add(input, output),
SetOperation::Smembers => self.set_members(output, resp_protocol_version),
SetOperation::Sismember => self.set_is_member(input, output, resp_protocol_version),
SetOperation::Smismember => self.set_multi_is_member(input, output, resp_protocol_version),
SetOperation::Srem => self.set_remove(input, output),
SetOperation::Scard => self.set_length(output),
SetOperation::Spop => self.set_pop(input, output, resp_protocol_version),
SetOperation::Srandmember => self.set_random_member(input, output, resp_protocol_version),
SetOperation::Sscan => {
self.scan_operate(input, output);
}
SetOperation::Smove
| SetOperation::Sunion
| SetOperation::Sunionstore
| SetOperation::Sdiff
| SetOperation::Sdiffstore
| SetOperation::Sinter
| SetOperation::Sinterstore => {
output.write_error(RESP_ERR_UNSUPPORTED_OPERATION.as_bytes());
}
}
if self.set.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;
}
}
pub fn scan(&self, start: i64, count: i64, pattern: &[u8]) -> (Vec<Vec<u8>>, i64) {
let mut items: Vec<Vec<u8>> = Vec::new();
let mut cursor = start;
if (self.set.len() as i64) < start {
cursor = 0;
return (items, cursor);
}
let mut index = 0_i64;
for item in &self.set {
if index < start {
index += 1;
continue;
}
if pattern.is_empty() || glob_match(pattern, item) {
items.push(item.clone());
}
cursor += 1;
if items.len() as i64 == count {
break;
}
}
if cursor == self.set.len() as i64 {
cursor = 0;
}
(items, cursor)
}
}
#[inline]
pub fn set_op_from_header(input: &ObjectInput) -> Option<SetOperation> {
SetOperation::try_from(input.header.sub_id()).ok()
}
impl SetObject {
pub(crate) fn scan_operate(&mut self, input: &ObjectInput, output: &mut ObjectOutput) {
scan_operate_shared(input, output, |cursor, count, pattern, _is_no_value| {
self.scan(cursor, count, pattern)
});
}
}
#[cfg(test)]
mod tests {
use super::*;
fn obj_with(members: &[&str]) -> SetObject {
let mut obj = SetObject::new();
for m in members {
obj.set.insert(m.as_bytes().to_vec());
obj.update_size(m.as_bytes(), true);
}
obj
}
#[test]
fn serde_round_trip() {
let obj = obj_with(&["a", "b", "c"]);
let mut bytes = Vec::new();
obj.serialize(&mut bytes).unwrap();
let restored = SetObject::deserialize(&mut io::Cursor::new(&bytes)).unwrap();
assert_eq!(restored.set.len(), 3);
assert!(restored.set.contains(b"a".as_slice()));
}
#[test]
fn members_round_trip_bitcode_path() {
let obj = obj_with(&["x", "y"]);
let restored = SetObject::from_members(obj.to_members());
assert_eq!(restored.set.len(), 2);
}
#[test]
fn scan_single_item_form() {
let obj = obj_with(&["a", "b", "c", "d"]);
let (items, cursor) = obj.scan(0, 2, b"");
assert_eq!(items.len(), 2);
assert_eq!(cursor, 2);
let (items, cursor) = obj.scan(2, 2, b"");
assert_eq!(items.len(), 2);
assert_eq!(cursor, 0);
let (items, _) = obj.scan(0, 10, b"a*");
assert_eq!(items, [b"a".to_vec()]);
let (items, cursor) = obj.scan(99, 10, b"");
assert!(items.is_empty());
assert_eq!(cursor, 0);
let (items, cursor) = obj.scan(0, 0, b"");
assert_eq!(items.len(), 4);
assert_eq!(cursor, 0);
}
}