use super::args::{self, Args, is, syntax};
use super::notify::{self, class};
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::strings::check_len;
use yo_kv::{Compare, Db, 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: &Db,
on: usize,
spec: &Spec,
args: Args<'_>,
out: &mut Out,
) -> Result<()> {
match spec.name {
"mset" => return mset(db, on, args, out),
"msetnx" => return msetnx(db, on, args, out),
"mget" => return mget(db, args, out),
"msetex" => return msetex(db, on, args, out),
"lcs" => return lcs(db, args, out),
_ => {}
}
let mut held = db.hold(args.get(1));
let db = &mut *held;
match spec.name {
"get" => match db.get(args.get(1))? {
Some(v) => write_str(out, v),
None => out.nil(),
},
"set" => set(db, on, 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();
}
notify::fire(on, class::STRING, "set", args.get(1));
}
"getdel" => {
if db.getdel_with(args.get(1), |v| write_str(out, v))? {
notify::fire(on, class::GENERIC, "del", args.get(1));
} else {
out.nil();
}
}
"getex" => getex(db, on, args, out)?,
"setnx" => {
let stored = db.setnx(args.get(1), args.get(2))?;
if stored {
notify::fire(on, class::STRING, "set", args.get(1));
}
out.int(i64::from(stored));
}
"setex" => {
db.setex(args.get(1), args.int(2)?, args.get(3))?;
timed_set(on, args.get(1));
out.ok();
}
"psetex" => {
db.psetex(args.get(1), args.int(2)?, args.get(3))?;
timed_set(on, args.get(1));
out.ok();
}
"append" => {
out.int(count(db.append(args.get(1), args.get(2))?));
notify::fire(on, class::STRING, "append", args.get(1));
}
"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))?));
if !args.get(3).is_empty() {
notify::fire(on, class::STRING, "setrange", args.get(1));
}
}
"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))?);
notify::fire(on, class::STRING, "incrby", args.get(1));
}
"decr" => {
out.int(db.decr(args.get(1))?);
notify::fire(on, class::STRING, "incrby", args.get(1));
}
"incrby" => {
out.int(db.incrby(args.get(1), args.int(2)?)?);
notify::fire(on, class::STRING, "incrby", args.get(1));
}
"decrby" => {
out.int(db.decrby(args.get(1), args.int(2)?)?);
notify::fire(on, class::STRING, "incrby", args.get(1));
}
"incrbyfloat" => {
out.human_double(db.incrbyfloat(args.get(1), args.float(2)?)?);
notify::fire(on, class::STRING, "incrbyfloat", args.get(1));
}
"delex" => delex(db, on, args, out)?,
"digest" => match db.digest(args.get(1))? {
Some(h) => out.bulk(&xxh3::hex(h)),
None => out.nil(),
},
"increx" => increx(db, on, 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)))
}
fn timed_set(on: usize, key: &[u8]) {
notify::fire(on, class::STRING, "set", key);
notify::fire(on, class::GENERIC, "expire", key);
}
#[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, on: usize, 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 done.stored {
if expire.is_some() {
timed_set(on, key);
} else {
notify::fire(on, class::STRING, "set", key);
}
}
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, on: usize, 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,
};
let had_deadline = wanted == Expire::Clear && matches!(db.deadline_of(key), yo_kv::Ask::At(_));
match db.getex(key, wanted)? {
Some(v) => write_str(out, v),
None => out.nil(),
}
match wanted {
Expire::At(_) => notify::fire(on, class::GENERIC, "expire", key),
Expire::Clear if had_deadline => notify::fire(on, class::GENERIC, "persist", key),
_ => {}
}
Ok(())
}
fn msetex(db: &Db, on: usize, 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.now_ms(), "msetex")?);
}
for (k, v) in pairs(args, 2, n) {
check_len(k, v.len())?;
}
let mut held = db.hold_keys(pairs(args, 2, n).map(|(k, _)| k));
let mut asked = |k: &[u8]| held.stripe_mut(db.stripe_of(k)).exists(k);
let allowed = match exists {
Exists::Always => true,
Exists::IfMissing => pairs(args, 2, n).all(|(k, _)| !asked(k)),
Exists::IfPresent => pairs(args, 2, n).all(|(k, _)| asked(k)),
};
if !allowed {
out.int(0);
return Ok(());
}
for (k, v) in pairs(args, 2, n) {
held.stripe_mut(db.stripe_of(k)).msetex(
core::iter::once((k, v)),
Exists::Always,
expire,
)?;
if let Expire::At(_) = expire {
notify::fire(on, class::GENERIC, "expire", k);
}
notify::fire(on, class::STRING, "set", k);
}
out.int(1);
Ok(())
}
fn mset(db: &Db, on: usize, args: Args<'_>, out: &mut Out) -> Result<()> {
let n = pair_count(args, "mset")?;
for (k, v) in pairs(args, 1, n) {
check_len(k, v.len())?;
}
let mut held = db.hold_keys(pairs(args, 1, n).map(|(k, _)| k));
for (k, v) in pairs(args, 1, n) {
held.stripe_mut(db.stripe_of(k))
.mset(core::iter::once((k, v)))?;
notify::fire(on, class::STRING, "set", k);
}
out.ok();
Ok(())
}
fn msetnx(db: &Db, on: usize, args: Args<'_>, out: &mut Out) -> Result<()> {
let n = pair_count(args, "msetnx")?;
for (k, v) in pairs(args, 1, n) {
check_len(k, v.len())?;
}
let mut held = db.hold_keys(pairs(args, 1, n).map(|(k, _)| k));
if pairs(args, 1, n).any(|(k, _)| held.stripe_mut(db.stripe_of(k)).exists(k)) {
out.int(0);
return Ok(());
}
for (k, v) in pairs(args, 1, n) {
held.stripe_mut(db.stripe_of(k))
.mset(core::iter::once((k, v)))?;
notify::fire(on, class::STRING, "set", k);
}
out.int(1);
Ok(())
}
fn mget(db: &Db, args: Args<'_>, out: &mut Out) -> Result<()> {
let mut held = db.hold_keys((1..args.len()).map(|i| args.get(i)));
out.array(args.len() - 1);
for i in 1..args.len() {
let key = args.get(i);
match held.stripe_mut(db.stripe_of(key)).mget_one(key) {
Some(v) => write_str(out, v),
None => out.nil(),
}
}
Ok(())
}
fn delex(db: &mut Keyspace, on: usize, 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
};
let gone = db.delex(args.get(1), compare);
if gone {
notify::fire(on, class::GENERIC, "del", args.get(1));
}
out.int(i64::from(gone));
Ok(())
}
fn increx(db: &mut Keyspace, on: usize, 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 kept =
seen & ENX != 0 && notify::armed() && matches!(db.deadline_of(key), yo_kv::Ask::At(_));
let done = db.increx(key, opts)?;
notify::fire(
on,
class::STRING,
if int_kind { "incrby" } else { "incrbyfloat" },
key,
);
if at.is_some() && !kept {
notify::fire(on, class::GENERIC, "expire", key);
} else if seen & bits::PERSIST != 0 {
notify::fire(on, class::GENERIC, "persist", key);
}
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: &Db, 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));
}
let x = db.hold(a).string_copy(a)?;
let y = db.hold(b).string_copy(b)?;
if want_idx {
let idx = yo_kv::lcs::idx(&x, &y, 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(yo_kv::lcs::len(&x, &y)?));
} else {
out.bulk(&yo_kv::lcs::string(&x, &y)?);
}
Ok(())
}