From 623174db24ed2118b0431ba9617bba1c99f7caff Mon Sep 17 00:00:00 2001 From: Patrick Reynolds Date: Thu, 1 Oct 2026 17:39:14 -0400 Subject: [PATCH 1/2] Implement stemmers --- Cargo.lock | 19 +- Cargo.toml | 2 +- README.md | 4 +- postgres/src/analysis.rs | 569 +++++++++++++++++++++++++ postgres/src/highlight.rs | 18 +- postgres/src/highlight_udfs.rs | 510 ++++++++++++++++++---- postgres/src/lib.rs | 126 ++++++ postgres/src/operator.rs | 26 +- postgres/src/options.rs | 49 ++- postgres/src/score.rs | 54 ++- postgres/src/tinql.rs | 5 - postgres/src/udfs.rs | 92 +++- private-regress.manifest | 7 + tinql/src/runtime/subtokenize.rs | 225 ++++++++-- tokenizer/Cargo.toml | 1 + tokenizer/src/compiled.rs | 73 ++-- tokenizer/src/lib.rs | 28 +- tokenizer/src/long_tokens.rs | 12 +- tokenizer/src/normalizer.rs | 80 ++++ tokenizer/src/source_spans.rs | 38 +- tokenizer/src/spec.rs | 12 + tokenizer/src/stemmer.rs | 124 ++++++ tokenizer/src/tokenizers/unicode.rs | 15 +- tokenizer/src/tokenizers/whitespace.rs | 8 +- tokenizer/tests/pipeline_properties.rs | 475 +++++++++++++++++++++ tokenizer/tests/stemming.rs | 286 +++++++++++++ 26 files changed, 2641 insertions(+), 217 deletions(-) create mode 100644 postgres/src/analysis.rs create mode 100644 tokenizer/src/normalizer.rs create mode 100644 tokenizer/src/stemmer.rs create mode 100644 tokenizer/tests/pipeline_properties.rs create mode 100644 tokenizer/tests/stemming.rs diff --git a/Cargo.lock b/Cargo.lock index 7a4cbfd..7a07a50 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -108,7 +108,7 @@ dependencies = [ [[package]] name = "boldi-vigna" -version = "1.0.3" +version = "1.0.4" dependencies = [ "proptest", "rustc-hash", @@ -1473,6 +1473,16 @@ version = "0.8.11" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d6f6ff9a378485b298a5286656da665ba74413d36db0979633275d2e708145d4" +[[package]] +name = "rust-stemmers" +version = "1.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e46a2036019fdb888131db7a4c847a1063a7493f971ed94ea82c67eada63ca54" +dependencies = [ + "serde", + "serde_derive", +] + [[package]] name = "rustc-hash" version = "2.1.3" @@ -1775,7 +1785,7 @@ dependencies = [ [[package]] name = "tin" -version = "1.0.3" +version = "1.0.4" dependencies = [ "boldi-vigna", "pgrx", @@ -1789,7 +1799,7 @@ dependencies = [ [[package]] name = "tinql" -version = "1.0.3" +version = "1.0.4" dependencies = [ "boldi-vigna", "pest", @@ -1828,10 +1838,11 @@ checksum = "1f3ccbac311fea05f86f61904b462b55fb3df8837a366dfc601a0161d0532f20" [[package]] name = "tokenizer" -version = "1.0.3" +version = "1.0.4" dependencies = [ "emojis", "proptest", + "rust-stemmers", "thiserror", "tinyvec", "unicode-normalization", diff --git a/Cargo.toml b/Cargo.toml index e166e09..cef84e9 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -3,7 +3,7 @@ resolver = "3" members = ["boldi-vigna", "postgres", "tinql", "tokenizer"] [workspace.package] -version = "1.0.3" +version = "1.0.4" edition = "2024" authors = ["PlanetScale"] license = "AGPL-3.0-or-later" diff --git a/README.md b/README.md index 0699d70..8e22b8c 100644 --- a/README.md +++ b/README.md @@ -24,7 +24,9 @@ Then run `CREATE EXTENSION tin` in the database. Lead loads on demand and does n ## Compatibility boundary -Lead provides the `tin` access method, the `==>` operator, TINQL parsing, tokenizer and index reloptions, and the scoring functions `tin.score`, `tin.full_score`, `tin.max_score`, and `tin.score_inspect`, plus explicit and implicitly bound `tin.highlight` and `tin.highlight_ansi`. Postgres 17 and 18 are build targets. Search results are exact because the access method returns whole-page candidates and Postgres evaluates `==>` against each visible heap tuple, including expression and partial-index rechecks. +Lead provides the `tin` access method, the `==>` operator, TINQL parsing, tokenizer and index reloptions (including the Snowball `stemmer` option), and the scoring functions `tin.score`, `tin.full_score`, `tin.max_score`, and `tin.score_inspect`, plus explicit and implicitly bound `tin.highlight` and `tin.highlight_ansi`. Postgres 17 and 18 are build targets. Search results are exact because the access method returns whole-page candidates and Postgres evaluates `==>` against each visible heap tuple, including expression and partial-index rechecks. + +The planner binds each `==>` and each `tin.highlight` or `tin.highlight_ansi` call to the analysis of the tin index that covers its column, so a stemmed index matches and highlights inflected words in queries over that column, including joins, partitions, and cached plans. `body ==> ANY(...)` is the exception: it uses the default analysis. `tin.tokenize`, `tin.ql_parse`, `tin.highlight`, and `tin.highlight_ansi` take a trailing `stemmer` argument as in TIN 1.0.4; their 1.0.3 signatures remain as `tin.tokenize_v1_0_3`, `tin.ql_parse_v1_0_3`, `tin.highlight_v1_0_3`, and `tin.highlight_ansi_v1_0_3`. As in TIN, changing `stemmer` on an existing index requires `REINDEX`. Scoring deliberately rescans and retokenizes the visible indexed column or expression, once per statement and search, under that statement's snapshot. A score call must be in the same query level as the matching `==>` predicate, and a partial tin index binds only when the query's quals imply its `WHERE` condition, as in tin. Implicit highlighting binds the same way and, like tin, returns the text unmarked when no index binds a search; passing its `query` argument explicitly works without a bound predicate. diff --git a/postgres/src/analysis.rs b/postgres/src/analysis.rs new file mode 100644 index 0000000..2f41cfb --- /dev/null +++ b/postgres/src/analysis.rs @@ -0,0 +1,569 @@ +// Copyright (C) 2026 PlanetScale +// +// This program is free software: you can redistribute it and/or modify +// it under the terms of the GNU Affero General Public License as published by +// the Free Software Foundation, either version 3 of the License, or +// (at your option) any later version. +// +// This program is distributed in the hope that it will be useful, +// but WITHOUT ANY WARRANTY; without even the implied warranty of +// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +// GNU Affero General Public License for more details. +// +// You should have received a copy of the GNU Affero General Public License +// along with this program. If not, see . +// +// The full license text is available in LICENSE. +use crate::udfs::TokenizeOptions; +use pgrx::{Internal, IntoDatum, PgList, Spi, pg_extern, pg_sys}; +use std::cell::RefCell; +use std::ffi::{CStr, CString}; +use std::rc::Rc; +use tokenizer::{ + CompiledTokenizerPipeline, Folding, GraphemeMode, LongTokenMode, PositionGapMode, + TokenizerPipelineSpec, TokenizerSpec, +}; + +/// Marks a `==>` search expression that carries the analysis configuration of +/// the index it was planned against. Snowball stemming is not idempotent, so +/// the expression keeps the user's raw text and the configuration travels next +/// to it instead of being applied to the query early. +const TAG: &str = "\u{1}tin-analysis\u{1}"; +const SEPARATOR: char = '\u{1}'; + +pub(crate) fn encode(spec: &TokenizerPipelineSpec) -> String { + let folding = |value| match value { + Folding::Fold => "fold", + Folding::Preserve => "preserve", + }; + [ + match spec.tokenizer { + TokenizerSpec::Unicode => "unicode", + TokenizerSpec::Whitespace => "whitespace", + } + .to_owned(), + folding(spec.case_folding).to_owned(), + folding(spec.accent_folding).to_owned(), + match spec.long_tokens.mode { + LongTokenMode::Truncate => "truncate", + LongTokenMode::Discard => "discard", + LongTokenMode::Split => "split", + } + .to_owned(), + spec.long_tokens.max_bytes.to_string(), + match spec.graphemes { + GraphemeMode::Discard => "discard", + GraphemeMode::Emoji => "emoji", + GraphemeMode::Retain => "retain", + } + .to_owned(), + match spec.position_gaps { + PositionGapMode::Collapse => "collapse", + PositionGapMode::Preserve => "preserve", + } + .to_owned(), + spec.stemmer + .map(|stemmer| stemmer.as_str().to_owned()) + .unwrap_or_default(), + ] + .join(",") +} + +pub(crate) fn decode(config: &str) -> Result { + let fields: Vec<&str> = config.split(',').collect(); + let [ + tokenizer, + case_folding, + accent_folding, + long_tokens, + max_token_bytes, + graphemes, + position_gaps, + stemmer, + ] = fields[..] + else { + return Err(format!("invalid analysis configuration {config:?}")); + }; + TokenizeOptions { + tokenizer, + case_folding, + accent_folding, + long_tokens, + max_token_bytes: max_token_bytes + .parse() + .map_err(|_| format!("invalid analysis configuration {config:?}"))?, + graphemes, + position_gaps, + stemmer: (!stemmer.is_empty()).then_some(stemmer), + } + .into_spec() +} + +pub(crate) fn tag(spec: &TokenizerPipelineSpec, raw: &str) -> String { + format!("{TAG}{}{SEPARATOR}{raw}", encode(spec)) +} + +/// Splits a possibly tagged `==>` search expression into the analysis +/// configuration it was bound to and the user's raw query text. +pub(crate) fn split_tag(text: &str) -> (Option<&str>, &str) { + match text + .strip_prefix(TAG) + .and_then(|rest| rest.split_once(SEPARATOR)) + { + Some((config, raw)) => (Some(config), raw), + None => (None, text), + } +} + +pub(crate) fn raw_query(text: &str) -> &str { + split_tag(text).1 +} + +thread_local! { + static PIPELINE: RefCell)>> = + const { RefCell::new(None) }; +} + +/// Compiles the pipeline named by a tag, reusing the previous one: a scan +/// evaluates the same configuration for every row. +pub(crate) fn pipeline_for(config: &str) -> Result, String> { + PIPELINE.with_borrow_mut(|slot| { + if let Some((cached, pipeline)) = slot + && cached == config + { + return Ok(Rc::clone(pipeline)); + } + let pipeline = Rc::new( + decode(config)? + .compile() + .map_err(|error| error.to_string())?, + ); + *slot = Some((config.to_owned(), Rc::clone(&pipeline))); + Ok(pipeline) + }) +} + +#[pg_extern(immutable, parallel_safe)] +fn bind_query_analysis(query: &str, analysis: &str) -> String { + format!("{TAG}{analysis}{SEPARATOR}{}", raw_query(query)) +} + +pub(crate) unsafe fn sibling_function( + funcid: pg_sys::Oid, + name: &str, + types: &[pg_sys::Oid], +) -> pg_sys::Oid { + let schema = unsafe { pg_sys::get_namespace_name(pg_sys::get_func_namespace(funcid)) }; + let name = CString::new(name).expect("function names contain no NUL"); + let qualified = unsafe { pg_sys::quote_qualified_identifier(schema, name.as_ptr()) }; + let names = unsafe { pg_sys::stringToQualifiedNameList(qualified, std::ptr::null_mut()) }; + unsafe { pg_sys::LookupFuncName(names, types.len() as i32, types.as_ptr(), false) } +} + +pub(crate) unsafe fn text_const(value: &str) -> *mut pg_sys::Const { + let datum = value.into_datum().expect("&str is never NULL"); + unsafe { + pg_sys::makeConst( + pg_sys::TEXTOID, + -1, + pg_sys::DEFAULT_COLLATION_OID, + -1, + datum, + false, + false, + ) + } +} + +unsafe fn text_of(node: *mut pg_sys::Node) -> Option { + if node.is_null() || unsafe { (*node).type_ } != pg_sys::NodeTag::T_Const { + return None; + } + let value = unsafe { &*node.cast::() }; + if value.constisnull || value.consttype != pg_sys::TEXTOID { + return None; + } + unsafe { ::from_datum(value.constvalue, false) } +} + +unsafe fn is_bind_call(node: *mut pg_sys::Node) -> bool { + if node.is_null() || unsafe { (*node).type_ } != pg_sys::NodeTag::T_FuncExpr { + return false; + } + let function = unsafe { (*node.cast::()).funcid }; + let name = unsafe { pg_sys::get_func_name(function) }; + !name.is_null() && unsafe { CStr::from_ptr(name) }.to_bytes() == b"bind_query_analysis" +} + +/// The user's search expression with any analysis binding removed. +pub(crate) unsafe fn unbound_query(node: *mut pg_sys::Node) -> *mut pg_sys::Node { + if unsafe { is_bind_call(node) } { + return unsafe { pg_sys::list_nth((*node.cast::()).args, 0).cast() }; + } + match unsafe { text_of(node) } { + Some(text) if split_tag(&text).0.is_some() => { + unsafe { text_const(raw_query(&text)) }.cast() + } + _ => node, + } +} + +unsafe fn is_bound_query(node: *mut pg_sys::Node) -> bool { + (unsafe { is_bind_call(node) }) + || unsafe { text_of(node) }.is_some_and(|text| split_tag(&text).0.is_some()) +} + +#[derive(Default)] +struct Leaves { + specs: Vec, + indexes: Vec, +} + +pub(crate) enum Resolution { + /// The operand is not covered by any tin index. + Unindexed, + Bound { + spec: TokenizerPipelineSpec, + indexes: Vec, + }, + /// Members of a partitioned, inherited, or UNION ALL relation analyze + /// the operand differently. + Conflict, +} + +/// Finds the analysis configuration of the tin index that covers `operand`, +/// looking through partitions, inheritance children, and UNION ALL arms. +pub(crate) unsafe fn resolve( + root: *mut pg_sys::PlannerInfo, + operand: *mut pg_sys::Node, +) -> Resolution { + let mut leaves = Leaves::default(); + unsafe { collect(root, operand, &mut leaves) }; + let Some(&first) = leaves.specs.first() else { + return Resolution::Unindexed; + }; + if leaves.specs.iter().any(|spec| *spec != first) { + return Resolution::Conflict; + } + if leaves.indexes.is_empty() { + return Resolution::Unindexed; + } + Resolution::Bound { + spec: first, + indexes: leaves.indexes, + } +} + +unsafe fn collect(root: *mut pg_sys::PlannerInfo, operand: *mut pg_sys::Node, leaves: &mut Leaves) { + let plain = |leaves: &mut Leaves| leaves.specs.push(TokenizerPipelineSpec::tin_default()); + let Some(varno) = (unsafe { crate::score::single_varno(operand) }) else { + return plain(leaves); + }; + let parse = unsafe { (*root).parse }; + if varno < 1 || varno > unsafe { pg_sys::list_length((*parse).rtable) } { + return plain(leaves); + } + let rte = + unsafe { pg_sys::list_nth((*parse).rtable, varno - 1) }.cast::(); + match unsafe { (*rte).rtekind } { + pg_sys::RTEKind::RTE_RELATION => unsafe { + let relid = (*rte).relid; + let members = if (*rte).inh { + descendants(relid) + } else { + vec![relid] + }; + for member in members { + match leaf_index((*root).parse, member, relid, varno, operand) { + Some((index, spec)) => { + leaves.indexes.push(index); + leaves.specs.push(spec); + } + None => plain(leaves), + } + } + }, + pg_sys::RTEKind::RTE_SUBQUERY if unsafe { (*rte).inh } => unsafe { + let arms = PgList::::from_pg((*root).append_rel_list); + let mut found = false; + for arm in arms.iter_ptr() { + if (*arm).parent_relid != varno as pg_sys::Index { + continue; + } + found = true; + let mut appinfo = arm; + let translated = pg_sys::adjust_appendrel_attrs( + root, + pg_sys::copyObjectImpl(operand.cast()).cast(), + 1, + &raw mut appinfo, + ); + collect(root, translated, leaves); + } + if !found { + plain(leaves); + } + }, + _ => plain(leaves), + } +} + +/// The relation and, for inheritance trees, every descendant that stores rows. +unsafe fn descendants(relid: pg_sys::Oid) -> Vec { + let has_children = unsafe { + let class = pg_sys::SearchSysCache1( + pg_sys::SysCacheIdentifier::RELOID as i32, + pg_sys::Datum::from(relid), + ); + if class.is_null() { + return vec![relid]; + } + let form = pg_sys::GETSTRUCT(class).cast::(); + let flag = (*form).relhassubclass; + pg_sys::ReleaseSysCache(class); + flag + }; + if !has_children { + return vec![relid]; + } + let sql = format!( + "WITH RECURSIVE tree(oid) AS ( + SELECT {}::pg_catalog.oid + UNION ALL + SELECT i.inhrelid FROM pg_catalog.pg_inherits i JOIN tree ON i.inhparent = tree.oid) + SELECT tree.oid FROM tree JOIN pg_catalog.pg_class c ON c.oid = tree.oid + WHERE c.relkind <> 'p'", + relid.to_u32() + ); + Spi::connect(|client| { + client + .select(&sql, None, &[]) + .expect("inheritance tree lookup") + .filter_map(|row| row.get::(1).ok().flatten()) + .collect() + }) +} + +/// The tin index on `member` that covers `operand`, with its analysis. +/// `operand` is expressed in the columns of `parent`; inheritance children +/// may order their columns differently. +unsafe fn leaf_index( + parse: *mut pg_sys::Query, + member: pg_sys::Oid, + parent: pg_sys::Oid, + varno: i32, + operand: *mut pg_sys::Node, +) -> Option<(pg_sys::Oid, TokenizerPipelineSpec)> { + let index = if member == parent { + unsafe { crate::score::find_matching_tin_index(parse, member, varno, operand) } + } else { + unsafe { translate_to_member(member, parent, varno, operand) }.and_then( + |translated| unsafe { + crate::score::find_matching_tin_index(std::ptr::null_mut(), member, 1, translated) + }, + ) + }?; + let spec = unsafe { + let relation = pg_sys::index_open(index, pg_sys::AccessShareLock as _); + let spec = crate::options::tokenizer_spec(relation); + pg_sys::index_close(relation, pg_sys::AccessShareLock as _); + spec + }; + Some((index, spec)) +} + +unsafe fn translate_to_member( + member: pg_sys::Oid, + parent: pg_sys::Oid, + varno: i32, + operand: *mut pg_sys::Node, +) -> Option<*mut pg_sys::Node> { + unsafe { + let parent_rel = pg_sys::table_open(parent, pg_sys::AccessShareLock as _); + let member_rel = pg_sys::table_open(member, pg_sys::AccessShareLock as _); + // Maps parent column numbers to the member's. + let map = + pg_sys::build_attrmap_by_name_if_req((*member_rel).rd_att, (*parent_rel).rd_att, false); + let normalized = pg_sys::copyObjectImpl(operand.cast()).cast::(); + pg_sys::ChangeVarNodes(normalized, varno, 1, 0); + let mut whole_row = false; + let translated = if map.is_null() { + normalized + } else { + pg_sys::map_variable_attnos( + normalized, + 1, + 0, + map, + (*(*member_rel).rd_rel).reltype, + &mut whole_row, + ) + }; + pg_sys::table_close(member_rel, pg_sys::AccessShareLock as _); + pg_sys::table_close(parent_rel, pg_sys::AccessShareLock as _); + (!whole_row).then_some(translated) + } +} + +/// Keeps cached plans honest: changing or rebuilding a covering index +/// invalidates any plan that bound its analysis. +pub(crate) unsafe fn depend_on_indexes(root: *mut pg_sys::PlannerInfo, indexes: &[pg_sys::Oid]) { + unsafe { + let glob = (*root).glob; + for &index in indexes { + if !pg_sys::list_member_oid((*glob).relationOids, index) { + (*glob).relationOids = pg_sys::lappend_oid((*glob).relationOids, index); + } + } + } +} + +pub(crate) unsafe fn conflict_scope( + root: *mut pg_sys::PlannerInfo, + operand: *mut pg_sys::Node, +) -> String { + unsafe { + let varno = crate::score::single_varno(operand).unwrap_or(1); + let rte = + pg_sys::list_nth((*(*root).parse).rtable, varno - 1).cast::(); + let name = if (*rte).rtekind == pg_sys::RTEKind::RTE_RELATION { + pg_sys::get_rel_name((*rte).relid) + } else if !(*rte).eref.is_null() { + (*(*rte).eref).aliasname + } else { + std::ptr::null_mut() + }; + if name.is_null() { + String::new() + } else { + CStr::from_ptr(name).to_string_lossy().into_owned() + } + } +} + +fn unhandled() -> Internal { + Internal::from(Some(pg_sys::Datum::from(0_usize))) +} + +/// Planner support for `==>`: binds the analysis of the searched index to the +/// right-hand side so every evaluation of the operator, whether a bitmap +/// recheck, a filter, or a join qual, analyzes text the way the index does. +#[pg_extern(immutable, parallel_unsafe)] +fn tin_text_support(request: Internal) -> Internal { + let Some(datum) = request.into_datum() else { + return unhandled(); + }; + unsafe { + let node = datum.cast_mut_ptr::(); + if node.is_null() || (*node).type_ != pg_sys::NodeTag::T_SupportRequestSimplify { + return unhandled(); + } + let request = &*node.cast::(); + if request.root.is_null() + || request.fcall.is_null() + || pg_sys::list_length((*request.fcall).args) != 2 + { + return unhandled(); + } + let left = pg_sys::list_nth((*request.fcall).args, 0).cast::(); + let right = pg_sys::list_nth((*request.fcall).args, 1).cast::(); + if right.is_null() || is_bound_query(right) { + return unhandled(); + } + if (*right).type_ == pg_sys::NodeTag::T_Const + && (*right.cast::()).constisnull + { + return unhandled(); + } + if crate::score::single_varno(left).is_none() { + return unhandled(); + } + let (spec, indexes) = match resolve(request.root, left) { + Resolution::Unindexed => return unhandled(), + Resolution::Conflict => pgrx::error!( + "tin index tokenization differs across the members of \"{}\" for this search", + conflict_scope(request.root, left) + ), + Resolution::Bound { spec, indexes } => (spec, indexes), + }; + depend_on_indexes(request.root, &indexes); + if spec == TokenizerPipelineSpec::tin_default() { + return unhandled(); + } + let bound_right: *mut pg_sys::Node = match text_of(right) { + Some(text) => text_const(&tag(&spec, &text)).cast(), + None => { + let function = sibling_function( + (*request.fcall).funcid, + "bind_query_analysis", + &[pg_sys::TEXTOID, pg_sys::TEXTOID], + ); + let mut args = PgList::::new(); + args.push(right); + args.push(text_const(&encode(&spec)).cast()); + pg_sys::makeFuncExpr( + function, + pg_sys::TEXTOID, + args.into_pg(), + pg_sys::DEFAULT_COLLATION_OID, + pg_sys::DEFAULT_COLLATION_OID, + pg_sys::CoercionForm::COERCE_EXPLICIT_CALL, + ) + .cast() + } + }; + let operator = operator_oid(); + let replacement = pg_sys::make_opclause( + operator, + pg_sys::BOOLOID, + false, + left.cast(), + bound_right.cast(), + pg_sys::InvalidOid, + (*request.fcall).inputcollid, + ); + pg_sys::set_opfuncid(replacement.cast()); + Internal::from(Some(pg_sys::Datum::from(replacement as usize))) + } +} + +unsafe fn operator_oid() -> pg_sys::Oid { + let mut name = PgList::::new(); + name.push(unsafe { pg_sys::makeString(c"pg_catalog".as_ptr().cast_mut()) }.cast()); + name.push(unsafe { pg_sys::makeString(c"==>".as_ptr().cast_mut()) }.cast()); + unsafe { + pg_sys::LookupOperName( + std::ptr::null_mut(), + name.into_pg(), + pg_sys::TEXTOID, + pg_sys::TEXTOID, + false, + -1, + ) + } +} + +pgrx::extension_sql!( + r#" +ALTER FUNCTION @extschema@.tin_text_cmpfunc(pg_catalog.text, pg_catalog.text) + SUPPORT @extschema@.tin_text_support; +"#, + name = "tin_text_support_binding", + requires = ["tin_text_operator", tin_text_support, bind_query_analysis] +); + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn tags_round_trip_the_configuration_and_raw_text() { + let mut spec = TokenizerPipelineSpec::tin_default(); + spec.stemmer = Some(tokenizer::Stemmer::French); + let tagged = tag(&spec, "créées OR run*"); + let (config, raw) = split_tag(&tagged); + assert_eq!(raw, "créées OR run*"); + assert_eq!(decode(config.unwrap()).unwrap(), spec); + assert_eq!(split_tag("plain text"), (None, "plain text")); + } +} diff --git a/postgres/src/highlight.rs b/postgres/src/highlight.rs index eb7bf62..bab9edd 100644 --- a/postgres/src/highlight.rs +++ b/postgres/src/highlight.rs @@ -19,7 +19,7 @@ use rustc_hash::FxHasher; use std::borrow::Cow; use std::hash::{Hash, Hasher}; use thiserror::Error; -use tokenizer::Tokenizer; +use tokenizer::{CompiledTokenizerPipeline, Tokenizer}; const ANSI_RESET: &str = "\x1b[0m"; const ANSI_TERM_FALLBACK_PALETTE: [&str; 8] = [ @@ -104,13 +104,14 @@ enum HighlightStyle<'a> { } pub(crate) fn highlight_text( + pipeline: &CompiledTokenizerPipeline, text: &str, begin_tag: &str, end_tag: &str, positions: &[MatchPosition], ) -> Result { highlight_text_with_tokenizer( - &tokenizer::presets::default_pipeline().source_spans(), + &pipeline.source_spans(), text, HighlightStyle::Html { begin_tag, end_tag }, positions, @@ -118,11 +119,12 @@ pub(crate) fn highlight_text( } pub(crate) fn highlight_text_ansi( + pipeline: &CompiledTokenizerPipeline, text: &str, positions: &[MatchPosition], ) -> Result { highlight_text_with_tokenizer( - &tokenizer::presets::default_pipeline().source_spans(), + &pipeline.source_spans(), text, HighlightStyle::Ansi, positions, @@ -160,9 +162,13 @@ pub(crate) fn rewrap_text(text: &str, wrap_to: usize) -> String { /// Evaluates a parsed query against the document text to produce match /// positions inline. /// -/// Uses the built-in default analyzer pipeline for the document. -pub(crate) fn query_positions(query: &tinql::runtime::Query, text: &str) -> Vec { - let doc = tinql::runtime::tokenize_doc(text, tokenizer::presets::default_pipeline()); +/// Uses the given analyzer pipeline for the document. +pub(crate) fn query_positions( + pipeline: &CompiledTokenizerPipeline, + query: &tinql::runtime::Query, + text: &str, +) -> Vec { + let doc = tinql::runtime::tokenize_doc(text, pipeline); let matches = tinql::runtime::evaluate_for_highlight(query, &doc); matches .into_iter() diff --git a/postgres/src/highlight_udfs.rs b/postgres/src/highlight_udfs.rs index 5ba285b..17c70fa 100644 --- a/postgres/src/highlight_udfs.rs +++ b/postgres/src/highlight_udfs.rs @@ -14,108 +14,284 @@ // along with this program. If not, see . // // The full license text is available in LICENSE. +use crate::analysis::{Resolution, encode, resolve, sibling_function, text_const, unbound_query}; use crate::highlight::{highlight_text, highlight_text_ansi, query_positions, rewrap_text}; -use crate::tinql::parse_tinql_to_query_default; -use pgrx::{Internal, IntoDatum, PgList, default, pg_extern, pg_guard, pg_sys}; +use crate::tinql::parse_tinql_to_query; +use crate::udfs::TokenizeOptions; +use pgrx::{FromDatum, Internal, IntoDatum, PgList, default, pg_extern, pg_guard, pg_sys}; use std::borrow::Cow; use std::ffi::{CStr, c_void}; use tinql::runtime::Query; -use tokenizer::presets::default_pipeline; +use tokenizer::CompiledTokenizerPipeline; -fn render_highlight(text: &str, begin_tag: &str, end_tag: &str, query: Option<&Query>) -> String { +const ANALYSIS_ARGS: usize = 8; + +#[derive(Clone, Copy)] +struct AnalysisArgs<'a> { + tokenizer: Option<&'a str>, + case_folding: Option<&'a str>, + accent_folding: Option<&'a str>, + long_tokens: Option<&'a str>, + max_token_bytes: Option, + graphemes: Option<&'a str>, + position_gaps: Option<&'a str>, + stemmer: Option<&'a str>, +} + +impl AnalysisArgs<'_> { + /// Unset arguments mean the default analysis, or the searched index's + /// analysis once the planner has bound the call to one. + fn pipeline(self) -> CompiledTokenizerPipeline { + TokenizeOptions { + tokenizer: self.tokenizer.unwrap_or("unicode"), + case_folding: self.case_folding.unwrap_or("fold"), + accent_folding: self.accent_folding.unwrap_or("fold"), + long_tokens: self.long_tokens.unwrap_or("split"), + max_token_bytes: self.max_token_bytes.unwrap_or(256), + graphemes: self.graphemes.unwrap_or("emoji"), + position_gaps: self.position_gaps.unwrap_or("preserve"), + stemmer: self.stemmer, + } + .into_spec() + .and_then(|spec| spec.compile().map_err(|error| error.to_string())) + .unwrap_or_else(|error| pgrx::error!("{error}")) + } +} + +fn render_highlight( + pipeline: &CompiledTokenizerPipeline, + text: &str, + begin_tag: &str, + end_tag: &str, + query: Option<&Query>, +) -> String { // Without an explicit query or a tin index to bind one from, tin leaves // the text unmarked. let positions = query - .map(|query| query_positions(query, text)) + .map(|query| query_positions(pipeline, query, text)) .unwrap_or_default(); - highlight_text(text, begin_tag, end_tag, &positions) + highlight_text(pipeline, text, begin_tag, end_tag, &positions) .unwrap_or_else(|error| pgrx::error!("{error}")) } -fn render_highlight_ansi(text: &str, wrap_to: Option, query: Option<&Query>) -> String { +fn render_highlight_ansi( + pipeline: &CompiledTokenizerPipeline, + text: &str, + wrap_to: Option, + query: Option<&Query>, +) -> String { let text = match wrap_to { Some(width) if width <= 0 => pgrx::error!("wrap_to must be positive"), Some(width) => Cow::Owned(rewrap_text(text, width as usize)), None => Cow::Borrowed(text), }; let positions = query - .map(|query| query_positions(query, text.as_ref())) + .map(|query| query_positions(pipeline, query, text.as_ref())) .unwrap_or_default(); if positions.is_empty() { return text.into_owned(); } - highlight_text_ansi(text.as_ref(), &positions).unwrap_or_else(|error| pgrx::error!("{error}")) + highlight_text_ansi(pipeline, text.as_ref(), &positions) + .unwrap_or_else(|error| pgrx::error!("{error}")) } /// Parses an explicit highlight query. A query that does not parse marks /// nothing. -fn explicit_query(query: Option<&str>) -> Option { - parse_tinql_to_query_default(query?).ok() +fn explicit_query(pipeline: &CompiledTokenizerPipeline, query: Option<&str>) -> Option { + parse_tinql_to_query(query?, pipeline).ok() } /// Combines the search texts that `highlight_support` bound. A NULL text /// matches no rows, so it marks nothing. -fn bound_query(queries: Vec>) -> Option { - let texts = queries.into_iter().collect::>>()?; - Some(crate::operator::parse_searches(&texts, default_pipeline())) +fn bound_query( + pipeline: &CompiledTokenizerPipeline, + queries: Vec>, +) -> Option { + let texts = queries + .into_iter() + .map(|text| text.map(|text| crate::analysis::raw_query(&text).to_owned())) + .collect::>>()?; + Some(crate::operator::parse_searches(&texts, pipeline)) +} + +#[pg_extern(name = "highlight_v1_0_3", immutable, parallel_safe)] +fn highlight_v1_0_3( + text: Option<&str>, + begin_tag: default!(&str, "''"), + end_tag: default!(&str, "''"), + query: default!(Option<&str>, "NULL"), +) -> Option { + let pipeline = tokenizer::presets::default_pipeline(); + Some(render_highlight( + pipeline, + text?, + begin_tag, + end_tag, + explicit_query(pipeline, query).as_ref(), + )) +} + +#[pg_extern(name = "highlight_ansi_v1_0_3", immutable, parallel_safe)] +fn highlight_ansi_v1_0_3( + text: Option<&str>, + wrap_to: default!(Option, "NULL"), + query: default!(Option<&str>, "NULL"), +) -> Option { + let pipeline = tokenizer::presets::default_pipeline(); + Some(render_highlight_ansi( + pipeline, + text?, + wrap_to, + explicit_query(pipeline, query).as_ref(), + )) } #[pg_extern(name = "highlight", immutable, parallel_safe)] +#[expect(clippy::too_many_arguments, reason = "TIN-compatible SQL signature")] fn highlight( text: Option<&str>, begin_tag: default!(&str, "''"), end_tag: default!(&str, "''"), query: default!(Option<&str>, "NULL"), + tokenizer: default!(Option<&str>, "NULL"), + case_folding: default!(Option<&str>, "NULL"), + accent_folding: default!(Option<&str>, "NULL"), + long_tokens: default!(Option<&str>, "NULL"), + max_token_bytes: default!(Option, "NULL"), + graphemes: default!(Option<&str>, "NULL"), + position_gaps: default!(Option<&str>, "NULL"), + stemmer: default!(Option<&str>, "NULL"), ) -> Option { + let pipeline = AnalysisArgs { + tokenizer, + case_folding, + accent_folding, + long_tokens, + max_token_bytes, + graphemes, + position_gaps, + stemmer, + } + .pipeline(); Some(render_highlight( + &pipeline, text?, begin_tag, end_tag, - explicit_query(query).as_ref(), + explicit_query(&pipeline, query).as_ref(), )) } #[pg_extern(name = "highlight_ansi", immutable, parallel_safe)] +#[expect(clippy::too_many_arguments, reason = "TIN-compatible SQL signature")] fn highlight_ansi( text: Option<&str>, wrap_to: default!(Option, "NULL"), query: default!(Option<&str>, "NULL"), + tokenizer: default!(Option<&str>, "NULL"), + case_folding: default!(Option<&str>, "NULL"), + accent_folding: default!(Option<&str>, "NULL"), + long_tokens: default!(Option<&str>, "NULL"), + max_token_bytes: default!(Option, "NULL"), + graphemes: default!(Option<&str>, "NULL"), + position_gaps: default!(Option<&str>, "NULL"), + stemmer: default!(Option<&str>, "NULL"), ) -> Option { + let pipeline = AnalysisArgs { + tokenizer, + case_folding, + accent_folding, + long_tokens, + max_token_bytes, + graphemes, + position_gaps, + stemmer, + } + .pipeline(); Some(render_highlight_ansi( + &pipeline, text?, wrap_to, - explicit_query(query).as_ref(), + explicit_query(&pipeline, query).as_ref(), )) } /// `highlight` with the search texts of the quals `highlight_support` bound. #[pg_extern(immutable, parallel_safe)] +#[expect( + clippy::too_many_arguments, + reason = "SQL signature used by the support function" +)] fn highlight_bound( text: Option<&str>, begin_tag: &str, end_tag: &str, queries: Vec>, + tokenizer: Option<&str>, + case_folding: Option<&str>, + accent_folding: Option<&str>, + long_tokens: Option<&str>, + max_token_bytes: Option, + graphemes: Option<&str>, + position_gaps: Option<&str>, + stemmer: Option<&str>, ) -> Option { + let pipeline = AnalysisArgs { + tokenizer, + case_folding, + accent_folding, + long_tokens, + max_token_bytes, + graphemes, + position_gaps, + stemmer, + } + .pipeline(); Some(render_highlight( + &pipeline, text?, begin_tag, end_tag, - bound_query(queries).as_ref(), + bound_query(&pipeline, queries).as_ref(), )) } /// `highlight_ansi` with the search texts of the quals `highlight_support` /// bound. #[pg_extern(immutable, parallel_safe)] +#[expect( + clippy::too_many_arguments, + reason = "SQL signature used by the support function" +)] fn highlight_ansi_bound( text: Option<&str>, wrap_to: Option, queries: Vec>, + tokenizer: Option<&str>, + case_folding: Option<&str>, + accent_folding: Option<&str>, + long_tokens: Option<&str>, + max_token_bytes: Option, + graphemes: Option<&str>, + position_gaps: Option<&str>, + stemmer: Option<&str>, ) -> Option { + let pipeline = AnalysisArgs { + tokenizer, + case_folding, + accent_folding, + long_tokens, + max_token_bytes, + graphemes, + position_gaps, + stemmer, + } + .pipeline(); Some(render_highlight_ansi( + &pipeline, text?, wrap_to, - bound_query(queries).as_ref(), + bound_query(&pipeline, queries).as_ref(), )) } @@ -134,9 +310,9 @@ unsafe extern "C-unwind" fn collect_queries(node: *mut pg_sys::Node, context: *m } let context = unsafe { &mut *context.cast::() }; if let Some((document, query)) = unsafe { crate::operator::search_operands(node) } - && unsafe { pg_sys::equal(document.cast(), context.document.cast()) } + && unsafe { crate::score::same_operand(document, context.document) } { - context.queries.push(query); + context.queries.push(unsafe { unbound_query(query) }); } unsafe { pg_sys::expression_tree_walker( @@ -151,6 +327,116 @@ fn unhandled() -> Internal { Internal::from(Some(pg_sys::Datum::from(0_usize))) } +struct TargetSearch { + document: *mut pg_sys::Node, + found: bool, +} + +#[pg_guard] +unsafe extern "C-unwind" fn find_highlight_call( + node: *mut pg_sys::Node, + context: *mut c_void, +) -> bool { + if node.is_null() || unsafe { (*node).type_ } == pg_sys::NodeTag::T_Query { + return false; + } + let search = unsafe { &mut *context.cast::() }; + if unsafe { (*node).type_ } == pg_sys::NodeTag::T_FuncExpr { + let call = node.cast::(); + let name = unsafe { pg_sys::get_func_name((*call).funcid) }; + let first = unsafe { pg_sys::list_nth((*call).args, 0).cast::() }; + if !name.is_null() + && unsafe { CStr::from_ptr(name) } + .to_bytes() + .starts_with(b"highlight") + && unsafe { crate::score::same_operand(first, search.document) } + { + search.found = true; + return true; + } + } + unsafe { pg_sys::expression_tree_walker(node, Some(find_highlight_call), context) } +} + +/// A query that reads other relations can only be injected where the scan's +/// output is formed. Join and ordering expressions are planned before then and +/// have to receive the query explicitly. +unsafe fn query_needs_other_relations( + root: *mut pg_sys::PlannerInfo, + call: *mut pg_sys::FuncExpr, + document: *mut pg_sys::Node, + queries: &[*mut pg_sys::Node], +) -> bool { + unsafe { + let own = pg_sys::pull_varnos(root, call.cast()); + let needed = queries + .iter() + .any(|&query| !pg_sys::bms_is_subset(pg_sys::pull_varnos(root, query), own)); + if !needed { + return false; + } + let mut search = TargetSearch { + document, + found: false, + }; + let targets = (*(*root).parse).targetList.cast::(); + find_highlight_call(targets, (&raw mut search).cast()); + !search.found + } +} + +const CONFLICT: &str = "tin.highlight() analysis conflicts with the searched index; matching \ + columns, including UNION ALL and partition children, must use the same analysis options"; + +/// Whether a supplied analysis argument leaves the index's analysis in force. +/// NULL inherits it; any other value has to be the same constant. +unsafe fn agrees(supplied: *mut pg_sys::Node, configured: *mut pg_sys::Node) -> bool { + unsafe { + if supplied.is_null() || (*supplied).type_ != pg_sys::NodeTag::T_Const { + return false; + } + let supplied = &*supplied.cast::(); + let configured = &*configured.cast::(); + if supplied.constisnull { + return true; + } + if configured.constisnull || supplied.consttype != configured.consttype { + return false; + } + if supplied.consttype == pg_sys::INT4OID { + i32::from_datum(supplied.constvalue, false) + == i32::from_datum(configured.constvalue, false) + } else { + <&str>::from_datum(supplied.constvalue, false) + == <&str>::from_datum(configured.constvalue, false) + } + } +} + +unsafe fn configuration_args(spec: &tokenizer::TokenizerPipelineSpec) -> Vec<*mut pg_sys::Node> { + let encoded = encode(spec); + let fields: Vec<&str> = encoded.split(',').collect(); + assert_eq!( + fields.len(), + ANALYSIS_ARGS, + "encoded analysis has eight fields" + ); + fields + .into_iter() + .enumerate() + .map(|(position, field)| unsafe { + match position { + 4 => crate::score::make_int4_const( + field.parse().expect("encoded token limit is an integer"), + ) + .cast(), + 7 if field.is_empty() => crate::score::make_null_const(pg_sys::TEXTOID).cast(), + _ => text_const(field).cast(), + } + }) + .collect() +} + #[pg_extern(immutable, parallel_unsafe)] fn highlight_support(request: Internal) -> Internal { let Some(datum) = request.into_datum() else { @@ -170,95 +456,141 @@ fn highlight_support(request: Internal) -> Internal { return unhandled(); } let name = CStr::from_ptr(function_name).to_bytes(); - let (query_position, bound, types): (_, _, &[pg_sys::Oid]) = if name == b"highlight" { - ( - 3, - c"tin.highlight_bound", - &[ - pg_sys::TEXTOID, - pg_sys::TEXTOID, - pg_sys::TEXTOID, - pg_sys::TEXTARRAYOID, - ], - ) - } else if name == b"highlight_ansi" { - ( - 2, - c"tin.highlight_ansi_bound", - &[pg_sys::TEXTOID, pg_sys::INT4OID, pg_sys::TEXTARRAYOID], - ) - } else { - return unhandled(); + let (query_position, legacy) = match name { + b"highlight" => (3, false), + b"highlight_v1_0_3" => (3, true), + b"highlight_ansi" => (2, false), + b"highlight_ansi_v1_0_3" => (2, true), + _ => return unhandled(), }; - if pg_sys::list_length((*request.fcall).args) <= query_position { - return unhandled(); - } - let supplied_query = - pg_sys::list_nth((*request.fcall).args, query_position).cast::(); - if supplied_query.is_null() || (*supplied_query).type_ != pg_sys::NodeTag::T_Const { + let args = PgList::::from_pg((*request.fcall).args); + let expected = query_position + 1 + if legacy { 0 } else { ANALYSIS_ARGS }; + if args.len() != expected { return unhandled(); } - if !(*supplied_query.cast::()).constisnull { + let document = args.get_ptr(0).expect("highlight has a text argument"); + if crate::score::single_varno(document).is_none() { return unhandled(); } - let document = pg_sys::list_nth((*request.fcall).args, 0).cast::(); - let Some(varno) = crate::score::single_varno(document) else { - return unhandled(); + let (spec, indexes) = match resolve(request.root, document) { + Resolution::Unindexed => return unhandled(), + Resolution::Conflict => pgrx::error!("{CONFLICT}"), + Resolution::Bound { spec, indexes } => (spec, indexes), }; - let parse = (*request.root).parse; - let rte = pg_sys::list_nth((*parse).rtable, varno - 1).cast::(); - if rte.is_null() - || (*rte).rtekind != pg_sys::RTEKind::RTE_RELATION - || crate::score::find_matching_tin_index(parse, (*rte).relid, varno, document).is_none() - { - return unhandled(); - } let mut binding = QueryContext { document, queries: Vec::new(), }; // Pulled-up subqueries leave their quals in nested FromExpr nodes. collect_queries( - (*parse).jointree.cast::(), + (*(*request.root).parse).jointree.cast::(), (&mut binding as *mut QueryContext).cast(), ); if binding.queries.is_empty() { return unhandled(); } - // Each search text is parsed on its own when the plan runs, since - // parameters and other run-time expressions have no text until then. - let mut queries = PgList::::new(); - for &query in &binding.queries { - queries.push(query); + crate::analysis::depend_on_indexes(request.root, &indexes); + let supplied_query = args + .get_ptr(query_position) + .expect("highlight has a query argument"); + let implicit = (*supplied_query).type_ == pg_sys::NodeTag::T_Const + && (*supplied_query.cast::()).constisnull; + if implicit + && query_needs_other_relations(request.root, request.fcall, document, &binding.queries) + { + pgrx::error!( + "tin.highlight() cannot infer its runtime query in this join or ordering \ + expression: the query requires additional relations; pass the tinql query \ + text explicitly as the highlight() query argument" + ); + } + let configured = configuration_args(&spec); + let mut rewritten = PgList::::new(); + for position in 0..query_position { + rewritten.push(args.get_ptr(position).expect("highlight argument exists")); + } + rewritten.push(if implicit { + // Each search text is parsed on its own when the plan runs, since + // parameters and other run-time expressions have no text until then. + let mut queries = PgList::::new(); + for &query in &binding.queries { + queries.push(pg_sys::copyObjectImpl(query.cast()).cast()); + } + crate::score::make_text_array(queries) + } else { + supplied_query + }); + for (offset, configured) in configured.into_iter().enumerate() { + if !legacy { + let supplied = args + .get_ptr(query_position + 1 + offset) + .expect("highlight analysis argument exists"); + if !agrees(supplied, configured) { + pgrx::error!("{CONFLICT}"); + } + } + rewritten.push(configured); + } + let ansi = query_position == 2; + let mut types = vec![pg_sys::TEXTOID; query_position + 1]; + if ansi { + types[1] = pg_sys::INT4OID; + } + if implicit { + types[query_position] = pg_sys::TEXTARRAYOID; + } + types.extend([ + pg_sys::TEXTOID, + pg_sys::TEXTOID, + pg_sys::TEXTOID, + pg_sys::TEXTOID, + pg_sys::INT4OID, + pg_sys::TEXTOID, + pg_sys::TEXTOID, + pg_sys::TEXTOID, + ]); + if !implicit + && !legacy + && pg_sys::equal((*request.fcall).args.cast(), rewritten.as_ptr().cast()) + { + return unhandled(); } - let query = crate::score::make_text_array(queries); let replacement = pg_sys::copyObjectImpl(request.fcall.cast()).cast::(); - (*replacement).funcid = crate::score::lookup_bound_function(bound, types); - let mut args = PgList::::new(); - for position in 0..pg_sys::list_length((*request.fcall).args) { - let argument = if position == query_position { - query + if implicit { + let bound = if ansi { + c"tin.highlight_ansi_bound" } else { - pg_sys::list_nth((*request.fcall).args, position).cast::() + c"tin.highlight_bound" }; - args.push(pg_sys::copyObjectImpl(argument.cast()).cast()); + (*replacement).funcid = crate::score::lookup_bound_function(bound, &types); + } else if legacy { + let base = if ansi { "highlight_ansi" } else { "highlight" }; + (*replacement).funcid = sibling_function((*request.fcall).funcid, base, &types); } - (*replacement).args = args.into_pg(); + (*replacement).args = rewritten.into_pg(); Internal::from(Some(pg_sys::Datum::from(replacement as usize))) } } pgrx::extension_sql!( r#" -ALTER FUNCTION @extschema@.highlight(pg_catalog.text, pg_catalog.text, pg_catalog.text, pg_catalog.text) +ALTER FUNCTION @extschema@.highlight_v1_0_3(pg_catalog.text, pg_catalog.text, pg_catalog.text, pg_catalog.text) SUPPORT @extschema@.highlight_support; -ALTER FUNCTION @extschema@.highlight_ansi(pg_catalog.text, pg_catalog.int4, pg_catalog.text) +ALTER FUNCTION @extschema@.highlight_ansi_v1_0_3(pg_catalog.text, pg_catalog.int4, pg_catalog.text) + SUPPORT @extschema@.highlight_support; +ALTER FUNCTION @extschema@.highlight(pg_catalog.text, pg_catalog.text, pg_catalog.text, pg_catalog.text, + pg_catalog.text, pg_catalog.text, pg_catalog.text, pg_catalog.text, pg_catalog.int4, pg_catalog.text, pg_catalog.text, pg_catalog.text) + SUPPORT @extschema@.highlight_support; +ALTER FUNCTION @extschema@.highlight_ansi(pg_catalog.text, pg_catalog.int4, pg_catalog.text, + pg_catalog.text, pg_catalog.text, pg_catalog.text, pg_catalog.text, pg_catalog.int4, pg_catalog.text, pg_catalog.text, pg_catalog.text) SUPPORT @extschema@.highlight_support; "#, name = "highlight_support_bindings", requires = [ highlight, highlight_ansi, + highlight_v1_0_3, + highlight_ansi_v1_0_3, highlight_bound, highlight_ansi_bound, highlight_support @@ -390,10 +722,36 @@ mod tests { #[pg_test] fn explicit_html_and_ansi_highlighting_render_matches() { assert_eq!( - highlight(Some("Hi there"), "", "", Some("hi")), + highlight( + Some("Hi there"), + "", + "", + Some("hi"), + None, + None, + None, + None, + None, + None, + None, + None + ), Some("Hi there".into()) ); - let ansi = highlight_ansi(Some("hi there"), None, Some("hi")).unwrap(); + let ansi = highlight_ansi( + Some("hi there"), + None, + Some("hi"), + None, + None, + None, + None, + None, + None, + None, + None, + ) + .unwrap(); assert!(ansi.contains("\x1b[")); assert!(ansi.contains("hi")); } diff --git a/postgres/src/lib.rs b/postgres/src/lib.rs index 283030f..adc8f4a 100644 --- a/postgres/src/lib.rs +++ b/postgres/src/lib.rs @@ -19,6 +19,7 @@ use pgrx::pg_guard; ::pgrx::pg_module_magic!(name); mod am; +mod analysis; mod bm25; mod highlight; mod highlight_udfs; @@ -1170,6 +1171,131 @@ mod tests { assert!(ansi.contains("\x1b[")); assert!(ansi.contains("Beer")); } + + #[pg_test] + fn stemmed_index_matches_inflected_words_everywhere() { + Spi::run( + "CREATE TABLE lite_stem (id int, body text); + INSERT INTO lite_stem VALUES + (1, 'molded'), (2, 'MOLDS'), (3, 'moldy'), (4, 'wine'); + CREATE INDEX lite_stem_idx ON lite_stem USING tin (body) + WITH (stemmer = 'en');", + ) + .unwrap(); + for setting in ["on", "off"] { + Spi::run(&format!("SET LOCAL enable_seqscan = {setting}")).unwrap(); + assert_eq!( + Spi::get_one::>( + "SELECT array_agg(id ORDER BY id) FROM lite_stem WHERE body ==> 'mold'" + ) + .unwrap(), + Some(vec![1, 2]), + "enable_seqscan = {setting}" + ); + } + assert_eq!( + Spi::get_one::( + "SELECT count(*) FROM lite_stem a JOIN lite_stem b ON a.id = b.id + WHERE b.body ==> 'molding'" + ) + .unwrap(), + Some(2) + ); + } + + #[pg_test] + fn cached_plans_follow_the_index_stemmer() { + Spi::run( + "CREATE TABLE lite_stem_cached (id int, body text); + INSERT INTO lite_stem_cached VALUES (1, 'molded'), (2, 'mold'); + CREATE INDEX lite_stem_cached_idx ON lite_stem_cached USING tin (body) + WITH (stemmer = 'en'); + SET LOCAL plan_cache_mode = force_generic_plan; + PREPARE lite_stem_query(text) AS + SELECT array_agg(id ORDER BY id) FROM lite_stem_cached WHERE body ==> $1;", + ) + .unwrap(); + let run = || { + Spi::get_one::>("EXECUTE lite_stem_query('molds')") + .unwrap() + .unwrap_or_default() + }; + assert_eq!(run(), vec![1, 2]); + Spi::run( + "ALTER INDEX lite_stem_cached_idx RESET (stemmer); + REINDEX INDEX lite_stem_cached_idx;", + ) + .unwrap(); + assert!(run().is_empty()); + } + + #[pg_test(error = "unknown stemmer language code: bogus")] + fn unknown_stemmer_languages_are_rejected() { + Spi::run( + "CREATE TABLE lite_stem_unknown (body text); + CREATE INDEX ON lite_stem_unknown USING tin (body) WITH (stemmer = 'bogus');", + ) + .unwrap(); + } + + #[pg_test(error = "stemming requires case_folding = fold")] + fn stemming_requires_case_folding() { + Spi::run( + "CREATE TABLE lite_stem_case (body text); + CREATE INDEX ON lite_stem_case USING tin (body) + WITH (case_folding = preserve); + ALTER INDEX lite_stem_case_body_idx SET (stemmer = 'en');", + ) + .unwrap(); + } + + #[pg_test] + fn highlights_bind_the_index_analysis() { + Spi::run( + "CREATE TABLE lite_stem_highlight (body text); + INSERT INTO lite_stem_highlight VALUES ('Running and RUNS'); + CREATE INDEX ON lite_stem_highlight USING tin (body) WITH (stemmer = 'en');", + ) + .unwrap(); + for call in [ + "tin.highlight(body)", + "tin.highlight(body, query => 'runs')", + "tin.highlight_v1_0_3(body, query => 'running')", + "tin.highlight(body, stemmer => 'en')", + ] { + assert_eq!( + Spi::get_one::(&format!( + "SELECT {call} FROM lite_stem_highlight WHERE body ==> 'run'" + )) + .unwrap(), + Some("Running and RUNS".into()), + "{call}" + ); + } + assert_eq!( + Spi::get_one::( + "SELECT tin.highlight('Running and RUNS', query => 'runs', stemmer => 'en')" + ) + .unwrap(), + Some("Running and RUNS".into()) + ); + } + + #[pg_test( + error = "tin.highlight() analysis conflicts with the searched index; matching columns, including UNION ALL and partition children, must use the same analysis options" + )] + fn highlights_reject_analysis_that_conflicts_with_the_index() { + Spi::run( + "CREATE TABLE lite_stem_conflict (body text); + CREATE INDEX ON lite_stem_conflict USING tin (body) WITH (stemmer = 'en');", + ) + .unwrap(); + Spi::run( + "SELECT tin.highlight(body, stemmer => 'fr') + FROM lite_stem_conflict WHERE body ==> 'runs'", + ) + .unwrap(); + } } /// Scoring checks that span several transactions or sessions, which a diff --git a/postgres/src/operator.rs b/postgres/src/operator.rs index 6f35ee9..90dadc7 100644 --- a/postgres/src/operator.rs +++ b/postgres/src/operator.rs @@ -96,7 +96,21 @@ pub(crate) fn parse_searches(texts: &[String], tokenizer: &T) -> Q } fn evaluate_text(document: &str, query_text: &str) -> Result { - let pipeline = default_pipeline(); + let (config, query_text) = crate::analysis::split_tag(query_text); + match config { + Some(config) => { + let pipeline = crate::analysis::pipeline_for(config)?; + evaluate_with(&pipeline, document, query_text) + } + None => evaluate_with(default_pipeline(), document, query_text), + } +} + +fn evaluate_with( + pipeline: &tokenizer::CompiledTokenizerPipeline, + document: &str, + query_text: &str, +) -> Result { let query = parse_search(query_text, pipeline)?; let document = tokenize_doc(document, pipeline); evaluate(&query, &document) @@ -150,6 +164,16 @@ mod tests { assert!(!evaluate_text("...", "*").unwrap()); } + #[pg_test] + fn bound_analysis_stems_the_raw_query_once() { + let mut spec = tokenizer::TokenizerPipelineSpec::tin_default(); + spec.stemmer = Some(tokenizer::Stemmer::English); + let bound = crate::analysis::tag(&spec, "accidental"); + assert!(evaluate_text("an accidental", &bound).unwrap()); + assert!(!evaluate_text("an accident", &bound).unwrap()); + assert!(!evaluate_text("an accidental", "accident").unwrap()); + } + #[pg_test] fn invalid_queries_are_reported() { assert!(evaluate_text("beer", "beer OR").is_err()); diff --git a/postgres/src/options.rs b/postgres/src/options.rs index 279ee60..885cd6d 100644 --- a/postgres/src/options.rs +++ b/postgres/src/options.rs @@ -101,6 +101,8 @@ struct IndexOptions { k1: f64, b: f64, score_stop_words: i32, + /// String offset; zero disables stemming. + stemmer: i32, } pub fn init() { @@ -208,6 +210,15 @@ pub fn init() { None, lock, ); + pg_sys::add_string_reloption( + kind, + c"stemmer".as_ptr(), + c"Snowball stemmer language code (unset disables stemming; changes require REINDEX)" + .as_ptr(), + std::ptr::null(), + None, + lock, + ); OPTION_KIND.store(kind, Ordering::Relaxed); } } @@ -287,8 +298,13 @@ pub unsafe extern "C-unwind" fn amoptions( pg_sys::relopt_type::RELOPT_TYPE_STRING, std::mem::offset_of!(IndexOptions, score_stop_words), ), + parse_entry( + c"stemmer".as_ptr(), + pg_sys::relopt_type::RELOPT_TYPE_STRING, + std::mem::offset_of!(IndexOptions, stemmer), + ), ]; - unsafe { + let packed = unsafe { pg_sys::build_reloptions( reloptions, validate, @@ -297,8 +313,16 @@ pub unsafe extern "C-unwind" fn amoptions( entries.as_ptr(), entries.len() as i32, ) - .cast() + .cast::() + }; + if validate { + // ALTER SET/RESET supplies the merged options, so the stemmer and + // case folding are checked together even when only one changed. + spec_from(unsafe { packed.as_ref() }) + .validate() + .unwrap_or_else(|error| pgrx::error!("{error}")); } + packed.cast() } unsafe fn parsed(index: pg_sys::Relation) -> Option<&'static IndexOptions> { @@ -309,7 +333,11 @@ unsafe fn parsed(index: pg_sys::Relation) -> Option<&'static IndexOptions> { } pub unsafe fn tokenizer_spec(index: pg_sys::Relation) -> TokenizerPipelineSpec { - let Some(options) = (unsafe { parsed(index) }) else { + spec_from(unsafe { parsed(index) }) +} + +fn spec_from(options: Option<&IndexOptions>) -> TokenizerPipelineSpec { + let Some(options) = options else { return TokenizerPipelineSpec::tin_default(); }; TokenizerPipelineSpec { @@ -336,9 +364,24 @@ pub unsafe fn tokenizer_spec(index: pg_sys::Relation) -> TokenizerPipelineSpec { GAPS_COLLAPSE => PositionGapMode::Collapse, _ => PositionGapMode::Preserve, }, + stemmer: stemmer(options), } } +fn stemmer(options: &IndexOptions) -> Option { + let offset = usize::try_from(options.stemmer).ok().filter(|&n| n != 0)?; + let ptr = std::ptr::from_ref(options).cast::(); + // build_reloptions owns the trailing, NUL-terminated string. + let language = unsafe { CStr::from_ptr(ptr.add(offset).cast()) } + .to_str() + .expect("stemmer reloption must be valid UTF-8"); + Some( + language + .parse() + .unwrap_or_else(|error| pgrx::error!("{error}")), + ) +} + fn decode_folding(value: i32) -> Folding { if value == FOLDING_PRESERVE { Folding::Preserve diff --git a/postgres/src/score.rs b/postgres/src/score.rs index 7c9c164..940dc90 100644 --- a/postgres/src/score.rs +++ b/postgres/src/score.rs @@ -226,7 +226,7 @@ fn score_bound( .take() .zip(search) .map(|(mut queries, search)| { - queries.push(search); + queries.push(crate::analysis::raw_query(&search).to_owned()); queries }); } @@ -780,7 +780,11 @@ unsafe fn restriction_clauses(parse: *mut pg_sys::Query, varno: i32) -> *mut pg_ if qual.is_null() { continue; } - let qual = pg_sys::copyObjectImpl(qual.cast()).cast::(); + let mut qual = pg_sys::copyObjectImpl(qual.cast()).cast::(); + // Quals the planner already preprocessed are implicit-AND lists. + if (*qual).type_ == pg_sys::NodeTag::T_List { + qual = pg_sys::make_ands_explicit(qual.cast()).cast(); + } let qual = pg_sys::remove_nulling_relids(qual, every_relation, std::ptr::null()); // Without a planner root, support functions decline to rewrite, // so this cannot reenter scoring support. @@ -792,11 +796,36 @@ unsafe fn restriction_clauses(parse: *mut pg_sys::Query, varno: i32) -> *mut pg_ } } +#[pg_guard] +unsafe extern "C-unwind" fn clear_nulling(node: *mut pg_sys::Node, context: *mut c_void) -> bool { + if node.is_null() { + return false; + } + if unsafe { (*node).type_ } == pg_sys::NodeTag::T_Var { + unsafe { (*node.cast::()).varnullingrels = std::ptr::null_mut() }; + return false; + } + unsafe { pg_sys::expression_tree_walker(node, Some(clear_nulling), context) } +} + +/// Compares two operands while ignoring outer-join nulling metadata, which +/// differs between a join's own clauses and expressions above the join. +pub(crate) unsafe fn same_operand(left: *mut pg_sys::Node, right: *mut pg_sys::Node) -> bool { + unsafe { + let left = pg_sys::copyObjectImpl(left.cast()).cast::(); + let right = pg_sys::copyObjectImpl(right.cast()).cast::(); + clear_nulling(left, std::ptr::null_mut()); + clear_nulling(right, std::ptr::null_mut()); + pg_sys::equal(left.cast(), right.cast()) + } +} + /// Returns the tin index that scoring and highlighting bind `operand` to, the /// way tin selects one: a partial index qualifies only when the query's quals /// on the relation imply its predicate. Lead's indexes store no pages, so of /// the qualifying indexes the newest partial one wins, then the newest full -/// one. +/// one. With a null `parse`, only full indexes qualify: the query's quals are +/// expressed in the columns of another relation. pub(crate) unsafe fn find_matching_tin_index( parse: *mut pg_sys::Query, heap_oid: pg_sys::Oid, @@ -812,6 +841,8 @@ pub(crate) unsafe fn find_matching_tin_index( let tin_am = unsafe { pg_sys::get_index_am_oid(tin_name.as_ptr(), false) }; let normalized = unsafe { pg_sys::copyObjectImpl(operand.cast()).cast::() }; unsafe { pg_sys::ChangeVarNodes(normalized, query_varno, 1, 0) }; + // Outer-join nulling metadata does not change which index covers a column. + unsafe { clear_nulling(normalized, std::ptr::null_mut()) }; let normalized = unsafe { pg_sys::strip_implicit_coercions(normalized) }; let heap = unsafe { pg_sys::table_open(heap_oid, pg_sys::AccessShareLock as _) }; let indexes = unsafe { PgList::::from_pg(pg_sys::RelationGetIndexList(heap)) }; @@ -854,12 +885,13 @@ pub(crate) unsafe fn find_matching_tin_index( let partial = !predicate.is_null(); let eligible = matches && (!partial - || unsafe { - pg_sys::ChangeVarNodes(predicate.cast(), 1, query_varno, 0); - let clauses = - *clauses.get_or_insert_with(|| restriction_clauses(parse, query_varno)); - pg_sys::predicate_implied_by(predicate, clauses, false) - }); + || !parse.is_null() + && unsafe { + pg_sys::ChangeVarNodes(predicate.cast(), 1, query_varno, 0); + let clauses = + *clauses.get_or_insert_with(|| restriction_clauses(parse, query_varno)); + pg_sys::predicate_implied_by(predicate, clauses, false) + }); unsafe { pg_sys::index_close(index, pg_sys::AccessShareLock as _) }; let choice = (partial, index_oid.to_u32()); if eligible && matched.is_none_or(|best| choice > best) { @@ -1433,7 +1465,7 @@ pub(crate) unsafe fn make_text_array(elements: PgList) -> *mut pg_ array.into_pg().cast() } -unsafe fn make_int4_const(value: i32) -> *mut pg_sys::Const { +pub(crate) unsafe fn make_int4_const(value: i32) -> *mut pg_sys::Const { unsafe { pg_sys::makeConst( pg_sys::INT4OID, @@ -1447,7 +1479,7 @@ unsafe fn make_int4_const(value: i32) -> *mut pg_sys::Const { } } -unsafe fn make_null_const(type_oid: pg_sys::Oid) -> *mut pg_sys::Const { +pub(crate) unsafe fn make_null_const(type_oid: pg_sys::Oid) -> *mut pg_sys::Const { unsafe { pg_sys::makeConst( type_oid, diff --git a/postgres/src/tinql.rs b/postgres/src/tinql.rs index 82d2a2b..ccc9f23 100644 --- a/postgres/src/tinql.rs +++ b/postgres/src/tinql.rs @@ -74,8 +74,3 @@ pub(crate) fn parse_tinql_to_query( ) -> Result { Ok(lower(&sub_tokenize(parse(query)?, tokenizer)?)?) } - -/// [`parse_tinql_to_query`] with the extension's default tokenization pipeline. -pub(crate) fn parse_tinql_to_query_default(query: &str) -> Result { - parse_tinql_to_query(query, tokenizer::presets::default_pipeline()) -} diff --git a/postgres/src/udfs.rs b/postgres/src/udfs.rs index e337555..44e147f 100644 --- a/postgres/src/udfs.rs +++ b/postgres/src/udfs.rs @@ -23,18 +23,19 @@ use tokenizer::{ pub const MAX_TOKEN_BYTES: usize = 2_692; #[derive(Debug, Clone, Copy)] -struct TokenizeOptions<'a> { - tokenizer: &'a str, - case_folding: &'a str, - accent_folding: &'a str, - long_tokens: &'a str, - max_token_bytes: i32, - graphemes: &'a str, - position_gaps: &'a str, +pub(crate) struct TokenizeOptions<'a> { + pub(crate) tokenizer: &'a str, + pub(crate) case_folding: &'a str, + pub(crate) accent_folding: &'a str, + pub(crate) long_tokens: &'a str, + pub(crate) max_token_bytes: i32, + pub(crate) graphemes: &'a str, + pub(crate) position_gaps: &'a str, + pub(crate) stemmer: Option<&'a str>, } impl TokenizeOptions<'_> { - fn into_spec(self) -> Result { + pub(crate) fn into_spec(self) -> Result { let max_bytes = usize::try_from(self.max_token_bytes) .ok() .filter(|n| (tokenizer::MIN_TOKEN_BYTES..=MAX_TOKEN_BYTES).contains(n)) @@ -45,7 +46,7 @@ impl TokenizeOptions<'_> { ) })?; - Ok(TokenizerPipelineSpec { + let spec = TokenizerPipelineSpec { tokenizer: parse_tokenizer(self.tokenizer)?, case_folding: parse_folding("case_folding", self.case_folding)?, accent_folding: parse_folding("accent_folding", self.accent_folding)?, @@ -55,7 +56,14 @@ impl TokenizeOptions<'_> { }, graphemes: parse_graphemes(self.graphemes)?, position_gaps: parse_position_gaps(self.position_gaps)?, - }) + stemmer: self + .stemmer + .map(str::parse) + .transpose() + .map_err(|error: tokenizer::TokenizerPipelineSpecError| error.to_string())?, + }; + spec.validate().map_err(|error| error.to_string())?; + Ok(spec) } } @@ -127,7 +135,7 @@ fn collect_tokens(text: &str, spec: TokenizerPipelineSpec) -> Vec { .collect() } -#[pg_extern(immutable, parallel_safe)] +#[pg_extern(name = "tokenize_v1_0_3", immutable, parallel_safe)] #[expect(clippy::too_many_arguments, reason = "TIN-compatible SQL signature")] pub fn tokenize<'a>( text: Option<&'a str>, @@ -138,6 +146,32 @@ pub fn tokenize<'a>( max_token_bytes: default!(i32, 256), graphemes: default!(&str, "'emoji'"), position_gaps: default!(&str, "'preserve'"), +) -> SetOfIterator<'a, String> { + tokenize_with_stemmer( + text, + tokenizer, + case_folding, + accent_folding, + long_tokens, + max_token_bytes, + graphemes, + position_gaps, + None, + ) +} + +#[pg_extern(name = "tokenize", immutable, parallel_safe)] +#[expect(clippy::too_many_arguments, reason = "TIN-compatible SQL signature")] +pub fn tokenize_with_stemmer<'a>( + text: Option<&'a str>, + tokenizer: default!(&str, "'unicode'"), + case_folding: default!(&str, "'fold'"), + accent_folding: default!(&str, "'fold'"), + long_tokens: default!(&str, "'split'"), + max_token_bytes: default!(i32, 256), + graphemes: default!(&str, "'emoji'"), + position_gaps: default!(&str, "'preserve'"), + stemmer: default!(Option<&str>, "NULL"), ) -> SetOfIterator<'a, String> { let pipeline = compile_options(TokenizeOptions { tokenizer, @@ -147,6 +181,7 @@ pub fn tokenize<'a>( max_token_bytes, graphemes, position_gaps, + stemmer, }); match text { Some(text) => SetOfIterator::new( @@ -164,7 +199,7 @@ pub fn maybe_quote(text: Option<&str>) -> Option { text.map(|text| tinql::maybe_quote(text).into_owned()) } -#[pg_extern(immutable, parallel_safe)] +#[pg_extern(name = "ql_parse_v1_0_3", immutable, parallel_safe)] #[expect(clippy::too_many_arguments, reason = "TIN-compatible SQL signature")] pub fn ql_parse( query: Option<&str>, @@ -176,6 +211,34 @@ pub fn ql_parse( max_token_bytes: default!(i32, 256), graphemes: default!(&str, "'emoji'"), position_gaps: default!(&str, "'preserve'"), +) -> Option { + ql_parse_with_stemmer( + query, + surface, + tokenizer, + case_folding, + accent_folding, + long_tokens, + max_token_bytes, + graphemes, + position_gaps, + None, + ) +} + +#[pg_extern(name = "ql_parse", immutable, parallel_safe)] +#[expect(clippy::too_many_arguments, reason = "TIN-compatible SQL signature")] +pub fn ql_parse_with_stemmer( + query: Option<&str>, + surface: default!(bool, true), + tokenizer: default!(&str, "'unicode'"), + case_folding: default!(&str, "'fold'"), + accent_folding: default!(&str, "'fold'"), + long_tokens: default!(&str, "'split'"), + max_token_bytes: default!(i32, 256), + graphemes: default!(&str, "'emoji'"), + position_gaps: default!(&str, "'preserve'"), + stemmer: default!(Option<&str>, "NULL"), ) -> Option { let query = query?; let pipeline = compile_options(TokenizeOptions { @@ -186,6 +249,7 @@ pub fn ql_parse( max_token_bytes, graphemes, position_gaps, + stemmer, }); let parsed = crate::tinql::parse(query).unwrap_or_else(|error| pgrx::error!("{error}")); let analyzed = tinql::runtime::subtokenize::sub_tokenize(parsed, &pipeline) @@ -215,6 +279,7 @@ mod tests { max_token_bytes: 32, graphemes: "retain", position_gaps: "collapse", + stemmer: None, } .into_spec() .unwrap(); @@ -247,6 +312,7 @@ mod tests { max_token_bytes: 3, graphemes: "emoji", position_gaps: "preserve", + stemmer: None, }; assert!(options.into_spec().is_err()); } diff --git a/private-regress.manifest b/private-regress.manifest index d7197f1..44df0f4 100644 --- a/private-regress.manifest +++ b/private-regress.manifest @@ -28,3 +28,10 @@ highlight_nested_relation_undermark highlight_relation_witness_overmark highlight_ilike_filter or_near_highlight +stemming +stemming_canonical_unicode +highlight_index_analysis +highlight_expression_analysis +highlight_having_analysis +highlight_join_analysis +highlight_null_analysis_defaults diff --git a/tinql/src/runtime/subtokenize.rs b/tinql/src/runtime/subtokenize.rs index 7c31a82..6c453b1 100644 --- a/tinql/src/runtime/subtokenize.rs +++ b/tinql/src/runtime/subtokenize.rs @@ -17,7 +17,7 @@ //! Sub-tokenization rewrite pass for the shared runtime. use crate::*; -use tokenizer::Tokenizer; +use tokenizer::{TokenStream, Tokenizer}; #[derive(Debug, thiserror::Error)] pub enum SubTokenizeError { @@ -41,9 +41,8 @@ struct AnalyzedToken { came_from_split: bool, } -fn analyze(text: &str, tokenizer: &T) -> Vec { - tokenizer - .tokenize(text) +fn analyze<'text>(tokens: impl Iterator>) -> Vec { + tokens .map(|token| { let pos = token.pos; let came_from_split = token.came_from_split(); @@ -64,7 +63,7 @@ fn rewrite(expr: Expr, tokenizer: &T) -> Result { - let tokens = analyze(&term, tokenizer); + let tokens = analyze(tokenizer.tokenize_literal(&term)); if tokens.iter().any(|token| token.came_from_split) { return Err(SubTokenizeError::SplitFuzzyTerm { term }); } @@ -200,7 +199,7 @@ fn rewrite_vec( } fn rewrite_term(s: &str, tokenizer: &T) -> Expr { - let tokens = analyze(s, tokenizer); + let tokens = analyze(tokenizer.tokenize(s)); match tokens.len() { // A term whose analyzed token stream is empty (bare punctuation, // emoji) can never match a stored token: it matches nothing. Failing @@ -286,7 +285,7 @@ fn flush_phrase_term_run( return; } - let text_bytes = terms.iter().map(String::len).sum::() + terms.len() + 1; + let text_bytes = terms.iter().map(String::len).sum::() + terms.len() - 1; let mut text = String::with_capacity(text_bytes); for term in terms.drain(..) { if !text.is_empty() { @@ -294,32 +293,22 @@ fn flush_phrase_term_run( } text.push_str(&term); } - // A one-byte word survives every valid pipeline. Its final position makes - // positions consumed by discarded trailing terms observable without - // widening Tokenizer's streaming interface. - text.push_str(" a"); - - let mut tokens = analyze(&text, tokenizer); - let sentinel = tokens - .pop() - .expect("the phrase position sentinel must survive tokenization"); - debug_assert_eq!(sentinel.text, "a"); - debug_assert!(!sentinel.came_from_split); - + let mut tokens = tokenizer.tokenize(&text); let mut previous_pos = None; - for token in tokens { + for token in tokens.by_ref() { let gap = previous_pos.map_or(token.pos, |previous: u32| { token.pos.saturating_sub(previous).saturating_sub(1) }); *pending_gap = pending_gap.saturating_add(gap); push_pending_gap(out, pending_gap, *has_anchor); previous_pos = Some(token.pos); - out.push(PhraseElement::Term(token.text)); + out.push(PhraseElement::Term(token.text.into_owned())); *has_anchor = true; } - let trailing_gap = previous_pos.map_or(sentinel.pos, |previous| { - sentinel.pos.saturating_sub(previous).saturating_sub(1) + let positions = tokens.positions_consumed(); + let trailing_gap = previous_pos.map_or(positions, |previous| { + positions.saturating_sub(previous).saturating_sub(1) }); *pending_gap = pending_gap.saturating_add(trailing_gap); } @@ -363,7 +352,7 @@ fn rewrite_wildcard( for part in parts { match part { WildcardPart::Literal(s) => { - let tokens = analyze(&s, tokenizer); + let tokens = analyze(tokenizer.tokenize_literal(&s)); if tokens.iter().any(|token| token.came_from_split) { return Err(SubTokenizeError::SplitWildcardLiteral { literal: s }); } @@ -457,7 +446,7 @@ fn normalize_range_bound( match bound { RangeBound::Open => Ok(RangeBound::Open), RangeBound::Term(s) => { - let tokens = analyze(&s, tokenizer); + let tokens = analyze(tokenizer.tokenize_literal(&s)); if tokens.iter().any(|token| token.came_from_split) { return Err(SubTokenizeError::SplitRangeBound { bound: s }); } @@ -867,26 +856,71 @@ mod tests { #[test] fn quoted_phrase_carries_a_trailing_discard_gap_across_alternatives() { - let tokenizer = whitespace_pipeline(LongTokenMode::Discard, 4, PositionGapMode::Preserve); - let phrase = Expr::Phrase { - elements: vec![ - PhraseElement::Term("aa".into()), - PhraseElement::Term("toolong".into()), - PhraseElement::Alternatives(vec![Expr::Term("cc".into())]), - ], - slop: None, - }; - assert_eq!( - sub_tokenize(phrase, &tokenizer).unwrap(), - Expr::Phrase { - elements: vec![ - PhraseElement::Term("aa".into()), - PhraseElement::Gap(1), - PhraseElement::Alternatives(vec![Expr::Term("cc".into())]), - ], - slop: None, + for (text, mode, stemmer, surviving) in [ + ("aa toolong", LongTokenMode::Discard, None, "aa"), + ("toolong", LongTokenMode::Discard, None, ""), + ("aa \u{0301}", LongTokenMode::Split, None, "aa"), + ("\u{0301}", LongTokenMode::Split, None, ""), + ("aa 👨‍👩‍👧‍👦", LongTokenMode::Truncate, None, "aa"), + ( + "abcdefghij \u{0301}", + LongTokenMode::Split, + None, + "abcd efgh ij", + ), + ( + "aa !!!", + LongTokenMode::Split, + Some(tokenizer::Stemmer::Arabic), + "aa", + ), + ( + "!!!", + LongTokenMode::Split, + Some(tokenizer::Stemmer::Arabic), + "", + ), + ] { + for gaps in [PositionGapMode::Preserve, PositionGapMode::Collapse] { + let mut spec = *whitespace_pipeline(mode, 4, gaps).spec(); + spec.case_folding = Folding::Fold; + spec.accent_folding = Folding::Fold; + spec.stemmer = stemmer; + let tokenizer = spec.compile().unwrap(); + for explicit_gap in [0, 2] { + let before = PhraseElement::Alternatives(vec![Expr::Term("zz".into())]); + let after = PhraseElement::Alternatives(vec![Expr::Term("cc".into())]); + let phrase = Expr::Phrase { + elements: vec![ + before.clone(), + PhraseElement::Term(text.into()), + PhraseElement::Gap(explicit_gap), + after.clone(), + ], + slop: None, + }; + let mut elements = vec![before]; + elements.extend( + surviving + .split_whitespace() + .map(|s| PhraseElement::Term(s.into())), + ); + let gap = explicit_gap + u32::from(gaps == PositionGapMode::Preserve); + if gap != 0 { + elements.push(PhraseElement::Gap(gap)); + } + elements.push(after); + assert_eq!( + sub_tokenize(phrase, &tokenizer).unwrap(), + Expr::Phrase { + elements, + slop: None + }, + "{text:?}, {mode:?}, {stemmer:?}, {gaps:?}, gap={explicit_gap}" + ); + } } - ); + } } #[test] @@ -948,4 +982,107 @@ mod tests { assert_eq!(rewritten, Expr::Regex("Beer.*".into())); } + + #[test] + fn stemming_applies_to_terms_and_phrases_but_not_pattern_literals() { + let mut spec = TokenizerPipelineSpec::tin_default(); + spec.stemmer = Some(tokenizer::Stemmer::English); + let pipeline = spec.compile().unwrap(); + for (query, expected) in [ + (Expr::Term("RUNNING".into()), Expr::Term("run".into())), + ( + Expr::Wildcard(vec![ + WildcardPart::Literal("RÚNNING".into()), + WildcardPart::Any, + ]), + Expr::Regex("running.*".into()), + ), + ( + Expr::Fuzzy { + term: "RÚNNING".into(), + prefix: 1, + distance: 2, + }, + Expr::Fuzzy { + term: "running".into(), + prefix: 1, + distance: 2, + }, + ), + ( + Expr::Range { + lower: RangeBound::Term("RÚNNING".into()), + upper: RangeBound::Term("RUNS".into()), + }, + Expr::Range { + lower: RangeBound::Term("running".into()), + upper: RangeBound::Term("runs".into()), + }, + ), + ( + Expr::Phrase { + elements: vec![ + PhraseElement::Term("RUNNING".into()), + PhraseElement::Term("RUNS".into()), + ], + slop: None, + }, + Expr::Phrase { + elements: vec![ + PhraseElement::Term("run".into()), + PhraseElement::Term("run".into()), + ], + slop: None, + }, + ), + ] { + assert_eq!(sub_tokenize(query, &pipeline).unwrap(), expected); + } + } + + #[test] + fn numeric_phrase_terms_survive_every_stemmer() { + for language in [ + "ar", "da", "nl", "en", "fi", "fr", "de", "el", "hu", "it", "no", "pt", "ro", "ru", + "es", "sv", "ta", "tr", + ] { + let mut spec = TokenizerPipelineSpec::tin_default(); + spec.stemmer = Some(language.parse().unwrap()); + let pipeline = spec.compile().unwrap(); + let phrase = Expr::Phrase { + elements: vec![PhraseElement::Term("0".into())], + slop: None, + }; + assert_eq!( + sub_tokenize(phrase.clone(), &pipeline).unwrap(), + phrase, + "{language}" + ); + } + } + + #[test] + fn terms_inside_phrase_alternatives_are_stemmed_once() { + let mut spec = TokenizerPipelineSpec::tin_default(); + spec.stemmer = Some(tokenizer::Stemmer::English); + let pipeline = spec.compile().unwrap(); + // Snowball stems accidental -> accident -> accid across two passes. + let phrase = Expr::Phrase { + elements: vec![ + PhraseElement::Term("accidental".into()), + PhraseElement::Alternatives(vec![Expr::Term("accidental".into())]), + ], + slop: None, + }; + assert_eq!( + sub_tokenize(phrase, &pipeline).unwrap(), + Expr::Phrase { + elements: vec![ + PhraseElement::Term("accident".into()), + PhraseElement::Alternatives(vec![Expr::Term("accident".into())]), + ], + slop: None, + } + ); + } } diff --git a/tokenizer/Cargo.toml b/tokenizer/Cargo.toml index 3e215e1..71f5348 100644 --- a/tokenizer/Cargo.toml +++ b/tokenizer/Cargo.toml @@ -7,6 +7,7 @@ license.workspace = true [dependencies] emojis = "0.6" +rust-stemmers = "=1.2.0" thiserror.workspace = true # tinyvec 1.13.0 fails to compile with unicode-normalization's alloc-only features. tinyvec = "=1.12.0" diff --git a/tokenizer/src/compiled.rs b/tokenizer/src/compiled.rs index 248a555..7cc7107 100644 --- a/tokenizer/src/compiled.rs +++ b/tokenizer/src/compiled.rs @@ -16,11 +16,12 @@ // The full license text is available in LICENSE. use crate::folder::{CompiledFolder, lowercase_ascii_from}; use crate::long_tokens::LongTokenIter; +use crate::normalizer::{CompiledNormalizer, NormalizerSpec}; use crate::spec::{GraphemeMode, TokenizerPipelineSpec, TokenizerSpec}; use crate::tokenizers::{ DiscardGraphemes, EmojiGraphemes, RetainGraphemes, UnicodeIter, WhitespaceIter, }; -use crate::{Classification, Token, Tokenizer}; +use crate::{Classification, Token, TokenStream, Tokenizer}; use unicode_normalization::char::is_combining_mark; use unicode_segmentation::UnicodeSegmentation; @@ -35,7 +36,11 @@ impl CompiledTokenizerPipeline { CompiledPipelineKind::TinDefault } else { let stages = CompiledStages { - folder: CompiledFolder::new(spec.case_folding, spec.accent_folding), + normalizer: NormalizerSpec::new( + spec.case_folding, + spec.accent_folding, + spec.stemmer, + ), long_tokens: spec.long_tokens, position_gaps: spec.position_gaps, }; @@ -68,11 +73,23 @@ impl Tokenizer for CompiledTokenizerPipeline { ) -> Self::Iter<'tokenizer, 'text> { CompiledTokenIter::new(&self.kind, text) } + + fn tokenize_literal<'tokenizer, 'text>( + &'tokenizer self, + text: &'text str, + ) -> Self::Iter<'tokenizer, 'text> { + if self.spec.stemmer.is_none() { + return self.tokenize(text); + } + let mut spec = self.spec; + spec.stemmer = None; + CompiledTokenIter::new(&Self::from_validated_spec(spec).kind, text) + } } #[derive(Clone, Copy)] struct CompiledStages { - folder: CompiledFolder, + normalizer: NormalizerSpec, long_tokens: crate::LongTokenSpec, position_gaps: crate::PositionGapMode, } @@ -85,14 +102,17 @@ enum CompiledPipelineKind { Whitespace(CompiledStages), } -struct FolderIter { +struct NormalizerIter { inner: I, - folder: CompiledFolder, + normalizer: CompiledNormalizer, } -impl FolderIter { - fn new(inner: I, folder: CompiledFolder) -> Self { - Self { inner, folder } +impl NormalizerIter { + fn new(inner: I, normalizer: NormalizerSpec) -> Self { + Self { + inner, + normalizer: normalizer.compile(), + } } fn inner(&self) -> &I { @@ -100,7 +120,7 @@ impl FolderIter { } } -impl<'text, I, const DEFAULT: bool> Iterator for FolderIter +impl<'text, I, const DEFAULT: bool> Iterator for NormalizerIter where I: Iterator>, { @@ -113,7 +133,7 @@ where if DEFAULT { CompiledFolder::apply_case_and_accent(&mut token); } else { - self.folder.apply(&mut token); + self.normalizer.apply(&mut token); } if !token.text.is_empty() { return Some(token); @@ -126,7 +146,7 @@ struct PipelineIter<'text, I, const DEFAULT: bool> where I: Iterator>, { - inner: LongTokenIter<'text, FolderIter, DEFAULT>, + inner: LongTokenIter<'text, NormalizerIter, DEFAULT>, } impl<'text, I, const DEFAULT: bool> PipelineIter<'text, I, DEFAULT> @@ -136,7 +156,7 @@ where fn new(inner: I, stages: CompiledStages) -> Self { Self { inner: LongTokenIter::new( - FolderIter::new(inner, stages.folder), + NormalizerIter::new(inner, stages.normalizer), stages.long_tokens, stages.position_gaps, ), @@ -144,19 +164,11 @@ where } } -impl<'text> PipelineIter<'text, UnicodeIter<'text, EmojiGraphemes>, true> { - /// Positions the exhausted pipeline consumed from its input: the - /// tokenizer's position counter (which already counts tokens the folder - /// later drops, so fold-to-empty gaps are included) plus one extra - /// position per 256-byte split continuation. Only meaningful after the - /// iterator returns `None`. +impl<'text, I: TokenStream<'text>, const DEFAULT: bool> PipelineIter<'text, I, DEFAULT> { + // Read after exhaustion so the base counter includes discarded trailing tokens. fn positions_consumed(&self) -> u32 { self.inner - .inner() - .inner() - .positions_consumed() - .checked_add(self.inner.split_position_offset()) - .expect("token position overflow") + .positions_consumed(self.inner.inner().inner().positions_consumed()) } } @@ -205,7 +217,7 @@ impl<'text> CompiledTokenIter<'text> { const fn default_stages() -> CompiledStages { CompiledStages { - folder: CompiledFolder::CaseAndAccent, + normalizer: NormalizerSpec::Fold(CompiledFolder::CaseAndAccent), long_tokens: crate::LongTokenSpec { mode: crate::LongTokenMode::Split, max_bytes: 256, @@ -720,6 +732,19 @@ impl<'text> Iterator for CompiledTokenIter<'text> { } } +impl<'text> TokenStream<'text> for CompiledTokenIter<'text> { + fn positions_consumed(&self) -> u32 { + match &self.inner { + CompiledTokenIterKind::TinDefaultAscii(iter) => iter.next_position, + CompiledTokenIterKind::TinDefaultMixed(iter) => iter.position_base, + CompiledTokenIterKind::UnicodeDiscard(iter) => iter.positions_consumed(), + CompiledTokenIterKind::UnicodeEmoji(iter) => iter.positions_consumed(), + CompiledTokenIterKind::UnicodeRetain(iter) => iter.positions_consumed(), + CompiledTokenIterKind::Whitespace(iter) => iter.positions_consumed(), + } + } +} + enum CompiledTokenIterKind<'text> { TinDefaultAscii(AsciiDefaultIter<'text>), TinDefaultMixed(MixedDefaultIter<'text>), diff --git a/tokenizer/src/lib.rs b/tokenizer/src/lib.rs index 204e255..49103ac 100644 --- a/tokenizer/src/lib.rs +++ b/tokenizer/src/lib.rs @@ -20,9 +20,11 @@ use std::fmt::Display; mod compiled; mod folder; mod long_tokens; +mod normalizer; pub mod presets; mod source_spans; mod spec; +mod stemmer; pub mod tokenizers; pub use compiled::CompiledTokenizerPipeline; @@ -31,6 +33,7 @@ pub use spec::{ Folding, GraphemeMode, LongTokenMode, LongTokenSpec, MIN_TOKEN_BYTES, PositionGapMode, TokenizerPipelineSpec, TokenizerPipelineSpecError, TokenizerSpec, }; +pub use stemmer::Stemmer; #[derive(Debug, Default, Clone, Copy, PartialEq, Eq)] #[repr(u8)] @@ -88,9 +91,17 @@ impl<'a> Token<'a> { } } +/// A token iterator that accounts for positions consumed by discarded tokens. +pub trait TokenStream<'text>: Iterator> { + /// Final position count, including preserved gaps and split continuations. + /// Call only after `next()` returns `None`; intermediate counts are unspecified. + /// With collapsed gaps, this equals the number of emitted tokens. + fn positions_consumed(&self) -> u32; +} + /// A [`Tokenizer`] converts a string into a stream of [`Token`]s. pub trait Tokenizer { - type Iter<'tokenizer, 'text>: Iterator> + type Iter<'tokenizer, 'text>: TokenStream<'text> where Self: 'tokenizer; @@ -99,6 +110,14 @@ pub trait Tokenizer { &'tokenizer self, text: &'text str, ) -> Self::Iter<'tokenizer, 'text>; + + /// Normalize query pattern literals, preserving their morphology. + fn tokenize_literal<'tokenizer, 'text>( + &'tokenizer self, + text: &'text str, + ) -> Self::Iter<'tokenizer, 'text> { + self.tokenize(text) + } } impl Tokenizer for &T { @@ -113,4 +132,11 @@ impl Tokenizer for &T { ) -> Self::Iter<'tokenizer, 'text> { (*self).tokenize(text) } + + fn tokenize_literal<'tokenizer, 'text>( + &'tokenizer self, + text: &'text str, + ) -> Self::Iter<'tokenizer, 'text> { + (*self).tokenize_literal(text) + } } diff --git a/tokenizer/src/long_tokens.rs b/tokenizer/src/long_tokens.rs index 8affb0f..5b6b0de 100644 --- a/tokenizer/src/long_tokens.rs +++ b/tokenizer/src/long_tokens.rs @@ -38,10 +38,14 @@ impl<'text, I, const DEFAULT: bool> LongTokenIter<'text, I, DEFAULT> where I: Iterator>, { - /// Extra positions created by split continuations: each continuation - /// chunk after the first occupies one additional position. - pub(crate) fn split_position_offset(&self) -> u32 { - self.split_position_offset + pub(crate) fn positions_consumed(&self, base_positions: u32) -> u32 { + if DEFAULT || self.position_gaps == PositionGapMode::Preserve { + base_positions + .checked_add(self.split_position_offset) + .expect("token position overflow") + } else { + self.next_collapsed_pos + } } pub(crate) fn inner(&self) -> &I { diff --git a/tokenizer/src/normalizer.rs b/tokenizer/src/normalizer.rs new file mode 100644 index 0000000..94022d2 --- /dev/null +++ b/tokenizer/src/normalizer.rs @@ -0,0 +1,80 @@ +// Copyright (C) 2026 PlanetScale +// +// This program is free software: you can redistribute it and/or modify +// it under the terms of the GNU Affero General Public License as published by +// the Free Software Foundation, either version 3 of the License, or +// (at your option) any later version. +// +// This program is distributed in the hope that it will be useful, +// but WITHOUT ANY WARRANTY; without even the implied warranty of +// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +// GNU Affero General Public License for more details. +// +// You should have received a copy of the GNU Affero General Public License +// along with this program. If not, see . +// +// The full license text is available in LICENSE. +use crate::folder::CompiledFolder; +use crate::{Folding, Stemmer, Token}; +use std::borrow::Cow; +use unicode_normalization::{UnicodeNormalization, is_nfc}; + +#[derive(Clone, Copy)] +pub(crate) enum NormalizerSpec { + Fold(CompiledFolder), + Stem(Stemmer, Folding), +} + +impl NormalizerSpec { + pub(crate) fn new(case: Folding, accent: Folding, stemmer: Option) -> Self { + match stemmer { + Some(stemmer) => Self::Stem(stemmer, accent), + None => Self::Fold(CompiledFolder::new(case, accent)), + } + } + + pub(crate) fn compile(self) -> CompiledNormalizer { + match self { + Self::Fold(folder) => CompiledNormalizer::Fold(folder), + Self::Stem(stemmer, accent) => CompiledNormalizer::Stem { + stemmer: stemmer.compile(), + accent: CompiledFolder::new(Folding::Preserve, accent), + }, + } + } +} + +pub(crate) enum CompiledNormalizer { + Fold(CompiledFolder), + Stem { + stemmer: rust_stemmers::Stemmer, + accent: CompiledFolder, + }, +} + +impl CompiledNormalizer { + pub(crate) fn apply(&self, token: &mut Token<'_>) { + match self { + Self::Fold(folder) => folder.apply(token), + Self::Stem { stemmer, accent } => { + CompiledFolder::Case.apply(token); + // Snowball suffix rules must see the same accents for canonically + // equivalent words; keep accent removal after stemming. + if !is_nfc(&token.text) { + token.text = Cow::Owned(token.text.nfc().collect()); + } + match &mut token.text { + Cow::Borrowed(text) => token.text = stemmer.stem(text), + Cow::Owned(text) => { + // A borrowed stem is unchanged. Keep the existing + // lowercase allocation instead of copying it again. + if let Cow::Owned(stemmed) = stemmer.stem(text) { + *text = stemmed; + } + } + } + accent.apply(token); + } + } + } +} diff --git a/tokenizer/src/source_spans.rs b/tokenizer/src/source_spans.rs index dfc71f9..3d0ba88 100644 --- a/tokenizer/src/source_spans.rs +++ b/tokenizer/src/source_spans.rs @@ -25,16 +25,17 @@ //! position. //! //! [`SourceSpanIter`] resolves that by walking the same base tokenizer and -//! folding each token only to *drive policy decisions*, while always emitting +//! normalizing each token only to *drive policy decisions*, while always emitting //! borrowed slices of the source. Its contract: one token per position the //! pipeline consumes, in position order, with `pos` equal to that position — //! under `position_gaps = preserve` that includes a whole-token placeholder //! for each gap position (folded-to-empty, discarded, or truncated-to-empty //! tokens), and a split token yields exactly as many slices as the pipeline //! emits chunks, cut at source grapheme boundaries mapped through folding. +//! With stemming, every resulting chunk instead spans the entire original word. -use crate::folder::CompiledFolder; use crate::long_tokens::{split_end, truncate_end}; +use crate::normalizer::{CompiledNormalizer, NormalizerSpec}; use crate::spec::{ GraphemeMode, LongTokenMode, LongTokenSpec, PositionGapMode, TokenizerPipelineSpec, TokenizerSpec, @@ -42,7 +43,7 @@ use crate::spec::{ use crate::tokenizers::{ DiscardGraphemes, EmojiGraphemes, RetainGraphemes, UnicodeIter, WhitespaceIter, }; -use crate::{CompiledTokenizerPipeline, Token, Tokenizer}; +use crate::{CompiledTokenizerPipeline, Token, TokenStream, Tokenizer}; use std::borrow::Cow; use std::collections::VecDeque; use unicode_segmentation::UnicodeSegmentation; @@ -95,7 +96,7 @@ impl<'text> Iterator for BaseIter<'text> { pub struct SourceSpanIter<'text> { base: BaseIter<'text>, - folder: CompiledFolder, + normalizer: CompiledNormalizer, long_tokens: LongTokenSpec, preserve_gaps: bool, pending: VecDeque>, @@ -114,7 +115,8 @@ impl<'text> SourceSpanIter<'text> { }; Self { base, - folder: CompiledFolder::new(spec.case_folding, spec.accent_folding), + normalizer: NormalizerSpec::new(spec.case_folding, spec.accent_folding, spec.stemmer) + .compile(), long_tokens: spec.long_tokens, preserve_gaps: spec.position_gaps == PositionGapMode::Preserve, pending: VecDeque::new(), @@ -131,13 +133,6 @@ impl<'text> SourceSpanIter<'text> { token } - /// The folded byte length of one source slice, without keeping the fold. - fn folded_len(&self, source: &str) -> usize { - let mut probe = Token::new(source, 0); - self.folder.apply(&mut probe); - probe.text.len() - } - /// Queue one source slice per chunk the pipeline emits for this token. /// Chunk-end targets come from running the production splitter over the /// actually-folded text, so the slot count matches the pipeline exactly. @@ -170,13 +165,22 @@ impl<'text> SourceSpanIter<'text> { chunk }; + let CompiledNormalizer::Fold(folder) = &self.normalizer else { + // A stem can change anywhere in the word. Each post-stem chunk + // refers to that entire source word, including removed suffixes. + self.pending.extend(targets.iter().map(|_| chunk(source))); + return; + }; + // Source cluster boundaries with the folded byte length of everything // before them, closed by the end-of-token boundary. let mut folded_prefix = 0usize; let mut boundaries = Vec::new(); for (offset, grapheme) in source.grapheme_indices(true) { boundaries.push((offset, folded_prefix)); - folded_prefix += self.folded_len(grapheme); + let mut probe = Token::new(grapheme, 0); + folder.apply(&mut probe); + folded_prefix += probe.text.len(); } boundaries.push((source.len(), folded_prefix)); @@ -222,7 +226,7 @@ impl<'text> Iterator for SourceSpanIter<'text> { }; let mut folded = Token::with_classification(source, 0, token.classification); - self.folder.apply(&mut folded); + self.normalizer.apply(&mut folded); let consumes_gap_position = if folded.text.is_empty() { true } else if folded.text.len() <= self.long_tokens.max_bytes { @@ -257,6 +261,12 @@ impl<'text> Iterator for SourceSpanIter<'text> { } } +impl<'text> TokenStream<'text> for SourceSpanIter<'text> { + fn positions_consumed(&self) -> u32 { + self.next_pos + } +} + #[cfg(test)] mod tests { use super::*; diff --git a/tokenizer/src/spec.rs b/tokenizer/src/spec.rs index f39dcd8..7391f01 100644 --- a/tokenizer/src/spec.rs +++ b/tokenizer/src/spec.rs @@ -27,6 +27,7 @@ pub struct TokenizerPipelineSpec { pub long_tokens: LongTokenSpec, pub graphemes: GraphemeMode, pub position_gaps: PositionGapMode, + pub stemmer: Option, } impl Default for TokenizerPipelineSpec { @@ -47,6 +48,7 @@ impl TokenizerPipelineSpec { }, graphemes: GraphemeMode::Emoji, position_gaps: PositionGapMode::Preserve, + stemmer: None, } } @@ -64,10 +66,14 @@ impl TokenizerPipelineSpec { }, graphemes: GraphemeMode::Discard, position_gaps: PositionGapMode::Collapse, + stemmer: None, } } pub fn validate(self) -> Result<(), TokenizerPipelineSpecError> { + if self.stemmer.is_some() && self.case_folding == Folding::Preserve { + return Err(TokenizerPipelineSpecError::StemmerRequiresCaseFolding); + } if self.long_tokens.max_bytes < MIN_TOKEN_BYTES { return Err(TokenizerPipelineSpecError::MaxTokenBytesTooSmall); } @@ -120,6 +126,10 @@ pub enum PositionGapMode { #[derive(Debug, thiserror::Error)] pub enum TokenizerPipelineSpecError { + #[error("stemming requires case_folding = fold")] + StemmerRequiresCaseFolding, + #[error("unknown stemmer language code: {0}")] + UnknownStemmer(String), #[error("max_token_bytes must be at least {MIN_TOKEN_BYTES}")] MaxTokenBytesTooSmall, } @@ -142,6 +152,7 @@ mod tests { }, graphemes: GraphemeMode::Emoji, position_gaps: PositionGapMode::Preserve, + stemmer: None, } ); } @@ -170,6 +181,7 @@ mod tests { }, graphemes: GraphemeMode::Discard, position_gaps: PositionGapMode::Collapse, + stemmer: None, } ); } diff --git a/tokenizer/src/stemmer.rs b/tokenizer/src/stemmer.rs new file mode 100644 index 0000000..6ddaf54 --- /dev/null +++ b/tokenizer/src/stemmer.rs @@ -0,0 +1,124 @@ +// Copyright (C) 2026 PlanetScale +// +// This program is free software: you can redistribute it and/or modify +// it under the terms of the GNU Affero General Public License as published by +// the Free Software Foundation, either version 3 of the License, or +// (at your option) any later version. +// +// This program is distributed in the hope that it will be useful, +// but WITHOUT ANY WARRANTY; without even the implied warranty of +// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +// GNU Affero General Public License for more details. +// +// You should have received a copy of the GNU Affero General Public License +// along with this program. If not, see . +// +// The full license text is available in LICENSE. +use crate::TokenizerPipelineSpecError; +use rust_stemmers::Algorithm; +use std::{fmt, str::FromStr}; + +/// A Snowball language selected by its ISO 639-1 code. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum Stemmer { + Arabic, + Danish, + Dutch, + English, + Finnish, + French, + German, + Greek, + Hungarian, + Italian, + Norwegian, + Portuguese, + Romanian, + Russian, + Spanish, + Swedish, + Tamil, + Turkish, +} + +impl Stemmer { + pub const fn as_str(self) -> &'static str { + match self { + Self::Arabic => "ar", + Self::Danish => "da", + Self::Dutch => "nl", + Self::English => "en", + Self::Finnish => "fi", + Self::French => "fr", + Self::German => "de", + Self::Greek => "el", + Self::Hungarian => "hu", + Self::Italian => "it", + Self::Norwegian => "no", + Self::Portuguese => "pt", + Self::Romanian => "ro", + Self::Russian => "ru", + Self::Spanish => "es", + Self::Swedish => "sv", + Self::Tamil => "ta", + Self::Turkish => "tr", + } + } + + pub(crate) fn compile(self) -> rust_stemmers::Stemmer { + rust_stemmers::Stemmer::create(match self { + Self::Arabic => Algorithm::Arabic, + Self::Danish => Algorithm::Danish, + Self::Dutch => Algorithm::Dutch, + Self::English => Algorithm::English, + Self::Finnish => Algorithm::Finnish, + Self::French => Algorithm::French, + Self::German => Algorithm::German, + Self::Greek => Algorithm::Greek, + Self::Hungarian => Algorithm::Hungarian, + Self::Italian => Algorithm::Italian, + Self::Norwegian => Algorithm::Norwegian, + Self::Portuguese => Algorithm::Portuguese, + Self::Romanian => Algorithm::Romanian, + Self::Russian => Algorithm::Russian, + Self::Spanish => Algorithm::Spanish, + Self::Swedish => Algorithm::Swedish, + Self::Tamil => Algorithm::Tamil, + Self::Turkish => Algorithm::Turkish, + }) + } +} + +impl FromStr for Stemmer { + type Err = TokenizerPipelineSpecError; + + fn from_str(code: &str) -> Result { + Ok(match code { + "ar" => Self::Arabic, + "da" => Self::Danish, + "nl" => Self::Dutch, + "en" => Self::English, + "fi" => Self::Finnish, + "fr" => Self::French, + "de" => Self::German, + "el" => Self::Greek, + "hu" => Self::Hungarian, + "it" => Self::Italian, + "no" => Self::Norwegian, + "pt" => Self::Portuguese, + "ro" => Self::Romanian, + "ru" => Self::Russian, + "es" => Self::Spanish, + "sv" => Self::Swedish, + "ta" => Self::Tamil, + "tr" => Self::Turkish, + _ => return Err(TokenizerPipelineSpecError::UnknownStemmer(code.to_owned())), + }) + } +} + +impl fmt::Display for Stemmer { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str(self.as_str()) + } +} diff --git a/tokenizer/src/tokenizers/unicode.rs b/tokenizer/src/tokenizers/unicode.rs index 4231b4c..2c04d0b 100644 --- a/tokenizer/src/tokenizers/unicode.rs +++ b/tokenizer/src/tokenizers/unicode.rs @@ -14,7 +14,7 @@ // along with this program. If not, see . // // The full license text is available in LICENSE. -use crate::{Classification, Token}; +use crate::{Classification, Token, TokenStream}; use std::marker::PhantomData; use unicode_properties::{GeneralCategoryGroup, UnicodeGeneralCategory}; use unicode_segmentation::{GraphemeIndices, UnicodeSegmentation, UnicodeWords}; @@ -137,13 +137,6 @@ impl<'a, P> UnicodeIter<'a, P> { } } - /// Positions handed out so far, including tokens a later stage drops - /// (e.g. a combining-mark run that folds to empty keeps its gap - /// position). Final once the iterator returns `None`. - pub(crate) fn positions_consumed(&self) -> u32 { - self.pos - } - #[inline] fn emit(&mut self, text: &'a str, classification: Classification) -> Token<'a> { let token = Token::with_classification(text, self.pos, classification); @@ -152,6 +145,12 @@ impl<'a, P> UnicodeIter<'a, P> { } } +impl<'a, P: GraphemePolicy> TokenStream<'a> for UnicodeIter<'a, P> { + fn positions_consumed(&self) -> u32 { + self.pos + } +} + impl<'a, P: GraphemePolicy> Iterator for UnicodeIter<'a, P> { type Item = Token<'a>; diff --git a/tokenizer/src/tokenizers/whitespace.rs b/tokenizer/src/tokenizers/whitespace.rs index 45e2513..bd11664 100644 --- a/tokenizer/src/tokenizers/whitespace.rs +++ b/tokenizer/src/tokenizers/whitespace.rs @@ -14,7 +14,7 @@ // along with this program. If not, see . // // The full license text is available in LICENSE. -use crate::{Token, Tokenizer}; +use crate::{Token, TokenStream, Tokenizer}; pub struct Whitespace; @@ -57,6 +57,12 @@ impl<'a> Iterator for WhitespaceIter<'a> { } } +impl<'a> TokenStream<'a> for WhitespaceIter<'a> { + fn positions_consumed(&self) -> u32 { + self.pos + } +} + #[cfg(test)] mod tests { use crate::tokenizers::Whitespace; diff --git a/tokenizer/tests/pipeline_properties.rs b/tokenizer/tests/pipeline_properties.rs new file mode 100644 index 0000000..3a7dbb5 --- /dev/null +++ b/tokenizer/tests/pipeline_properties.rs @@ -0,0 +1,475 @@ +// Copyright (C) 2026 PlanetScale +// +// This program is free software: you can redistribute it and/or modify +// it under the terms of the GNU Affero General Public License as published by +// the Free Software Foundation, either version 3 of the License, or +// (at your option) any later version. +// +// This program is distributed in the hope that it will be useful, +// but WITHOUT ANY WARRANTY; without even the implied warranty of +// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +// GNU Affero General Public License for more details. +// +// You should have received a copy of the GNU Affero General Public License +// along with this program. If not, see . +// +// The full license text is available in LICENSE. +use proptest::prelude::*; +use tokenizer::{ + Classification, Folding, GraphemeMode, LongTokenMode, LongTokenSpec, PositionGapMode, + TokenStream, Tokenizer, TokenizerPipelineSpec, TokenizerSpec, +}; +use unicode_normalization::UnicodeNormalization; +use unicode_normalization::char::is_combining_mark; +use unicode_properties::{GeneralCategoryGroup, UnicodeGeneralCategory}; +use unicode_segmentation::UnicodeSegmentation; + +#[derive(Debug, PartialEq, Eq)] +struct ReferenceToken { + text: String, + pos: u32, + classification: Classification, + split: bool, +} + +fn reference_pipeline(text: &str, spec: TokenizerPipelineSpec) -> (Vec, u32) { + let base = match spec.tokenizer { + TokenizerSpec::Unicode => reference_unicode(text, spec.graphemes), + TokenizerSpec::Whitespace => text + .split_whitespace() + .enumerate() + .map(|(pos, text)| ReferenceToken { + text: text.to_owned(), + pos: pos as u32, + classification: Classification::Unknown, + split: false, + }) + .collect(), + }; + + let base_count = base.len() as u32; + let mut output = Vec::new(); + let mut split_offset = 0u32; + for mut token in base { + token.text = reference_fold(&token.text, spec.case_folding, spec.accent_folding); + if token.text.is_empty() { + continue; + } + let chunks = reference_long_token(&token.text, spec.long_tokens); + let was_split = spec.long_tokens.mode == LongTokenMode::Split && chunks.len() > 1; + let chunk_count = chunks.len(); + for (continuation, chunk) in chunks.into_iter().enumerate() { + let pos = match spec.position_gaps { + PositionGapMode::Collapse => output.len() as u32, + PositionGapMode::Preserve => token.pos + split_offset + continuation as u32, + }; + output.push(ReferenceToken { + text: chunk, + pos, + classification: token.classification, + split: was_split, + }); + } + if was_split { + split_offset += chunk_count as u32 - 1; + } + } + let positions_consumed = match spec.position_gaps { + PositionGapMode::Collapse => output.len() as u32, + PositionGapMode::Preserve => base_count + split_offset, + }; + (output, positions_consumed) +} + +fn reference_unicode(text: &str, graphemes: GraphemeMode) -> Vec { + let mut output = Vec::new(); + for (_, segment) in text.split_word_bound_indices() { + // unicode-segmentation can attach delimiter whitespace to a rare + // combining-mark-only word boundary; production trims that edge, so + // the reference must classify the trimmed segment too. + let segment = if segment.chars().next().is_some_and(char::is_whitespace) { + segment.trim_matches(char::is_whitespace) + } else { + segment + }; + if segment.is_empty() { + continue; + } + if graphemes != GraphemeMode::Discard + && segment.as_bytes().first().is_some_and(u8::is_ascii_digit) + && !segment.is_ascii() + && emojis::get(segment).is_some() + { + push_reference(&mut output, segment, Classification::Emoji); + continue; + } + if segment.chars().any(char::is_alphanumeric) { + push_reference(&mut output, segment, Classification::Word); + continue; + } + if graphemes == GraphemeMode::Discard { + continue; + } + for grapheme in segment.trim_matches(char::is_whitespace).graphemes(true) { + let classification = if emojis::get(grapheme).is_some() { + Some(Classification::Emoji) + } else if graphemes == GraphemeMode::Retain { + reference_retained_grapheme(grapheme) + } else { + None + }; + if let Some(classification) = classification { + push_reference(&mut output, grapheme, classification); + } + } + } + output +} + +fn push_reference(output: &mut Vec, text: &str, classification: Classification) { + output.push(ReferenceToken { + text: text.to_owned(), + pos: output.len() as u32, + classification, + split: false, + }); +} + +fn reference_retained_grapheme(grapheme: &str) -> Option { + let mut mark = false; + let mut symbol = false; + for c in grapheme.chars() { + if matches!(c, '\u{200d}' | '\u{fe0e}' | '\u{fe0f}') { + continue; + } + match c.general_category_group() { + GeneralCategoryGroup::Symbol => symbol = true, + GeneralCategoryGroup::Mark => mark = true, + GeneralCategoryGroup::Letter | GeneralCategoryGroup::Number => {} + GeneralCategoryGroup::Punctuation + | GeneralCategoryGroup::Separator + | GeneralCategoryGroup::Other => return None, + } + } + if symbol { + Some(Classification::Symbol) + } else { + mark.then_some(Classification::Grapheme) + } +} + +fn reference_fold(text: &str, case: Folding, accent: Folding) -> String { + let case_folded = match case { + Folding::Preserve => text.to_owned(), + Folding::Fold => text.chars().flat_map(char::to_lowercase).collect(), + }; + match accent { + Folding::Preserve => case_folded, + Folding::Fold => case_folded + .chars() + .nfd() + .filter(|c| !is_combining_mark(*c)) + .nfc() + .collect(), + } +} + +fn reference_long_token(text: &str, spec: LongTokenSpec) -> Vec { + if text.len() <= spec.max_bytes { + return vec![text.to_owned()]; + } + match spec.mode { + LongTokenMode::Discard => Vec::new(), + LongTokenMode::Truncate => { + let mut output = String::new(); + for grapheme in text.graphemes(true) { + if output.len() + grapheme.len() > spec.max_bytes { + break; + } + output.push_str(grapheme); + } + (!output.is_empty()).then_some(output).into_iter().collect() + } + LongTokenMode::Split => { + let mut chunks = Vec::new(); + let mut chunk = String::new(); + for grapheme in text.graphemes(true) { + if grapheme.len() > spec.max_bytes { + if !chunk.is_empty() { + chunks.push(std::mem::take(&mut chunk)); + } + for scalar in grapheme.chars() { + if !chunk.is_empty() && chunk.len() + scalar.len_utf8() > spec.max_bytes { + chunks.push(std::mem::take(&mut chunk)); + } + chunk.push(scalar); + } + } else { + if !chunk.is_empty() && chunk.len() + grapheme.len() > spec.max_bytes { + chunks.push(std::mem::take(&mut chunk)); + } + chunk.push_str(grapheme); + } + } + if !chunk.is_empty() { + chunks.push(chunk); + } + chunks + } + } +} + +fn targeted_text() -> impl Strategy { + prop::collection::vec( + prop_oneof![ + "[A-Za-z0-9_]{0,16}", + Just("can't 29.3".to_owned()), + Just("CAFÉ Straße İ".to_owned()), + Just("東京 مرحبا क्\u{200d}ष".to_owned()), + Just("a\u{0301}\u{0327}".to_owned()), + Just("\u{0301}\u{0327}".to_owned()), + Just(" \u{0351}\u{034c}\u{0369}\u{0314}\u{0357}\u{0305} ".to_owned()), + Just("😀 👨‍👩‍👧‍👦 👍🏽 🇺🇳 1️⃣".to_owned()), + Just(" +—©™!\t\n".to_owned()), + Just("\u{0000}\u{200d}\u{fe0f}".to_owned()), + Just("x² 3½ a¹b ¼¾".to_owned()), + Just("\u{1100}\u{1161}\u{11a8} 한글".to_owned()), + ], + 0..30, + ) + .prop_map(|parts| parts.concat()) +} + +fn spec_from_indexes( + tokenizer_index: u8, + case_index: u8, + accent_index: u8, + long_index: u8, + max_bytes: usize, + grapheme_index: u8, + position_index: u8, +) -> TokenizerPipelineSpec { + TokenizerPipelineSpec { + stemmer: None, + tokenizer: if tokenizer_index == 0 { + TokenizerSpec::Unicode + } else { + TokenizerSpec::Whitespace + }, + case_folding: if case_index == 0 { + Folding::Preserve + } else { + Folding::Fold + }, + accent_folding: if accent_index == 0 { + Folding::Preserve + } else { + Folding::Fold + }, + long_tokens: LongTokenSpec { + mode: match long_index { + 0 => LongTokenMode::Truncate, + 1 => LongTokenMode::Discard, + _ => LongTokenMode::Split, + }, + max_bytes, + }, + graphemes: match grapheme_index { + 0 => GraphemeMode::Discard, + 1 => GraphemeMode::Emoji, + _ => GraphemeMode::Retain, + }, + position_gaps: if position_index == 0 { + PositionGapMode::Collapse + } else { + PositionGapMode::Preserve + }, + } +} + +proptest! { + #![proptest_config(ProptestConfig::with_cases(500))] + + // The fixed default tuple compiles to its own specialized iterators, so + // it needs the same independent-oracle coverage as the configured + // pipelines: the generated specs below never land on the default tuple + // (their max_bytes stays under 17), which would otherwise leave the + // specialized default path tested only against in-crate references that + // share its folding code. + #[test] + fn default_pipeline_matches_independent_materialized_reference(text in targeted_text()) { + let spec = TokenizerPipelineSpec::tin_default(); + let expected = reference_pipeline(&text, spec); + let pipeline = spec.compile().unwrap(); + let mut stream = pipeline.tokenize(&text); + let actual = stream + .by_ref() + .map(|token| { + let split = token.came_from_split(); + ReferenceToken { + text: token.text.into_owned(), + pos: token.pos, + classification: token.classification, + split, + } + }) + .collect::>(); + prop_assert_eq!((actual, stream.positions_consumed()), expected); + } + + // The source-span view walks the general tokenizer, so its alignment + // with the specialized default pipeline holds only while the two emit + // identical position streams; pin that alignment directly. + #[test] + fn source_spans_stay_position_aligned_with_the_default_pipeline(text in targeted_text()) { + let spec = TokenizerPipelineSpec::tin_default(); + let pipeline = spec.compile().unwrap(); + let source_spans = pipeline.source_spans(); + let mut spans = source_spans.tokenize(&text); + let slots = spans + .by_ref() + .map(|token| { + prop_assert!( + matches!(token.text, std::borrow::Cow::Borrowed(_)), + "source spans must borrow from the input" + ); + Ok(token.text.into_owned()) + }) + .collect::, TestCaseError>>()?; + let mut stream = pipeline.tokenize(&text); + for token in stream.by_ref() { + let slot = slots.get(token.pos as usize); + prop_assert!( + slot.is_some(), + "pipeline emitted position {} but only {} slots exist", + token.pos, + slots.len(), + ); + let folded_slot = + reference_fold(slot.unwrap(), spec.case_folding, spec.accent_folding); + prop_assert!( + folded_slot.contains(token.text.as_ref()), + "slot at position {} folds to {:?}, which does not cover pipeline token {:?}", + token.pos, + folded_slot, + token.text, + ); + } + prop_assert_eq!(spans.positions_consumed() as usize, slots.len()); + prop_assert_eq!(stream.positions_consumed(), spans.positions_consumed()); + } + + // The source-span view must stay position-aligned with the pipeline for + // every configuration: dense slot ordinals, monotonic borrowed spans, and + // for every emitted pipeline token a slot at its position whose folded + // source covers the token's text. + #[test] + fn source_spans_stay_position_aligned_with_the_pipeline( + text in targeted_text(), + tokenizer_index in 0u8..2, + case_index in 0u8..2, + accent_index in 0u8..2, + long_index in 0u8..3, + max_bytes in tokenizer::MIN_TOKEN_BYTES..17, + grapheme_index in 0u8..3, + position_index in 0u8..2, + ) { + let spec = spec_from_indexes( + tokenizer_index, + case_index, + accent_index, + long_index, + max_bytes, + grapheme_index, + position_index, + ); + let pipeline = spec.compile().unwrap(); + let mut stream = pipeline.tokenize(&text); + let tokens = stream + .by_ref() + .map(|token| (token.text.into_owned(), token.pos)) + .collect::>(); + + let base = text.as_ptr() as usize; + let mut spans = Vec::new(); + let source_spans = pipeline.source_spans(); + let mut span_stream = source_spans.tokenize(&text); + for (ordinal, token) in span_stream.by_ref().enumerate() { + prop_assert_eq!(token.pos as usize, ordinal, "slot ordinals must be dense"); + prop_assert!( + matches!(token.text, std::borrow::Cow::Borrowed(_)), + "source spans must borrow from the input" + ); + let slice = match &token.text { + std::borrow::Cow::Borrowed(slice) => *slice, + std::borrow::Cow::Owned(owned) => owned.as_str(), + }; + let start = slice.as_ptr() as usize - base; + prop_assert!(start + slice.len() <= text.len(), "span must stay inside the input"); + // Starts are non-decreasing; slots may overlap only when folding + // fissions one indivisible source cluster into several chunks. + if let Some(&(_, previous_start)) = spans.last() { + prop_assert!(start >= previous_start, "span starts must be monotonic"); + } + spans.push((slice.to_owned(), start)); + } + prop_assert_eq!(span_stream.positions_consumed() as usize, spans.len()); + prop_assert_eq!(stream.positions_consumed(), span_stream.positions_consumed()); + + for (token_text, pos) in &tokens { + prop_assert!( + (*pos as usize) < spans.len(), + "pipeline emitted position {} but only {} slots exist", + pos, + spans.len(), + ); + let (span_text, _) = &spans[*pos as usize]; + let folded_span = reference_fold(span_text, spec.case_folding, spec.accent_folding); + prop_assert!( + folded_span.contains(token_text.as_str()), + "slot at position {} folds to {:?}, which does not cover pipeline token {:?}", + pos, + folded_span, + token_text, + ); + } + } + + #[test] + fn compiled_pipeline_matches_independent_materialized_reference( + text in targeted_text(), + tokenizer_index in 0u8..2, + case_index in 0u8..2, + accent_index in 0u8..2, + long_index in 0u8..3, + max_bytes in tokenizer::MIN_TOKEN_BYTES..17, + grapheme_index in 0u8..3, + position_index in 0u8..2, + ) { + let spec = spec_from_indexes( + tokenizer_index, + case_index, + accent_index, + long_index, + max_bytes, + grapheme_index, + position_index, + ); + let expected = reference_pipeline(&text, spec); + let pipeline = spec.compile().unwrap(); + let mut stream = pipeline.tokenize(&text); + let actual = stream + .by_ref() + .map(|token| { + let split = token.came_from_split(); + ReferenceToken { + text: token.text.into_owned(), + pos: token.pos, + classification: token.classification, + split, + } + }) + .collect::>(); + prop_assert_eq!((actual, stream.positions_consumed()), expected); + } +} diff --git a/tokenizer/tests/stemming.rs b/tokenizer/tests/stemming.rs new file mode 100644 index 0000000..c5f1a59 --- /dev/null +++ b/tokenizer/tests/stemming.rs @@ -0,0 +1,286 @@ +// Copyright (C) 2026 PlanetScale +// +// This program is free software: you can redistribute it and/or modify +// it under the terms of the GNU Affero General Public License as published by +// the Free Software Foundation, either version 3 of the License, or +// (at your option) any later version. +// +// This program is distributed in the hope that it will be useful, +// but WITHOUT ANY WARRANTY; without even the implied warranty of +// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +// GNU Affero General Public License for more details. +// +// You should have received a copy of the GNU Affero General Public License +// along with this program. If not, see . +// +// The full license text is available in LICENSE. +use proptest::prelude::*; +use std::borrow::Cow; +use tokenizer::{ + Folding, LongTokenMode, PositionGapMode, Stemmer, Tokenizer, TokenizerPipelineSpec, + TokenizerPipelineSpecError, +}; +use unicode_normalization::UnicodeNormalization; +use unicode_normalization::char::is_combining_mark; + +const LANGUAGES: [(&str, rust_stemmers::Algorithm); 18] = [ + ("ar", rust_stemmers::Algorithm::Arabic), + ("da", rust_stemmers::Algorithm::Danish), + ("nl", rust_stemmers::Algorithm::Dutch), + ("en", rust_stemmers::Algorithm::English), + ("fi", rust_stemmers::Algorithm::Finnish), + ("fr", rust_stemmers::Algorithm::French), + ("de", rust_stemmers::Algorithm::German), + ("el", rust_stemmers::Algorithm::Greek), + ("hu", rust_stemmers::Algorithm::Hungarian), + ("it", rust_stemmers::Algorithm::Italian), + ("no", rust_stemmers::Algorithm::Norwegian), + ("pt", rust_stemmers::Algorithm::Portuguese), + ("ro", rust_stemmers::Algorithm::Romanian), + ("ru", rust_stemmers::Algorithm::Russian), + ("es", rust_stemmers::Algorithm::Spanish), + ("sv", rust_stemmers::Algorithm::Swedish), + ("ta", rust_stemmers::Algorithm::Tamil), + ("tr", rust_stemmers::Algorithm::Turkish), +]; + +fn spec(stemmer: Stemmer, mode: LongTokenMode, max_bytes: usize) -> TokenizerPipelineSpec { + let mut spec = TokenizerPipelineSpec::tin_default(); + spec.stemmer = Some(stemmer); + spec.long_tokens.mode = mode; + spec.long_tokens.max_bytes = max_bytes; + spec +} + +#[test] +fn language_codes_and_normalization_match_stock_snowball() { + let words = [ + "RUNNING", + "runner", + "généreusement", + "créées", + "canción", + "hablábamos", + "HÄUSERN", + "красивыми", + "kitapları", + "0", + ]; + for (code, algorithm) in LANGUAGES { + let language = code.parse::().unwrap(); + assert_eq!(language.as_str(), code); + assert_eq!(language.to_string(), code); + let stock = rust_stemmers::Stemmer::create(algorithm); + for accent_folding in [Folding::Fold, Folding::Preserve] { + let mut spec = spec(language, LongTokenMode::Split, 256); + spec.accent_folding = accent_folding; + let pipeline = spec.compile().unwrap(); + for word in words { + let lowercase = word.to_lowercase().nfc().collect::(); + let stemmed = stock.stem(&lowercase); + let expected = match accent_folding { + Folding::Preserve => stemmed.into_owned(), + Folding::Fold => stemmed + .nfd() + .filter(|c| !is_combining_mark(*c)) + .nfc() + .collect(), + }; + let actual: Vec<_> = pipeline + .tokenize(word) + .map(|t| t.text.into_owned()) + .collect(); + let expected: Vec<_> = (!expected.is_empty()) + .then_some(expected) + .into_iter() + .collect(); + assert_eq!(actual, expected, "language={code} word={word:?}"); + } + } + } + for code in ["", "english", "EN", "xx"] { + assert!(code.parse::().is_err(), "{code:?}"); + } +} + +#[test] +fn canonical_spellings_share_stems_and_preserve_original_spans() { + for (language, text) in [ + (Stemmer::French, "créées mangées café"), + (Stemmer::Spanish, "hablábamos canción"), + (Stemmer::Turkish, "görüşler kitapları"), + ] { + for accent_folding in [Folding::Fold, Folding::Preserve] { + for max_bytes in [4, 256] { + let mut spec = spec(language, LongTokenMode::Split, max_bytes); + spec.accent_folding = accent_folding; + let pipeline = spec.compile().unwrap(); + let expected: Vec<_> = pipeline.tokenize(text).collect(); + for spelling in [text.to_owned(), text.nfd().collect()] { + let actual: Vec<_> = pipeline.tokenize(&spelling).collect(); + assert_eq!(actual, expected, "{language:?} {accent_folding:?}"); + let spans: Vec<_> = pipeline.source_spans().tokenize(&spelling).collect(); + for token in &actual { + let source = &spans[token.pos as usize]; + assert_eq!(source.pos, token.pos); + assert!(matches!(source.text, Cow::Borrowed(_))); + assert!(spelling.split_whitespace().any(|word| { + word.as_ptr() == source.text.as_ptr() && word.len() == source.text.len() + })); + assert!( + pipeline + .tokenize(&source.text) + .any(|word| word.text == token.text) + ); + } + } + } + } + } +} + +#[test] +fn already_normalized_unicode_keeps_borrowed_text() { + let mut spec = spec(Stemmer::English, LongTokenMode::Split, 256); + spec.accent_folding = Folding::Preserve; + let pipeline = spec.compile().unwrap(); + // The combining mark makes NFC's quick check inconclusive, but no + // composition exists for q + acute and the exact check preserves the borrow. + for word in ["café", "q\u{0301}"] { + let token = pipeline.tokenize(word).next().unwrap(); + assert_eq!(token.text, word); + assert!(matches!(token.text, Cow::Borrowed(_))); + assert_eq!(token.text.as_ptr(), word.as_ptr()); + } +} + +#[test] +fn stemming_rejects_preserved_case_and_defaults_stay_disabled() { + let mut spec = spec(Stemmer::English, LongTokenMode::Split, 256); + spec.case_folding = Folding::Preserve; + assert!(matches!( + spec.validate(), + Err(TokenizerPipelineSpecError::StemmerRequiresCaseFolding) + )); + assert_eq!(TokenizerPipelineSpec::tin_default().stemmer, None); + assert_eq!(TokenizerPipelineSpec::legacy_tin_default().stemmer, None); +} + +#[test] +fn borrowing_and_literal_analysis_keep_their_ownership_and_semantics() { + let pipeline = spec(Stemmer::English, LongTokenMode::Split, 256) + .compile() + .unwrap(); + assert!(matches!( + pipeline.tokenize("run").next().unwrap().text, + Cow::Borrowed("run") + )); + let actual: Vec<_> = pipeline + .tokenize("runs running runner") + .map(|t| t.text.into_owned()) + .collect(); + assert_eq!(actual, ["run", "run", "runner"]); + + // Exercise the &T blanket implementation as well as the concrete method. + fn literals(pipeline: T) -> Vec { + pipeline + .tokenize_literal("RÚNNING runners") + .map(|t| t.text.into_owned()) + .collect() + } + assert_eq!(literals(&pipeline), ["running", "runners"]); +} + +#[test] +fn source_spans_follow_stemming_before_long_token_policy() { + for mode in [ + LongTokenMode::Split, + LongTokenMode::Discard, + LongTokenMode::Truncate, + ] { + for position_gaps in [PositionGapMode::Preserve, PositionGapMode::Collapse] { + let mut spec = spec(Stemmer::English, mode, 4); + spec.position_gaps = position_gaps; + let pipeline = spec.compile().unwrap(); + let text = "RUNNING relational cat"; + let tokens: Vec<_> = pipeline.tokenize(text).collect(); + let spans: Vec<_> = pipeline.source_spans().tokenize(text).collect(); + assert_eq!(tokens[0].text, "run"); + assert_eq!(spans[0].text, "RUNNING"); + for token in tokens { + let span = &spans[token.pos as usize]; + assert_eq!(span.pos, token.pos); + assert!(matches!(span.text, Cow::Borrowed(_))); + let source = if token.text == "run" { + "RUNNING" + } else if token.text == "cat" { + "cat" + } else { + "relational" + }; + assert_eq!(span.text, source, "{mode:?} {position_gaps:?}"); + } + } + } +} + +#[test] +fn empty_normalized_tokens_preserve_or_collapse_positions_with_stemming() { + let word = "\u{0301}"; + for gaps in [PositionGapMode::Preserve, PositionGapMode::Collapse] { + let mut spec = spec(Stemmer::English, LongTokenMode::Split, 256); + spec.tokenizer = tokenizer::TokenizerSpec::Whitespace; + spec.position_gaps = gaps; + let pipeline = spec.compile().unwrap(); + let text = format!("{word} 0"); + let tokens: Vec<_> = pipeline.tokenize(&text).collect(); + let spans: Vec<_> = pipeline.source_spans().tokenize(&text).collect(); + assert_eq!(tokens.len(), 1); + assert_eq!(spans[tokens[0].pos as usize].text, "0"); + } +} + +proptest! { + #![proptest_config(ProptestConfig::with_cases(500))] + + #[test] + fn stemmed_positions_and_full_spans_match_independent_ascii_policy( + words in prop::collection::vec("[A-Za-z]{1,40}", 0..20), + max_bytes in 4usize..32, + mode_index in 0u8..3, + preserve in any::(), + ) { + let mode = [LongTokenMode::Split, LongTokenMode::Truncate, LongTokenMode::Discard][mode_index as usize]; + let mut spec = spec(Stemmer::English, mode, max_bytes); + spec.position_gaps = if preserve { PositionGapMode::Preserve } else { PositionGapMode::Collapse }; + let pipeline = spec.compile().unwrap(); + let stock = rust_stemmers::Stemmer::create(rust_stemmers::Algorithm::English); + let text = words.join(" "); + let spans: Vec<_> = pipeline.source_spans().tokenize(&text).collect(); + let actual: Vec<_> = pipeline.tokenize(&text).map(|token| { + let source = &spans[token.pos as usize]; + (token.text.into_owned(), token.pos, source.text.to_string()) + }).collect(); + let mut expected = Vec::new(); + let mut next_position = 0; + for word in &words { + let lowercase = word.to_lowercase(); + let stemmed = stock.stem(&lowercase); + let chunks: Vec<_> = match mode { + LongTokenMode::Split => stemmed.as_bytes().chunks(max_bytes).map(|chunk| std::str::from_utf8(chunk).unwrap()).collect(), + LongTokenMode::Truncate => vec![&stemmed[..stemmed.len().min(max_bytes)]], + LongTokenMode::Discard if stemmed.len() > max_bytes => vec![], + LongTokenMode::Discard => vec![stemmed.as_ref()], + }; + let emitted = chunks.len(); + for chunk in chunks { + expected.push((chunk.to_owned(), next_position, word.clone())); + next_position += 1; + } + if emitted == 0 && preserve { + next_position += 1; + } + } + prop_assert_eq!(actual, expected); + } +} From a03e6822abffc10b50e3a1ad6230b41cc050ace1 Mon Sep 17 00:00:00 2001 From: Patrick Reynolds Date: Tue, 6 Oct 2026 23:26:33 -0400 Subject: [PATCH 2/2] Support stemming for `body ==> ANY(...)` This commit hooks `get_relation_info_hook` to decorate any string literals on the right side of `==>`, even in `ANY(...)`, with information about the index so stemming matches, `tin.score`, and `tin.highlight` know which index to use. --- README.md | 2 +- postgres/src/analysis.rs | 57 +++++++--- postgres/src/array_analysis.rs | 197 +++++++++++++++++++++++++++++++++ postgres/src/highlight_udfs.rs | 24 +++- postgres/src/lib.rs | 196 ++++++++++++++++++++++++++++++++ postgres/src/operator.rs | 2 +- 6 files changed, 461 insertions(+), 17 deletions(-) create mode 100644 postgres/src/array_analysis.rs diff --git a/README.md b/README.md index 8e22b8c..2ba30da 100644 --- a/README.md +++ b/README.md @@ -26,7 +26,7 @@ Then run `CREATE EXTENSION tin` in the database. Lead loads on demand and does n Lead provides the `tin` access method, the `==>` operator, TINQL parsing, tokenizer and index reloptions (including the Snowball `stemmer` option), and the scoring functions `tin.score`, `tin.full_score`, `tin.max_score`, and `tin.score_inspect`, plus explicit and implicitly bound `tin.highlight` and `tin.highlight_ansi`. Postgres 17 and 18 are build targets. Search results are exact because the access method returns whole-page candidates and Postgres evaluates `==>` against each visible heap tuple, including expression and partial-index rechecks. -The planner binds each `==>` and each `tin.highlight` or `tin.highlight_ansi` call to the analysis of the tin index that covers its column, so a stemmed index matches and highlights inflected words in queries over that column, including joins, partitions, and cached plans. `body ==> ANY(...)` is the exception: it uses the default analysis. `tin.tokenize`, `tin.ql_parse`, `tin.highlight`, and `tin.highlight_ansi` take a trailing `stemmer` argument as in TIN 1.0.4; their 1.0.3 signatures remain as `tin.tokenize_v1_0_3`, `tin.ql_parse_v1_0_3`, `tin.highlight_v1_0_3`, and `tin.highlight_ansi_v1_0_3`. As in TIN, changing `stemmer` on an existing index requires `REINDEX`. +The planner binds each `==>` and each `tin.highlight` or `tin.highlight_ansi` call to the analysis of the tin index that covers its column, so a stemmed index matches and highlights inflected words in queries over that column, including joins, partitions, and cached plans. That includes `body ==> ANY(...)`, whose array elements are bound to the index analysis too; `tin.highlight` marks only the implicit queries of plain `==>` quals, as TIN does, though an `ANY(...)` search binds the analysis of an explicit `query`. `tin.tokenize`, `tin.ql_parse`, `tin.highlight`, and `tin.highlight_ansi` take a trailing `stemmer` argument as in TIN 1.0.4; their 1.0.3 signatures remain as `tin.tokenize_v1_0_3`, `tin.ql_parse_v1_0_3`, `tin.highlight_v1_0_3`, and `tin.highlight_ansi_v1_0_3`. As in TIN, changing `stemmer` on an existing index requires `REINDEX`. Scoring deliberately rescans and retokenizes the visible indexed column or expression, once per statement and search, under that statement's snapshot. A score call must be in the same query level as the matching `==>` predicate, and a partial tin index binds only when the query's quals imply its `WHERE` condition, as in tin. Implicit highlighting binds the same way and, like tin, returns the text unmarked when no index binds a search; passing its `query` argument explicitly works without a bound predicate. diff --git a/postgres/src/analysis.rs b/postgres/src/analysis.rs index 2f41cfb..674a276 100644 --- a/postgres/src/analysis.rs +++ b/postgres/src/analysis.rs @@ -143,9 +143,27 @@ pub(crate) fn pipeline_for(config: &str) -> Result }) } +/// Binds an encoded analysis configuration to one search text. +pub(crate) fn bind_text(analysis: &str, text: &str) -> String { + format!("{TAG}{analysis}{SEPARATOR}{}", raw_query(text)) +} + #[pg_extern(immutable, parallel_safe)] fn bind_query_analysis(query: &str, analysis: &str) -> String { - format!("{TAG}{analysis}{SEPARATOR}{}", raw_query(query)) + bind_text(analysis, query) +} + +/// Binds an encoded analysis configuration to every search text of an array. +pub(crate) fn bind_texts(analysis: &str, queries: Vec>) -> Vec> { + queries + .into_iter() + .map(|text| text.map(|text| bind_text(analysis, &text))) + .collect() +} + +#[pg_extern(immutable, parallel_safe)] +fn bind_query_array_analysis(queries: Vec>, analysis: &str) -> Vec> { + bind_texts(analysis, queries) } pub(crate) unsafe fn sibling_function( @@ -441,6 +459,29 @@ pub(crate) unsafe fn conflict_scope( } } +/// Finds the analysis a search over `operand` has to use when it differs from +/// the default, and makes the plan depend on the indexes that decided it. +/// Raises an error when the members of a partitioned, inherited, or UNION ALL +/// relation analyze the operand differently. +pub(crate) unsafe fn bound_analysis( + root: *mut pg_sys::PlannerInfo, + operand: *mut pg_sys::Node, +) -> Option { + unsafe { + crate::score::single_varno(operand)?; + let (spec, indexes) = match resolve(root, operand) { + Resolution::Unindexed => return None, + Resolution::Conflict => pgrx::error!( + "tin index tokenization differs across the members of \"{}\" for this search", + conflict_scope(root, operand) + ), + Resolution::Bound { spec, indexes } => (spec, indexes), + }; + depend_on_indexes(root, &indexes); + (spec != TokenizerPipelineSpec::tin_default()).then_some(spec) + } +} + fn unhandled() -> Internal { Internal::from(Some(pg_sys::Datum::from(0_usize))) } @@ -475,21 +516,9 @@ fn tin_text_support(request: Internal) -> Internal { { return unhandled(); } - if crate::score::single_varno(left).is_none() { + let Some(spec) = bound_analysis(request.root, left) else { return unhandled(); - } - let (spec, indexes) = match resolve(request.root, left) { - Resolution::Unindexed => return unhandled(), - Resolution::Conflict => pgrx::error!( - "tin index tokenization differs across the members of \"{}\" for this search", - conflict_scope(request.root, left) - ), - Resolution::Bound { spec, indexes } => (spec, indexes), }; - depend_on_indexes(request.root, &indexes); - if spec == TokenizerPipelineSpec::tin_default() { - return unhandled(); - } let bound_right: *mut pg_sys::Node = match text_of(right) { Some(text) => text_const(&tag(&spec, &text)).cast(), None => { diff --git a/postgres/src/array_analysis.rs b/postgres/src/array_analysis.rs new file mode 100644 index 0000000..992a705 --- /dev/null +++ b/postgres/src/array_analysis.rs @@ -0,0 +1,197 @@ +// Copyright (C) 2026 PlanetScale +// +// This program is free software: you can redistribute it and/or modify +// it under the terms of the GNU Affero General Public License as published by +// the Free Software Foundation, either version 3 of the License, or +// (at your option) any later version. +// +// This program is distributed in the hope that it will be useful, +// but WITHOUT ANY WARRANTY; without even the implied warranty of +// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +// GNU Affero General Public License for more details. +// +// You should have received a copy of the GNU Affero General Public License +// along with this program. If not, see . +// +// The full license text is available in LICENSE. +//! Binds index analysis to `==> ANY(...)`. +//! +//! Postgres never asks a support function to simplify a `ScalarArrayOpExpr`, +//! so `tin_text_support` cannot reach `body ==> ANY(ARRAY[...])`. This module +//! covers it from `get_relation_info_hook`. The planner has already +//! preprocessed the query's expressions when it first opens a relation, and +//! has not yet distributed the quals to relations, so the hook can still +//! rewrite the array of every search in place. The hook runs for the first +//! query of a session too: Postgres reads it after opening the relation's +//! indexes, which is when it loads this library. + +use crate::analysis::{bind_texts, bound_analysis, encode, sibling_function}; +use pgrx::{FromDatum, IntoDatum, PgList, pg_guard, pg_sys}; +use std::cell::RefCell; +use std::collections::HashSet; +use std::ffi::{CStr, c_void}; + +static mut PREVIOUS_HOOK: pg_sys::get_relation_info_hook_type = None; + +thread_local! { + static PROCESSED: RefCell> = RefCell::new(HashSet::new()); +} + +/// Forgets a planner root when the memory context that holds it goes away, +/// since a later statement can allocate a new root at the same address. +struct ProcessedRoot(usize); + +impl Drop for ProcessedRoot { + fn drop(&mut self) { + PROCESSED.with_borrow_mut(|processed| processed.remove(&self.0)); + } +} + +pub(crate) fn init() { + unsafe { + PREVIOUS_HOOK = pg_sys::get_relation_info_hook; + pg_sys::get_relation_info_hook = Some(relation_info_hook); + } +} + +#[pg_guard] +unsafe extern "C-unwind" fn relation_info_hook( + root: *mut pg_sys::PlannerInfo, + relation_oid: pg_sys::Oid, + inhparent: bool, + rel: *mut pg_sys::RelOptInfo, +) { + unsafe { + if let Some(previous) = PREVIOUS_HOOK { + previous(root, relation_oid, inhparent, rel); + } + bind_arrays(root); + } +} + +/// Binds the arrays of every `==> ANY(...)` in `root`'s query, once per root. +unsafe fn bind_arrays(root: *mut pg_sys::PlannerInfo) { + if root.is_null() || unsafe { (*root).parse.is_null() } { + return; + } + if !PROCESSED.with_borrow_mut(|processed| processed.insert(root as usize)) { + return; + } + unsafe { + let home = pg_sys::GetMemoryChunkContext(root.cast()); + pgrx::PgMemoryContexts::For(home).leak_and_drop_on_delete(ProcessedRoot(root as usize)); + let operator = crate::operator::search_operator(); + if operator == pg_sys::InvalidOid { + return; + } + let mut context = Context { root, operator }; + // Subqueries that survive pull-up are planned with their own roots. + pg_sys::query_tree_walker_impl((*root).parse, Some(walk), (&raw mut context).cast(), 0); + } +} + +struct Context { + root: *mut pg_sys::PlannerInfo, + operator: pg_sys::Oid, +} + +#[pg_guard] +unsafe extern "C-unwind" fn walk(node: *mut pg_sys::Node, context: *mut c_void) -> bool { + if node.is_null() || unsafe { (*node).type_ } == pg_sys::NodeTag::T_Query { + return false; + } + if unsafe { (*node).type_ } == pg_sys::NodeTag::T_ScalarArrayOpExpr { + unsafe { bind_search(node.cast(), &*context.cast::()) }; + } + unsafe { pg_sys::expression_tree_walker(node, Some(walk), context) } +} + +unsafe fn bind_search(saop: *mut pg_sys::ScalarArrayOpExpr, context: &Context) { + unsafe { + if (*saop).opno != context.operator || pg_sys::list_length((*saop).args) != 2 { + return; + } + let left = pg_sys::list_nth((*saop).args, 0).cast::(); + let array = pg_sys::list_nth((*saop).args, 1).cast::(); + if left.is_null() || array.is_null() || is_bound_array(array) { + return; + } + if (*array).type_ == pg_sys::NodeTag::T_Const + && (*array.cast::()).constisnull + { + return; + } + let Some(spec) = bound_analysis(context.root, left) else { + return; + }; + let analysis = encode(&spec); + let bound = if (*array).type_ == pg_sys::NodeTag::T_Const { + let Some(bound) = bind_constant(array.cast(), &analysis) else { + return; + }; + bound + } else { + pg_sys::set_sa_opfuncid(saop); + bind_expression(array, &analysis, (*saop).opfuncid) + }; + (*pg_sys::list_nth_cell((*saop).args, 1)).ptr_value = bound.cast(); + } +} + +unsafe fn is_bound_array(node: *mut pg_sys::Node) -> bool { + if unsafe { (*node).type_ } != pg_sys::NodeTag::T_FuncExpr { + return false; + } + let function = unsafe { (*node.cast::()).funcid }; + let name = unsafe { pg_sys::get_func_name(function) }; + !name.is_null() && unsafe { CStr::from_ptr(name) }.to_bytes() == b"bind_query_array_analysis" +} + +/// The constant array with every text bound, or `None` when nothing changes. +unsafe fn bind_constant(array: *mut pg_sys::Const, analysis: &str) -> Option<*mut pg_sys::Node> { + unsafe { + let texts = Vec::>::from_datum((*array).constvalue, false)?; + let bound = bind_texts(analysis, texts.clone()); + if bound == texts { + return None; + } + let constant = pg_sys::makeConst( + (*array).consttype, + (*array).consttypmod, + (*array).constcollid, + (*array).constlen, + bound.into_datum()?, + false, + (*array).constbyval, + ); + Some(constant.cast()) + } +} + +/// Calls `bind_query_array_analysis` on an array that has no value until the +/// plan runs, such as a parameter. +unsafe fn bind_expression( + array: *mut pg_sys::Node, + analysis: &str, + operator_function: pg_sys::Oid, +) -> *mut pg_sys::Node { + unsafe { + let function = sibling_function( + operator_function, + "bind_query_array_analysis", + &[pg_sys::TEXTARRAYOID, pg_sys::TEXTOID], + ); + let mut args = PgList::::new(); + args.push(array); + args.push(crate::analysis::text_const(analysis).cast()); + pg_sys::makeFuncExpr( + function, + pg_sys::TEXTARRAYOID, + args.into_pg(), + pg_sys::DEFAULT_COLLATION_OID, + pg_sys::DEFAULT_COLLATION_OID, + pg_sys::CoercionForm::COERCE_EXPLICIT_CALL, + ) + .cast() + } +} diff --git a/postgres/src/highlight_udfs.rs b/postgres/src/highlight_udfs.rs index 17c70fa..f58516d 100644 --- a/postgres/src/highlight_udfs.rs +++ b/postgres/src/highlight_udfs.rs @@ -298,6 +298,10 @@ fn highlight_ansi_bound( struct QueryContext { document: *mut pg_sys::Node, queries: Vec<*mut pg_sys::Node>, + /// Whether a `==> ANY(...)` qual searches the document. Like tin, only + /// `==>` quals supply an implicit query, but any search on the document + /// binds the analysis of the index that covers it. + array_search: bool, } #[pg_guard] @@ -314,6 +318,20 @@ unsafe extern "C-unwind" fn collect_queries(node: *mut pg_sys::Node, context: *m { context.queries.push(unsafe { unbound_query(query) }); } + if unsafe { (*node).type_ } == pg_sys::NodeTag::T_ScalarArrayOpExpr { + let saop = node.cast::(); + if unsafe { (*saop).opno } == crate::operator::search_operator() + && unsafe { pg_sys::list_length((*saop).args) } == 2 + && unsafe { + crate::score::same_operand( + pg_sys::list_nth((*saop).args, 0).cast(), + context.document, + ) + } + { + context.array_search = true; + } + } unsafe { pg_sys::expression_tree_walker( node, @@ -480,13 +498,14 @@ fn highlight_support(request: Internal) -> Internal { let mut binding = QueryContext { document, queries: Vec::new(), + array_search: false, }; // Pulled-up subqueries leave their quals in nested FromExpr nodes. collect_queries( (*(*request.root).parse).jointree.cast::(), (&mut binding as *mut QueryContext).cast(), ); - if binding.queries.is_empty() { + if binding.queries.is_empty() && !binding.array_search { return unhandled(); } crate::analysis::depend_on_indexes(request.root, &indexes); @@ -495,6 +514,9 @@ fn highlight_support(request: Internal) -> Internal { .expect("highlight has a query argument"); let implicit = (*supplied_query).type_ == pg_sys::NodeTag::T_Const && (*supplied_query.cast::()).constisnull; + if implicit && binding.queries.is_empty() { + return unhandled(); + } if implicit && query_needs_other_relations(request.root, request.fcall, document, &binding.queries) { diff --git a/postgres/src/lib.rs b/postgres/src/lib.rs index adc8f4a..7b493d4 100644 --- a/postgres/src/lib.rs +++ b/postgres/src/lib.rs @@ -20,6 +20,7 @@ use pgrx::pg_guard; mod am; mod analysis; +mod array_analysis; mod bm25; mod highlight; mod highlight_udfs; @@ -34,6 +35,7 @@ mod udfs; #[pg_guard] pub extern "C-unwind" fn _PG_init() { options::init(); + array_analysis::init(); } #[cfg(test)] @@ -1229,6 +1231,200 @@ mod tests { assert!(run().is_empty()); } + fn ids(sql: &str) -> Vec { + Spi::get_one::>(sql).unwrap().unwrap_or_default() + } + + #[pg_test] + fn any_searches_bind_the_index_analysis() { + Spi::run( + "CREATE TABLE lite_stem_any (id int, body text); + INSERT INTO lite_stem_any VALUES + (1, 'molded'), (2, 'wines'), (3, 'moldy'), (4, 'beer'); + CREATE INDEX lite_stem_any_idx ON lite_stem_any USING tin (body) + WITH (stemmer = 'en');", + ) + .unwrap(); + for setting in ["on", "off"] { + Spi::run(&format!("SET LOCAL enable_seqscan = {setting}")).unwrap(); + for (array, expected) in [ + ("ARRAY['mold']", vec![1]), + ("ARRAY['mold', NULL, 'wine']", vec![1, 2]), + ("'{mold,beers}'", vec![1, 4]), + ("ARRAY[]::text[]", vec![]), + ] { + assert_eq!( + ids(&format!( + "SELECT array_agg(id ORDER BY id) FROM lite_stem_any + WHERE body ==> ANY({array})" + )), + expected, + "{array}, enable_seqscan = {setting}" + ); + } + } + assert_eq!( + Spi::get_one::( + "SELECT body ==> ANY(ARRAY['mold']) FROM lite_stem_any WHERE id = 1" + ) + .unwrap(), + Some(true) + ); + assert_eq!( + ids("SELECT array_agg(id ORDER BY id) FROM lite_stem_any + WHERE body ==> ANY(ARRAY['mold']) OR body ==> 'wine'"), + vec![1, 2] + ); + assert_eq!( + ids("SELECT array_agg(id ORDER BY id) FROM lite_stem_any + WHERE body ==> ALL(ARRAY['mold', 'molds'])"), + vec![1] + ); + } + + #[pg_test] + fn any_searches_stay_correct_in_cached_plans() { + Spi::run( + "CREATE TABLE lite_stem_any_cached (id int, body text); + INSERT INTO lite_stem_any_cached VALUES (1, 'molded'), (2, 'wines'); + CREATE INDEX lite_stem_any_cached_idx ON lite_stem_any_cached USING tin (body) + WITH (stemmer = 'en'); + PREPARE lite_stem_any_array(text[]) AS + SELECT array_agg(id ORDER BY id) FROM lite_stem_any_cached + WHERE body ==> ANY($1); + PREPARE lite_stem_any_const AS + SELECT array_agg(id ORDER BY id) FROM lite_stem_any_cached + WHERE body ==> ANY(ARRAY['mold']);", + ) + .unwrap(); + for mode in ["force_generic_plan", "force_custom_plan"] { + Spi::run(&format!("SET LOCAL plan_cache_mode = {mode}")).unwrap(); + for _ in 0..7 { + assert_eq!( + ids("EXECUTE lite_stem_any_array(ARRAY['mold', 'wine'])"), + vec![1, 2], + "{mode}" + ); + assert_eq!(ids("EXECUTE lite_stem_any_const"), vec![1], "{mode}"); + } + } + Spi::run( + "ALTER INDEX lite_stem_any_cached_idx RESET (stemmer); + REINDEX INDEX lite_stem_any_cached_idx;", + ) + .unwrap(); + assert!(ids("EXECUTE lite_stem_any_const").is_empty()); + assert!(ids("EXECUTE lite_stem_any_array(ARRAY['mold'])").is_empty()); + } + + #[pg_test] + fn any_searches_bind_expression_indexes() { + Spi::run( + "CREATE TABLE lite_stem_any_expr (id int, body text); + INSERT INTO lite_stem_any_expr VALUES (1, 'Molded'), (2, 'beer'); + CREATE INDEX ON lite_stem_any_expr USING tin ((lower(body))) + WITH (stemmer = 'en');", + ) + .unwrap(); + assert_eq!( + ids("SELECT array_agg(id ORDER BY id) FROM lite_stem_any_expr + WHERE lower(body) ==> ANY(ARRAY['mold'])"), + vec![1] + ); + assert!( + ids("SELECT array_agg(id ORDER BY id) FROM lite_stem_any_expr + WHERE body ==> ANY(ARRAY['molds'])") + .is_empty() + ); + } + + #[pg_test] + fn any_searches_bind_partitioned_indexes() { + Spi::run( + "CREATE TABLE lite_stem_any_parts (part int, id int, body text) + PARTITION BY LIST (part); + CREATE TABLE lite_stem_any_part1 PARTITION OF lite_stem_any_parts + FOR VALUES IN (1); + CREATE TABLE lite_stem_any_part2 PARTITION OF lite_stem_any_parts + FOR VALUES IN (2); + INSERT INTO lite_stem_any_parts VALUES + (1, 4, 'molded'), (2, 5, 'molds'), (2, 6, 'beer'); + CREATE INDEX ON lite_stem_any_parts USING tin (body) WITH (stemmer = 'en');", + ) + .unwrap(); + assert_eq!( + ids("SELECT array_agg(id ORDER BY id) FROM lite_stem_any_parts + WHERE body ==> ANY(ARRAY['mold'])"), + vec![4, 5] + ); + } + + #[pg_test( + error = "tin index tokenization differs across the members of \"lite_stem_any_mixed\" for this search" + )] + fn any_searches_reject_conflicting_partition_analysis() { + Spi::run( + "CREATE TABLE lite_stem_any_mixed (part int, body text) PARTITION BY LIST (part); + CREATE TABLE lite_stem_any_mixed1 PARTITION OF lite_stem_any_mixed FOR VALUES IN (1); + CREATE TABLE lite_stem_any_mixed2 PARTITION OF lite_stem_any_mixed FOR VALUES IN (2); + CREATE INDEX ON lite_stem_any_mixed1 USING tin (body) WITH (stemmer = 'en'); + CREATE INDEX ON lite_stem_any_mixed2 USING tin (body);", + ) + .unwrap(); + Spi::run("SELECT 1 FROM lite_stem_any_mixed WHERE body ==> ANY(ARRAY['mold'])").unwrap(); + } + + #[pg_test] + fn any_searches_bind_join_quals() { + Spi::run( + "CREATE TABLE lite_stem_any_join (id int, body text); + INSERT INTO lite_stem_any_join VALUES (1, 'molded'), (2, 'beer'); + CREATE INDEX ON lite_stem_any_join USING tin (body) WITH (stemmer = 'en'); + CREATE TABLE lite_stem_any_probe (id int, word text); + INSERT INTO lite_stem_any_probe VALUES (1, 'mold'), (2, 'wine');", + ) + .unwrap(); + assert_eq!( + ids("SELECT array_agg(p.id ORDER BY p.id) + FROM lite_stem_any_probe p JOIN lite_stem_any_join j + ON j.body ==> ANY(ARRAY[p.word])"), + vec![1] + ); + assert_eq!( + ids("SELECT array_agg(j.id ORDER BY j.id) + FROM lite_stem_any_probe p LEFT JOIN lite_stem_any_join j + ON j.body ==> ANY(ARRAY['mold']) AND j.id = p.id + WHERE j.id IS NOT NULL"), + vec![1] + ); + } + + #[pg_test] + fn explicit_highlights_bind_the_analysis_of_any_searches() { + Spi::run( + "CREATE TABLE lite_stem_any_highlight (body text); + INSERT INTO lite_stem_any_highlight VALUES ('Running and RUNS'); + CREATE INDEX ON lite_stem_any_highlight USING tin (body) WITH (stemmer = 'en');", + ) + .unwrap(); + assert_eq!( + Spi::get_one::( + "SELECT tin.highlight(body, query => 'runs') FROM lite_stem_any_highlight + WHERE body ==> ANY(ARRAY['run'])" + ) + .unwrap(), + Some("Running and RUNS".into()) + ); + assert_eq!( + Spi::get_one::( + "SELECT tin.highlight(body) FROM lite_stem_any_highlight + WHERE body ==> ANY(ARRAY['run'])" + ) + .unwrap(), + Some("Running and RUNS".into()) + ); + } + #[pg_test(error = "unknown stemmer language code: bogus")] fn unknown_stemmer_languages_are_rejected() { Spi::run( diff --git a/postgres/src/operator.rs b/postgres/src/operator.rs index 90dadc7..4178069 100644 --- a/postgres/src/operator.rs +++ b/postgres/src/operator.rs @@ -33,7 +33,7 @@ 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 { +pub(crate) fn search_operator() -> pg_sys::Oid { unsafe { let mut names = std::ptr::null_mut(); for name in [c"pg_catalog", c"==>"] {