use std::collections::{BTreeMap, BTreeSet};
use ordered_float::OrderedFloat;
use serde::{Serialize, de::DeserializeOwned};
use crate::codec::{Bytes, Codec};
use crate::error::Error;
use crate::keys::matches_pattern_for_internal_use as matches_pattern;
use crate::keys::{ScanCursor, ScanPage};
use crate::store::Store;
type ScoreKey = OrderedFloat<f64>;
type ScoreIndex = BTreeMap<ScoreKey, BTreeSet<Bytes>>;
pub struct ZSetRef<'a, C: Codec> {
store: &'a Store<C>,
key: &'a str,
}
impl<'a, C: Codec> ZSetRef<'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)
}
#[inline]
fn index_insert(score_to_members: &mut ScoreIndex, score: f64, member: Bytes) {
score_to_members
.entry(OrderedFloat(score))
.or_default()
.insert(member);
}
#[inline]
fn index_remove(score_to_members: &mut ScoreIndex, score: f64, member: &Bytes) {
let key = OrderedFloat(score);
if let Some(set) = score_to_members.get_mut(&key) {
set.remove(member);
if set.is_empty() {
score_to_members.remove(&key);
}
}
}
pub fn zadd<T: Serialize>(&self, score: f64, member: &T) -> Result<bool, Error> {
let m = self.enc(member)?;
self.store.with_zset_mut(self.key, |ze| {
if let Some(old) = ze.member_to_score.insert(m.clone(), score) {
if old != score {
Self::index_remove(&mut ze.score_to_members, old, &m);
Self::index_insert(&mut ze.score_to_members, score, m);
}
Ok(false)
} else {
Self::index_insert(&mut ze.score_to_members, score, m);
Ok(true)
}
})
}
pub fn zrem<T: Serialize>(&self, member: &T) -> Result<bool, Error> {
let m = self.enc(member)?;
self.store.with_zset_mut(self.key, |ze| {
let Some(old) = ze.member_to_score.remove(&m) else {
return Ok(false);
};
Self::index_remove(&mut ze.score_to_members, old, &m);
Ok(true)
})
}
pub fn zscore<T: Serialize>(&self, member: &T) -> Result<Option<f64>, Error> {
let m = self.enc(member)?;
self.store.with_zset_read(self.key, |opt| {
let Some(ze) = opt else { return Ok(None) };
Ok(ze.member_to_score.get(&m).copied())
})
}
pub fn zcard(&self) -> Result<usize, Error> {
self.store.with_zset_read(self.key, |opt| {
Ok(opt.map(|ze| ze.member_to_score.len()).unwrap_or(0))
})
}
pub fn zrange<T: DeserializeOwned>(&self, start: isize, stop: isize) -> Result<Vec<T>, Error> {
self.range_by_rank::<T>(start, stop, false)
}
pub fn zrevrange<T: DeserializeOwned>(
&self,
start: isize,
stop: isize,
) -> Result<Vec<T>, Error> {
self.range_by_rank::<T>(start, stop, true)
}
fn range_by_rank<T: DeserializeOwned>(
&self,
start: isize,
stop: isize,
rev: bool,
) -> Result<Vec<T>, Error> {
self.store.with_zset_read(self.key, |opt| {
let Some(ze) = opt else { return Ok(vec![]) };
let n = ze.member_to_score.len() as isize;
if n == 0 {
return Ok(vec![]);
}
let mut s = if start < 0 { n + start } else { start };
let mut e = if stop < 0 { n + stop } else { stop };
if s < 0 {
s = 0;
}
if e < 0 {
return Ok(vec![]);
}
if s >= n {
return Ok(vec![]);
}
if e >= n {
e = n - 1;
}
if e < s {
return Ok(vec![]);
}
let want_start = s as usize;
let want_end = e as usize; let mut out = Vec::with_capacity(want_end - want_start + 1);
let iter_scores: Box<dyn Iterator<Item = (&ScoreKey, &BTreeSet<Bytes>)>> = if rev {
Box::new(ze.score_to_members.iter().rev())
} else {
Box::new(ze.score_to_members.iter())
};
let mut idx = 0usize;
for (_score, members) in iter_scores {
let iter_members: Box<dyn Iterator<Item = &Bytes>> = if rev {
Box::new(members.iter().rev())
} else {
Box::new(members.iter())
};
for m in iter_members {
if idx >= want_start && idx <= want_end {
out.push(self.dec::<T>(m)?);
}
if idx > want_end {
return Ok(out);
}
idx += 1;
}
}
Ok(out)
})
}
pub fn zrangebyscore<T: DeserializeOwned>(&self, min: f64, max: f64) -> Result<Vec<T>, Error> {
let lo = OrderedFloat(min);
let hi = OrderedFloat(max);
self.store.with_zset_read(self.key, |opt| {
let Some(ze) = opt else { return Ok(vec![]) };
let mut out = Vec::new();
for (_score, members) in ze.score_to_members.range(lo..=hi) {
for m in members.iter() {
out.push(self.dec::<T>(m)?);
}
}
Ok(out)
})
}
pub fn zrank<T: Serialize>(&self, member: &T) -> Result<Option<usize>, Error> {
self.rank_of(member, false)
}
pub fn zrevrank<T: Serialize>(&self, member: &T) -> Result<Option<usize>, Error> {
self.rank_of(member, true)
}
fn rank_of<T: Serialize>(&self, member: &T, rev: bool) -> Result<Option<usize>, Error> {
let m = self.enc(member)?;
self.store.with_zset_read(self.key, |opt| {
let Some(ze) = opt else { return Ok(None) };
let Some(score) = ze.member_to_score.get(&m).copied() else {
return Ok(None);
};
let mut idx = 0usize;
let iter_scores: Box<dyn Iterator<Item = (&ScoreKey, &BTreeSet<Bytes>)>> = if rev {
Box::new(ze.score_to_members.iter().rev())
} else {
Box::new(ze.score_to_members.iter())
};
for (s, members) in iter_scores {
let iter_members: Box<dyn Iterator<Item = &Bytes>> = if rev {
Box::new(members.iter().rev())
} else {
Box::new(members.iter())
};
for mm in iter_members {
if *s == OrderedFloat(score) && *mm == m {
return Ok(Some(idx));
}
idx += 1;
}
}
Ok(None)
})
}
pub fn zincrby<T: Serialize>(&self, member: &T, delta: f64) -> Result<f64, Error> {
let m = self.enc(member)?;
self.store.with_zset_mut(self.key, |ze| {
let old = ze.member_to_score.get(&m).copied().unwrap_or(0.0);
let next = old + delta;
if ze.member_to_score.contains_key(&m) {
Self::index_remove(&mut ze.score_to_members, old, &m);
}
ze.member_to_score.insert(m.clone(), next);
Self::index_insert(&mut ze.score_to_members, next, m);
Ok(next)
})
}
pub fn zpopmin<T: DeserializeOwned>(&self) -> Result<Option<(T, f64)>, Error> {
self.pop_extreme::<T>( false)
}
pub fn zpopmax<T: DeserializeOwned>(&self) -> Result<Option<(T, f64)>, Error> {
self.pop_extreme::<T>( true)
}
fn pop_extreme<T: DeserializeOwned>(&self, max: bool) -> Result<Option<(T, f64)>, Error> {
self.store.with_zset_mut(self.key, |ze| {
let (score_key, members_snapshot) = if max {
match ze.score_to_members.iter().next_back() {
None => return Ok(None),
Some((s, m)) => (*s, m.clone()),
}
} else {
match ze.score_to_members.iter().next() {
None => return Ok(None),
Some((s, m)) => (*s, m.clone()),
}
};
let member = members_snapshot
.iter()
.next()
.cloned()
.expect("members not empty");
let score = score_key.0;
ze.member_to_score.remove(&member);
Self::index_remove(&mut ze.score_to_members, score, &member);
let decoded = self.dec::<T>(&member)?;
Ok(Some((decoded, score)))
})
}
pub fn zremrangebyscore(&self, min: f64, max: f64) -> Result<usize, Error> {
let lo = OrderedFloat(min);
let hi = OrderedFloat(max);
self.store.with_zset_mut(self.key, |ze| {
let mut to_remove: Vec<(Bytes, f64)> = Vec::new();
for (s, members) in ze.score_to_members.range(lo..=hi) {
for m in members.iter() {
to_remove.push((m.clone(), s.0));
}
}
for (m, s) in &to_remove {
ze.member_to_score.remove(m);
Self::index_remove(&mut ze.score_to_members, *s, m);
}
Ok(to_remove.len())
})
}
pub fn zremrangebyrank(&self, start: isize, stop: isize) -> Result<usize, Error> {
self.store.with_zset_mut(self.key, |ze| {
let n = ze.member_to_score.len() as isize;
if n == 0 {
return Ok(0);
}
let mut s = if start < 0 { n + start } else { start };
let mut e = if stop < 0 { n + stop } else { stop };
if s < 0 {
s = 0;
}
if e < 0 {
return Ok(0);
}
if s >= n {
return Ok(0);
}
if e >= n {
e = n - 1;
}
if e < s {
return Ok(0);
}
let want_start = s as usize;
let want_end = e as usize;
let mut idx = 0usize;
let mut targets: Vec<(Bytes, f64)> = Vec::new();
for (score_key, members) in ze.score_to_members.iter() {
for m in members.iter() {
if idx >= want_start && idx <= want_end {
targets.push((m.clone(), score_key.0));
}
if idx > want_end {
break;
}
idx += 1;
}
if idx > want_end {
break;
}
}
for (m, sc) in &targets {
ze.member_to_score.remove(m);
Self::index_remove(&mut ze.score_to_members, *sc, m);
}
Ok(targets.len())
})
}
pub fn zscan<T: DeserializeOwned>(
&self,
cursor: ScanCursor,
pattern: Option<&str>,
count: usize,
) -> Result<ScanPage<(T, f64)>, Error> {
let mut items: Vec<(Bytes, f64)> = self.store.with_zset_read(self.key, |opt| {
let Some(ze) = opt else { return Ok(vec![]) };
let mut out = Vec::with_capacity(ze.member_to_score.len());
for (m, s) in ze.member_to_score.iter() {
out.push((m.clone(), *s));
}
Ok(out)
})?;
if let Some(pat) = pattern {
let mut filtered = Vec::new();
for (m, s) in items.into_iter() {
if let Ok(as_str) = self.dec::<String>(&m)
&& matches_pattern(&as_str, pat)
{
filtered.push((m, s));
}
}
items = filtered;
}
items.sort_by(|a, b| a.0.cmp(&b.0));
let len = items.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 (m, s) in items.into_iter().skip(start).take(end - start) {
out.push((self.dec::<T>(&m)?, s));
}
Ok(ScanPage {
cursor: ScanCursor(next),
items: out,
})
}
}
#[cfg(test)]
mod tests {
use crate::Store;
#[test]
fn zset_basic_ops() {
let store = Store::new();
let z = store.zset("rank");
assert_eq!(z.zcard().unwrap(), 0);
assert!(z.zadd(10.0, &"a").unwrap());
assert!(!z.zadd(12.0, &"a").unwrap());
assert!(z.zadd(11.0, &"b").unwrap());
assert_eq!(z.zcard().unwrap(), 2);
assert_eq!(z.zscore(&"a").unwrap(), Some(12.0));
let r: Vec<String> = z.zrange(0, -1).unwrap();
assert_eq!(r, vec!["b".to_string(), "a".to_string()]);
}
#[test]
fn zset_rank_pop_remove_ranges() {
let store = Store::new();
let z = store.zset("z2");
z.zadd(1.0, &"a").unwrap();
z.zadd(2.0, &"b").unwrap();
z.zadd(3.0, &"c").unwrap();
assert_eq!(z.zrank(&"a").unwrap(), Some(0));
assert_eq!(z.zrevrank(&"a").unwrap(), Some(2));
let p: Option<(String, f64)> = z.zpopmin().unwrap();
assert_eq!(p.unwrap().0, "a".to_string());
let removed = z.zremrangebyscore(2.0, 3.0).unwrap();
assert_eq!(removed, 2);
assert_eq!(z.zcard().unwrap(), 0);
}
#[test]
fn zset_incr_and_revrange() {
let store = Store::new();
let z = store.zset("z3");
z.zadd(1.0, &"a").unwrap();
let v = z.zincrby(&"a", 2.5).unwrap();
assert!((v - 3.5).abs() < 1e-9);
let rr: Vec<String> = z.zrevrange(0, -1).unwrap();
assert_eq!(rr, vec!["a".to_string()]);
}
#[test]
fn zset_update_score_reorders() {
let store = Store::new();
let z = store.zset("z");
z.zadd(1.0, &"a").unwrap();
z.zadd(2.0, &"b").unwrap();
let r1: Vec<String> = z.zrange(0, -1).unwrap();
assert_eq!(r1, vec!["a", "b"]);
z.zadd(3.0, &"a").unwrap();
let r2: Vec<String> = z.zrange(0, -1).unwrap();
assert_eq!(r2, vec!["b", "a"]);
}
#[test]
fn zset_edge_cases() {
let store = Store::new();
let z = store.zset("z");
assert_eq!(z.zpopmax::<String>().unwrap(), None);
assert!(!z.zrem(&"missing").unwrap());
assert_eq!(z.zrank(&"missing").unwrap(), None);
z.zadd(1.0, &"a").unwrap();
z.zadd(2.0, &"b").unwrap();
let empty: Vec<String> = z.zrangebyscore(10.0, 20.0).unwrap();
assert!(empty.is_empty());
let rr: Vec<String> = z.zrevrange(0, 0).unwrap();
assert_eq!(rr, vec!["b".to_string()]);
let p = z.zpopmax::<String>().unwrap();
assert_eq!(p.unwrap().0, "b".to_string());
}
}