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 88cb361..6cb2a63 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(); @@ -747,6 +747,319 @@ 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: + // 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, std::f32::consts::LN_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 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 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( + "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 9c14604..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, } @@ -67,14 +69,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 +94,54 @@ 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() + .fold(0.0_f64, |total, score| total + f64::from(score)) as f32 +} + type CorpusMemo = FxHashMap; fn score_context_error(function: &str) -> ! { @@ -149,10 +193,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, @@ -160,17 +206,32 @@ 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}")); - // 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), @@ -183,10 +244,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 { - corpus.score(document) + 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 @@ -212,17 +275,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 +361,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 +378,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 +400,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 +418,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 +429,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 +503,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 +630,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,19 +666,8 @@ 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 { - 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::() }; - let right = unsafe { pg_sys::list_nth((*op).args, 1).cast::() }; - if !left.is_null() { - binding.matches.push((left, right)); - } - } + 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) } } @@ -792,6 +902,187 @@ 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() + || (!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 +} + +/// 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], +) { + 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, _) = crate::operator::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. @@ -839,57 +1130,136 @@ 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, - document: *mut pg_sys::Node, + 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 { - let mode = pg_sys::list_nth(function.args, 4).cast::(); + 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.document.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)] fn score_support(request: Internal) -> Internal { let unhandled = || Internal::from(Some(pg_sys::Datum::from(0_usize))); @@ -935,55 +1305,88 @@ 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(); - let mode = if fname.as_ref() == "full_score" { - ScoreMode::FullScore - } else if fname.as_ref() == "max_score" { - let mut binding = FullScoreBinding { - ctid, - document, - 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)" ); - 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(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); @@ -996,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 { @@ -1027,24 +1437,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 +1508,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 +1551,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 = [