Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
77 changes: 75 additions & 2 deletions postgres/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -980,12 +980,17 @@ mod tests {
max(", max(tin.score(ctid, 0.05))", "body ==> 'ruby'"),
Some(0.0)
);
// tin computes no maximum under the score policy for several scans.
// Several scans report the summed maximum under either policy.
for quals in [
"title ==> 'gems' OR body ==> 'ruby'",
"body ==> 'ruby' OR id = 5",
] {
assert_eq!(max(", max(tin.score(ctid))", quals), None, "{quals}");
let score = Spi::get_one::<f32>(&format!(
"SELECT max(tin.score(ctid)) FROM lite_max_policy WHERE {quals}"
))
.unwrap();
assert!(score.is_some_and(|score| score > 0.0), "{quals}");
assert_eq!(max(", max(tin.score(ctid))", quals), score, "{quals}");
assert!(max("", quals).is_some_and(|max| max > 0.0), "{quals}");
}
assert!(
Expand All @@ -997,6 +1002,74 @@ mod tests {
);
}

#[pg_test]
fn max_score_reports_the_in_use_scorers_maximum_on_every_shape() {
Spi::run(
"CREATE TABLE lite_max_shapes (id int PRIMARY KEY, title text, body text);
INSERT INTO lite_max_shapes
SELECT g,
CASE WHEN g IN (1, 2) THEN 'zebra word' WHEN g = 3 THEN 'common'
ELSE 'filler ' || g END,
CASE WHEN g IN (2, 5) THEN 'zebra zebra text' WHEN g = 6 THEN 'common'
ELSE 'padding ' || g END
FROM generate_series(1, 40) AS g;
CREATE INDEX ON lite_max_shapes USING tin (title);
CREATE INDEX ON lite_max_shapes USING tin (body);",
)
.unwrap();
// Each maximum is the highest score of the in-use scoring function
// over the rows the quals admit. tin's expected values: 2.7828705
// (title) + 3.3875334 (body) for row 2, which matches both. Lead
// scores a row by every search it matches, not by the arm that
// admits it, so the split-arm OR is checked against the control only.
let both = 6.170_404;
let body = 3.387_533_4;
for (quals, expected) in [
("title ==> 'zebra' OR body ==> 'zebra'", Some(both)),
("title ==> 'zebra' AND body ==> 'zebra'", Some(both)),
("body ==> 'zebra' OR id = 1", Some(body)),
("(body ==> 'zebra' OR id = 1) AND id < 10", Some(body)),
("(title ==> 'zebra' AND id < 2) OR body ==> 'zebra'", None),
(
"(title ==> 'zebra' OR id = 3) AND body ==> 'zebra'",
Some(both),
),
] {
// Without a scoring function, max_score follows tin.full_score.
let policies = [
("tin.score(ctid)", true),
("tin.score(ctid, 0.04)", true),
("tin.full_score(ctid)", true),
("tin.full_score(ctid)", false),
];
for (scorer, beside) in policies {
let control = Spi::get_one::<f32>(&format!(
"SELECT max(s) FROM (SELECT {scorer} AS s FROM lite_max_shapes
WHERE {quals}) AS scored"
))
.unwrap();
// The aggregate over the scorer keeps it in the query.
let scored = if beside {
format!("max({scorer})")
} else {
"NULL::real".into()
};
let (low, high) = Spi::get_three::<f32, f32, f32>(&format!(
"SELECT min(tin.max_score(ctid)), max(tin.max_score(ctid)), {scored}
FROM lite_max_shapes WHERE {quals}"
))
.map(|(low, high, _)| (low, high))
.unwrap();
let context = format!("{scorer} beside={beside}: {quals}");
assert!(control.is_some(), "{context}");
assert_eq!((low, high), (control, control), "{context}");
if scorer != "tin.score(ctid, 0.04)" && expected.is_some() {
assert_eq!(control, expected, "{context}");
}
}
}
}

#[pg_test(
error = "tin.score() and tin.full_score() cannot be combined on one scanned relation; use one scoring function per relation (tin.max_score() adapts to either)"
)]
Expand Down
46 changes: 0 additions & 46 deletions postgres/src/score.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1222,45 +1222,6 @@ unsafe fn expanded_arguments(function: &pg_sys::FuncExpr) -> *mut pg_sys::List {
}
}

/// Reports whether a disjunction in the quals restricting relation `varno`
/// combines one of its searches with a qual that is not purely searches, as
/// in `body ==> 'a' OR id = 1`. tin scans such a disjunction with several
/// scans, and computes no `max_score` for them under the `tin.score`
/// policy.
unsafe fn has_mixed_disjunction(parse: *mut pg_sys::Query, varno: i32) -> bool {
let is_search = |node| unsafe {
crate::operator::search_operands(node)
.is_some_and(|(document, _)| single_varno(document) == Some(varno))
};
let mut mixed = false;
let clauses = unsafe { restriction_clauses(parse, varno) };
for clause in unsafe { PgList::<pg_sys::Node>::from_pg(clauses) }.iter_ptr() {
visit_disjunctions(clause, &mut |arms| {
let searches = arms.iter().any(|&arm| {
let mut found = false;
visit_searches(arm, &mut |node| found |= is_search(node));
found
});
mixed |= searches && !arms.iter().all(|&arm| only_searches(arm, &is_search));
});
}
mixed
}

/// Reports whether `node` combines nothing but searches with `AND` and `OR`.
fn only_searches(node: *mut pg_sys::Node, is_search: &impl Fn(*mut pg_sys::Node) -> bool) -> bool {
if node.is_null() || is_negation(node) {
return false;
}
if unsafe { (*node).type_ } != pg_sys::NodeTag::T_BoolExpr {
return is_search(node);
}
let expression = unsafe { &*node.cast::<pg_sys::BoolExpr>() };
unsafe { PgList::<pg_sys::Node>::from_pg(expression.args) }
.iter_ptr()
.all(|arm| only_searches(arm, is_search))
}

#[pg_extern(immutable, parallel_unsafe)]
fn score_support(request: Internal) -> Internal {
let unhandled = || Internal::from(Some(pg_sys::Datum::from(0_usize)));
Expand Down Expand Up @@ -1354,13 +1315,6 @@ fn score_support(request: Internal) -> Internal {
"max_score" => ScoreMode::MaxScore,
_ => ScoreMode::Score,
};
if mode == ScoreMode::MaxScore
&& (groups.len() > 1 || has_mixed_disjunction(parse, ctid.varno))
{
return Internal::from(Some(pg_sys::Datum::from(
make_null_const(pg_sys::FLOAT4OID) as usize,
)));
}
if mode.is_max() {
mark_required_groups(parse, rte, ctid.varno, &mut groups);
}
Expand Down