use yo_common::num::parse_f64;
use yo_common::{Code, Error, Result};
use yo_kv::{Db, Foreign, Keyspace};
use yo_search::suggest::Suggestions;
use super::args::{self, Args};
use super::table::Spec;
use crate::reply::Out;
const DEFAULT_MAX: u64 = 5;
const MAX_MAX: u64 = u32::MAX as u64;
const LONGEST_ADD: usize = 7;
const BAD_SCORE: &str = "invalid score";
const UNKNOWN_ADD: &[u8] = b"Unknown argument `";
const NO_PAYLOAD: &[u8] = b"Invalid payload: Expected an argument, but none provided";
const UNKNOWN_GET: &[u8] = b"SEARCH_PARSE_ARGS Unrecognized argument: ";
const MAX_RANGE: &[u8] = b"SEARCH_PARSE_ARGS MAX: Value is outside acceptable bounds";
const MAX_KIND: &[u8] = b"SEARCH_PARSE_ARGS MAX: Could not convert argument to expected type";
const MAX_MISSING: &[u8] = b"SEARCH_PARSE_ARGS MAX: Expected an argument, but none provided";
const WRONG_KIND: &str = "Operation against a key holding the wrong kind of value";
#[derive(Debug)]
pub(super) struct SugBody {
s: Suggestions,
}
impl Foreign for SugBody {
fn type_name(&self) -> &'static str {
"trietype0"
}
fn encoding(&self) -> &'static str {
"raw"
}
fn memory_bytes(&self) -> usize {
self.s.memory_bytes()
}
fn is_empty(&self) -> bool {
self.s.is_empty()
}
}
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 {
"FT.SUGADD" => add(db, args, out),
"FT.SUGGET" => get(db, args, out),
"FT.SUGDEL" => del(db, args, out),
"FT.SUGLEN" => len(db, args, out),
other => unreachable!("{other} is not a suggestion command"),
}
}
fn add(db: &mut Keyspace, args: Args<'_>, out: &mut Out) -> Result<()> {
if args.len() > LONGEST_ADD {
return Err(args::wrong_arity("FT.SUGADD"));
}
let mut incr = false;
let mut payload = None;
let mut at = 4;
while at < args.len() {
let word = args.get(at);
if word.eq_ignore_ascii_case(b"incr") {
incr = true;
at += 1;
} else if word.eq_ignore_ascii_case(b"payload") {
let Some(bytes) = args.opt(at + 1) else {
out.error(NO_PAYLOAD);
return Ok(());
};
if !bytes.is_empty() {
payload = Some(bytes);
}
at += 2;
} else {
let mut line = UNKNOWN_ADD.to_vec();
line.extend_from_slice(word);
line.push(b'`');
out.error(&line);
return Ok(());
}
}
let Some(score) = double(args.get(3)) else {
return Err(Error::new(Code::Invalid, BAD_SCORE));
};
let key = args.get(1);
if mine(db.foreign(key)?)?.is_none() {
db.put_foreign(
key,
Box::new(SugBody {
s: Suggestions::new(),
}),
);
}
let body = mine_mut(db.foreign_mut(key)?.expect("there or just made"))?;
body.s.add(args.get(2), score, incr, payload);
out.uint(body.s.len() as u64);
Ok(())
}
fn get(db: &mut Keyspace, args: Args<'_>, out: &mut Out) -> Result<()> {
let mut fuzzy = false;
let mut scores = false;
let mut payloads = false;
let mut max = DEFAULT_MAX;
let mut at = 3;
while at < args.len() {
let word = args.get(at);
at += 1;
if word.eq_ignore_ascii_case(b"fuzzy") {
fuzzy = true;
} else if word.eq_ignore_ascii_case(b"withscores") {
scores = true;
} else if word.eq_ignore_ascii_case(b"withpayloads") {
payloads = true;
} else if word.eq_ignore_ascii_case(b"max") {
let Some(count) = args.opt(at) else {
out.error(MAX_MISSING);
return Ok(());
};
at += 1;
match whole(count) {
Ok(count) => max = count,
Err(line) => {
out.error(line);
return Ok(());
}
}
} else {
let mut line = UNKNOWN_GET.to_vec();
line.extend_from_slice(word);
out.error(&line);
return Ok(());
}
}
let Some(body) = mine(db.foreign(args.get(1))?)? else {
out.array(0);
return Ok(());
};
let hits = body.s.best(args.get(2), fuzzy, max as usize);
let each = 1 + usize::from(scores) + usize::from(payloads);
out.array(hits.len() * each);
for hit in hits {
out.bulk(hit.term);
if scores {
out.double(hit.score);
}
if payloads {
match hit.payload {
Some(bytes) => out.bulk(bytes),
None => out.nil(),
}
}
}
Ok(())
}
fn del(db: &mut Keyspace, args: Args<'_>, out: &mut Out) -> Result<()> {
let key = args.get(1);
let gone = match mine_mut_opt(db.foreign_mut(key)?)? {
Some(body) => body.s.remove(args.get(2)),
None => false,
};
out.uint(u64::from(gone));
db.reap_foreign(key);
Ok(())
}
fn len(db: &mut Keyspace, args: Args<'_>, out: &mut Out) -> Result<()> {
let held = mine(db.foreign(args.get(1))?)?;
out.uint(held.map_or(0, |body| body.s.len() as u64));
Ok(())
}
fn double(arg: &[u8]) -> Option<f64> {
let n = parse_f64(arg)?;
let text = core::str::from_utf8(arg).ok()?;
let word = text.trim_start_matches(['+', '-']);
if n.is_infinite()
&& !word.eq_ignore_ascii_case("inf")
&& !word.eq_ignore_ascii_case("infinity")
{
return None;
}
if n == 0.0 && text.bytes().any(|b| (b'1'..=b'9').contains(&b)) {
return None;
}
Some(n)
}
fn whole(arg: &[u8]) -> core::result::Result<u64, &'static [u8]> {
if let Some(n) = yo_common::num::parse_i64(arg) {
return u64::try_from(n)
.ok()
.filter(|n| (1..=MAX_MAX).contains(n))
.ok_or(MAX_RANGE);
}
let n = double(arg).ok_or(MAX_KIND)?.trunc();
if n < 1.0 {
return Err(MAX_KIND);
}
#[allow(clippy::cast_precision_loss)]
if n > MAX_MAX as f64 {
return Err(MAX_RANGE);
}
#[allow(clippy::cast_possible_truncation, clippy::cast_sign_loss)]
Ok(n as u64)
}
fn mine(body: Option<&dyn Foreign>) -> Result<Option<&SugBody>> {
match body {
Some(body) => match body.downcast_ref::<SugBody>() {
Some(body) => Ok(Some(body)),
None => Err(Error::new(Code::WrongType, WRONG_KIND)),
},
None => Ok(None),
}
}
fn mine_mut(body: &mut dyn Foreign) -> Result<&mut SugBody> {
body.downcast_mut::<SugBody>()
.ok_or_else(|| Error::new(Code::WrongType, WRONG_KIND))
}
fn mine_mut_opt(body: Option<&mut dyn Foreign>) -> Result<Option<&mut SugBody>> {
match body {
Some(body) => mine_mut(body).map(Some),
None => Ok(None),
}
}