// Copyright (c) 2023-2026 ParadeDB, Inc. // // This file is part of ParadeDB - Postgres for Search and Analytics // // 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 . use tests::fixtures::querygen::distinctgen::arb_distinct_mode; use tests::fixtures::querygen::groupbygen::arb_group_by; use tests::fixtures::querygen::joingen::{arb_joins, arb_semi_joins, JoinType}; use tests::fixtures::querygen::numericgen::arb_numeric_expr; use tests::fixtures::querygen::orderbygen::arb_joinscan_order_parts; use tests::fixtures::querygen::pagegen::arb_paging_exprs; use tests::fixtures::querygen::pdbagggen::{arb_pdb_agg_join, arb_pdb_agg_single_table}; use tests::fixtures::querygen::wheregen::arb_wheres; use tests::fixtures::querygen::wheregen::Expr as WhereExpr; use tests::fixtures::querygen::{ arb_joins_and_wheres, compare, compare_on, compare_with_side, generated_queries_setup, Column, IndexExpression, PgGucs, Sides, }; use tests::fixtures::*; use futures::executor::block_on; use lockfree_object_pool::MutexObjectPool; use proptest::prelude::*; use rstest::*; use serde_json::Value; use sqlx::{PgConnection, Row}; const COLUMNS: &[Column] = &[ Column::new("id", "SERIAL8", "'4'") .primary_key() .groupable({ true }), Column::new("uuid", "UUID", "'550e8400-e29b-41d4-a716-446655440000'") .groupable({ true }) .bm25_text_field(r#""uuid": { "tokenizer": { "type": "keyword" }, "fast": true }"#) .random_generator_sql("rpad(lpad((random() * 2147483647)::integer::text, 10, '0'), 32, '0')::uuid"), Column::new("name", "TEXT", "'bob'") .bm25_text_field(r#""name": { "tokenizer": { "type": "keyword" }, "fast": true }"#) .random_generator_sql( "(ARRAY ['alice', 'bob', 'cloe', 'sally', 'brandy', 'brisket', 'anchovy']::text[])[(floor(random() * 7) + 1)::int]" ), Column::new("color", "VARCHAR", "'blue'") .whereable(true) .bm25_text_field(r#""color": { "tokenizer": { "type": "keyword" }, "fast": true }"#) .random_generator_sql( "(ARRAY ['red', 'green', 'blue', 'orange', 'purple', 'pink', 'yellow', NULL]::text[])[(floor(random() * 8) + 1)::int]" ), Column::new("age", "INTEGER", "'20'") .bm25_numeric_field(r#""age": { "fast": true }"#) .random_generator_sql("(floor(random() * 100) + 1)"), Column::new("quantity", "INTEGER", "'7'") .whereable(true) .bm25_numeric_field(r#""quantity": { "fast": true }"#) .random_generator_sql("CASE WHEN random() < 0.1 THEN NULL ELSE (floor(random() * 100) + 1)::int END"), Column::new("price", "NUMERIC(10,2)", "'99.99'") .groupable({ // TODO: Grouping on a float fails to ORDER BY (even in cases without an ORDER BY): // ``` // Cannot ORDER BY OrderByInfo // ``` false }) .bm25_numeric_field(r#""price": { "fast": true }"#) .random_generator_sql("(random() * 1000 + 10)::numeric(10,2)"), // Additional NUMERIC columns for testing Numeric64 vs NumericBytes storage Column::new("small_numeric", "NUMERIC(5,2)", "'12.34'") .groupable(false) .bm25_numeric_field(r#""small_numeric": { "fast": true }"#) .random_generator_sql("(random() * 100)::numeric(5,2)"), Column::new("int_numeric", "NUMERIC(10,0)", "'12345'") .groupable(false) .bm25_numeric_field(r#""int_numeric": { "fast": true }"#) .random_generator_sql("(floor(random() * 1000000))::numeric(10,0)"), Column::new("high_scale", "NUMERIC(18,6)", "'123.456789'") .groupable(false) .bm25_numeric_field(r#""high_scale": { "fast": true }"#) .random_generator_sql("(random() * 10000)::numeric(18,6)"), Column::new("big_numeric", "NUMERIC", "'12345.67890'") .groupable(false) // Cannot aggregate NumericBytes .bm25_numeric_field(r#""big_numeric": { "fast": true }"#) .random_generator_sql("(random() * 100000)::numeric"), Column::new("rating", "INTEGER", "'4'") .indexed({ // Marked un-indexed in order to test heap-filter pushdown. false }) .groupable({ true }) .bm25_numeric_field(r#""rating": { "fast": true }"#) .random_generator_sql("(floor(random() * 5) + 1)::int"), Column::new("category", "TEXT", "'electronics'") .whereable(false) .bm25_v2_expression(IndexExpression::Upper) .random_generator_sql( "(ARRAY ['electronics', 'clothing', 'food', 'books', 'toys', 'sports', 'home']::text[])[(floor(random() * 7) + 1)::int]" ), Column::new("literal_normalized", "TEXT", "'Hello World'") .whereable({ // literal_normalized lowercases text, so BM25 @@@ would match case-insensitively // while PostgreSQL = does exact matching. This causes test failures when comparing // results, so we exclude it from WHERE clause testing. false }) .groupable(false) .bm25_v2_expression(IndexExpression::LiteralNormalized) .random_generator_sql( "(ARRAY ['Hello World', 'HELLO WORLD', 'hello world', 'HeLLo WoRLD', 'GOODBYE WORLD', 'goodbye world']::text[])[(floor(random() * 6) + 1)::int]" ), Column::new("metadata", "JSONB", "'{\"brand\": \"apple\", \"rating\": 4}'") .whereable(false) .groupable(false) .bm25_json_field(r#""metadata": { "fast": true }"#) .random_generator_sql( "jsonb_build_object( 'brand', (ARRAY ['apple', 'samsung', 'sony', 'lg']::text[])[(floor(random() * 4) + 1)::int], 'rating', (floor(random() * 5) + 1)::int )" ), Column::new("tags", "TEXT[]", "ARRAY['alpha', 'beta']::text[]") .whereable(false) .groupable(false) .bm25_text_field(r#""tags": { "tokenizer": { "type": "keyword" }, "fast": true }"#) .random_generator_sql( "(CASE (floor(random() * 5) + 1)::int \ WHEN 1 THEN ARRAY['alpha', 'beta']::text[] \ WHEN 2 THEN ARRAY['gamma']::text[] \ WHEN 3 THEN ARRAY['delta', 'epsilon', 'zeta']::text[] \ WHEN 4 THEN NULL \ ELSE ARRAY[]::text[] \ END)", ), ]; fn columns_named(names: Vec<&'static str>) -> Vec { COLUMNS .iter() .filter(|c| names.contains(&c.name)) .cloned() .collect() } #[derive(Clone, Copy, Debug)] enum SubqueryKind { Exists, In, } #[derive(Clone, Copy, Debug)] enum SubqueryPolarity { Positive, Negated, } #[derive(Clone, Debug)] struct GeneratedSubquery { kind: SubqueryKind, polarity: SubqueryPolarity, column: &'static str, inner_where_expr: WhereExpr, paging_exprs: String, } impl GeneratedSubquery { fn to_sql(&self, op: &str, outer_table_name: &str, inner_table_name: &str) -> String { let base = match self.kind { SubqueryKind::Exists => format!( "EXISTS (\ SELECT 1 FROM {inner_table_name} \ WHERE {inner_table_name}.{column} = {outer_table_name}.{column} \ AND {} {}\ )", self.inner_where_expr.to_sql(op), self.paging_exprs, column = self.column, ), SubqueryKind::In => format!( "{outer_table_name}.{column} IN (\ SELECT {column} FROM {inner_table_name} WHERE {} {}\ )", self.inner_where_expr.to_sql(op), self.paging_exprs, column = self.column, ), }; match (self.kind, self.polarity) { (_, SubqueryPolarity::Positive) => base, (SubqueryKind::Exists, SubqueryPolarity::Negated) => format!("NOT {base}"), (SubqueryKind::In, SubqueryPolarity::Negated) => format!("NOT ({base})"), } } } /// /// Tests all join configurations against small tables (important for joins that produce /// cartesian products or expansive results). /// /// When the generated query is within the subset guaranteed to be plannable by ParadeDB's /// custom scan (INNER/OUTER joins without Cartesian products, LIMIT present, and GUC enabled), /// this test explicitly verifies via EXPLAIN that PostgreSQL selected ParadeDB Join Scan /// (or Aggregate Scan). For all queries, it verifies exact result correctness and parity /// against PostgreSQL. /// #[rstest] #[tokio::test] async fn generated_joins_small(database: Db) { let pool = MutexObjectPool::::new( move || { block_on(async { { database.connection().await } }) }, |_| {}, ); let tables_and_sizes = [("users", 10), ("products", 10), ("orders", 10)]; let tables = tables_and_sizes .iter() .map(|(table, _)| table) .collect::>(); let setup_sql = generated_queries_setup(&mut pool.pull(), &tables_and_sizes, COLUMNS); let where_and_join_columns = columns_named(vec!["id", "name", "color", "age", "uuid", "tags"]); proptest!(|( ((join, where_expr, cross_rel), distinct_mode, mut order_parts) in arb_joins_and_wheres( any::(), tables, &where_and_join_columns, ).prop_flat_map(|(join, where_expr, cross_rel)| { let used_tables = join .used_tables() .into_iter() .map(str::to_string) .collect::>(); ( Just((join, where_expr, cross_rel)), arb_distinct_mode(used_tables.clone(), COLUMNS), arb_joinscan_order_parts(used_tables, false), ) }), limit in proptest::option::of(1..=50usize), gucs in any::(), )| { let join_clause = join.to_sql(); let used_tables = join.used_tables(); let mut target_cols = vec![ format!("{}.id", used_tables[0]), format!("{}.name", used_tables[0]), ]; if distinct_mode.is_distinct() { for table in &used_tables[1..] { let col = format!("{}.id", table); if !target_cols.contains(&col) { target_cols.push(col); } } if let Some(expr) = distinct_mode.expression() { let expr_sql = expr.to_sql(); if !target_cols.contains(&expr_sql) { target_cols.push(expr_sql); } } // Ensure every expression in order_parts appears in SELECT DISTINCT // to satisfy PostgreSQL's requirement that ORDER BY expressions must appear in the SELECT list. for part in &order_parts { if !target_cols.contains(part) { target_cols.push(part.clone()); } } } for alias in join.unnest_aliases() { target_cols.push(alias.to_string()); } let target_cols_str = target_cols.join(", "); let distinct_kw = if distinct_mode.is_distinct() { "DISTINCT " } else { "" }; let from = format!("SELECT {distinct_kw}{target_cols_str} {join_clause}"); let (where_pg, where_bm25) = match &cross_rel { Some(cr) => ( format!("({}) AND ({})", where_expr.to_sql(" = "), cr.to_sql()), format!("({}) AND ({})", where_expr.to_sql("@@@"), cr.to_sql()), ), None => (where_expr.to_sql(" = "), where_expr.to_sql("@@@")), }; for alias in join.unnest_aliases() { order_parts.push(alias.to_string()); } let order_by = order_parts.join(", "); let limit_clause = match limit { Some(l) => format!("LIMIT {l}"), None => "".to_string(), }; let pg_query = format!("{from} WHERE {where_pg} ORDER BY {order_by} {limit_clause}"); let bm25_query = format!("{from} WHERE {where_bm25} ORDER BY {order_by} {limit_clause}"); // Assert that JoinScan or AggregateScan was actually used whenever the query is within // the subset guaranteed to be plannable by ParadeDB's custom scan. let is_supported_join = (|| { // DISTINCT expressions (e.g. `col * 10`) fall back to PostgreSQL because // LIMIT cannot be pushed down below upper deduplication. if distinct_mode.expression().is_some() { return false; } // Cross joins (Cartesian products) intentionally fall back to PostgreSQL. if !join.has_no_cross() || cross_rel.is_some() { return false; } // Multi-unnest queries: PostgreSQL's optimizer can pair unnest function scans together // into intermediate unnest sub-joins (e.g. `users_tags JOIN products_tags`), which // JoinScan cannot absorb because neither side is a base table provider containing the array source. // TODO: https://github.com/paradedb/paradedb/pull/6239 if join.unnest_aliases().len() > 1 { return false; } let has_null_ordering = order_parts .iter() .any(|p| p.contains("IS NULL") || p.contains("IS NOT NULL")); // When DISTINCT is active, ORDER BY expressions are projected into SELECT DISTINCT. // Null-predicate ordering (`col IS (NOT) NULL`) alongside base column `col` creates // derived expressions in DISTINCT that cannot push down LIMIT. if distinct_mode.is_distinct() && has_null_ordering { return false; } // INNER joins support arbitrary WHERE expressions, cross-relation predicates, and expression ORDER BY. if join.has_only_inner() { return true; } // Outer joins (LEFT, RIGHT, FULL): cross-table OR predicates cannot be pushed down. if where_expr.has_cross_table_or() { return false; } // Outer joins: expression or NULL-predicate ORDER BY (`upper()`, `IS NULL`) cannot be // guaranteed across outer join boundaries. let has_expr_ordering = has_null_ordering || order_parts.iter().any(|p| p.contains("upper(")); if has_expr_ordering { return false; } // Outer joins: null-testing predicates (`IS NULL`, `IS NOT NULL`) in WHERE can cause // PostgreSQL to simplify outer joins into anti-joins. If subsequent joins reference // columns from the pruned relation, JoinScan intentionally declines because the join keys // cannot be resolved to output-visible equivalents. // TODO: https://github.com/paradedb/paradedb/pull/6239 if where_expr.has_null_predicate() { return false; } true })(); let expect_custom_scan = gucs.join_custom_scan && limit.is_some() && is_supported_join; if expect_custom_scan { let conn = &mut pool.pull(); gucs.set().execute(conn); let explain_query = format!("EXPLAIN (FORMAT JSON) {bm25_query}"); let (plan,): (Value,) = explain_query.fetch_one(conn); let plan_str = format!("{plan:#?}"); prop_assert!( plan_str.contains("ParadeDB Join Scan") || plan_str.contains("ParadeDB Aggregate Scan"), "Query should use ParadeDB Join Scan or Aggregate Scan but got plan: {plan_str}\nQuery: {bm25_query}", ); } compare( &pg_query, &bm25_query, &gucs, &mut pool.pull(), &setup_sql, |query, conn| { "SET work_mem TO '16MB';".execute(conn); let rows = query.fetch_dynamic(conn); let mut row_strings: Vec = rows .into_iter() .map(|row| { use sqlx::Row; let id: i64 = row.try_get(0).unwrap_or(0); format!("{:020}|{:?}", id, row) }) .collect(); row_strings.sort(); row_strings }, )?; }); } #[rstest] #[tokio::test] async fn generated_single_relation(database: Db) { let pool = MutexObjectPool::::new( move || { block_on(async { { database.connection().await } }) }, |_| {}, ); let table_name = "users"; let setup_sql = generated_queries_setup(&mut pool.pull(), &[(table_name, 10)], COLUMNS); proptest!(|( where_expr in arb_wheres( vec![table_name], COLUMNS, ), gucs in any::(), target in prop_oneof![Just("COUNT(*)"), Just("id")], )| { compare( &format!("SELECT {target} FROM {table_name} WHERE {}", where_expr.to_sql(" = ")), &format!("SELECT {target} FROM {table_name} WHERE {}", where_expr.to_sql("@@@")), &gucs, &mut pool.pull(), &setup_sql, |query, conn| { let mut rows = query.fetch::<(i64,)>(conn); rows.sort(); rows } )?; }); } /// /// Property test for GROUP BY aggregates with ORDER BY and LIMIT/OFFSET /// - ensures equivalence between PostgreSQL and bm25 behavior /// #[rstest] #[tokio::test] async fn generated_group_by_aggregates(database: Db) { let pool = MutexObjectPool::::new( move || { block_on(async { { database.connection().await } }) }, |_| {}, ); let table_name = "users"; let setup_sql = generated_queries_setup(&mut pool.pull(), &[(table_name, 50)], COLUMNS); // Columns that can be used for grouping (must have fast: true in index) let columns: Vec<_> = COLUMNS .iter() .filter(|col| col.is_groupable && col.is_whereable) .cloned() .collect(); let grouping_columns: Vec<_> = columns.iter().map(|col| col.name).collect(); proptest!(|( text_where_expr in arb_wheres( vec![table_name], &columns, ), numeric_where_expr in arb_wheres( vec![table_name], &columns_named(vec!["age", "price", "rating"]), ), group_by_expr in arb_group_by(grouping_columns.to_vec(), vec!["COUNT(*)", "SUM(price)", "AVG(price)", "MIN(rating)", "MAX(rating)", "SUM(age)", "AVG(age)"]), limit in prop::option::of(5..21_usize), offset in prop::option::of(0..4_usize), gucs in any::(), )| { let select_list = group_by_expr.to_select_list(); let group_by_clause = group_by_expr.to_sql(); let order_by_and_offset_clause: String = if group_by_expr.group_by_columns.is_empty() { String::new() } else { let order_by_items: Vec = group_by_expr.group_by_columns .iter() .map(|item| {format!("{item} ASC NULLS LAST")}) .collect(); // only apply OFFSET when there are grouping columns, otherwise // we'd offset the single aggregate row let offset_clause = offset .map(|value| format!(" OFFSET {value}")) .unwrap_or_default(); format!("ORDER BY {}{offset_clause}", order_by_items.join(", ")) }; let limit_clause = limit .map(|value| format!(" LIMIT {value}")) .unwrap_or_default(); // Create combined WHERE clause for PostgreSQL using = operator let pg_where_clause = format!( "({}) AND ({})", text_where_expr.to_sql(" = "), numeric_where_expr.to_sql(" < ") ); // Create combined WHERE clause for BM25 using appropriate operators let bm25_where_clause = format!( "({}) AND ({})", text_where_expr.to_sql("@@@"), numeric_where_expr.to_sql(" < ") ); let pg_query = format!( "SELECT {select_list} FROM {table_name} WHERE {pg_where_clause} {group_by_clause} {order_by_and_offset_clause}{limit_clause}", ); let bm25_query = format!( "SELECT {select_list} FROM {table_name} WHERE {bm25_where_clause} {group_by_clause} {order_by_and_offset_clause}{limit_clause}", ); // Custom result comparator for GROUP BY results let compare_results = |query: &str, conn: &mut PgConnection| -> Vec { // Fetch all rows as dynamic results and convert to string representation let rows = query.fetch_dynamic(conn); let string_rows: Vec = rows .into_iter() .map(|row| { // Convert entire row to a string representation for comparison let mut row_string = String::new(); for i in 0..row.len() { if i > 0 { row_string.push('|'); } // Try to get value as different types, converting to string let value_str = if let Ok(val) = row.try_get::(i) { val.to_string() } else if let Ok(val) = row.try_get::(i) { val.to_string() } else if let Ok(val) = row.try_get::(i) { val } else { "NULL".to_string() }; row_string.push_str(&value_str); } row_string }) .collect(); string_rows }; compare(&pg_query, &bm25_query, &gucs, &mut pool.pull(), &setup_sql, compare_results)?; }); } #[rstest] #[tokio::test] async fn generated_paging_small(database: Db) { let pool = MutexObjectPool::::new( move || { block_on(async { { database.connection().await } }) }, |_| {}, ); let table_name = "users"; let setup_sql = generated_queries_setup(&mut pool.pull(), &[(table_name, 1000)], COLUMNS); proptest!(|( where_expr in arb_wheres(vec![table_name], &columns_named(vec!["name"])), paging_exprs in arb_paging_exprs(table_name, vec!["name", "color", "age", "quantity"], vec!["id", "uuid"]), gucs in any::(), )| { compare( &format!("SELECT id FROM {table_name} WHERE {} {paging_exprs}", where_expr.to_sql(" = ")), &format!("SELECT id FROM {table_name} WHERE {} {paging_exprs}", where_expr.to_sql("@@@")), &gucs, &mut pool.pull(), &setup_sql, |query, conn| query.fetch::<(i64,)>(conn), )?; }); } /// Generates paging expressions on a large table, which was necessary to reproduce /// https://github.com/paradedb/tantivy/pull/51 /// /// TODO: Explore whether this could use https://github.com/paradedb/paradedb/pull/2681 /// to use a large segment count rather than a large table size. #[rstest] #[tokio::test] async fn generated_paging_large(database: Db) { let pool = MutexObjectPool::::new( move || { block_on(async { { database.connection().await } }) }, |_| {}, ); let table_name = "users"; let setup_sql = generated_queries_setup(&mut pool.pull(), &[(table_name, 100000)], COLUMNS); proptest!(|( paging_exprs in arb_paging_exprs(table_name, vec![], vec!["uuid"]), gucs in any::(), )| { compare( &format!("SELECT uuid::text FROM {table_name} WHERE name = 'bob' {paging_exprs}"), &format!("SELECT uuid::text FROM {table_name} WHERE name @@@ 'bob' {paging_exprs}"), &gucs, &mut pool.pull(), &setup_sql, |query, conn| query.fetch::<(String,)>(conn), )?; }); } #[rstest] #[tokio::test] async fn generated_subquery(database: Db) { let pool = MutexObjectPool::::new( move || { block_on(async { { database.connection().await } }) }, |_| {}, ); let outer_table_name = "products"; let inner_table_name = "orders"; let setup_sql = generated_queries_setup( &mut pool.pull(), &[(outer_table_name, 10), (inner_table_name, 10)], COLUMNS, ); proptest!(|( outer_where_expr in arb_wheres( vec![outer_table_name], COLUMNS, ), inner_where_expr in arb_wheres( vec![inner_table_name], COLUMNS, ), subquery_column in proptest::sample::select(&["name", "color", "age"]), subquery_kind in prop_oneof![Just(SubqueryKind::Exists), Just(SubqueryKind::In)], subquery_polarity in prop_oneof![ Just(SubqueryPolarity::Positive), Just(SubqueryPolarity::Negated), ], paging_exprs in arb_paging_exprs(inner_table_name, vec!["name", "color", "age"], vec!["id", "uuid"]), gucs in any::(), )| { let subquery = GeneratedSubquery { kind: subquery_kind, polarity: subquery_polarity, column: subquery_column, inner_where_expr, paging_exprs, }; let pg = format!( "SELECT COUNT(*) FROM {outer_table_name} \ WHERE {} AND {}", subquery.to_sql(" = ", outer_table_name, inner_table_name), outer_where_expr.to_sql(" = "), ); let bm25 = format!( "SELECT COUNT(*) FROM {outer_table_name} \ WHERE {} AND {}", subquery.to_sql("@@@", outer_table_name, inner_table_name), outer_where_expr.to_sql("@@@"), ); compare( &pg, &bm25, &gucs, &mut pool.pull(), &setup_sql, |query, conn| query.fetch_one::<(i64,)>(conn), )?; }); } /// Property test for aggregate-on-join via DataFusion — ensures equivalence between /// PostgreSQL native aggregation and ParadeDB's DataFusion aggregate backend. /// /// This test randomly combines: /// - 2 or 3 table INNER joins /// - BM25 predicates (@@@ on outer table) /// - GROUP BY with 0-2 grouping columns /// - Aggregate functions: COUNT(*), SUM, AVG, MIN, MAX /// /// Verifies that the DataFusion aggregate path produces the same results as /// PostgreSQL's native hash/sort aggregate on top of nested loop joins. #[rstest] #[tokio::test] async fn generated_aggregate_join(database: Db) { let pool = MutexObjectPool::::new( move || { block_on(async { { database.connection().await } }) }, |_| {}, ); // Three tables for 2-way and 3-way join testing let tables_and_sizes = [("users", 50), ("products", 50), ("orders", 50)]; let all_tables: Vec<&str> = tables_and_sizes.iter().map(|(table, _)| *table).collect(); let setup_sql = generated_queries_setup(&mut pool.pull(), &tables_and_sizes, COLUMNS); // Text columns for BM25 WHERE clauses let text_columns = columns_named(vec!["name"]); // Columns for join keys let join_key_columns = columns_named(vec!["id", "age"]); // Columns for GROUP BY (must be fast fields) let mut grouping_columns: Vec = COLUMNS .iter() .filter(|col| col.is_groupable && col.is_whereable) .map(|col| format!("{}.{}", all_tables[0], col.name)) .collect(); grouping_columns.extend(vec![ // Group by text representation format!("{}.metadata->>'brand'", all_tables[0]), // Group by JSONB representation format!("{}.metadata->'brand'", all_tables[0]), ]); proptest!(|( num_tables in 2..=3usize, // Outer table BM25 predicate outer_bm25 in arb_wheres(vec![all_tables[0]], &text_columns), // GROUP BY + aggregates group_by_expr in arb_group_by( grouping_columns.clone(), vec![ "COUNT(*)", "SUM(users.age)", "AVG(users.age)", "MIN(users.rating)", "MAX(users.rating)", ], ), mut gucs in any::(), )| { // Build join with selected number of tables let tables_for_join: Vec<&str> = all_tables[..num_tables].to_vec(); // Generate join expression (include LEFT/FULL to cover outer-join aggregate paths) let join = arb_joins( prop_oneof![Just(JoinType::Inner), Just(JoinType::Left), Just(JoinType::Full)], tables_for_join.clone(), &join_key_columns, ); let join_expr = { use proptest::strategy::ValueTree; use proptest::test_runner::TestRunner; let mut runner = TestRunner::default(); join.new_tree(&mut runner).unwrap().current() }; let join_clause = join_expr.to_sql(); let select_list = group_by_expr.to_select_list(); let group_by_clause = group_by_expr.to_sql(); // Build WHERE clauses let bm25_where = outer_bm25.to_sql("@@@"); let pg_where = outer_bm25.to_sql(" = "); // PostgreSQL native query let pg_query = format!( "SELECT {select_list} {join_clause} WHERE {pg_where} {group_by_clause}" ); // BM25 query with aggregate custom scan enabled let bm25_query = format!( "SELECT {select_list} {join_clause} WHERE {bm25_where} {group_by_clause}" ); // GUCs: enable both join and aggregate custom scans gucs.aggregate_custom_scan = true; gucs.join_custom_scan = true; gucs.custom_scan = true; compare( &pg_query, &bm25_query, &gucs, &mut pool.pull(), &setup_sql, |query, conn| { let rows = query.fetch_dynamic(conn); let mut string_rows: Vec = rows .into_iter() .map(|row| { let mut row_string = String::new(); for i in 0..row.len() { if i > 0 { row_string.push('|'); } let value_str = if let Ok(val) = row.try_get::(i) { val.to_string() } else if let Ok(val) = row.try_get::(i) { val.to_string() } else if let Ok(val) = row.try_get::(i) { format!("{:.6}", val) } else if let Ok(val) = row.try_get::(i) { val } else { "NULL".to_string() }; row_string.push_str(&value_str); } row_string }) .collect(); string_rows.sort(); string_rows }, )?; }); } /// Property test for DISTINCT aggregates on multi-table joins — ensures /// SUM(DISTINCT), COUNT(DISTINCT), AVG(DISTINCT) produce the same results /// via DataFusion aggregate pushdown as native PostgreSQL. #[rstest] #[tokio::test] async fn generated_aggregate_join_distinct(database: Db) { let pool = MutexObjectPool::::new( move || block_on(async { database.connection().await }), |_| {}, ); let tables_and_sizes = [("users", 50), ("products", 50), ("orders", 50)]; let all_tables: Vec<&str> = tables_and_sizes.iter().map(|(table, _)| *table).collect(); let setup_sql = generated_queries_setup(&mut pool.pull(), &tables_and_sizes, COLUMNS); let text_columns = columns_named(vec!["name"]); let join_key_columns = columns_named(vec!["id", "age"]); let grouping_columns: Vec<&str> = COLUMNS .iter() .filter(|col| col.is_groupable && col.is_whereable) .map(|col| col.name) .collect(); proptest!(|( num_tables in 2..=3usize, outer_bm25 in arb_wheres(vec![all_tables[0]], &text_columns), group_by_expr in arb_group_by( grouping_columns.iter().map(|c| format!("{}.{}", all_tables[0], c)).collect::>(), vec![ "COUNT(DISTINCT users.age)", "SUM(DISTINCT users.age)", "AVG(DISTINCT users.age)", ], ), mut gucs in any::(), )| { let tables_for_join: Vec<&str> = all_tables[..num_tables].to_vec(); let join = arb_joins( prop_oneof![Just(JoinType::Inner), Just(JoinType::Left)], tables_for_join.clone(), &join_key_columns, ); let join_expr = { use proptest::strategy::ValueTree; use proptest::test_runner::TestRunner; let mut runner = TestRunner::default(); join.new_tree(&mut runner).unwrap().current() }; let join_clause = join_expr.to_sql(); let select_list = group_by_expr.to_select_list(); let group_by_clause = group_by_expr.to_sql(); let bm25_where = outer_bm25.to_sql("@@@"); let pg_where = outer_bm25.to_sql(" = "); let pg_query = format!( "SELECT {select_list} {join_clause} WHERE {pg_where} {group_by_clause}" ); let bm25_query = format!( "SELECT {select_list} {join_clause} WHERE {bm25_where} {group_by_clause}" ); gucs.aggregate_custom_scan = true; gucs.join_custom_scan = true; gucs.custom_scan = true; compare( &pg_query, &bm25_query, &gucs, &mut pool.pull(), &setup_sql, |query, conn| { let rows = query.fetch_dynamic(conn); let mut string_rows: Vec = rows .into_iter() .map(|row| { let mut row_string = String::new(); for i in 0..row.len() { if i > 0 { row_string.push('|'); } let value_str = if let Ok(val) = row.try_get::(i) { val.to_string() } else if let Ok(val) = row.try_get::(i) { val.to_string() } else if let Ok(val) = row.try_get::(i) { format!("{:.6}", val) } else if let Ok(val) = row.try_get::(i) { val } else { "NULL".to_string() }; row_string.push_str(&value_str); } row_string }) .collect(); string_rows.sort(); string_rows }, )?; }); } /// /// Property test for STDDEV/VARIANCE aggregates — ensures equivalence between /// PostgreSQL native and ParadeDB's DataFusion aggregate backend. /// /// STDDEV and VARIANCE return FLOAT8 which cannot be compared exactly due to /// floating-point precision differences between DataFusion and PostgreSQL. /// Results are rounded to 6 decimal places before comparison. /// #[rstest] #[tokio::test] async fn generated_group_by_stddev(database: Db) { let pool = MutexObjectPool::::new( move || { block_on(async { { database.connection().await } }) }, |_| {}, ); let table_name = "users"; let setup_sql = generated_queries_setup(&mut pool.pull(), &[(table_name, 50)], COLUMNS); // Columns that can be used for grouping (must have fast: true in index) let columns: Vec<_> = COLUMNS .iter() .filter(|col| col.is_groupable && col.is_whereable) .cloned() .collect(); let grouping_columns: Vec<_> = columns.iter().map(|col| col.name).collect(); proptest!(|( text_where_expr in arb_wheres( vec![table_name], &columns, ), numeric_where_expr in arb_wheres( vec![table_name], &columns_named(vec!["age", "price", "rating"]), ), group_by_expr in arb_group_by(grouping_columns.to_vec(), vec!["STDDEV(price)", "VARIANCE(price)", "STDDEV(age)", "VARIANCE(age)"]), gucs in any::(), )| { let select_list = group_by_expr.to_select_list(); let group_by_clause = group_by_expr.to_sql(); // Create combined WHERE clause for PostgreSQL using = operator let pg_where_clause = format!( "({}) AND ({})", text_where_expr.to_sql(" = "), numeric_where_expr.to_sql(" < ") ); // Create combined WHERE clause for BM25 using appropriate operators let bm25_where_clause = format!( "({}) AND ({})", text_where_expr.to_sql("@@@"), numeric_where_expr.to_sql(" < ") ); let pg_query = format!( "SELECT {select_list} FROM {table_name} WHERE {pg_where_clause} {group_by_clause}", ); let bm25_query = format!( "SELECT {select_list} FROM {table_name} WHERE {bm25_where_clause} {group_by_clause}", ); // Custom result comparator that rounds f64 values to 6 decimal places let compare_results = |query: &str, conn: &mut PgConnection| -> Vec { let rows = query.fetch_dynamic(conn); let mut string_rows: Vec = rows .into_iter() .map(|row| { let mut row_string = String::new(); for i in 0..row.len() { if i > 0 { row_string.push('|'); } let value_str = if let Ok(val) = row.try_get::(i) { // Round to 6 decimal places to absorb floating-point differences format!("{:.6}", (val * 1_000_000.0).round() / 1_000_000.0) } else if let Ok(val) = row.try_get::(i) { val.to_string() } else if let Ok(val) = row.try_get::(i) { val.to_string() } else if let Ok(val) = row.try_get::(i) { val } else { "NULL".to_string() }; row_string.push_str(&value_str); } row_string }) .collect(); // Sort for consistent comparison string_rows.sort(); string_rows }; compare(&pg_query, &bm25_query, &gucs, &mut pool.pull(), &setup_sql, compare_results)?; }); } /// /// Property test for aggregate-on-join — ensures equivalence between PostgreSQL /// native aggregation with joins and ParadeDB's DataFusion aggregate backend. /// /// Combines INNER JOINs with GROUP BY aggregates (COUNT, SUM, AVG, MIN, MAX) /// and verifies the DataFusion aggregate path matches PostgreSQL when both /// `enable_aggregate_custom_scan` and `enable_join_custom_scan` are enabled. /// /// Uses only 2 tables and INNER JOIN to keep the test focused and fast. /// #[rstest] #[tokio::test] async fn generated_join_aggregates(database: Db) { let pool = MutexObjectPool::::new( move || { block_on(async { { database.connection().await } }) }, |_| {}, ); // Two tables for join testing, small sizes to keep tests fast let tables_and_sizes = [("users", 30), ("products", 30)]; let all_tables: Vec<&str> = tables_and_sizes.iter().map(|(table, _)| *table).collect(); let setup_sql = generated_queries_setup(&mut pool.pull(), &tables_and_sizes, COLUMNS); // Text columns for BM25 WHERE clauses let text_columns = columns_named(vec!["name"]); // Columns for join keys let join_key_columns = columns_named(vec!["id", "age"]); // Columns for GROUP BY (must be fast fields, qualified with table name) let grouping_columns: Vec = COLUMNS .iter() .filter(|col| col.is_groupable && col.is_whereable) .map(|col| format!("{}.{}", all_tables[0], col.name)) .collect(); proptest!(|( // Outer table BM25 predicate outer_bm25 in arb_wheres(vec![all_tables[0]], &text_columns), // GROUP BY + aggregates group_by_expr in arb_group_by( grouping_columns.clone(), vec!["COUNT(*)", "SUM(users.age)", "AVG(users.age)", "MIN(users.rating)", "MAX(users.rating)"], ), mut gucs in any::(), )| { // Generate join expression (INNER JOIN only) let join = arb_joins( Just(JoinType::Inner), all_tables.clone(), &join_key_columns, ); let join_expr = { use proptest::strategy::ValueTree; use proptest::test_runner::TestRunner; let mut runner = TestRunner::default(); join.new_tree(&mut runner).unwrap().current() }; let join_clause = join_expr.to_sql(); let select_list = group_by_expr.to_select_list(); let group_by_clause = group_by_expr.to_sql(); // Build WHERE clauses let bm25_where = outer_bm25.to_sql("@@@"); let pg_where = outer_bm25.to_sql(" = "); // PostgreSQL native query let pg_query = format!( "SELECT {select_list} {join_clause} WHERE {pg_where} {group_by_clause}" ); // BM25 query with aggregate custom scan enabled let bm25_query = format!( "SELECT {select_list} {join_clause} WHERE {bm25_where} {group_by_clause}" ); // GUCs: enable both join and aggregate custom scans gucs.aggregate_custom_scan = true; gucs.join_custom_scan = true; gucs.custom_scan = true; compare( &pg_query, &bm25_query, &gucs, &mut pool.pull(), &setup_sql, |query, conn| { let rows = query.fetch_dynamic(conn); let mut string_rows: Vec = rows .into_iter() .map(|row| { let mut row_string = String::new(); for i in 0..row.len() { if i > 0 { row_string.push('|'); } let value_str = if let Ok(val) = row.try_get::(i) { val.to_string() } else if let Ok(val) = row.try_get::(i) { val.to_string() } else if let Ok(val) = row.try_get::(i) { format!("{:.6}", val) } else if let Ok(val) = row.try_get::(i) { val } else { "NULL".to_string() }; row_string.push_str(&value_str); } row_string }) .collect(); string_rows.sort(); string_rows }, )?; }); } /// /// Property test for numeric pushdown - ensures equivalence between PostgreSQL and BM25 behavior /// for numeric comparison operators (=, <, <=, >, >=, BETWEEN). /// /// Tests both Numeric64 (precision <= 18) and NumericBytes (unlimited precision) storage types. /// #[rstest] #[tokio::test] async fn generated_numeric_pushdown(database: Db) { let pool = MutexObjectPool::::new( move || { block_on(async { { database.connection().await } }) }, |_| {}, ); let table_name = "users"; // Use more rows to get better coverage of value ranges let setup_sql = generated_queries_setup(&mut pool.pull(), &[(table_name, 100)], COLUMNS); // Numeric columns for testing - includes both Numeric64 and NumericBytes storage types let numeric_columns = columns_named(vec![ "price", // NUMERIC(10,2) - Numeric64 "small_numeric", // NUMERIC(5,2) - Numeric64 "int_numeric", // NUMERIC(10,0) - Numeric64 (integer-like) "high_scale", // NUMERIC(18,6) - Numeric64 with high scale "big_numeric", // NUMERIC - NumericBytes (unlimited precision) "age", // INTEGER - for comparison ]); proptest!(|( numeric_expr in arb_numeric_expr(vec![table_name], &numeric_columns), gucs in any::(), )| { // Both queries use the same SQL since numeric comparison operators // are handled identically - the pushdown happens internally in BM25 let where_clause = numeric_expr.to_sql(); // We need a BM25 predicate to trigger the custom scan // Use an OR clause to match all possible name values in the test data let bm25_predicate = format!( "{table_name}.name @@@ pdb.all()" ); // PostgreSQL query: uses only the numeric predicate let pg_query = format!( "SELECT id FROM {table_name} WHERE {where_clause} ORDER BY id" ); // BM25 query: combines BM25 predicate with numeric pushdown let bm25_query = format!( "SELECT id FROM {table_name} WHERE {bm25_predicate} AND {where_clause} ORDER BY id" ); compare( &pg_query, &bm25_query, &gucs, &mut pool.pull(), &setup_sql, |query, conn| { let mut rows = query.fetch::<(i64,)>(conn); rows.sort(); rows }, )?; }); } /// /// Property test for JoinScan SEMI and ANTI joins. /// /// Fuzzes between: /// - SEMI join via `IN (SELECT ...)` subquery /// - ANTI join via `NOT EXISTS (SELECT ... WHERE correlated)` with `IS NOT NULL` /// /// PostgreSQL only triggers an anti join plan with `NOT EXISTS`, not `NOT IN` /// (due to NULL semantics). An `IS NOT NULL` condition on the join column is /// also required for the anti join optimization. /// /// This complements `generated_joins_small` and verifies both: /// - `paradedb.enable_join_custom_scan = false`: no ParadeDB Join Scan is used /// - `paradedb.enable_join_custom_scan = true`: ParadeDB Join Scan is used #[rstest] #[tokio::test] async fn generated_join_semi_like(database: Db) { let pool = MutexObjectPool::::new( move || { block_on(async { { database.connection().await } }) }, |_| {}, ); // Use varied table sizes to test both when the left side is the largest source // and when the right side is the largest source (which now forces the partition // to the left side anyway for SEMI/ANTI correctness). let tables_and_sizes = [ ("users", 500), ("products", 120), ("orders", 40), ("logs", 1000), ]; let setup_sql = generated_queries_setup(&mut pool.pull(), &tables_and_sizes, COLUMNS); let all_tables = vec!["users", "products", "orders", "logs"]; let join_key_columns = vec!["id", "age", "uuid"]; let search_terms = vec![ "alice", "bob", "cloe", "sally", "brandy", "brisket", "anchovy", ]; proptest!(|( semi_join in arb_semi_joins(all_tables.clone(), join_key_columns.clone()), inner_term in proptest::sample::select(search_terms.clone()), is_anti_join in proptest::bool::ANY, limit in 1..=50usize, nested_join in proptest::option::of(arb_semi_joins(all_tables.clone(), join_key_columns.clone())), nested_term in proptest::sample::select(search_terms.clone()), nested_is_anti in proptest::bool::ANY, mut gucs in any::(), )| { let outer = semi_join.outer_table(); let inner = semi_join.inner_table(); let join_col = semi_join.join_column(); // Skip nested cases where the nested inner table collides with the outer tables, // since that would create a self-join which changes the semantics. let nested = nested_join.as_ref().and_then(|nj| { let nested_inner = nj.inner_table(); if nested_inner != outer && nested_inner != inner && nj.outer_table() != nj.inner_table() { Some((nj, &nested_term, nested_is_anti)) } else { None } }); // Build the subquery clause parameterized by operator. // SEMI: IN (SELECT ...), ANTI: IS NOT NULL AND NOT EXISTS (SELECT 1 ... WHERE correlated) // PostgreSQL only uses an anti join plan with NOT EXISTS (not NOT IN) and // requires IS NOT NULL on the join column. let subquery_clause = |op: &str| { let nested_clause = if let Some((nj, nterm, nis_anti)) = &nested { let mid = inner; let deep = nj.inner_table(); let ncol = nj.join_column(); if *nis_anti { format!( " AND {mid}.{ncol} IS NOT NULL AND NOT EXISTS (\ SELECT 1 FROM {deep} \ WHERE {deep}.{ncol} = {mid}.{ncol} \ AND {deep}.name {op} '{nterm}'\ )" ) } else { format!( " AND {mid}.{ncol} IN (\ SELECT {deep}.{ncol} FROM {deep} \ WHERE {deep}.name {op} '{nterm}'\ )" ) } } else { String::new() }; if is_anti_join { format!( "{outer}.{join_col} IS NOT NULL AND NOT EXISTS (\ SELECT 1 FROM {inner} \ WHERE {inner}.{join_col} = {outer}.{join_col} \ AND {inner}.name {op} '{inner_term}'\ {nested_clause}\ )" ) } else { format!( "{outer}.{join_col} IN (\ SELECT {inner}.{join_col} FROM {inner} \ WHERE {inner}.name {op} '{inner_term}'\ {nested_clause}\ )" ) } }; let pg_where = format!("TRUE AND {}", subquery_clause(" = ")); let bm25_where = format!("{outer}.id @@@ pdb.all() AND {}", subquery_clause("@@@")); let pg_query = format!( "SELECT {outer}.id, {outer}.name \ FROM {outer} \ WHERE {pg_where} \ ORDER BY {outer}.id \ LIMIT {limit}" ); let bm25_query = format!( "SELECT {outer}.id, {outer}.name \ FROM {outer} \ WHERE {bm25_where} \ ORDER BY {outer}.id \ LIMIT {limit}" ); for join_custom_scan in [false, true] { gucs.join_custom_scan = join_custom_scan; compare( &pg_query, &bm25_query, &gucs, &mut pool.pull(), &setup_sql, |query, conn| query.fetch::<(i64, String)>(conn), )?; } }); } /// /// Property test for numeric precision preservation. /// /// Tests that high-precision numeric values (which would lose precision if converted to f64) /// are correctly matched in BM25 queries. This specifically tests the Numeric64 storage /// type with values that have more than 15-16 significant digits (f64's precision limit). /// /// Example: 123456789012345678 and 123456789012345679 are distinct in NUMERIC(18,0) /// but would be indistinguishable if converted to f64. /// #[rstest] #[tokio::test] async fn generated_numeric_precision(database: Db) { let pool = MutexObjectPool::::new( move || { block_on(async { { database.connection().await } }) }, |_| {}, ); let table_name = "precision_test"; // Custom setup for precision testing - uses NUMERIC(18,0) which stores as Numeric64 // but with values that exceed f64's precision let precision_columns: &[Column] = &[ Column::new("id", "SERIAL8", "'1'") .primary_key() .groupable(true), Column::new("name", "TEXT", "'test'") .bm25_text_field(r#""name": { "tokenizer": { "type": "keyword" }, "fast": true }"#) .random_generator_sql("'test'"), Column::new("big_int", "NUMERIC(18,0)", "'123456789012345678'") .groupable(false) .bm25_numeric_field(r#""big_int": { "fast": true }"#) // Generate high-precision values that differ only in lower digits // These values would collide if converted to f64 .random_generator_sql( "(ARRAY [123456789012345678, 123456789012345679, 123456789012345680, 999999999999999998, 999999999999999999]::numeric[])[(floor(random() * 5) + 1)::int]" ), ]; let setup_sql = generated_queries_setup(&mut pool.pull(), &[(table_name, 50)], precision_columns); // High-precision test values that would be indistinguishable in f64 let precision_test_values = vec![ "123456789012345678", "123456789012345679", "123456789012345680", "999999999999999998", "999999999999999999", ]; proptest!(|( test_value in proptest::sample::select(precision_test_values), gucs in any::(), )| { // PostgreSQL query - should find exact matches only let pg_query = format!( "SELECT COUNT(*) FROM {table_name} WHERE big_int = {test_value}" ); // BM25 query - should produce identical results // Use 'test' as the name value since all rows have name = 'test' let bm25_query = format!( "SELECT COUNT(*) FROM {table_name} WHERE name @@@ 'test' AND big_int = {test_value}" ); compare( &pg_query, &bm25_query, &gucs, &mut pool.pull(), &setup_sql, |query, conn| query.fetch_one::<(i64,)>(conn).0, )?; }); } /// /// Property test for numeric range queries with precision preservation. /// /// Tests that range queries (>, <, >=, <=, BETWEEN) on high-precision numeric values /// produce correct results without precision loss. /// #[rstest] #[tokio::test] async fn generated_numeric_range_precision(database: Db) { let pool = MutexObjectPool::::new( move || { block_on(async { { database.connection().await } }) }, |_| {}, ); let table_name = "range_precision_test"; // Custom setup for range precision testing let precision_columns: &[Column] = &[ Column::new("id", "SERIAL8", "'1'") .primary_key() .groupable(true), Column::new("name", "TEXT", "'test'") .bm25_text_field(r#""name": { "tokenizer": { "type": "keyword" }, "fast": true }"#) .random_generator_sql("'test'"), Column::new("big_int", "NUMERIC(18,0)", "'100'") .groupable(false) .bm25_numeric_field(r#""big_int": { "fast": true }"#) // Generate sequential high-precision values .random_generator_sql("(floor(random() * 100) + 123456789012345600)::numeric(18,0)"), ]; let setup_sql = generated_queries_setup(&mut pool.pull(), &[(table_name, 100)], precision_columns); // Range boundaries that would collide in f64 let range_bounds = vec![ ("123456789012345650", "123456789012345660"), ("123456789012345670", "123456789012345680"), ("123456789012345690", "123456789012345700"), ]; proptest!(|( (low, high) in proptest::sample::select(range_bounds), gucs in any::(), )| { // PostgreSQL query - range filter let pg_query = format!( "SELECT COUNT(*) FROM {table_name} WHERE big_int >= {low} AND big_int < {high}" ); // BM25 query - should produce identical results // Use 'test' as the name value since all rows have name = 'test' let bm25_query = format!( "SELECT COUNT(*) FROM {table_name} WHERE name @@@ 'test' AND big_int >= {low} AND big_int < {high}" ); compare( &pg_query, &bm25_query, &gucs, &mut pool.pull(), &setup_sql, |query, conn| query.fetch_one::<(i64,)>(conn).0, )?; }); } /// Property test for `pdb.agg()` over joins: the buckets and metrics the DataFusion backend /// assembles must equal the rows of the equivalent SQL `GROUP BY` run by PostgreSQL, which is /// the oracle since `pdb.agg()` itself has no native fallback. Covers nested `terms`, a /// `size` cut under the default count order, NULL buckets, NUMERIC metrics, `cardinality`, /// a SQL `GROUP BY` beside the call, and MPP when the parallel GUCs are on. /// TODO: Consider merging this property test with other "aggregate over join" tests /// (such as `generated_aggregate_join` and `generated_join_aggregates`) in the future. #[rstest] #[tokio::test] async fn generated_pdb_agg_join(database: Db) { let pool = MutexObjectPool::::new( move || block_on(async { database.connection().await }), |_| {}, ); let tables_and_sizes = [("users", 50), ("products", 50), ("orders", 50)]; let all_tables: Vec = tables_and_sizes .iter() .map(|(table, _)| table.to_string()) .collect(); let setup_sql = generated_queries_setup(&mut pool.pull(), &tables_and_sizes, COLUMNS); let where_columns = columns_named(vec!["name", "color"]); let join_key_columns = columns_named(vec!["id", "age"]); proptest!(|( (join_expr, agg, wheres) in arb_pdb_agg_join(all_tables.clone(), &join_key_columns, &where_columns), mut gucs in any::(), )| { let join_clause = join_expr.to_sql(); let pg_query = agg.pg_query(&join_clause, &wheres.pg_where()); let bm25_query = agg.pdb_query(&join_clause, &wheres.bm25_where()); // `pdb.agg()` over a join only runs on the DataFusion backend. gucs.aggregate_custom_scan = true; gucs.join_custom_scan = true; gucs.custom_scan = true; if !agg.outer_aggs.is_empty() { let pg_outer_query = agg.pg_outer_query(&join_clause, &wheres.pg_where()); compare_with_side( &pg_outer_query, &bm25_query, &gucs, &mut pool.pull(), &setup_sql, |query, side, conn| { "SET work_mem TO '64MB';".execute(conn); let mut rows = agg .outer_rows(query.fetch_dynamic(conn), side.is_candidate()) .unwrap(); rows.sort(); rows }, )?; } compare( &pg_query, &bm25_query, &gucs, &mut pool.pull(), &setup_sql, |query, conn| { // A keyless join under three bucket keys makes tens of thousands of // buckets, and the DataFusion aggregate cannot spill past `work_mem`. "SET work_mem TO '64MB';".execute(conn); let mut rows = agg.rows(query.fetch_dynamic(conn)).unwrap(); rows.sort(); rows }, )?; }); } /// The same single-table `pdb.agg()` answered by Tantivy and by DataFusion. The documents must /// be equal as they are: bucket order, `sum_other_doc_count`, NULL buckets, and metric values /// alike. A grouped query sorted by an aggregate under a `LIMIT` is routed to DataFusion whatever /// the planner's group estimate, and the limit sits above any group count, so it cuts nothing. #[rstest] #[tokio::test] async fn generated_pdb_agg_single_table(database: Db) { let pool = MutexObjectPool::::new( move || block_on(async { database.connection().await }), |_| {}, ); let setup_sql = generated_queries_setup(&mut pool.pull(), &[("users", 50)], COLUMNS); let text_columns = columns_named(vec!["name"]); proptest!(|( outer_bm25 in arb_wheres(vec!["users".to_string()], &text_columns), agg in arb_pdb_agg_single_table(), mut gucs in any::(), )| { let group = agg.outer_group.as_deref().expect("a single-table spec sits beside a GROUP BY"); let tantivy_query = format!( "SELECT {group}, COUNT(*), {} FROM users WHERE {} GROUP BY {group}", agg.call(), outer_bm25.to_sql("@@@") ); let datafusion_query = format!("{tantivy_query} ORDER BY COUNT(*) DESC LIMIT 1000"); gucs.aggregate_custom_scan = true; gucs.custom_scan = true; let sides = Sides { baseline: gucs.set(), candidate: gucs.set(), }; compare_on( &sides, &tantivy_query, &datafusion_query, &gucs, &mut pool.pull(), &setup_sql, |query, conn| { let mut documents = agg.documents(query.fetch_dynamic(conn)).unwrap(); documents.sort_by(|a, b| a.0.cmp(&b.0)); documents }, )?; }); }