use yo_common::num::{parse_f64, parse_i64};
use yo_common::{Code, Error, Result};
use yo_kv::{Db, Foreign, Keyspace};
use yo_sketch::topk::TopK;
use super::args::{self, Args};
use super::table::Spec;
use crate::reply::Out;
const DEFAULT_WIDTH: u32 = 8;
const DEFAULT_DEPTH: u32 = 7;
const DEFAULT_DECAY: f64 = 0.9;
const MAX_INCREMENT: i64 = 100_000;
const EXISTS: &[u8] = b"TopK: key already exists";
const MISSING: &[u8] = b"TopK: key does not exist";
const BAD_K: &[u8] = b"TopK: invalid k";
const BAD_WIDTH: &[u8] = b"TopK: invalid width";
const BAD_DEPTH: &[u8] = b"TopK: invalid depth";
const BAD_DECAY: &[u8] = b"TopK: invalid decay value. must be '<= 1' & '> 0'";
const NO_MEMORY: &[u8] = b"ERR Insufficient memory to create topk data structure";
const KEYWORD: &[u8] = b"WITHCOUNT keyword expected";
const BAD_INCREMENT: &[u8] = b"TopK: increment must be an integer greater or equal to 0 and smaller or equal to 100,000";
const WRONG_KIND: &str = "Operation against a key holding the wrong kind of value";
#[derive(Debug)]
pub(super) struct TopKBody {
t: TopK,
}
impl Foreign for TopKBody {
fn type_name(&self) -> &'static str {
"TopK-TYPE"
}
fn encoding(&self) -> &'static str {
"raw"
}
fn memory_bytes(&self) -> usize {
self.t.memory_bytes()
}
fn is_empty(&self) -> bool {
false
}
}
pub(super) fn execute(db: &Db, spec: &Spec, args: Args<'_>, out: &mut Out) -> Result<()> {
let mut held = db.hold(args.get(1));
let db = &mut *held;
match spec.name {
"topk.reserve" => reserve(db, args, out),
"topk.add" => add(db, args, out),
"topk.incrby" => incrby(db, args, out),
"topk.query" => query(db, args, out),
"topk.count" => count(db, args, out),
"topk.list" => list(db, args, out),
"topk.info" => info(db, args, out),
other => unreachable!("{other} is not a top k command"),
}
}
fn reserve(db: &mut Keyspace, args: Args<'_>, out: &mut Out) -> Result<()> {
if args.len() != 3 && args.len() != 6 {
return Err(args::wrong_arity("topk.reserve"));
}
let key = args.get(1);
if db.kind_of(key).is_some() {
out.error(EXISTS);
return Ok(());
}
let Some(k) = size(args.get(2)) else {
out.error(BAD_K);
return Ok(());
};
let (mut width, mut depth, mut decay) = (DEFAULT_WIDTH, DEFAULT_DEPTH, DEFAULT_DECAY);
if args.len() == 6 {
let Some(w) = size(args.get(3)) else {
out.error(BAD_WIDTH);
return Ok(());
};
let Some(d) = size(args.get(4)) else {
out.error(BAD_DEPTH);
return Ok(());
};
let Some(rate) = parse_f64(args.get(5)).filter(|&n| n > 0.0 && n <= 1.0) else {
out.error(BAD_DECAY);
return Ok(());
};
(width, depth, decay) = (w, d, rate);
}
let Some(t) = TopK::new(k, width, depth, decay) else {
out.error(NO_MEMORY);
return Ok(());
};
db.put_foreign(key, Box::new(TopKBody { t }));
out.ok();
Ok(())
}
fn add(db: &mut Keyspace, args: Args<'_>, out: &mut Out) -> Result<()> {
let Some(body) = write(db, args.get(1))? else {
out.error(MISSING);
return Ok(());
};
out.array(args.len() - 2);
for i in 2..args.len() {
expelled(body.t.add(args.get(i), 1), out);
}
Ok(())
}
fn incrby(db: &mut Keyspace, args: Args<'_>, out: &mut Out) -> Result<()> {
if !args.len().is_multiple_of(2) {
return Err(args::wrong_arity("topk.incrby"));
}
let Some(body) = write(db, args.get(1))? else {
out.error(MISSING);
return Ok(());
};
let start = out.len();
let mut written = 0;
for i in (2..args.len()).step_by(2) {
match parse_i64(args.get(i + 1)).filter(|&n| (0..=MAX_INCREMENT).contains(&n)) {
Some(by) => expelled(body.t.add(args.get(i), by as u32), out),
None => {
out.error(BAD_INCREMENT);
written += 1;
break;
}
}
written += 1;
}
out.close_array(start, written);
Ok(())
}
fn expelled(item: Option<Box<[u8]>>, out: &mut Out) {
match item {
Some(item) => out.bulk(&item),
None => out.nil(),
}
}
fn query(db: &mut Keyspace, args: Args<'_>, out: &mut Out) -> Result<()> {
let Some(body) = read(db, args.get(1))? else {
out.error(MISSING);
return Ok(());
};
out.array(args.len() - 2);
for i in 2..args.len() {
out.bool(body.t.query(args.get(i)));
}
Ok(())
}
fn count(db: &mut Keyspace, args: Args<'_>, out: &mut Out) -> Result<()> {
let Some(body) = read(db, args.get(1))? else {
out.error(MISSING);
return Ok(());
};
out.array(args.len() - 2);
for i in 2..args.len() {
out.uint(u64::from(body.t.count_of(args.get(i))));
}
Ok(())
}
fn list(db: &mut Keyspace, args: Args<'_>, out: &mut Out) -> Result<()> {
if args.len() != 2 && args.len() != 3 {
return Err(args::wrong_arity("topk.list"));
}
let mut counts = false;
if let Some(word) = args.opt(2) {
if !is_prefix_of(word, b"withcount") {
out.error(KEYWORD);
return Ok(());
}
counts = true;
}
let Some(body) = read(db, args.get(1))? else {
out.error(MISSING);
return Ok(());
};
let kept = body.t.list();
out.array(kept.len() * if counts { 2 } else { 1 });
for (item, n) in kept {
out.bulk(item);
if counts {
out.uint(u64::from(n));
}
}
Ok(())
}
fn info(db: &mut Keyspace, args: Args<'_>, out: &mut Out) -> Result<()> {
let Some(body) = read(db, args.get(1))? else {
out.error(MISSING);
return Ok(());
};
out.map(4);
out.simple(b"k");
out.uint(u64::from(body.t.k()));
out.simple(b"width");
out.uint(u64::from(body.t.width()));
out.simple(b"depth");
out.uint(u64::from(body.t.depth()));
out.simple(b"decay");
out.double(body.t.decay());
Ok(())
}
fn size(arg: &[u8]) -> Option<u32> {
parse_i64(arg)
.filter(|&n| n >= 1 && n <= i64::from(u32::MAX))
.map(|n| n as u32)
}
fn is_prefix_of(word: &[u8], full: &[u8]) -> bool {
word.len() <= full.len() && word.eq_ignore_ascii_case(&full[..word.len()])
}
fn write<'d>(db: &'d mut Keyspace, key: &[u8]) -> Result<Option<&'d mut TopKBody>> {
match db.foreign_mut(key)? {
Some(body) => match body.downcast_mut::<TopKBody>() {
Some(body) => Ok(Some(body)),
None => Err(Error::new(Code::WrongType, WRONG_KIND)),
},
None => Ok(None),
}
}
fn read<'d>(db: &'d mut Keyspace, key: &[u8]) -> Result<Option<&'d TopKBody>> {
match db.foreign(key)? {
Some(body) => match body.downcast_ref::<TopKBody>() {
Some(body) => Ok(Some(body)),
None => Err(Error::new(Code::WrongType, WRONG_KIND)),
},
None => Ok(None),
}
}