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..2ba30da 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. 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
new file mode 100644
index 0000000..674a276
--- /dev/null
+++ b/postgres/src/analysis.rs
@@ -0,0 +1,598 @@
+// 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)
+ })
+}
+
+/// 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 {
+ 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(
+ 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()
+ }
+ }
+}
+
+/// 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)))
+}
+
+/// 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();
+ }
+ let Some(spec) = bound_analysis(request.root, left) else {
+ 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/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.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..f58516d 100644
--- a/postgres/src/highlight_udfs.rs
+++ b/postgres/src/highlight_udfs.rs
@@ -14,114 +14,294 @@
// 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(),
))
}
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]
@@ -134,9 +314,23 @@ 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) });
+ }
+ 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(
@@ -151,6 +345,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 +474,145 @@ 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(),
+ array_search: false,
};
// 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() {
+ if binding.queries.is_empty() && !binding.array_search {
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 && binding.queries.is_empty() {
+ return unhandled();
+ }
+ 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 +744,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..7b493d4 100644
--- a/postgres/src/lib.rs
+++ b/postgres/src/lib.rs
@@ -19,6 +19,8 @@ use pgrx::pg_guard;
::pgrx::pg_module_magic!(name);
mod am;
+mod analysis;
+mod array_analysis;
mod bm25;
mod highlight;
mod highlight_udfs;
@@ -33,6 +35,7 @@ mod udfs;
#[pg_guard]
pub extern "C-unwind" fn _PG_init() {
options::init();
+ array_analysis::init();
}
#[cfg(test)]
@@ -1170,6 +1173,325 @@ 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());
+ }
+
+ 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(
+ "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..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"==>"] {
@@ -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);
+ }
+}