use yo_common::lock::Held;
use yo_common::num::DIGITS_MAX;
use yo_common::num::i64_digits;
use yo_common::{Error, Result, glob_matches, parse_i64};
use yo_kv::lookups;
use yo_kv::{Aggregate, Db, Member, Query, ZAdd, ZBound, ZEnd, ZOp};
use super::args::{self, Args};
use super::notify::{self, class};
use super::scan;
use super::table::Spec;
use crate::reply::Out;
const BAD_RANGE: &str = "min or max is not a float";
const BAD_LEX: &str = "min or max not valid string range item";
const NX_AND_XX: &str = "XX and NX options at the same time are not compatible";
const NX_AND_GT_LT: &str = "GT, LT, and/or NX options at the same time are not compatible";
const ONE_PAIR: &str = "INCR option supports a single increment-element pair";
const LIMIT_NEEDS_BY: &str =
"syntax error, LIMIT is only supported in combination with either BYSCORE or BYLEX";
const SCORES_NOT_BYLEX: &str = "syntax error, WITHSCORES not supported in combination with BYLEX";
const BAD_WEIGHT: &str = "weight value is not a float";
const BAD_LIMIT: &str = "LIMIT can't be negative";
const NOVALUES_IS_HSCAN: &str = "NOVALUES option can only be used in HSCAN";
const BAD_POP_COUNT: &str = "value is out of range, must be positive";
pub(super) const BAD_NUMKEYS: &str = "numkeys should be greater than 0";
pub(super) const BAD_MPOP_COUNT: &str = "count should be greater than 0";
pub(super) fn execute(
db: &Db,
on: usize,
spec: &Spec,
args: Args<'_>,
out: &mut Out,
) -> Result<()> {
match spec.name {
"zadd" => zadd(db, on, args, out)?,
"zincrby" => {
let by = score(args.get(2))?;
let (key, m) = (args.get(1), args.get(3));
let mut stripe = db.hold(key);
let before = moved_from(&mut stripe, key, m)?;
let now = stripe
.zincrby(key, m, by, ZAdd::default())?
.expect("ZINCRBY has no gate that can refuse a member");
drop(stripe);
out.double(now);
if before != Some(now) {
notify::fire(on, class::ZSET, "zincr", key);
}
}
"zcard" => {
let key = args.get(1);
out.int(count(db.hold(key).zcard(key)?));
}
"zscore" => {
let key = args.get(1);
match db.hold(key).zscore(key, args.get(2))? {
Some(s) => out.double(s),
None => out.nil(),
}
}
"zmscore" => {
out.array(args.len() - 2);
let key = args.get(1);
let mut stripe = db.hold(key);
let mut quiet = None;
for i in 2..args.len() {
match stripe.zscore(key, args.get(i))? {
Some(s) => out.double(s),
None => out.nil(),
}
quiet.get_or_insert_with(lookups::quiet);
}
}
"zrem" => {
let key = args.get(1);
let gone = db.hold(key).zrem(key, members(args))?;
out.int(count(gone));
if gone > 0 {
notify::fire(on, class::ZSET, "zrem", key);
notify::emptied(db, on, key);
}
}
"zrank" | "zrevrank" => rank(db, spec.name == "zrevrank", args, out)?,
"zcount" => {
let q = Query::score(bound(args.get(2))?, bound(args.get(3))?);
let key = args.get(1);
out.int(count(db.hold(key).zcount(key, &q)?));
}
"zlexcount" => {
let q = Query::lex(lex(args.get(2))?, lex(args.get(3))?);
let key = args.get(1);
out.int(count(db.hold(key).zcount(key, &q)?));
}
"zrange" | "zrevrange" | "zrangebyscore" | "zrevrangebyscore" | "zrangebylex"
| "zrevrangebylex" => {
let form = Form::of(spec.name);
let (q, withscores) = parse_range(form, args, 1)?;
let key = args.get(1);
let mut stripe = db.hold(key);
let w = stripe.zwindow(key, &q)?;
let _quiet = lookups::quiet();
let nested = withscores && out.proto().is_resp3();
out.array(if withscores && !nested {
w.count * 2
} else {
w.count
});
stripe.zwalk(key, w, |m, sc| {
if nested {
out.array(2);
}
write_member(out, m);
if withscores {
out.double(sc);
}
})?;
}
"zrangestore" => {
let (q, _) = parse_range(Form::Store, args, 2)?;
let dst = args.get(1);
let had = notify::armed() && db.hold(dst).exists(dst);
let stored = db.zrangestore(dst, args.get(2), &q)?;
out.int(count(stored));
if stored > 0 {
notify::fire(on, class::ZSET, "zrangestore", dst);
} else if had {
notify::fire(on, class::GENERIC, "del", dst);
}
}
"zremrangebyrank" | "zremrangebyscore" | "zremrangebylex" => {
let q = match spec.name {
"zremrangebyrank" => Query::rank(args.int(2)?, args.int(3)?),
"zremrangebyscore" => Query::score(bound(args.get(2))?, bound(args.get(3))?),
_ => Query::lex(lex(args.get(2))?, lex(args.get(3))?),
};
let key = args.get(1);
let gone = db.hold(key).zremrange(key, &q)?;
out.int(count(gone));
if gone > 0 {
notify::fire(on, class::ZSET, spec.name, key);
notify::emptied(db, on, key);
}
}
"zunion" | "zinter" | "zdiff" => {
let op = op_of(spec.name);
let a = Algebra::parse(spec.name, op, args, 1, Scores::Allowed)?;
let nested = a.withscores && out.proto().is_resp3();
let start = out.len();
let mut n = 0;
db.zsetop(op, a.keys(args), &a.weights, a.agg, |m, sc| {
if nested {
out.array(2);
}
write_member(out, m);
if a.withscores {
out.double(sc);
}
n += 1;
})?;
out.close_array(start, if nested || !a.withscores { n } else { n * 2 });
}
"zunionstore" | "zinterstore" | "zdiffstore" => {
let op = op_of(spec.name);
let a = Algebra::parse(spec.name, op, args, 2, Scores::Refused)?;
let dst = args.get(1);
let had = notify::armed() && db.hold(dst).exists(dst);
let got = db.zsetop_store(op, dst, a.keys(args), &a.weights, a.agg)?;
out.int(count(got));
if got > 0 {
notify::fire(on, class::ZSET, spec.name, dst);
} else if had {
notify::fire(on, class::GENERIC, "del", dst);
}
}
"zintercard" => out.int(count(intercard(db, args)?)),
"zrandmember" => randmember(db, args, out)?,
"zscan" => zscan(db, args, out)?,
"zpopmin" | "zpopmax" => pop(db, on, spec.name, end_of_name(spec.name), args, out)?,
"zmpop" => mpop(db, on, args, out)?,
other => unreachable!("{other} is not a sorted set command"),
}
Ok(())
}
fn zadd(db: &Db, on: usize, args: Args<'_>, out: &mut Out) -> Result<()> {
let (mut nx, mut xx, mut gt, mut lt, mut incr) = (false, false, false, false, false);
let mut opts = ZAdd::default();
let mut at = 2;
while at < args.len() {
let arg = args.get(at);
if args::is(arg, b"nx") {
nx = true;
} else if args::is(arg, b"xx") {
xx = true;
} else if args::is(arg, b"gt") {
gt = true;
} else if args::is(arg, b"lt") {
lt = true;
} else if args::is(arg, b"ch") {
opts.changed = true;
} else if args::is(arg, b"incr") {
incr = true;
} else {
break;
}
at += 1;
}
let left = args.len() - at;
if left == 0 || !left.is_multiple_of(2) {
return Err(args::syntax());
}
if nx && xx {
return Err(Error::new(yo_common::Code::Invalid, NX_AND_XX));
}
if (nx && (gt || lt)) || (gt && lt) {
return Err(Error::new(yo_common::Code::Invalid, NX_AND_GT_LT));
}
if incr && left != 2 {
return Err(Error::new(yo_common::Code::Invalid, ONE_PAIR));
}
opts.gate = if nx {
yo_kv::Gate::IfMissing
} else if xx {
yo_kv::Gate::IfPresent
} else {
yo_kv::Gate::Always
};
opts.only = if gt {
yo_kv::Move::Up
} else if lt {
yo_kv::Move::Down
} else {
yo_kv::Move::Any
};
for i in (at..args.len()).step_by(2) {
score(args.get(i))?;
}
let key = args.get(1);
if incr {
let m = args.get(at + 1);
let mut stripe = db.hold(key);
let before = moved_from(&mut stripe, key, m)?;
let done = stripe.zincrby(key, m, score(args.get(at))?, opts)?;
drop(stripe);
match done {
Some(now) => {
out.double(now);
if before != Some(now) {
notify::fire(on, class::ZSET, "zincr", key);
}
}
None => out.nil(),
}
return Ok(());
}
let pairs = (at..args.len())
.step_by(2)
.map(|i| (score(args.get(i)).unwrap_or(0.0), args.get(i + 1)));
let (added, changed) = db.hold(key).zadd_counts(key, pairs, opts)?;
out.int(count(if opts.changed { added + changed } else { added }));
if added + changed > 0 {
notify::fire(on, class::ZSET, "zadd", key);
}
Ok(())
}
fn moved_from(
stripe: &mut Held<'_, yo_kv::Keyspace>,
key: &[u8],
member: &[u8],
) -> Result<Option<f64>> {
if !notify::armed() {
return Ok(None);
}
stripe.zscore(key, member)
}
fn rank(db: &Db, rev: bool, args: Args<'_>, out: &mut Out) -> Result<()> {
let withscore = match args.len() {
3 => false,
4 if args::is(args.get(3), b"withscore") => true,
4 => return Err(args::syntax()),
_ => return Err(args::wrong_arity(if rev { "zrevrank" } else { "zrank" })),
};
let key = args.get(1);
match db.hold(key).zrank(key, args.get(2), rev)? {
Some((at, s)) => {
if withscore {
out.array(2);
}
out.int(count(at));
if withscore {
out.double(s);
}
}
None if withscore => out.nil_array(),
None => out.nil(),
}
Ok(())
}
fn score(arg: &[u8]) -> Result<f64> {
yo_common::num::parse_f64(arg)
.ok_or_else(|| Error::new(yo_common::Code::Invalid, args::NOT_A_FLOAT))
}
fn bound(arg: &[u8]) -> Result<ZBound> {
let (open, digits) = match arg.split_first() {
Some((b'(', rest)) => (true, rest),
_ => (false, arg),
};
let Some(at) = yo_common::num::parse_f64(digits) else {
return Err(Error::new(yo_common::Code::Invalid, BAD_RANGE));
};
Ok(if open {
ZBound::open(at)
} else {
ZBound::closed(at)
})
}
fn lex(arg: &[u8]) -> Result<yo_kv::Lex<'_>> {
match arg.split_first() {
Some((b'-', b"")) => Ok(yo_kv::Lex::Min),
Some((b'+', b"")) => Ok(yo_kv::Lex::Max),
Some((b'[', rest)) => Ok(yo_kv::Lex::Incl(rest)),
Some((b'(', rest)) => Ok(yo_kv::Lex::Excl(rest)),
_ => Err(Error::new(yo_common::Code::Invalid, BAD_LEX)),
}
}
fn members(args: Args<'_>) -> impl Iterator<Item = &[u8]> + Clone {
(2..args.len()).map(move |i| args.get(i))
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Form {
Range,
Store,
RevRange,
ByScore { rev: bool },
ByLex { rev: bool },
}
impl Form {
fn of(name: &str) -> Form {
match name {
"zrangestore" => Form::Store,
"zrevrange" => Form::RevRange,
"zrangebyscore" => Form::ByScore { rev: false },
"zrevrangebyscore" => Form::ByScore { rev: true },
"zrangebylex" => Form::ByLex { rev: false },
"zrevrangebylex" => Form::ByLex { rev: true },
_ => Form::Range,
}
}
fn takes_mode(self) -> bool {
matches!(self, Form::Range | Form::Store)
}
}
fn parse_range<'a>(form: Form, args: Args<'a>, key: usize) -> Result<(Query<'a>, bool)> {
let (mut lo, mut hi) = (key + 1, key + 2);
let (mut byscore, mut bylex, mut rev) = (false, false, false);
match form {
Form::Range | Form::Store => {}
Form::RevRange => rev = true,
Form::ByScore { rev: r } => {
byscore = true;
rev = r;
}
Form::ByLex { rev: r } => {
bylex = true;
rev = r;
}
}
let mut withscores = false;
let mut limit: Option<(i64, i64)> = None;
let mut at = key + 3;
while at < args.len() {
let arg = args.get(at);
if form != Form::Store && args::is(arg, b"withscores") {
withscores = true;
} else if args::is(arg, b"limit") && at + 2 < args.len() {
limit = Some((args.int(at + 1)?, args.int(at + 2)?));
at += 2;
} else if form.takes_mode() && args::is(arg, b"byscore") {
byscore = true;
} else if form.takes_mode() && args::is(arg, b"bylex") {
bylex = true;
} else if form.takes_mode() && args::is(arg, b"rev") {
rev = true;
} else {
return Err(args::syntax());
}
at += 1;
}
if byscore && bylex {
return Err(args::syntax());
}
if rev && (byscore || bylex) {
core::mem::swap(&mut lo, &mut hi);
}
if withscores && bylex {
return Err(Error::new(yo_common::Code::Invalid, SCORES_NOT_BYLEX));
}
if limit.is_some() && !byscore && !bylex {
return Err(Error::new(yo_common::Code::Invalid, LIMIT_NEEDS_BY));
}
let mut q = if byscore {
Query::score(bound(args.get(lo))?, bound(args.get(hi))?)
} else if bylex {
Query::lex(lex(args.get(lo))?, lex(args.get(hi))?)
} else {
Query::rank(args.int(lo)?, args.int(hi)?)
}
.rev(rev);
if let Some((offset, take)) = limit {
q = q.limit(
usize::try_from(offset).unwrap_or(usize::MAX),
usize::try_from(take).ok(),
);
}
Ok((q, withscores))
}
fn op_of(name: &str) -> ZOp {
match name {
"zunion" | "zunionstore" => ZOp::Union,
"zinter" | "zinterstore" => ZOp::Inter,
_ => ZOp::Diff,
}
}
#[derive(PartialEq, Eq)]
enum Scores {
Allowed,
Refused,
}
struct Algebra {
from: usize,
to: usize,
weights: Vec<f64>,
agg: Aggregate,
withscores: bool,
}
impl Algebra {
fn parse(name: &str, op: ZOp, args: Args<'_>, at: usize, scores: Scores) -> Result<Algebra> {
let numkeys = args.int(at)?;
if numkeys < 1 {
return Err(Error::fmt(
yo_common::Code::Invalid,
format_args!("at least 1 input key is needed for '{name}' command"),
));
}
let numkeys = usize::try_from(numkeys).unwrap_or(usize::MAX);
let from = at + 1;
if numkeys > args.len() - from {
return Err(args::syntax());
}
let to = from + numkeys;
let mut got = Algebra {
from,
to,
weights: Vec::new(),
agg: Aggregate::Sum,
withscores: false,
};
let combines = op != ZOp::Diff;
let mut i = to;
while i < args.len() {
let rest = args.len() - i;
if combines && args::is(args.get(i), b"weights") && rest > numkeys {
got.weights.reserve_exact(numkeys);
for k in 0..numkeys {
let arg = args.get(i + 1 + k);
let Some(w) = yo_common::num::parse_f64(arg) else {
return Err(Error::new(yo_common::Code::Invalid, BAD_WEIGHT));
};
got.weights.push(w);
}
i += 1 + numkeys;
} else if combines && args::is(args.get(i), b"aggregate") && rest >= 2 {
let word = args.get(i + 1);
got.agg = if args::is(word, b"sum") {
Aggregate::Sum
} else if args::is(word, b"min") {
Aggregate::Min
} else if args::is(word, b"max") {
Aggregate::Max
} else {
return Err(args::syntax());
};
i += 2;
} else if scores == Scores::Allowed && args::is(args.get(i), b"withscores") {
got.withscores = true;
i += 1;
} else {
return Err(args::syntax());
}
}
Ok(got)
}
fn keys<'a>(&self, args: Args<'a>) -> impl Iterator<Item = &'a [u8]> + Clone {
(self.from..self.to).map(move |i| args.get(i))
}
}
fn intercard(db: &Db, args: Args<'_>) -> Result<usize> {
let numkeys = args.int(1)?;
if numkeys < 1 {
return Err(Error::new(
yo_common::Code::Invalid,
"at least 1 input key is needed for 'zintercard' command",
));
}
let numkeys = usize::try_from(numkeys).unwrap_or(usize::MAX);
if numkeys > args.len() - 2 {
return Err(args::syntax());
}
let end = 2 + numkeys;
let mut limit = 0usize;
if end < args.len() {
if args.len() != end + 2 || !args::is(args.get(end), b"limit") {
return Err(args::syntax());
}
limit = match parse_i64(args.get(end + 1)) {
Some(n) if n >= 0 => usize::try_from(n).unwrap_or(usize::MAX),
_ => return Err(Error::new(yo_common::Code::Invalid, BAD_LIMIT)),
};
}
db.zintercard((2..end).map(|i| args.get(i)), limit)
}
fn randmember(db: &Db, args: Args<'_>, out: &mut Out) -> Result<()> {
let key = args.get(1);
if args.len() == 2 {
let mut wrote = false;
db.hold(key).zrandmember(key, 1, |m, _| {
write_member(out, m);
wrote = true;
})?;
if !wrote {
out.nil();
}
return Ok(());
}
let withscores = match args.len() {
3 => false,
4 if args::is(args.get(3), b"withscores") => true,
_ => return Err(args::syntax()),
};
let want = args.int(2)?;
let nested = withscores && out.proto().is_resp3();
let start = out.len();
let mut n = 0;
db.hold(key).zrandmember(key, want, |m, sc| {
if nested {
out.array(2);
}
write_member(out, m);
if withscores {
out.double(sc);
}
n += 1;
})?;
out.close_array(start, if nested || !withscores { n } else { n * 2 });
Ok(())
}
fn zscan(db: &Db, args: Args<'_>, out: &mut Out) -> Result<()> {
let cursor = scan::parse_cursor(args.get(2))?;
let mut pattern = None;
let mut count = scan::COUNT;
let mut i = 3;
while i < args.len() {
let rest = args.len() - i;
if args::is(args.get(i), b"match") && rest >= 2 {
pattern = Some(args.get(i + 1));
} else if args::is(args.get(i), b"count") && rest >= 2 {
count = match args.int(i + 1)? {
n if n >= 1 => usize::try_from(n).unwrap_or(usize::MAX),
_ => return Err(args::syntax()),
};
} else if args::is(args.get(i), b"novalues") {
return Err(Error::new(yo_common::Code::Invalid, NOVALUES_IS_HSCAN));
} else {
return Err(args::syntax());
}
i += 2;
}
let key = args.get(1);
scan::reply(out, |out| {
let mut n = 0;
let next = db.hold(key).zscan(key, cursor, count, |m, sc| {
if !matches(pattern, m) {
return;
}
write_member(out, m);
out.bulk_double(sc);
n += 2;
})?;
Ok((next, n))
})
}
fn pop(
db: &Db,
on: usize,
name: &'static str,
end: ZEnd,
args: Args<'_>,
out: &mut Out,
) -> Result<()> {
let want = match args.len() {
2 => None,
3 => match args.int(2) {
Ok(n) if n >= 0 => Some(usize::try_from(n).unwrap_or(usize::MAX)),
_ => return Err(Error::new(yo_common::Code::Invalid, BAD_POP_COUNT)),
},
_ => return Err(args::syntax()),
};
let nested = want.is_some() && out.proto().is_resp3();
let start = out.len();
let mut n = 0;
let key = args.get(1);
db.hold(key).zpop(key, end, want.unwrap_or(1), |m, sc| {
if nested {
out.array(2);
}
write_member(out, m);
out.double(sc);
n += 1;
})?;
out.close_array(start, if nested { n } else { n * 2 });
if n > 0 {
notify::fire(on, class::ZSET, name, key);
notify::emptied(db, on, key);
}
Ok(())
}
fn mpop(db: &Db, on: usize, args: Args<'_>, out: &mut Out) -> Result<()> {
let (end, from, to, want) = parse_mpop(args, 1)?;
let mut held = db.hold_keys((from..to).map(|i| args.get(i)));
for i in from..to {
let key = args.get(i);
let stripe = held.stripe_mut(db.stripe_of(key));
if stripe.zcard(key)? == 0 {
continue;
}
out.array(2);
out.bulk(key);
let mark = out.len();
let n = stripe.zpop(key, end, want, |m, sc| {
out.array(2);
write_member(out, m);
out.double(sc);
})?;
let empty = notify::armed() && stripe.zcard(key)? == 0;
out.close_array(mark, n);
notify::fire(on, class::ZSET, popped(end), key);
if empty {
notify::fire(on, class::GENERIC, "del", key);
}
return Ok(());
}
out.nil_array();
Ok(())
}
pub(super) fn parse_mpop(args: Args<'_>, at: usize) -> Result<(ZEnd, usize, usize, usize)> {
let numkeys = match args.int(at) {
Ok(n) if n > 0 => usize::try_from(n).unwrap_or(usize::MAX),
_ => return Err(Error::new(yo_common::Code::Invalid, BAD_NUMKEYS)),
};
let from = at + 1;
if numkeys >= args.len() - from {
return Err(args::syntax());
}
let to = from + numkeys;
let end = end_of(args.get(to))?;
let mut want = 1usize;
if to + 1 < args.len() {
if args.len() != to + 3 || !args::is(args.get(to + 1), b"count") {
return Err(args::syntax());
}
want = match args.int(to + 2) {
Ok(n) if n > 0 => usize::try_from(n).unwrap_or(usize::MAX),
_ => return Err(Error::new(yo_common::Code::Invalid, BAD_MPOP_COUNT)),
};
}
Ok((end, from, to, want))
}
pub(super) fn end_of(arg: &[u8]) -> Result<ZEnd> {
if args::is(arg, b"min") {
Ok(ZEnd::Min)
} else if args::is(arg, b"max") {
Ok(ZEnd::Max)
} else {
Err(args::syntax())
}
}
pub(super) fn end_of_name(name: &str) -> ZEnd {
if name.ends_with("min") {
ZEnd::Min
} else {
ZEnd::Max
}
}
pub(super) const fn popped(end: ZEnd) -> &'static str {
match end {
ZEnd::Min => "zpopmin",
ZEnd::Max => "zpopmax",
}
}
#[inline]
fn matches(pattern: Option<&[u8]>, m: Member<'_>) -> bool {
let Some(pattern) = pattern else {
return true;
};
match m {
Member::Str(s) => glob_matches(pattern, s),
Member::Int(n) => {
let mut buf = [0u8; DIGITS_MAX];
glob_matches(pattern, i64_digits(&mut buf, n))
}
}
}
#[inline]
fn count(n: usize) -> i64 {
i64::try_from(n).unwrap_or(i64::MAX)
}
#[inline]
fn write_member(out: &mut Out, m: Member<'_>) {
match m {
Member::Int(n) => out.bulk_int(n),
Member::Str(s) => out.bulk(s),
}
}