From 99614a3166d9dab2f78b3cb9012507961088fe78 Mon Sep 17 00:00:00 2001 From: Claude Date: Tue, 6 Oct 2026 18:07:04 +0000 Subject: [PATCH 1/6] fix(score): sum scores across indexed columns Scoring bound only the first searched column that had a tin index, so a search over several indexed columns scored one of them and the result depended on the order of the predicates. Group the searches by the tin index each binds to, score every group, and sum the groups that match the row, as tin does. max_score takes the highest summed score among documents that match every group the quals require. Co-Authored-By: Claude Opus 5.5 (1M context) Claude-Session: https://claude.ai/code/session_012mKGtrg7DEhVPgqYrJtHGg --- postgres/src/lib.rs | 4 +- postgres/src/score.rs | 509 ++++++++++++++++++++++++++++++------------ 2 files changed, 372 insertions(+), 141 deletions(-) diff --git a/postgres/src/lib.rs b/postgres/src/lib.rs index 88cb361..27b035c 100644 --- a/postgres/src/lib.rs +++ b/postgres/src/lib.rs @@ -419,8 +419,8 @@ mod tests { Spi::run( "CREATE TABLE lite_stale (title text); CREATE INDEX ON lite_stale USING tin (title); - ALTER FUNCTION tin.score_bound(text, text[], int4, int4, int4, - float4, float4, float4, text[], text[]) RENAME TO score_bound_stale;", + ALTER FUNCTION tin.score_bound(text[], text[], int4[], int4, int4[], bool[], + int4, float4, float4, float4, text[], text[]) RENAME TO score_bound_stale;", ) .unwrap(); Spi::run("SELECT tin.score(ctid) FROM lite_stale WHERE title ==> 'lorem'").unwrap(); diff --git a/postgres/src/score.rs b/postgres/src/score.rs index 9c14604..c8200d4 100644 --- a/postgres/src/score.rs +++ b/postgres/src/score.rs @@ -67,14 +67,23 @@ impl TryFrom for ScoreMode { } } -/// Identifies one search that a `score_bound` call site scores rows against. -/// The search texts can vary per row, as in `t.body ==> q.text` driven by a -/// join, so a call site may hold several of these. +/// One tin index that a `score_bound` call site scores rows against, with the +/// searches bound to it. `None` stands for a NULL search expression, which +/// matches no rows. A required group's searches restrict every output row. +#[derive(Clone, Debug, Eq, Hash, PartialEq)] +struct GroupKey { + index_oid: u32, + queries: Option>, + required: bool, +} + +/// Identifies the searches that a `score_bound` call site scores rows +/// against. The search texts can vary per row, as in `t.body ==> q.text` +/// driven by a join, so a call site may hold several of these. #[derive(Clone, Debug, Eq, Hash, PartialEq)] struct CorpusKey { heap_oid: u32, - index_oid: u32, - queries: Vec, + groups: Vec, mode: ScoreMode, dense: u32, k1: Option, @@ -83,21 +92,52 @@ struct CorpusKey { replace: Option>, } -/// Corpus-wide statistics for one search: the retained terms with N, avgdl, -/// and df baked into their scorers, and the best score among matching -/// documents for the `max_score` modes. -struct ScoreCorpus { +/// Corpus-wide statistics for the searches bound to one tin index: the +/// retained terms with N, avgdl, and df baked into their scorers. +struct GroupCorpus { tokenizer: CompiledTokenizerPipeline, + query: Query, scorers: Vec<(String, TermScorer)>, +} + +impl GroupCorpus { + /// Scores `document` against this index's searches. A document they do + /// not match, which another index's searches admitted, gets no score. + fn score(&self, document: &str) -> Option { + evaluate(&self.query, &tokenize_doc(document, &self.tokenizer)) + .unwrap_or_else(|error| pgrx::error!("tin score query evaluation failed: {error}")) + .matched + .then(|| score_tokens(&self.scorers, &tokenize(&self.tokenizer, document))) + } +} + +/// The groups of one call site, aligned with its `documents` argument, and +/// the best score among matching rows for the `max_score` modes. +struct ScoreCorpus { + groups: Vec>, max: f32, } impl ScoreCorpus { - fn score(&self, document: &str) -> f32 { - score_tokens(&self.scorers, &tokenize(&self.tokenizer, document)) + /// Returns the score of each group whose searches match the row. + fn group_scores(&self, documents: &[Option]) -> Vec<(usize, f32)> { + self.groups + .iter() + .zip(documents) + .enumerate() + .filter_map(|(position, (group, document))| { + Some((position, group.as_ref()?.score(document.as_deref()?)?)) + }) + .collect() } } +/// Sums per-index scores the way tin does: exactly, rounding once to `f32`, +/// so the total does not depend on the order of the indexes. +fn sum_group_scores(scores: impl IntoIterator) -> f32 { + scores.into_iter().map(f64::from).sum::() as f32 +} + type CorpusMemo = FxHashMap; fn score_context_error(function: &str) -> ! { @@ -149,10 +189,12 @@ fn bits(value: Option) -> Option { reason = "SQL signature used by the scoring support function" )] fn score_bound( - document: &str, + documents: Vec>, query: Vec>, + query_groups: Vec, heap_oid: i32, - index_oid: i32, + index_oids: Vec, + required: Vec, mode: i32, dense_ratio: Option, k1: Option, @@ -163,14 +205,29 @@ fn score_bound( ) -> f32 { let mode = ScoreMode::try_from(mode) .unwrap_or_else(|mode| pgrx::error!("tin.score_bound(): unknown score mode {mode}")); - // A NULL search expression matches no rows, so nothing needs a score. - let Some(queries) = query.into_iter().collect::>>() else { - return 0.0; - }; + let mut groups = index_oids + .iter() + .zip(&required) + .map(|(&index_oid, &required)| GroupKey { + index_oid: index_oid as u32, + queries: Some(Vec::new()), + required, + }) + .collect::>(); + for (search, group) in query.into_iter().zip(query_groups) { + let group = &mut groups[group as usize]; + group.queries = group + .queries + .take() + .zip(search) + .map(|(mut queries, search)| { + queries.push(search); + queries + }); + } let key = CorpusKey { heap_oid: heap_oid as u32, - index_oid: index_oid as u32, - queries, + groups, mode, dense: dense_ratio.unwrap_or(DenseRatio::DEFAULT).to_bits(), k1: bits(k1), @@ -185,7 +242,12 @@ fn score_bound( if mode.is_max() { corpus.max } else { - corpus.score(document) + sum_group_scores( + corpus + .group_scores(&documents) + .into_iter() + .map(|(_, score)| score), + ) } } @@ -212,17 +274,81 @@ fn build_corpus( term_add: Option>, term_replace: Option>, ) -> ScoreCorpus { - let full = key.mode.is_full(); let heap_oid = pg_sys::Oid::from(key.heap_oid); - let index = unsafe { - PgRelation::with_lock( - pg_sys::Oid::from(key.index_oid), - pg_sys::AccessShareLock as _, - ) - }; - if unsafe { pg_sys::IndexGetRelation(index.oid(), false) } != heap_oid { - pgrx::error!("tin score index no longer belongs to the scored relation"); + let indexes = key + .groups + .iter() + .map(|group| { + let index = unsafe { + PgRelation::with_lock( + pg_sys::Oid::from(group.index_oid), + pg_sys::AccessShareLock as _, + ) + }; + if unsafe { pg_sys::IndexGetRelation(index.oid(), false) } != heap_oid { + pgrx::error!("tin score index no longer belongs to the scored relation"); + } + index + }) + .collect::>(); + let rows = load_documents( + heap_oid, + &indexes.iter().map(PgRelation::oid).collect::>(), + ); + let groups = key + .groups + .iter() + .zip(&indexes) + .enumerate() + .map(|(position, (group, index))| { + let queries = group.queries.as_deref()?; + let documents = rows.iter().filter_map(|row| row[position].as_deref()); + Some(build_group( + key, + index, + queries, + documents, + k1, + b, + term_add.clone(), + term_replace.clone(), + )) + }) + .collect::>(); + let mut corpus = ScoreCorpus { groups, max: 0.0 }; + if key.mode.is_max() { + for (row, documents) in rows.iter().enumerate() { + if row.is_multiple_of(10) { + pgrx::check_for_interrupts!(); + } + let scores = corpus.group_scores(documents); + let complete = key.groups.iter().enumerate().all(|(position, group)| { + !group.required || scores.iter().any(|&(matched, _)| matched == position) + }); + if complete && !scores.is_empty() { + let score = sum_group_scores(scores.into_iter().map(|(_, score)| score)); + corpus.max = corpus.max.max(score); + } + } } + corpus +} + +#[expect( + clippy::too_many_arguments, + reason = "one index's share of the scoring call's arguments" +)] +fn build_group<'a>( + key: &CorpusKey, + index: &PgRelation, + queries: &[String], + documents: impl Iterator, + k1: Option, + b: Option, + term_add: Option>, + term_replace: Option>, +) -> GroupCorpus { + let full = key.mode.is_full(); let tokenizer = unsafe { crate::options::tokenizer(index.as_ptr()) }; let defaults = unsafe { crate::options::bm25(index.as_ptr()) }; let stop_csv = unsafe { crate::options::score_stop_words(index.as_ptr()) }; @@ -234,7 +360,7 @@ fn build_corpus( if !full && !dense.is_valid() { pgrx::error!("dense_ratio must be finite and non-negative"); } - let query = crate::operator::parse_searches(&key.queries, &tokenizer); + let query = crate::operator::parse_searches(queries, &tokenizer); let mut inputs = Vec::new(); collect_score_terms(&query, 1.0, false, &mut inputs); let edit = TermSetEdit::from_bound_arrays(term_add, term_replace) @@ -251,8 +377,7 @@ fn build_corpus( stop_csv.as_deref().and_then(ScoreStopWords::from_csv) }; let terms = compile_scoring_terms(inputs, &edit, stop.as_ref()); - let documents = load_documents(heap_oid, index.oid()); - let tokenized = tokenize_documents(&documents, &tokenizer); + let tokenized = tokenize_documents(documents, &tokenizer); let total_docs = tokenized.len() as u64; let average_length = if total_docs == 0 { 1.0 @@ -274,21 +399,10 @@ fn build_corpus( .unwrap_or_else(|error| pgrx::error!("tin score parameters: {error}")); scorers.push((term.text().to_owned(), scorer)); } - let mut max = 0.0_f32; - if key.mode.is_max() { - for (document, tokens) in documents.iter().zip(&tokenized) { - let matched = evaluate(&query, &tokenize_doc(document, &tokenizer)) - .unwrap_or_else(|error| pgrx::error!("tin score query evaluation failed: {error}")) - .matched; - if matched { - max = max.max(score_tokens(&scorers, tokens)); - } - } - } - ScoreCorpus { + GroupCorpus { tokenizer, + query, scorers, - max, } } @@ -303,7 +417,10 @@ fn score_tokens(scorers: &[(String, TermScorer)], tokens: &[String]) -> f32 { })) } -fn load_documents(heap_oid: pg_sys::Oid, index_oid: pg_sys::Oid) -> Vec { +/// Reads the documents of the given tin indexes on `heap_oid`, one row per +/// heap row that any of them covers, with each index's document in its +/// position, or `None` where the row is NULL or outside that index. +fn load_documents(heap_oid: pg_sys::Oid, index_oids: &[pg_sys::Oid]) -> Vec>> { unsafe { let relname = pg_sys::get_rel_name(heap_oid); let namespace = pg_sys::get_namespace_name(pg_sys::get_rel_namespace(heap_oid)); @@ -311,43 +428,47 @@ fn load_documents(heap_oid: pg_sys::Oid, index_oid: pg_sys::Oid) -> Vec pgrx::error!("tin score relation no longer exists"); } let qualified = pg_sys::quote_qualified_identifier(namespace, relname); - let index_sql = format!( - "SELECT CASE WHEN i.indkey[0] = 0 \ - THEN pg_catalog.pg_get_expr(i.indexprs, i.indrelid) \ - ELSE pg_catalog.quote_ident(a.attname) END, \ - pg_catalog.pg_get_expr(i.indpred, i.indrelid) \ - FROM pg_catalog.pg_index i \ - LEFT JOIN pg_catalog.pg_attribute a \ - ON a.attrelid=i.indrelid AND a.attnum=i.indkey[0] \ - WHERE i.indexrelid={}::oid AND i.indrelid={}::oid", - index_oid.to_u32(), - heap_oid.to_u32(), - ); - let mut columns = select_read_only(&index_sql, "tin score index lookup failed") - .pop() - .unwrap_or_else(|| pgrx::error!("tin score index lookup failed: index not found")) - .into_iter(); - let expression = columns - .next() - .flatten() - .unwrap_or_else(|| pgrx::error!("tin score index expression no longer exists")); - let predicate = columns - .next() - .flatten() - .map(|predicate| format!(" AND ({predicate})")) - .unwrap_or_default(); + let mut columns = Vec::new(); + let mut conditions = Vec::new(); + for index_oid in index_oids { + let index_sql = format!( + "SELECT CASE WHEN i.indkey[0] = 0 \ + THEN pg_catalog.pg_get_expr(i.indexprs, i.indrelid) \ + ELSE pg_catalog.quote_ident(a.attname) END, \ + pg_catalog.pg_get_expr(i.indpred, i.indrelid) \ + FROM pg_catalog.pg_index i \ + LEFT JOIN pg_catalog.pg_attribute a \ + ON a.attrelid=i.indrelid AND a.attnum=i.indkey[0] \ + WHERE i.indexrelid={}::oid AND i.indrelid={}::oid", + index_oid.to_u32(), + heap_oid.to_u32(), + ); + let mut definition = select_read_only(&index_sql, "tin score index lookup failed") + .pop() + .unwrap_or_else(|| pgrx::error!("tin score index lookup failed: index not found")) + .into_iter(); + let expression = definition + .next() + .flatten() + .unwrap_or_else(|| pgrx::error!("tin score index expression no longer exists")); + let predicate = definition + .next() + .flatten() + .map(|predicate| format!(" AND ({predicate})")) + .unwrap_or_default(); + let condition = format!("({expression}) IS NOT NULL{predicate}"); + columns.push(format!( + "CASE WHEN {condition} THEN ({expression})::text END" + )); + conditions.push(format!("({condition})")); + } let sql = format!( - "SELECT ({expression})::text FROM {} WHERE ({expression}) IS NOT NULL{predicate}", + "SELECT {} FROM {} WHERE {}", + columns.join(", "), CStr::from_ptr(qualified).to_string_lossy(), + conditions.join(" OR "), ); select_read_only(&sql, "tin score corpus scan failed") - .into_iter() - .map(|mut row| { - row.pop() - .flatten() - .expect("corpus query excludes null documents") - }) - .collect() } } @@ -381,12 +502,11 @@ fn select_read_only(sql: &str, failure: &str) -> Vec>> { }) } -fn tokenize_documents( - documents: &[String], +fn tokenize_documents<'a>( + documents: impl Iterator, tokenizer: &CompiledTokenizerPipeline, ) -> Vec> { documents - .iter() .enumerate() .map(|(row, document)| { if row.is_multiple_of(10) { @@ -509,8 +629,8 @@ fn score_inspect( if !ratio.is_valid() { pgrx::error!("dense_ratio must be finite and non-negative"); } - let docs = load_documents(heap_oid, index.oid()); - let tokenized = tokenize_documents(&docs, &tokenizer); + let docs = load_documents(heap_oid, &[index.oid()]); + let tokenized = tokenize_documents(docs.iter().filter_map(|row| row[0].as_deref()), &tokenizer); let n = tokenized.len() as u64; let rows = terms .into_iter() @@ -545,21 +665,32 @@ unsafe extern "C-unwind" fn find_qual(node: *mut pg_sys::Node, context: *mut c_v return false; } let binding = unsafe { &mut *context.cast::() }; - if unsafe { (*node).type_ } == pg_sys::NodeTag::T_OpExpr { + if let Some(search) = unsafe { search_operands(node) } { + binding.matches.push(search); + } + unsafe { pg_sys::expression_tree_walker(node, Some(find_qual), context) } +} + +/// Returns the document and query operands of a `==>` search. +unsafe fn search_operands( + node: *mut pg_sys::Node, +) -> Option<(*mut pg_sys::Node, *mut pg_sys::Node)> { + unsafe { + if (*node).type_ != pg_sys::NodeTag::T_OpExpr { + return None; + } let op = node.cast::(); - let name = unsafe { pg_sys::get_opname((*op).opno) }; - if !name.is_null() - && unsafe { CStr::from_ptr(name) }.to_bytes() == b"==>" - && unsafe { pg_sys::list_length((*op).args) } == 2 + let name = pg_sys::get_opname((*op).opno); + if name.is_null() + || CStr::from_ptr(name).to_bytes() != b"==>" + || pg_sys::list_length((*op).args) != 2 { - let left = unsafe { pg_sys::list_nth((*op).args, 0).cast::() }; - let right = unsafe { pg_sys::list_nth((*op).args, 1).cast::() }; - if !left.is_null() { - binding.matches.push((left, right)); - } + return None; } + let left = pg_sys::list_nth((*op).args, 0).cast::(); + let right = pg_sys::list_nth((*op).args, 1).cast::(); + (!left.is_null()).then_some((left, right)) } - unsafe { pg_sys::expression_tree_walker(node, Some(find_qual), context) } } struct VarnoBinding { @@ -792,6 +923,87 @@ unsafe fn refuse_unindexed_scoring( unreachable!("ERROR reports do not return") } +/// The searches on the scored relation that bind to one tin index. +struct SearchGroup { + index_oid: pg_sys::Oid, + document: *mut pg_sys::Node, + queries: Vec<*mut pg_sys::Node>, + /// Whether every output row satisfies one of these searches. + required: bool, +} + +/// Groups the relation's searches by the tin index each binds to, in index +/// OID order, so that a row's score sums the same per-index scores whatever +/// the order of the quals. Searches without a tin index do not score, and +/// scoring is refused when none of them has one, as tin does. +unsafe fn bind_search_groups( + parse: *mut pg_sys::Query, + rte: *mut pg_sys::RangeTblEntry, + varno: i32, + searches: &[(*mut pg_sys::Node, *mut pg_sys::Node)], +) -> Vec { + let mut bound: Vec<(*mut pg_sys::Node, Option)> = Vec::new(); + let mut groups: Vec = Vec::new(); + let mut unindexed = Vec::new(); + for &(document, query) in searches { + let index_oid = match bound + .iter() + .find(|(operand, _)| unsafe { pg_sys::equal((*operand).cast(), document.cast()) }) + { + Some(&(_, index_oid)) => index_oid, + None => { + let index_oid = + unsafe { find_matching_tin_index(parse, (*rte).relid, varno, document) }; + bound.push((document, index_oid)); + index_oid + } + }; + let Some(index_oid) = index_oid else { + unindexed.push(document); + continue; + }; + match groups.iter_mut().find(|group| group.index_oid == index_oid) { + Some(group) => group.queries.push(query), + None => groups.push(SearchGroup { + index_oid, + document, + queries: vec![query], + required: false, + }), + } + } + if groups.is_empty() { + unsafe { refuse_unindexed_scoring(rte, varno, &unindexed) }; + } + groups.sort_by_key(|group| group.index_oid.to_u32()); + groups +} + +/// Marks the groups with a search among the quals that restrict every output +/// row. `max_score` considers only documents that match every such group. +unsafe fn mark_required_groups( + parse: *mut pg_sys::Query, + rte: *mut pg_sys::RangeTblEntry, + varno: i32, + groups: &mut [SearchGroup], +) { + unsafe { + let clauses = restriction_clauses(parse, varno); + for clause in PgList::::from_pg(clauses).iter_ptr() { + let Some((document, _)) = search_operands(clause) else { + continue; + }; + let Some(index_oid) = find_matching_tin_index(parse, (*rte).relid, varno, document) + else { + continue; + }; + for group in groups.iter_mut() { + group.required |= group.index_oid == index_oid; + } + } + } +} + /// Renders a search expression on relation `varno` against the relation's /// own name. Returns `None` for expressions a single-relation deparse context /// cannot describe. @@ -841,7 +1053,7 @@ unsafe fn deparse_search_expression( struct FullScoreBinding { ctid: *const pg_sys::Var, - document: *mut pg_sys::Node, + documents: *mut pg_sys::Node, support: pg_sys::Oid, bound: pg_sys::Oid, } @@ -857,10 +1069,10 @@ unsafe extern "C-unwind" fn has_full_score(node: *mut pg_sys::Node, context: *mu let function = &*node.cast::(); // Earlier query clauses may already contain the rewritten scorer. if function.funcid == binding.bound { - let mode = pg_sys::list_nth(function.args, 4).cast::(); + let mode = pg_sys::list_nth(function.args, 6).cast::(); if (*mode).xpr.type_ == pg_sys::NodeTag::T_Const && (*mode).constvalue.value() == ScoreMode::FullScore as usize - && pg_sys::equal(pg_sys::list_nth(function.args, 0), binding.document.cast()) + && pg_sys::equal(pg_sys::list_nth(function.args, 0), binding.documents.cast()) { return true; } @@ -935,18 +1147,12 @@ fn score_support(request: Internal) -> Internal { if searches.is_empty() { return unhandled(); } - let Some((document, first_query, index_oid)) = - searches.iter().find_map(|&(document, query)| { - find_matching_tin_index(parse, (*rte).relid, ctid.varno, document) - .map(|index_oid| (document, query, index_oid)) - }) - else { - let operands = searches - .iter() - .map(|&(document, _)| document) - .collect::>(); - refuse_unindexed_scoring(rte, ctid.varno, &operands); - }; + let mut groups = bind_search_groups(parse, rte, ctid.varno, &searches); + let mut documents = PgList::::new(); + for group in &groups { + documents.push(pg_sys::copyObjectImpl(group.document.cast()).cast()); + } + let documents = make_text_array(documents); let original_nargs = pg_sys::list_length((*request.fcall).args); let function_name = pg_sys::get_func_name((*request.fcall).funcid); let fname = CStr::from_ptr(function_name).to_string_lossy(); @@ -955,7 +1161,7 @@ fn score_support(request: Internal) -> Internal { } else if fname.as_ref() == "max_score" { let mut binding = FullScoreBinding { ctid, - document, + documents, support: pg_sys::get_func_support((*request.fcall).funcid), bound: lookup_score_bound(), }; @@ -965,6 +1171,7 @@ fn score_support(request: Internal) -> Internal { (&mut binding as *mut FullScoreBinding).cast(), pg_sys::QTW_IGNORE_RC_SUBQUERIES as i32, ); + mark_required_groups(parse, rte, ctid.varno, &mut groups); if full { ScoreMode::MaxFullScore } else { @@ -974,16 +1181,29 @@ fn score_support(request: Internal) -> Internal { ScoreMode::Score }; let mut args = PgList::::new(); - args.push(pg_sys::copyObjectImpl(document.cast()).cast()); - let same_expression = binding - .matches - .iter() - .copied() - .filter(|(candidate, _)| pg_sys::equal((*candidate).cast(), document.cast())) - .collect::>(); - args.push(make_query_array(&same_expression, first_query)); + args.push(documents); + let mut queries = PgList::::new(); + let mut query_groups = Vec::new(); + for (position, group) in groups.iter().enumerate() { + for query in group_queries(&group.queries) { + queries.push(query); + query_groups.push(position as i32); + } + } + args.push(make_text_array(queries)); + args.push(make_array_const(query_groups, pg_sys::INT4ARRAYOID)); args.push(make_int4_const((*rte).relid.to_u32() as i32).cast()); - args.push(make_int4_const(index_oid.to_u32() as i32).cast()); + args.push(make_array_const( + groups + .iter() + .map(|group| group.index_oid.to_u32() as i32) + .collect(), + pg_sys::INT4ARRAYOID, + )); + args.push(make_array_const( + groups.iter().map(|group| group.required).collect(), + pg_sys::BOOLARRAYOID, + )); args.push(make_int4_const(mode as i32).cast()); let null_float = || make_null_const(pg_sys::FLOAT4OID); let null_array = || make_null_const(pg_sys::TEXTARRAYOID); @@ -1027,24 +1247,33 @@ fn score_support(request: Internal) -> Internal { } } -/// Builds the `text[]` of search expressions that the scorer combines. Passing -/// the expressions as an array keeps parameters and other non-constant nodes, -/// which cannot be combined at plan time, contributing to the scores. -unsafe fn make_query_array( - matches: &[(*mut pg_sys::Node, *mut pg_sys::Node)], - first_query: *mut pg_sys::Node, -) -> *mut pg_sys::Node { - let mut elements = PgList::::new(); - for &(_, node) in matches { - if node.is_null() || unsafe { pg_sys::exprType(node) } != pg_sys::TEXTOID { - continue; - } - elements.push(unsafe { pg_sys::copyObjectImpl(node.cast()).cast() }); - } +/// Returns copies of a group's search expressions, which the scorer combines. +/// Passing the expressions as an array keeps parameters and other +/// non-constant nodes, which cannot be combined at plan time, contributing to +/// the scores. +unsafe fn group_queries(queries: &[*mut pg_sys::Node]) -> Vec<*mut pg_sys::Node> { + let copy = |node: *mut pg_sys::Node| unsafe { pg_sys::copyObjectImpl(node.cast()).cast() }; + let mut elements = queries + .iter() + .copied() + .filter(|&node| !node.is_null() && unsafe { pg_sys::exprType(node) } == pg_sys::TEXTOID) + .map(copy) + .collect::>(); if elements.is_empty() { - elements.push(unsafe { pg_sys::copyObjectImpl(first_query.cast()).cast() }); + elements.push(copy(queries[0])); } - unsafe { make_text_array(elements) } + elements +} + +/// Builds an array constant of `values`. +unsafe fn make_array_const( + values: Vec, + array_type: pg_sys::Oid, +) -> *mut pg_sys::Node { + let datum = values + .into_datum() + .expect("an array of non-null values is not NULL"); + unsafe { pg_sys::makeConst(array_type, -1, pg_sys::InvalidOid, -1, datum, false, false).cast() } } /// Builds a `text[]` expression from `text` expressions. @@ -1089,10 +1318,12 @@ unsafe fn make_null_const(type_oid: pg_sys::Oid) -> *mut pg_sys::Const { unsafe fn lookup_score_bound() -> pg_sys::Oid { let types = [ - pg_sys::TEXTOID, pg_sys::TEXTARRAYOID, + pg_sys::TEXTARRAYOID, + pg_sys::INT4ARRAYOID, pg_sys::INT4OID, - pg_sys::INT4OID, + pg_sys::INT4ARRAYOID, + pg_sys::BOOLARRAYOID, pg_sys::INT4OID, pg_sys::FLOAT4OID, pg_sys::FLOAT4OID, @@ -1130,7 +1361,7 @@ ALTER FUNCTION @extschema@.full_score(pg_catalog.tid) SUPPORT @extschema@.score_ ALTER FUNCTION @extschema@.full_score(pg_catalog.tid, pg_catalog.float4, pg_catalog.float4) SUPPORT @extschema@.score_support; ALTER FUNCTION @extschema@.score(pg_catalog.tid, pg_catalog.float4, pg_catalog.float4, pg_catalog.float4, pg_catalog.text[], pg_catalog.text[]) SUPPORT @extschema@.score_support; ALTER FUNCTION @extschema@.max_score(pg_catalog.tid) SUPPORT @extschema@.score_support; -REVOKE ALL ON FUNCTION @extschema@.score_bound(pg_catalog.text, pg_catalog.text[], pg_catalog.int4, pg_catalog.int4, pg_catalog.int4, pg_catalog.float4, pg_catalog.float4, pg_catalog.float4, pg_catalog.text[], pg_catalog.text[]) FROM PUBLIC; +REVOKE ALL ON FUNCTION @extschema@.score_bound(pg_catalog.text[], pg_catalog.text[], pg_catalog.int4[], pg_catalog.int4, pg_catalog.int4[], pg_catalog.bool[], pg_catalog.int4, pg_catalog.float4, pg_catalog.float4, pg_catalog.float4, pg_catalog.text[], pg_catalog.text[]) FROM PUBLIC; "#, name = "score_support_bindings", requires = [ From ef53f6280cf52ea443890390231857ae50543c5e Mon Sep 17 00:00:00 2001 From: Claude Date: Tue, 6 Oct 2026 18:17:07 +0000 Subject: [PATCH 2/6] fix(score): refuse scoring where tin cannot scan an unindexed search tin cannot scan a disjunction in which a search without a tin index is an alternative to an indexed one, so it refuses to score that query and names the unindexed search. Do the same, require a group for max_score whenever the quals imply a match of it, and keep an empty sum at +0. Co-Authored-By: Claude Opus 5.5 (1M context) Claude-Session: https://claude.ai/code/session_012mKGtrg7DEhVPgqYrJtHGg --- postgres/src/score.rs | 136 ++++++++++++++++++++++++++++++++++++------ 1 file changed, 119 insertions(+), 17 deletions(-) diff --git a/postgres/src/score.rs b/postgres/src/score.rs index c8200d4..027d8c2 100644 --- a/postgres/src/score.rs +++ b/postgres/src/score.rs @@ -135,7 +135,9 @@ impl ScoreCorpus { /// Sums per-index scores the way tin does: exactly, rounding once to `f32`, /// so the total does not depend on the order of the indexes. fn sum_group_scores(scores: impl IntoIterator) -> f32 { - scores.into_iter().map(f64::from).sum::() as f32 + scores + .into_iter() + .fold(0.0_f64, |total, score| total + f64::from(score)) as f32 } type CorpusMemo = FxHashMap; @@ -972,38 +974,138 @@ unsafe fn bind_search_groups( }), } } - if groups.is_empty() { + if groups.is_empty() + || (!unindexed.is_empty() && unsafe { has_unscannable_search(parse, rte, varno) }) + { unsafe { refuse_unindexed_scoring(rte, varno, &unindexed) }; } groups.sort_by_key(|group| group.index_oid.to_u32()); groups } -/// Marks the groups with a search among the quals that restrict every output -/// row. `max_score` considers only documents that match every such group. +/// Reports whether a search without a tin index is an alternative to one +/// with a tin index, as in `title ==> 'a' OR body ==> 'b'` with only `body` +/// indexed. tin cannot scan such a disjunction, so it cannot score the rows +/// either alternative admits. A search without an index elsewhere only +/// filters the rows that the indexed searches find. +unsafe fn has_unscannable_search( + parse: *mut pg_sys::Query, + rte: *mut pg_sys::RangeTblEntry, + varno: i32, +) -> bool { + let indexed = + |node| unsafe { search_index(parse, rte, varno, node).map(|index| index.is_some()) }; + let searches_index = |node: *mut pg_sys::Node| { + let mut found = false; + visit_searches(node, &mut |search| found |= indexed(search) == Some(true)); + found + }; + let mut unscannable = false; + let clauses = unsafe { restriction_clauses(parse, varno) }; + for clause in unsafe { PgList::::from_pg(clauses) }.iter_ptr() { + visit_disjunctions(clause, &mut |arms| { + unscannable |= arms.iter().enumerate().any(|(position, &arm)| { + indexed(arm) == Some(false) + && arms.iter().enumerate().any(|(other, &alternative)| { + other != position && searches_index(alternative) + }) + }); + }); + } + unscannable +} + +/// Calls `visit` with the arms of every `OR` in `node` outside a `NOT`. +fn visit_disjunctions(node: *mut pg_sys::Node, visit: &mut impl FnMut(&[*mut pg_sys::Node])) { + if node.is_null() + || is_negation(node) + || unsafe { (*node).type_ } != pg_sys::NodeTag::T_BoolExpr + { + return; + } + let expression = unsafe { &*node.cast::() }; + let arms = unsafe { PgList::::from_pg(expression.args) } + .iter_ptr() + .collect::>(); + if expression.boolop == pg_sys::BoolExprType::OR_EXPR { + visit(&arms); + } + for arm in arms { + visit_disjunctions(arm, visit); + } +} + +/// Calls `visit` with every node in `node` outside a `NOT`. +fn visit_searches(node: *mut pg_sys::Node, visit: &mut impl FnMut(*mut pg_sys::Node)) { + if node.is_null() || is_negation(node) { + return; + } + visit(node); + if unsafe { (*node).type_ } == pg_sys::NodeTag::T_BoolExpr { + let expression = unsafe { &*node.cast::() }; + for arm in unsafe { PgList::::from_pg(expression.args) }.iter_ptr() { + visit_searches(arm, visit); + } + } +} + +/// Marks the groups that the quals restricting every output row require a +/// match of. `max_score` considers only documents that match every such +/// group. unsafe fn mark_required_groups( parse: *mut pg_sys::Query, rte: *mut pg_sys::RangeTblEntry, varno: i32, groups: &mut [SearchGroup], ) { - unsafe { - let clauses = restriction_clauses(parse, varno); - for clause in PgList::::from_pg(clauses).iter_ptr() { - let Some((document, _)) = search_operands(clause) else { - continue; - }; - let Some(index_oid) = find_matching_tin_index(parse, (*rte).relid, varno, document) - else { - continue; - }; - for group in groups.iter_mut() { - group.required |= group.index_oid == index_oid; - } + let bound = |node| unsafe { search_index(parse, rte, varno, node) }; + let clauses = unsafe { restriction_clauses(parse, varno) }; + for clause in unsafe { PgList::::from_pg(clauses) }.iter_ptr() { + for group in groups.iter_mut() { + group.required |= requires_index(clause, group.index_oid, &bound); } } } +/// Returns the tin index that a search on relation `varno` binds to, with +/// `Some(None)` for a search without one, or `None` when `node` is not a +/// search on the relation. +unsafe fn search_index( + parse: *mut pg_sys::Query, + rte: *mut pg_sys::RangeTblEntry, + varno: i32, + node: *mut pg_sys::Node, +) -> Option> { + unsafe { + let (document, _) = search_operands(node)?; + (single_varno(document) == Some(varno)) + .then(|| find_matching_tin_index(parse, (*rte).relid, varno, document)) + } +} + +/// Reports whether every row that satisfies `node` matches a search bound to +/// `index_oid`. +fn requires_index( + node: *mut pg_sys::Node, + index_oid: pg_sys::Oid, + bound: &impl Fn(*mut pg_sys::Node) -> Option>, +) -> bool { + if node.is_null() || is_negation(node) { + return false; + } + if unsafe { (*node).type_ } != pg_sys::NodeTag::T_BoolExpr { + return bound(node) == Some(Some(index_oid)); + } + let expression = unsafe { &*node.cast::() }; + let arms = unsafe { PgList::::from_pg(expression.args) }; + let mut arms = arms.iter_ptr(); + if expression.boolop == pg_sys::BoolExprType::OR_EXPR { + arms.all(|arm| requires_index(arm, index_oid, bound)) + } else { + arms.any(|arm| requires_index(arm, index_oid, bound)) + } +} + /// Renders a search expression on relation `varno` against the relation's /// own name. Returns `None` for expressions a single-relation deparse context /// cannot describe. From dfe0d1c70aaf1f7ca63ae37fedea6ea7d3da2451 Mon Sep 17 00:00:00 2001 From: Claude Date: Tue, 6 Oct 2026 18:19:10 +0000 Subject: [PATCH 3/6] test(score): cover scoring across several indexed columns Pin the scores tin returns for the order-swapped searches, the summed scores and max_score of rows matching several columns, and the shapes with a partial index that the quals do not imply. Co-Authored-By: Claude Opus 5.5 (1M context) Claude-Session: https://claude.ai/code/session_012mKGtrg7DEhVPgqYrJtHGg --- postgres/src/lib.rs | 162 ++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 162 insertions(+) diff --git a/postgres/src/lib.rs b/postgres/src/lib.rs index 27b035c..f0c1774 100644 --- a/postgres/src/lib.rs +++ b/postgres/src/lib.rs @@ -747,6 +747,168 @@ mod tests { ); } + #[pg_test] + fn scores_across_indexed_columns_do_not_depend_on_predicate_order() { + Spi::run( + "CREATE TABLE lite_fields (id int, title text, body text); + INSERT INTO lite_fields VALUES (1, 'quiet', 'ruby ruby'), (2, 'ruby', 'quiet'); + CREATE INDEX ON lite_fields USING tin (title); + CREATE INDEX ON lite_fields USING tin (body);", + ) + .unwrap(); + let full = "SELECT array_agg(tin.full_score(ctid) ORDER BY id) FROM lite_fields WHERE"; + // tin scores each row by the one column it matches, in either order. + for quals in [ + "title ==> 'ruby' OR body ==> 'ruby'", + "body ==> 'ruby' OR title ==> 'ruby'", + ] { + assert_eq!( + scores_by_id(&format!("{full} {quals}")), + [0.871_385_04, 0.693_147_2], + "{quals}" + ); + } + let title = scores_by_id(&format!("{full} title ==> 'quiet'")); + let body = scores_by_id(&format!("{full} body ==> 'ruby'")); + for quals in [ + "title ==> 'quiet' AND body ==> 'ruby'", + "body ==> 'ruby' AND title ==> 'quiet'", + ] { + assert_eq!( + scores_by_id(&format!("{full} {quals}")), + [title[0] + body[0]], + "{quals}" + ); + } + } + + #[pg_test] + fn scores_across_indexed_columns_sum_per_column_scores() { + Spi::run( + "CREATE TABLE lite_field_sums (id int, title text, body text); + INSERT INTO lite_field_sums VALUES + (1, 'alpha title', 'beta body only'), + (2, 'beta title', 'alpha body only'), + (3, 'alpha beta', 'alpha beta both'), + (4, 'gamma', 'delta'); + CREATE INDEX ON lite_field_sums USING tin (title); + CREATE INDEX ON lite_field_sums USING tin (body);", + ) + .unwrap(); + let scores = |quals: &str| { + scores_by_id(&format!( + "SELECT array_agg(tin.full_score(ctid) ORDER BY id) + || array_agg(tin.max_score(ctid) ORDER BY id) + FROM lite_field_sums WHERE {quals}" + )) + }; + let title = scores("title ==> 'alpha'"); + let body = scores("body ==> 'alpha'"); + // tin's multi-index scan returns these scores and this maximum. + let (title_only, body_only, both) = (0.654_875_3, 0.640_724_24, 1.295_599_5); + assert_eq!( + [title_only, body_only, both], + [title[0], body[0], title[1] + body[1]] + ); + assert_eq!( + scores("title ==> 'alpha' OR body ==> 'alpha'"), + [title_only, body_only, both, both, both, both] + ); + assert_eq!( + scores("body ==> 'alpha' AND title ==> 'alpha'"), + [both, both] + ); + } + + #[pg_test] + fn max_score_across_indexed_columns_counts_rows_the_quals_admit() { + Spi::run( + "CREATE TABLE lite_field_max (id int, title text, body text); + INSERT INTO lite_field_max VALUES + (1, 'ruby ruby ruby', 'nothing here'), + (2, 'ruby and many other title words', 'ruby with plenty of other body words'), + (3, 'gems', 'gems'), (4, NULL, 'ruby gems'), (5, 'rails', NULL); + CREATE INDEX ON lite_field_max USING tin (title); + CREATE INDEX ON lite_field_max USING tin (body);", + ) + .unwrap(); + let scores = |quals: &str| { + scores_by_id(&format!( + "SELECT array_agg(tin.full_score(ctid) ORDER BY id) || max(tin.max_score(ctid)) + FROM lite_field_max WHERE {quals}" + )) + }; + // The single-column match in row 1 outscores row 2, which alone + // matches both columns when the quals require both. + assert_eq!( + scores("title ==> 'ruby' OR body ==> 'ruby'"), + [1.068_417_9, 0.915_753_84, 0.802_591_5, 1.068_417_9] + ); + for quals in [ + "body ==> 'ruby' AND title ==> 'ruby'", + "title ==> 'ruby' AND (body ==> 'ruby' OR body ==> 'gems')", + ] { + assert_eq!(scores(quals), [0.915_753_84, 0.915_753_84], "{quals}"); + } + // NULL columns score nothing. + assert_eq!( + scores("title ==> 'ruby OR rails' OR body ==> 'gems'"), + [ + 1.068_417_9, + 0.467_246_86, + 0.953_077_44, + 0.802_591_5, + 1.627_717_5, + 1.627_717_5 + ] + ); + } + + #[pg_test] + fn scoring_skips_columns_whose_partial_index_the_quals_do_not_imply() { + Spi::run( + "CREATE TABLE lite_field_partial (id int, title text, body text, active boolean); + INSERT INTO lite_field_partial VALUES + (1, 'ruby', 'ruby', true), (2, 'ruby', 'quiet', false), + (3, 'quiet', 'ruby ruby', true), (4, 'other', 'words', true); + CREATE INDEX ON lite_field_partial USING tin (title) WHERE active; + CREATE INDEX ON lite_field_partial USING tin (body);", + ) + .unwrap(); + let full = + "SELECT array_agg(tin.full_score(ctid) ORDER BY id) FROM lite_field_partial WHERE"; + let body = scores_by_id(&format!("{full} body ==> 'ruby'")); + assert_eq!(body, [0.754_912_8, 0.815_467_3]); + // Without `active`, only the body search can score and the title + // search filters, as in tin. + assert_eq!( + scores_by_id(&format!("{full} title ==> 'ruby' AND body ==> 'ruby'")), + body[..1] + ); + assert_eq!( + scores_by_id(&format!( + "{full} active AND (title ==> 'ruby' OR body ==> 'ruby')" + )), + [1.735_742, 0.815_467_3] + ); + } + + #[pg_test(error = "cannot compute scores for this query")] + fn scoring_refuses_an_unindexed_alternative_to_an_indexed_search() { + Spi::run( + "CREATE TABLE lite_field_unscannable (id int, title text, body text, active boolean); + INSERT INTO lite_field_unscannable VALUES (1, 'ruby', 'ruby', true); + CREATE INDEX ON lite_field_unscannable USING tin (title) WHERE active; + CREATE INDEX ON lite_field_unscannable USING tin (body);", + ) + .unwrap(); + Spi::run( + "SELECT tin.full_score(ctid) FROM lite_field_unscannable + WHERE title ==> 'ruby' OR body ==> 'ruby'", + ) + .unwrap(); + } + #[pg_test] fn highlighting_supports_explicit_and_implicit_queries() { assert_eq!( From 912a7d90f879bd614839f72f86c4a9b1833e3a41 Mon Sep 17 00:00:00 2001 From: Claude Date: Tue, 6 Oct 2026 18:22:20 +0000 Subject: [PATCH 4/6] fix(score,highlight): recognize searches by Lead's ==> operator OID Scoring and highlighting treated any operator named ==> with two arguments as a search, so a ==> from another schema or for other types bound to them. A ==>(text, int) put an int query into the text[] of searches and crashed the backend. Compare the operator OID with Lead's pg_catalog.==>(text, text) instead, looked up on each use. Co-Authored-By: Claude Opus 5.5 (1M context) Claude-Session: https://claude.ai/code/session_012mKGtrg7DEhVPgqYrJtHGg --- postgres/src/highlight_udfs.rs | 18 +++----------- postgres/src/lib.rs | 45 ++++++++++++++++++++++++++++++++-- postgres/src/operator.rs | 41 ++++++++++++++++++++++++++++++- postgres/src/score.rs | 26 ++------------------ 4 files changed, 89 insertions(+), 41 deletions(-) diff --git a/postgres/src/highlight_udfs.rs b/postgres/src/highlight_udfs.rs index 8b43aec..fc186e3 100644 --- a/postgres/src/highlight_udfs.rs +++ b/postgres/src/highlight_udfs.rs @@ -132,20 +132,10 @@ unsafe extern "C-unwind" fn collect_queries(node: *mut pg_sys::Node, context: *m return false; } let context = unsafe { &mut *context.cast::() }; - if unsafe { (*node).type_ } == pg_sys::NodeTag::T_OpExpr { - let op = node.cast::(); - let name = unsafe { pg_sys::get_opname((*op).opno) }; - if !name.is_null() - && unsafe { CStr::from_ptr(name) }.to_bytes() == b"==>" - && unsafe { pg_sys::list_length((*op).args) } == 2 - { - let left = unsafe { pg_sys::list_nth((*op).args, 0).cast::() }; - if unsafe { pg_sys::equal(left.cast(), context.document.cast()) } { - context - .queries - .push(unsafe { pg_sys::list_nth((*op).args, 1).cast::() }); - } - } + if let Some((document, query)) = unsafe { crate::operator::search_operands(node) } + && unsafe { pg_sys::equal(document.cast(), context.document.cast()) } + { + context.queries.push(query); } unsafe { pg_sys::expression_tree_walker( diff --git a/postgres/src/lib.rs b/postgres/src/lib.rs index f0c1774..c54355c 100644 --- a/postgres/src/lib.rs +++ b/postgres/src/lib.rs @@ -757,14 +757,15 @@ mod tests { ) .unwrap(); let full = "SELECT array_agg(tin.full_score(ctid) ORDER BY id) FROM lite_fields WHERE"; - // tin scores each row by the one column it matches, in either order. + // tin scores each row by the one column it matches, in either order: + // 0.87138504 and 0.6931472, which is ln 2. for quals in [ "title ==> 'ruby' OR body ==> 'ruby'", "body ==> 'ruby' OR title ==> 'ruby'", ] { assert_eq!( scores_by_id(&format!("{full} {quals}")), - [0.871_385_04, 0.693_147_2], + [0.871_385_04, std::f32::consts::LN_2], "{quals}" ); } @@ -909,6 +910,46 @@ mod tests { .unwrap(); } + #[pg_test] + fn other_search_operators_do_not_bind_scoring_or_highlighting() { + Spi::run( + "CREATE TABLE lite_other_operator (id int, body text); + INSERT INTO lite_other_operator VALUES (1, 'beer wine'), (2, 'beer'), (3, 'cider'); + CREATE INDEX ON lite_other_operator USING tin (body); + CREATE SCHEMA lite_other; + CREATE FUNCTION lite_other.longer_than(text, int) RETURNS boolean + LANGUAGE sql IMMUTABLE AS 'SELECT length($1) > $2'; + CREATE OPERATOR lite_other.==> ( + LEFTARG = text, RIGHTARG = int, FUNCTION = lite_other.longer_than);", + ) + .unwrap(); + let rows = |select: &str, quals: &str| { + Spi::get_one::>(&format!( + "SELECT array_agg(({select})::text ORDER BY id) + FROM lite_other_operator WHERE {quals}" + )) + .unwrap() + .unwrap() + }; + let other = "body OPERATOR(lite_other.==>) 3"; + for select in [ + "tin.score(ctid)", + "tin.full_score(ctid)", + "tin.max_score(ctid)", + "tin.highlight(body)", + ] { + assert_eq!( + rows(select, &format!("body ==> 'beer' AND {other}")), + rows(select, "body ==> 'beer'"), + "{select}" + ); + } + assert_eq!( + rows("tin.highlight(body)", other), + ["beer wine", "beer", "cider"] + ); + } + #[pg_test] fn highlighting_supports_explicit_and_implicit_queries() { assert_eq!( diff --git a/postgres/src/operator.rs b/postgres/src/operator.rs index 3a61702..aca3727 100644 --- a/postgres/src/operator.rs +++ b/postgres/src/operator.rs @@ -16,7 +16,7 @@ // The full license text is available in LICENSE. #[allow(unused_imports)] use crate::am::amhandler; -use pgrx::{extension_sql, pg_extern}; +use pgrx::{extension_sql, pg_extern, pg_sys}; use tinql::runtime::{ Query, SimplificationProfile, evaluate, lower::lower, simplify, subtokenize::sub_tokenize, tokenize_doc, @@ -30,6 +30,45 @@ fn parse_search(query_text: &str, tokenizer: &T) -> Result(text, text)` operator, or `InvalidOid` if +/// it does not exist. It is looked up on each call because the extension can +/// be dropped and recreated. +fn search_operator() -> pg_sys::Oid { + unsafe { + let mut names = std::ptr::null_mut(); + for name in [c"pg_catalog", c"==>"] { + names = pg_sys::lappend( + names, + pg_sys::makeString(pg_sys::pstrdup(name.as_ptr())).cast(), + ); + } + pg_sys::OpernameGetOprid(names, pg_sys::TEXTOID, pg_sys::TEXTOID) + } +} + +/// Returns the document and query operands of `node` when it is a search +/// with Lead's `==>` operator. Operators of the same name in other schemas +/// or for other types are not searches. +/// +/// # Safety +/// `node` must be a valid expression node. +pub(crate) unsafe fn search_operands( + node: *mut pg_sys::Node, +) -> Option<(*mut pg_sys::Node, *mut pg_sys::Node)> { + unsafe { + if (*node).type_ != pg_sys::NodeTag::T_OpExpr { + return None; + } + let op = &*node.cast::(); + if op.opno != search_operator() || pg_sys::list_length(op.args) != 2 { + return None; + } + let left = pg_sys::list_nth(op.args, 0).cast::(); + let right = pg_sys::list_nth(op.args, 1).cast::(); + (!left.is_null()).then_some((left, right)) + } +} + fn invalid_search(error: String) -> ! { pgrx::error!("invalid ==> query: {error}") } diff --git a/postgres/src/score.rs b/postgres/src/score.rs index 027d8c2..27eb941 100644 --- a/postgres/src/score.rs +++ b/postgres/src/score.rs @@ -667,34 +667,12 @@ unsafe extern "C-unwind" fn find_qual(node: *mut pg_sys::Node, context: *mut c_v return false; } let binding = unsafe { &mut *context.cast::() }; - if let Some(search) = unsafe { search_operands(node) } { + if let Some(search) = unsafe { crate::operator::search_operands(node) } { binding.matches.push(search); } unsafe { pg_sys::expression_tree_walker(node, Some(find_qual), context) } } -/// Returns the document and query operands of a `==>` search. -unsafe fn search_operands( - node: *mut pg_sys::Node, -) -> Option<(*mut pg_sys::Node, *mut pg_sys::Node)> { - unsafe { - if (*node).type_ != pg_sys::NodeTag::T_OpExpr { - return None; - } - let op = node.cast::(); - let name = pg_sys::get_opname((*op).opno); - if name.is_null() - || CStr::from_ptr(name).to_bytes() != b"==>" - || pg_sys::list_length((*op).args) != 2 - { - return None; - } - let left = pg_sys::list_nth((*op).args, 0).cast::(); - let right = pg_sys::list_nth((*op).args, 1).cast::(); - (!left.is_null()).then_some((left, right)) - } -} - struct VarnoBinding { varno: i32, seen: bool, @@ -1077,7 +1055,7 @@ unsafe fn search_index( node: *mut pg_sys::Node, ) -> Option> { unsafe { - let (document, _) = search_operands(node)?; + let (document, _) = crate::operator::search_operands(node)?; (single_varno(document) == Some(varno)) .then(|| find_matching_tin_index(parse, (*rte).relid, varno, document)) } From 17a014da213e5c63c60f48ba85c11fdbf46d810a Mon Sep 17 00:00:00 2001 From: Claude Date: Tue, 6 Oct 2026 18:40:29 +0000 Subject: [PATCH 5/6] fix(score): leave rows that no search admits unscored A row that only a non-search qual admits, such as the id = 8 of body ==> 'gems' OR id = 8, scored 0. tin returns NULL for it, so return NULL when none of the row's indexed columns matches its searches. Co-Authored-By: Claude Opus 5.5 (1M context) Claude-Session: https://claude.ai/code/session_012mKGtrg7DEhVPgqYrJtHGg --- postgres/src/lib.rs | 38 ++++++++++++++++++++++++++++++++++++++ postgres/src/score.rs | 15 ++++++--------- 2 files changed, 44 insertions(+), 9 deletions(-) diff --git a/postgres/src/lib.rs b/postgres/src/lib.rs index c54355c..025b670 100644 --- a/postgres/src/lib.rs +++ b/postgres/src/lib.rs @@ -910,6 +910,44 @@ mod tests { .unwrap(); } + #[pg_test] + fn rows_that_no_search_admits_have_no_score() { + Spi::run( + "CREATE TABLE lite_unsearched (id int, body text, active boolean); + INSERT INTO lite_unsearched VALUES + (1, 'gems', false), (2, 'gems gems', false), + (3, 'words', true), (4, 'other words', false); + CREATE INDEX ON lite_unsearched USING tin (body);", + ) + .unwrap(); + let scores = |select: &str, quals: &str| { + Spi::get_one::>>(&format!( + "SELECT {select} FROM lite_unsearched WHERE {quals}" + )) + .unwrap() + .unwrap() + }; + // The scores tin returns: the row only `id = 4` or `active` admits + // has none, and max_score still reports the searched maximum. + for quals in ["body ==> 'gems' OR id = 4", "body ==> 'gems' OR active"] { + let max = Some(0.871_385_04); + assert_eq!( + scores( + "array_agg(tin.full_score(ctid) ORDER BY id) + || array_agg(tin.max_score(ctid) ORDER BY id)", + quals + ), + [Some(0.802_591_5), max, None, max, max, max], + "{quals}" + ); + assert_eq!( + scores("array_agg(tin.score(ctid) ORDER BY id)", quals), + [Some(0.0), Some(0.0), None], + "{quals}" + ); + } + } + #[pg_test] fn other_search_operators_do_not_bind_scoring_or_highlighting() { Spi::run( diff --git a/postgres/src/score.rs b/postgres/src/score.rs index 27eb941..acb35af 100644 --- a/postgres/src/score.rs +++ b/postgres/src/score.rs @@ -204,7 +204,7 @@ fn score_bound( term_add: Option>, term_replace: Option>, fcinfo: pg_sys::FunctionCallInfo, -) -> f32 { +) -> Option { let mode = ScoreMode::try_from(mode) .unwrap_or_else(|mode| pgrx::error!("tin.score_bound(): unknown score mode {mode}")); let mut groups = index_oids @@ -242,15 +242,12 @@ fn score_bound( .entry(key) .or_insert_with_key(|key| build_corpus(key, k1, b, term_add, term_replace)); if mode.is_max() { - corpus.max - } else { - sum_group_scores( - corpus - .group_scores(&documents) - .into_iter() - .map(|(_, score)| score), - ) + return Some(corpus.max); } + // Like tin, a row that no search admits, only another qual such as the + // `id = 8` of `body ==> 'beer' OR id = 8`, has no score. + let scores = corpus.group_scores(&documents); + (!scores.is_empty()).then(|| sum_group_scores(scores.into_iter().map(|(_, score)| score))) } /// Returns the corpora memoized on this call site's `FmgrInfo`. The executor From 0a8122becc3d5e0247018be13b0e9308dda2b770 Mon Sep 17 00:00:00 2001 From: Claude Date: Tue, 6 Oct 2026 18:53:35 +0000 Subject: [PATCH 6/6] fix(score): pick max_score's policy from the relation's score calls tin scores a relation with one policy. max_score takes the policy of the relation's tin.score calls, sharing their dense_ratio, term_add, and term_replace, and otherwise reports the full_score maximum, including when it is the only scoring call. Under the tin.score policy tin computes no maximum for several indexes or a disjunction with other quals, and returns NULL. It refuses tin.score beside tin.full_score on one relation, and tin.score calls whose scan arguments differ. Do the same. Co-Authored-By: Claude Opus 5.5 (1M context) Claude-Session: https://claude.ai/code/session_012mKGtrg7DEhVPgqYrJtHGg --- postgres/src/lib.rs | 72 ++++++++++++++ postgres/src/score.rs | 211 ++++++++++++++++++++++++++++++++---------- 2 files changed, 234 insertions(+), 49 deletions(-) diff --git a/postgres/src/lib.rs b/postgres/src/lib.rs index 025b670..6cb2a63 100644 --- a/postgres/src/lib.rs +++ b/postgres/src/lib.rs @@ -948,6 +948,78 @@ mod tests { } } + #[pg_test] + fn max_score_adapts_to_the_relations_score_calls() { + Spi::run( + "CREATE TABLE lite_max_policy (id int, title text, body text); + INSERT INTO lite_max_policy SELECT g, 'filler ' || g, 'padding ' || g + FROM generate_series(1, 40) AS g; + UPDATE lite_max_policy SET body = 'ruby ruby ruby' WHERE id = 1; + UPDATE lite_max_policy SET body = 'ruby and other body words' WHERE id = 2; + UPDATE lite_max_policy SET body = 'ruby gems', title = 'gems' WHERE id = 3; + UPDATE lite_max_policy SET body = 'gems' WHERE id = 4; + CREATE INDEX ON lite_max_policy USING tin (title); + CREATE INDEX ON lite_max_policy USING tin (body);", + ) + .unwrap(); + let max = |select: &str, quals: &str| { + Spi::get_one::(&format!( + "SELECT max(tin.max_score(ctid)){select} FROM lite_max_policy WHERE {quals}" + )) + .unwrap() + }; + let full = max(", max(tin.full_score(ctid))", "body ==> 'ruby'"); + assert!(full.is_some_and(|full| full > 0.0)); + // Alone, max_score reports the full_score maximum, as tin does. + assert_eq!(max("", "body ==> 'ruby'"), full); + // Beside tin.score it takes that call's policy, which elides ruby, + // in 3 of 40 documents, at a dense_ratio of 0.05. + assert_eq!(max(", max(tin.score(ctid))", "body ==> 'ruby'"), full); + assert_eq!( + max(", max(tin.score(ctid, 0.05))", "body ==> 'ruby'"), + Some(0.0) + ); + // tin computes no maximum under the score policy for several scans. + for quals in [ + "title ==> 'gems' OR body ==> 'ruby'", + "body ==> 'ruby' OR id = 5", + ] { + assert_eq!(max(", max(tin.score(ctid))", quals), None, "{quals}"); + assert!(max("", quals).is_some_and(|max| max > 0.0), "{quals}"); + } + assert!( + max( + ", max(tin.score(ctid))", + "body ==> 'ruby' OR body ==> 'gems'" + ) + .is_some_and(|max| max > 0.0) + ); + } + + #[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)" + )] + fn score_and_full_score_cannot_score_one_relation() { + Spi::run( + "CREATE TABLE lite_mixed_family (body text); + CREATE INDEX ON lite_mixed_family USING tin (body); + SELECT tin.score(ctid), tin.full_score(ctid) FROM lite_mixed_family + WHERE body ==> 'ruby';", + ) + .unwrap(); + } + + #[pg_test(error = "tin.score() calls on one relation must use identical dense_ratio arguments")] + fn score_calls_on_one_relation_share_their_scan_arguments() { + Spi::run( + "CREATE TABLE lite_mixed_ratio (body text); + CREATE INDEX ON lite_mixed_ratio USING tin (body); + SELECT tin.score(ctid), tin.score(ctid, 0.5) FROM lite_mixed_ratio + WHERE body ==> 'ruby';", + ) + .unwrap(); + } + #[pg_test] fn other_search_operators_do_not_bind_scoring_or_highlighting() { Spi::run( diff --git a/postgres/src/score.rs b/postgres/src/score.rs index acb35af..1dd9ecc 100644 --- a/postgres/src/score.rs +++ b/postgres/src/score.rs @@ -37,9 +37,11 @@ enum ScoreMode { Score = 0, /// `tin.full_score`: every query term scores. FullScore = 1, - /// `tin.max_score` under the `tin.score` policy. + /// `tin.max_score` on a relation that `tin.score` also scores, under that + /// call's policy. MaxScore = 2, - /// `tin.max_score` in a query that also calls `tin.full_score`. + /// `tin.max_score` on any other relation, under the `tin.full_score` + /// policy. MaxFullScore = 3, } @@ -1128,55 +1130,134 @@ unsafe fn deparse_search_expression( } } -struct FullScoreBinding { +/// The `dense_ratio`, `term_add`, and `term_replace` arguments of a +/// `tin.score` call, which a scan shares between its `tin.score` calls and +/// the `tin.max_score` that adapts to them. +const SCAN_ARGUMENTS: [(i32, &str); 3] = [(1, "dense_ratio"), (4, "term_add"), (5, "term_replace")]; + +/// The scoring calls on one relation, found in the query either as written +/// or already rewritten into `score_bound` by earlier clauses. +struct ScoringCalls { ctid: *const pg_sys::Var, documents: *mut pg_sys::Node, support: pg_sys::Oid, bound: pg_sys::Oid, + full: bool, + /// The scan arguments of each `tin.score` call, simplified. + score: Vec<[*mut pg_sys::Node; 3]>, } #[pg_guard] -unsafe extern "C-unwind" fn has_full_score(node: *mut pg_sys::Node, context: *mut c_void) -> bool { +unsafe extern "C-unwind" fn collect_scoring_calls( + node: *mut pg_sys::Node, + context: *mut c_void, +) -> bool { unsafe { if node.is_null() || (*node).type_ == pg_sys::NodeTag::T_Query { return false; } - let binding = &*context.cast::(); + let calls = &mut *context.cast::(); if (*node).type_ == pg_sys::NodeTag::T_FuncExpr { let function = &*node.cast::(); - // Earlier query clauses may already contain the rewritten scorer. - if function.funcid == binding.bound { + if function.funcid == calls.bound { let mode = pg_sys::list_nth(function.args, 6).cast::(); if (*mode).xpr.type_ == pg_sys::NodeTag::T_Const - && (*mode).constvalue.value() == ScoreMode::FullScore as usize - && pg_sys::equal(pg_sys::list_nth(function.args, 0), binding.documents.cast()) + && pg_sys::equal(pg_sys::list_nth(function.args, 0), calls.documents.cast()) { - return true; - } - } else if pg_sys::get_func_support(function.funcid) == binding.support - && CStr::from_ptr(pg_sys::get_func_name(function.funcid)).to_bytes() - == b"full_score" - { - for position in 0..pg_sys::list_length(function.args) { - let mut argument = - pg_sys::list_nth(function.args, position).cast::(); - if (*argument).type_ == pg_sys::NodeTag::T_NamedArgExpr { - let named = &*argument.cast::(); - if named.argnumber != 0 { - continue; + let argument = |position| pg_sys::list_nth(function.args, position).cast(); + match ScoreMode::try_from((*mode).constvalue.value() as i32) { + Ok(ScoreMode::FullScore) => calls.full = true, + Ok(ScoreMode::Score) => { + calls.score.push([argument(7), argument(10), argument(11)]) } - argument = named.arg.cast(); - } else if position != 0 { - continue; + _ => {} } - if pg_sys::equal(argument.cast(), binding.ctid.cast()) { - return true; + } + } else if pg_sys::get_func_support(function.funcid) == calls.support { + let name = CStr::from_ptr(pg_sys::get_func_name(function.funcid)).to_bytes(); + if name == b"score" || name == b"full_score" { + let args = expanded_arguments(function); + if pg_sys::equal(pg_sys::list_nth(args, 0), calls.ctid.cast()) { + if name == b"full_score" { + calls.full = true; + } else { + calls.score.push(SCAN_ARGUMENTS.map(|(position, _)| { + pg_sys::eval_const_expressions( + std::ptr::null_mut(), + pg_sys::list_nth(args, position).cast(), + ) + })); + } } } } } - pg_sys::expression_tree_walker(node, Some(has_full_score), context) + pg_sys::expression_tree_walker(node, Some(collect_scoring_calls), context) + } +} + +/// Returns a call's arguments in positional order with defaults filled in, +/// the form the planner gives the support function. +unsafe fn expanded_arguments(function: &pg_sys::FuncExpr) -> *mut pg_sys::List { + unsafe { + let tuple = pg_sys::SearchSysCache1( + pg_sys::SysCacheIdentifier::PROCOID as i32, + pg_sys::Datum::from(function.funcid), + ); + if tuple.is_null() { + pgrx::error!( + "cache lookup failed for function {}", + function.funcid.to_u32() + ); + } + let args = pg_sys::expand_function_arguments( + pg_sys::copyObjectImpl(function.args.cast()).cast(), + false, + function.funcresulttype, + tuple, + ); + pg_sys::ReleaseSysCache(tuple); + args + } +} + +/// 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::::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::() }; + unsafe { PgList::::from_pg(expression.args) } + .iter_ptr() + .all(|arm| only_searches(arm, is_search)) } #[pg_extern(immutable, parallel_unsafe)] @@ -1233,30 +1314,55 @@ fn score_support(request: Internal) -> Internal { let original_nargs = pg_sys::list_length((*request.fcall).args); let function_name = pg_sys::get_func_name((*request.fcall).funcid); let fname = CStr::from_ptr(function_name).to_string_lossy(); - let mode = if fname.as_ref() == "full_score" { - ScoreMode::FullScore - } else if fname.as_ref() == "max_score" { - let mut binding = FullScoreBinding { - ctid, - documents, - support: pg_sys::get_func_support((*request.fcall).funcid), - bound: lookup_score_bound(), - }; - let full = pg_sys::query_tree_walker( - parse, - Some(has_full_score), - (&mut binding as *mut FullScoreBinding).cast(), - pg_sys::QTW_IGNORE_RC_SUBQUERIES as i32, + let mut calls = ScoringCalls { + ctid, + documents, + support: pg_sys::get_func_support((*request.fcall).funcid), + bound: lookup_score_bound(), + full: false, + score: Vec::new(), + }; + pg_sys::query_tree_walker( + parse, + Some(collect_scoring_calls), + (&mut calls as *mut ScoringCalls).cast(), + pg_sys::QTW_IGNORE_RC_SUBQUERIES as i32, + ); + if calls.full && !calls.score.is_empty() { + pgrx::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)" ); - mark_required_groups(parse, rte, ctid.varno, &mut groups); - if full { - ScoreMode::MaxFullScore - } else { - ScoreMode::MaxScore + } + for (index, (_, name)) in SCAN_ARGUMENTS.iter().enumerate() { + if calls + .score + .iter() + .any(|call| !pg_sys::equal(call[index].cast(), calls.score[0][index].cast())) + { + pgrx::error!( + "tin.score() calls on one relation must use identical {name} arguments" + ); } - } else { - ScoreMode::Score + } + // tin.max_score adapts to the relation's tin.score calls, and + // otherwise reports the full_score maximum. + let mode = match fname.as_ref() { + "full_score" => ScoreMode::FullScore, + "max_score" if calls.score.is_empty() => ScoreMode::MaxFullScore, + "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); + } let mut args = PgList::::new(); args.push(documents); let mut queries = PgList::::new(); @@ -1293,6 +1399,13 @@ fn score_support(request: Internal) -> Internal { .cast(), ); } + } else if mode == ScoreMode::MaxScore { + let [dense_ratio, term_add, term_replace] = calls.score[0]; + args.push(pg_sys::copyObjectImpl(dense_ratio.cast()).cast()); + args.push(null_float().cast()); + args.push(null_float().cast()); + args.push(pg_sys::copyObjectImpl(term_add.cast()).cast()); + args.push(pg_sys::copyObjectImpl(term_replace.cast()).cast()); } else { args.push(null_float().cast()); if mode == ScoreMode::FullScore && original_nargs == 3 {