use crate::bits::{self, Field, Op, Overflow};
use crate::db::Db;
use crate::keyspace::Keyspace;
use crate::strings::{STRING_MAX, check_len};
use crate::value::{self, Kind, Str};
use yo_common::num::{self, DIGITS_MAX};
use yo_common::{Code, Error, Result};
use yo_index::RawMap;
const BAD_BIT_OFFSET: &str = "bit offset is not an integer or out of range";
const TOO_LONG: &str = "string exceeds maximum allowed size (proto-max-bulk-len)";
pub const BIT_OFFSET_MAX: u64 = 4 * 1024 * 1024 * 1024 - 1;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum Unit {
#[default]
Byte,
Bit,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Sub {
pub op: SubOp,
pub field: Field,
pub at: u64,
pub on: Overflow,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum SubOp {
Get,
Set(i64),
Incr(i64),
}
impl SubOp {
const fn writes(self) -> bool {
!matches!(self, SubOp::Get)
}
}
impl Keyspace {
pub fn getbit(&mut self, key: &[u8], offset: u64) -> Result<bool> {
if offset > BIT_OFFSET_MAX {
return Err(Error::new(Code::Invalid, BAD_BIT_OFFSET));
}
self.reap(key);
self.string_only(key)?;
self.warm(key)?;
let mut digits = [0u8; DIGITS_MAX];
let bytes = self.bitmap(key, &mut digits);
let byte = (offset / 8) as usize;
Ok(bytes.get(byte).is_some_and(|b| b & mask(offset) != 0))
}
pub fn setbit(&mut self, key: &[u8], offset: u64, bit: bool) -> Result<bool> {
if offset > BIT_OFFSET_MAX {
return Err(Error::new(Code::Invalid, BAD_BIT_OFFSET));
}
let byte = (offset / 8) as usize;
check_len(key, byte + 1)?;
self.thaw(key)?;
let now = self.clock.now_ms();
let hash = RawMap::hash_of(key);
let mut dead = false;
if let Some(rec) = self.map.value_mut_hashed(hash, key) {
if value::kind(rec) != Kind::String {
return Err(crate::keyspace::wrong_type());
}
if value::is_expired(rec, now) {
dead = true;
} else if let Some(b) = value::raw_in_place(rec).and_then(|it| it.get_mut(byte)) {
let had = *b & mask(offset) != 0;
if bit {
*b |= mask(offset);
} else {
*b &= !mask(offset);
}
return Ok(had);
}
}
if dead {
self.drop_key(key);
self.expired += 1;
}
let mut bytes = std::mem::take(&mut self.scratch);
bytes.clear();
let deadline = match self.map.get(key) {
Some(rec) => {
value::read(rec).write_to(&mut bytes);
value::expire_at(rec)
}
None => None,
};
if bytes.len() <= byte {
bytes.resize(byte + 1, 0);
}
let had = bytes[byte] & mask(offset) != 0;
if bit {
bytes[byte] |= mask(offset);
} else {
bytes[byte] &= !mask(offset);
}
self.store_raw(key, &bytes, deadline);
self.scratch = bytes;
Ok(had)
}
pub fn bitcount(&mut self, key: &[u8], range: Option<(i64, i64, Unit)>) -> Result<u64> {
self.reap(key);
self.string_only(key)?;
self.warm(key)?;
let mut digits = [0u8; DIGITS_MAX];
let bytes = self.bitmap(key, &mut digits);
let Some((start, end, unit)) = range else {
return Ok(bits::count(bytes));
};
match window(bytes.len(), start, end, unit) {
Some((from, to)) => Ok(bits::count_range(bytes, from, to)),
None => Ok(0),
}
}
pub fn bitpos(
&mut self,
key: &[u8],
bit: bool,
start: Option<i64>,
end: Option<i64>,
unit: Unit,
) -> Result<i64> {
self.reap(key);
self.string_only(key)?;
self.warm(key)?;
let here = self.map.get(key).is_some();
let mut digits = [0u8; DIGITS_MAX];
let bytes = self.bitmap(key, &mut digits);
if bytes.is_empty() {
return Ok(if !bit && !here { 0 } else { -1 });
}
let all = bytes.len() as u64 * 8;
let (from, to) = match (start, end) {
(None, _) => (0, all),
(Some(s), None) => match window(bytes.len(), s, -1, unit) {
Some(r) => r,
None => return Ok(-1),
},
(Some(s), Some(e)) => match window(bytes.len(), s, e, unit) {
Some(r) => r,
None => return Ok(-1),
},
};
match bits::find(bytes, bit, from, to) {
Some(at) => Ok(at as i64),
None if !bit && end.is_none() => Ok(all as i64),
None => Ok(-1),
}
}
pub fn bitop<'k, I>(&mut self, op: Op, dest: &[u8], srcs: I) -> Result<usize>
where
I: Iterator<Item = &'k [u8]> + Clone,
{
for src in srcs.clone() {
self.reap(src);
self.string_only(src)?;
self.thaw(src)?;
}
let mut flat = std::mem::take(&mut self.scratch);
let mut ends = std::mem::take(&mut self.rows);
flat.clear();
ends.clear();
let mut digits = [0u8; DIGITS_MAX];
for src in srcs.clone() {
let bytes = self.bitmap(src, &mut digits);
flat.extend_from_slice(bytes);
ends.push(flat.len());
}
let len = bits::width(parts(&flat, &ends));
if len > STRING_MAX {
self.scratch = flat;
self.rows = ends;
return Err(Error::new(Code::Invalid, TOO_LONG));
}
let split = flat.len();
flat.resize(split + len, 0);
let (read, write) = flat.split_at_mut(split);
bits::combine(op, parts(read, &ends), write);
let outcome = if len == 0 {
self.del(dest);
Ok(0)
} else {
self.reap(dest);
match self.string_only(dest) {
Ok(()) => {
self.store_raw(dest, &flat[split..], None);
Ok(len)
}
Err(e) => Err(e),
}
};
self.scratch = flat;
self.rows = ends;
outcome
}
pub fn bitfield(&mut self, key: &[u8], ops: &[Sub]) -> Result<Vec<Option<i64>>> {
let grow = ops.iter().filter(|s| s.op.writes()).map(reach).max();
self.bitfield_with(key, grow, |bytes| {
ops.iter().map(|&sub| apply(bytes, sub)).collect()
})
}
pub fn bitfield_with<T>(
&mut self,
key: &[u8],
grow: Option<usize>,
run: impl FnOnce(&mut [u8]) -> T,
) -> Result<T> {
self.reap(key);
self.string_only(key)?;
self.thaw(key)?;
let need = grow.unwrap_or(0);
check_len(key, need)?;
let mut bytes = std::mem::take(&mut self.scratch);
bytes.clear();
let deadline = match self.map.get(key) {
Some(rec) => {
value::read(rec).write_to(&mut bytes);
value::expire_at(rec)
}
None => None,
};
if bytes.len() < need {
bytes.resize(need, 0);
}
let out = run(&mut bytes);
if grow.is_some() {
self.store_raw(key, &bytes, deadline);
}
self.scratch = bytes;
Ok(out)
}
fn bitmap<'a>(&'a self, key: &[u8], digits: &'a mut [u8; DIGITS_MAX]) -> &'a [u8] {
match self.peek(key) {
None => &[],
Some(Str::Bytes(b)) => b,
Some(Str::Int(n)) => num::i64_digits(digits, n),
}
}
}
impl Db {
pub fn bitop<'k, I>(&mut self, op: Op, dest: &'k [u8], srcs: I) -> Result<usize>
where
I: Iterator<Item = &'k [u8]> + Clone,
{
if let Some(home) = self.one_stripe(std::iter::once(dest).chain(srcs.clone())) {
return self.stripe_mut(home).bitop(op, dest, srcs);
}
for src in srcs.clone() {
let stripe = self.at(src);
stripe.reap(src);
stripe.string_only(src)?;
stripe.thaw(src)?;
}
let (mut flat, mut ends) = self.take_scratch();
flat.clear();
ends.clear();
let mut digits = [0u8; DIGITS_MAX];
for src in srcs.clone() {
let bytes = self.at_ref(src).bitmap(src, &mut digits);
flat.extend_from_slice(bytes);
ends.push(flat.len());
}
let len = bits::width(parts(&flat, &ends));
if len > STRING_MAX {
self.put_scratch(flat, ends);
return Err(Error::new(Code::Invalid, TOO_LONG));
}
let split = flat.len();
flat.resize(split + len, 0);
let (read, write) = flat.split_at_mut(split);
bits::combine(op, parts(read, &ends), write);
let outcome = if len == 0 {
self.at(dest).del(dest);
Ok(0)
} else {
let stripe = self.at(dest);
stripe.reap(dest);
match stripe.string_only(dest) {
Ok(()) => {
stripe.store_raw(dest, &flat[split..], None);
Ok(len)
}
Err(e) => Err(e),
}
};
self.put_scratch(flat, ends);
outcome
}
}
fn parts<'a>(flat: &'a [u8], ends: &'a [usize]) -> impl Iterator<Item = &'a [u8]> + Clone {
std::iter::once(0)
.chain(ends.iter().copied())
.zip(ends.iter().copied())
.map(|(from, to)| &flat[from..to])
}
#[must_use]
pub fn apply(bytes: &mut [u8], sub: Sub) -> Option<i64> {
let had = bits::get(bytes, sub.at, sub.field);
match sub.op {
SubOp::Get => Some(had),
SubOp::Set(val) => bits::setting(sub.field, val, sub.on).map(|next| {
bits::set(bytes, sub.at, sub.field, next);
had
}),
SubOp::Incr(by) => bits::adding(sub.field, had, by, sub.on).inspect(|&next| {
bits::set(bytes, sub.at, sub.field, next);
}),
}
}
#[must_use]
pub const fn reach(sub: &Sub) -> usize {
(sub.field.last_bit(sub.at) / 8 + 1) as usize
}
#[inline]
const fn mask(offset: u64) -> u8 {
0x80 >> (offset % 8)
}
fn window(len: usize, start: i64, end: i64, unit: Unit) -> Option<(u64, u64)> {
let items = match unit {
Unit::Byte => len as i64,
Unit::Bit => (len as i64).checked_mul(8)?,
};
if items == 0 {
return None;
}
let back = |i: i64| if i < 0 { (items + i).max(0) } else { i };
let (from, to) = (back(start), back(end).min(items - 1));
if from > to {
return None;
}
let scale = match unit {
Unit::Byte => 8,
Unit::Bit => 1,
};
Some(((from * scale) as u64, ((to + 1) * scale) as u64))
}
#[must_use]
pub const fn max_bits() -> u64 {
STRING_MAX as u64 * 8
}
#[cfg(test)]
mod tests {
use super::*;
use crate::keyspace::Keyspace;
fn db() -> Keyspace {
Keyspace::new()
}
fn keys<'k>(names: &'k [&'k [u8]]) -> impl Iterator<Item = &'k [u8]> + Clone {
names.iter().copied()
}
#[test]
fn a_bit_is_set_and_read_back() {
let mut db = db();
assert!(!db.setbit(b"k", 7, true).expect("a bit"));
assert!(db.getbit(b"k", 7).expect("a bit"));
assert!(!db.getbit(b"k", 6).expect("a bit"));
assert_eq!(db.strlen(b"k").expect("a length"), 1);
assert_eq!(
db.get(b"k").expect("a value").expect("bytes").to_vec(),
b"\x01"
);
assert!(db.setbit(b"k", 7, false).expect("a bit"));
assert!(!db.setbit(b"k", 7, false).expect("a bit"));
}
#[test]
fn a_write_creates_and_pads_even_when_the_bit_is_zero() {
let mut db = db();
assert!(!db.setbit(b"k", 0, false).expect("a bit"));
assert!(db.exists(b"k"));
assert_eq!(db.strlen(b"k").expect("a length"), 1);
db.setbit(b"k", 40, true).expect("a bit");
assert_eq!(db.strlen(b"k").expect("a length"), 6);
}
#[test]
fn a_write_leaves_the_value_raw_and_a_read_does_not() {
let mut db = db();
db.set_plain(b"n", b"12345").expect("a set");
assert_eq!(db.encoding(b"n"), Some(value::Encoding::Int));
assert!(db.getbit(b"n", 3).expect("a bit"));
assert_eq!(db.encoding(b"n"), Some(value::Encoding::Int));
assert!(!db.setbit(b"n", 0, false).expect("a bit"));
assert_eq!(db.encoding(b"n"), Some(value::Encoding::Raw));
assert_eq!(
db.get(b"n").expect("a value").expect("bytes").to_vec(),
b"12345"
);
}
#[test]
fn a_write_keeps_the_deadline() {
let mut db = db();
db.setex(b"k", 100, b"abc").expect("a set");
db.setbit(b"k", 40, true).expect("a bit");
assert_eq!(db.strlen(b"k").expect("a length"), 6);
assert!(db.expire_at(b"k").is_some());
db.setbit(b"k", 1, true).expect("a bit");
assert!(db.expire_at(b"k").is_some());
}
#[test]
fn counting_takes_the_ranges_a_real_server_takes() {
let mut db = db();
db.set_plain(b"k", b"foobar").expect("a set");
let count = |db: &mut Keyspace, r| db.bitcount(b"k", r).expect("a count");
assert_eq!(count(&mut db, None), 26);
assert_eq!(count(&mut db, Some((0, 0, Unit::Byte))), 4);
assert_eq!(count(&mut db, Some((1, 1, Unit::Byte))), 6);
assert_eq!(count(&mut db, Some((0, -5, Unit::Byte))), 10);
assert_eq!(count(&mut db, Some((5, 30, Unit::Bit))), 17);
assert_eq!(count(&mut db, Some((0, -5, Unit::Bit))), 25);
assert_eq!(count(&mut db, Some((-100, 100, Unit::Byte))), 26);
assert_eq!(count(&mut db, Some((2, 1, Unit::Byte))), 0);
assert_eq!(count(&mut db, Some((5, 3, Unit::Bit))), 0);
assert_eq!(count(&mut db, Some((10, 20, Unit::Byte))), 0);
assert_eq!(db.bitcount(b"gone", None).expect("a count"), 0);
}
#[test]
fn searching_takes_the_ranges_a_real_server_takes() {
let mut db = db();
db.set_plain(b"ones", b"\xff\xff\xff").expect("a set");
db.set_plain(b"mix", b"\x00\xff\x00").expect("a set");
let pos = |db: &mut Keyspace, k: &[u8], bit, s, e| {
db.bitpos(k, bit, s, e, Unit::Byte).expect("a position")
};
assert_eq!(pos(&mut db, b"mix", true, None, None), 8);
assert_eq!(pos(&mut db, b"mix", false, None, None), 0);
assert_eq!(pos(&mut db, b"mix", true, Some(2), None), -1);
assert_eq!(pos(&mut db, b"mix", true, Some(-1), Some(-1)), -1);
assert_eq!(pos(&mut db, b"mix", false, Some(-100), None), 0);
assert_eq!(pos(&mut db, b"ones", false, None, None), 24);
assert_eq!(pos(&mut db, b"ones", false, Some(-1), None), 24);
assert_eq!(pos(&mut db, b"ones", false, Some(0), Some(-1)), -1);
assert_eq!(pos(&mut db, b"ones", false, Some(0), Some(100)), -1);
assert_eq!(pos(&mut db, b"ones", false, Some(10), None), -1);
assert_eq!(pos(&mut db, b"ones", false, Some(3), None), -1);
assert_eq!(pos(&mut db, b"ones", true, Some(10), None), -1);
assert_eq!(pos(&mut db, b"ones", false, Some(2), Some(1)), -1);
assert_eq!(
db.bitpos(b"ones", false, Some(5), Some(20), Unit::Bit)
.expect("a position"),
-1
);
}
#[test]
fn searching_an_absent_or_empty_key() {
let mut db = db();
let pos = |db: &mut Keyspace, k: &[u8], bit| {
db.bitpos(k, bit, None, None, Unit::Byte)
.expect("a position")
};
assert_eq!(pos(&mut db, b"gone", false), 0);
assert_eq!(pos(&mut db, b"gone", true), -1);
db.set_plain(b"empty", b"").expect("a set");
assert_eq!(pos(&mut db, b"empty", false), -1);
assert_eq!(pos(&mut db, b"empty", true), -1);
assert_eq!(
db.bitcount(b"empty", Some((0, -1, Unit::Byte)))
.expect("a count"),
0
);
}
#[test]
fn combining_writes_a_destination_and_deletes_an_empty_one() {
let mut db = db();
db.set_plain(b"a", b"\xf0\x0f\xff").expect("a set");
db.set_plain(b"b", b"\xff\x00").expect("a set");
let n = db
.bitop(Op::And, b"d", keys(&[b"a", b"b"]))
.expect("a length");
assert_eq!(n, 3);
assert_eq!(
db.get(b"d").expect("a value").expect("bytes").to_vec(),
b"\xf0\x00\x00"
);
db.set_plain(b"z", b"\x00\x00").expect("a set");
let n = db
.bitop(Op::And, b"d", keys(&[b"a", b"z"]))
.expect("a length");
assert_eq!(n, 3);
assert!(db.exists(b"d"));
let n = db
.bitop(Op::Or, b"d", keys(&[b"no1", b"no2"]))
.expect("a length");
assert_eq!(n, 0);
assert!(!db.exists(b"d"));
}
#[test]
fn combining_reads_an_int_key_as_its_digits() {
let mut db = db();
db.set_plain(b"n", b"12345").expect("a set");
db.bitop(Op::Or, b"d", keys(&[b"n"])).expect("a length");
assert_eq!(
db.get(b"d").expect("a value").expect("bytes").to_vec(),
b"12345"
);
}
#[test]
fn a_field_is_read_written_and_incremented() {
let mut db = db();
let u8f = Field::new(false, 8).expect("a width");
let sub = |op, at| Sub {
op,
field: u8f,
at,
on: Overflow::Wrap,
};
let out = db
.bitfield(b"k", &[sub(SubOp::Set(255), 0), sub(SubOp::Get, 0)])
.expect("replies");
assert_eq!(out, vec![Some(0), Some(255)]);
assert_eq!(db.strlen(b"k").expect("a length"), 1);
let out = db
.bitfield(b"k", &[sub(SubOp::Incr(10), 0)])
.expect("replies");
assert_eq!(out, vec![Some(9)], "wrapped round");
let fail = Sub {
on: Overflow::Fail,
..sub(SubOp::Incr(250), 0)
};
let out = db
.bitfield(b"k", &[fail, sub(SubOp::Get, 0)])
.expect("replies");
assert_eq!(out, vec![None, Some(9)]);
}
#[test]
fn a_read_only_bitfield_creates_nothing_and_re_encodes_nothing() {
let mut db = db();
let f = Field::new(true, 16).expect("a width");
let get = Sub {
op: SubOp::Get,
field: f,
at: 0,
on: Overflow::Wrap,
};
assert_eq!(
db.bitfield(b"gone", &[get]).expect("replies"),
vec![Some(0)]
);
assert!(!db.exists(b"gone"));
db.set_plain(b"s", b"hello").expect("a set");
assert_eq!(db.encoding(b"s"), Some(value::Encoding::Embstr));
db.bitfield(b"s", &[get]).expect("replies");
assert_eq!(
db.encoding(b"s"),
Some(value::Encoding::Embstr),
"still short"
);
}
#[test]
fn a_write_grows_the_value_even_when_every_write_fails() {
let mut db = db();
let f = Field::new(false, 8).expect("a width");
let sub = Sub {
op: SubOp::Set(300),
field: f,
at: 64,
on: Overflow::Fail,
};
assert_eq!(db.bitfield(b"k", &[sub]).expect("replies"), vec![None]);
assert_eq!(db.strlen(b"k").expect("a length"), 9);
}
#[test]
fn a_bit_command_on_the_wrong_type_says_so() {
let mut db = db();
let member: &[u8] = b"x";
db.sadd(b"s", std::iter::once(member)).expect("a member");
assert!(db.getbit(b"s", 0).is_err());
assert!(db.setbit(b"s", 0, true).is_err());
assert!(db.bitcount(b"s", None).is_err());
assert!(db.bitpos(b"s", true, None, None, Unit::Byte).is_err());
assert!(db.bitop(Op::Or, b"d", keys(&[b"s"])).is_err());
let f = Field::new(false, 8).expect("a width");
let sub = Sub {
op: SubOp::Get,
field: f,
at: 0,
on: Overflow::Wrap,
};
assert!(db.bitfield(b"s", &[sub]).is_err());
}
#[test]
fn an_offset_past_the_end_of_the_world_is_refused() {
let mut db = db();
assert!(db.setbit(b"k", BIT_OFFSET_MAX + 1, true).is_err());
assert!(db.getbit(b"k", BIT_OFFSET_MAX + 1).is_err());
assert!(db.setbit(b"k", BIT_OFFSET_MAX, true).is_err());
assert!(max_bits() < BIT_OFFSET_MAX);
}
}