use std::collections::HashSet;
use std::time::{SystemTime, UNIX_EPOCH};
use serde::{Serialize, de::DeserializeOwned};
use crate::codec::{Bytes, Codec};
use crate::error::Error;
use crate::keys::{ScanCursor, ScanPage, matches_pattern_for_internal_use as matches_pattern};
use crate::store::Store;
pub struct SetRef<'a, C: Codec> {
store: &'a Store<C>,
key: &'a str,
}
impl<'a, C: Codec> SetRef<'a, C> {
pub(crate) fn new(store: &'a Store<C>, key: &'a str) -> Self {
Self { store, key }
}
#[inline]
fn enc<T: Serialize>(&self, v: &T) -> Result<Bytes, Error> {
self.store.codec().encode(v)
}
#[inline]
fn dec<T: DeserializeOwned>(&self, b: &[u8]) -> Result<T, Error> {
self.store.codec().decode(b)
}
fn pseudo_rand_index(len: usize) -> usize {
if len == 0 {
return 0;
}
let nanos = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.subsec_nanos() as usize;
nanos % len
}
pub fn sadd<T: Serialize>(&self, member: &T) -> Result<bool, Error> {
let m = self.enc(member)?;
self.store.with_set_mut(self.key, |set| Ok(set.insert(m)))
}
pub fn srem<T: Serialize>(&self, member: &T) -> Result<bool, Error> {
let m = self.enc(member)?;
self.store.with_set_mut(self.key, |set| Ok(set.remove(&m)))
}
pub fn smembers<T: DeserializeOwned>(&self) -> Result<Vec<T>, Error> {
self.store.with_set_read(self.key, |opt| {
let Some(set) = opt else {
return Ok(vec![]);
};
let mut out = Vec::with_capacity(set.len());
for b in set.iter() {
out.push(self.dec::<T>(b)?);
}
Ok(out)
})
}
pub fn sismember<T: Serialize>(&self, member: &T) -> Result<bool, Error> {
let m = self.enc(member)?;
self.store.with_set_read(self.key, |opt| {
let Some(set) = opt else {
return Ok(false);
};
Ok(set.contains(&m))
})
}
pub fn scard(&self) -> Result<usize, Error> {
self.store
.with_set_read(self.key, |opt| Ok(opt.map(|s| s.len()).unwrap_or(0)))
}
pub fn is_empty(&self) -> Result<bool, Error> {
Ok(self.scard()? == 0)
}
pub fn srandmember<T: DeserializeOwned>(&self) -> Result<Option<T>, Error> {
self.store.with_set_read(self.key, |opt| {
let Some(set) = opt else {
return Ok(None);
};
if set.is_empty() {
return Ok(None);
}
let idx = Self::pseudo_rand_index(set.len());
let b = set.iter().nth(idx).expect("idx < len and set not empty");
Ok(Some(self.dec::<T>(b)?))
})
}
pub fn spop<T: DeserializeOwned>(&self) -> Result<Option<T>, Error> {
self.store.with_set_mut(self.key, |set| {
if set.is_empty() {
return Ok(None);
}
let idx = Self::pseudo_rand_index(set.len());
let picked = set
.iter()
.nth(idx)
.cloned()
.expect("idx < len and set not empty");
set.remove(&picked);
Ok(Some(self.dec::<T>(&picked)?))
})
}
fn read_set_snapshot(&self, key: &str) -> Result<HashSet<Bytes>, Error> {
self.store.with_set_read(key, |opt| {
Ok(opt
.map(|s| s.iter().cloned().collect::<HashSet<_>>())
.unwrap_or_default())
})
}
pub fn sunion<T: DeserializeOwned>(&self, others: &[&str]) -> Result<Vec<T>, Error> {
let mut acc = self.read_set_snapshot(self.key)?;
for k in others {
let s = self.read_set_snapshot(k)?;
for x in s {
acc.insert(x);
}
}
let mut out = Vec::with_capacity(acc.len());
for b in acc {
out.push(self.dec::<T>(&b)?);
}
Ok(out)
}
pub fn sinter<T: DeserializeOwned>(&self, others: &[&str]) -> Result<Vec<T>, Error> {
let mut acc = self.read_set_snapshot(self.key)?;
for k in others {
let s = self.read_set_snapshot(k)?;
acc = acc.intersection(&s).cloned().collect();
}
let mut out = Vec::with_capacity(acc.len());
for b in acc {
out.push(self.dec::<T>(&b)?);
}
Ok(out)
}
pub fn sdiff<T: DeserializeOwned>(&self, others: &[&str]) -> Result<Vec<T>, Error> {
let mut acc = self.read_set_snapshot(self.key)?;
for k in others {
let s = self.read_set_snapshot(k)?;
for x in s {
acc.remove(&x);
}
}
let mut out = Vec::with_capacity(acc.len());
for b in acc {
out.push(self.dec::<T>(&b)?);
}
Ok(out)
}
fn write_dest_set(&self, dest: &str, values: HashSet<Bytes>) -> Result<usize, Error> {
self.store.purge_if_expired(dest);
self.store.with_map_write(|m| match m.get_mut(dest) {
None => {
let n = values.len();
m.insert(
dest.to_string(),
crate::entry::Entry::Set(crate::entry::SetEntry {
meta: crate::entry::Meta::new(crate::entry::ValueType::Set),
set: values,
}),
);
Ok(n)
}
Some(entry) => match entry {
crate::entry::Entry::Set(se) => {
let n = values.len();
se.set = values;
Ok(n)
}
other => Err(Error::WrongType {
expected: crate::entry::ValueType::Set.as_str(),
got: other.value_type().as_str(),
}),
},
})
}
pub fn sunionstore(&self, dest: &str, others: &[&str]) -> Result<usize, Error> {
let mut acc = self.read_set_snapshot(self.key)?;
for k in others {
let s = self.read_set_snapshot(k)?;
for x in s {
acc.insert(x);
}
}
self.write_dest_set(dest, acc)
}
pub fn sinterstore(&self, dest: &str, others: &[&str]) -> Result<usize, Error> {
let mut acc = self.read_set_snapshot(self.key)?;
for k in others {
let s = self.read_set_snapshot(k)?;
acc = acc.intersection(&s).cloned().collect();
}
self.write_dest_set(dest, acc)
}
pub fn sdiffstore(&self, dest: &str, others: &[&str]) -> Result<usize, Error> {
let mut acc = self.read_set_snapshot(self.key)?;
for k in others {
let s = self.read_set_snapshot(k)?;
for x in s {
acc.remove(&x);
}
}
self.write_dest_set(dest, acc)
}
pub fn sscan<T: DeserializeOwned>(
&self,
cursor: ScanCursor,
pattern: Option<&str>,
count: usize,
) -> Result<ScanPage<T>, Error> {
let mut bytes = self.store.with_set_read(self.key, |opt| {
let Some(set) = opt else {
return Ok(vec![]);
};
Ok(set.iter().cloned().collect::<Vec<_>>())
})?;
if bytes.is_empty() {
return Ok(ScanPage {
cursor: ScanCursor(0),
items: vec![],
});
}
bytes.sort_by(|a, b| a.as_ref().cmp(b.as_ref()));
let bytes = if let Some(pat) = pattern {
let mut filtered = Vec::with_capacity(bytes.len());
for b in bytes {
if let Ok(s) = self.dec::<String>(&b)
&& matches_pattern(&s, pat)
{
filtered.push(b);
}
}
filtered
} else {
bytes
};
let len = bytes.len();
if len == 0 {
return Ok(ScanPage {
cursor: ScanCursor(0),
items: vec![],
});
}
let start = cursor.0 as usize;
if start >= len {
return Ok(ScanPage {
cursor: ScanCursor(0),
items: vec![],
});
}
let take = count.max(1);
let end = (start + take).min(len);
let next = if end >= len { 0 } else { end as u64 };
let mut out = Vec::with_capacity(end - start);
for b in bytes.iter().take(end).skip(start) {
out.push(self.dec::<T>(b)?);
}
Ok(ScanPage {
cursor: ScanCursor(next),
items: out,
})
}
pub fn sclear(&self) -> Result<(), Error> {
self.store.with_set_mut(self.key, |set| {
set.clear();
Ok(())
})
}
}
#[cfg(test)]
mod tests {
use crate::{Store, keys};
#[test]
fn set_basic_ops() {
let store = Store::new();
let s = store.set("s");
assert_eq!(s.scard().unwrap(), 0);
assert!(s.sadd(&"a").unwrap());
assert!(!s.sadd(&"a").unwrap());
assert!(s.sismember(&"a").unwrap());
assert_eq!(s.scard().unwrap(), 1);
let members: Vec<String> = s.smembers().unwrap();
assert_eq!(members.len(), 1);
assert!(s.srem(&"a").unwrap());
assert!(!s.sismember(&"a").unwrap());
}
#[test]
fn set_union_inter_diff() {
let store = Store::new();
let a = store.set("a");
let b = store.set("b");
a.sadd(&"x").unwrap();
a.sadd(&"y").unwrap();
b.sadd(&"y").unwrap();
b.sadd(&"z").unwrap();
let mut u: Vec<String> = a.sunion(&["b"]).unwrap();
u.sort();
assert_eq!(u, vec!["x", "y", "z"]);
let mut i: Vec<String> = a.sinter(&["b"]).unwrap();
i.sort();
assert_eq!(i, vec!["y"]);
let mut d: Vec<String> = a.sdiff(&["b"]).unwrap();
d.sort();
assert_eq!(d, vec!["x"]);
}
#[test]
fn set_store_variants() {
let store = Store::new();
let a = store.set("a2");
let b = store.set("b2");
a.sadd(&"x").unwrap();
a.sadd(&"y").unwrap();
b.sadd(&"y").unwrap();
b.sadd(&"z").unwrap();
let n = a.sunionstore("u2", &["b2"]).unwrap();
assert_eq!(n, 3);
let u2: Vec<String> = store.set("u2").smembers().unwrap();
assert_eq!(u2.len(), 3);
}
#[test]
fn set_pop_random() {
let store = Store::new();
let s = store.set("p");
s.sadd(&"a").unwrap();
s.sadd(&"b").unwrap();
let _any: Option<String> = s.srandmember().unwrap();
let popped: Option<String> = s.spop().unwrap();
assert!(popped.is_some());
assert_eq!(s.scard().unwrap(), 1);
}
#[test]
fn set_sscan_basic_paging() {
let store = Store::new();
let s = store.set("scan");
for i in 0..15 {
s.sadd(&format!("k{i}")).unwrap();
}
let p1 = s
.sscan::<String>(keys::ScanCursor(0), Some("k*"), 10)
.unwrap();
assert_eq!(p1.items.len(), 10);
assert_ne!(p1.cursor.0, 0);
let p2 = s.sscan::<String>(p1.cursor, Some("k*"), 10).unwrap();
assert_eq!(p2.cursor.0, 0);
}
}