diff --git a/postgres/src/highlight_udfs.rs b/postgres/src/highlight_udfs.rs index fc186e3..5ba285b 100644 --- a/postgres/src/highlight_udfs.rs +++ b/postgres/src/highlight_udfs.rs @@ -15,10 +15,11 @@ // // The full license text is available in LICENSE. 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 std::borrow::Cow; use std::ffi::{CStr, c_void}; -use tinql::runtime::{Query, parse_tinql_to_query_default}; +use tinql::runtime::Query; use tokenizer::presets::default_pipeline; fn render_highlight(text: &str, begin_tag: &str, end_tag: &str, query: Option<&Query>) -> String { diff --git a/postgres/src/lib.rs b/postgres/src/lib.rs index 6cb2a63..e0a5e0b 100644 --- a/postgres/src/lib.rs +++ b/postgres/src/lib.rs @@ -27,6 +27,7 @@ mod operator; pub(crate) mod options; mod score; mod tf_bucket; +mod tinql; mod udfs; #[pg_guard] @@ -1290,6 +1291,155 @@ mod session_tests { } } + #[test] + fn overlong_and_overnested_queries_fail_every_parse_like_tin() { + let mut client = session(); + client + .batch_execute( + "CREATE TEMP TABLE lead_query_limits (id integer, body text); + INSERT INTO lead_query_limits VALUES + (1, 'alpha beta'), (2, 'beta gamma'), (3, 'delta'); + CREATE INDEX lead_query_limits_idx ON lead_query_limits USING tin (body);", + ) + .unwrap(); + // at: 2048 bytes. over: 2049 bytes. wide: 1028 characters, 2050 bytes. + // nested(8) and nested(64) stay within the 64-level depth limit; + // nested(65) is one level past it and well under the byte limit, so + // depth is what rejects it. + let at = format!("{:<2047})", "alpha OR (gamma"); + let over = format!("{:<2048})", "alpha OR (gamma"); + let wide = format!("alpha {}", "é".repeat(1022)); + let nested = |depth: usize| format!("{}delta{}", "(".repeat(depth), ")".repeat(depth)); + let deep = nested(65); + assert_eq!( + ( + at.len(), + over.len(), + wide.chars().count(), + wide.len(), + deep.len() + ), + (2048, 2049, 1028, 2050, 135) + ); + + let search = "SELECT id FROM lead_query_limits WHERE body ==> $1 ORDER BY id"; + let ids = |client: &mut postgres::Client, query: &str| { + client + .query(search, &[&query]) + .unwrap() + .iter() + .map(|row| row.get::<_, i32>(0)) + .collect::>() + }; + assert_eq!(ids(&mut client, &at), [1, 2]); + assert_eq!(ids(&mut client, &nested(8)), [3]); + assert_eq!(ids(&mut client, &nested(64)), [3]); + let row = client + .query_one( + "SELECT tin.ql_parse($1), tin.ql_parse($1, false), + tin.highlight('alpha beta gamma', query => $1), + (SELECT count(*) FROM tin.score_inspect('lead_query_limits_idx', $1, 1.0)), + (SELECT count(tin.score(ctid)) FROM lead_query_limits + WHERE body ==> $1)", + &[&at], + ) + .unwrap(); + assert_eq!( + ( + row.get::<_, String>(0), + row.get::<_, String>(1), + row.get::<_, String>(2), + row.get::<_, i64>(3), + row.get::<_, i64>(4), + ), + ( + "alpha OR gamma".into(), + "OR(alpha, gamma)".into(), + "alpha beta gamma".into(), + 2, + 2 + ) + ); + + // Every row matches the first search, so only the scoring and + // highlighting binds parse the second. + let bound = "FROM lead_query_limits + WHERE body ==> 'alpha OR beta OR delta' OR body ==> $1"; + let entry_points = [ + search.to_owned(), + "SELECT 'alpha beta' ==> $1".to_owned(), + "SELECT tin.ql_parse($1)".to_owned(), + "SELECT tin.ql_parse($1, false)".to_owned(), + "SELECT tin.highlight('alpha beta gamma', query => $1)".to_owned(), + "SELECT * FROM tin.score_inspect('lead_query_limits_idx', $1)".to_owned(), + format!("SELECT tin.score(ctid) {bound}"), + format!("SELECT tin.full_score(ctid) {bound}"), + format!("SELECT tin.max_score(ctid) {bound}"), + format!("SELECT tin.highlight(body) {bound}"), + format!("SELECT tin.highlight_ansi(body) {bound}"), + ]; + let too_long = |bytes: usize| { + ( + "54000".to_owned(), + "tinql query is too long".to_owned(), + Some(format!( + "The query is {bytes} bytes; the limit is 2048 bytes." + )), + ) + }; + let too_deep = ( + "54000".to_owned(), + "tinql query is nested too deeply".to_owned(), + Some("The query nests 65 levels deep; the limit is 64.".to_owned()), + ); + for (query, expected) in [ + (&over, too_long(2049)), + (&wide, too_long(2050)), + (&deep, too_deep), + ] { + for sql in &entry_points { + let error = client.query(sql.as_str(), &[query]).unwrap_err(); + let error = error + .as_db_error() + .unwrap_or_else(|| panic!("{sql}: {error}")); + assert_eq!( + ( + error.code().code().to_owned(), + error.message().to_owned(), + error.detail().map(str::to_owned) + ), + expected, + "{sql}" + ); + } + } + + client + .batch_execute( + "SET plan_cache_mode = force_generic_plan; + PREPARE lead_query_limit_search(text) AS + SELECT id FROM lead_query_limits WHERE body ==> $1 ORDER BY id;", + ) + .unwrap(); + let mut execute = + |query: &str| client.query(&format!("EXECUTE lead_query_limit_search('{query}')"), &[]); + let rows = execute(&at).unwrap(); + assert_eq!( + rows.iter() + .map(|row| row.get::<_, i32>(0)) + .collect::>(), + [1, 2] + ); + for query in [&over, &deep] { + let error = execute(query).unwrap_err(); + assert_eq!( + error.code().map(|code| code.code()), + Some("54000"), + "{error}" + ); + } + } + #[test] fn unproved_partial_indexes_refuse_scoring_like_tin() { let mut client = session(); diff --git a/postgres/src/operator.rs b/postgres/src/operator.rs index aca3727..6f35ee9 100644 --- a/postgres/src/operator.rs +++ b/postgres/src/operator.rs @@ -25,7 +25,7 @@ use tokenizer::Tokenizer; use tokenizer::presets::default_pipeline; fn parse_search(query_text: &str, tokenizer: &T) -> Result { - let parsed = tinql::parse(query_text, tinql::ImplicitOp::And).map_err(|e| e.to_string())?; + let parsed = crate::tinql::parse(query_text).map_err(|e| e.to_string())?; let analyzed = sub_tokenize(parsed, tokenizer).map_err(|e| e.to_string())?; lower(&analyzed).map_err(|e| e.to_string()) } @@ -125,30 +125,32 @@ CREATE OPERATOR CLASS @extschema@.tin_text_ops DEFAULT FOR TYPE pg_catalog.text requires = [amhandler, tin_text_cmpfunc] ); -#[cfg(test)] +#[cfg(feature = "pg_test")] +#[pgrx::pg_schema] mod tests { use super::evaluate_text; + use pgrx::pg_test; - #[test] + #[pg_test] fn boolean_and_positional_queries_are_exact() { assert!(evaluate_text("A craft beer bar", "craft AND beer").unwrap()); assert!(evaluate_text("A craft beer bar", "\"craft beer\"").unwrap()); assert!(!evaluate_text("Beer for craft fans", "\"craft beer\"").unwrap()); } - #[test] + #[pg_test] fn expansions_use_the_document_term_universe() { assert!(evaluate_text("brewhouse", "brew*").unwrap()); assert!(evaluate_text("jalapeno", "jalapeño~1").unwrap()); assert!(!evaluate_text("winery", "brew*").unwrap()); } - #[test] + #[pg_test] fn empty_documents_do_not_match_match_all() { assert!(!evaluate_text("...", "*").unwrap()); } - #[test] + #[pg_test] fn invalid_queries_are_reported() { assert!(evaluate_text("beer", "beer OR").is_err()); } diff --git a/postgres/src/score.rs b/postgres/src/score.rs index 1dd9ecc..7f7d6cf 100644 --- a/postgres/src/score.rs +++ b/postgres/src/score.rs @@ -18,6 +18,7 @@ use crate::bm25::{ Bm25Overrides, DenseRatio, ScoreStopWords, ScoringTermInput, TermScorer, TermSetEdit, compile_scoring_terms, sum_scores_in_order, }; +use crate::tinql::parse_tinql_to_query; use pgrx::iter::TableIterator; use pgrx::{ FromDatum, Internal, IntoDatum, PgBox, PgList, PgMemoryContexts, PgRelation, Spi, default, @@ -25,7 +26,7 @@ use pgrx::{ }; use rustc_hash::FxHashMap; use std::ffi::{CStr, CString, c_void}; -use tinql::runtime::{Query, SpanTermSlot, evaluate, parse_tinql_to_query, tokenize_doc}; +use tinql::runtime::{Query, SpanTermSlot, evaluate, tokenize_doc}; use tokenizer::{CompiledTokenizerPipeline, Tokenizer}; /// Which scorer `score_support` rewrote into `score_bound`. It crosses the SQL diff --git a/postgres/src/tinql.rs b/postgres/src/tinql.rs new file mode 100644 index 0000000..82d2a2b --- /dev/null +++ b/postgres/src/tinql.rs @@ -0,0 +1,81 @@ +// 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. +//! The backend's only entry into TINQL parsing. +//! +//! Every parse of query text in this crate goes through [`parse`], which +//! rejects text that is longer than [`MAX_QUERY_BYTES`] or nested deeper than +//! [`MAX_NESTING_DEPTH`] with SQLSTATE 54000 (`program_limit_exceeded`) before +//! the parser runs. + +use ::tinql::runtime::{Query, QueryError, lower::lower, subtokenize::sub_tokenize}; +use tokenizer::Tokenizer; + +/// Longest accepted TINQL query text, in bytes. +pub(crate) const MAX_QUERY_BYTES: usize = 2 * 1024; + +/// Deepest `(`/`[` nesting the parser will accept. It recurses one stack +/// frame per enclosing bracket, and an unoptimized build overflows an 8 MiB +/// stack near 450 levels; this caps well below that, leaving margin for +/// smaller stacks, and no real query nests anywhere near this deep. The byte +/// cap alone cannot stand in for this: a backend running on a smaller stack +/// could overflow within [`MAX_QUERY_BYTES`]. +pub(crate) const MAX_NESTING_DEPTH: usize = 64; + +/// Parse TINQL text with implicit AND, raising `program_limit_exceeded` when +/// it is longer than [`MAX_QUERY_BYTES`] or nested deeper than +/// [`MAX_NESTING_DEPTH`]. +pub(crate) fn parse(query: &str) -> Result<::tinql::Expr, ::tinql::ParseError> { + if query.len() > MAX_QUERY_BYTES { + pgrx::ereport!( + ERROR, + pgrx::PgSqlErrorCode::ERRCODE_PROGRAM_LIMIT_EXCEEDED, + "tinql query is too long", + format!( + "The query is {} bytes; the limit is {MAX_QUERY_BYTES} bytes.", + query.len() + ), + ); + } + // Nesting depth never exceeds the byte length (each level opens a + // bracket), so a query no longer than the limit cannot reach it; scan only + // the ones that could. + if query.len() > MAX_NESTING_DEPTH { + let depth = ::tinql::max_bracket_depth(query); + if depth > MAX_NESTING_DEPTH { + pgrx::ereport!( + ERROR, + pgrx::PgSqlErrorCode::ERRCODE_PROGRAM_LIMIT_EXCEEDED, + "tinql query is nested too deeply", + format!("The query nests {depth} levels deep; the limit is {MAX_NESTING_DEPTH}."), + ); + } + } + ::tinql::parse(query, ::tinql::ImplicitOp::And) +} + +/// Parse a tinql string, normalize it through the tokenizer, and lower it to a query AST. +pub(crate) fn parse_tinql_to_query( + query: &str, + tokenizer: &T, +) -> 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 2165e65..e337555 100644 --- a/postgres/src/udfs.rs +++ b/postgres/src/udfs.rs @@ -187,8 +187,7 @@ pub fn ql_parse( graphemes, position_gaps, }); - let parsed = - tinql::parse(query, tinql::ImplicitOp::And).unwrap_or_else(|error| pgrx::error!("{error}")); + let parsed = crate::tinql::parse(query).unwrap_or_else(|error| pgrx::error!("{error}")); let analyzed = tinql::runtime::subtokenize::sub_tokenize(parsed, &pipeline) .unwrap_or_else(|error| pgrx::error!("{error}")); if surface { diff --git a/tinql/src/lib.rs b/tinql/src/lib.rs index 0947acd..ef7bf35 100644 --- a/tinql/src/lib.rs +++ b/tinql/src/lib.rs @@ -55,3 +55,81 @@ pub fn parse(input: &str, implicit_op: ImplicitOp) -> Result { } parser::pest_parser::parse(input, implicit_op) } + +/// The deepest `(`/`[` nesting reached in `input`, counting grouping and +/// alternatives but not the atomic contents of a double-quoted phrase (a +/// `\` there escapes the next character). The grammar recurses one level per +/// enclosing bracket, so this bounds the parser's recursion depth from above, +/// and since each level opens a bracket it never exceeds `input.len()`. An +/// unmatched closer floors the running depth at zero rather than going +/// negative, so the result is an upper bound on any valid parse's depth and a +/// cheap pre-parse guard against stack-overflowing deeply nested input. +pub fn max_bracket_depth(input: &str) -> usize { + let mut depth: usize = 0; + let mut max: usize = 0; + let mut in_phrase = false; + let mut chars = input.chars(); + while let Some(c) = chars.next() { + if in_phrase { + match c { + '\\' => { + chars.next(); + } + '"' => in_phrase = false, + _ => {} + } + } else { + match c { + '"' => in_phrase = true, + '(' | '[' => { + depth += 1; + max = max.max(depth); + } + ')' | ']' => depth = depth.saturating_sub(1), + _ => {} + } + } + } + max +} + +#[cfg(test)] +mod depth_tests { + use super::max_bracket_depth; + + #[test] + fn counts_nested_grouping_and_alternatives() { + assert_eq!(max_bracket_depth(""), 0); + assert_eq!(max_bracket_depth("beer"), 0); + assert_eq!(max_bracket_depth("(beer OR wine)"), 1); + assert_eq!(max_bracket_depth("((a))"), 2); + assert_eq!(max_bracket_depth("[a [b [c]]]"), 3); + // Mixed brackets nest together. + assert_eq!(max_bracket_depth("([a])"), 2); + // The deepest point wins, not the last. + assert_eq!(max_bracket_depth("((a)) (b)"), 2); + } + + #[test] + fn ignores_brackets_inside_a_phrase() { + assert_eq!(max_bracket_depth("\"[[[[\""), 0); + // An escaped quote stays inside the phrase; the trailing group counts. + assert_eq!(max_bracket_depth("\"a \\\" b\" (c)"), 1); + // An unterminated phrase swallows the rest, so nothing after counts. + assert_eq!(max_bracket_depth("(a) \"[[["), 1); + } + + #[test] + fn unmatched_closers_floor_at_zero() { + assert_eq!(max_bracket_depth(")))"), 0); + assert_eq!(max_bracket_depth("a) (b"), 1); + assert_eq!(max_bracket_depth("((("), 3); + } + + #[test] + fn depth_never_exceeds_length() { + for q in ["", "beer", "(((x)))", "[a, b, c]", "\"[[[\" (())"] { + assert!(max_bracket_depth(q) <= q.len()); + } + } +}