use fastrand::Rng;
use super::set_object::SetObject;
use crate::{
hash::hash_object::{pick_k_random_indexes, pick_random_index},
types::{ObjectInput, object_output::ObjectOutput},
};
#[inline]
fn arg(input: &ObjectInput, i: usize) -> &[u8] {
input.arg(i)
}
pub const NO_COUNT: i32 = i32::MIN;
impl SetObject {
pub(crate) fn set_add(&mut self, input: &ObjectInput, output: &mut ObjectOutput) {
let mut added = 0_i64;
for i in 0..input.parse_state.count {
let member = arg(input, i);
if self.set.insert(member.to_vec()) {
added += 1;
self.update_size(member, true);
}
}
output.result1 = added;
}
pub(crate) fn set_members(&mut self, output: &mut ObjectOutput, resp_protocol_version: u8) {
write_set_length(output, self.set.len(), resp_protocol_version);
let mut written = 0_i64;
for item in &self.set {
output.write_bulk_string(item);
written += 1;
}
output.result1 = written;
}
pub(crate) fn set_is_member(&mut self, input: &ObjectInput, output: &mut ObjectOutput) {
let member = arg(input, 0);
let is_member = self.set.contains(member);
output.write_int64(i64::from(is_member));
output.result1 = 1;
}
pub(crate) fn set_multi_is_member(&mut self, input: &ObjectInput, output: &mut ObjectOutput) {
output.write_array_length(input.parse_state.count);
for i in 0..input.parse_state.count {
let member = arg(input, i);
output.write_int64(i64::from(self.set.contains(member)));
}
output.result1 = input.parse_state.count as i64;
}
pub(crate) fn set_remove(&mut self, input: &ObjectInput, output: &mut ObjectOutput) {
let mut removed = 0_i64;
for i in 0..input.parse_state.count {
let member = arg(input, i);
if self.set.remove(member) {
removed += 1;
self.update_size(member, false);
}
}
output.result1 = removed;
}
pub(crate) fn set_length(&mut self, output: &mut ObjectOutput) {
output.result1 = self.set.len() as i64;
}
pub(crate) fn set_pop(
&mut self,
input: &ObjectInput,
output: &mut ObjectOutput,
resp_protocol_version: u8,
) {
let count = input.arg1;
let mut count_done = 0_i64;
if count >= 1 {
let count_parameter = (count as usize).min(self.set.len());
let mut rng = Rng::new();
write_set_length(output, count_parameter, resp_protocol_version);
for _ in 0..count_parameter {
if self.set.is_empty() {
break;
}
let index = rng.usize(..self.set.len());
let Some(item) = self.set.iter().nth(index).cloned() else {
break;
};
self.set.remove(&item);
self.update_size(&item, false);
output.write_bulk_string(&item);
count_done += 1;
}
count_done += i64::from(count) - count_done;
} else if count == NO_COUNT {
if !self.set.is_empty() {
let index = fastrand::usize(..self.set.len());
let item = self.set.iter().nth(index).cloned().unwrap();
self.set.remove(&item);
self.update_size(&item, false);
output.write_bulk_string(&item);
} else {
output.write_null(resp_protocol_version);
}
count_done += 1;
}
output.result1 = count_done;
}
pub(crate) fn set_random_member(
&mut self,
input: &ObjectInput,
output: &mut ObjectOutput,
resp_protocol_version: u8,
) {
let count = input.arg1;
let seed = input.arg2;
let mut count_done = 0_i64;
if count > 0 {
let count_parameter = (count as usize).min(self.set.len());
let indexes = pick_k_random_indexes(self.set.len(), count_parameter, seed, true);
write_set_length(output, count_parameter, resp_protocol_version);
for index in indexes {
let Some(element) = self.set.iter().nth(index).cloned() else {
continue;
};
output.write_bulk_string(&element);
count_done += 1;
}
count_done += i64::from(count) - count_parameter as i64;
} else if count == NO_COUNT {
if !self.set.is_empty() {
let index = pick_random_index(self.set.len(), seed);
if let Some(item) = self.set.iter().nth(index).cloned() {
output.write_bulk_string(&item);
}
} else {
output.write_null(resp_protocol_version);
}
count_done += 1;
} else {
let count_parameter = count.unsigned_abs() as usize;
let indexes = pick_k_random_indexes(self.set.len(), count_parameter, seed, false);
if !self.set.is_empty() {
output.write_array_length(count_parameter);
for index in indexes {
let Some(element) = self.set.iter().nth(index).cloned() else {
continue;
};
output.write_bulk_string(&element);
count_done += 1;
}
} else {
output.write_null(resp_protocol_version);
}
}
output.result1 = count_done;
}
}
#[inline]
fn write_set_length(output: &mut ObjectOutput, len: usize, resp_protocol_version: u8) {
output.write_set_length(len, resp_protocol_version);
}