use std::io::{self, Read, Write};
use gxhash::HashSet;
use wbase::glob::glob_match;
use wresp::cmd_strings::RESP_ERR_GENERIC_UNSUPPORTED_OPERATION as RESP_ERR_UNSUPPORTED_OPERATION;
use wval::GarnetObjectType;
use crate::{
hash::hash_object::scan_operate_shared,
object_store_utils::GarnetObjectPayload,
types::{
ObjectInput,
object_output::{ObjectOutput, ObjectOutputFlags},
},
};
#[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()
}
#[inline]
pub fn len(&self) -> usize {
self.set.len()
}
#[inline]
pub fn is_empty(&self) -> bool {
self.set.is_empty()
}
pub fn deserialize_from_slice(slice: &[u8]) -> io::Result<Self> {
let set: HashSet<Vec<u8>> =
bitcode::decode(slice).map_err(|e| io::Error::new(io::ErrorKind::InvalidData, e))?;
let mut obj = Self::new();
for item in &set {
obj.update_size(item, true);
}
obj.set = set;
Ok(obj)
}
pub fn deserialize<R: Read>(reader: &mut R) -> io::Result<Self> {
let mut buf = Vec::new();
reader.read_to_end(&mut buf)?;
Self::deserialize_from_slice(&buf)
}
#[inline]
pub fn serialize_to_vec(&self) -> Vec<u8> {
bitcode::encode(&self.set)
}
pub fn serialize<W: Write>(&self, writer: &mut W) -> io::Result<()> {
writer.write_all(&self.serialize_to_vec())
}
pub fn count(&self) -> usize {
self.set.len()
}
pub fn add(&mut self, member: &[u8]) -> bool {
if self.set.insert(member.to_vec()) {
self.update_size(member, true);
true
} else {
false
}
}
pub fn remove(&mut self, member: &[u8]) -> bool {
if self.set.remove(member) {
self.update_size(member, false);
true
} else {
false
}
}
pub fn contains(&self, member: &[u8]) -> bool {
self.set.contains(member)
}
pub fn pop(&mut self) -> Option<Vec<u8>> {
let item = self.set.iter().next().cloned()?;
self.set.remove(&item);
self.update_size(&item, false);
Some(item)
}
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),
SetOperation::Smismember => self.set_multi_is_member(input, output),
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, _| {
self.scan(cursor, count, pattern)
});
}
}
impl GarnetObjectPayload for SetObject {
const OBJECT_TAG: GarnetObjectType = GarnetObjectType::Set;
#[inline]
fn from_blob(raw: &[u8]) -> Self {
Self::deserialize_from_slice(raw).unwrap_or_default()
}
#[inline]
fn to_blob(&self) -> Vec<u8> {
self.serialize_to_vec()
}
#[inline]
fn is_empty(&self) -> bool {
self.set.is_empty()
}
}