use yo_common::num::{DIGITS_MAX, i64_digits};
use yo_common::{Code, Error, Result, glob_matches, parse_i64};
use yo_kv::{Db, Member};
use super::args::{self, Args};
use super::notify::{self, class};
use super::scan;
use super::table::Spec;
use crate::reply::Out;
const BAD_POP_COUNT: &str = "value is out of range, must be positive";
const BAD_NUMKEYS: &str = "numkeys should be greater than 0";
const TOO_MANY_KEYS: &str = "Number of keys can't be greater than number of args";
const BAD_LIMIT: &str = "LIMIT can't be negative";
pub(super) fn execute(
db: &Db,
on: usize,
spec: &Spec,
args: Args<'_>,
out: &mut Out,
) -> Result<()> {
let key = args.get(1);
match spec.name {
"sadd" => {
let added = db.hold(key).sadd(key, members(args))?;
out.int(count(added));
if added > 0 {
notify::fire(on, class::SET, "sadd", key);
}
}
"srem" => {
let gone = db.hold(key).srem(key, members(args))?;
out.int(count(gone));
if gone > 0 {
notify::fire(on, class::SET, "srem", key);
notify::emptied(db, on, key);
}
}
"scard" => out.int(count(db.hold(key).scard(key)?)),
"sismember" => out.int(i64::from(db.hold(key).sismember(key, args.get(2))?)),
"smismember" => db.hold(key).with_set(key, |set| {
out.array(args.len() - 2);
for m in members(args) {
out.int(i64::from(set.is_some_and(|s| s.contains(m))));
}
})?,
"smembers" => db.hold(key).with_set(key, |set| match set {
Some(s) => {
out.set(s.len());
for m in s.iter() {
write_member(out, m);
}
}
None => out.set(0),
})?,
"spop" => match args.len() {
2 => {
let start = out.len();
let mut got = false;
db.hold(key).spop_into(key, 1, |m| {
write_member(out, m);
got = true;
})?;
if !got {
out.nil();
}
debug_assert!(out.len() > start, "a reply went out either way");
if got {
notify::fire(on, class::SET, "spop", key);
notify::emptied(db, on, key);
}
}
3 => {
let want = pop_count(args.get(2))?;
let start = out.len();
let mut n = 0;
db.hold(key).spop_into(key, want, |m| {
write_member(out, m);
n += 1;
})?;
out.close_set(start, n);
if n > 0 {
notify::fire(on, class::SET, "spop", key);
notify::emptied(db, on, key);
}
}
_ => return Err(args::syntax()),
},
"srandmember" => match args.len() {
2 => db.hold(key).srandmember(key, |m| match m {
Some(m) => write_member(out, m),
None => out.nil(),
})?,
3 => {
let count = args.int(2)?;
let start = out.len();
let mut n = 0;
db.hold(key).srandmember_n(key, count, |m| {
write_member(out, m);
n += 1;
})?;
out.close_array(start, n);
}
_ => return Err(args::syntax()),
},
"smove" => {
let (dst, m) = (args.get(2), args.get(3));
let ask = notify::armed() && key != dst;
let had = ask && db.hold(dst).sismember(dst, m).unwrap_or(false);
let done = db.smove(key, dst, m)?;
out.int(i64::from(done));
if done && ask {
notify::fire(on, class::SET, "srem", key);
notify::emptied(db, on, key);
if !had {
notify::fire(on, class::SET, "sadd", dst);
}
}
}
"sscan" => scan(db, args, out)?,
"sinter" | "sunion" | "sdiff" => {
let start = out.len();
let mut n = 0;
let keys = keys(args, 1);
let mut take = |m: &[u8]| {
out.bulk(m);
n += 1;
};
match spec.name {
"sinter" => db.sinter(keys, 0, &mut take)?,
"sunion" => db.sunion(keys, 0, &mut take)?,
_ => db.sdiff(keys, 0, &mut take)?,
};
out.close_set(start, n);
}
"sintercard" | "sunioncard" | "sdiffcard" => {
let (end, limit) = cardinality(args)?;
let keys = rest(args, 2, end);
out.int(count(match spec.name {
"sintercard" => db.sintercard(keys, limit)?,
"sunioncard" => db.sunioncard(keys, limit)?,
_ => db.sdiffcard(keys, limit)?,
}));
}
"sinterstore" | "sunionstore" | "sdiffstore" => {
let had = notify::armed() && db.hold(key).exists(key);
let rest = keys(args, 2);
let stored = match spec.name {
"sinterstore" => db.sinterstore(key, rest)?,
"sunionstore" => db.sunionstore(key, rest)?,
_ => db.sdiffstore(key, rest)?,
};
out.int(count(stored));
if stored > 0 {
notify::fire(on, class::SET, spec.name, key);
} else if had {
notify::fire(on, class::GENERIC, "del", key);
}
}
other => unreachable!("{other} is not a set command"),
}
Ok(())
}
fn scan(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 {
return Err(args::syntax());
}
i += 2;
}
scan::reply(out, |out| {
let mut n = 0;
let key = args.get(1);
let next = db.hold(key).sscan(key, cursor, count, |m| {
if matches(pattern, m) {
write_member(out, m);
n += 1;
}
})?;
Ok((next, n))
})
}
fn pop_count(arg: &[u8]) -> Result<usize> {
match parse_i64(arg) {
Some(n) if n >= 0 => Ok(usize::try_from(n).unwrap_or(usize::MAX)),
_ => Err(Error::new(Code::Invalid, BAD_POP_COUNT)),
}
}
#[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))
}
}
}
fn cardinality(args: Args<'_>) -> Result<(usize, usize)> {
let numkeys = match parse_i64(args.get(1)) {
Some(n) if n > 0 => usize::try_from(n).unwrap_or(usize::MAX),
_ => return Err(Error::new(Code::Invalid, BAD_NUMKEYS)),
};
if numkeys > args.len() - 2 {
return Err(Error::new(Code::Invalid, TOO_MANY_KEYS));
}
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(Code::Invalid, BAD_LIMIT)),
};
}
Ok((end, limit))
}
#[inline]
fn members(args: Args<'_>) -> impl Iterator<Item = &[u8]> + Clone {
rest(args, 2, args.len())
}
#[inline]
fn keys(args: Args<'_>, from: usize) -> impl Iterator<Item = &[u8]> + Clone {
rest(args, from, args.len())
}
#[inline]
fn rest(args: Args<'_>, from: usize, end: usize) -> impl Iterator<Item = &[u8]> + Clone {
(from..end).map(move |i| args.get(i))
}
#[inline]
fn write_member(out: &mut Out, m: Member<'_>) {
match m {
Member::Int(n) => out.bulk_int(n),
Member::Str(s) => out.bulk(s),
}
}
#[inline]
fn count(n: usize) -> i64 {
i64::try_from(n).unwrap_or(i64::MAX)
}