use yo_kv::value::Kind;
use yo_search::Registry;
use yo_search::expr::{Expr, Value};
use yo_search::field::Kind as Type;
use super::super::Server;
use super::super::args::{self, Args};
use super::super::indexing::{self, Change, Document};
use super::super::table::Spec;
use super::{Fail, LANGUAGES, MISSING, NOT_LOADED, QUOTE_END, line};
use crate::reply::Out;
use yo_common::Result;
const SCORE_FIELD: &[u8] = b"__score";
const LANGUAGE_FIELD: &[u8] = b"__language";
const PAYLOAD_FIELD: &[u8] = b"__payload";
const PLAIN: f64 = 1.0;
const NO_FIELDS: &str = "SEARCH_ADD_ARGS No field list found";
const ODD_FIELDS: &str = "SEARCH_ADD_ARGS Fields must be specified in FIELD VALUE pairs";
const BAD_SCORE: &str = "SEARCH_ADD_ARGS Could not parse document score";
const SCORE_RANGE: &str = "SEARCH_ADD_ARGS Score must be between 0 and 1";
const BAD_LANGUAGE: &str = "SEARCH_ADD_ARGS Unsupported language";
const UNKNOWN_WORD: &str = "SEARCH_ADD_ARGS Unknown keyword `";
const UNKNOWN_END: &str = "` provided";
const NO_ARGUMENT: &str = "SEARCH_ADD_ARGS Parsing error for document option ";
const NO_ARGUMENT_END: &str = ": Expected an argument, but none provided";
const HERE_ALREADY: &str = "SEARCH_DOCUMENT_EXISTS Document already exists";
const BAD_KEY: &str = "SEARCH_REDIS_KEY_TYPE_BAD Invalid Redis key";
const NO_DOCUMENT: &str = "SEARCH_DOC_NOT_FOUND ";
struct Adding<'a> {
replace: bool,
partial: bool,
language: Option<&'a [u8]>,
payload: Option<&'a [u8]>,
condition: Option<&'a [u8]>,
at: usize,
}
pub(in crate::dispatch) fn execute(
server: &Server,
db: usize,
spec: &Spec,
args: Args<'_>,
out: &mut Out,
) -> Result<()> {
match spec.name {
"FT.ADD" | "FT.SAFEADD" => add(server, db, spec, args, out),
"FT.GET" => get(server, db, spec, args, out),
"FT.MGET" => mget(server, db, spec, args, out),
"FT.DEL" => del(server, db, spec, args, out),
other => unreachable!("{other} is not a document command"),
}
}
fn add(server: &Server, db: usize, spec: &Spec, args: Args<'_>, out: &mut Out) -> Result<()> {
if args.len() < 4 {
return Err(args::wrong_arity(spec.name));
}
let name = args.get(1);
let key = args.get(2);
let Some(canon) = resolved(&mut server.search.lock(), name) else {
Fail::naming(MISSING, name).write(out);
return Ok(());
};
let worth = args.get(3);
let Some(score) = yo_common::num::parse_f64(worth) else {
Fail::plain(BAD_SCORE).write(out);
return Ok(());
};
if !(0.0..=PLAIN).contains(&score) {
Fail::plain(SCORE_RANGE).write(out);
return Ok(());
}
let asked = match options(args) {
Ok(asked) => asked,
Err(fail) => {
fail.write(out);
return Ok(());
}
};
let kind = server.dbs[db].hold(key).kind_of(key);
if !matches!(kind, None | Some(Kind::Hash)) {
Fail::plain(BAD_KEY).write(out);
return Ok(());
}
if kind.is_some() {
if !asked.replace {
Fail::plain(HERE_ALREADY).write(out);
return Ok(());
}
if let Some(src) = asked.condition {
match keeps(server, db, &canon, key, src) {
Ok(true) => {}
Ok(false) => {
out.simple(b"NOADD");
return Ok(());
}
Err(text) => {
out.error(&text);
return Ok(());
}
}
}
}
let mut pairs: Vec<(&[u8], &[u8])> = Vec::new();
let mut at = asked.at;
while let (Some(field), Some(value)) = (args.opt(at), args.opt(at + 1)) {
pairs.push((field, value));
at += 2;
}
let recorded = asked.partial || score != PLAIN;
if recorded {
pairs.push((SCORE_FIELD, worth));
}
if let Some(language) = asked.language {
pairs.push((LANGUAGE_FIELD, language));
}
if let Some(payload) = asked.payload {
pairs.push((PAYLOAD_FIELD, payload));
}
{
let mut reg = server.search.lock();
if let Some(index) = reg.get_mut(&canon) {
if recorded {
index.definition.score_field = Some(SCORE_FIELD.into());
}
if asked.language.is_some() {
index.definition.language_field = Some(LANGUAGE_FIELD.into());
}
if asked.payload.is_some() {
index.definition.payload_field = Some(PAYLOAD_FIELD.into());
}
}
}
{
let mut held = server.dbs[db].hold(key);
if !asked.partial {
held.del(key);
}
if !pairs.is_empty() {
held.hset(key, pairs.iter().copied())?;
}
}
indexing::changed(server, db, key, Change::Key);
out.ok();
Ok(())
}
fn options<'a>(args: Args<'a>) -> core::result::Result<Adding<'a>, Fail<'a>> {
let mut asked = Adding {
replace: false,
partial: false,
language: None,
payload: None,
condition: None,
at: 0,
};
let mut at = 4;
loop {
let Some(word) = args.opt(at) else {
return Err(Fail::plain(NO_FIELDS));
};
at += 1;
if args::is(word, b"FIELDS") {
break;
}
if args::is(word, b"REPLACE") {
asked.replace = true;
continue;
}
if args::is(word, b"PARTIAL") {
asked.partial = true;
continue;
}
let taking = match () {
() if args::is(word, b"LANGUAGE") => "LANGUAGE",
() if args::is(word, b"PAYLOAD") => "PAYLOAD",
() if args::is(word, b"IF") => "IF",
() => return Err(Fail::about(UNKNOWN_WORD, word, UNKNOWN_END)),
};
let Some(value) = args.opt(at) else {
return Err(Fail::about(NO_ARGUMENT, taking.as_bytes(), NO_ARGUMENT_END));
};
at += 1;
match taking {
"LANGUAGE" => {
if !LANGUAGES.iter().any(|known| args::is(value, known)) {
return Err(Fail::plain(BAD_LANGUAGE));
}
asked.language = Some(value);
}
"PAYLOAD" => asked.payload = Some(value),
_ => asked.condition = Some(value),
}
}
if !(args.len() - at).is_multiple_of(2) {
return Err(Fail::plain(ODD_FIELDS));
}
asked.at = at;
Ok(asked)
}
fn keeps(
server: &Server,
db: usize,
canon: &[u8],
key: &[u8],
src: &[u8],
) -> core::result::Result<bool, Vec<u8>> {
{
let reg = server.search.lock();
let held = reg
.named(canon)
.and_then(|index| index.held.docs.id(key))
.is_some();
if !held {
return Err(NO_DOCUMENT.as_bytes().to_vec());
}
}
let doc = indexing::read(&server.dbs[db], key).unwrap_or_default();
let reg = server.search.lock();
let Some(index) = reg.named(canon) else {
return Err(NO_DOCUMENT.as_bytes().to_vec());
};
let mut row: Vec<Value> = Vec::new();
let mut names: Vec<Box<[u8]>> = Vec::new();
let mut expr = Expr::parse(src)?;
expr.bind(&mut |name| {
if let Some(at) = names.iter().position(|held| **held == *name) {
return Some(at);
}
let field = index.field(name)?;
names.push(name.into());
row.push(valued(&doc, &field.identifier, &field.kind));
Some(row.len() - 1)
})
.map_err(|missing| line(NOT_LOADED, &missing.0, QUOTE_END))?;
Ok(expr.eval(&row)?.truth())
}
fn valued(doc: &Document, from: &[u8], kind: &Type) -> Value {
let Some(value) = doc.held(from) else {
return Value::Missing;
};
match kind {
Type::Numeric => match yo_common::num::parse_f64(value) {
Some(number) => Value::Number(number),
None => Value::Text(value.into()),
},
_ => Value::Text(value.into()),
}
}
fn get(server: &Server, db: usize, spec: &Spec, args: Args<'_>, out: &mut Out) -> Result<()> {
if args.len() != 3 {
return Err(args::wrong_arity(spec.name));
}
let name = args.get(1);
let Some(canon) = resolved(&mut server.search.lock(), name) else {
Fail::naming(MISSING, name).write(out);
return Ok(());
};
shown(server, db, &canon, args.get(2), out);
Ok(())
}
fn mget(server: &Server, db: usize, spec: &Spec, args: Args<'_>, out: &mut Out) -> Result<()> {
if args.len() < 3 {
return Err(args::wrong_arity(spec.name));
}
let name = args.get(1);
let Some(canon) = resolved(&mut server.search.lock(), name) else {
Fail::naming(MISSING, name).write(out);
return Ok(());
};
out.array(args.len() - 2);
for at in 2..args.len() {
shown(server, db, &canon, args.get(at), out);
}
Ok(())
}
fn shown(server: &Server, db: usize, canon: &[u8], key: &[u8], out: &mut Out) {
let recorded = {
let reg = server.search.lock();
let held = reg
.named(canon)
.filter(|index| index.held.docs.id(key).is_some());
let Some(index) = held else {
out.nil();
return;
};
[
index.definition.score_field.clone(),
index.definition.language_field.clone(),
index.definition.payload_field.clone(),
]
};
let Some(doc) = indexing::read(&server.dbs[db], key) else {
out.array(0);
return;
};
let pairs: Vec<(&[u8], &[u8])> = doc
.pairs()
.into_iter()
.filter(|(field, _)| {
!recorded
.iter()
.flatten()
.any(|name| name.as_ref() == *field)
})
.collect();
out.array(pairs.len() * 2);
for (field, value) in pairs {
out.bulk(field);
out.bulk(value);
}
}
fn del(server: &Server, db: usize, spec: &Spec, args: Args<'_>, out: &mut Out) -> Result<()> {
if args.len() != 3 && args.len() != 4 {
return Err(args::wrong_arity(spec.name));
}
let name = args.get(1);
if resolved(&mut server.search.lock(), name).is_none() {
Fail::naming(MISSING, name).write(out);
return Ok(());
}
let key = args.get(2);
let gone = server.dbs[db].hold(key).del(key);
if gone {
indexing::changed(server, db, key, Change::Key);
}
out.int(i64::from(gone));
Ok(())
}
fn resolved(reg: &mut Registry, name: &[u8]) -> Option<Box<[u8]>> {
reg.open(name).map(|index| index.name.clone())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::dispatch::tests::encode;
use crate::proto::Limits;
use crate::request::Argv;
fn read(words: &[&[u8]], check: impl FnOnce(core::result::Result<Adding<'_>, Fail<'_>>)) {
let wire = encode(words);
let mut argv = Argv::new();
argv.decode(&wire, &Limits::default()).expect("it decodes");
check(options(Args::new(&argv, &wire)));
}
#[test]
fn the_option_words_come_in_any_order_and_any_case() {
read(
&[
b"FT.ADD",
b"i",
b"k",
b"1.0",
b"partial",
b"LaNgUaGe",
b"FrEnCh",
b"REPLACE",
b"PAYLOAD",
b"pp",
b"IF",
b"@n>1",
b"FIELDS",
b"t",
b"a",
],
|got| {
let asked = got.ok().expect("it reads");
assert!(asked.replace);
assert!(asked.partial);
assert_eq!(asked.language, Some(b"FrEnCh".as_slice()));
assert_eq!(asked.payload, Some(b"pp".as_slice()));
assert_eq!(asked.condition, Some(b"@n>1".as_slice()));
assert_eq!(asked.at, 13);
},
);
}
#[test]
fn fields_ends_the_option_list() {
read(
&[b"FT.ADD", b"i", b"k", b"1.0", b"FIELDS", b"REPLACE", b"x"],
|got| {
let asked = got.ok().expect("it reads");
assert!(!asked.replace);
assert_eq!(asked.at, 5);
},
);
}
#[test]
fn an_option_with_nothing_after_it_names_itself_in_capitals() {
for (word, named) in [
(b"language".as_slice(), "LANGUAGE"),
(b"PaYlOaD", "PAYLOAD"),
(b"if", "IF"),
] {
read(&[b"FT.ADD", b"i", b"k", b"1.0", word], |got| {
let fail = got.err().expect("it refuses");
assert_eq!(fail.head, NO_ARGUMENT);
assert_eq!(fail.word, named.as_bytes());
assert_eq!(fail.tail, NO_ARGUMENT_END);
});
}
}
#[test]
fn an_unknown_word_comes_back_as_the_client_wrote_it() {
for word in [b"bogus".as_slice(), b"NOSAVE"] {
read(
&[b"FT.ADD", b"i", b"k", b"1.0", word, b"FIELDS", b"t", b"a"],
|got| {
let fail = got.err().expect("it refuses");
assert_eq!(fail.head, UNKNOWN_WORD);
assert_eq!(fail.word, word);
assert_eq!(fail.tail, UNKNOWN_END);
},
);
}
}
#[test]
fn a_field_list_is_required_and_comes_in_pairs() {
read(&[b"FT.ADD", b"i", b"k", b"1.0", b"REPLACE"], |got| {
assert_eq!(got.err().expect("it refuses").head, NO_FIELDS);
});
read(&[b"FT.ADD", b"i", b"k", b"1.0", b"FIELDS", b"t"], |got| {
assert_eq!(got.err().expect("it refuses").head, ODD_FIELDS);
});
read(&[b"FT.ADD", b"i", b"k", b"1.0", b"FIELDS"], |got| {
assert_eq!(got.ok().expect("it reads").at, 5);
});
}
#[test]
fn only_a_language_the_module_knows() {
read(
&[
b"FT.ADD",
b"i",
b"k",
b"1.0",
b"LANGUAGE",
b"klingon",
b"FIELDS",
b"t",
b"a",
],
|got| {
assert_eq!(got.err().expect("it refuses").head, BAD_LANGUAGE);
},
);
}
}