use crate::resp::{
hyperloglog::hyper_log_log::{HyperLogLog, SPARSE_SIZE_MAX_CAP},
resp_server_session::RespServerSession,
};
const RESP_ERR_WRONG_TYPE_HLL: &[u8] =
b"-WRONGTYPE Key is not a valid HyperLogLog string value.\r\n";
fn load_hll<'s>(
hll: &HyperLogLog,
store: &wkv::BatchStoreSession<'s, impl wdev::Device>,
key: &[u8],
) -> Result<Option<Vec<u8>>, ()> {
match store.try_read_sync(key, |v| v.to_vec()) {
Ok(Some(Some(raw))) => {
if hll.is_valid_hyll_len(&raw, raw.len()) {
Ok(Some(raw))
} else {
Err(())
}
}
Ok(Some(None)) | Ok(None) => Ok(None),
Err(_) => Err(()),
}
}
fn store_hll(
hll: &HyperLogLog,
store: &wkv::BatchStoreSession<impl wdev::Device>,
key: &[u8],
blob: &[u8],
) {
let len = if hll.is_sparse(blob) {
hll
.sparse_current_size_in_bytes(blob)
.max(hll.sparse_bytes())
} else {
hll.dense_bytes()
};
let _ = store.try_upsert_sync(key, &blob[..len]);
}
impl RespServerSession {
pub fn hyper_log_log_add<'a, D: wdev::Device>(
&mut self,
parse_state: &[&[u8]],
store: &wkv::BatchStoreSession<'a, D>,
output: &mut Vec<u8>,
) -> wresp::Result<bool> {
if parse_state.is_empty() {
output.extend_from_slice(b"-ERR wrong number of arguments for 'PFADD' command\r\n");
return Ok(true);
}
let key = parse_state[0];
let elements = &parse_state[1..];
let hll = HyperLogLog::new();
let existing = match load_hll(&hll, store, key) {
Ok(existing) => existing,
Err(()) => {
output.extend_from_slice(RESP_ERR_WRONG_TYPE_HLL);
return Ok(true);
}
};
let mut updated = false;
match existing {
None => {
let initial = hll.sparse_initial_length(elements.len());
let mut blob = vec![0_u8; initial];
hll.init(elements, &mut blob);
updated = true;
store_hll(&hll, store, key, &blob);
}
Some(raw) => {
if hll.is_dense(&raw) {
let mut blob = raw;
hll.update(elements, &mut blob, &mut updated);
if updated {
store_hll(&hll, store, key, &blob);
}
} else {
let mut blob = raw;
if !hll.update(elements, &mut blob, &mut updated) {
let new_len = hll.update_grow(elements.len(), &blob);
let mut grown = vec![0_u8; new_len];
hll.copy_update(elements, &blob, &mut grown);
updated = true;
store_hll(&hll, store, key, &grown);
} else if updated {
store_hll(&hll, store, key, &blob);
}
}
}
}
output.extend_from_slice(if updated { b":1\r\n" } else { b":0\r\n" });
Ok(true)
}
pub fn hyper_log_log_length<'a, D: wdev::Device>(
&mut self,
parse_state: &[&[u8]],
store: &wkv::BatchStoreSession<'a, D>,
output: &mut Vec<u8>,
) -> wresp::Result<bool> {
if parse_state.is_empty() {
output.extend_from_slice(b"-ERR wrong number of arguments for 'PFCOUNT' command\r\n");
return Ok(true);
}
let hll = HyperLogLog::new();
let mut dense: Option<Vec<u8>> = None;
for key in parse_state {
match load_hll(&hll, store, key) {
Err(()) => {
output.extend_from_slice(RESP_ERR_WRONG_TYPE_HLL);
return Ok(true);
}
Ok(None) => continue,
Ok(Some(raw)) => {
let mut view = vec![0_u8; hll.dense_bytes()];
if hll.is_sparse(&raw) {
hll.init_dense(&mut view);
hll.sparse_to_dense(&raw, &mut view);
} else {
view.copy_from_slice(&raw[..hll.dense_bytes()]);
}
match &mut dense {
None => dense = Some(view),
Some(dst) => {
hll.dense_to_dense(&view, dst);
}
}
}
}
}
let card = dense.map(|mut d| hll.count(&mut d)).unwrap_or(0);
let mut buf = itoa::Buffer::new();
output.push(b':');
output.extend_from_slice(buf.format(card).as_bytes());
output.extend_from_slice(b"\r\n");
Ok(true)
}
pub fn hyper_log_log_merge<'a, D: wdev::Device>(
&mut self,
parse_state: &[&[u8]],
store: &wkv::BatchStoreSession<'a, D>,
output: &mut Vec<u8>,
) -> wresp::Result<bool> {
if parse_state.len() < 2 {
output.extend_from_slice(b"-ERR wrong number of arguments for 'PFMERGE' command\r\n");
return Ok(true);
}
let dest = parse_state[0];
let sources = &parse_state[1..];
let hll = HyperLogLog::new();
let mut dst = match load_hll(&hll, store, dest) {
Err(()) => {
output.extend_from_slice(RESP_ERR_WRONG_TYPE_HLL);
return Ok(true);
}
Ok(Some(raw)) => raw,
Ok(None) => {
let mut blob = vec![0_u8; SPARSE_SIZE_MAX_CAP];
hll.init_sparse(&mut blob);
blob[..hll.sparse_bytes()].to_vec()
}
};
for src_key in sources {
let src = match load_hll(&hll, store, src_key) {
Err(()) => {
output.extend_from_slice(RESP_ERR_WRONG_TYPE_HLL);
return Ok(true);
}
Ok(Some(raw)) => raw,
Ok(None) => continue,
};
let new_len = hll.merge_grow(&src, &dst);
if new_len != dst.len() {
let mut grown = vec![0_u8; new_len];
let old_len = dst.len();
hll.copy_update_merge(&src, &dst, &mut grown, old_len, new_len);
dst = grown;
} else {
hll.merge(&src, &mut dst);
hll.set_card(&mut dst, i64::MIN);
}
}
store_hll(&hll, store, dest, &dst);
output.extend_from_slice(b"+OK\r\n");
Ok(true)
}
}
#[cfg(test)]
mod tests {
use std::{iter::once, sync::Arc};
use tempfile::{TempDir, tempdir};
use wdev::SegmentedDevice;
use wkv::{StoreConfig, WedbStore};
use super::*;
type TestSession = wkv::StoreSession<SegmentedDevice>;
fn fixture(tag: &str) -> (TempDir, Arc<WedbStore<SegmentedDevice>>, TestSession) {
let dir = tempdir().unwrap();
let device = Arc::new(SegmentedDevice::single_file(dir.path().join(tag)).unwrap());
let config = StoreConfig::new(16384, 65536, 64, 0.5).unwrap();
let store = Arc::new(WedbStore::open(config, device).unwrap());
let session = store.new_session().unwrap();
(dir, store, session)
}
#[test]
fn pfadd_pfcount_pfmerge_flow() {
let (_dir, _store, session) = fixture("hll.db");
let batch = session.enter_batch();
let mut sess = RespServerSession::default();
let mut out = Vec::new();
sess
.hyper_log_log_add(&[b"h1", b"a", b"b", b"c"], &batch, &mut out)
.unwrap();
assert_eq!(out, b":1\r\n");
out.clear();
sess
.hyper_log_log_add(&[b"h1", b"a", b"b", b"c"], &batch, &mut out)
.unwrap();
assert_eq!(out, b":0\r\n");
out.clear();
sess
.hyper_log_log_add(&[b"h1", b"d"], &batch, &mut out)
.unwrap();
assert_eq!(out, b":1\r\n");
out.clear();
sess
.hyper_log_log_length(&[b"h1"], &batch, &mut out)
.unwrap();
assert_eq!(out, b":4\r\n");
out.clear();
sess
.hyper_log_log_length(&[b"missing"], &batch, &mut out)
.unwrap();
assert_eq!(out, b":0\r\n");
out.clear();
sess
.hyper_log_log_merge(&[b"h2", b"h1"], &batch, &mut out)
.unwrap();
assert_eq!(out, b"+OK\r\n");
out.clear();
sess
.hyper_log_log_length(&[b"h1", b"h2"], &batch, &mut out)
.unwrap();
assert_eq!(out, b":4\r\n");
let _ = batch.try_upsert_sync(b"bad", b"plain-string-value");
out.clear();
sess
.hyper_log_log_add(&[b"bad", b"x"], &batch, &mut out)
.unwrap();
assert_eq!(out, RESP_ERR_WRONG_TYPE_HLL);
out.clear();
sess
.hyper_log_log_length(&[b"bad"], &batch, &mut out)
.unwrap();
assert_eq!(out, RESP_ERR_WRONG_TYPE_HLL);
out.clear();
sess
.hyper_log_log_merge(&[b"dest", b"bad"], &batch, &mut out)
.unwrap();
assert_eq!(out, RESP_ERR_WRONG_TYPE_HLL);
out.clear();
sess.hyper_log_log_add(&[], &batch, &mut out).unwrap();
assert_eq!(
out,
b"-ERR wrong number of arguments for 'PFADD' command\r\n"
);
out.clear();
sess
.hyper_log_log_merge(&[b"dest"], &batch, &mut out)
.unwrap();
assert_eq!(
out,
b"-ERR wrong number of arguments for 'PFMERGE' command\r\n"
);
}
#[test]
fn pfmerge_empty_source_no_underflow() {
let (_dir, _store, session) = fixture("hll3.db");
let batch = session.enter_batch();
let mut sess = RespServerSession::default();
let mut out = Vec::new();
sess
.hyper_log_log_merge(&[b"empty", b"no-such-1", b"no-such-2"], &batch, &mut out)
.unwrap();
assert_eq!(out, b"+OK\r\n");
out.clear();
sess
.hyper_log_log_length(&[b"empty"], &batch, &mut out)
.unwrap();
assert_eq!(out, b":0\r\n");
out.clear();
sess
.hyper_log_log_merge(&[b"copy", b"empty"], &batch, &mut out)
.unwrap();
assert_eq!(out, b"+OK\r\n");
out.clear();
sess
.hyper_log_log_merge(&[b"empty", b"empty"], &batch, &mut out)
.unwrap();
assert_eq!(out, b"+OK\r\n");
out.clear();
sess
.hyper_log_log_length(&[b"copy", b"empty"], &batch, &mut out)
.unwrap();
assert_eq!(out, b":0\r\n");
let mut out2 = Vec::new();
sess
.hyper_log_log_add(&[b"real", b"x", b"y"], &batch, &mut out2)
.unwrap();
out.clear();
sess
.hyper_log_log_merge(&[b"real", b"empty"], &batch, &mut out)
.unwrap();
assert_eq!(out, b"+OK\r\n");
out.clear();
sess
.hyper_log_log_length(&[b"real"], &batch, &mut out)
.unwrap();
assert_eq!(out, b":2\r\n");
}
#[test]
fn dense_upgrade_and_merge_monotonic() {
let (_dir, _store, session) = fixture("hll2.db");
let batch = session.enter_batch();
let mut sess = RespServerSession::default();
let elements: Vec<Vec<u8>> = (0..1000)
.map(|i| format!("elem-{i}").into_bytes())
.collect();
let args: Vec<&[u8]> = once(b"big".as_slice())
.chain(elements.iter().map(|e| e.as_slice()))
.collect();
let mut out = Vec::new();
sess.hyper_log_log_add(&args, &batch, &mut out).unwrap();
assert_eq!(out, b":1\r\n");
out.clear();
sess
.hyper_log_log_length(&[b"big"], &batch, &mut out)
.unwrap();
let card: i64 = String::from_utf8_lossy(&out[1..out.len() - 2])
.parse()
.unwrap();
assert!((880..=1120).contains(&card), "card = {card}");
out.clear();
sess
.hyper_log_log_merge(&[b"copy", b"big"], &batch, &mut out)
.unwrap();
out.clear();
sess
.hyper_log_log_length(&[b"copy"], &batch, &mut out)
.unwrap();
let card2: i64 = String::from_utf8_lossy(&out[1..out.len() - 2])
.parse()
.unwrap();
assert_eq!(card, card2);
}
}