use super::args::{self, Args, is};
use super::notify::{self, class};
use super::table::Spec;
use crate::reply::Out;
use yo_common::num::parse_i64;
use yo_common::{Code, Error, Result};
use yo_kv::bitmaps::{Sub, SubOp, Unit};
use yo_kv::bits::{Field, Op, Overflow};
use yo_kv::{Db, bitmaps};
const BAD_OFFSET: &str = "bit offset is not an integer or out of range";
const BAD_BIT: &str = "bit is not an integer or out of range";
const BAD_SEARCH_BIT: &str = "The bit argument must be 1 or 0.";
const BAD_TYPE: &str =
"Invalid bitfield type. Use something like i16 u8. Note that u64 is not supported but i64 is.";
const BAD_OVERFLOW: &str = "Invalid OVERFLOW type specified";
const RO_GET_ONLY: &str = "BITFIELD_RO only supports the GET subcommand";
pub(super) fn execute(
db: &Db,
on: usize,
spec: &Spec,
args: Args<'_>,
out: &mut Out,
) -> Result<()> {
match spec.name {
"setbit" => {
let offset = offset(args.get(2))?;
let bit = match parse_i64(args.get(3)) {
Some(0) => false,
Some(1) => true,
_ => return Err(Error::new(Code::Invalid, BAD_BIT)),
};
let key = args.get(1);
out.int(i64::from(db.hold(key).setbit(key, offset, bit)?));
notify::fire(on, class::STRING, "setbit", key);
}
"getbit" => {
let offset = offset(args.get(2))?;
let key = args.get(1);
out.int(i64::from(db.hold(key).getbit(key, offset)?));
}
"bitcount" => {
let range = range(args, 2)?;
let key = args.get(1);
let set = db.hold(key).bitcount(key, range)?;
out.int(i64::try_from(set).unwrap_or(i64::MAX));
}
"bitpos" => bitpos(db, args, out)?,
"bitop" => bitop(db, on, args, out)?,
"bitfield" => bitfield(db, on, args, out, false)?,
"bitfield_ro" => bitfield(db, on, args, out, true)?,
_ => return Err(args::syntax()),
}
Ok(())
}
fn bitpos(db: &Db, args: Args<'_>, out: &mut Out) -> Result<()> {
let bit = match args.int(2)? {
0 => false,
1 => true,
_ => return Err(Error::new(Code::Invalid, BAD_SEARCH_BIT)),
};
let (start, end, unit) = match args.len() {
3 => (None, None, Unit::Byte),
4 => (Some(args.int(3)?), None, Unit::Byte),
_ => {
let (start, end, unit) = range(args, 3)?.ok_or_else(args::syntax)?;
(Some(start), Some(end), unit)
}
};
let key = args.get(1);
out.int(db.hold(key).bitpos(key, bit, start, end, unit)?);
Ok(())
}
fn bitop(db: &Db, on: usize, args: Args<'_>, out: &mut Out) -> Result<()> {
let op = Op::parse(args.get(1)).ok_or_else(args::syntax)?;
let sources = args.len() - 3;
if op == Op::Not && sources != 1 {
return Err(bad(op, "must be called with a single source key."));
}
if sources < 2 && matches!(op, Op::Diff | Op::Diff1 | Op::AndOr) {
return Err(bad(op, "must be called with at least two source keys."));
}
let dst = args.get(2);
let had = notify::armed() && db.hold(dst).exists(dst);
let srcs = (3..args.len()).map(|i| args.get(i));
let len = db.bitop(op, dst, srcs)?;
if len > 0 {
notify::fire(on, class::STRING, "set", dst);
} else if had {
notify::fire(on, class::GENERIC, "del", dst);
}
out.int(count(len));
Ok(())
}
fn bad(op: Op, tail: &str) -> Error {
Error::fmt(Code::Invalid, format_args!("BITOP {} {tail}", op.name()))
}
fn bitfield(db: &Db, on: usize, args: Args<'_>, out: &mut Out, readonly: bool) -> Result<()> {
let mut grow: Option<usize> = None;
let mut at = 2;
let mut n = 0;
let mut over = Overflow::Wrap;
while at < args.len() {
let (sub, next) = parse(args, at, &mut over, readonly)?;
if let Some(sub) = sub {
n += 1;
if sub.op != SubOp::Get {
let need = bitmaps::reach(&sub);
grow = Some(grow.map_or(need, |had: usize| had.max(need)));
}
}
at = next;
}
out.array(n);
let key = args.get(1);
db.hold(key).bitfield_with(key, grow, |bytes| {
let mut at = 2;
let mut over = Overflow::Wrap;
while at < args.len() {
let (sub, next) = match parse(args, at, &mut over, readonly) {
Ok(step) => step,
Err(_) => break,
};
if let Some(sub) = sub {
match bitmaps::apply(bytes, sub) {
Some(n) => out.int(n),
None => out.nil(),
}
}
at = next;
}
})?;
if grow.is_some() {
notify::fire(on, class::STRING, "setbit", key);
}
Ok(())
}
fn parse(
args: Args<'_>,
at: usize,
on: &mut Overflow,
readonly: bool,
) -> Result<(Option<Sub>, usize)> {
let word = args.get(at);
if is(word, b"overflow") {
let arg = args.opt(at + 1).ok_or_else(args::syntax)?;
*on = Overflow::parse(arg).ok_or_else(|| Error::new(Code::Invalid, BAD_OVERFLOW))?;
return Ok((None, at + 2));
}
let get = is(word, b"get");
let set = is(word, b"set");
let incr = is(word, b"incrby");
if !get && !set && !incr {
return Err(args::syntax());
}
if readonly && !get {
return Err(Error::new(Code::Unsupported, RO_GET_ONLY));
}
let words = if get { 3 } else { 4 };
if at + words > args.len() {
return Err(args::syntax());
}
let field =
Field::parse(args.get(at + 1)).ok_or_else(|| Error::new(Code::Invalid, BAD_TYPE))?;
let bit = at_bit(args.get(at + 2), field)?;
let op = if get {
SubOp::Get
} else {
let n = args.int(at + 3)?;
if set { SubOp::Set(n) } else { SubOp::Incr(n) }
};
Ok((
Some(Sub {
op,
field,
at: bit,
on: *on,
}),
at + words,
))
}
fn at_bit(arg: &[u8], field: Field) -> Result<u64> {
let bad = || Error::new(Code::Invalid, BAD_OFFSET);
let (digits, scale) = match arg.split_first() {
Some((b'#', rest)) => (rest, u64::from(field.bits())),
_ => (arg, 1),
};
let n = parse_i64(digits).ok_or_else(bad)?;
let n = u64::try_from(n).map_err(|_| bad())?;
let bit = n.checked_mul(scale).ok_or_else(bad)?;
if field.last_bit(bit) > bitmaps::BIT_OFFSET_MAX {
return Err(bad());
}
Ok(bit)
}
fn offset(arg: &[u8]) -> Result<u64> {
match parse_i64(arg) {
Some(n) if n >= 0 && n as u64 <= bitmaps::BIT_OFFSET_MAX => Ok(n as u64),
_ => Err(Error::new(Code::Invalid, BAD_OFFSET)),
}
}
fn range(args: Args<'_>, at: usize) -> Result<Option<(i64, i64, Unit)>> {
if args.len() <= at {
return Ok(None);
}
if args.len() < at + 2 || args.len() > at + 3 {
return Err(args::syntax());
}
let start = args.int(at)?;
let end = args.int(at + 1)?;
let unit = match args.opt(at + 2) {
None => Unit::Byte,
Some(w) if is(w, b"byte") => Unit::Byte,
Some(w) if is(w, b"bit") => Unit::Bit,
Some(_) => return Err(args::syntax()),
};
Ok(Some((start, end, unit)))
}
fn count(n: usize) -> i64 {
i64::try_from(n).unwrap_or(i64::MAX)
}