use super::{
const_f64, const_value, lower_bayesian_evidence_fusion, lower_bayesian_match_with_prior,
lower_calibrated_vector_match, lower_graph_function, lower_learned_fusion,
lower_multi_field_match, lower_operator_arg, lower_positive_evidence_pool,
lower_staged_retrieval, try_lower_attention_fusion, try_lower_fts_match, try_lower_knn_match,
try_lower_text_match, BinaryOp, Predicate, RetrievalConstants, RetrievalExpr, ScalarExpr,
TextScoringMode,
};
pub(super) fn lower_document_boolean(
mut children: Vec<RetrievalExpr>,
union: bool,
) -> RetrievalExpr {
let graph_children = children
.iter()
.filter(|child| tree_returns_graph(child))
.count();
if graph_children > 0 && graph_children < children.len() {
children = children
.into_iter()
.map(|child| {
if tree_returns_graph(&child) {
RetrievalExpr::EncodeGraphPosting {
source: Box::new(child),
}
} else {
child
}
})
.collect();
}
if union {
RetrievalExpr::Union(children)
} else {
RetrievalExpr::Intersect(children)
}
}
pub(super) fn tree_returns_graph(tree: &RetrievalExpr) -> bool {
match tree {
RetrievalExpr::Traverse { .. }
| RetrievalExpr::RegularPathQuery { .. }
| RetrievalExpr::PageRank { .. }
| RetrievalExpr::HITS { .. }
| RetrievalExpr::BetweennessCentrality { .. }
| RetrievalExpr::TemporalTraverse { .. } => true,
RetrievalExpr::Intersect(children) | RetrievalExpr::Union(children) => {
!children.is_empty() && children.iter().all(tree_returns_graph)
}
RetrievalExpr::Composed(children) => children.last().is_some_and(tree_returns_graph),
_ => false,
}
}
pub(super) fn lower_function(
name: &str,
args: &[ScalarExpr],
constants: &RetrievalConstants<'_>,
) -> Option<RetrievalExpr> {
let lower = name.to_ascii_lowercase();
match lower.as_str() {
"text_match" => {
try_lower_text_match("text_match", args, constants, TextScoringMode::BM25).ok()
}
"bayesian_match" => try_lower_text_match(
"bayesian_match",
args,
constants,
TextScoringMode::BayesianBM25,
)
.ok(),
"fts_match" => try_lower_fts_match(args, constants).ok(),
"bayesian_match_with_prior" => lower_bayesian_match_with_prior(args, constants),
"calibrated_vector_match" => lower_calibrated_vector_match(args, constants),
"knn_match" => try_lower_knn_match(args, constants).ok(),
"fuse_bayesian_evidence" | "fuse_log_odds" => {
lower_bayesian_evidence_fusion(args, constants)
}
"pool_positive_evidence" => lower_positive_evidence_pool(args, constants),
"multi_field_match" => lower_multi_field_match(args, constants),
"staged_retrieval" => lower_staged_retrieval(args, constants),
"attention" | "fuse_attention" | "fuse_multihead" => {
try_lower_attention_fusion(&lower, args, constants).ok()
}
"learned_fusion" | "fuse_learned" => lower_learned_fusion(args, constants),
"sparse_threshold" => {
if args.len() != 2 {
return None;
}
let source = lower_operator_arg(args.first()?, constants)?;
let threshold = const_f64(args.get(1)?, constants)?;
Some(RetrievalExpr::SparseThreshold {
source: Box::new(source),
threshold,
})
}
_ => lower_graph_function(&lower, args, constants),
}
}
pub(super) fn lower_comparison(
op: BinaryOp,
lhs: &ScalarExpr,
rhs: &ScalarExpr,
constants: &RetrievalConstants<'_>,
) -> Option<RetrievalExpr> {
let (col_expr, val_expr, swap) = match (column_name(lhs), column_name(rhs)) {
(Some(_), _) => (lhs, rhs, false),
(None, Some(_)) => (rhs, lhs, true),
_ => return None,
};
let field = column_name(col_expr)?;
let value = const_value(val_expr, constants)?;
let predicate = match (op, swap) {
(BinaryOp::Equal, _) => Predicate::Equals(value),
(BinaryOp::NotEqual, _) => Predicate::NotEquals(value),
(BinaryOp::Less, false) | (BinaryOp::Greater, true) => Predicate::LessThan(value),
(BinaryOp::LessEqual, false) | (BinaryOp::GreaterEqual, true) => {
Predicate::LessThanOrEqual(value)
}
(BinaryOp::Greater, false) | (BinaryOp::Less, true) => Predicate::GreaterThan(value),
(BinaryOp::GreaterEqual, false) | (BinaryOp::LessEqual, true) => {
Predicate::GreaterThanOrEqual(value)
}
_ => return None,
};
Some(RetrievalExpr::Filter {
field,
predicate,
source: None,
})
}
pub(super) fn column_name(expr: &ScalarExpr) -> Option<String> {
match expr {
ScalarExpr::Column(name) => Some(name.clone()),
ScalarExpr::QualifiedColumn { column, .. } => Some(column.clone()),
_ => None,
}
}