use super::args::{self, Args, is, syntax};
use super::table::Spec;
use crate::reply::Out;
use yo_common::num::{parse_f64, parse_i64};
use yo_common::{Code, Error, Result, xxh3};
use yo_kv::{Compare, Exists, Expire, IncrEx, IncrExpire, Keyspace, Num, SetOptions, Str};
const BAD_DIGEST: &str = "must be exactly 16 hexadecimal characters";
const BAD_OFFSET: &str = "offset is out of range";
const BAD_NUMKEYS: &str = "invalid numkeys value";
const BAD_PAIRS: &str = "wrong number of key-value pairs";
const LEN_AND_IDX: &str = "If you want both the length and indexes, please just use IDX.";
pub(super) fn execute(db: &mut Keyspace, spec: &Spec, args: Args<'_>, out: &mut Out) -> Result<()> {
match spec.name {
"get" => match db.get(args.get(1))? {
Some(v) => write_str(out, v),
None => out.nil(),
},
"set" => set(db, args, out)?,
"getset" => {
let mut had = false;
db.set_with(
args.get(1),
args.get(2),
SetOptions::PLAIN.returning(),
|v| {
had = true;
write_str(out, v);
},
)?;
if !had {
out.nil();
}
}
"getdel" => {
if !db.getdel_with(args.get(1), |v| write_str(out, v))? {
out.nil();
}
}
"getex" => getex(db, args, out)?,
"setnx" => out.int(i64::from(db.setnx(args.get(1), args.get(2))?)),
"setex" => {
db.setex(args.get(1), args.int(2)?, args.get(3))?;
out.ok();
}
"psetex" => {
db.psetex(args.get(1), args.int(2)?, args.get(3))?;
out.ok();
}
"mset" => {
let n = pair_count(args, "mset")?;
db.mset(pairs(args, 1, n))?;
out.ok();
}
"msetnx" => {
let n = pair_count(args, "msetnx")?;
out.int(i64::from(db.msetnx(pairs(args, 1, n))?));
}
"mget" => {
out.array(args.len() - 1);
for i in 1..args.len() {
match db.mget_one(args.get(i)) {
Some(v) => write_str(out, v),
None => out.nil(),
}
}
}
"append" => out.int(count(db.append(args.get(1), args.get(2))?)),
"strlen" => out.int(count(db.strlen(args.get(1))?)),
"setrange" => {
let offset =
usize::try_from(args.int(2)?).map_err(|_| Error::new(Code::Invalid, BAD_OFFSET))?;
out.int(count(db.setrange(args.get(1), offset, args.get(3))?));
}
"getrange" | "substr" => {
let (start, end) = (args.int(2)?, args.int(3)?);
out.bulk(&db.getrange(args.get(1), start, end)?);
}
"incr" => out.int(db.incr(args.get(1))?),
"decr" => out.int(db.decr(args.get(1))?),
"incrby" => out.int(db.incrby(args.get(1), args.int(2)?)?),
"decrby" => out.int(db.decrby(args.get(1), args.int(2)?)?),
"incrbyfloat" => out.human_double(db.incrbyfloat(args.get(1), args.float(2)?)?),
"lcs" => lcs(db, args, out)?,
"msetex" => msetex(db, args, out)?,
"delex" => delex(db, args, out)?,
"digest" => match db.digest(args.get(1))? {
Some(h) => out.bulk(&xxh3::hex(h)),
None => out.nil(),
},
"increx" => increx(db, args, out)?,
_ => return Err(args::unknown_command(args)),
}
Ok(())
}
fn write_str(out: &mut Out, v: Str<'_>) {
match v {
Str::Int(n) => out.bulk_int(n),
Str::Bytes(b) => out.bulk(b),
}
}
fn count(n: usize) -> i64 {
i64::try_from(n).unwrap_or(i64::MAX)
}
fn pair_count(args: Args<'_>, name: &str) -> Result<usize> {
if args.len() % 2 != 1 {
return Err(args::wrong_arity(name));
}
Ok((args.len() - 1) / 2)
}
fn pairs<'a>(
args: Args<'a>,
from: usize,
count: usize,
) -> impl Iterator<Item = (&'a [u8], &'a [u8])> + Clone {
(0..count).map(move |i| (args.get(from + 2 * i), args.get(from + 2 * i + 1)))
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Unit {
Sec,
Millis,
SecAt,
MillisAt,
}
impl Unit {
fn parse(arg: &[u8]) -> Option<Unit> {
if is(arg, b"EX") {
Some(Unit::Sec)
} else if is(arg, b"PX") {
Some(Unit::Millis)
} else if is(arg, b"EXAT") {
Some(Unit::SecAt)
} else if is(arg, b"PXAT") {
Some(Unit::MillisAt)
} else {
None
}
}
fn is_seconds(self) -> bool {
matches!(self, Unit::Sec | Unit::SecAt)
}
fn is_relative(self) -> bool {
matches!(self, Unit::Sec | Unit::Millis)
}
}
fn deadline(unit: Unit, n: i64, now: u64, name: &str) -> Result<u64> {
if n <= 0 || (unit.is_seconds() && n > i64::MAX / 1000) {
return Err(args::invalid_expire(name));
}
let ms = if unit.is_seconds() { n * 1000 } else { n };
let ms = ms as u64;
Ok(if unit.is_relative() { now + ms } else { ms })
}
mod bits {
pub const NX: u16 = 1 << 0;
pub const XX: u16 = 1 << 1;
pub const EX: u16 = 1 << 2;
pub const PX: u16 = 1 << 3;
pub const EXAT: u16 = 1 << 4;
pub const PXAT: u16 = 1 << 5;
pub const KEEPTTL: u16 = 1 << 6;
pub const PERSIST: u16 = 1 << 7;
pub const IF: u16 = 1 << 8;
pub const ANY_EXPIRE: u16 = EX | PX | EXAT | PXAT;
}
fn unit_bit(unit: Unit) -> u16 {
match unit {
Unit::Sec => bits::EX,
Unit::Millis => bits::PX,
Unit::SecAt => bits::EXAT,
Unit::MillisAt => bits::PXAT,
}
}
fn set(db: &mut Keyspace, args: Args<'_>, out: &mut Out) -> Result<()> {
let (key, val) = (args.get(1), args.get(2));
let mut opts = SetOptions::PLAIN;
let mut seen = 0u16;
let mut expire: Option<(Unit, usize)> = None;
let mut cond: Option<(&[u8], usize)> = None;
let mut i = 3;
while i < args.len() {
let o = args.get(i);
if is(o, b"NX") && seen & (bits::XX | bits::IF) == 0 {
seen |= bits::NX;
opts = opts.if_missing();
i += 1;
} else if is(o, b"XX") && seen & (bits::NX | bits::IF) == 0 {
seen |= bits::XX;
opts = opts.if_present();
i += 1;
} else if is(o, b"GET") {
opts = opts.returning();
i += 1;
} else if is(o, b"KEEPTTL") && seen & bits::ANY_EXPIRE == 0 {
seen |= bits::KEEPTTL;
opts = opts.expiring(Expire::Keep);
i += 1;
} else if let Some(u) = Unit::parse(o)
&& seen & (bits::KEEPTTL | (bits::ANY_EXPIRE & !unit_bit(u))) == 0
&& i + 1 < args.len()
{
seen |= unit_bit(u);
expire = Some((u, i + 1));
i += 2;
} else if is_condition(o)
&& seen & (bits::NX | bits::XX) == 0
&& cond.is_none_or(|(k, _)| is(o, k))
&& i + 1 < args.len()
{
seen |= bits::IF;
cond = Some((o, i + 1));
i += 2;
} else {
return Err(syntax());
}
}
if let Some((u, at)) = expire {
let ms = deadline(u, args.int(at)?, db.clock().now_ms(), "set")?;
opts = opts.expiring(Expire::At(ms));
}
if let Some((keyword, at)) = cond {
opts.compare = Some(condition(keyword, args.get(at))?);
}
let mut had = false;
let done = db.set_with(key, val, opts, |v| {
had = true;
write_str(out, v);
})?;
if opts.get {
if !had {
out.nil();
}
} else if done.stored {
out.ok();
} else {
out.nil();
}
Ok(())
}
fn is_condition(arg: &[u8]) -> bool {
is(arg, b"IFEQ") || is(arg, b"IFNE") || is(arg, b"IFDEQ") || is(arg, b"IFDNE")
}
fn condition<'a>(keyword: &[u8], arg: &'a [u8]) -> Result<Compare<'a>> {
if is(keyword, b"IFEQ") {
Ok(Compare::Equal(arg))
} else if is(keyword, b"IFNE") {
Ok(Compare::NotEqual(arg))
} else {
let d = xxh3::from_hex(arg).ok_or_else(|| Error::new(Code::Invalid, BAD_DIGEST))?;
if is(keyword, b"IFDEQ") {
Ok(Compare::DigestEqual(d))
} else {
Ok(Compare::DigestNotEqual(d))
}
}
}
fn getex(db: &mut Keyspace, args: Args<'_>, out: &mut Out) -> Result<()> {
let key = args.get(1);
let mut seen = 0u16;
let mut expire: Option<(Unit, usize)> = None;
let mut i = 2;
while i < args.len() {
let o = args.get(i);
if is(o, b"PERSIST") && seen & bits::ANY_EXPIRE == 0 {
seen |= bits::PERSIST;
i += 1;
} else if let Some(u) = Unit::parse(o)
&& seen & (bits::PERSIST | (bits::ANY_EXPIRE & !unit_bit(u))) == 0
&& i + 1 < args.len()
{
seen |= unit_bit(u);
expire = Some((u, i + 1));
i += 2;
} else {
return Err(syntax());
}
}
if !db.exists(key) {
out.nil();
return Ok(());
}
let wanted = match expire {
Some((u, at)) => Expire::At(deadline(u, args.int(at)?, db.clock().now_ms(), "getex")?),
None if seen & bits::PERSIST != 0 => Expire::Clear,
None => Expire::Keep,
};
match db.getex(key, wanted)? {
Some(v) => write_str(out, v),
None => out.nil(),
}
Ok(())
}
fn msetex(db: &mut Keyspace, args: Args<'_>, out: &mut Out) -> Result<()> {
let n = parse_i64(args.get(1))
.filter(|&n| n > 0)
.and_then(|n| usize::try_from(n).ok())
.ok_or_else(|| Error::new(Code::Invalid, BAD_NUMKEYS))?;
let end = n
.checked_mul(2)
.and_then(|pairs| pairs.checked_add(2))
.filter(|&end| end <= args.len())
.ok_or_else(|| Error::new(Code::Invalid, BAD_PAIRS))?;
let mut seen = 0u16;
let mut exists = Exists::Always;
let mut expire = Expire::Clear;
let mut at: Option<(Unit, usize)> = None;
let mut i = end;
while i < args.len() {
let o = args.get(i);
if is(o, b"NX") && seen & bits::XX == 0 {
seen |= bits::NX;
exists = Exists::IfMissing;
i += 1;
} else if is(o, b"XX") && seen & bits::NX == 0 {
seen |= bits::XX;
exists = Exists::IfPresent;
i += 1;
} else if is(o, b"KEEPTTL") && seen & bits::ANY_EXPIRE == 0 {
seen |= bits::KEEPTTL;
expire = Expire::Keep;
i += 1;
} else if let Some(u) = Unit::parse(o)
&& seen & (bits::KEEPTTL | (bits::ANY_EXPIRE & !unit_bit(u))) == 0
&& i + 1 < args.len()
{
seen |= unit_bit(u);
at = Some((u, i + 1));
i += 2;
} else {
return Err(syntax());
}
}
if let Some((u, pos)) = at {
expire = Expire::At(deadline(u, args.int(pos)?, db.clock().now_ms(), "msetex")?);
}
out.int(i64::from(db.msetex(pairs(args, 2, n), exists, expire)?));
Ok(())
}
fn delex(db: &mut Keyspace, args: Args<'_>, out: &mut Out) -> Result<()> {
if args.len() != 2 && args.len() != 4 {
return Err(args::wrong_arity("delex"));
}
let compare = if args.len() == 4 {
if !is_condition(args.get(2)) {
return Err(syntax());
}
Some(condition(args.get(2), args.get(3))?)
} else {
None
};
out.int(i64::from(db.delex(args.get(1), compare)));
Ok(())
}
fn increx(db: &mut Keyspace, args: Args<'_>, out: &mut Out) -> Result<()> {
const BY: u16 = 1 << 9;
const SATURATE: u16 = 1 << 10;
const LBOUND: u16 = 1 << 11;
const UBOUND: u16 = 1 << 12;
const ENX: u16 = 1 << 13;
let key = args.get(1);
let mut seen = 0u16;
let mut opts = IncrEx::PLAIN;
let mut int_kind = true;
let mut by: Option<usize> = None;
let mut lower: Option<usize> = None;
let mut upper: Option<usize> = None;
let mut at: Option<(Unit, usize)> = None;
let mut i = 2;
while i < args.len() {
let o = args.get(i);
if (is(o, b"BYINT") || is(o, b"BYFLOAT")) && seen & BY == 0 && i + 1 < args.len() {
seen |= BY;
int_kind = is(o, b"BYINT");
by = Some(i + 1);
i += 2;
} else if is(o, b"SATURATE") && seen & SATURATE == 0 {
seen |= SATURATE;
opts = opts.saturating();
i += 1;
} else if is(o, b"LBOUND") && seen & LBOUND == 0 && i + 1 < args.len() {
seen |= LBOUND;
lower = Some(i + 1);
i += 2;
} else if is(o, b"UBOUND") && seen & UBOUND == 0 && i + 1 < args.len() {
seen |= UBOUND;
upper = Some(i + 1);
i += 2;
} else if is(o, b"PERSIST") && seen & (bits::ANY_EXPIRE | bits::PERSIST) == 0 {
seen |= bits::PERSIST;
i += 1;
} else if is(o, b"ENX") && seen & ENX == 0 {
seen |= ENX;
i += 1;
} else if let Some(u) = Unit::parse(o)
&& seen & (bits::PERSIST | bits::ANY_EXPIRE) == 0
&& i + 1 < args.len()
{
seen |= unit_bit(u);
at = Some((u, i + 1));
i += 2;
} else {
return Err(syntax());
}
}
if seen & ENX != 0 && at.is_none() {
return Err(Error::new(Code::Invalid, "ENX flag requires an expiration"));
}
if let Some(pos) = by {
opts = opts.by(number(args.get(pos), int_kind, "Increment")?);
}
let lower = lower
.map(|pos| number(args.get(pos), int_kind, "LBOUND"))
.transpose()?;
let upper = upper
.map(|pos| number(args.get(pos), int_kind, "UBOUND"))
.transpose()?;
opts = opts.between(lower, upper);
if let Some((u, pos)) = at {
let ms = deadline(u, args.int(pos)?, db.clock().now_ms(), "increx")?;
opts = opts.expiring(if seen & ENX != 0 {
IncrExpire::AtIfNone(ms)
} else {
IncrExpire::At(ms)
});
} else if seen & bits::PERSIST != 0 {
opts = opts.expiring(IncrExpire::Persist);
}
let done = db.increx(key, opts)?;
out.array(2);
write_num(out, done.value);
write_num(out, done.applied);
Ok(())
}
fn number(arg: &[u8], int_kind: bool, what: &str) -> Result<Num> {
if int_kind {
parse_i64(arg).map(Num::Int).ok_or_else(|| {
Error::fmt(
Code::Invalid,
format_args!("{what} is not an integer or out of range"),
)
})
} else {
parse_f64(arg)
.map(Num::Float)
.ok_or_else(|| Error::fmt(Code::Invalid, format_args!("{what} is not a valid float")))
}
}
fn write_num(out: &mut Out, n: Num) {
match n {
Num::Int(v) => out.int(v),
Num::Float(v) => out.double(v),
}
}
fn lcs(db: &mut Keyspace, args: Args<'_>, out: &mut Out) -> Result<()> {
let (a, b) = (args.get(1), args.get(2));
let (mut want_len, mut want_idx, mut with_len) = (false, false, false);
let mut minmatchlen = 0u32;
let mut i = 3;
while i < args.len() {
let o = args.get(i);
if is(o, b"LEN") {
want_len = true;
i += 1;
} else if is(o, b"IDX") {
want_idx = true;
i += 1;
} else if is(o, b"WITHMATCHLEN") {
with_len = true;
i += 1;
} else if is(o, b"MINMATCHLEN") && i + 1 < args.len() {
minmatchlen = u32::try_from(args.int(i + 1)?.max(0)).unwrap_or(u32::MAX);
i += 2;
} else {
return Err(syntax());
}
}
if want_len && want_idx {
return Err(Error::new(Code::Invalid, LEN_AND_IDX));
}
if want_idx {
let idx = db.lcs_idx(a, b, minmatchlen)?;
out.map(2);
out.bulk(b"matches");
out.array(idx.matches.len());
for m in &idx.matches {
out.array(if with_len { 3 } else { 2 });
out.array(2);
out.int(i64::from(m.a.0));
out.int(i64::from(m.a.1));
out.array(2);
out.int(i64::from(m.b.0));
out.int(i64::from(m.b.1));
if with_len {
out.int(i64::from(m.len));
}
}
out.bulk(b"len");
out.int(count(idx.len));
} else if want_len {
out.int(count(db.lcs_len(a, b)?));
} else {
out.bulk(&db.lcs(a, b)?);
}
Ok(())
}