use std::time::Instant;
use yo_common::Result;
use yo_common::num::{DOUBLE_MAX, parse_f64, write_g17};
use yo_common::parse_i64;
use yo_search::Index;
use yo_search::explain::Note;
use yo_search::expr::Value;
use yo_search::field::Kind;
use yo_search::query::{self, Ask, Node, What, Yield};
use yo_search::score::Scorer;
use super::aggregate::{self, Reads};
use super::cursor::{Kept, Made};
use super::{Args, Asked, Order, Row, Rows, Watch};
use crate::dispatch::Server;
use crate::dispatch::args;
use crate::reply::Out;
const WINDOW: usize = 20;
const CONSTANT: f64 = 60.0;
const ALPHA: f64 = 0.3;
const BETA: f64 = 0.7;
const AWAY: &[u8] = b"__yo_hybrid_distance";
fn unknown(name: &[u8]) -> Vec<u8> {
about(name, b"Unknown argument")
}
fn twice(name: &[u8]) -> Vec<u8> {
about(name, b"Argument specified multiple times")
}
fn convert(name: &[u8]) -> Vec<u8> {
about(name, b"Could not convert argument to expected type")
}
fn about(name: &[u8], why: &[u8]) -> Vec<u8> {
let mut out = b"SEARCH_PARSE_ARGS ".to_vec();
out.extend_from_slice(name);
out.extend_from_slice(b": ");
out.extend_from_slice(why);
out
}
fn inside(name: &[u8], section: &[u8]) -> Vec<u8> {
let mut out = b"SEARCH_ARG_UNRECOGNIZED Unknown argument `".to_vec();
out.extend_from_slice(name);
out.extend_from_slice(b"` in ");
out.extend_from_slice(section);
out
}
fn plain(text: &str) -> Vec<u8> {
text.as_bytes().to_vec()
}
fn value(name: &[u8]) -> Vec<u8> {
let mut out = b"SEARCH_SYNTAX Invalid ".to_vec();
out.extend_from_slice(name);
out.extend_from_slice(b" value");
out
}
const LONGEST: i64 = 60_000;
const CAPPED: &str = "Query TIMEOUT exceeded the configured maximum (search-_max-foreground-timeout-limit) while search-workers is disabled; effective timeout was capped";
const NO_WORKING: &str = "SEARCH_PARSE_ARGS EXPLAINSCORE is not supported with GROUPBY";
fn g17(value: f64) -> String {
let mut buf = [0u8; DOUBLE_MAX];
String::from_utf8_lossy(write_g17(&mut buf, value)).into_owned()
}
enum Combine {
Rrf { constant: f64, window: usize },
Linear {
alpha: f64,
beta: f64,
window: usize,
},
}
impl Combine {
const fn window(&self) -> usize {
match self {
Combine::Rrf { window, .. } | Combine::Linear { window, .. } => *window,
}
}
}
impl Default for Combine {
fn default() -> Combine {
Combine::Rrf {
constant: CONSTANT,
window: WINDOW,
}
}
}
enum Reach {
Knn(u64),
Range(Box<[u8]>),
}
#[derive(Clone, Copy)]
struct Held<'a> {
front: &'a [Yield],
back: &'a [Yield],
defaults: bool,
}
struct Asks<'a> {
text: &'a [u8],
scorer: Scorer,
scored: Option<Box<[u8]>>,
field: &'a [u8],
param: &'a [u8],
reach: Option<Reach>,
options: Vec<(&'a [u8], &'a [u8])>,
filter: Option<&'a [u8]>,
nears: Option<Box<[u8]>>,
combine: Combine,
asked: Asked<'a>,
loaded: bool,
once: Vec<&'static [u8]>,
warn: Option<&'static str>,
explaining: bool,
grouped: bool,
}
pub(crate) fn hybrid(server: &Server, db: usize, args: Args<'_>, out: &mut Out) -> Result<()> {
let name = args.get(1);
let clock = Instant::now();
let mut reg = server.search.lock();
let Some(index) = reg.open(name) else {
super::Fail::naming(super::MISSING, name).write(out);
return Ok(());
};
let read = match reads(args, index, None) {
Ok(read) => read,
Err(text) => {
out.error(&text);
return Ok(());
}
};
let (text, vector) = match branches(&read, index) {
Ok(pair) => pair,
Err(text) => {
out.error(&text);
return Ok(());
}
};
let front: Vec<Yield> = read.scored.iter().map(|name| named(name)).collect();
let back: Vec<Yield> = read.nears.iter().map(|name| named(name)).collect();
let held = Held {
front: &front,
back: &back,
defaults: !read.loaded,
};
let asks = match reads(args, index, Some(held)) {
Ok(asks) => asks,
Err(text) => {
out.error(&text);
return Ok(());
}
};
let canon = index.name.clone();
if asks.asked.cursor.is_some()
&& let Err(fail) = super::cursor::room(server, &canon, 2)
{
fail.write(out);
return Ok(());
}
let (found, mut texts, mut nears) = walked(index, &asks, text, vector);
if let Some(want) = asks.asked.cursor {
let (Ok(mut mine), Ok(mut theirs)) = (
reads(args, index, Some(held)),
reads(args, index, Some(held)),
) else {
return Ok(());
};
let told = asks.explaining;
alone(
&mut mine.asked,
asks.scored.as_deref(),
asks.nears.as_deref(),
told,
);
alone(
&mut theirs.asked,
asks.nears.as_deref(),
asks.scored.as_deref(),
told,
);
drop(reg);
soloed(&mut texts, asks.scored.as_deref(), false);
soloed(&mut nears, asks.nears.as_deref(), true);
let said = asks.warn.map(|warn| warn.as_bytes().to_vec());
let one = branch(server, db, found, &texts, &mine.asked, said);
let two = branch(server, db, nears.len(), &nears, &theirs.asked, None);
let one = super::cursor::hold(server, &canon, one, want);
let two = super::cursor::hold(server, &canon, two, want);
match out.proto().is_resp3() {
true => out.map(3),
false => out.array(6),
}
out.bulk(b"SEARCH");
out.int(one as i64);
out.bulk(b"VSIM");
out.int(two as i64);
out.bulk(b"warnings");
out.array(0);
return Ok(());
}
let (total, rows) = merged(&asks, &texts, &nears);
drop(reg);
let spent = clock.elapsed();
writes(server, db, total, &rows, &asks, spent, out);
Ok(())
}
fn alone(asked: &mut Asked<'_>, mine: Option<&[u8]>, theirs: Option<&[u8]>, explaining: bool) {
let loaded: Vec<&[u8]> = asked.pipe.load.iter().map(|(_, name)| *name).collect();
asked.pipe.base.retain(|(name, from)| match from {
Reads::Score => false,
Reads::Field(..) => loaded.contains(&&**name),
_ => theirs != Some(&**name),
});
if let Some(mine) = mine
&& let Some(at) = asked.pipe.base.iter().position(|(name, _)| &**name == mine)
{
let held = asked.pipe.base.remove(at);
asked.pipe.base.insert(0, held);
}
asked.pipe.steps.clear();
asked.pipe.arrange = None;
asked.pipe.stage = None;
asked.rows.offset = 0;
asked.rows.count = usize::MAX;
asked.rows.scores = explaining;
}
fn soloed(rows: &mut [Row], name: Option<&[u8]>, near: bool) {
for row in rows {
let value = match near {
true => 1.0 / (1.0 + row.away(AWAY).unwrap_or(f64::INFINITY)),
false => row.score,
};
if near {
row.score = value;
}
if let Some(name) = name {
row.dists.push((name.into(), value));
}
}
}
fn branch(
server: &Server,
db: usize,
total: usize,
rows: &[Row],
asked: &Asked<'_>,
said: Option<Vec<u8>>,
) -> Kept {
let watch: Option<&mut Watch> = None;
let made = aggregate::runs(server, db, total, rows, asked, watch);
let total = made.start - made.gone.min(made.start);
let held: Vec<(Option<Row>, Vec<Value>)> = made
.table
.into_iter()
.map(|held| (held.from.map(|at| rows[at].clone()), held.values))
.collect();
Kept {
made: Made::Piped {
names: made.names,
rows: held,
sorted: made.sorted,
warning: made.warning.or(said),
},
walk: Vec::new(),
shows: asked.rolls(),
total,
whole: true,
loader: asked.pipe.loader,
offset: 0,
window: usize::MAX,
}
}
fn named(name: &[u8]) -> Yield {
Yield {
name: name.into(),
field: Box::default(),
asked: Box::default(),
ordered: false,
}
}
fn reads<'a>(
args: Args<'a>,
index: &Index,
held: Option<Held<'_>>,
) -> core::result::Result<Asks<'a>, Vec<u8>> {
let mut asked = Asked::default();
asked.rows.count = usize::MAX;
if let Some(held) = held {
asked.pipe.binding = true;
for want in held.front {
asked.pipe.base.push((want.name.clone(), Reads::Distance));
}
if held.defaults {
asked.pipe.base.push((b"__key".to_vec().into(), Reads::Key));
asked
.pipe
.base
.push((b"__score".to_vec().into(), Reads::Score));
}
for want in held.back {
asked.pipe.base.push((want.name.clone(), Reads::Distance));
}
}
let mut at = 2;
if !args::is(args.get(at), b"SEARCH") {
let Some(count) = counting(args.get(at)) else {
return Err(plain(
"SEARCH_SYNTAX Invalid subqueries count: expected an unsigned integer",
));
};
if count != 2 {
return Err(plain(
"SEARCH_PARSE_ARGS FT.HYBRID currently supports only two subqueries",
));
}
at += 1;
if !args::is(args.get(at), b"SEARCH") {
return Err(plain("SEARCH_PARSE_ARGS Missing required argument SEARCH"));
}
}
let text = args.get(at + 1);
at += 2;
let mut named_scorer = None;
let mut scored = None;
while at < args.len() && !args::is(args.get(at), b"VSIM") {
let word = args.get(at);
if args::is(word, b"SCORER") {
let Some(name) = args.opt(at + 1) else {
return Err(inside(word, b"SEARCH"));
};
named_scorer = Some(name);
at += 2;
continue;
}
if args::is(word, b"YIELD_SCORE_AS") {
let Some(name) = args.opt(at + 1) else {
return Err(inside(word, b"SEARCH"));
};
scored = Some(name.into());
at += 2;
continue;
}
return Err(inside(word, b"SEARCH"));
}
let mut scorer = Scorer::default_scorer();
if let Some(name) = named_scorer {
let Some(found) = Scorer::named(name) else {
let mut out = b"SEARCH_QUERY_BAD No such scorer ".to_vec();
out.extend_from_slice(name);
return Err(out);
};
scorer = found;
}
if at >= args.len() {
return Err(plain("SEARCH_PARSE_ARGS Missing required argument VSIM"));
}
at += 1;
let field = args.get(at);
let Some(field) = field.strip_prefix(b"@") else {
return Err(plain(
"SEARCH_SYNTAX Missing @ prefix for vector field name",
));
};
let Some(param) = args.opt(at + 1).and_then(|word| word.strip_prefix(b"$")) else {
return Err(plain(
"SEARCH_SYNTAX Invalid vector argument, expected a parameter name starting with $",
));
};
match index.field(field).map(|held| &held.kind) {
Some(Kind::Vector(_)) => {}
Some(_) => {
let mut out = b"SEARCH_SYNTAX Expected a VECTOR field `".to_vec();
out.extend_from_slice(field);
out.push(b'`');
return Err(out);
}
None => {
let mut out = b"SEARCH_SYNTAX Unknown field `".to_vec();
out.extend_from_slice(field);
out.push(b'`');
return Err(out);
}
}
at += 2;
let mut reach = None;
let mut options = Vec::new();
let mut nears = None;
let mut filter = None;
let mut stage = 0;
while at < args.len() {
let word = args.get(at);
let knn = args::is(word, b"KNN");
if stage < 1 && (knn || args::is(word, b"RANGE")) {
let section: &[u8] = match knn {
true => b"KNN",
false => b"RANGE",
};
at = neighbours(args, at, section, &mut reach, &mut options)?;
stage = 1;
continue;
}
if stage < 2 && args::is(word, b"FILTER") {
let Some(src) = args.opt(at + 1) else {
return Err(inside(word, b"VSIM"));
};
filter = Some(src);
at += 2;
stage = 2;
continue;
}
if stage < 3 && args::is(word, b"YIELD_SCORE_AS") {
let Some(name) = args.opt(at + 1) else {
return Err(inside(word, b"VSIM"));
};
nears = Some(name.into());
at += 2;
stage = 3;
continue;
}
break;
}
let mut asks = Asks {
text,
scorer,
scored,
field,
param,
reach,
options,
filter,
nears,
combine: Combine::default(),
asked,
loaded: false,
once: Vec::new(),
warn: None,
explaining: false,
grouped: false,
};
pipeline(args, at, index, &mut asks)?;
for (name, from) in &mut asks.asked.pipe.base {
if !matches!(from, Reads::Field(..)) {
continue;
}
if &**name == b"__key" {
*from = Reads::Key;
} else if &**name == b"__score" {
*from = Reads::Score;
}
}
Ok(asks)
}
fn pipeline<'a>(
args: Args<'a>,
from: usize,
index: &Index,
asks: &mut Asks<'a>,
) -> core::result::Result<(), Vec<u8>> {
let mut at = from;
while at < args.len() {
let word = args.get(at);
if args::is(word, b"COMBINE") {
if asks.only(b"COMBINE").is_err() {
at += 1;
continue;
}
at = combined(args, at, asks)?;
continue;
}
if args::is(word, b"DIALECT") {
return Err(plain(
"SEARCH_PARSE_ARGS DIALECT is not supported in FT.HYBRID or any of its subqueries. Please check the documentation on search-default-dialect configuration.",
));
}
if args::is(word, b"GROUPBY") {
asks.grouped = true;
if args.opt(at + 1).and_then(counting).is_none() {
return Err(about(b"GROUPBY", b"Invalid argument count"));
}
}
if let Some(next) = super::step(args, at, &mut asks.asked, index)? {
asks.loaded |= args::is(word, b"LOAD");
at = next;
continue;
}
if args::is(word, b"PARAMS") {
asks.only(b"PARAMS")?;
if args.opt(at + 1).and_then(counting).is_none() {
return Err(about(b"PARAMS", b"Invalid argument count"));
}
at = super::params(args, at, &mut asks.asked)?;
continue;
}
if args::is(word, b"LIMIT") {
asks.only(b"LIMIT")?;
}
if args::is(word, b"TIMEOUT") {
asks.only(b"TIMEOUT")?;
let Some(value) = args.opt(at + 1) else {
return Err(convert(b"TIMEOUT"));
};
let Some(read) = parse_i64(value) else {
return Err(convert(b"TIMEOUT"));
};
if read <= 0 || read > LONGEST {
asks.warn = Some(CAPPED);
}
at += 2;
continue;
}
if args::is(word, b"EXPLAINSCORE") {
asks.only(b"EXPLAINSCORE")?;
asks.explaining = true;
at += 1;
continue;
}
if args::is(word, b"WITHCURSOR") {
asks.only(b"WITHCURSOR")?;
}
if let Some(next) = super::plan(args, at, &mut asks.asked, super::Mode::Aggregate)? {
at = next;
continue;
}
return Err(unknown(word));
}
if asks.explaining && asks.grouped {
return Err(plain(NO_WORKING));
}
Ok(())
}
impl Asks<'_> {
fn only(&mut self, name: &'static [u8]) -> core::result::Result<(), Vec<u8>> {
if self.once.contains(&name) {
return Err(twice(name));
}
self.once.push(name);
Ok(())
}
}
fn neighbours<'a>(
args: Args<'a>,
at: usize,
section: &[u8],
reach: &mut Option<Reach>,
options: &mut Vec<(&'a [u8], &'a [u8])>,
) -> core::result::Result<usize, Vec<u8>> {
let knn = section == b"KNN";
let Some(count) = args.opt(at + 1).and_then(counting_zero) else {
return Err(plain(
"SEARCH_PARSE_ARGS Invalid argument count: expected an unsigned integer",
));
};
if count == 0 || count % 2 != 0 {
let mut out = b"SEARCH_SYNTAX Invalid argument count: ".to_vec();
out.extend_from_slice(count.to_string().as_bytes());
out.extend_from_slice(b" (must be a positive even number for key/value pairs)");
return Err(out);
}
let count = usize::try_from(count).unwrap_or(0);
let (want, extra): (&[u8], &[u8]) = match knn {
true => (b"K", b"EF_RUNTIME"),
false => (b"RADIUS", b"EPSILON"),
};
let mut found = None;
let mut next = at + 2;
let end = next + count;
while next < end {
let (Some(name), Some(held)) = (args.opt(next), args.opt(next + 1)) else {
return Err(plain(
"SEARCH_PARSE_ARGS Invalid argument count: expected an unsigned integer",
));
};
if args::is(name, want) {
found = Some(held);
} else if args::is(name, extra) {
let fine = match knn {
true => counting(held).is_some(),
false => parse_f64(held).is_some_and(|read| read > 0.0),
};
if !fine {
return Err(value(extra));
}
options.push((name, held));
} else {
return Err(inside(name, section));
}
next += 2;
}
let Some(found) = found else {
let mut out = b"SEARCH_PARSE_ARGS Missing required argument ".to_vec();
out.extend_from_slice(want);
return Err(out);
};
*reach = Some(match knn {
true => match counting(found) {
Some(k) => Reach::Knn(k),
None => return Err(value(b"K")),
},
false => match parse_f64(found).filter(|read| *read >= 0.0) {
Some(_) => Reach::Range(found.into()),
None => return Err(value(b"RADIUS")),
},
});
Ok(end)
}
fn combined(
args: Args<'_>,
at: usize,
asks: &mut Asks<'_>,
) -> core::result::Result<usize, Vec<u8>> {
let algo = args.get(at + 1);
let rrf = args::is(algo, b"RRF");
if !rrf && !args::is(algo, b"LINEAR") {
return Err(about(b"COMBINE", b"Invalid value for argument"));
}
let name: &[u8] = match rrf {
true => b"RRF",
false => b"LINEAR",
};
let Some(count) = args.opt(at + 2).and_then(counting_zero) else {
let mut out = b"SEARCH_PARSE_ARGS Invalid ".to_vec();
out.extend_from_slice(name);
out.extend_from_slice(
b" argument count, error: Could not convert argument to expected type",
);
return Err(out);
};
if count % 2 != 0 {
let mut out = b"SEARCH_PARSE_ARGS ".to_vec();
out.extend_from_slice(name);
out.extend_from_slice(
b" expects pairs of key value arguments, argument count must be an even number",
);
return Err(out);
}
let count = usize::try_from(count).unwrap_or(0);
let mut constant = CONSTANT;
let mut window = WINDOW;
let mut alpha = None;
let mut beta = None;
let mut next = at + 3;
let end = next + count;
while next < end {
let held = args.get(next);
let Some(value) = args.opt(next + 1) else {
let mut out = b"SEARCH_SYNTAX Missing value for ".to_vec();
out.extend_from_slice(held);
return Err(out);
};
let number = |name: &[u8]| parse_f64(value).ok_or_else(|| convert(name));
if args::is(held, b"WINDOW") {
let Some(read) = counting(value) else {
return match parse_i64(value).is_some() {
true => Err(about(b"WINDOW", b"Value below minimum")),
false => Err(convert(b"WINDOW")),
};
};
window = usize::try_from(read).unwrap_or(usize::MAX);
} else if rrf && args::is(held, b"CONSTANT") {
constant = number(b"CONSTANT")?;
} else if !rrf && args::is(held, b"ALPHA") {
alpha = Some(number(b"ALPHA")?);
} else if !rrf && args::is(held, b"BETA") {
beta = Some(number(b"BETA")?);
} else {
return Err(unknown(held));
}
next += 2;
}
if !rrf && (alpha.is_none() || beta.is_none()) {
let missing: &[u8] = match alpha.is_none() {
true => b"ALPHA",
false => b"BETA",
};
let mut out = b"SEARCH_SYNTAX Missing value for ".to_vec();
out.extend_from_slice(missing);
return Err(out);
}
asks.combine = match rrf {
true => Combine::Rrf { constant, window },
false => Combine::Linear {
alpha: alpha.unwrap_or(ALPHA),
beta: beta.unwrap_or(BETA),
window,
},
};
Ok(end)
}
fn counting(src: &[u8]) -> Option<u64> {
counting_zero(src).filter(|held| *held > 0)
}
fn counting_zero(src: &[u8]) -> Option<u64> {
let text = core::str::from_utf8(src).ok()?;
let body = text.strip_prefix('+').unwrap_or(text);
match body.strip_prefix("0x").or_else(|| body.strip_prefix("0X")) {
Some(hex) => u64::from_str_radix(hex, 16).ok(),
None => body.parse::<u64>().ok(),
}
}
fn branches(asks: &Asks<'_>, index: &Index) -> core::result::Result<(Node, Node), Vec<u8>> {
let ask = Ask {
dialect: 2,
params: &asks.asked.params,
verbatim: asks.asked.verbatim,
stopwords: asks.asked.stopwords,
};
let text = query::parse(asks.text, index, &ask).map_err(|bad| super::refused(&bad))?;
let mut src: Vec<u8> = Vec::new();
let wide = Reach::Knn(asks.combine.window() as u64);
match asks.reach.as_ref().unwrap_or(&wide) {
Reach::Knn(k) => {
src.extend_from_slice(b"*=>[KNN ");
src.extend_from_slice(k.to_string().as_bytes());
src.extend_from_slice(b" @");
src.extend_from_slice(asks.field);
src.extend_from_slice(b" $");
src.extend_from_slice(asks.param);
for (name, value) in &asks.options {
src.push(b' ');
src.extend_from_slice(name);
src.push(b' ');
src.extend_from_slice(value);
}
src.extend_from_slice(b" AS ");
src.extend_from_slice(AWAY);
src.push(b']');
}
Reach::Range(radius) => {
src.push(b'@');
src.extend_from_slice(asks.field);
src.extend_from_slice(b":[VECTOR_RANGE ");
src.extend_from_slice(radius);
src.extend_from_slice(b" $");
src.extend_from_slice(asks.param);
src.extend_from_slice(b"]=>{$yield_distance_as: ");
src.extend_from_slice(AWAY);
for (name, value) in &asks.options {
src.extend_from_slice(b"; $");
src.extend_from_slice(&name.to_ascii_lowercase());
src.extend_from_slice(b": ");
src.extend_from_slice(value);
}
src.push(b'}');
}
}
let mut vector = query::parse(&src, index, &ask).map_err(|bad| super::refused(&bad))?;
if let Some(src) = asks.filter {
let over = query::parse(src, index, &ask).map_err(|bad| super::refused(&bad))?;
match &mut vector.what {
What::Vector(held) if held.k.is_some() => held.over = Some(Box::new(over)),
_ => vector = Node::new(What::Intersect(vec![over, vector])),
}
}
Ok((text, vector))
}
#[derive(Default)]
struct Parts {
text: Option<(usize, f64, Option<Note>)>,
near: Option<(usize, f64)>,
}
fn walked(index: &Index, asks: &Asks<'_>, text: Node, vector: Node) -> (usize, Vec<Row>, Vec<Row>) {
let window = asks.combine.window();
let mut want = Rows {
scorer: asks.scorer,
count: usize::MAX,
explaining: asks.explaining,
..Rows::default()
};
let shaped = super::shape(text, index, &want);
let (found, mut texts, _) = super::gather(index, shaped, &want, Order::Ranked, true, false);
want.explaining = false;
let held = query::yields(&vector);
want.nearest = held
.iter()
.find(|held| held.ordered)
.map(|held| held.name.clone());
want.distance = held;
let shaped = super::shape(vector, index, &want);
let (_, mut nears, _) = super::gather(index, shaped, &want, Order::Forwards, true, false);
nears.sort_by(|left, right| {
let held = |row: &Row| row.away(AWAY).unwrap_or(f64::INFINITY);
held(left)
.partial_cmp(&held(right))
.unwrap_or(core::cmp::Ordering::Equal)
});
let reach = match &asks.reach {
Some(Reach::Knn(k)) => usize::try_from(*k).unwrap_or(usize::MAX).min(window),
Some(Reach::Range(_)) | None => window,
};
texts.truncate(window);
nears.truncate(reach);
(found, texts, nears)
}
fn merged(asks: &Asks<'_>, texts: &[Row], nears: &[Row]) -> (usize, Vec<Row>) {
let mut rows: Vec<Row> = Vec::new();
let mut parts: Vec<Parts> = Vec::new();
let mut where_at: Vec<(Box<[u8]>, usize)> = Vec::new();
let place = |rows: &mut Vec<Row>,
parts: &mut Vec<Parts>,
where_at: &mut Vec<(Box<[u8]>, usize)>,
key: &[u8]| match where_at.iter().find(|(held, _)| **held == *key)
{
Some((_, at)) => *at,
None => {
rows.push(Row {
key: key.into(),
score: 0.0,
payload: None,
note: None,
sort: None,
dists: Vec::new(),
});
parts.push(Parts::default());
where_at.push((key.into(), rows.len() - 1));
rows.len() - 1
}
};
for (rank, row) in texts.iter().enumerate() {
let at = place(&mut rows, &mut parts, &mut where_at, &row.key);
rows[at].score += match &asks.combine {
Combine::Rrf { constant, .. } => 1.0 / (constant + (rank + 1) as f64),
Combine::Linear { alpha, .. } => alpha * row.score,
};
if let Some(name) = &asks.scored {
rows[at].dists.push((name.clone(), row.score));
}
if asks.explaining {
parts[at].text = Some((rank + 1, row.score, row.note.clone()));
}
}
for (rank, row) in nears.iter().enumerate() {
let away = row.away(AWAY).unwrap_or(f64::INFINITY);
let close = 1.0 / (1.0 + away);
let at = place(&mut rows, &mut parts, &mut where_at, &row.key);
rows[at].score += match &asks.combine {
Combine::Rrf { constant, .. } => 1.0 / (constant + (rank + 1) as f64),
Combine::Linear { beta, .. } => beta * close,
};
if let Some(name) = &asks.nears {
rows[at].dists.push((name.clone(), close));
}
if asks.explaining {
parts[at].near = Some((rank + 1, close));
}
}
if asks.explaining {
for (at, row) in rows.iter_mut().enumerate() {
row.note = Some(working(asks, &parts[at], row.score));
}
}
rows.sort_by(|left, right| {
right
.score
.partial_cmp(&left.score)
.unwrap_or(core::cmp::Ordering::Equal)
.then_with(|| left.key.cmp(&right.key))
});
let total = rows.len();
(total, rows)
}
fn working(asks: &Asks<'_>, parts: &Parts, score: f64) -> Note {
let (head, inside) = match &asks.reach {
Some(Reach::Range(radius)) => (
format!(
"vector branch (RANGE: radius={:.4})",
parse_f64(radius).unwrap_or_default()
),
Some(Note::Line(format!(
"matched within radius = {}",
parts.near.is_some()
))),
),
Some(Reach::Knn(_)) | None => ("vector branch (KNN)".to_owned(), None),
};
let scorer = String::from_utf8_lossy(asks.scorer.name()).into_owned();
let (top, text, mut near) = match &asks.combine {
Combine::Rrf { constant, window } => {
let rank = |place: usize| format!("1 / (constant {constant:.2} + rank {place})");
let text = match &parts.text {
Some((place, _, note)) => (
rank(*place),
Note::Under(
format!("text rank = {place}"),
vec![Note::Under(
format!("Text scorer: {scorer}"),
note.clone().into_iter().collect(),
)],
),
),
None => (
"0 [text: no match]".to_owned(),
Note::Line("text rank = <no match>".to_owned()),
),
};
let near = match &parts.near {
Some((place, _)) => (
rank(*place),
vec![Note::Line(format!("vector rank = {place}"))],
),
None => (
"0 [vector: no match]".to_owned(),
vec![Note::Line("vector rank = <no match>".to_owned())],
),
};
(
format!("Hybrid score (RRF: window={window}, constant={constant:.2})"),
text,
near,
)
}
Combine::Linear {
alpha,
beta,
window,
} => {
let text = match &parts.text {
Some((_, held, note)) => {
let mut under = vec![Note::Line(format!("normalized text score = {held:.4}"))];
under.extend(note.clone());
(
format!("{alpha:.4} * {held:.4}"),
Note::Under(
format!(
"text contribution = {alpha:.4} * {held:.4} = {:.4}",
alpha * held
),
vec![Note::Under(format!("Text scorer: {scorer}"), under)],
),
)
}
None => (
"0 [text: no match]".to_owned(),
Note::Line("text contribution = <no match>".to_owned()),
),
};
let near = match &parts.near {
Some((_, close)) => (
format!("{beta:.4} * {close:.4}"),
vec![
Note::Line(format!(
"vector contribution = {beta:.4} * {close:.4} = {:.4}",
beta * close
)),
Note::Line(format!("normalized vector score = {close:.4}")),
],
),
None => (
"0 [vector: no match]".to_owned(),
vec![Note::Line("vector contribution = <no match>".to_owned())],
),
};
(
format!("Hybrid score (LINEAR: alpha={alpha:.4}, beta={beta:.4}, window={window})"),
text,
near,
)
}
};
if let Some(line) = inside {
near.1.insert(0, line);
}
Note::Under(
format!("final score: {} + {} = {}", text.0, near.0, g17(score)),
vec![Note::Under(top, vec![text.1, Note::Under(head, near.1)])],
)
}
fn writes(
server: &Server,
db: usize,
total: usize,
rows: &[Row],
asks: &Asks<'_>,
spent: core::time::Duration,
out: &mut Out,
) {
let watch: Option<&mut Watch> = None;
let made = aggregate::runs(server, db, total, rows, &asks.asked, watch);
let shown: Vec<(Option<&Row>, &Vec<yo_search::expr::Value>)> = made
.table
.iter()
.map(|held| (held.from.map(|at| &rows[at]), &held.values))
.collect();
let deep = out.proto().is_resp3();
match deep {
true => out.map(4),
false => out.array(8),
}
out.bulk(b"total_results");
out.int((made.start - made.gone.min(made.start)) as i64);
out.bulk(b"results");
out.array(shown.len());
for (row, values) in &shown {
let told = match asks.explaining {
true => row.and_then(|row| row.note.as_ref().map(|note| (row.score, note))),
false => None,
};
aggregate::mapped(&made.names, values, None, told, out);
}
out.bulk(b"warnings");
let said = made.warning.as_deref().or(asks.warn.map(str::as_bytes));
match said {
Some(warning) => {
out.array(1);
out.bulk(warning);
}
None => out.array(0),
}
out.bulk(b"execution_time");
let ms = spent.as_secs_f64() * 1000.0;
match deep {
true => out.double(ms),
false => out.bulk(format!("{ms:.6}").as_bytes()),
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::dispatch::search::aggregate::Shape;
fn base() -> Asked<'static> {
let mut asked = Asked::default();
asked.pipe.base = vec![
(b"ts".to_vec().into(), Reads::Distance),
(b"__key".to_vec().into(), Reads::Key),
(b"__score".to_vec().into(), Reads::Score),
(b"vs".to_vec().into(), Reads::Distance),
];
asked
}
fn names(asked: &Asked<'_>) -> Vec<String> {
asked
.pipe
.base
.iter()
.map(|(name, _)| String::from_utf8_lossy(name).into_owned())
.collect()
}
#[test]
fn a_branch_keeps_its_own_yield_and_puts_it_first() {
let mut asked = base();
alone(&mut asked, Some(b"vs"), Some(b"ts"), false);
assert_eq!(names(&asked), vec!["vs", "__key"]);
let mut asked = base();
alone(&mut asked, Some(b"ts"), Some(b"vs"), false);
assert_eq!(names(&asked), vec!["ts", "__key"]);
}
#[test]
fn a_field_a_step_named_goes_and_a_field_a_load_named_stays() {
let mut asked = base();
asked.pipe.load = vec![(b"t", b"t")];
asked.pipe.base.push((
b"t".to_vec().into(),
Reads::Field(b"t".to_vec().into(), Shape::Words),
));
asked.pipe.base.push((
b"n".to_vec().into(),
Reads::Field(b"n".to_vec().into(), Shape::Number),
));
alone(&mut asked, Some(b"ts"), Some(b"vs"), false);
assert_eq!(names(&asked), vec!["ts", "__key", "t"]);
}
#[test]
fn a_branch_runs_no_step_and_takes_no_window() {
let mut asked = base();
asked.rows.offset = 2;
asked.rows.count = 3;
asked.pipe.arrange = Some(0);
alone(&mut asked, None, None, true);
assert!(asked.pipe.steps.is_empty());
assert!(asked.pipe.arrange.is_none());
assert_eq!(asked.rows.offset, 0);
assert_eq!(asked.rows.count, usize::MAX);
assert!(asked.rows.scores);
}
#[test]
fn the_vector_branch_scores_a_row_by_how_close_it_came() {
let row = |away: f64| Row {
key: b"k".to_vec().into(),
score: 0.0,
payload: None,
note: None,
sort: None,
dists: vec![(AWAY.to_vec().into(), away)],
};
let mut rows = vec![row(0.0), row(1.0), row(4.0)];
soloed(&mut rows, Some(b"vs"), true);
let scores: Vec<f64> = rows.iter().map(|row| row.score).collect();
assert_eq!(scores, vec![1.0, 0.5, 0.2]);
assert_eq!(rows[2].away(b"vs"), Some(0.2));
}
#[test]
fn the_text_branch_leaves_the_score_where_the_scorer_put_it() {
let mut rows = vec![Row {
key: b"k".to_vec().into(),
score: 0.75,
payload: None,
note: None,
sort: None,
dists: Vec::new(),
}];
soloed(&mut rows, Some(b"ts"), false);
assert_eq!(rows[0].score, 0.75);
assert_eq!(rows[0].away(b"ts"), Some(0.75));
}
}