use yo_common::{Code, Error, Result, glob_matches, parse_i64};
use yo_kv::hash::Text;
use yo_kv::{Ask, Cond, Db, Exists, Expire, Keyspace, MAX_AT};
use super::args::{self, Args};
use super::scan;
use super::table::Spec;
use crate::reply::Out;
const NOT_AN_INT: &str = "value is not an integer or out of range";
const NOT_A_FLOAT: &str = "value is not a valid float";
const BAD_EXPIRE: &str = "invalid expire time, must be >= 0";
const FIELD_COUNT: &str = "wrong number of arguments";
const BAD_NUMFIELDS: &str = "Parameter `numFields` should be greater than 0";
const SECOND: i64 = 1000;
const DEL_BAD_COUNT: &str = "Number of fields must be a positive integer";
const DEL_MISMATCH: &str = "The `numfields` parameter must match the number of arguments";
const DEL_NO_FIELDS: &str = "Mandatory argument FIELDS is missing or not at the right position";
const EX_BAD_COUNT: &str = "invalid number of fields";
const EX_MISMATCH: &str = "wrong number of arguments";
const GETEX_ONE_OF: &str = "Only one of EX, PX, EXAT, PXAT or PERSIST arguments can be specified";
const SETEX_ONE_OF: &str = "Only one of EX, PX, EXAT, PXAT or KEEPTTL arguments can be specified";
const SETEX_ONE_COND: &str = "Only one of FXX or FNX arguments can be specified";
pub(super) fn execute(db: &Db, spec: &Spec, args: Args<'_>, out: &mut Out) -> Result<()> {
let mut held = db.hold(args.get(1));
let db = &mut *held;
match spec.name {
"hset" | "hmset" => {
if args.len() < 4 || !args.len().is_multiple_of(2) {
return Err(args::wrong_arity(spec.name));
}
let added = db.hset(args.get(1), pairs(args))?;
if spec.name == "hset" {
out.int(count(added));
} else {
out.ok();
}
}
"hsetnx" => out.int(i64::from(db.hsetnx(
args.get(1),
args.get(2),
args.get(3),
)?)),
"hget" => db.hget(args.get(1), args.get(2), |t| match t {
Some(t) => write_text(out, t),
None => out.nil(),
})?,
"hdel" => out.int(count(db.hdel(args.get(1), fields(args, 2))?)),
"hlen" => out.int(count(db.hlen(args.get(1))?)),
"hexists" => out.int(i64::from(db.hexists(args.get(1), args.get(2))?)),
"hstrlen" => out.int(count(db.hstrlen(args.get(1), args.get(2))?)),
"hmget" => {
out.array(args.len() - 2);
db.hmget(args.get(1), fields(args, 2), |t| match t {
Some(t) => write_text(out, t),
None => out.nil(),
})?;
}
"hgetall" => db.with_hash(args.get(1), |hash| match hash {
Some(h) => {
out.map(h.len());
for (field, value) in h.iter() {
write_text(out, field);
write_text(out, value);
}
}
None => out.map(0),
})?,
"hkeys" | "hvals" => {
let want_keys = spec.name == "hkeys";
db.with_hash(args.get(1), |hash| match hash {
Some(h) => {
out.array(h.len());
for (field, value) in h.iter() {
write_text(out, if want_keys { field } else { value });
}
}
None => out.array(0),
})?;
}
"hincrby" => out.int(db.hincrby(args.get(1), args.get(2), incr_int(args.get(3))?)?),
"hincrbyfloat" => {
let by = incr_float(args.get(3))?;
out.human_double(db.hincrbyfloat(args.get(1), args.get(2), by)?);
}
"hrandfield" => randfield(db, args, out)?,
"hscan" => scan(db, args, out)?,
"hexpire" | "hpexpire" | "hexpireat" | "hpexpireat" => expire(db, spec.name, args, out)?,
"httl" | "hpttl" | "hexpiretime" | "hpexpiretime" | "hpersist" => {
ask(db, spec.name, args, out)?;
}
"hgetdel" => getdel(db, args, out)?,
"hgetex" => getex(db, args, out)?,
"hsetex" => setex(db, args, out)?,
other => unreachable!("{other} is not a hash command"),
}
Ok(())
}
fn randfield(db: &mut Keyspace, args: Args<'_>, out: &mut Out) -> Result<()> {
match args.len() {
2 => db.hrandfield(args.get(1), |pair| match pair {
Some((field, _)) => write_text(out, field),
None => out.nil(),
})?,
3 | 4 => {
let with_values = match args.len() {
4 if args::is(args.get(3), b"withvalues") => true,
4 => return Err(args::syntax()),
_ => false,
};
let n = args.int(2)?;
let nested = with_values && out.proto().is_resp3();
let start = out.len();
let mut written = 0;
db.hrandfield_n(args.get(1), n, |field, value| {
if nested {
out.array(2);
}
write_text(out, field);
written += 1;
if with_values {
write_text(out, value);
if !nested {
written += 1;
}
}
})?;
out.close_array(start, written);
}
_ => return Err(args::syntax()),
}
Ok(())
}
fn scan(db: &mut Keyspace, 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 novalues = false;
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") {
novalues = true;
i += 1;
continue;
} else {
return Err(args::syntax());
}
i += 2;
}
scan::reply(out, |out| {
let mut n = 0;
let next = db.hscan(args.get(1), cursor, count, |field, value| {
if !matches(pattern, field) {
return;
}
write_text(out, field);
n += 1;
if !novalues {
write_text(out, value);
n += 1;
}
})?;
Ok((next, n))
})
}
fn expire(db: &mut Keyspace, name: &str, args: Args<'_>, out: &mut Out) -> Result<()> {
let relative = matches!(name, "hexpire" | "hpexpire");
let scale = if matches!(name, "hpexpire" | "hpexpireat") {
1
} else {
SECOND
};
let at = moment(args.int(2)?, scale, relative, name, db.clock().now_ms())?;
let (cond, from) = condition(args, 3)?;
let fields = field_list(args, from, name)?;
out.array(fields.len());
db.hexpire(args.get(1), at, cond, fields.iter(args), |applied| {
out.int(applied as i64);
})
}
fn moment(by: i64, scale: i64, relative: bool, name: &str, now: u64) -> Result<u64> {
if by < 0 {
return Err(Error::new(Code::Invalid, BAD_EXPIRE));
}
by.checked_mul(scale)
.and_then(|ms| {
if relative {
ms.checked_add(now as i64)
} else {
Some(ms)
}
})
.and_then(|ms| u64::try_from(ms).ok())
.filter(|&ms| ms <= MAX_AT)
.ok_or_else(|| out_of_range(name))
}
fn ask(db: &mut Keyspace, name: &str, args: Args<'_>, out: &mut Out) -> Result<()> {
let fields = field_list(args, 2, name)?;
out.array(fields.len());
let now = db.clock().now_ms();
if name == "hpersist" {
return db.hpersist(args.get(1), fields.iter(args), |asked| {
out.int(match asked {
Ask::Missing => -2,
Ask::NoDeadline => -1,
Ask::At(_) => 1,
});
});
}
let left = name == "httl" || name == "hpttl";
let millis = name == "hpttl" || name == "hpexpiretime";
db.httl(args.get(1), fields.iter(args), |asked| {
let ms = if left {
asked.remaining_ms(now)
} else {
match asked {
Ask::Missing => -2,
Ask::NoDeadline => -1,
Ask::At(at) => at as i64,
}
};
out.int(if millis || ms < 0 {
ms
} else {
(ms + SECOND - 1) / SECOND
});
})
}
fn getdel(db: &mut Keyspace, args: Args<'_>, out: &mut Out) -> Result<()> {
if !args::is(args.get(2), b"fields") {
return Err(Error::new(Code::Invalid, DEL_NO_FIELDS));
}
let fields = ex_field_list(args, 2, 1, DEL_BAD_COUNT, DEL_MISMATCH)?;
out.array(fields.len());
db.hgetdel(args.get(1), fields.iter(args), |t| match t {
Some(t) => write_text(out, t),
None => out.nil(),
})
}
fn getex(db: &mut Keyspace, args: Args<'_>, out: &mut Out) -> Result<()> {
let opts = options(db.clock().now_ms(), args, "hgetex")?;
let fields = ex_field_list(args, opts.fields_at, 1, EX_BAD_COUNT, EX_MISMATCH)?;
out.array(fields.len());
db.hgetex(args.get(1), opts.expire, fields.iter(args), |t| match t {
Some(t) => write_text(out, t),
None => out.nil(),
})
}
fn setex(db: &mut Keyspace, args: Args<'_>, out: &mut Out) -> Result<()> {
let opts = options(db.clock().now_ms(), args, "hsetex")?;
let fields = ex_field_list(args, opts.fields_at, 2, EX_BAD_COUNT, EX_MISMATCH)?;
let wrote = db.hsetex(args.get(1), opts.exists, opts.expire, fields.pairs(args))?;
out.int(i64::from(wrote));
Ok(())
}
#[derive(Debug, Clone, Copy)]
struct Options {
exists: Exists,
expire: Expire,
fields_at: usize,
}
fn options(now: u64, args: Args<'_>, name: &str) -> Result<Options> {
let setting = name == "hsetex";
let mut opts = Options {
exists: Exists::Always,
expire: if setting { Expire::Clear } else { Expire::Keep },
fields_at: 0,
};
let mut had_expire = false;
let mut had_exists = false;
let one_of = if setting { SETEX_ONE_OF } else { GETEX_ONE_OF };
let mut i = 2;
while i < args.len() {
let arg = args.get(i);
if args::is(arg, b"fields") {
opts.fields_at = i;
return Ok(opts);
}
let clause = match arg {
a if args::is(a, b"ex") => Some((SECOND, true)),
a if args::is(a, b"px") => Some((1, true)),
a if args::is(a, b"exat") => Some((SECOND, false)),
a if args::is(a, b"pxat") => Some((1, false)),
_ => None,
};
if let Some((scale, relative)) = clause {
if had_expire {
return Err(Error::new(Code::Invalid, one_of));
}
had_expire = true;
opts.expire = Expire::At(moment(args.int(i + 1)?, scale, relative, name, now)?);
i += 2;
continue;
}
if (setting && args::is(arg, b"keepttl")) || (!setting && args::is(arg, b"persist")) {
if had_expire {
return Err(Error::new(Code::Invalid, one_of));
}
had_expire = true;
opts.expire = if setting { Expire::Keep } else { Expire::Clear };
i += 1;
continue;
}
if setting && (args::is(arg, b"fnx") || args::is(arg, b"fxx")) {
if had_exists {
return Err(Error::new(Code::Invalid, SETEX_ONE_COND));
}
had_exists = true;
opts.exists = if args::is(arg, b"fnx") {
Exists::IfMissing
} else {
Exists::IfPresent
};
i += 1;
continue;
}
return Err(unknown(arg));
}
Err(args::wrong_arity(name))
}
fn unknown(arg: &[u8]) -> Error {
Error::fmt(
Code::Invalid,
format_args!("unknown argument: {}", String::from_utf8_lossy(arg)),
)
}
fn ex_field_list(
args: Args<'_>,
at: usize,
step: usize,
bad_count: &str,
mismatch: &str,
) -> Result<Fields> {
let Some(len) = parse_i64(args.get(at + 1))
.filter(|&n| n > 0)
.and_then(|n| usize::try_from(n).ok())
else {
return Err(Error::new(Code::Invalid, bad_count));
};
let from = at + 2;
if len.checked_mul(step).and_then(|w| from.checked_add(w)) != Some(args.len()) {
return Err(Error::new(Code::Invalid, mismatch));
}
Ok(Fields { from, len })
}
#[derive(Debug, Clone, Copy)]
struct Fields {
from: usize,
len: usize,
}
impl Fields {
#[inline]
const fn len(self) -> usize {
self.len
}
#[inline]
fn iter(self, args: Args<'_>) -> impl Iterator<Item = &[u8]> {
(self.from..self.from + self.len).map(move |i| args.get(i))
}
#[inline]
fn pairs(self, args: Args<'_>) -> impl Iterator<Item = (&[u8], &[u8])> + Clone {
(0..self.len).map(move |k| {
let at = self.from + k * 2;
(args.get(at), args.get(at + 1))
})
}
}
fn field_list(args: Args<'_>, at: usize, name: &str) -> Result<Fields> {
if !args::is(args.get(at), b"fields") {
return Err(args::wrong_arity(name));
}
let n = args.int(at + 1)?;
if n < 1 {
return Err(Error::new(Code::Invalid, BAD_NUMFIELDS));
}
let len = usize::try_from(n).unwrap_or(usize::MAX);
let from = at + 2;
if args.len() != from + len {
return Err(Error::new(Code::Invalid, FIELD_COUNT));
}
Ok(Fields { from, len })
}
fn condition(args: Args<'_>, at: usize) -> Result<(Cond, usize)> {
let cond = match args.get(at) {
a if args::is(a, b"nx") => Cond::NotSet,
a if args::is(a, b"xx") => Cond::AlreadySet,
a if args::is(a, b"gt") => Cond::Greater,
a if args::is(a, b"lt") => Cond::Less,
_ => return Ok((Cond::Always, at)),
};
Ok((cond, at + 1))
}
fn out_of_range(name: &str) -> Error {
Error::new(
Code::Invalid,
format!("invalid expire time in '{name}' command"),
)
}
#[inline]
fn pairs(args: Args<'_>) -> impl Iterator<Item = (&[u8], &[u8])> + Clone {
(3..args.len())
.step_by(2)
.map(move |i| (args.get(i - 1), args.get(i)))
}
#[inline]
fn fields(args: Args<'_>, from: usize) -> impl Iterator<Item = &[u8]> {
(from..args.len()).map(move |i| args.get(i))
}
#[inline]
fn write_text(out: &mut Out, t: Text<'_>) {
match t {
Text::Int(n) => out.bulk_int(n),
Text::Str(s) => out.bulk(s),
}
}
fn incr_int(arg: &[u8]) -> Result<i64> {
parse_i64(arg).ok_or_else(|| Error::new(Code::Invalid, NOT_AN_INT))
}
fn incr_float(arg: &[u8]) -> Result<f64> {
yo_common::num::parse_f64(arg).ok_or_else(|| Error::new(Code::Invalid, NOT_A_FLOAT))
}
#[inline]
fn matches(pattern: Option<&[u8]>, field: Text<'_>) -> bool {
let Some(p) = pattern else {
return true;
};
match field {
Text::Str(s) => glob_matches(p, s),
Text::Int(n) => {
let mut buf = [0u8; yo_common::num::DIGITS_MAX];
glob_matches(p, yo_common::num::i64_digits(&mut buf, n))
}
}
}
#[inline]
fn count(n: usize) -> i64 {
i64::try_from(n).unwrap_or(i64::MAX)
}