From a01eb6eb0cdbcd122b8b77e16ac155c06458dcc2 Mon Sep 17 00:00:00 2001 From: Nupur Agrawal Date: Mon, 14 Sep 2026 22:41:20 +0530 Subject: [PATCH 01/14] refactor: unify temporary projection rewrites --- .../src/backend/pool/connection/aggregate.rs | 14 +-- pgdog/src/backend/pool/connection/buffer.rs | 10 +- .../pool/connection/multi_shard/mod.rs | 4 +- .../frontend/router/parser/query/select.rs | 2 +- .../rewrite/statement/aggregate/engine.rs | 37 +++--- .../parser/rewrite/statement/aggregate/mod.rs | 23 +++- .../rewrite/statement/aggregate/plan.rs | 112 ------------------ .../router/parser/rewrite/statement/mod.rs | 1 + .../router/parser/rewrite/statement/plan.rs | 9 +- .../parser/rewrite/statement/projection.rs | 95 +++++++++++++++ pgdog/src/frontend/router/parser/route.rs | 16 ++- 11 files changed, 162 insertions(+), 161 deletions(-) delete mode 100644 pgdog/src/frontend/router/parser/rewrite/statement/aggregate/plan.rs create mode 100644 pgdog/src/frontend/router/parser/rewrite/statement/projection.rs diff --git a/pgdog/src/backend/pool/connection/aggregate.rs b/pgdog/src/backend/pool/connection/aggregate.rs index c06bf594a..171d8bbfa 100644 --- a/pgdog/src/backend/pool/connection/aggregate.rs +++ b/pgdog/src/backend/pool/connection/aggregate.rs @@ -6,7 +6,7 @@ use std::mem; use crate::{ frontend::router::parser::{ Aggregate, AggregateFunction, AggregateTarget, - rewrite::statement::aggregate::{AggregateRewritePlan, HelperKind}, + rewrite::statement::{aggregate::HelperKind, projection::ProjectionRewritePlan}, }, net::{ Decoder, @@ -245,7 +245,7 @@ impl<'a> Aggregates<'a> { rows: &'a VecDeque, decoder: &'a Decoder, aggregate: &'a Aggregate, - plan: &AggregateRewritePlan, + plan: &ProjectionRewritePlan, ) -> Option { let mut helper_columns: HashMap = HashMap::new(); @@ -262,7 +262,7 @@ impl<'a> Aggregates<'a> { } } - for helper in plan.helpers() { + for helper in plan.aggregate_helpers() { let Some(index) = decoder.row_description().field_index(&helper.alias) else { continue; }; @@ -523,7 +523,7 @@ mod test { shard1.add("3"); rows.push_back(shard1); - let plan = AggregateRewritePlan::default(); + let plan = ProjectionRewritePlan::default(); let mut result = Aggregates::new(&rows, &decoder, &aggregate, &plan) .unwrap() .aggregate() @@ -557,7 +557,7 @@ mod test { &rows, &decoder, &aggregate, - &AggregateRewritePlan::default(), + &ProjectionRewritePlan::default(), ) .unwrap() .aggregate() @@ -602,7 +602,7 @@ mod test { &rows, &decoder, &aggregate, - &AggregateRewritePlan::default(), + &ProjectionRewritePlan::default(), ) .unwrap() .aggregate() @@ -649,7 +649,7 @@ mod test { &rows, &decoder, &aggregate, - &AggregateRewritePlan::default(), + &ProjectionRewritePlan::default(), ) .unwrap() .aggregate() diff --git a/pgdog/src/backend/pool/connection/buffer.rs b/pgdog/src/backend/pool/connection/buffer.rs index d8ea35a4d..df7a63ca1 100644 --- a/pgdog/src/backend/pool/connection/buffer.rs +++ b/pgdog/src/backend/pool/connection/buffer.rs @@ -9,7 +9,7 @@ use std::{ use crate::{ frontend::router::parser::{ Aggregate, DistinctBy, DistinctColumn, Limit, OrderBy, - rewrite::statement::aggregate::AggregateRewritePlan, + rewrite::statement::projection::ProjectionRewritePlan, }, net::{ Decoder, @@ -140,7 +140,7 @@ impl Buffer { &mut self, aggregate: &Aggregate, decoder: &Decoder, - plan: &AggregateRewritePlan, + plan: &ProjectionRewritePlan, ) -> Result<(), super::Error> { let buffer: VecDeque = std::mem::take(&mut self.buffer); let mut rows = if aggregate.is_empty() { @@ -157,7 +157,7 @@ impl Buffer { Ok(()) } - fn drop_helper_columns(rows: &mut VecDeque, plan: &AggregateRewritePlan) { + fn drop_helper_columns(rows: &mut VecDeque, plan: &ProjectionRewritePlan) { if plan.is_noop() { return; } @@ -285,7 +285,7 @@ mod test { buf.add(dr.message()).unwrap(); } - buf.aggregate(&agg, &Decoder::from(rd), &AggregateRewritePlan::default()) + buf.aggregate(&agg, &Decoder::from(rd), &ProjectionRewritePlan::default()) .unwrap(); buf.mark_full(); @@ -312,7 +312,7 @@ mod test { } } - buf.aggregate(&agg, &Decoder::from(rd), &AggregateRewritePlan::default()) + buf.aggregate(&agg, &Decoder::from(rd), &ProjectionRewritePlan::default()) .unwrap(); buf.mark_full(); diff --git a/pgdog/src/backend/pool/connection/multi_shard/mod.rs b/pgdog/src/backend/pool/connection/multi_shard/mod.rs index ee4ebaca6..13372c5a4 100644 --- a/pgdog/src/backend/pool/connection/multi_shard/mod.rs +++ b/pgdog/src/backend/pool/connection/multi_shard/mod.rs @@ -264,7 +264,7 @@ impl MultiShard { .aggregate( self.route.aggregate(), &self.decoder, - self.route.aggregate_rewrite_plan(), + self.route.projection_rewrite_plan(), ) .map_err(Error::from)?; @@ -311,7 +311,7 @@ impl MultiShard { { // Only send it to the client once all shards sent it, // so we don't get early requests from clients. - let plan = self.route.aggregate_rewrite_plan(); + let plan = self.route.projection_rewrite_plan(); if plan.is_noop() { forward = Some(message); } else { diff --git a/pgdog/src/frontend/router/parser/query/select.rs b/pgdog/src/frontend/router/parser/query/select.rs index 5dceda733..c552f72ee 100644 --- a/pgdog/src/frontend/router/parser/query/select.rs +++ b/pgdog/src/frontend/router/parser/query/select.rs @@ -279,7 +279,7 @@ impl QueryParser { // Only rewrite if query is cross-shard. if query.is_cross_shard() && context.shards > 1 { - query.set_rewrite_plan(cached_ast.rewrite_plan.aggregates.clone()); + query.set_projection_rewrite_plan(cached_ast.rewrite_plan.projection.clone()); } Ok(Command::Query( diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/aggregate/engine.rs b/pgdog/src/frontend/router/parser/rewrite/statement/aggregate/engine.rs index 2a2ef2fb8..d4b062a8d 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/aggregate/engine.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/aggregate/engine.rs @@ -3,7 +3,10 @@ use crate::frontend::router::parser::aggregate::{Aggregate, AggregateFunction}; use itertools::*; use pg_raw_parse::{Node, make, nodes}; -use super::{AggregateRewritePlan, HelperKind, HelperMapping, RewriteOutput}; +use super::{AggregateHelper, HelperKind}; +use crate::frontend::router::parser::rewrite::statement::projection::{ + ProjectionRewritePlan, RewriteOutput, +}; /// Query rewrite engine. Currently supports injecting helper aggregates for AVG and /// variance-related functions that require additional helper aggregates when run @@ -18,7 +21,7 @@ impl AggregatesRewrite { mem: make::MemoryToken<'a>, aggregate: &Aggregate, ) -> RewriteOutput { - let mut plan = AggregateRewritePlan::new(); + let mut plan = ProjectionRewritePlan::default(); let helper_nodes = aggregate .targets() @@ -43,9 +46,9 @@ impl AggregatesRewrite { format!("__pgdog_{}_col{}", kind.alias_suffix(), target.column()); let node = mem.make_res_target(Some(&helper_alias), mem.empty(), func.uncast()); - plan.add_helper(HelperMapping { + plan.add_aggregate_helper(AggregateHelper { target_column: target.column(), - helper_column: select.target_list().len() + idx, + projected_column: select.target_list().len() + idx, distinct: target.is_distinct(), kind, alias: helper_alias, @@ -196,10 +199,10 @@ mod tests { let (ast, output) = rewrite("SELECT AVG(price) FROM menu"); assert!(!output.plan.is_noop()); assert_eq!(output.plan.drop_columns().collect::>(), &[1]); - assert_eq!(output.plan.helpers().len(), 1); - let helper = &output.plan.helpers()[0]; + assert_eq!(output.plan.aggregate_helpers().len(), 1); + let helper = &output.plan.aggregate_helpers()[0]; assert_eq!(helper.target_column, 0); - assert_eq!(helper.helper_column, 1); + assert_eq!(helper.projected_column, 1); assert!(!helper.distinct); assert!(matches!(helper.kind, HelperKind::Count)); @@ -217,10 +220,10 @@ mod tests { fn rewrite_engine_handles_mismatched_pair() { let (ast, output) = rewrite("SELECT COUNT(price::numeric), AVG(price) FROM menu"); assert_eq!(output.plan.drop_columns().collect::>(), &[2]); - assert_eq!(output.plan.helpers().len(), 1); - let helper = &output.plan.helpers()[0]; + assert_eq!(output.plan.aggregate_helpers().len(), 1); + let helper = &output.plan.aggregate_helpers()[0]; assert_eq!(helper.target_column, 1); - assert_eq!(helper.helper_column, 2); + assert_eq!(helper.projected_column, 2); assert!(!helper.distinct); assert!(matches!(helper.kind, HelperKind::Count)); @@ -240,16 +243,16 @@ mod tests { fn rewrite_engine_multiple_avg_helpers() { let (ast, output) = rewrite("SELECT AVG(price), AVG(discount) FROM menu"); assert_eq!(output.plan.drop_columns().collect::>(), &[2, 3]); - assert_eq!(output.plan.helpers().len(), 2); + assert_eq!(output.plan.aggregate_helpers().len(), 2); - let helper_price = &output.plan.helpers()[0]; + let helper_price = &output.plan.aggregate_helpers()[0]; assert_eq!(helper_price.target_column, 0); - assert_eq!(helper_price.helper_column, 2); + assert_eq!(helper_price.projected_column, 2); assert!(matches!(helper_price.kind, HelperKind::Count)); - let helper_discount = &output.plan.helpers()[1]; + let helper_discount = &output.plan.aggregate_helpers()[1]; assert_eq!(helper_discount.target_column, 1); - assert_eq!(helper_discount.helper_column, 3); + assert_eq!(helper_discount.projected_column, 3); assert!(matches!(helper_discount.kind, HelperKind::Count)); let aggregate = Aggregate::parse(&ast, &Default::default()); @@ -269,11 +272,11 @@ mod tests { let (ast, output) = rewrite("SELECT STDDEV(price) FROM menu"); assert!(!output.plan.is_noop()); assert_eq!(output.plan.drop_columns().collect::>(), &[1, 2, 3]); - assert_eq!(output.plan.helpers().len(), 3); + assert_eq!(output.plan.aggregate_helpers().len(), 3); let kinds: Vec = output .plan - .helpers() + .aggregate_helpers() .iter() .map(|helper| { assert_eq!(helper.target_column, 0); diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/aggregate/mod.rs b/pgdog/src/frontend/router/parser/rewrite/statement/aggregate/mod.rs index 0ae8e04a5..f03c7e3a4 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/aggregate/mod.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/aggregate/mod.rs @@ -1,13 +1,30 @@ mod engine; -mod plan; +pub(crate) use super::projection::AggregateHelper; use super::{Error, RewritePlan, StatementRewrite}; use crate::backend::schema::Schema; use crate::frontend::router::parser::aggregate::Aggregate; use pg_raw_parse::{make::MemoryToken, nodes::SelectStmtMut}; pub(crate) use engine::AggregatesRewrite; -pub(crate) use plan::{AggregateRewritePlan, HelperKind, HelperMapping, RewriteOutput}; + +/// Type of aggregate function added to the result set. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum HelperKind { + Count, + Sum, + SumSquares, +} + +impl HelperKind { + pub(crate) fn alias_suffix(self) -> &'static str { + match self { + Self::Count => "count", + Self::Sum => "sum", + Self::SumSquares => "sumsq", + } + } +} impl StatementRewrite<'_> { /// Add missing COUNT(*) and other helps when using aggregates. @@ -32,7 +49,7 @@ impl StatementRewrite<'_> { return Ok(()); } - plan.aggregates = output.plan; + plan.projection = output.plan; self.rewritten = true; Ok(()) } diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/aggregate/plan.rs b/pgdog/src/frontend/router/parser/rewrite/statement/aggregate/plan.rs deleted file mode 100644 index bffd185d3..000000000 --- a/pgdog/src/frontend/router/parser/rewrite/statement/aggregate/plan.rs +++ /dev/null @@ -1,112 +0,0 @@ -/// Type of aggregate function added to the result set. -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub(crate) enum HelperKind { - /// COUNT(*) or COUNT(name) - Count, - /// SUM(column) - Sum, - /// SUM(POWER(column, 2)) - SumSquares, -} - -impl HelperKind { - /// Suffix for the aggregate function. - pub(crate) fn alias_suffix(&self) -> &'static str { - match self { - HelperKind::Count => "count", - HelperKind::Sum => "sum", - HelperKind::SumSquares => "sumsq", - } - } -} - -/// Context on the aggregate function column added to the result set. -#[derive(Debug, Clone, PartialEq)] -pub(crate) struct HelperMapping { - pub(crate) target_column: usize, - pub(crate) helper_column: usize, - pub(crate) distinct: bool, - pub(crate) kind: HelperKind, - pub(crate) alias: String, -} - -/// Plan describing how the proxy rewrites a query and its results. -#[derive(Debug, Clone, Default, PartialEq)] -pub(crate) struct AggregateRewritePlan { - helpers: Vec, -} - -impl AggregateRewritePlan { - /// Create new no-op aggregate rewrite plan. - pub(crate) fn new() -> Self { - Self { - helpers: Vec::new(), - } - } - - /// Is this plan a no-op? Doesn't do anything. - pub(crate) fn is_noop(&self) -> bool { - self.helpers.is_empty() - } - - pub(crate) fn drop_columns(&self) -> impl Iterator + '_ { - self.helpers.iter().map(|h| h.helper_column) - } - - pub(crate) fn helpers(&self) -> &[HelperMapping] { - &self.helpers - } - - pub(crate) fn add_helper(&mut self, mapping: HelperMapping) { - self.helpers.push(mapping); - } -} - -#[derive(Debug, Default, Clone)] -pub(crate) struct RewriteOutput { - pub(crate) plan: AggregateRewritePlan, -} - -impl RewriteOutput { - pub(crate) fn new(plan: AggregateRewritePlan) -> Self { - Self { plan } - } -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn rewrite_plan_noop() { - let plan = AggregateRewritePlan::new(); - assert!(plan.is_noop()); - assert!(plan.drop_columns().count() == 0); - assert!(plan.helpers().is_empty()); - } - - #[test] - fn rewrite_plan_helpers() { - let mut plan = AggregateRewritePlan::new(); - plan.add_helper(HelperMapping { - target_column: 0, - helper_column: 1, - distinct: false, - kind: HelperKind::Count, - alias: "__pgdog_count_expr7_col0".into(), - }); - assert_eq!(plan.helpers().len(), 1); - let helper = &plan.helpers()[0]; - assert_eq!(helper.target_column, 0); - assert_eq!(helper.helper_column, 1); - assert!(!helper.distinct); - assert!(matches!(helper.kind, HelperKind::Count)); - assert_eq!(helper.alias, "__pgdog_count_expr7_col0"); - } - - #[test] - fn rewrite_output_defaults() { - let output = RewriteOutput::default(); - assert!(output.plan.is_noop()); - } -} diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs b/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs index 813f0be50..db67ba6cc 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs @@ -14,6 +14,7 @@ pub(crate) mod insert; pub(crate) mod nextval; pub(crate) mod offset; pub(crate) mod plan; +pub(crate) mod projection; pub(crate) mod simple_prepared; pub(crate) mod unique_id; pub(crate) mod update; diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs b/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs index 94bac8bce..2ed998446 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs @@ -8,7 +8,7 @@ use super::insert::{build_resolved_split_requests, build_split_requests}; use super::nextval::SequenceCall; use super::offset::OffsetPlan; use super::{ - Error, InsertSplit, PrepareExecute, ShardingKeyUpdate, aggregate::AggregateRewritePlan, + Error, InsertSplit, PrepareExecute, ShardingKeyUpdate, projection::ProjectionRewritePlan, }; #[derive(Debug, Clone, PartialEq, Eq)] @@ -50,9 +50,8 @@ pub(crate) struct RewritePlan { /// multiple queries. pub(crate) insert_split: Vec, - /// Position in the result where the count(*) or count(name) - /// functions are added. - pub(crate) aggregates: AggregateRewritePlan, + /// Temporary result columns needed while merging cross-shard results. + pub(crate) projection: ProjectionRewritePlan, /// Sharding key is being updated, we need to execute /// a multi-step plan. @@ -91,7 +90,7 @@ impl RewritePlan { && self.stmt.is_none() && self.prepare_rewrites.is_empty() && self.insert_split.is_empty() - && self.aggregates.is_noop() + && self.projection.is_noop() && self.sharding_key_update.is_none() && self.offset.is_none() } diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs b/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs new file mode 100644 index 000000000..3ce95be35 --- /dev/null +++ b/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs @@ -0,0 +1,95 @@ +use super::aggregate::HelperKind; + +/// Aggregate function projected temporarily for cross-shard merging. +#[derive(Debug, Clone, PartialEq)] +pub(crate) struct AggregateHelper { + pub(crate) target_column: usize, + pub(crate) projected_column: usize, + pub(crate) distinct: bool, + pub(crate) kind: HelperKind, + pub(crate) alias: String, +} + +/// Column projected temporarily for cross-shard ordering. +#[derive(Debug, Clone, PartialEq)] +pub(crate) struct OrderByHelper { + pub(crate) sort_position: usize, + pub(crate) projected_column: usize, +} + +/// Temporary result columns required while merging cross-shard results. +#[derive(Debug, Clone, Default, PartialEq)] +pub(crate) struct ProjectionRewritePlan { + aggregate_helpers: Vec, + order_by_helpers: Vec, +} + +impl ProjectionRewritePlan { + pub(crate) fn is_noop(&self) -> bool { + self.aggregate_helpers.is_empty() && self.order_by_helpers.is_empty() + } + + pub(crate) fn drop_columns(&self) -> impl Iterator + '_ { + self.aggregate_helpers + .iter() + .map(|helper| helper.projected_column) + .chain( + self.order_by_helpers + .iter() + .map(|helper| helper.projected_column), + ) + } + + pub(crate) fn aggregate_helpers(&self) -> &[AggregateHelper] { + &self.aggregate_helpers + } + + pub(crate) fn order_by_helpers(&self) -> &[OrderByHelper] { + &self.order_by_helpers + } + + pub(crate) fn add_aggregate_helper(&mut self, helper: AggregateHelper) { + self.aggregate_helpers.push(helper); + } + + pub(crate) fn add_order_by_helper(&mut self, helper: OrderByHelper) { + self.order_by_helpers.push(helper); + } +} + +#[derive(Debug, Default, Clone)] +pub(crate) struct RewriteOutput { + pub(crate) plan: ProjectionRewritePlan, +} + +impl RewriteOutput { + pub(crate) fn new(plan: ProjectionRewritePlan) -> Self { + Self { plan } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn projection_plan_tracks_helpers() { + let mut plan = ProjectionRewritePlan::default(); + plan.add_aggregate_helper(AggregateHelper { + target_column: 0, + projected_column: 1, + distinct: false, + kind: HelperKind::Count, + alias: "__pgdog_count_col0".into(), + }); + plan.add_order_by_helper(OrderByHelper { + sort_position: 0, + projected_column: 2, + }); + + assert!(!plan.is_noop()); + assert_eq!(plan.drop_columns().collect::>(), [1, 2]); + assert_eq!(plan.aggregate_helpers().len(), 1); + assert_eq!(plan.order_by_helpers().len(), 1); + } +} diff --git a/pgdog/src/frontend/router/parser/route.rs b/pgdog/src/frontend/router/parser/route.rs index 94bb9ab24..727aaaf69 100644 --- a/pgdog/src/frontend/router/parser/route.rs +++ b/pgdog/src/frontend/router/parser/route.rs @@ -2,7 +2,7 @@ use std::{fmt::Display, ops::Deref}; use super::{ Aggregate, DistinctBy, Limit, OrderBy, explain_trace::ExplainTrace, - rewrite::statement::aggregate::AggregateRewritePlan, statement::AdvisoryLocks, + rewrite::statement::projection::ProjectionRewritePlan, statement::AdvisoryLocks, }; use crate::frontend::{client::query_engine::TempTableChange, router::sharding::PendingLookup}; use lazy_static::lazy_static; @@ -110,10 +110,8 @@ pub(crate) struct Route { advisory_locks: AdvisoryLocks, /// `DISTINCT` clause, if set. distinct: Option, - /// Rewrites performed by the aggregate rewriter; adds - /// helper columns to this query so we can compute things - /// like avg() or variance(). - rewrite_plan: AggregateRewritePlan, + /// Temporary columns projected for cross-shard result processing. + projection_rewrite_plan: ProjectionRewritePlan, /// Our query explain plan. We attach /// this to the `EXPLAIN` output. explain: Option, @@ -402,12 +400,12 @@ impl Route { self.is_cross_shard() && self.is_write() } - pub(crate) fn aggregate_rewrite_plan(&self) -> &AggregateRewritePlan { - &self.rewrite_plan + pub(crate) fn projection_rewrite_plan(&self) -> &ProjectionRewritePlan { + &self.projection_rewrite_plan } - pub(crate) fn set_rewrite_plan(&mut self, plan: AggregateRewritePlan) { - self.rewrite_plan = plan; + pub(crate) fn set_projection_rewrite_plan(&mut self, plan: ProjectionRewritePlan) { + self.projection_rewrite_plan = plan; } pub(super) fn with_temp_table_change(mut self, temp_table: Option) -> Self { From a6e65327df1f6990ac3d00f97813835a90057023 Mon Sep 17 00:00:00 2001 From: Nupur Agrawal Date: Mon, 14 Sep 2026 23:04:35 +0530 Subject: [PATCH 02/14] refactor: finalize shard rewrites after routing --- .../client/query_engine/route_query.rs | 11 +- .../frontend/client/query_engine/test/mod.rs | 1 + .../query_engine/test/rewrite_offset.rs | 41 +- .../query_engine/test/rewrite_projection.rs | 416 ++++++++++++++++++ .../prepared_statements/global_cache.rs | 106 ++++- pgdog/src/frontend/router/parser/cache/ast.rs | 6 + .../frontend/router/parser/query/select.rs | 9 +- .../parser/rewrite/statement/aggregate/mod.rs | 33 -- .../router/parser/rewrite/statement/mod.rs | 3 +- .../router/parser/rewrite/statement/offset.rs | 112 ++--- .../router/parser/rewrite/statement/plan.rs | 19 +- .../parser/rewrite/statement/projection.rs | 211 ++++++++- pgdog/src/frontend/router/parser/route.rs | 4 + pgdog/src/net/messages/bind.rs | 26 -- 14 files changed, 819 insertions(+), 179 deletions(-) create mode 100644 pgdog/src/frontend/client/query_engine/test/rewrite_projection.rs diff --git a/pgdog/src/frontend/client/query_engine/route_query.rs b/pgdog/src/frontend/client/query_engine/route_query.rs index a8f563a6f..032a322e6 100644 --- a/pgdog/src/frontend/client/query_engine/route_query.rs +++ b/pgdog/src/frontend/client/query_engine/route_query.rs @@ -4,6 +4,7 @@ use tracing::trace; use crate::frontend::router::Error as RouterError; use crate::frontend::router::parser::Error as ParserError; use crate::frontend::router::parser::rewrite::statement::plan::RewriteResult; +use crate::frontend::router::parser::rewrite::statement::projection; use crate::frontend::router::sharding::lookup; use crate::util::safe_timeout; @@ -147,9 +148,15 @@ impl QueryEngine { context.client_request.messages, command, ); - // Apply post-parser rewrites, e.g. offset/limit. + projection::finalize_after_route( + context.client_request, + &cluster.schema(), + rewrite_result.and_then(RewriteResult::offset_plan), + )?; + + // Resolve route-dependent values, e.g. offset/limit. if let Some(rewrite_result) = rewrite_result { - rewrite_result.apply_after_parser(context.client_request)?; + rewrite_result.apply_after_route(context.client_request)?; } // Only validate shard placement for requests that actually execute diff --git a/pgdog/src/frontend/client/query_engine/test/mod.rs b/pgdog/src/frontend/client/query_engine/test/mod.rs index 791312a6a..9c749911e 100644 --- a/pgdog/src/frontend/client/query_engine/test/mod.rs +++ b/pgdog/src/frontend/client/query_engine/test/mod.rs @@ -31,6 +31,7 @@ mod replicas; mod rewrite_extended; mod rewrite_insert_split; mod rewrite_offset; +mod rewrite_projection; mod rewrite_simple_prepared; mod schema_changed; mod set; diff --git a/pgdog/src/frontend/client/query_engine/test/rewrite_offset.rs b/pgdog/src/frontend/client/query_engine/test/rewrite_offset.rs index b21c29045..3fc55a53f 100644 --- a/pgdog/src/frontend/client/query_engine/test/rewrite_offset.rs +++ b/pgdog/src/frontend/client/query_engine/test/rewrite_offset.rs @@ -1,6 +1,7 @@ +use crate::backend::schema::Schema; use crate::frontend::router::parser::Limit; use crate::frontend::router::parser::rewrite::statement::{ - offset::OffsetPlan, plan::RewriteResult, + offset::OffsetPlan, plan::RewriteResult, projection, }; use crate::frontend::router::parser::route::{Route, Shard, ShardWithPriority}; @@ -143,12 +144,18 @@ async fn test_offset_with_unique_id_simple() { "should have bigint cast: {rewritten_sql}" ); - // apply_after_parser with a cross-shard route. + // Finalize with a cross-shard route. context.client_request.route = Some(cross_shard_route()); + projection::finalize_after_route( + context.client_request, + &Schema::default(), + rewrite_result.as_ref().and_then(RewriteResult::offset_plan), + ) + .unwrap(); rewrite_result .as_ref() .unwrap() - .apply_after_parser(context.client_request) + .apply_after_route(context.client_request) .unwrap(); let final_sql = match &context.client_request.messages[0] { @@ -159,7 +166,7 @@ async fn test_offset_with_unique_id_simple() { // unique_id rewrite must survive. assert!( !final_sql.contains("pgdog.unique_id"), - "unique_id rewrite must survive apply_after_parser: {final_sql}" + "unique_id rewrite must survive post-route finalization: {final_sql}" ); assert!( final_sql.contains("::bigint"), @@ -167,8 +174,8 @@ async fn test_offset_with_unique_id_simple() { ); // LIMIT/OFFSET must be rewritten for cross-shard. assert!( - final_sql.contains("LIMIT 15"), - "LIMIT should be 10+5=15: {final_sql}" + final_sql.contains("LIMIT 10 + 5"), + "LIMIT should request limit+offset rows: {final_sql}" ); assert!( !final_sql.contains("OFFSET"), @@ -210,29 +217,35 @@ async fn test_offset_with_unique_id_extended() { "SELECT $4::bigint, $1 FROM test LIMIT $2 OFFSET $3" ); - // apply_after_parser with cross-shard route should only rewrite Bind params. + // Post-route finalization rewrites the SQL without changing Bind values. context.client_request.route = Some(cross_shard_route()); + projection::finalize_after_route( + context.client_request, + &Schema::default(), + rewrite_result.as_ref().and_then(RewriteResult::offset_plan), + ) + .unwrap(); rewrite_result .as_ref() .unwrap() - .apply_after_parser(context.client_request) + .apply_after_route(context.client_request) .unwrap(); - // SQL unchanged (all limit/offset are params). + // SQL uses a stable expression suitable for prepared-statement caching. let final_sql = match &context.client_request.messages[0] { ProtocolMessage::Parse(p) => p.query().to_owned(), _ => panic!("expected Parse"), }; assert_eq!( - final_sql, "SELECT $4::bigint, $1 FROM test LIMIT $2 OFFSET $3", - "SQL must be unchanged for all-param case" + final_sql, "SELECT $4::bigint, $1 FROM test LIMIT $2 + $3", + "SQL must push down limit+offset" ); - // Bind params: $1=hello unchanged, $2=limit rewritten to 15, $3=offset rewritten to 0. + // Bind parameters retain the client values used by the SQL expression. if let ProtocolMessage::Bind(bind) = &context.client_request.messages[1] { assert_eq!(bind.params_raw()[0].data.as_ref(), b"hello"); - assert_eq!(bind.params_raw()[1].data.as_ref(), b"15"); - assert_eq!(bind.params_raw()[2].data.as_ref(), b"0"); + assert_eq!(bind.params_raw()[1].data.as_ref(), b"10"); + assert_eq!(bind.params_raw()[2].data.as_ref(), b"5"); } else { panic!("expected Bind"); } diff --git a/pgdog/src/frontend/client/query_engine/test/rewrite_projection.rs b/pgdog/src/frontend/client/query_engine/test/rewrite_projection.rs new file mode 100644 index 000000000..52c8ee634 --- /dev/null +++ b/pgdog/src/frontend/client/query_engine/test/rewrite_projection.rs @@ -0,0 +1,416 @@ +use crate::backend::schema::Schema; +use crate::frontend::router::parser::rewrite::statement::plan::RewriteResult; +use crate::frontend::router::parser::rewrite::statement::projection; +use crate::frontend::router::parser::route::{Route, Shard, ShardWithPriority}; +use crate::frontend::{ + PreparedStatements, + router::parser::{Limit, OrderBy}, +}; + +use super::prelude::*; +use super::test_sharded_client; + +fn route(shard: Shard) -> Route { + Route::select( + ShardWithPriority::new_table(shard), + vec![], + Default::default(), + Limit::default(), + None, + ) +} + +#[tokio::test] +async fn direct_aggregate_keeps_base_sql() { + let sql = "SELECT AVG(price) FROM products"; + let mut client = test_sharded_client(); + client.client_request = ClientRequest::from(vec![ProtocolMessage::Query(Query::new(sql))]); + + let mut engine = QueryEngine::from_client(&client).unwrap(); + let mut context = QueryEngineContext::new(&mut client); + let result = engine.parse_and_rewrite(&mut context).await.unwrap(); + + let query = match &context.client_request.messages[0] { + ProtocolMessage::Query(query) => query, + _ => panic!("expected Query"), + }; + assert_eq!(query.query(), sql, "pre-route phase must not add helpers"); + + context.client_request.route = Some(route(Shard::Direct(0))); + projection::finalize_after_route( + context.client_request, + &Schema::default(), + result.as_ref().and_then(RewriteResult::offset_plan), + ) + .unwrap(); + + let query = match &context.client_request.messages[0] { + ProtocolMessage::Query(query) => query, + _ => panic!("expected Query"), + }; + assert_eq!(query.query(), sql); + assert!( + context + .client_request + .route() + .projection_rewrite_plan() + .is_noop() + ); +} + +#[tokio::test] +async fn cross_shard_aggregate_adds_and_tracks_helpers() { + let mut client = test_sharded_client(); + client.client_request = ClientRequest::from(vec![ProtocolMessage::Query(Query::new( + "SELECT AVG(price) FROM products", + ))]); + + let mut engine = QueryEngine::from_client(&client).unwrap(); + let mut context = QueryEngineContext::new(&mut client); + let result = engine.parse_and_rewrite(&mut context).await.unwrap(); + context.client_request.route = Some(route(Shard::All)); + + projection::finalize_after_route( + context.client_request, + &Schema::default(), + result.as_ref().and_then(RewriteResult::offset_plan), + ) + .unwrap(); + + let query = match &context.client_request.messages[0] { + ProtocolMessage::Query(query) => query, + _ => panic!("expected Query"), + }; + assert!(query.query().contains("__pgdog_count_col0")); + assert_eq!( + context + .client_request + .route() + .projection_rewrite_plan() + .aggregate_helpers() + .len(), + 1 + ); +} + +#[tokio::test] +async fn named_prepared_aggregate_uses_cross_shard_variant() { + let mut client = test_sharded_client(); + client.client_request = ClientRequest::from(vec![ + ProtocolMessage::Parse(Parse::named( + "avg_measurement", + "SELECT AVG(value) FROM measurements", + )), + ProtocolMessage::Bind(Bind::new_params("avg_measurement", &[])), + ProtocolMessage::Execute(Execute::new()), + ProtocolMessage::Sync(Sync), + ]); + + let mut engine = QueryEngine::from_client(&client).unwrap(); + let mut context = QueryEngineContext::new(&mut client); + engine.rewrite_extended(&mut context).unwrap(); + let base = match &context.client_request.messages[1] { + ProtocolMessage::Bind(bind) => bind.statement().to_owned(), + _ => panic!("expected Bind"), + }; + let result = engine.parse_and_rewrite(&mut context).await.unwrap(); + context.client_request.route = Some(route(Shard::All)); + + projection::finalize_after_route( + context.client_request, + &Schema::default(), + result.as_ref().and_then(RewriteResult::offset_plan), + ) + .unwrap(); + + let variant = format!("{base}_cross_shard"); + match &context.client_request.messages[0] { + ProtocolMessage::Parse(parse) => { + assert_eq!(parse.name(), variant); + assert!(parse.query().contains("__pgdog_count_col0")); + } + _ => panic!("expected Parse"), + } + match &context.client_request.messages[1] { + ProtocolMessage::Bind(bind) => assert_eq!(bind.statement(), variant), + _ => panic!("expected Bind"), + } + + let cache = PreparedStatements::global(); + let cache = cache.read(); + assert!( + !cache + .rewritten_parse(&base) + .unwrap() + .query() + .contains("__pgdog_") + ); + assert!( + cache + .rewritten_parse(&variant) + .unwrap() + .query() + .contains("__pgdog_count_col0") + ); +} + +#[tokio::test] +async fn named_prepared_direct_aggregate_keeps_base_variant() { + let mut client = test_sharded_client(); + client.client_request = ClientRequest::from(vec![ + ProtocolMessage::Parse(Parse::named( + "direct_avg", + "SELECT AVG(value) FROM direct_measurements", + )), + ProtocolMessage::Bind(Bind::new_params("direct_avg", &[])), + ProtocolMessage::Execute(Execute::new()), + ProtocolMessage::Sync(Sync), + ]); + + let mut engine = QueryEngine::from_client(&client).unwrap(); + let mut context = QueryEngineContext::new(&mut client); + engine.rewrite_extended(&mut context).unwrap(); + let base = match &context.client_request.messages[1] { + ProtocolMessage::Bind(bind) => bind.statement().to_owned(), + _ => panic!("expected Bind"), + }; + let result = engine.parse_and_rewrite(&mut context).await.unwrap(); + context.client_request.route = Some(route(Shard::Direct(0))); + + projection::finalize_after_route( + context.client_request, + &Schema::default(), + result.as_ref().and_then(RewriteResult::offset_plan), + ) + .unwrap(); + + match &context.client_request.messages[0] { + ProtocolMessage::Parse(parse) => { + assert_eq!(parse.name(), base); + assert!(!parse.query().contains("__pgdog_")); + } + _ => panic!("expected Parse"), + } + match &context.client_request.messages[1] { + ProtocolMessage::Bind(bind) => assert_eq!(bind.statement(), base), + _ => panic!("expected Bind"), + } + assert!( + PreparedStatements::global() + .read() + .rewritten_parse(&format!("{base}_cross_shard")) + .is_none() + ); +} + +#[tokio::test] +async fn cross_shard_order_by_projects_missing_sort_column() { + let sql = "SELECT id FROM products ORDER BY price"; + let mut client = test_sharded_client(); + client.client_request = ClientRequest::from(vec![ProtocolMessage::Query(Query::new(sql))]); + + let mut engine = QueryEngine::from_client(&client).unwrap(); + let mut context = QueryEngineContext::new(&mut client); + let result = engine.parse_and_rewrite(&mut context).await.unwrap(); + context.client_request.route = Some(Route::select( + ShardWithPriority::new_table(Shard::All), + vec![OrderBy::AscColumn("price".into())], + Default::default(), + Limit::default(), + None, + )); + + projection::finalize_after_route( + context.client_request, + &Schema::default(), + result.as_ref().and_then(RewriteResult::offset_plan), + ) + .unwrap(); + + let query = match &context.client_request.messages[0] { + ProtocolMessage::Query(query) => query, + _ => panic!("expected Query"), + }; + assert!(query.query().contains("price AS __pgdog_order_col0")); + assert_eq!( + context.client_request.route().order_by(), + &[OrderBy::Asc(2)] + ); + assert_eq!( + context + .client_request + .route() + .projection_rewrite_plan() + .order_by_helpers() + .len(), + 1 + ); +} + +#[tokio::test] +async fn aggregate_order_by_and_offset_compose_after_route() { + let mut client = test_sharded_client(); + client.client_request = ClientRequest::from(vec![ProtocolMessage::Query(Query::new( + "SELECT AVG(value) FROM measurements ORDER BY created_at LIMIT 10 OFFSET 5", + ))]); + + let mut engine = QueryEngine::from_client(&client).unwrap(); + let mut context = QueryEngineContext::new(&mut client); + let result = engine.parse_and_rewrite(&mut context).await.unwrap(); + context.client_request.route = Some(Route::select( + ShardWithPriority::new_table(Shard::All), + vec![OrderBy::AscColumn("created_at".into())], + Default::default(), + Limit::default(), + None, + )); + + projection::finalize_after_route( + context.client_request, + &Schema::default(), + result.as_ref().and_then(RewriteResult::offset_plan), + ) + .unwrap(); + result + .as_ref() + .unwrap() + .apply_after_route(context.client_request) + .unwrap(); + + let query = match &context.client_request.messages[0] { + ProtocolMessage::Query(query) => query, + _ => panic!("expected Query"), + }; + assert!(query.query().contains("__pgdog_count_col0")); + assert!(query.query().contains("created_at AS __pgdog_order_col0")); + assert!(query.query().contains("LIMIT 10 + 5")); + assert!(!query.query().contains("OFFSET")); + + let route = context.client_request.route(); + assert_eq!(route.order_by(), &[OrderBy::Asc(3)]); + assert_eq!( + route + .projection_rewrite_plan() + .drop_columns() + .collect::>(), + [1, 2] + ); + assert_eq!( + route.limit(), + &Limit { + limit: Some(10), + offset: Some(5), + } + ); +} + +#[tokio::test] +async fn split_anonymous_prepare_finalizes_saved_parse_on_execute() { + let mut client = test_sharded_client(); + client.client_request = ClientRequest::default(); + client + .client_request + .push(ProtocolMessage::Parse(Parse::new_anonymous( + "SELECT AVG(value) FROM split_measurements", + ))); + client + .client_request + .push(ProtocolMessage::Describe(Describe::new_statement(""))); + client.client_request.push(Flush.into()); + client.client_request.clear(); + client + .client_request + .push(ProtocolMessage::Bind(Bind::new_params("", &[]))); + client + .client_request + .push(ProtocolMessage::Execute(Execute::new())); + client.client_request.push(ProtocolMessage::Sync(Sync)); + + let mut engine = QueryEngine::from_client(&client).unwrap(); + let mut context = QueryEngineContext::new(&mut client); + let result = engine.parse_and_rewrite(&mut context).await.unwrap(); + context.client_request.route = Some(route(Shard::All)); + + projection::finalize_after_route( + context.client_request, + &Schema::default(), + result.as_ref().and_then(RewriteResult::offset_plan), + ) + .unwrap(); + + assert!( + context + .client_request + .last_parse + .as_ref() + .unwrap() + .query() + .contains("__pgdog_count_col0") + ); + assert!( + context + .client_request + .messages + .iter() + .all(|message| !matches!(message, ProtocolMessage::Parse(_))) + ); +} + +#[tokio::test] +async fn named_statement_can_switch_from_direct_to_cross_shard_variant() { + let mut client = test_sharded_client(); + client.client_request = ClientRequest::from(vec![ + ProtocolMessage::Parse(Parse::named( + "route_switch", + "SELECT AVG(value) FROM route_switch_measurements", + )), + ProtocolMessage::Sync(Sync), + ]); + + let base = { + let mut engine = QueryEngine::from_client(&client).unwrap(); + let mut context = QueryEngineContext::new(&mut client); + engine.rewrite_extended(&mut context).unwrap(); + let base = match &context.client_request.messages[0] { + ProtocolMessage::Parse(parse) => parse.name().to_owned(), + _ => panic!("expected Parse"), + }; + let result = engine.parse_and_rewrite(&mut context).await.unwrap(); + context.client_request.route = Some(route(Shard::Direct(0))); + projection::finalize_after_route( + context.client_request, + &Schema::default(), + result.as_ref().and_then(RewriteResult::offset_plan), + ) + .unwrap(); + base + }; + + client.client_request.clear(); + client + .client_request + .push(ProtocolMessage::Bind(Bind::new_params("route_switch", &[]))); + client + .client_request + .push(ProtocolMessage::Execute(Execute::new())); + client.client_request.push(ProtocolMessage::Sync(Sync)); + + let mut engine = QueryEngine::from_client(&client).unwrap(); + let mut context = QueryEngineContext::new(&mut client); + engine.rewrite_extended(&mut context).unwrap(); + let result = engine.parse_and_rewrite(&mut context).await.unwrap(); + context.client_request.route = Some(route(Shard::All)); + projection::finalize_after_route( + context.client_request, + &Schema::default(), + result.as_ref().and_then(RewriteResult::offset_plan), + ) + .unwrap(); + + match &context.client_request.messages[0] { + ProtocolMessage::Bind(bind) => { + assert_eq!(bind.statement(), format!("{base}_cross_shard")); + } + _ => panic!("expected Bind"), + } +} diff --git a/pgdog/src/frontend/prepared_statements/global_cache.rs b/pgdog/src/frontend/prepared_statements/global_cache.rs index 4908b10db..1138fb53d 100644 --- a/pgdog/src/frontend/prepared_statements/global_cache.rs +++ b/pgdog/src/frontend/prepared_statements/global_cache.rs @@ -30,6 +30,7 @@ use super::*; pub(crate) struct GlobalCache { statements: HashMap, names: HashMap, + cross_shard_variants: HashMap, unused: HashSet, counter: Counter, } @@ -39,6 +40,7 @@ impl MemoryUsage for GlobalCache { fn memory_usage(&self) -> usize { self.statements.memory_usage() + self.names.memory_usage() + + self.cross_shard_variants.memory_usage() + self.counter.memory_usage() + self.unused.capacity() * 1usize.memory_usage() } @@ -137,10 +139,44 @@ impl GlobalCache { } } + /// Create, or retrieve, the server-side prepared statement variant used + /// for cross-shard execution. Variants are derived data owned by the base + /// statement and therefore do not have an independent usage counter. + pub(crate) fn cross_shard_variant(&mut self, name: &str, query: &str) -> Option { + let variant_name = format!("{name}_cross_shard"); + if self.cross_shard_variants.contains_key(&variant_name) { + return Some(variant_name); + } + + let mut parse = self.rewritten_parse(name)?; + parse.rename(&variant_name); + parse.set_query(query); + let cache_key = CacheKey::Extended { + query: parse.query_ref(), + data_types: parse.data_types_ref(), + }; + self.cross_shard_variants.insert( + variant_name.clone(), + Statement { + stmt: StatementType::Parse { + parse, + rewrite: None, + }, + row_description: None, + cache_key, + }, + ); + + Some(variant_name) + } + /// Client sent a Describe for a prepared statement and received a RowDescription. /// We record the RowDescription for later use by the results decoder. pub(crate) fn insert_row_description(&mut self, name: &str, row_description: RowDescription) { - if let Some(entry) = self.names.get_mut(name) + if let Some(entry) = self + .names + .get_mut(name) + .or_else(|| self.cross_shard_variants.get_mut(name)) && entry.row_description.is_none() { entry.row_description = Some(row_description); @@ -176,8 +212,9 @@ impl GlobalCache { /// Used for preparing this statement on a server connection. /// pub(crate) fn rewritten_parse(&self, name: &str) -> Option { - self.names + self.cross_shard_variants .get(name) + .or_else(|| self.names.get(name)) .and_then(|p| p.rewritten_parse().clone().or(p.parse())) } @@ -195,7 +232,10 @@ impl GlobalCache { /// It can be used to decode results received from executing the prepared /// statement. pub(crate) fn row_description(&self, name: &str) -> Option { - self.names.get(name).and_then(|p| p.row_description.clone()) + self.cross_shard_variants + .get(name) + .or_else(|| self.names.get(name)) + .and_then(|p| p.row_description.clone()) } /// Number of prepared statements in the local cache. @@ -267,6 +307,8 @@ impl GlobalCache { fn remove(&mut self, name: &str) { if let Some(stmt) = self.names.remove(name) { self.statements.remove(stmt.cache_key()); + self.cross_shard_variants + .remove(&format!("{name}_cross_shard")); } } @@ -303,6 +345,7 @@ impl GlobalCache { #[cfg(test)] mod test { use super::*; + use crate::net::messages::Field; impl GlobalCache { /// Get the query string stored in the global cache @@ -344,6 +387,63 @@ mod test { assert_ne!(owned.as_ptr(), source.query_ref().as_ptr()); } + #[test] + fn cross_shard_variant_is_owned_by_base_statement() { + let mut cache = GlobalCache::default(); + let (_, base) = cache.insert(&Parse::named( + "client", + "SELECT AVG(value) FROM measurements", + )); + + let variant = cache + .cross_shard_variant( + &base, + "SELECT AVG(value), COUNT(value) AS __pgdog_count_col0 FROM measurements", + ) + .unwrap(); + assert_eq!(variant, format!("{base}_cross_shard")); + assert_eq!( + cache.rewritten_parse(&base).unwrap().query(), + "SELECT AVG(value) FROM measurements" + ); + assert!( + cache + .rewritten_parse(&variant) + .unwrap() + .query() + .contains("__pgdog_count_col0") + ); + assert_eq!(cache.len(), 1, "variant is not a second logical statement"); + + cache.close(&base); + assert_eq!(cache.close_unused(0), 1); + assert!(cache.rewritten_parse(&variant).is_none()); + } + + #[test] + fn cross_shard_variant_has_separate_row_description() { + let mut cache = GlobalCache::default(); + let (_, base) = cache.insert(&Parse::named( + "client", + "SELECT AVG(value) FROM measurements", + )); + let variant = cache + .cross_shard_variant( + &base, + "SELECT AVG(value), COUNT(value) AS __pgdog_count_col0 FROM measurements", + ) + .unwrap(); + + cache.insert_row_description(&base, RowDescription::new(&[Field::double("avg")])); + cache.insert_row_description( + &variant, + RowDescription::new(&[Field::double("avg"), Field::bigint("__pgdog_count_col0")]), + ); + + assert_eq!(cache.row_description(&base).unwrap().len(), 1); + assert_eq!(cache.row_description(&variant).unwrap().len(), 2); + } + #[test] fn test_prep_stmt_cache_close() { let mut cache = GlobalCache::default(); diff --git a/pgdog/src/frontend/router/parser/cache/ast.rs b/pgdog/src/frontend/router/parser/cache/ast.rs index 3a94302eb..5061b2738 100644 --- a/pgdog/src/frontend/router/parser/cache/ast.rs +++ b/pgdog/src/frontend/router/parser/cache/ast.rs @@ -4,6 +4,7 @@ use std::fmt::Debug; use std::ops::Deref; use std::time::Instant; +use once_cell::sync::OnceCell; use parking_lot::Mutex; use std::sync::Arc; use tracing::warn; @@ -14,6 +15,7 @@ use crate::backend::schema::Schema; use crate::frontend::PreparedStatements; use crate::frontend::router::parser::cache::AstQuery; use crate::frontend::router::parser::rewrite::statement::RewritePlan; +use crate::frontend::router::parser::rewrite::statement::projection::PostRouteRewrite; use crate::frontend::router::sharding::ShardOrLookup; use crate::net::parameter::ParameterValue; use crate::{backend::ShardingSchema, config::Role}; @@ -42,6 +44,8 @@ pub(crate) struct AstInner { pub(crate) stats: Mutex, /// Rewrite plan. pub(crate) rewrite_plan: RewritePlan, + /// Lazily generated SQL and response metadata for cross-shard execution. + pub(crate) post_route_rewrite: OnceCell>, /// Original query. pub(crate) query_without_comment: Arc, } @@ -53,6 +57,7 @@ impl AstInner { ast, stats: Mutex::new(Stats::new()), rewrite_plan: RewritePlan::default(), + post_route_rewrite: OnceCell::new(), query_without_comment: "".into(), } } @@ -124,6 +129,7 @@ impl Ast { stats: Mutex::new(stats), ast, rewrite_plan, + post_route_rewrite: OnceCell::new(), query_without_comment: query.query_without_comment.into(), }), }) diff --git a/pgdog/src/frontend/router/parser/query/select.rs b/pgdog/src/frontend/router/parser/query/select.rs index c552f72ee..b2b56bd84 100644 --- a/pgdog/src/frontend/router/parser/query/select.rs +++ b/pgdog/src/frontend/router/parser/query/select.rs @@ -16,7 +16,7 @@ impl QueryParser { /// pub(super) fn select( &mut self, - cached_ast: &Ast, + _cached_ast: &Ast, stmt: &nodes::SelectStmt, context: &mut QueryParserContext, ) -> Result { @@ -269,7 +269,7 @@ impl QueryParser { } } - let mut query = Route::select( + let query = Route::select( context.shards_calculator.shard().clone(), order_by, aggregates, @@ -277,11 +277,6 @@ impl QueryParser { distinct, ); - // Only rewrite if query is cross-shard. - if query.is_cross_shard() && context.shards > 1 { - query.set_projection_rewrite_plan(cached_ast.rewrite_plan.projection.clone()); - } - Ok(Command::Query( query .with_read(!writes) diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/aggregate/mod.rs b/pgdog/src/frontend/router/parser/rewrite/statement/aggregate/mod.rs index f03c7e3a4..06639ccbd 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/aggregate/mod.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/aggregate/mod.rs @@ -1,10 +1,6 @@ mod engine; pub(crate) use super::projection::AggregateHelper; -use super::{Error, RewritePlan, StatementRewrite}; -use crate::backend::schema::Schema; -use crate::frontend::router::parser::aggregate::Aggregate; -use pg_raw_parse::{make::MemoryToken, nodes::SelectStmtMut}; pub(crate) use engine::AggregatesRewrite; @@ -25,32 +21,3 @@ impl HelperKind { } } } - -impl StatementRewrite<'_> { - /// Add missing COUNT(*) and other helps when using aggregates. - pub(super) fn rewrite_aggregates<'a>( - &mut self, - select: &mut SelectStmtMut<'a, '_>, - mem: MemoryToken<'a>, - plan: &mut RewritePlan, - schema: &Schema, - ) -> Result<(), Error> { - if self.schema.shards == 1 { - return Ok(()); - } - - let aggregate = Aggregate::parse(select, schema); - if aggregate.is_empty() { - return Ok(()); - } - - let output = AggregatesRewrite::rewrite_select(select, mem, &aggregate); - if output.plan.is_noop() { - return Ok(()); - } - - plan.projection = output.plan; - self.rewritten = true; - Ok(()) - } -} diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs b/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs index db67ba6cc..075c1d760 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs @@ -181,8 +181,7 @@ impl<'a> StatementRewrite<'a> { return Err(err); } - if let NodeMut::SelectStmt(mut select) = stmt.stmt_mut() { - self.rewrite_aggregates(&mut select, mem, &mut plan, self.db_schema)?; + if let NodeMut::SelectStmt(select) = stmt.stmt_mut() { self.limit_offset(&select, &mut plan); } diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/offset.rs b/pgdog/src/frontend/router/parser/rewrite/statement/offset.rs index 8ce84b0d2..fd11fcf34 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/offset.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/offset.rs @@ -1,11 +1,10 @@ use std::ops::Deref; -use pg_raw_parse::{ConstValue, Node, Owned, StmtList, deparse, nodes}; +use pg_raw_parse::{ConstValue, Node, deparse, make, nodes}; use crate::frontend::ClientRequest; use crate::frontend::router::parser::Limit; use crate::net::ProtocolMessage; -use crate::net::messages::bind::{Format, Parameter}; use super::*; @@ -18,7 +17,7 @@ pub(crate) struct OffsetPlan { } impl OffsetPlan { - pub(super) fn apply_after_parser(&self, request: &mut ClientRequest) -> Result<(), Error> { + pub(super) fn apply_after_route(&self, request: &mut ClientRequest) -> Result<(), Error> { let route = match request.route.as_mut() { Some(route) => route, None => return Ok(()), @@ -59,53 +58,10 @@ impl OffsetPlan { ); } - let new_limit = limit_val.unwrap_or(0) + offset_val.unwrap_or(0); - - // Overwrite parameterized limit. - if self.limit.limit.is_none() { - let idx = self.limit_param - 1; - let fmt = bind.parameter_format(idx)?; - let param = match fmt { - Format::Binary => Parameter::new(&(new_limit as i64).to_be_bytes()), - Format::Text => { - Parameter::new(itoa::Buffer::new().format(new_limit).as_bytes()) - } - }; - bind.set_param(idx, param); - } - - // Overwrite parameterized offset. - if self.limit.offset.is_none() { - let idx = self.offset_param - 1; - let fmt = bind.parameter_format(idx)?; - let param = match fmt { - Format::Binary => Parameter::new(&0i64.to_be_bytes()), - Format::Text => Parameter::new(b"0"), - }; - bind.set_param(idx, param); - } break; } } - // Rewrite SQL if any value was a literal. - if self.limit.limit.is_some() || self.limit.offset.is_some() { - let new_limit = (limit_val.unwrap_or(0) + offset_val.unwrap_or(0)) as i32; - let ast = request.ast.as_ref().ok_or(Error::MissingAst)?; - - if let Some(rewritten) = rewrite_ast_limit_offset(&ast.ast, new_limit) { - let result = pg_raw_parse::deparse(&*rewritten)?; - let new_sql = result.as_str(); - for message in request.messages.iter_mut() { - match message { - ProtocolMessage::Query(q) => q.set_query(new_sql), - ProtocolMessage::Parse(p) => p.set_query(new_sql), - _ => {} - } - } - } - } - route.set_limit(Limit { limit: limit_val, offset: offset_val, @@ -114,7 +70,7 @@ impl OffsetPlan { Ok(()) } - /// `apply_after_parser` helper method for handling Prepare + Execute cases, where + /// `apply_after_route` helper method for handling Prepare + Execute cases, where /// we need to re-write limit / offset for multi-shard queries upon execution. fn handle_prepare_execute(&self, request: &mut ClientRequest) -> Result<(), Error> { // Assert expectations of what should've happened before this method was called @@ -220,19 +176,20 @@ fn extract_limit_value(node: Node<'_>) -> Option { } } -fn rewrite_ast_limit_offset(ast: &StmtList, new_limit: i32) -> Option> { - let Some(Node::SelectStmt(select)) = ast.stmts().next() else { - return None; - }; - - Some(make::owned(|mem| { - let mut select = mem.make_unique(select); - select - .as_mut() - .set_limit_count(mem.make_a_const(ConstValue::Integer(new_limit)).uncast()); - select.as_mut().set_limit_offset(mem.none()); - select - })) +pub(super) fn rewrite_select<'a>( + select: &mut nodes::SelectStmtMut<'a, '_>, + mem: make::MemoryToken<'a>, +) { + let limit = select.limit_count(); + let offset = select.limit_offset(); + let combined = mem.make_a_expr( + nodes::A_Expr_Kind::AEXPR_OP, + mem.make_list(&[mem.make_string(Some("+")).uncast()]), + mem.make_unique(limit), + mem.make_unique(offset), + ); + select.set_limit_count(combined.uncast()); + select.set_limit_offset(mem.none()); } impl StatementRewrite<'_> { @@ -267,7 +224,6 @@ mod tests { use crate::backend::schema::Schema; use crate::frontend::PreparedStatements; use crate::frontend::router::parser::StatementRewriteContext; - use crate::frontend::router::parser::cache::ast::Ast; use crate::frontend::router::parser::route::{Route, Shard, ShardWithPriority}; use crate::net::Parse; use crate::net::messages::Query; @@ -312,10 +268,6 @@ mod tests { ) } - fn make_ast(sql: &str) -> Ast { - Ast::new_record(sql).unwrap() - } - fn run_limit_offset(sql: &str, schema: &ShardingSchema) -> RewritePlan { let stmt = pg_raw_parse::parse(sql).unwrap(); let db_schema = Schema::default(); @@ -396,7 +348,7 @@ mod tests { } #[test] - fn test_apply_after_parser_literals_cross_shard() { + fn test_apply_after_route_literals_cross_shard() { let plan = OffsetPlan { limit: Limit { limit: Some(10), @@ -410,15 +362,13 @@ mod tests { "SELECT * FROM t LIMIT 10 OFFSET 5", ))]); request.route = Some(cross_shard_route()); - request.ast = Some(make_ast("SELECT * FROM t LIMIT 10 OFFSET 5")); - - plan.apply_after_parser(&mut request).unwrap(); + plan.apply_after_route(&mut request).unwrap(); let query = match &request.messages[0] { ProtocolMessage::Query(q) => q.query().to_owned(), _ => panic!("expected Query"), }; - assert_eq!(query, "SELECT * FROM t LIMIT 15"); + assert_eq!(query, "SELECT * FROM t LIMIT 10 OFFSET 5"); let route = request.route.unwrap(); assert_eq!(route.limit().limit, Some(10)); @@ -426,7 +376,7 @@ mod tests { } #[test] - fn test_apply_after_parser_params_cross_shard() { + fn test_apply_after_route_params_cross_shard() { let plan = OffsetPlan { limit: Limit { limit: None, @@ -442,11 +392,11 @@ mod tests { ))]); request.route = Some(cross_shard_route()); - plan.apply_after_parser(&mut request).unwrap(); + plan.apply_after_route(&mut request).unwrap(); if let ProtocolMessage::Bind(bind) = &request.messages[0] { - assert_eq!(bind.params_raw()[0].data.as_ref(), b"15"); - assert_eq!(bind.params_raw()[1].data.as_ref(), b"0"); + assert_eq!(bind.params_raw()[0].data.as_ref(), b"10"); + assert_eq!(bind.params_raw()[1].data.as_ref(), b"5"); } else { panic!("expected Bind"); } @@ -457,7 +407,7 @@ mod tests { } #[test] - fn test_apply_after_parser_single_shard_noop() { + fn test_apply_after_route_single_shard_noop() { let plan = OffsetPlan { limit: Limit { limit: Some(10), @@ -472,7 +422,7 @@ mod tests { ))]); request.route = Some(single_shard_route()); - plan.apply_after_parser(&mut request).unwrap(); + plan.apply_after_route(&mut request).unwrap(); let query = match &request.messages[0] { ProtocolMessage::Query(q) => q.query().to_owned(), @@ -482,7 +432,7 @@ mod tests { } #[test] - fn test_apply_after_parser_mixed_limit_literal_offset_param() { + fn test_apply_after_route_mixed_limit_literal_offset_param() { let plan = OffsetPlan { limit: Limit { limit: Some(10), @@ -497,12 +447,10 @@ mod tests { ProtocolMessage::Bind(Bind::new_params("s", &[Parameter::new(b"5")])), ]); request.route = Some(cross_shard_route()); - request.ast = Some(make_ast("SELECT * FROM t LIMIT 10 OFFSET $1")); - - plan.apply_after_parser(&mut request).unwrap(); + plan.apply_after_route(&mut request).unwrap(); if let ProtocolMessage::Bind(bind) = &request.messages[1] { - assert_eq!(bind.params_raw()[0].data.as_ref(), b"0"); + assert_eq!(bind.params_raw()[0].data.as_ref(), b"5"); } else { panic!("expected Bind"); } @@ -511,7 +459,7 @@ mod tests { ProtocolMessage::Parse(p) => p.query().to_owned(), _ => panic!("expected Parse"), }; - assert_eq!(sql, "SELECT * FROM t LIMIT 15"); + assert_eq!(sql, "SELECT * FROM t LIMIT 10 OFFSET $1"); let route = request.route.unwrap(); assert_eq!(route.limit().limit, Some(10)); diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs b/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs index 2ed998446..3898df9bb 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs @@ -7,9 +7,7 @@ use super::super::ee; use super::insert::{build_resolved_split_requests, build_split_requests}; use super::nextval::SequenceCall; use super::offset::OffsetPlan; -use super::{ - Error, InsertSplit, PrepareExecute, ShardingKeyUpdate, projection::ProjectionRewritePlan, -}; +use super::{Error, InsertSplit, PrepareExecute, ShardingKeyUpdate}; #[derive(Debug, Clone, PartialEq, Eq)] pub(crate) enum GeneratedId { @@ -50,9 +48,6 @@ pub(crate) struct RewritePlan { /// multiple queries. pub(crate) insert_split: Vec, - /// Temporary result columns needed while merging cross-shard results. - pub(crate) projection: ProjectionRewritePlan, - /// Sharding key is being updated, we need to execute /// a multi-step plan. pub(crate) sharding_key_update: Option, @@ -69,11 +64,18 @@ pub(crate) enum RewriteResult { } impl RewriteResult { - pub(crate) fn apply_after_parser(&self, request: &mut ClientRequest) -> Result<(), Error> { + pub(crate) fn offset_plan(&self) -> Option<&OffsetPlan> { + match self { + Self::InPlace { offset } => offset.as_ref(), + _ => None, + } + } + + pub(crate) fn apply_after_route(&self, request: &mut ClientRequest) -> Result<(), Error> { match self { Self::InPlace { offset: Some(offset), - } => offset.apply_after_parser(request), + } => offset.apply_after_route(request), _ => Ok(()), } } @@ -90,7 +92,6 @@ impl RewritePlan { && self.stmt.is_none() && self.prepare_rewrites.is_empty() && self.insert_split.is_empty() - && self.projection.is_noop() && self.sharding_key_update.is_none() && self.offset.is_none() } diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs b/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs index 3ce95be35..e465367be 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs @@ -1,4 +1,12 @@ -use super::aggregate::HelperKind; +use super::Error; +use super::aggregate::{AggregatesRewrite, HelperKind}; +use super::offset::{self, OffsetPlan}; +use crate::backend::schema::Schema; +use crate::frontend::router::parser::{Aggregate, OrderBy}; +use crate::frontend::{ClientRequest, PreparedStatements}; +use crate::net::ProtocolMessage; +use pg_raw_parse::{Node, StmtList, make}; +use std::sync::Arc; /// Aggregate function projected temporarily for cross-shard merging. #[derive(Debug, Clone, PartialEq)] @@ -62,12 +70,213 @@ pub(crate) struct RewriteOutput { pub(crate) plan: ProjectionRewritePlan, } +#[derive(Debug, Clone)] +pub(crate) struct PostRouteRewrite { + sql: Arc, + plan: ProjectionRewritePlan, +} + impl RewriteOutput { pub(crate) fn new(plan: ProjectionRewritePlan) -> Self { Self { plan } } } +/// Add temporary columns needed to merge a cross-shard SELECT. +/// +/// This deliberately operates on a copy of the cached AST. The cached AST is +/// the route-independent representation and must remain suitable for direct +/// execution. +pub(crate) fn finalize_after_route( + request: &mut ClientRequest, + schema: &Schema, + offset_plan: Option<&OffsetPlan>, +) -> Result<(), Error> { + if !request.route().is_cross_shard() { + return Ok(()); + } + + let Some(ast) = request.ast.as_ref() else { + return Ok(()); + }; + let rewrite_offset = offset_plan.is_some_and(|plan| !plan.prepare_execute); + let order_by = request.route().order_by(); + let Some(rewrite) = ast + .post_route_rewrite + .get_or_try_init(|| build(&ast.ast, schema, order_by, rewrite_offset))? + else { + return Ok(()); + }; + let base_name = request.messages.iter().find_map(|message| match message { + ProtocolMessage::Parse(parse) if !parse.anonymous() => Some(parse.name()), + ProtocolMessage::Bind(bind) if !bind.anonymous() => Some(bind.statement()), + ProtocolMessage::Describe(describe) if describe.is_statement() && !describe.anonymous() => { + Some(describe.statement()) + } + _ => None, + }); + let variant = base_name.and_then(|name| { + PreparedStatements::global() + .write() + .cross_shard_variant(name, &rewrite.sql) + }); + + for message in &mut request.messages { + match message { + ProtocolMessage::Query(query) => query.set_query(&rewrite.sql), + ProtocolMessage::Parse(parse) => { + parse.set_query(&rewrite.sql); + if let Some(variant) = &variant { + parse.rename(variant); + } + } + ProtocolMessage::Bind(bind) => { + if let Some(variant) = &variant { + bind.rename(variant); + } + } + ProtocolMessage::Describe(describe) if describe.is_statement() => { + if let Some(variant) = &variant { + describe.rename(variant); + } + } + _ => {} + } + } + if let Some(parse) = request.last_parse.as_mut() { + parse.set_query(&rewrite.sql); + } + if !rewrite.plan.is_noop() + && let Some(route) = request.route.as_mut() + { + route.set_projection_rewrite_plan(rewrite.plan.clone()); + let mut order_by = route.order_by().to_vec(); + for helper in rewrite.plan.order_by_helpers() { + let Some(sort) = order_by.get_mut(helper.sort_position) else { + continue; + }; + *sort = if sort.asc() { + OrderBy::Asc(helper.projected_column + 1) + } else { + OrderBy::Desc(helper.projected_column + 1) + }; + } + route.set_order_by(order_by); + } + + Ok(()) +} + +fn build( + ast: &StmtList, + schema: &Schema, + order_by: &[OrderBy], + rewrite_offset: bool, +) -> Result, Error> { + let Some(Node::SelectStmt(select)) = ast.stmts().next() else { + return Ok(None); + }; + + let aggregate = Aggregate::parse(select, schema); + if aggregate.is_empty() && order_by.is_empty() && !rewrite_offset { + return Ok(None); + } + + let mut plan = ProjectionRewritePlan::default(); + let rewritten = make::owned(|mem| { + let mut select = mem.make_unique(select); + if !aggregate.is_empty() { + plan = AggregatesRewrite::rewrite_select(&mut select.as_mut(), mem, &aggregate).plan; + } + rewrite_order_by(&mut select.as_mut(), mem, order_by, &mut plan); + if rewrite_offset { + offset::rewrite_select(&mut select.as_mut(), mem); + } + select + }); + if plan.is_noop() && !rewrite_offset { + return Ok(None); + } + let sql: Arc = pg_raw_parse::deparse(&*rewritten)?.as_str().into(); + + Ok(Some(PostRouteRewrite { sql, plan })) +} + +fn rewrite_order_by<'a>( + select: &mut pg_raw_parse::nodes::SelectStmtMut<'a, '_>, + mem: make::MemoryToken<'a>, + order_by: &[OrderBy], + plan: &mut ProjectionRewritePlan, +) { + let mut helpers = Vec::new(); + let mut sort_position = 0; + for sort in select.sort_clause() { + let node = sort.node(); + let Some(order) = order_by.get(sort_position) else { + break; + }; + let supported = matches!( + (node, order), + (Node::A_Const(_), OrderBy::Asc(_) | OrderBy::Desc(_)) + | ( + Node::ColumnRef(_), + OrderBy::AscColumn(_) | OrderBy::DescColumn(_) + ) + | (Node::A_Expr(_), OrderBy::AscVectorL2Column(_, _)) + ); + if !supported { + continue; + } + + let needs_helper = match node { + Node::ColumnRef(column) => { + let Some(name) = column + .fields() + .into_iter() + .next_back() + .and_then(Node::as_str) + else { + continue; + }; + !select.target_list().iter().any(|target| { + target.name() == Some(name) + || matches!( + target.val(), + Node::ColumnRef(projected) + if projected.fields().into_iter().next_back().and_then(Node::as_str) + == Some(name) + ) + }) + } + Node::A_Expr(_) => matches!(order, OrderBy::AscVectorL2Column(_, _)), + _ => false, + }; + let current_sort_position = sort_position; + sort_position += 1; + if !needs_helper { + continue; + } + + let projected_column = select.target_list().len() + helpers.len(); + let alias = format!("__pgdog_order_col{current_sort_position}"); + helpers.push(mem.make_res_target( + Some(&alias), + mem.empty(), + mem.make_unique(node).uncast(), + )); + plan.add_order_by_helper(OrderByHelper { + sort_position: current_sort_position, + projected_column, + }); + } + + if !helpers.is_empty() { + select + .target_list_mut() + .extend(mem, mem.make_list(&helpers)); + } +} + #[cfg(test)] mod tests { use super::*; diff --git a/pgdog/src/frontend/router/parser/route.rs b/pgdog/src/frontend/router/parser/route.rs index 727aaaf69..8f6c6b0f5 100644 --- a/pgdog/src/frontend/router/parser/route.rs +++ b/pgdog/src/frontend/router/parser/route.rs @@ -253,6 +253,10 @@ impl Route { &self.order_by } + pub(crate) fn set_order_by(&mut self, order_by: Vec) { + self.order_by = order_by; + } + pub(crate) fn aggregate(&self) -> &Aggregate { &self.aggregate } diff --git a/pgdog/src/net/messages/bind.rs b/pgdog/src/net/messages/bind.rs index 330fa9aab..e9b6ec0bb 100644 --- a/pgdog/src/net/messages/bind.rs +++ b/pgdog/src/net/messages/bind.rs @@ -302,18 +302,6 @@ impl Bind { self.original = None; } - /// Overwrite an existing parameter at the given index. - /// Returns `false` if the index is out of bounds. - pub(crate) fn set_param(&mut self, index: usize, param: Parameter) -> bool { - if let Some(slot) = self.params.get_mut(index) { - *slot = param; - self.original = None; - true - } else { - false - } - } - /// Get the effective format for new parameters. pub(crate) fn default_param_format(&self) -> Format { if self.codes.len() == 1 { @@ -554,20 +542,6 @@ mod test { } } - #[test] - fn test_set_param() { - let mut bind = Bind::new_params("test", &[Parameter::new(b"10"), Parameter::new(b"5")]); - assert_eq!(bind.params_raw()[0].data.as_ref(), b"10"); - assert_eq!(bind.params_raw()[1].data.as_ref(), b"5"); - - bind.set_param(0, Parameter::new(b"15")); - bind.set_param(1, Parameter::new(b"0")); - - assert_eq!(bind.params_raw()[0].data.as_ref(), b"15"); - assert_eq!(bind.params_raw()[1].data.as_ref(), b"0"); - assert_eq!(bind.params_raw().len(), 2); - } - #[test] fn test_large_parameter_count_round_trip() { let count = 35_000; From 06b3a3899833d24b37818644d9aa2d4d135cbfdd Mon Sep 17 00:00:00 2001 From: Nupur Agrawal Date: Mon, 14 Sep 2026 23:45:26 +0530 Subject: [PATCH 03/14] separate order_by --- .../router/parser/rewrite/statement/mod.rs | 1 + .../parser/rewrite/statement/order_by.rs | 131 ++++++++++++++++++ .../parser/rewrite/statement/projection.rs | 78 +---------- 3 files changed, 134 insertions(+), 76 deletions(-) create mode 100644 pgdog/src/frontend/router/parser/rewrite/statement/order_by.rs diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs b/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs index 075c1d760..98b00d0b8 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs @@ -13,6 +13,7 @@ pub(crate) mod error; pub(crate) mod insert; pub(crate) mod nextval; pub(crate) mod offset; +pub(crate) mod order_by; pub(crate) mod plan; pub(crate) mod projection; pub(crate) mod simple_prepared; diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/order_by.rs b/pgdog/src/frontend/router/parser/rewrite/statement/order_by.rs new file mode 100644 index 000000000..b5cb06f9b --- /dev/null +++ b/pgdog/src/frontend/router/parser/rewrite/statement/order_by.rs @@ -0,0 +1,131 @@ +use pg_raw_parse::{Node, make, nodes}; + +use crate::frontend::router::parser::OrderBy; + +use super::projection::{OrderByHelper, ProjectionRewritePlan}; + +/// Project ORDER BY expressions that are missing from the SELECT list +/// so cross-shard results can be sorted, then stripped. +pub(super) fn rewrite_select<'a>( + select: &mut nodes::SelectStmtMut<'a, '_>, + mem: make::MemoryToken<'a>, + order_by: &[OrderBy], + plan: &mut ProjectionRewritePlan, +) { + let mut helpers = Vec::new(); + let mut sort_position = 0; + for sort in select.sort_clause() { + let node = sort.node(); + let Some(order) = order_by.get(sort_position) else { + break; + }; + let supported = matches!( + (node, order), + (Node::A_Const(_), OrderBy::Asc(_) | OrderBy::Desc(_)) + | ( + Node::ColumnRef(_), + OrderBy::AscColumn(_) | OrderBy::DescColumn(_) + ) + | (Node::A_Expr(_), OrderBy::AscVectorL2Column(_, _)) + ); + if !supported { + continue; + } + + let needs_helper = match node { + Node::ColumnRef(column) => { + let Some(name) = column + .fields() + .into_iter() + .next_back() + .and_then(Node::as_str) + else { + continue; + }; + !select.target_list().iter().any(|target| { + target.name() == Some(name) + || matches!( + target.val(), + Node::ColumnRef(projected) + if projected.fields().into_iter().next_back().and_then(Node::as_str) + == Some(name) + ) + }) + } + Node::A_Expr(_) => matches!(order, OrderBy::AscVectorL2Column(_, _)), + _ => false, + }; + let current_sort_position = sort_position; + sort_position += 1; + if !needs_helper { + continue; + } + + let projected_column = select.target_list().len() + helpers.len(); + let alias = format!("__pgdog_order_col{current_sort_position}"); + helpers.push(mem.make_res_target( + Some(&alias), + mem.empty(), + mem.make_unique(node).uncast(), + )); + plan.add_order_by_helper(OrderByHelper { + sort_position: current_sort_position, + projected_column, + }); + } + + if !helpers.is_empty() { + select + .target_list_mut() + .extend(mem, mem.make_list(&helpers)); + } +} + +#[cfg(test)] +mod tests { + use super::*; + use pg_raw_parse::{Node, make}; + + fn rewrite(sql: &str, order_by: Vec) -> (String, ProjectionRewritePlan) { + let ast = pg_raw_parse::parse(sql).unwrap(); + let mut plan = ProjectionRewritePlan::default(); + let rewritten = make::owned(|mem| { + let Node::SelectStmt(select) = ast.stmts().next().unwrap() else { + panic!("expected SELECT"); + }; + let mut select = mem.make_unique(select); + rewrite_select(&mut select.as_mut(), mem, &order_by, &mut plan); + select + }); + ( + pg_raw_parse::deparse(&*rewritten) + .unwrap() + .as_str() + .to_owned(), + plan, + ) + } + + #[test] + fn projects_missing_sort_column() { + let (sql, plan) = rewrite( + "SELECT id FROM products ORDER BY price", + vec![OrderBy::AscColumn("price".into())], + ); + + assert!(sql.contains("price AS __pgdog_order_col0")); + assert_eq!(plan.order_by_helpers().len(), 1); + assert_eq!(plan.order_by_helpers()[0].projected_column, 1); + } + + #[test] + fn skips_already_projected_sort_column() { + let (sql, plan) = rewrite( + "SELECT id, price FROM products ORDER BY price", + vec![OrderBy::AscColumn("price".into())], + ); + + assert!(!sql.contains("__pgdog_order_col")); + assert!(plan.is_noop()); + } +} diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs b/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs index e465367be..51cb15cd5 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs @@ -1,6 +1,7 @@ use super::Error; use super::aggregate::{AggregatesRewrite, HelperKind}; use super::offset::{self, OffsetPlan}; +use super::order_by; use crate::backend::schema::Schema; use crate::frontend::router::parser::{Aggregate, OrderBy}; use crate::frontend::{ClientRequest, PreparedStatements}; @@ -188,7 +189,7 @@ fn build( if !aggregate.is_empty() { plan = AggregatesRewrite::rewrite_select(&mut select.as_mut(), mem, &aggregate).plan; } - rewrite_order_by(&mut select.as_mut(), mem, order_by, &mut plan); + order_by::rewrite_select(&mut select.as_mut(), mem, order_by, &mut plan); if rewrite_offset { offset::rewrite_select(&mut select.as_mut(), mem); } @@ -202,81 +203,6 @@ fn build( Ok(Some(PostRouteRewrite { sql, plan })) } -fn rewrite_order_by<'a>( - select: &mut pg_raw_parse::nodes::SelectStmtMut<'a, '_>, - mem: make::MemoryToken<'a>, - order_by: &[OrderBy], - plan: &mut ProjectionRewritePlan, -) { - let mut helpers = Vec::new(); - let mut sort_position = 0; - for sort in select.sort_clause() { - let node = sort.node(); - let Some(order) = order_by.get(sort_position) else { - break; - }; - let supported = matches!( - (node, order), - (Node::A_Const(_), OrderBy::Asc(_) | OrderBy::Desc(_)) - | ( - Node::ColumnRef(_), - OrderBy::AscColumn(_) | OrderBy::DescColumn(_) - ) - | (Node::A_Expr(_), OrderBy::AscVectorL2Column(_, _)) - ); - if !supported { - continue; - } - - let needs_helper = match node { - Node::ColumnRef(column) => { - let Some(name) = column - .fields() - .into_iter() - .next_back() - .and_then(Node::as_str) - else { - continue; - }; - !select.target_list().iter().any(|target| { - target.name() == Some(name) - || matches!( - target.val(), - Node::ColumnRef(projected) - if projected.fields().into_iter().next_back().and_then(Node::as_str) - == Some(name) - ) - }) - } - Node::A_Expr(_) => matches!(order, OrderBy::AscVectorL2Column(_, _)), - _ => false, - }; - let current_sort_position = sort_position; - sort_position += 1; - if !needs_helper { - continue; - } - - let projected_column = select.target_list().len() + helpers.len(); - let alias = format!("__pgdog_order_col{current_sort_position}"); - helpers.push(mem.make_res_target( - Some(&alias), - mem.empty(), - mem.make_unique(node).uncast(), - )); - plan.add_order_by_helper(OrderByHelper { - sort_position: current_sort_position, - projected_column, - }); - } - - if !helpers.is_empty() { - select - .target_list_mut() - .extend(mem, mem.make_list(&helpers)); - } -} - #[cfg(test)] mod tests { use super::*; From 4507ac56594f821bc50362b8814b122b40ef4062 Mon Sep 17 00:00:00 2001 From: Nupur Agrawal Date: Tue, 15 Sep 2026 00:13:50 +0530 Subject: [PATCH 04/14] fix --- pgdog/src/backend/pool/connection/buffer.rs | 13 ++- .../pool/connection/multi_shard/mod.rs | 12 ++- .../src/frontend/client/query_engine/query.rs | 2 +- .../rewrite/statement/aggregate/engine.rs | 25 +++--- .../router/parser/rewrite/statement/offset.rs | 21 ++++- .../parser/rewrite/statement/order_by.rs | 81 +++++++++++++------ .../parser/rewrite/statement/projection.rs | 68 +++++++++++++--- 7 files changed, 163 insertions(+), 59 deletions(-) diff --git a/pgdog/src/backend/pool/connection/buffer.rs b/pgdog/src/backend/pool/connection/buffer.rs index df7a63ca1..a0c16bb9c 100644 --- a/pgdog/src/backend/pool/connection/buffer.rs +++ b/pgdog/src/backend/pool/connection/buffer.rs @@ -151,18 +151,25 @@ impl Buffer { buffer }; - Self::drop_helper_columns(&mut rows, plan); + Self::drop_helper_columns(&mut rows, plan, decoder); self.buffer = rows; Ok(()) } - fn drop_helper_columns(rows: &mut VecDeque, plan: &ProjectionRewritePlan) { + fn drop_helper_columns( + rows: &mut VecDeque, + plan: &ProjectionRewritePlan, + decoder: &Decoder, + ) { if plan.is_noop() { return; } - let drop = plan.drop_columns().collect(); + let drop = plan.drop_columns(decoder.row_description()); + if drop.is_empty() { + return; + } for row in rows.iter_mut() { row.drop_columns(&drop); diff --git a/pgdog/src/backend/pool/connection/multi_shard/mod.rs b/pgdog/src/backend/pool/connection/multi_shard/mod.rs index 13372c5a4..03f981e66 100644 --- a/pgdog/src/backend/pool/connection/multi_shard/mod.rs +++ b/pgdog/src/backend/pool/connection/multi_shard/mod.rs @@ -1,6 +1,6 @@ //! Multi-shard connection state. -use std::collections::VecDeque; +use std::collections::{BTreeSet, VecDeque}; use crate::{ frontend::router::Route, @@ -312,11 +312,15 @@ impl MultiShard { // Only send it to the client once all shards sent it, // so we don't get early requests from clients. let plan = self.route.projection_rewrite_plan(); - if plan.is_noop() { + let drop = if plan.is_noop() { + BTreeSet::new() + } else { + plan.drop_columns(&rd) + }; + if drop.is_empty() { forward = Some(message); } else { - let client_rd = rd.drop_columns(plan.drop_columns()); - forward = Some(client_rd.message()); + forward = Some(rd.drop_columns(drop).message()); } // The next statement describes a different result set. diff --git a/pgdog/src/frontend/client/query_engine/query.rs b/pgdog/src/frontend/client/query_engine/query.rs index 838726bf0..c00719062 100644 --- a/pgdog/src/frontend/client/query_engine/query.rs +++ b/pgdog/src/frontend/client/query_engine/query.rs @@ -545,5 +545,5 @@ impl ExplainResponseState { pub(crate) fn should_emit(&self) -> bool { self.supported && !self.annotated - } + } } diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/aggregate/engine.rs b/pgdog/src/frontend/router/parser/rewrite/statement/aggregate/engine.rs index d4b062a8d..c5d57f6cb 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/aggregate/engine.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/aggregate/engine.rs @@ -40,15 +40,13 @@ impl AggregatesRewrite { .into_iter() .map(move |spec| (target, spec)) }) - .enumerate() - .map(|(idx, (target, HelperSpec { func, kind }))| { + .map(|(target, HelperSpec { func, kind })| { let helper_alias = format!("__pgdog_{}_col{}", kind.alias_suffix(), target.column()); let node = mem.make_res_target(Some(&helper_alias), mem.empty(), func.uncast()); plan.add_aggregate_helper(AggregateHelper { target_column: target.column(), - projected_column: select.target_list().len() + idx, distinct: target.is_distinct(), kind, alias: helper_alias, @@ -198,11 +196,13 @@ mod tests { fn rewrite_engine_adds_helper() { let (ast, output) = rewrite("SELECT AVG(price) FROM menu"); assert!(!output.plan.is_noop()); - assert_eq!(output.plan.drop_columns().collect::>(), &[1]); + assert_eq!( + output.plan.aliases().collect::>(), + ["__pgdog_count_col0"] + ); assert_eq!(output.plan.aggregate_helpers().len(), 1); let helper = &output.plan.aggregate_helpers()[0]; assert_eq!(helper.target_column, 0); - assert_eq!(helper.projected_column, 1); assert!(!helper.distinct); assert!(matches!(helper.kind, HelperKind::Count)); @@ -219,11 +219,13 @@ mod tests { #[test] fn rewrite_engine_handles_mismatched_pair() { let (ast, output) = rewrite("SELECT COUNT(price::numeric), AVG(price) FROM menu"); - assert_eq!(output.plan.drop_columns().collect::>(), &[2]); + assert_eq!( + output.plan.aliases().collect::>(), + ["__pgdog_count_col1"] + ); assert_eq!(output.plan.aggregate_helpers().len(), 1); let helper = &output.plan.aggregate_helpers()[0]; assert_eq!(helper.target_column, 1); - assert_eq!(helper.projected_column, 2); assert!(!helper.distinct); assert!(matches!(helper.kind, HelperKind::Count)); @@ -242,17 +244,18 @@ mod tests { #[test] fn rewrite_engine_multiple_avg_helpers() { let (ast, output) = rewrite("SELECT AVG(price), AVG(discount) FROM menu"); - assert_eq!(output.plan.drop_columns().collect::>(), &[2, 3]); + assert_eq!( + output.plan.aliases().collect::>(), + ["__pgdog_count_col0", "__pgdog_count_col1"] + ); assert_eq!(output.plan.aggregate_helpers().len(), 2); let helper_price = &output.plan.aggregate_helpers()[0]; assert_eq!(helper_price.target_column, 0); - assert_eq!(helper_price.projected_column, 2); assert!(matches!(helper_price.kind, HelperKind::Count)); let helper_discount = &output.plan.aggregate_helpers()[1]; assert_eq!(helper_discount.target_column, 1); - assert_eq!(helper_discount.projected_column, 3); assert!(matches!(helper_discount.kind, HelperKind::Count)); let aggregate = Aggregate::parse(&ast, &Default::default()); @@ -271,7 +274,7 @@ mod tests { fn rewrite_engine_stddev_helpers() { let (ast, output) = rewrite("SELECT STDDEV(price) FROM menu"); assert!(!output.plan.is_noop()); - assert_eq!(output.plan.drop_columns().collect::>(), &[1, 2, 3]); + assert_eq!(output.plan.aliases().count(), 3); assert_eq!(output.plan.aggregate_helpers().len(), 3); let kinds: Vec = output diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/offset.rs b/pgdog/src/frontend/router/parser/rewrite/statement/offset.rs index fd11fcf34..07ffd4492 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/offset.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/offset.rs @@ -176,17 +176,30 @@ fn extract_limit_value(node: Node<'_>) -> Option { } } +/// `$1 + $2` is ambiguous to Postgres when both sides are untyped parameters, +/// so spell out the type both operands would have had as LIMIT/OFFSET. +fn to_bigint<'a>(node: Node<'_>, mem: make::MemoryToken<'a>) -> make::Unique<'a, Node<'a>> { + mem.make_type_cast( + mem.make_unique(node).uncast(), + mem.make_list(&[ + mem.make_string(Some("pg_catalog")), + mem.make_string(Some("int8")), + ]), + ) + .uncast() +} + pub(super) fn rewrite_select<'a>( select: &mut nodes::SelectStmtMut<'a, '_>, mem: make::MemoryToken<'a>, ) { - let limit = select.limit_count(); - let offset = select.limit_offset(); + let limit = to_bigint(select.limit_count(), mem); + let offset = to_bigint(select.limit_offset(), mem); let combined = mem.make_a_expr( nodes::A_Expr_Kind::AEXPR_OP, mem.make_list(&[mem.make_string(Some("+")).uncast()]), - mem.make_unique(limit), - mem.make_unique(offset), + limit, + offset, ); select.set_limit_count(combined.uncast()); select.set_limit_offset(mem.none()); diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/order_by.rs b/pgdog/src/frontend/router/parser/rewrite/statement/order_by.rs index b5cb06f9b..7959e455c 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/order_by.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/order_by.rs @@ -4,6 +4,23 @@ use crate::frontend::router::parser::OrderBy; use super::projection::{OrderByHelper, ProjectionRewritePlan}; +/// A `*` in the select list already projects every column of its table, so +/// nothing sorted by those columns needs a helper. +fn projects_star(select: &nodes::SelectStmtMut<'_, '_>) -> bool { + select.target_list().iter().any(|target| { + matches!(target.val(), Node::ColumnRef(column) + if column.fields().into_iter().any(|field| matches!(field, Node::A_Star(_)))) + }) +} + +fn projects_column(select: &nodes::SelectStmtMut<'_, '_>, name: &str) -> bool { + select.target_list().iter().any(|target| { + target.name() == Some(name) + || matches!(target.val(), Node::ColumnRef(projected) + if projected.fields().into_iter().next_back().and_then(Node::as_str) == Some(name)) + }) +} + /// Project ORDER BY expressions that are missing from the SELECT list /// so cross-shard results can be sorted, then stripped. pub(super) fn rewrite_select<'a>( @@ -12,6 +29,10 @@ pub(super) fn rewrite_select<'a>( order_by: &[OrderBy], plan: &mut ProjectionRewritePlan, ) { + if projects_star(select) { + return; + } + let mut helpers = Vec::new(); let mut sort_position = 0; for sort in select.sort_clause() { @@ -32,36 +53,26 @@ pub(super) fn rewrite_select<'a>( continue; } + let current_sort_position = sort_position; + sort_position += 1; + let needs_helper = match node { - Node::ColumnRef(column) => { - let Some(name) = column - .fields() - .into_iter() - .next_back() - .and_then(Node::as_str) - else { - continue; - }; - !select.target_list().iter().any(|target| { - target.name() == Some(name) - || matches!( - target.val(), - Node::ColumnRef(projected) - if projected.fields().into_iter().next_back().and_then(Node::as_str) - == Some(name) - ) - }) - } + Node::ColumnRef(column) => match column + .fields() + .into_iter() + .next_back() + .and_then(Node::as_str) + { + Some(name) => !projects_column(select, name), + None => false, + }, Node::A_Expr(_) => matches!(order, OrderBy::AscVectorL2Column(_, _)), _ => false, }; - let current_sort_position = sort_position; - sort_position += 1; if !needs_helper { continue; } - let projected_column = select.target_list().len() + helpers.len(); let alias = format!("__pgdog_order_col{current_sort_position}"); helpers.push(mem.make_res_target( Some(&alias), @@ -70,7 +81,7 @@ pub(super) fn rewrite_select<'a>( )); plan.add_order_by_helper(OrderByHelper { sort_position: current_sort_position, - projected_column, + alias, }); } @@ -115,7 +126,7 @@ mod tests { assert!(sql.contains("price AS __pgdog_order_col0")); assert_eq!(plan.order_by_helpers().len(), 1); - assert_eq!(plan.order_by_helpers()[0].projected_column, 1); + assert_eq!(plan.order_by_helpers()[0].alias, "__pgdog_order_col0"); } #[test] @@ -128,4 +139,26 @@ mod tests { assert!(!sql.contains("__pgdog_order_col")); assert!(plan.is_noop()); } + + #[test] + fn skips_star_select() { + let (sql, plan) = rewrite( + "SELECT * FROM products ORDER BY id", + vec![OrderBy::AscColumn("id".into())], + ); + + assert!(!sql.contains("__pgdog_order_col")); + assert!(plan.is_noop()); + } + + #[test] + fn skips_qualified_star_select() { + let (sql, plan) = rewrite( + "SELECT products.* FROM products ORDER BY products.id", + vec![OrderBy::AscColumn("id".into())], + ); + + assert!(!sql.contains("__pgdog_order_col")); + assert!(plan.is_noop()); + } } diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs b/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs index 51cb15cd5..9b29aa346 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs @@ -6,14 +6,15 @@ use crate::backend::schema::Schema; use crate::frontend::router::parser::{Aggregate, OrderBy}; use crate::frontend::{ClientRequest, PreparedStatements}; use crate::net::ProtocolMessage; +use crate::net::messages::RowDescription; use pg_raw_parse::{Node, StmtList, make}; +use std::collections::BTreeSet; use std::sync::Arc; /// Aggregate function projected temporarily for cross-shard merging. #[derive(Debug, Clone, PartialEq)] pub(crate) struct AggregateHelper { pub(crate) target_column: usize, - pub(crate) projected_column: usize, pub(crate) distinct: bool, pub(crate) kind: HelperKind, pub(crate) alias: String, @@ -23,7 +24,7 @@ pub(crate) struct AggregateHelper { #[derive(Debug, Clone, PartialEq)] pub(crate) struct OrderByHelper { pub(crate) sort_position: usize, - pub(crate) projected_column: usize, + pub(crate) alias: String, } /// Temporary result columns required while merging cross-shard results. @@ -38,17 +39,29 @@ impl ProjectionRewritePlan { self.aggregate_helpers.is_empty() && self.order_by_helpers.is_empty() } - pub(crate) fn drop_columns(&self) -> impl Iterator + '_ { + /// Aliases of every column we added to the statement. + pub(crate) fn aliases(&self) -> impl Iterator { self.aggregate_helpers .iter() - .map(|helper| helper.projected_column) + .map(|helper| helper.alias.as_str()) .chain( self.order_by_helpers .iter() - .map(|helper| helper.projected_column), + .map(|helper| helper.alias.as_str()), ) } + /// Positions of our columns in the result set. + /// + /// Resolved against the `RowDescription` Postgres sent back: the position + /// in the statement is not the position in the result, because `*` expands + /// to however many columns the table has. + pub(crate) fn drop_columns(&self, row_description: &RowDescription) -> BTreeSet { + self.aliases() + .filter_map(|alias| row_description.field_index(alias)) + .collect() + } + pub(crate) fn aggregate_helpers(&self) -> &[AggregateHelper] { &self.aggregate_helpers } @@ -156,10 +169,14 @@ pub(crate) fn finalize_after_route( let Some(sort) = order_by.get_mut(helper.sort_position) else { continue; }; - *sort = if sort.asc() { - OrderBy::Asc(helper.projected_column + 1) - } else { - OrderBy::Desc(helper.projected_column + 1) + // Sort by name: the alias is resolved against the RowDescription, + // so it survives `*` expanding to any number of columns. + *sort = match &*sort { + OrderBy::AscVectorL2Column(_, vector) => { + OrderBy::AscVectorL2Column(helper.alias.clone(), vector.clone()) + } + sort if sort.asc() => OrderBy::AscColumn(helper.alias.clone()), + _ => OrderBy::DescColumn(helper.alias.clone()), }; } route.set_order_by(order_by); @@ -212,19 +229,46 @@ mod tests { let mut plan = ProjectionRewritePlan::default(); plan.add_aggregate_helper(AggregateHelper { target_column: 0, - projected_column: 1, distinct: false, kind: HelperKind::Count, alias: "__pgdog_count_col0".into(), }); plan.add_order_by_helper(OrderByHelper { sort_position: 0, - projected_column: 2, + alias: "__pgdog_order_col0".into(), }); assert!(!plan.is_noop()); - assert_eq!(plan.drop_columns().collect::>(), [1, 2]); + assert_eq!( + plan.aliases().collect::>(), + ["__pgdog_count_col0", "__pgdog_order_col0"] + ); assert_eq!(plan.aggregate_helpers().len(), 1); assert_eq!(plan.order_by_helpers().len(), 1); } + + /// `SELECT *` expands to N columns, so our helpers are not where the + /// statement says they are. Resolve them by name instead. + #[test] + fn drop_columns_resolves_positions_from_row_description() { + use crate::net::messages::Field; + + let mut plan = ProjectionRewritePlan::default(); + plan.add_order_by_helper(OrderByHelper { + sort_position: 0, + alias: "__pgdog_order_col0".into(), + }); + + let row_description = RowDescription::new(&[ + Field::bigint("id"), + Field::text("value"), + Field::bigint("__pgdog_order_col0"), + ]); + + assert_eq!( + plan.drop_columns(&row_description), + BTreeSet::from([2]), + "helper sits after the expanded columns" + ); + } } From ca39461121a43a6f317dab01765c842fccf6eeef Mon Sep 17 00:00:00 2001 From: Nupur Agrawal Date: Tue, 15 Sep 2026 00:18:05 +0530 Subject: [PATCH 05/14] fmt --- pgdog/src/frontend/client/query_engine/query.rs | 2 +- .../client/query_engine/test/rewrite_projection.rs | 11 +++++++---- 2 files changed, 8 insertions(+), 5 deletions(-) diff --git a/pgdog/src/frontend/client/query_engine/query.rs b/pgdog/src/frontend/client/query_engine/query.rs index c00719062..838726bf0 100644 --- a/pgdog/src/frontend/client/query_engine/query.rs +++ b/pgdog/src/frontend/client/query_engine/query.rs @@ -545,5 +545,5 @@ impl ExplainResponseState { pub(crate) fn should_emit(&self) -> bool { self.supported && !self.annotated - } + } } diff --git a/pgdog/src/frontend/client/query_engine/test/rewrite_projection.rs b/pgdog/src/frontend/client/query_engine/test/rewrite_projection.rs index 52c8ee634..6d4f8272a 100644 --- a/pgdog/src/frontend/client/query_engine/test/rewrite_projection.rs +++ b/pgdog/src/frontend/client/query_engine/test/rewrite_projection.rs @@ -283,17 +283,20 @@ async fn aggregate_order_by_and_offset_compose_after_route() { }; assert!(query.query().contains("__pgdog_count_col0")); assert!(query.query().contains("created_at AS __pgdog_order_col0")); - assert!(query.query().contains("LIMIT 10 + 5")); + assert!(query.query().contains("LIMIT 10::bigint + 5::bigint")); assert!(!query.query().contains("OFFSET")); let route = context.client_request.route(); - assert_eq!(route.order_by(), &[OrderBy::Asc(3)]); + assert_eq!( + route.order_by(), + &[OrderBy::AscColumn("__pgdog_order_col0".into())] + ); assert_eq!( route .projection_rewrite_plan() - .drop_columns() + .aliases() .collect::>(), - [1, 2] + ["__pgdog_count_col0", "__pgdog_order_col0"] ); assert_eq!( route.limit(), From d6075ddab4790e3797dcc62a44a24cffec0eee75 Mon Sep 17 00:00:00 2001 From: Nupur Agrawal Date: Tue, 15 Sep 2026 00:25:16 +0530 Subject: [PATCH 06/14] more fixes --- pgdog/src/backend/pool/connection/buffer.rs | 13 +--- .../pool/connection/multi_shard/mod.rs | 12 ++-- .../query_engine/test/rewrite_projection.rs | 9 +-- .../rewrite/statement/aggregate/engine.rs | 25 +++---- .../parser/rewrite/statement/order_by.rs | 5 +- .../parser/rewrite/statement/projection.rs | 72 +++++-------------- 6 files changed, 40 insertions(+), 96 deletions(-) diff --git a/pgdog/src/backend/pool/connection/buffer.rs b/pgdog/src/backend/pool/connection/buffer.rs index a0c16bb9c..df7a63ca1 100644 --- a/pgdog/src/backend/pool/connection/buffer.rs +++ b/pgdog/src/backend/pool/connection/buffer.rs @@ -151,25 +151,18 @@ impl Buffer { buffer }; - Self::drop_helper_columns(&mut rows, plan, decoder); + Self::drop_helper_columns(&mut rows, plan); self.buffer = rows; Ok(()) } - fn drop_helper_columns( - rows: &mut VecDeque, - plan: &ProjectionRewritePlan, - decoder: &Decoder, - ) { + fn drop_helper_columns(rows: &mut VecDeque, plan: &ProjectionRewritePlan) { if plan.is_noop() { return; } - let drop = plan.drop_columns(decoder.row_description()); - if drop.is_empty() { - return; - } + let drop = plan.drop_columns().collect(); for row in rows.iter_mut() { row.drop_columns(&drop); diff --git a/pgdog/src/backend/pool/connection/multi_shard/mod.rs b/pgdog/src/backend/pool/connection/multi_shard/mod.rs index 03f981e66..13372c5a4 100644 --- a/pgdog/src/backend/pool/connection/multi_shard/mod.rs +++ b/pgdog/src/backend/pool/connection/multi_shard/mod.rs @@ -1,6 +1,6 @@ //! Multi-shard connection state. -use std::collections::{BTreeSet, VecDeque}; +use std::collections::VecDeque; use crate::{ frontend::router::Route, @@ -312,15 +312,11 @@ impl MultiShard { // Only send it to the client once all shards sent it, // so we don't get early requests from clients. let plan = self.route.projection_rewrite_plan(); - let drop = if plan.is_noop() { - BTreeSet::new() - } else { - plan.drop_columns(&rd) - }; - if drop.is_empty() { + if plan.is_noop() { forward = Some(message); } else { - forward = Some(rd.drop_columns(drop).message()); + let client_rd = rd.drop_columns(plan.drop_columns()); + forward = Some(client_rd.message()); } // The next statement describes a different result set. diff --git a/pgdog/src/frontend/client/query_engine/test/rewrite_projection.rs b/pgdog/src/frontend/client/query_engine/test/rewrite_projection.rs index 6d4f8272a..4596f53ec 100644 --- a/pgdog/src/frontend/client/query_engine/test/rewrite_projection.rs +++ b/pgdog/src/frontend/client/query_engine/test/rewrite_projection.rs @@ -287,16 +287,13 @@ async fn aggregate_order_by_and_offset_compose_after_route() { assert!(!query.query().contains("OFFSET")); let route = context.client_request.route(); - assert_eq!( - route.order_by(), - &[OrderBy::AscColumn("__pgdog_order_col0".into())] - ); + assert_eq!(route.order_by(), &[OrderBy::Asc(3)]); assert_eq!( route .projection_rewrite_plan() - .aliases() + .drop_columns() .collect::>(), - ["__pgdog_count_col0", "__pgdog_order_col0"] + [1, 2] ); assert_eq!( route.limit(), diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/aggregate/engine.rs b/pgdog/src/frontend/router/parser/rewrite/statement/aggregate/engine.rs index c5d57f6cb..d4b062a8d 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/aggregate/engine.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/aggregate/engine.rs @@ -40,13 +40,15 @@ impl AggregatesRewrite { .into_iter() .map(move |spec| (target, spec)) }) - .map(|(target, HelperSpec { func, kind })| { + .enumerate() + .map(|(idx, (target, HelperSpec { func, kind }))| { let helper_alias = format!("__pgdog_{}_col{}", kind.alias_suffix(), target.column()); let node = mem.make_res_target(Some(&helper_alias), mem.empty(), func.uncast()); plan.add_aggregate_helper(AggregateHelper { target_column: target.column(), + projected_column: select.target_list().len() + idx, distinct: target.is_distinct(), kind, alias: helper_alias, @@ -196,13 +198,11 @@ mod tests { fn rewrite_engine_adds_helper() { let (ast, output) = rewrite("SELECT AVG(price) FROM menu"); assert!(!output.plan.is_noop()); - assert_eq!( - output.plan.aliases().collect::>(), - ["__pgdog_count_col0"] - ); + assert_eq!(output.plan.drop_columns().collect::>(), &[1]); assert_eq!(output.plan.aggregate_helpers().len(), 1); let helper = &output.plan.aggregate_helpers()[0]; assert_eq!(helper.target_column, 0); + assert_eq!(helper.projected_column, 1); assert!(!helper.distinct); assert!(matches!(helper.kind, HelperKind::Count)); @@ -219,13 +219,11 @@ mod tests { #[test] fn rewrite_engine_handles_mismatched_pair() { let (ast, output) = rewrite("SELECT COUNT(price::numeric), AVG(price) FROM menu"); - assert_eq!( - output.plan.aliases().collect::>(), - ["__pgdog_count_col1"] - ); + assert_eq!(output.plan.drop_columns().collect::>(), &[2]); assert_eq!(output.plan.aggregate_helpers().len(), 1); let helper = &output.plan.aggregate_helpers()[0]; assert_eq!(helper.target_column, 1); + assert_eq!(helper.projected_column, 2); assert!(!helper.distinct); assert!(matches!(helper.kind, HelperKind::Count)); @@ -244,18 +242,17 @@ mod tests { #[test] fn rewrite_engine_multiple_avg_helpers() { let (ast, output) = rewrite("SELECT AVG(price), AVG(discount) FROM menu"); - assert_eq!( - output.plan.aliases().collect::>(), - ["__pgdog_count_col0", "__pgdog_count_col1"] - ); + assert_eq!(output.plan.drop_columns().collect::>(), &[2, 3]); assert_eq!(output.plan.aggregate_helpers().len(), 2); let helper_price = &output.plan.aggregate_helpers()[0]; assert_eq!(helper_price.target_column, 0); + assert_eq!(helper_price.projected_column, 2); assert!(matches!(helper_price.kind, HelperKind::Count)); let helper_discount = &output.plan.aggregate_helpers()[1]; assert_eq!(helper_discount.target_column, 1); + assert_eq!(helper_discount.projected_column, 3); assert!(matches!(helper_discount.kind, HelperKind::Count)); let aggregate = Aggregate::parse(&ast, &Default::default()); @@ -274,7 +271,7 @@ mod tests { fn rewrite_engine_stddev_helpers() { let (ast, output) = rewrite("SELECT STDDEV(price) FROM menu"); assert!(!output.plan.is_noop()); - assert_eq!(output.plan.aliases().count(), 3); + assert_eq!(output.plan.drop_columns().collect::>(), &[1, 2, 3]); assert_eq!(output.plan.aggregate_helpers().len(), 3); let kinds: Vec = output diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/order_by.rs b/pgdog/src/frontend/router/parser/rewrite/statement/order_by.rs index 7959e455c..2c51d8211 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/order_by.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/order_by.rs @@ -73,6 +73,7 @@ pub(super) fn rewrite_select<'a>( continue; } + let projected_column = select.target_list().len() + helpers.len(); let alias = format!("__pgdog_order_col{current_sort_position}"); helpers.push(mem.make_res_target( Some(&alias), @@ -81,7 +82,7 @@ pub(super) fn rewrite_select<'a>( )); plan.add_order_by_helper(OrderByHelper { sort_position: current_sort_position, - alias, + projected_column, }); } @@ -126,7 +127,7 @@ mod tests { assert!(sql.contains("price AS __pgdog_order_col0")); assert_eq!(plan.order_by_helpers().len(), 1); - assert_eq!(plan.order_by_helpers()[0].alias, "__pgdog_order_col0"); + assert_eq!(plan.order_by_helpers()[0].projected_column, 1); } #[test] diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs b/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs index 9b29aa346..60d3b811b 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs @@ -6,15 +6,14 @@ use crate::backend::schema::Schema; use crate::frontend::router::parser::{Aggregate, OrderBy}; use crate::frontend::{ClientRequest, PreparedStatements}; use crate::net::ProtocolMessage; -use crate::net::messages::RowDescription; use pg_raw_parse::{Node, StmtList, make}; -use std::collections::BTreeSet; use std::sync::Arc; /// Aggregate function projected temporarily for cross-shard merging. #[derive(Debug, Clone, PartialEq)] pub(crate) struct AggregateHelper { pub(crate) target_column: usize, + pub(crate) projected_column: usize, pub(crate) distinct: bool, pub(crate) kind: HelperKind, pub(crate) alias: String, @@ -24,7 +23,7 @@ pub(crate) struct AggregateHelper { #[derive(Debug, Clone, PartialEq)] pub(crate) struct OrderByHelper { pub(crate) sort_position: usize, - pub(crate) alias: String, + pub(crate) projected_column: usize, } /// Temporary result columns required while merging cross-shard results. @@ -39,29 +38,21 @@ impl ProjectionRewritePlan { self.aggregate_helpers.is_empty() && self.order_by_helpers.is_empty() } - /// Aliases of every column we added to the statement. - pub(crate) fn aliases(&self) -> impl Iterator { + /// Positions of the columns we added to the select list. + /// + /// Helpers are only ever added to an explicit select list, so the position + /// in the statement is the position in the result set. + pub(crate) fn drop_columns(&self) -> impl Iterator + '_ { self.aggregate_helpers .iter() - .map(|helper| helper.alias.as_str()) + .map(|helper| helper.projected_column) .chain( self.order_by_helpers .iter() - .map(|helper| helper.alias.as_str()), + .map(|helper| helper.projected_column), ) } - /// Positions of our columns in the result set. - /// - /// Resolved against the `RowDescription` Postgres sent back: the position - /// in the statement is not the position in the result, because `*` expands - /// to however many columns the table has. - pub(crate) fn drop_columns(&self, row_description: &RowDescription) -> BTreeSet { - self.aliases() - .filter_map(|alias| row_description.field_index(alias)) - .collect() - } - pub(crate) fn aggregate_helpers(&self) -> &[AggregateHelper] { &self.aggregate_helpers } @@ -169,14 +160,10 @@ pub(crate) fn finalize_after_route( let Some(sort) = order_by.get_mut(helper.sort_position) else { continue; }; - // Sort by name: the alias is resolved against the RowDescription, - // so it survives `*` expanding to any number of columns. - *sort = match &*sort { - OrderBy::AscVectorL2Column(_, vector) => { - OrderBy::AscVectorL2Column(helper.alias.clone(), vector.clone()) - } - sort if sort.asc() => OrderBy::AscColumn(helper.alias.clone()), - _ => OrderBy::DescColumn(helper.alias.clone()), + *sort = if sort.asc() { + OrderBy::Asc(helper.projected_column + 1) + } else { + OrderBy::Desc(helper.projected_column + 1) }; } route.set_order_by(order_by); @@ -229,46 +216,19 @@ mod tests { let mut plan = ProjectionRewritePlan::default(); plan.add_aggregate_helper(AggregateHelper { target_column: 0, + projected_column: 1, distinct: false, kind: HelperKind::Count, alias: "__pgdog_count_col0".into(), }); plan.add_order_by_helper(OrderByHelper { sort_position: 0, - alias: "__pgdog_order_col0".into(), + projected_column: 2, }); assert!(!plan.is_noop()); - assert_eq!( - plan.aliases().collect::>(), - ["__pgdog_count_col0", "__pgdog_order_col0"] - ); + assert_eq!(plan.drop_columns().collect::>(), [1, 2]); assert_eq!(plan.aggregate_helpers().len(), 1); assert_eq!(plan.order_by_helpers().len(), 1); } - - /// `SELECT *` expands to N columns, so our helpers are not where the - /// statement says they are. Resolve them by name instead. - #[test] - fn drop_columns_resolves_positions_from_row_description() { - use crate::net::messages::Field; - - let mut plan = ProjectionRewritePlan::default(); - plan.add_order_by_helper(OrderByHelper { - sort_position: 0, - alias: "__pgdog_order_col0".into(), - }); - - let row_description = RowDescription::new(&[ - Field::bigint("id"), - Field::text("value"), - Field::bigint("__pgdog_order_col0"), - ]); - - assert_eq!( - plan.drop_columns(&row_description), - BTreeSet::from([2]), - "helper sits after the expanded columns" - ); - } } From 8757c755b276a1b2c25f9161eacf08476e742fcb Mon Sep 17 00:00:00 2001 From: Nupur Agrawal Date: Tue, 15 Sep 2026 01:11:27 +0530 Subject: [PATCH 07/14] describe croos shard prepared variants --- pgdog/src/backend/prepared_statements.rs | 122 +++++++++++++++++- pgdog/src/backend/server.rs | 10 +- .../client/query_engine/route_query.rs | 1 - .../query_engine/test/rewrite_offset.rs | 12 +- .../prepared_statements/global_cache.rs | 9 +- .../parser/rewrite/statement/aggregate/mod.rs | 1 - .../parser/rewrite/statement/order_by.rs | 4 - .../parser/rewrite/statement/projection.rs | 12 -- 8 files changed, 137 insertions(+), 34 deletions(-) diff --git a/pgdog/src/backend/prepared_statements.rs b/pgdog/src/backend/prepared_statements.rs index 14a93270d..bd55d3bdd 100644 --- a/pgdog/src/backend/prepared_statements.rs +++ b/pgdog/src/backend/prepared_statements.rs @@ -8,8 +8,8 @@ use std::{ use crate::{ frontend::{self, prepared_statements::GlobalCache}, net::{ - Close, CloseComplete, FromBytes, Message, ParseComplete, Protocol, ProtocolMessage, - ToBytes, + Close, CloseComplete, Describe, FromBytes, Message, ParseComplete, Protocol, + ProtocolMessage, ToBytes, messages::{ParameterDescription, RowDescription, parse::Parse}, }, state::State, @@ -60,6 +60,7 @@ pub(super) struct Prepare { /// Some if statement was prepared previously, but has expired since close: Option, parse: ProtocolMessage, + describe: Option, } impl Prepare { @@ -72,8 +73,15 @@ impl Prepare { &self.parse } + pub(super) fn describe(&self) -> Option<&ProtocolMessage> { + self.describe.as_ref() + } + fn anonymize(&mut self) { self.parse.anonymize(); + if let Some(describe) = &mut self.describe { + describe.anonymize(); + } } } @@ -88,6 +96,10 @@ pub(super) enum HandleResult { rewrite: ProtocolMessage, }, PrependProtocolMessage(ProtocolMessage), + PrependProtocolMessageRewrite { + prepend: ProtocolMessage, + rewrite: ProtocolMessage, + }, } /// Server-specific prepared statements. @@ -177,6 +189,7 @@ impl PreparedStatements { } self.state.add_ignore('1'); self.parses.push_back(bind.statement().to_string()); + message.describe = self.internal_describe(bind.statement()); self.state.add('2'); if self.config.level.rewrite_anonymous() { message.anonymize(); @@ -192,12 +205,25 @@ impl PreparedStatements { } None => { + let mut describe = self.internal_describe(bind.statement()); self.state.add('2'); if self.config.level.rewrite_anonymous() { let mut bind = bind.clone(); bind.anonymize(); + if let Some(describe) = &mut describe { + describe.anonymize(); + } + if let Some(describe) = describe { + return Ok(HandleResult::PrependProtocolMessageRewrite { + prepend: describe, + rewrite: ProtocolMessage::Bind(bind), + }); + } return Ok(HandleResult::Rewrite(ProtocolMessage::Bind(bind))); } + if let Some(describe) = describe { + return Ok(HandleResult::PrependProtocolMessage(describe)); + } } } } else { @@ -518,9 +544,27 @@ impl PreparedStatements { // it still holds. close: expired.then(|| ProtocolMessage::Close(Close::named(name))), parse: ProtocolMessage::Parse(parse), + describe: None, })) } + fn internal_describe(&mut self, name: &str) -> Option { + if self.describes.iter().any(|describe| describe == name) + || !self + .global_cache + .read() + .cross_shard_variant_needs_row_description(name) + { + return None; + } + + self.describes.push_back(name.to_owned()); + self.state.add_ignore(ExecutionCode::DescriptionOrNothing); + self.state.add_ignore(ExecutionCode::DescriptionOrNothing); + + Some(ProtocolMessage::Describe(Describe::new_statement(name))) + } + /// The server has prepared this statement already. pub(crate) fn contains(&mut self, name: &str) -> bool { self.local_cache.promote(name) @@ -677,8 +721,9 @@ pub(crate) mod test { use crate::frontend::PreparedStatements as FrontendPreparedStatements; use crate::net::{ Bind, CommandComplete, Describe, ErrorResponse, Execute, Message, Parse, - Prepare as SimplePrepare, ProtocolMessage, Query, Sync, bind::Parameter, - messages::ReadyForQuery, + Prepare as SimplePrepare, ProtocolMessage, Query, Sync, + bind::Parameter, + messages::{ReadyForQuery, row_description::Field}, }; use pgdog_config::PreparedStatementsLevel; @@ -762,6 +807,75 @@ pub(crate) mod test { assert_parse_without_close!(ps.handle(&bind(&name)).unwrap()); } + #[test] + fn bind_describes_a_pending_statement_when_result_shape_is_missing() { + let base = insert_global( + "internal_description", + "SELECT AVG(value) FROM measurements_internal_rd", + ); + let name = FrontendPreparedStatements::global() + .write() + .cross_shard_variant( + &base, + "SELECT AVG(value), COUNT(value) AS __pgdog_count_col0 \ + FROM measurements_internal_rd", + ) + .unwrap(); + let mut ps = new_extended(); + let parse = ProtocolMessage::Parse(ps.parse(&name).unwrap()); + + assert_eq!(ps.handle(&parse).unwrap(), HandleResult::Forward); + + let HandleResult::PrependProtocolMessage(ProtocolMessage::Describe(describe)) = + ps.handle(&bind(&name)).unwrap() + else { + panic!("expected an internal Describe before Bind"); + }; + assert_eq!(describe.statement(), name); + + let mut parse_complete = Message::new(ParseComplete.to_bytes()); + assert!(ps.forward(&mut parse_complete).unwrap()); + + let mut parameter_description = Message::new(ParameterDescription::empty().to_bytes()); + assert!(!ps.forward(&mut parameter_description).unwrap()); + + let row_description = + RowDescription::new(&[Field::double("avg"), Field::bigint("__pgdog_count_col0")]); + let mut message = Message::new(row_description.to_bytes()); + assert!(!ps.forward(&mut message).unwrap()); + + assert_eq!( + ps.global_cache.read().row_description(&name).unwrap().len(), + 2 + ); + } + + #[test] + fn bind_prepares_and_describes_an_unseen_cross_shard_variant() { + let base = insert_global( + "prepare_and_describe", + "SELECT AVG(value) FROM measurements_prepare_and_describe", + ); + let name = FrontendPreparedStatements::global() + .write() + .cross_shard_variant( + &base, + "SELECT AVG(value), COUNT(value) AS __pgdog_count_col0 \ + FROM measurements_prepare_and_describe", + ) + .unwrap(); + let mut ps = new_extended(); + + let HandleResult::Prepend(prepare) = ps.handle(&bind(&name)).unwrap() else { + panic!("expected Parse and internal Describe before Bind"); + }; + assert!(matches!(prepare.parse(), ProtocolMessage::Parse(parse) if parse.name() == name)); + assert!( + matches!(prepare.describe(), Some(ProtocolMessage::Describe(describe)) + if describe.statement() == name) + ); + } + #[test] fn bind_leaves_a_statement_within_its_ttl_alone() { let name = insert_global("ttl_fresh", "SELECT $1::bigint"); diff --git a/pgdog/src/backend/server.rs b/pgdog/src/backend/server.rs index 132ad3441..ce3ee920a 100644 --- a/pgdog/src/backend/server.rs +++ b/pgdog/src/backend/server.rs @@ -503,6 +503,10 @@ impl Server { self.send_stream(protocol_message).await?; self.send_stream(message).await?; } + HandleResult::PrependProtocolMessageRewrite { prepend, rewrite } => { + self.send_stream(prepend).await?; + self.send_stream(rewrite).await?; + } HandleResult::PrependRewrite { prepend, rewrite } => { self.send_prepare(prepend).await?; self.send_stream(rewrite).await?; @@ -518,7 +522,11 @@ impl Server { self.send_stream(close).await?; } - self.send_stream(prepare.parse()).await + self.send_stream(prepare.parse()).await?; + if let Some(describe) = prepare.describe() { + self.send_stream(describe).await?; + } + Ok(()) } /// Send a message to Postgres and force us to ignore its respose in [`Self::read`]. diff --git a/pgdog/src/frontend/client/query_engine/route_query.rs b/pgdog/src/frontend/client/query_engine/route_query.rs index 032a322e6..5b4b519df 100644 --- a/pgdog/src/frontend/client/query_engine/route_query.rs +++ b/pgdog/src/frontend/client/query_engine/route_query.rs @@ -154,7 +154,6 @@ impl QueryEngine { rewrite_result.and_then(RewriteResult::offset_plan), )?; - // Resolve route-dependent values, e.g. offset/limit. if let Some(rewrite_result) = rewrite_result { rewrite_result.apply_after_route(context.client_request)?; } diff --git a/pgdog/src/frontend/client/query_engine/test/rewrite_offset.rs b/pgdog/src/frontend/client/query_engine/test/rewrite_offset.rs index 3fc55a53f..6eea15746 100644 --- a/pgdog/src/frontend/client/query_engine/test/rewrite_offset.rs +++ b/pgdog/src/frontend/client/query_engine/test/rewrite_offset.rs @@ -144,7 +144,6 @@ async fn test_offset_with_unique_id_simple() { "should have bigint cast: {rewritten_sql}" ); - // Finalize with a cross-shard route. context.client_request.route = Some(cross_shard_route()); projection::finalize_after_route( context.client_request, @@ -169,12 +168,12 @@ async fn test_offset_with_unique_id_simple() { "unique_id rewrite must survive post-route finalization: {final_sql}" ); assert!( - final_sql.contains("::bigint"), - "bigint cast must survive: {final_sql}" + final_sql.contains(")::bigint FROM"), + "unique_id bigint cast must survive: {final_sql}" ); // LIMIT/OFFSET must be rewritten for cross-shard. assert!( - final_sql.contains("LIMIT 10 + 5"), + final_sql.contains("LIMIT 10::bigint + 5::bigint"), "LIMIT should request limit+offset rows: {final_sql}" ); assert!( @@ -217,7 +216,6 @@ async fn test_offset_with_unique_id_extended() { "SELECT $4::bigint, $1 FROM test LIMIT $2 OFFSET $3" ); - // Post-route finalization rewrites the SQL without changing Bind values. context.client_request.route = Some(cross_shard_route()); projection::finalize_after_route( context.client_request, @@ -231,17 +229,15 @@ async fn test_offset_with_unique_id_extended() { .apply_after_route(context.client_request) .unwrap(); - // SQL uses a stable expression suitable for prepared-statement caching. let final_sql = match &context.client_request.messages[0] { ProtocolMessage::Parse(p) => p.query().to_owned(), _ => panic!("expected Parse"), }; assert_eq!( - final_sql, "SELECT $4::bigint, $1 FROM test LIMIT $2 + $3", + final_sql, "SELECT $4::bigint, $1 FROM test LIMIT $2::bigint + $3::bigint", "SQL must push down limit+offset" ); - // Bind parameters retain the client values used by the SQL expression. if let ProtocolMessage::Bind(bind) = &context.client_request.messages[1] { assert_eq!(bind.params_raw()[0].data.as_ref(), b"hello"); assert_eq!(bind.params_raw()[1].data.as_ref(), b"10"); diff --git a/pgdog/src/frontend/prepared_statements/global_cache.rs b/pgdog/src/frontend/prepared_statements/global_cache.rs index 1138fb53d..e7015f89a 100644 --- a/pgdog/src/frontend/prepared_statements/global_cache.rs +++ b/pgdog/src/frontend/prepared_statements/global_cache.rs @@ -139,9 +139,6 @@ impl GlobalCache { } } - /// Create, or retrieve, the server-side prepared statement variant used - /// for cross-shard execution. Variants are derived data owned by the base - /// statement and therefore do not have an independent usage counter. pub(crate) fn cross_shard_variant(&mut self, name: &str, query: &str) -> Option { let variant_name = format!("{name}_cross_shard"); if self.cross_shard_variants.contains_key(&variant_name) { @@ -238,6 +235,12 @@ impl GlobalCache { .and_then(|p| p.row_description.clone()) } + pub(crate) fn cross_shard_variant_needs_row_description(&self, name: &str) -> bool { + self.cross_shard_variants + .get(name) + .is_some_and(|statement| statement.row_description.is_none()) + } + /// Number of prepared statements in the local cache. pub(crate) fn len(&self) -> usize { self.statements.len() diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/aggregate/mod.rs b/pgdog/src/frontend/router/parser/rewrite/statement/aggregate/mod.rs index 06639ccbd..07da4bb48 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/aggregate/mod.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/aggregate/mod.rs @@ -4,7 +4,6 @@ pub(crate) use super::projection::AggregateHelper; pub(crate) use engine::AggregatesRewrite; -/// Type of aggregate function added to the result set. #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub(crate) enum HelperKind { Count, diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/order_by.rs b/pgdog/src/frontend/router/parser/rewrite/statement/order_by.rs index 2c51d8211..e1f2cc4dd 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/order_by.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/order_by.rs @@ -4,8 +4,6 @@ use crate::frontend::router::parser::OrderBy; use super::projection::{OrderByHelper, ProjectionRewritePlan}; -/// A `*` in the select list already projects every column of its table, so -/// nothing sorted by those columns needs a helper. fn projects_star(select: &nodes::SelectStmtMut<'_, '_>) -> bool { select.target_list().iter().any(|target| { matches!(target.val(), Node::ColumnRef(column) @@ -21,8 +19,6 @@ fn projects_column(select: &nodes::SelectStmtMut<'_, '_>, name: &str) -> bool { }) } -/// Project ORDER BY expressions that are missing from the SELECT list -/// so cross-shard results can be sorted, then stripped. pub(super) fn rewrite_select<'a>( select: &mut nodes::SelectStmtMut<'a, '_>, mem: make::MemoryToken<'a>, diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs b/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs index 60d3b811b..2fecf16eb 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs @@ -9,7 +9,6 @@ use crate::net::ProtocolMessage; use pg_raw_parse::{Node, StmtList, make}; use std::sync::Arc; -/// Aggregate function projected temporarily for cross-shard merging. #[derive(Debug, Clone, PartialEq)] pub(crate) struct AggregateHelper { pub(crate) target_column: usize, @@ -19,14 +18,12 @@ pub(crate) struct AggregateHelper { pub(crate) alias: String, } -/// Column projected temporarily for cross-shard ordering. #[derive(Debug, Clone, PartialEq)] pub(crate) struct OrderByHelper { pub(crate) sort_position: usize, pub(crate) projected_column: usize, } -/// Temporary result columns required while merging cross-shard results. #[derive(Debug, Clone, Default, PartialEq)] pub(crate) struct ProjectionRewritePlan { aggregate_helpers: Vec, @@ -38,10 +35,6 @@ impl ProjectionRewritePlan { self.aggregate_helpers.is_empty() && self.order_by_helpers.is_empty() } - /// Positions of the columns we added to the select list. - /// - /// Helpers are only ever added to an explicit select list, so the position - /// in the statement is the position in the result set. pub(crate) fn drop_columns(&self) -> impl Iterator + '_ { self.aggregate_helpers .iter() @@ -87,11 +80,6 @@ impl RewriteOutput { } } -/// Add temporary columns needed to merge a cross-shard SELECT. -/// -/// This deliberately operates on a copy of the cached AST. The cached AST is -/// the route-independent representation and must remain suitable for direct -/// execution. pub(crate) fn finalize_after_route( request: &mut ClientRequest, schema: &Schema, From 0b8032025c1e685b57e0a950dfcc03755f982a33 Mon Sep 17 00:00:00 2001 From: Nupur Agrawal Date: Tue, 15 Sep 2026 15:52:23 +0530 Subject: [PATCH 08/14] edge cases in order_by --- pgdog/src/backend/pool/connection/buffer.rs | 7 +- .../pool/connection/multi_shard/mod.rs | 2 + .../pool/connection/multi_shard/test.rs | 62 ++++++++++++++- .../parser/rewrite/statement/order_by.rs | 79 ++++++++++++++----- 4 files changed, 124 insertions(+), 26 deletions(-) diff --git a/pgdog/src/backend/pool/connection/buffer.rs b/pgdog/src/backend/pool/connection/buffer.rs index df7a63ca1..40e697d80 100644 --- a/pgdog/src/backend/pool/connection/buffer.rs +++ b/pgdog/src/backend/pool/connection/buffer.rs @@ -143,7 +143,7 @@ impl Buffer { plan: &ProjectionRewritePlan, ) -> Result<(), super::Error> { let buffer: VecDeque = std::mem::take(&mut self.buffer); - let mut rows = if aggregate.is_empty() { + let rows = if aggregate.is_empty() { buffer } else if let Some(aggregates) = Aggregates::new(&buffer, decoder, aggregate, plan) { aggregates.aggregate()? @@ -151,20 +151,19 @@ impl Buffer { buffer }; - Self::drop_helper_columns(&mut rows, plan); self.buffer = rows; Ok(()) } - fn drop_helper_columns(rows: &mut VecDeque, plan: &ProjectionRewritePlan) { + pub(super) fn drop_helper_columns(&mut self, plan: &ProjectionRewritePlan) { if plan.is_noop() { return; } let drop = plan.drop_columns().collect(); - for row in rows.iter_mut() { + for row in self.buffer.iter_mut() { row.drop_columns(&drop); } } diff --git a/pgdog/src/backend/pool/connection/multi_shard/mod.rs b/pgdog/src/backend/pool/connection/multi_shard/mod.rs index 13372c5a4..c9eea3cf0 100644 --- a/pgdog/src/backend/pool/connection/multi_shard/mod.rs +++ b/pgdog/src/backend/pool/connection/multi_shard/mod.rs @@ -269,6 +269,8 @@ impl MultiShard { .map_err(Error::from)?; self.buffer.sort(self.route.order_by(), &self.decoder); + self.buffer + .drop_helper_columns(self.route.projection_rewrite_plan()); self.buffer.distinct(self.route.distinct(), &self.decoder); self.buffer.limit(self.route.limit()); } diff --git a/pgdog/src/backend/pool/connection/multi_shard/test.rs b/pgdog/src/backend/pool/connection/multi_shard/test.rs index 8cd70f231..2971d282e 100644 --- a/pgdog/src/backend/pool/connection/multi_shard/test.rs +++ b/pgdog/src/backend/pool/connection/multi_shard/test.rs @@ -1,5 +1,8 @@ use crate::{ - frontend::router::parser::{DistinctBy, Shard, ShardWithPriority}, + frontend::router::parser::{ + DistinctBy, OrderBy, Shard, ShardWithPriority, + rewrite::statement::projection::{OrderByHelper, ProjectionRewritePlan}, + }, net::{BindComplete, DataRow, Field, Format}, }; @@ -60,6 +63,63 @@ fn test_inconsistent_data_rows() { } } +#[test] +fn test_order_by_helper_is_dropped_after_sorting() { + let mut plan = ProjectionRewritePlan::default(); + plan.add_order_by_helper(OrderByHelper { + sort_position: 0, + projected_column: 1, + }); + let mut route = Route::select( + ShardWithPriority::new_default_unset(Shard::All), + vec![OrderBy::Asc(2)], + Default::default(), + Default::default(), + None, + ); + route.set_projection_rewrite_plan(plan); + let mut multi_shard = MultiShard::new(vec![0, 1], &route); + + let row_description = + RowDescription::new(&[Field::bigint("id"), Field::bigint("__pgdog_order_col0")]); + assert!( + multi_shard + .handle_server_message(row_description.message()) + .unwrap() + .is_none() + ); + let client_description = multi_shard + .handle_server_message(row_description.message()) + .unwrap() + .unwrap(); + assert_eq!( + RowDescription::from_bytes(client_description.to_bytes()) + .unwrap() + .len(), + 1 + ); + + let mut first = DataRow::new(); + first.add(1_i64).add(20_i64); + let mut second = DataRow::new(); + second.add(2_i64).add(10_i64); + multi_shard.handle_server_message(first.message()).unwrap(); + multi_shard.handle_server_message(second.message()).unwrap(); + + for _ in 0..2 { + multi_shard + .handle_server_message(CommandComplete::from_str("SELECT 1").message()) + .unwrap(); + } + + for expected in [2_i64, 1_i64] { + let message = multi_shard.get_server_message().unwrap(); + let row = DataRow::from_bytes(message.to_bytes()).unwrap(); + assert_eq!(row.len(), 1); + assert_eq!(row.get::(0, Format::Text).unwrap(), expected); + } +} + #[test] fn test_rd_before_dr() { let mut multi_shard = MultiShard::new( diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/order_by.rs b/pgdog/src/frontend/router/parser/rewrite/statement/order_by.rs index e1f2cc4dd..504f82ad8 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/order_by.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/order_by.rs @@ -4,18 +4,45 @@ use crate::frontend::router::parser::OrderBy; use super::projection::{OrderByHelper, ProjectionRewritePlan}; -fn projects_star(select: &nodes::SelectStmtMut<'_, '_>) -> bool { - select.target_list().iter().any(|target| { - matches!(target.val(), Node::ColumnRef(column) - if column.fields().into_iter().any(|field| matches!(field, Node::A_Star(_)))) - }) +fn same_column(left: &nodes::ColumnRef, right: &nodes::ColumnRef) -> bool { + left.fields() + .into_iter() + .map(Node::as_str) + .eq(right.fields().into_iter().map(Node::as_str)) +} + +fn star_covers(star: &nodes::ColumnRef, column: &nodes::ColumnRef) -> bool { + let star_fields = star.fields(); + let column_fields = column.fields(); + let star_len = star_fields.len(); + + matches!(star_fields.into_iter().next_back(), Some(Node::A_Star(_))) + && (star_len == 1 + || (star_len == column_fields.len() + && star + .fields() + .into_iter() + .take(star_len - 1) + .map(Node::as_str) + .eq(column_fields + .into_iter() + .take(star_len - 1) + .map(Node::as_str)))) } -fn projects_column(select: &nodes::SelectStmtMut<'_, '_>, name: &str) -> bool { +fn projects_column(select: &nodes::SelectStmtMut<'_, '_>, column: &nodes::ColumnRef) -> bool { + let fields = column.fields(); + let unqualified = fields.len() == 1; + let name = fields.into_iter().next_back().and_then(Node::as_str); + select.target_list().iter().any(|target| { - target.name() == Some(name) + (unqualified && target.name() == name) || matches!(target.val(), Node::ColumnRef(projected) - if projected.fields().into_iter().next_back().and_then(Node::as_str) == Some(name)) + if star_covers(projected, column) + || same_column(projected, column) + || (unqualified + && projected.fields().into_iter().next_back().and_then(Node::as_str) + == name)) }) } @@ -25,10 +52,6 @@ pub(super) fn rewrite_select<'a>( order_by: &[OrderBy], plan: &mut ProjectionRewritePlan, ) { - if projects_star(select) { - return; - } - let mut helpers = Vec::new(); let mut sort_position = 0; for sort in select.sort_clause() { @@ -53,15 +76,7 @@ pub(super) fn rewrite_select<'a>( sort_position += 1; let needs_helper = match node { - Node::ColumnRef(column) => match column - .fields() - .into_iter() - .next_back() - .and_then(Node::as_str) - { - Some(name) => !projects_column(select, name), - None => false, - }, + Node::ColumnRef(column) => !projects_column(select, column), Node::A_Expr(_) => matches!(order, OrderBy::AscVectorL2Column(_, _)), _ => false, }; @@ -158,4 +173,26 @@ mod tests { assert!(!sql.contains("__pgdog_order_col")); assert!(plan.is_noop()); } + + #[test] + fn projects_column_not_covered_by_qualified_star() { + let (sql, plan) = rewrite( + "SELECT a.* FROM a JOIN b ON a.id = b.a_id ORDER BY b.score", + vec![OrderBy::AscColumn("score".into())], + ); + + assert!(sql.contains("b.score AS __pgdog_order_col0")); + assert_eq!(plan.order_by_helpers().len(), 1); + } + + #[test] + fn distinguishes_same_named_columns_from_different_relations() { + let (sql, plan) = rewrite( + "SELECT a.price FROM a JOIN b ON a.id = b.a_id ORDER BY b.price", + vec![OrderBy::AscColumn("price".into())], + ); + + assert!(sql.contains("b.price AS __pgdog_order_col0")); + assert_eq!(plan.order_by_helpers().len(), 1); + } } From 8a87626c02bac55cdf91506dfea2fc871496b17d Mon Sep 17 00:00:00 2001 From: Nupur Agrawal Date: Tue, 15 Sep 2026 17:12:35 +0530 Subject: [PATCH 09/14] fix post route rewrites --- .../pool/connection/multi_shard/test.rs | 3 +- .../query_engine/test/rewrite_offset.rs | 75 +++++++++++++ .../query_engine/test/rewrite_projection.rs | 106 +++++++++++++++++- .../prepared_statements/global_cache.rs | 15 ++- pgdog/src/frontend/prepared_statements/mod.rs | 9 ++ .../parser/rewrite/statement/order_by.rs | 87 +++++++------- .../parser/rewrite/statement/projection.rs | 47 ++++++-- 7 files changed, 281 insertions(+), 61 deletions(-) diff --git a/pgdog/src/backend/pool/connection/multi_shard/test.rs b/pgdog/src/backend/pool/connection/multi_shard/test.rs index 2971d282e..17d62c7b3 100644 --- a/pgdog/src/backend/pool/connection/multi_shard/test.rs +++ b/pgdog/src/backend/pool/connection/multi_shard/test.rs @@ -1,7 +1,7 @@ use crate::{ frontend::router::parser::{ DistinctBy, OrderBy, Shard, ShardWithPriority, - rewrite::statement::projection::{OrderByHelper, ProjectionRewritePlan}, + rewrite::statement::projection::{OrderByHelper, OrderBySource, ProjectionRewritePlan}, }, net::{BindComplete, DataRow, Field, Format}, }; @@ -68,6 +68,7 @@ fn test_order_by_helper_is_dropped_after_sorting() { let mut plan = ProjectionRewritePlan::default(); plan.add_order_by_helper(OrderByHelper { sort_position: 0, + source: OrderBySource::Column("price".into()), projected_column: 1, }); let mut route = Route::select( diff --git a/pgdog/src/frontend/client/query_engine/test/rewrite_offset.rs b/pgdog/src/frontend/client/query_engine/test/rewrite_offset.rs index 6eea15746..65617cde5 100644 --- a/pgdog/src/frontend/client/query_engine/test/rewrite_offset.rs +++ b/pgdog/src/frontend/client/query_engine/test/rewrite_offset.rs @@ -246,3 +246,78 @@ async fn test_offset_with_unique_id_extended() { panic!("expected Bind"); } } + +#[tokio::test] +async fn split_anonymous_pagination_keeps_original_plan() { + let mut client = test_sharded_client(); + client.client_request = ClientRequest::default(); + client + .client_request + .push(ProtocolMessage::Parse(Parse::new_anonymous( + "SELECT * FROM test LIMIT 10 OFFSET 5", + ))); + client + .client_request + .push(ProtocolMessage::Describe(Describe::new_statement(""))); + client.client_request.push(Flush.into()); + + { + let mut engine = QueryEngine::from_client(&client).unwrap(); + let mut context = QueryEngineContext::new(&mut client); + let result = engine.parse_and_rewrite(&mut context).await.unwrap(); + context.client_request.route = Some(cross_shard_route()); + projection::finalize_after_route( + context.client_request, + &Schema::default(), + result.as_ref().and_then(RewriteResult::offset_plan), + ) + .unwrap(); + result + .as_ref() + .unwrap() + .apply_after_route(context.client_request) + .unwrap(); + } + + assert_eq!( + client.client_request.last_parse.as_ref().unwrap().query(), + "SELECT * FROM test LIMIT 10 OFFSET 5" + ); + + client.client_request.clear(); + client + .client_request + .push(ProtocolMessage::Bind(Bind::new_params("", &[]))); + client + .client_request + .push(ProtocolMessage::Execute(Execute::new())); + client.client_request.push(ProtocolMessage::Sync(Sync)); + + let mut engine = QueryEngine::from_client(&client).unwrap(); + let mut context = QueryEngineContext::new(&mut client); + let result = engine.parse_and_rewrite(&mut context).await.unwrap(); + context.client_request.route = Some(cross_shard_route()); + projection::finalize_after_route( + context.client_request, + &Schema::default(), + result.as_ref().and_then(RewriteResult::offset_plan), + ) + .unwrap(); + result + .as_ref() + .unwrap() + .apply_after_route(context.client_request) + .unwrap(); + + assert_eq!( + context.client_request.last_parse.as_ref().unwrap().query(), + "SELECT * FROM test LIMIT 10::bigint + 5::bigint" + ); + assert_eq!( + context.client_request.route().limit(), + &Limit { + limit: Some(10), + offset: Some(5), + } + ); +} diff --git a/pgdog/src/frontend/client/query_engine/test/rewrite_projection.rs b/pgdog/src/frontend/client/query_engine/test/rewrite_projection.rs index 4596f53ec..ed9e503d2 100644 --- a/pgdog/src/frontend/client/query_engine/test/rewrite_projection.rs +++ b/pgdog/src/frontend/client/query_engine/test/rewrite_projection.rs @@ -1,4 +1,5 @@ use crate::backend::schema::Schema; +use crate::frontend::router::parser::cache::ast::Ast; use crate::frontend::router::parser::rewrite::statement::plan::RewriteResult; use crate::frontend::router::parser::rewrite::statement::projection; use crate::frontend::router::parser::route::{Route, Shard, ShardWithPriority}; @@ -6,6 +7,7 @@ use crate::frontend::{ PreparedStatements, router::parser::{Limit, OrderBy}, }; +use pgdog_vector::Vector; use super::prelude::*; use super::test_sharded_client; @@ -247,6 +249,73 @@ async fn cross_shard_order_by_projects_missing_sort_column() { ); } +#[test] +fn cached_projection_does_not_depend_on_first_route_order() { + let sql = "SELECT id FROM products ORDER BY embedding <-> $1, price"; + let ast = Ast::new_record(sql).unwrap(); + let mut first = ClientRequest::from(vec![ProtocolMessage::Query(Query::new(sql))]); + first.ast = Some(ast.clone()); + first.route = Some(Route::select( + ShardWithPriority::new_table(Shard::All), + vec![OrderBy::AscColumn("price".into())], + Default::default(), + Limit::default(), + None, + )); + + projection::finalize_after_route(&mut first, &Schema::default(), None).unwrap(); + let first_query = match &first.messages[0] { + ProtocolMessage::Query(query) => query, + _ => panic!("expected Query"), + }; + assert!(first_query.query().contains("__pgdog_order_col0")); + assert!(first_query.query().contains("__pgdog_order_col1")); + assert_eq!(first.route().order_by(), &[OrderBy::Asc(3)]); + + let mut second = ClientRequest::from(vec![ProtocolMessage::Query(Query::new(sql))]); + second.ast = Some(ast); + second.route = Some(Route::select( + ShardWithPriority::new_table(Shard::All), + vec![ + OrderBy::AscVectorL2Column("embedding".into(), Vector::from(&[1.0, 2.0, 3.0][..])), + OrderBy::AscColumn("price".into()), + ], + Default::default(), + Limit::default(), + None, + )); + + projection::finalize_after_route(&mut second, &Schema::default(), None).unwrap(); + assert_eq!( + second.route().order_by(), + &[OrderBy::Asc(2), OrderBy::Asc(3)] + ); +} + +#[test] +fn helper_replaces_the_matching_duplicate_order_by_position() { + let sql = "SELECT a.price FROM a JOIN b ON a.id = b.a_id ORDER BY a.price, b.price"; + let mut request = ClientRequest::from(vec![ProtocolMessage::Query(Query::new(sql))]); + request.ast = Some(Ast::new_record(sql).unwrap()); + request.route = Some(Route::select( + ShardWithPriority::new_table(Shard::All), + vec![ + OrderBy::AscColumn("price".into()), + OrderBy::AscColumn("price".into()), + ], + Default::default(), + Limit::default(), + None, + )); + + projection::finalize_after_route(&mut request, &Schema::default(), None).unwrap(); + + assert_eq!( + request.route().order_by(), + &[OrderBy::AscColumn("price".into()), OrderBy::Asc(2)] + ); +} + #[tokio::test] async fn aggregate_order_by_and_offset_compose_after_route() { let mut client = test_sharded_client(); @@ -305,7 +374,7 @@ async fn aggregate_order_by_and_offset_compose_after_route() { } #[tokio::test] -async fn split_anonymous_prepare_finalizes_saved_parse_on_execute() { +async fn split_anonymous_prepare_rewrites_each_execution_once() { let mut client = test_sharded_client(); client.client_request = ClientRequest::default(); client @@ -317,6 +386,35 @@ async fn split_anonymous_prepare_finalizes_saved_parse_on_execute() { .client_request .push(ProtocolMessage::Describe(Describe::new_statement(""))); client.client_request.push(Flush.into()); + + { + let mut engine = QueryEngine::from_client(&client).unwrap(); + let mut context = QueryEngineContext::new(&mut client); + let result = engine.parse_and_rewrite(&mut context).await.unwrap(); + context.client_request.route = Some(route(Shard::All)); + projection::finalize_after_route( + context.client_request, + &Schema::default(), + result.as_ref().and_then(RewriteResult::offset_plan), + ) + .unwrap(); + } + + let parse = match &client.client_request.messages[0] { + ProtocolMessage::Parse(parse) => parse, + _ => panic!("expected Parse"), + }; + assert_eq!(parse.query().matches("__pgdog_count_col0").count(), 1); + assert!( + !client + .client_request + .last_parse + .as_ref() + .unwrap() + .query() + .contains("__pgdog_count_col0") + ); + client.client_request.clear(); client .client_request @@ -338,14 +436,16 @@ async fn split_anonymous_prepare_finalizes_saved_parse_on_execute() { ) .unwrap(); - assert!( + assert_eq!( context .client_request .last_parse .as_ref() .unwrap() .query() - .contains("__pgdog_count_col0") + .matches("__pgdog_count_col0") + .count(), + 1 ); assert!( context diff --git a/pgdog/src/frontend/prepared_statements/global_cache.rs b/pgdog/src/frontend/prepared_statements/global_cache.rs index e7015f89a..21c87c68a 100644 --- a/pgdog/src/frontend/prepared_statements/global_cache.rs +++ b/pgdog/src/frontend/prepared_statements/global_cache.rs @@ -47,6 +47,13 @@ impl MemoryUsage for GlobalCache { } impl GlobalCache { + pub(crate) fn cross_shard_variant_name(&self, name: &str) -> Option { + let variant_name = format!("{name}_cross_shard"); + self.cross_shard_variants + .contains_key(&variant_name) + .then_some(variant_name) + } + /// Record a Parse message with the global cache and return a globally unique /// name PgDog is using for that statement. /// @@ -140,10 +147,10 @@ impl GlobalCache { } pub(crate) fn cross_shard_variant(&mut self, name: &str, query: &str) -> Option { - let variant_name = format!("{name}_cross_shard"); - if self.cross_shard_variants.contains_key(&variant_name) { + if let Some(variant_name) = self.cross_shard_variant_name(name) { return Some(variant_name); } + let variant_name = format!("{name}_cross_shard"); let mut parse = self.rewritten_parse(name)?; parse.rename(&variant_name); @@ -405,6 +412,10 @@ mod test { ) .unwrap(); assert_eq!(variant, format!("{base}_cross_shard")); + assert_eq!( + cache.cross_shard_variant_name(&base).as_deref(), + Some(variant.as_str()) + ); assert_eq!( cache.rewritten_parse(&base).unwrap().query(), "SELECT AVG(value) FROM measurements" diff --git a/pgdog/src/frontend/prepared_statements/mod.rs b/pgdog/src/frontend/prepared_statements/mod.rs index 3e8b40654..861df8f9a 100644 --- a/pgdog/src/frontend/prepared_statements/mod.rs +++ b/pgdog/src/frontend/prepared_statements/mod.rs @@ -67,6 +67,15 @@ impl PreparedStatements { Self::new().global.clone() } + pub(crate) fn cross_shard_variant(name: &str, query: &str) -> Option { + let cache = Self::global(); + if let Some(variant) = cache.read().cross_shard_variant_name(name) { + return Some(variant); + } + + cache.write().cross_shard_variant(name, query) + } + /// Rewrite extended protocol messages to use global names. This allows multiple /// clients to re-use the same statement prepared on a Postgres server. /// diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/order_by.rs b/pgdog/src/frontend/router/parser/rewrite/statement/order_by.rs index 504f82ad8..a7007970d 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/order_by.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/order_by.rs @@ -1,8 +1,8 @@ -use pg_raw_parse::{Node, make, nodes}; +use pg_raw_parse::{ConstValue, Node, make, nodes}; -use crate::frontend::router::parser::OrderBy; +use crate::frontend::router::parser::Column; -use super::projection::{OrderByHelper, ProjectionRewritePlan}; +use super::projection::{OrderByHelper, OrderBySource, ProjectionRewritePlan}; fn same_column(left: &nodes::ColumnRef, right: &nodes::ColumnRef) -> bool { left.fields() @@ -49,25 +49,32 @@ fn projects_column(select: &nodes::SelectStmtMut<'_, '_>, column: &nodes::Column pub(super) fn rewrite_select<'a>( select: &mut nodes::SelectStmtMut<'a, '_>, mem: make::MemoryToken<'a>, - order_by: &[OrderBy], plan: &mut ProjectionRewritePlan, ) { let mut helpers = Vec::new(); let mut sort_position = 0; for sort in select.sort_clause() { let node = sort.node(); - let Some(order) = order_by.get(sort_position) else { - break; + let source = match node { + Node::ColumnRef(column) => column + .fields() + .into_iter() + .next_back() + .and_then(Node::as_str) + .map(|name| OrderBySource::Column(name.to_owned())), + Node::A_Expr(expr) + if expr.name().iter().next().and_then(Node::as_str) == Some("<->") => + { + [expr.lexpr(), expr.rexpr()] + .into_iter() + .find_map(|node| Column::try_from(node).ok()) + .map(|column| OrderBySource::Vector(column.name.to_owned())) + } + _ => None, }; - let supported = matches!( - (node, order), - (Node::A_Const(_), OrderBy::Asc(_) | OrderBy::Desc(_)) - | ( - Node::ColumnRef(_), - OrderBy::AscColumn(_) | OrderBy::DescColumn(_) - ) - | (Node::A_Expr(_), OrderBy::AscVectorL2Column(_, _)) - ); + let supported = source.is_some() + || matches!(node, Node::A_Const(constant) + if matches!(constant.val(), Some(ConstValue::Integer(_)))); if !supported { continue; } @@ -77,7 +84,7 @@ pub(super) fn rewrite_select<'a>( let needs_helper = match node { Node::ColumnRef(column) => !projects_column(select, column), - Node::A_Expr(_) => matches!(order, OrderBy::AscVectorL2Column(_, _)), + Node::A_Expr(_) => source.is_some(), _ => false, }; if !needs_helper { @@ -93,6 +100,7 @@ pub(super) fn rewrite_select<'a>( )); plan.add_order_by_helper(OrderByHelper { sort_position: current_sort_position, + source: source.expect("only columns and vector expressions need helpers"), projected_column, }); } @@ -109,7 +117,7 @@ mod tests { use super::*; use pg_raw_parse::{Node, make}; - fn rewrite(sql: &str, order_by: Vec) -> (String, ProjectionRewritePlan) { + fn rewrite(sql: &str) -> (String, ProjectionRewritePlan) { let ast = pg_raw_parse::parse(sql).unwrap(); let mut plan = ProjectionRewritePlan::default(); let rewritten = make::owned(|mem| { @@ -117,7 +125,7 @@ mod tests { panic!("expected SELECT"); }; let mut select = mem.make_unique(select); - rewrite_select(&mut select.as_mut(), mem, &order_by, &mut plan); + rewrite_select(&mut select.as_mut(), mem, &mut plan); select }); ( @@ -131,10 +139,7 @@ mod tests { #[test] fn projects_missing_sort_column() { - let (sql, plan) = rewrite( - "SELECT id FROM products ORDER BY price", - vec![OrderBy::AscColumn("price".into())], - ); + let (sql, plan) = rewrite("SELECT id FROM products ORDER BY price"); assert!(sql.contains("price AS __pgdog_order_col0")); assert_eq!(plan.order_by_helpers().len(), 1); @@ -143,10 +148,7 @@ mod tests { #[test] fn skips_already_projected_sort_column() { - let (sql, plan) = rewrite( - "SELECT id, price FROM products ORDER BY price", - vec![OrderBy::AscColumn("price".into())], - ); + let (sql, plan) = rewrite("SELECT id, price FROM products ORDER BY price"); assert!(!sql.contains("__pgdog_order_col")); assert!(plan.is_noop()); @@ -154,10 +156,7 @@ mod tests { #[test] fn skips_star_select() { - let (sql, plan) = rewrite( - "SELECT * FROM products ORDER BY id", - vec![OrderBy::AscColumn("id".into())], - ); + let (sql, plan) = rewrite("SELECT * FROM products ORDER BY id"); assert!(!sql.contains("__pgdog_order_col")); assert!(plan.is_noop()); @@ -165,10 +164,7 @@ mod tests { #[test] fn skips_qualified_star_select() { - let (sql, plan) = rewrite( - "SELECT products.* FROM products ORDER BY products.id", - vec![OrderBy::AscColumn("id".into())], - ); + let (sql, plan) = rewrite("SELECT products.* FROM products ORDER BY products.id"); assert!(!sql.contains("__pgdog_order_col")); assert!(plan.is_noop()); @@ -176,10 +172,7 @@ mod tests { #[test] fn projects_column_not_covered_by_qualified_star() { - let (sql, plan) = rewrite( - "SELECT a.* FROM a JOIN b ON a.id = b.a_id ORDER BY b.score", - vec![OrderBy::AscColumn("score".into())], - ); + let (sql, plan) = rewrite("SELECT a.* FROM a JOIN b ON a.id = b.a_id ORDER BY b.score"); assert!(sql.contains("b.score AS __pgdog_order_col0")); assert_eq!(plan.order_by_helpers().len(), 1); @@ -187,12 +180,20 @@ mod tests { #[test] fn distinguishes_same_named_columns_from_different_relations() { - let (sql, plan) = rewrite( - "SELECT a.price FROM a JOIN b ON a.id = b.a_id ORDER BY b.price", - vec![OrderBy::AscColumn("price".into())], - ); + let (sql, plan) = + rewrite("SELECT a.price FROM a JOIN b ON a.id = b.a_id ORDER BY a.price, b.price"); + + assert!(sql.contains("b.price AS __pgdog_order_col1")); + assert_eq!(plan.order_by_helpers().len(), 1); + assert_eq!(plan.order_by_helpers()[0].sort_position, 1); + } + + #[test] + fn projects_vector_distance_without_resolved_parameter() { + let (sql, plan) = rewrite("SELECT id FROM products ORDER BY embedding <-> $1 LIMIT 5"); - assert!(sql.contains("b.price AS __pgdog_order_col0")); + assert!(sql.contains("embedding <-> $1")); + assert!(sql.contains("AS __pgdog_order_col0")); assert_eq!(plan.order_by_helpers().len(), 1); } } diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs b/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs index 2fecf16eb..7e37fe2fe 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs @@ -21,9 +21,29 @@ pub(crate) struct AggregateHelper { #[derive(Debug, Clone, PartialEq)] pub(crate) struct OrderByHelper { pub(crate) sort_position: usize, + pub(crate) source: OrderBySource, pub(crate) projected_column: usize, } +#[derive(Debug, Clone, PartialEq)] +pub(crate) enum OrderBySource { + Column(String), + Vector(String), +} + +impl OrderByHelper { + fn matches(&self, order_by: &OrderBy) -> bool { + match (&self.source, order_by) { + (OrderBySource::Column(source), OrderBy::AscColumn(column)) + | (OrderBySource::Column(source), OrderBy::DescColumn(column)) + | (OrderBySource::Vector(source), OrderBy::AscVectorL2Column(column, _)) => { + source == column + } + _ => false, + } + } +} + #[derive(Debug, Clone, Default, PartialEq)] pub(crate) struct ProjectionRewritePlan { aggregate_helpers: Vec, @@ -93,10 +113,9 @@ pub(crate) fn finalize_after_route( return Ok(()); }; let rewrite_offset = offset_plan.is_some_and(|plan| !plan.prepare_execute); - let order_by = request.route().order_by(); let Some(rewrite) = ast .post_route_rewrite - .get_or_try_init(|| build(&ast.ast, schema, order_by, rewrite_offset))? + .get_or_try_init(|| build(&ast.ast, schema, rewrite_offset))? else { return Ok(()); }; @@ -108,11 +127,8 @@ pub(crate) fn finalize_after_route( } _ => None, }); - let variant = base_name.and_then(|name| { - PreparedStatements::global() - .write() - .cross_shard_variant(name, &rewrite.sql) - }); + let variant = + base_name.and_then(|name| PreparedStatements::cross_shard_variant(name, &rewrite.sql)); for message in &mut request.messages { match message { @@ -136,7 +152,9 @@ pub(crate) fn finalize_after_route( _ => {} } } - if let Some(parse) = request.last_parse.as_mut() { + if request.is_executable() + && let Some(parse) = request.last_parse.as_mut() + { parse.set_query(&rewrite.sql); } if !rewrite.plan.is_noop() @@ -145,7 +163,12 @@ pub(crate) fn finalize_after_route( route.set_projection_rewrite_plan(rewrite.plan.clone()); let mut order_by = route.order_by().to_vec(); for helper in rewrite.plan.order_by_helpers() { - let Some(sort) = order_by.get_mut(helper.sort_position) else { + let position = order_by + .get(helper.sort_position) + .filter(|sort| helper.matches(sort)) + .map(|_| helper.sort_position) + .or_else(|| order_by.iter().position(|sort| helper.matches(sort))); + let Some(sort) = position.and_then(|position| order_by.get_mut(position)) else { continue; }; *sort = if sort.asc() { @@ -163,7 +186,6 @@ pub(crate) fn finalize_after_route( fn build( ast: &StmtList, schema: &Schema, - order_by: &[OrderBy], rewrite_offset: bool, ) -> Result, Error> { let Some(Node::SelectStmt(select)) = ast.stmts().next() else { @@ -171,7 +193,7 @@ fn build( }; let aggregate = Aggregate::parse(select, schema); - if aggregate.is_empty() && order_by.is_empty() && !rewrite_offset { + if aggregate.is_empty() && select.sort_clause().is_empty() && !rewrite_offset { return Ok(None); } @@ -181,7 +203,7 @@ fn build( if !aggregate.is_empty() { plan = AggregatesRewrite::rewrite_select(&mut select.as_mut(), mem, &aggregate).plan; } - order_by::rewrite_select(&mut select.as_mut(), mem, order_by, &mut plan); + order_by::rewrite_select(&mut select.as_mut(), mem, &mut plan); if rewrite_offset { offset::rewrite_select(&mut select.as_mut(), mem); } @@ -211,6 +233,7 @@ mod tests { }); plan.add_order_by_helper(OrderByHelper { sort_position: 0, + source: OrderBySource::Column("created_at".into()), projected_column: 2, }); From 017fc1f698f30059bddf099d5dd4f3a20ed76453 Mon Sep 17 00:00:00 2001 From: Nupur Agrawal Date: Thu, 17 Sep 2026 21:41:54 +0530 Subject: [PATCH 10/14] add helper --- .../frontend/prepared_statements/global_cache.rs | 16 ++++++++++------ pgdog/src/frontend/prepared_statements/mod.rs | 2 +- 2 files changed, 11 insertions(+), 7 deletions(-) diff --git a/pgdog/src/frontend/prepared_statements/global_cache.rs b/pgdog/src/frontend/prepared_statements/global_cache.rs index 4b2699f87..52fc65a0c 100644 --- a/pgdog/src/frontend/prepared_statements/global_cache.rs +++ b/pgdog/src/frontend/prepared_statements/global_cache.rs @@ -13,6 +13,10 @@ use fnv::FnvHashSet as HashSet; use super::*; +fn cross_shard_variant_name(name: &str) -> String { + format!("{name}_cross_shard") +} + /// Global prepared statements cache. /// /// The cache contains two mappings: @@ -47,8 +51,8 @@ impl MemoryUsage for GlobalCache { } impl GlobalCache { - pub(crate) fn cross_shard_variant_name(&self, name: &str) -> Option { - let variant_name = format!("{name}_cross_shard"); + pub(crate) fn existing_cross_shard_variant_name(&self, name: &str) -> Option { + let variant_name = cross_shard_variant_name(name); self.cross_shard_variants .contains_key(&variant_name) .then_some(variant_name) @@ -151,11 +155,11 @@ impl GlobalCache { } pub(crate) fn cross_shard_variant(&mut self, name: &str, query: &str) -> Option { - if let Some(variant_name) = self.cross_shard_variant_name(name) { + if let Some(variant_name) = self.existing_cross_shard_variant_name(name) { return Some(variant_name); } let client_params = self.client_params(name); - let variant_name = format!("{name}_cross_shard"); + let variant_name = cross_shard_variant_name(name); let mut parse = self.rewritten_parse(name)?; parse.rename(&variant_name); @@ -328,7 +332,7 @@ impl GlobalCache { if let Some(stmt) = self.names.remove(name) { self.statements.remove(stmt.cache_key()); self.cross_shard_variants - .remove(&format!("{name}_cross_shard")); + .remove(&cross_shard_variant_name(name)); } } @@ -423,7 +427,7 @@ mod test { .unwrap(); assert_eq!(variant, format!("{base}_cross_shard")); assert_eq!( - cache.cross_shard_variant_name(&base).as_deref(), + cache.existing_cross_shard_variant_name(&base).as_deref(), Some(variant.as_str()) ); assert_eq!( diff --git a/pgdog/src/frontend/prepared_statements/mod.rs b/pgdog/src/frontend/prepared_statements/mod.rs index e752f05e4..804bc836e 100644 --- a/pgdog/src/frontend/prepared_statements/mod.rs +++ b/pgdog/src/frontend/prepared_statements/mod.rs @@ -72,7 +72,7 @@ impl PreparedStatements { pub(crate) fn cross_shard_variant(name: &str, query: &str) -> Option { let cache = Self::global(); - if let Some(variant) = cache.read().cross_shard_variant_name(name) { + if let Some(variant) = cache.read().existing_cross_shard_variant_name(name) { return Some(variant); } From f6550d476ecd96951a9f3c6d30893f7a0026be1d Mon Sep 17 00:00:00 2001 From: Nupur Agrawal Date: Thu, 17 Sep 2026 22:04:04 +0530 Subject: [PATCH 11/14] comments --- pgdog/src/backend/pool/connection/multi_shard/mod.rs | 2 ++ pgdog/src/backend/prepared_statements.rs | 2 ++ pgdog/src/frontend/prepared_statements/global_cache.rs | 2 ++ pgdog/src/frontend/prepared_statements/mod.rs | 2 ++ pgdog/src/frontend/router/parser/cache/ast.rs | 3 ++- .../frontend/router/parser/rewrite/statement/offset.rs | 5 +++++ .../router/parser/rewrite/statement/projection.rs | 10 ++++++++++ 7 files changed, 25 insertions(+), 1 deletion(-) diff --git a/pgdog/src/backend/pool/connection/multi_shard/mod.rs b/pgdog/src/backend/pool/connection/multi_shard/mod.rs index c9eea3cf0..d734da4d7 100644 --- a/pgdog/src/backend/pool/connection/multi_shard/mod.rs +++ b/pgdog/src/backend/pool/connection/multi_shard/mod.rs @@ -260,6 +260,8 @@ impl MultiShard { self.buffer.mark_full(); if !self.buffer.is_empty() { + // Helpers remain in the internal row through aggregation and + // sorting, then are removed before client-visible operations. self.buffer .aggregate( self.route.aggregate(), diff --git a/pgdog/src/backend/prepared_statements.rs b/pgdog/src/backend/prepared_statements.rs index dbed3808e..9684e0505 100644 --- a/pgdog/src/backend/prepared_statements.rs +++ b/pgdog/src/backend/prepared_statements.rs @@ -566,6 +566,8 @@ impl PreparedStatements { })) } + /// Named Bind execution does not return RowDescription, so describe a new + /// cross-shard variant internally before decoding its helper columns. fn internal_describe(&mut self, name: &str) -> Option { if self.describes.iter().any(|describe| describe == name) || !self diff --git a/pgdog/src/frontend/prepared_statements/global_cache.rs b/pgdog/src/frontend/prepared_statements/global_cache.rs index 52fc65a0c..b96a7f941 100644 --- a/pgdog/src/frontend/prepared_statements/global_cache.rs +++ b/pgdog/src/frontend/prepared_statements/global_cache.rs @@ -154,6 +154,8 @@ impl GlobalCache { } } + /// Keep helper-bearing SQL separate from the base statement so direct + /// executions and client-visible metadata retain the original shape. pub(crate) fn cross_shard_variant(&mut self, name: &str, query: &str) -> Option { if let Some(variant_name) = self.existing_cross_shard_variant_name(name) { return Some(variant_name); diff --git a/pgdog/src/frontend/prepared_statements/mod.rs b/pgdog/src/frontend/prepared_statements/mod.rs index 804bc836e..9ec0c3731 100644 --- a/pgdog/src/frontend/prepared_statements/mod.rs +++ b/pgdog/src/frontend/prepared_statements/mod.rs @@ -72,6 +72,8 @@ impl PreparedStatements { pub(crate) fn cross_shard_variant(name: &str, query: &str) -> Option { let cache = Self::global(); + // Existing variants are the hot path; only creation needs the global + // write lock. if let Some(variant) = cache.read().existing_cross_shard_variant_name(name) { return Some(variant); } diff --git a/pgdog/src/frontend/router/parser/cache/ast.rs b/pgdog/src/frontend/router/parser/cache/ast.rs index dbdd8d053..e680fa45d 100644 --- a/pgdog/src/frontend/router/parser/cache/ast.rs +++ b/pgdog/src/frontend/router/parser/cache/ast.rs @@ -42,7 +42,8 @@ pub(crate) struct AstInner { pub(crate) stats: Mutex, /// Rewrite plan. pub(crate) rewrite_plan: RewritePlan, - /// Lazily generated SQL and response metadata for cross-shard execution. + /// Lazily generated cross-shard SQL and response metadata. This is derived + /// only from the AST so Bind values cannot permanently change a cache entry. pub(crate) post_route_rewrite: OnceCell>, /// Original query. pub(crate) query_without_comment: Arc, diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/offset.rs b/pgdog/src/frontend/router/parser/rewrite/statement/offset.rs index 9416efd1d..8d35bea45 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/offset.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/offset.rs @@ -189,6 +189,11 @@ fn to_bigint<'a>(node: Node<'_>, mem: make::MemoryToken<'a>) -> make::Unique<'a, .uncast() } +/// Keep LIMIT and OFFSET as an expression instead of folding them into an +/// `A_Const`. One cached SQL form then works for literals and placeholders +/// without changing client Bind values or their text/binary encoding. Postgres +/// evaluates the per-shard fetch bound; the route keeps the original values for +/// final proxy-side pagination. pub(super) fn rewrite_select<'a>( select: &mut nodes::SelectStmtMut<'a, '_>, mem: make::MemoryToken<'a>, diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs b/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs index 7e37fe2fe..5c207f0f5 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs @@ -1,3 +1,9 @@ +//! Route-dependent SELECT rewrites. +//! +//! The cached AST stays in its base form so one statement can safely alternate +//! between direct and cross-shard execution. Cross-shard SQL is built from that +//! AST only after routing and cached independently from per-execution Bind values. + use super::Error; use super::aggregate::{AggregatesRewrite, HelperKind}; use super::offset::{self, OffsetPlan}; @@ -152,6 +158,8 @@ pub(crate) fn finalize_after_route( _ => {} } } + // Parse/Describe-only requests must keep the saved anonymous Parse in its + // base form; the later Bind/Execute request will finalize its injected copy. if request.is_executable() && let Some(parse) = request.last_parse.as_mut() { @@ -163,6 +171,8 @@ pub(crate) fn finalize_after_route( route.set_projection_rewrite_plan(rewrite.plan.clone()); let mut order_by = route.order_by().to_vec(); for helper in rewrite.plan.order_by_helpers() { + // Prefer the structural position. The source fallback handles a + // bind-dependent vector sort omitted from this execution's route. let position = order_by .get(helper.sort_position) .filter(|sort| helper.matches(sort)) From 054ca7231093a5408b5900cbec861a634e887a7d Mon Sep 17 00:00:00 2001 From: Nupur Agrawal Date: Tue, 22 Sep 2026 13:47:53 +0530 Subject: [PATCH 12/14] cleanup and fixes --- .../src/backend/pool/connection/aggregate.rs | 2 +- pgdog/src/backend/pool/connection/buffer.rs | 4 +- .../pool/connection/multi_shard/mod.rs | 8 +- .../pool/connection/multi_shard/test.rs | 43 +++++++---- .../query_engine/test/rewrite_projection.rs | 43 +++++++---- .../frontend/router/parser/query/explain.rs | 3 +- pgdog/src/frontend/router/parser/query/mod.rs | 4 +- .../frontend/router/parser/query/prepare.rs | 2 +- .../frontend/router/parser/query/select.rs | 3 - .../rewrite/statement/aggregate/engine.rs | 77 ++++++++----------- .../parser/rewrite/statement/aggregate/mod.rs | 5 ++ .../parser/rewrite/statement/order_by.rs | 32 +++++--- .../parser/rewrite/statement/projection.rs | 73 +++++++----------- pgdog/src/frontend/router/parser/route.rs | 10 +-- 14 files changed, 156 insertions(+), 153 deletions(-) diff --git a/pgdog/src/backend/pool/connection/aggregate.rs b/pgdog/src/backend/pool/connection/aggregate.rs index 171d8bbfa..e8bf3b7ad 100644 --- a/pgdog/src/backend/pool/connection/aggregate.rs +++ b/pgdog/src/backend/pool/connection/aggregate.rs @@ -262,7 +262,7 @@ impl<'a> Aggregates<'a> { } } - for helper in plan.aggregate_helpers() { + for helper in &plan.aggregate_helpers { let Some(index) = decoder.row_description().field_index(&helper.alias) else { continue; }; diff --git a/pgdog/src/backend/pool/connection/buffer.rs b/pgdog/src/backend/pool/connection/buffer.rs index 40e697d80..f523753f9 100644 --- a/pgdog/src/backend/pool/connection/buffer.rs +++ b/pgdog/src/backend/pool/connection/buffer.rs @@ -156,12 +156,12 @@ impl Buffer { Ok(()) } - pub(super) fn drop_helper_columns(&mut self, plan: &ProjectionRewritePlan) { + pub(super) fn drop_helper_columns(&mut self, plan: &ProjectionRewritePlan, decoder: &Decoder) { if plan.is_noop() { return; } - let drop = plan.drop_columns().collect(); + let drop = plan.drop_columns(decoder.row_description()); for row in self.buffer.iter_mut() { row.drop_columns(&drop); diff --git a/pgdog/src/backend/pool/connection/multi_shard/mod.rs b/pgdog/src/backend/pool/connection/multi_shard/mod.rs index d734da4d7..7dbd5d3bd 100644 --- a/pgdog/src/backend/pool/connection/multi_shard/mod.rs +++ b/pgdog/src/backend/pool/connection/multi_shard/mod.rs @@ -266,13 +266,13 @@ impl MultiShard { .aggregate( self.route.aggregate(), &self.decoder, - self.route.projection_rewrite_plan(), + &self.route.projection_rewrite_plan, ) .map_err(Error::from)?; self.buffer.sort(self.route.order_by(), &self.decoder); self.buffer - .drop_helper_columns(self.route.projection_rewrite_plan()); + .drop_helper_columns(&self.route.projection_rewrite_plan, &self.decoder); self.buffer.distinct(self.route.distinct(), &self.decoder); self.buffer.limit(self.route.limit()); } @@ -315,11 +315,11 @@ impl MultiShard { { // Only send it to the client once all shards sent it, // so we don't get early requests from clients. - let plan = self.route.projection_rewrite_plan(); + let plan = &self.route.projection_rewrite_plan; if plan.is_noop() { forward = Some(message); } else { - let client_rd = rd.drop_columns(plan.drop_columns()); + let client_rd = rd.drop_columns(plan.drop_columns(&rd)); forward = Some(client_rd.message()); } diff --git a/pgdog/src/backend/pool/connection/multi_shard/test.rs b/pgdog/src/backend/pool/connection/multi_shard/test.rs index 17d62c7b3..e2a73b901 100644 --- a/pgdog/src/backend/pool/connection/multi_shard/test.rs +++ b/pgdog/src/backend/pool/connection/multi_shard/test.rs @@ -64,25 +64,29 @@ fn test_inconsistent_data_rows() { } #[test] -fn test_order_by_helper_is_dropped_after_sorting() { +fn test_order_by_helper_after_star_expansion_is_dropped_after_sorting() { let mut plan = ProjectionRewritePlan::default(); - plan.add_order_by_helper(OrderByHelper { + plan.order_by_helpers.push(OrderByHelper { sort_position: 0, source: OrderBySource::Column("price".into()), - projected_column: 1, + alias: "__pgdog_order_col0".into(), }); let mut route = Route::select( ShardWithPriority::new_default_unset(Shard::All), - vec![OrderBy::Asc(2)], + vec![OrderBy::AscColumn("__pgdog_order_col0".into())], Default::default(), Default::default(), None, ); - route.set_projection_rewrite_plan(plan); + route.projection_rewrite_plan = plan; let mut multi_shard = MultiShard::new(vec![0, 1], &route); - let row_description = - RowDescription::new(&[Field::bigint("id"), Field::bigint("__pgdog_order_col0")]); + let row_description = RowDescription::new(&[ + Field::bigint("id"), + Field::text("value"), + Field::timestamp("created_at"), + Field::bigint("__pgdog_order_col0"), + ]); assert!( multi_shard .handle_server_message(row_description.message()) @@ -93,17 +97,28 @@ fn test_order_by_helper_is_dropped_after_sorting() { .handle_server_message(row_description.message()) .unwrap() .unwrap(); + let client_description = RowDescription::from_bytes(client_description.to_bytes()).unwrap(); assert_eq!( - RowDescription::from_bytes(client_description.to_bytes()) - .unwrap() - .len(), - 1 + client_description + .fields + .iter() + .map(|field| field.name.as_str()) + .collect::>(), + ["id", "value", "created_at"] ); let mut first = DataRow::new(); - first.add(1_i64).add(20_i64); + first + .add(1_i64) + .add("first") + .add("2026-01-01 00:00:00") + .add(20_i64); let mut second = DataRow::new(); - second.add(2_i64).add(10_i64); + second + .add(2_i64) + .add("second") + .add("2026-01-02 00:00:00") + .add(10_i64); multi_shard.handle_server_message(first.message()).unwrap(); multi_shard.handle_server_message(second.message()).unwrap(); @@ -116,7 +131,7 @@ fn test_order_by_helper_is_dropped_after_sorting() { for expected in [2_i64, 1_i64] { let message = multi_shard.get_server_message().unwrap(); let row = DataRow::from_bytes(message.to_bytes()).unwrap(); - assert_eq!(row.len(), 1); + assert_eq!(row.len(), 3); assert_eq!(row.get::(0, Format::Text).unwrap(), expected); } } diff --git a/pgdog/src/frontend/client/query_engine/test/rewrite_projection.rs b/pgdog/src/frontend/client/query_engine/test/rewrite_projection.rs index ed9e503d2..533efb683 100644 --- a/pgdog/src/frontend/client/query_engine/test/rewrite_projection.rs +++ b/pgdog/src/frontend/client/query_engine/test/rewrite_projection.rs @@ -55,7 +55,7 @@ async fn direct_aggregate_keeps_base_sql() { context .client_request .route() - .projection_rewrite_plan() + .projection_rewrite_plan .is_noop() ); } @@ -88,8 +88,8 @@ async fn cross_shard_aggregate_adds_and_tracks_helpers() { context .client_request .route() - .projection_rewrite_plan() - .aggregate_helpers() + .projection_rewrite_plan + .aggregate_helpers .len(), 1 ); @@ -236,14 +236,14 @@ async fn cross_shard_order_by_projects_missing_sort_column() { assert!(query.query().contains("price AS __pgdog_order_col0")); assert_eq!( context.client_request.route().order_by(), - &[OrderBy::Asc(2)] + &[OrderBy::AscColumn("__pgdog_order_col0".into())] ); assert_eq!( context .client_request .route() - .projection_rewrite_plan() - .order_by_helpers() + .projection_rewrite_plan + .order_by_helpers .len(), 1 ); @@ -270,7 +270,10 @@ fn cached_projection_does_not_depend_on_first_route_order() { }; assert!(first_query.query().contains("__pgdog_order_col0")); assert!(first_query.query().contains("__pgdog_order_col1")); - assert_eq!(first.route().order_by(), &[OrderBy::Asc(3)]); + assert_eq!( + first.route().order_by(), + &[OrderBy::AscColumn("__pgdog_order_col1".into())] + ); let mut second = ClientRequest::from(vec![ProtocolMessage::Query(Query::new(sql))]); second.ast = Some(ast); @@ -288,7 +291,10 @@ fn cached_projection_does_not_depend_on_first_route_order() { projection::finalize_after_route(&mut second, &Schema::default(), None).unwrap(); assert_eq!( second.route().order_by(), - &[OrderBy::Asc(2), OrderBy::Asc(3)] + &[ + OrderBy::AscColumn("__pgdog_order_col0".into()), + OrderBy::AscColumn("__pgdog_order_col1".into()) + ] ); } @@ -312,7 +318,10 @@ fn helper_replaces_the_matching_duplicate_order_by_position() { assert_eq!( request.route().order_by(), - &[OrderBy::AscColumn("price".into()), OrderBy::Asc(2)] + &[ + OrderBy::AscColumn("price".into()), + OrderBy::AscColumn("__pgdog_order_col1".into()) + ] ); } @@ -356,13 +365,17 @@ async fn aggregate_order_by_and_offset_compose_after_route() { assert!(!query.query().contains("OFFSET")); let route = context.client_request.route(); - assert_eq!(route.order_by(), &[OrderBy::Asc(3)]); assert_eq!( - route - .projection_rewrite_plan() - .drop_columns() - .collect::>(), - [1, 2] + route.order_by(), + &[OrderBy::AscColumn("__pgdog_order_col0".into())] + ); + assert_eq!( + route.projection_rewrite_plan.aggregate_helpers[0].alias, + "__pgdog_count_col0" + ); + assert_eq!( + route.projection_rewrite_plan.order_by_helpers[0].alias, + "__pgdog_order_col0" ); assert_eq!( route.limit(), diff --git a/pgdog/src/frontend/router/parser/query/explain.rs b/pgdog/src/frontend/router/parser/query/explain.rs index f011fe2d9..b3f687966 100644 --- a/pgdog/src/frontend/router/parser/query/explain.rs +++ b/pgdog/src/frontend/router/parser/query/explain.rs @@ -4,7 +4,6 @@ use pg_raw_parse::nodes; impl QueryParser { pub(super) fn explain( &mut self, - cached_ast: &Ast, stmt: &nodes::ExplainStmt, context: &mut QueryParserContext, ) -> Result { @@ -19,7 +18,7 @@ impl QueryParser { } let result = match query { - Node::SelectStmt(stmt) => self.select(cached_ast, stmt, context), + Node::SelectStmt(stmt) => self.select(stmt, context), Node::InsertStmt(stmt) => self.insert(stmt.into(), context), Node::UpdateStmt(stmt) => self.update(stmt.into(), context), Node::DeleteStmt(stmt) => self.delete(stmt.into(), context), diff --git a/pgdog/src/frontend/router/parser/query/mod.rs b/pgdog/src/frontend/router/parser/query/mod.rs index 6ffde6ede..c946b5342 100644 --- a/pgdog/src/frontend/router/parser/query/mod.rs +++ b/pgdog/src/frontend/router/parser/query/mod.rs @@ -386,7 +386,7 @@ impl QueryParser { ShardWithPriority::new_override_canonical_schema_info(Shard::Direct(0)), ))); } else { - self.select(statement, stmt, context) + self.select(stmt, context) } } @@ -442,7 +442,7 @@ impl QueryParser { Node::ExecuteStmt(stmt) => self.execute(stmt, context), - Node::ExplainStmt(stmt) => self.explain(statement, stmt, context), + Node::ExplainStmt(stmt) => self.explain(stmt, context), Node::DiscardStmt(stmt) => { let target = match stmt.target { diff --git a/pgdog/src/frontend/router/parser/query/prepare.rs b/pgdog/src/frontend/router/parser/query/prepare.rs index f570282e3..9adfde889 100644 --- a/pgdog/src/frontend/router/parser/query/prepare.rs +++ b/pgdog/src/frontend/router/parser/query/prepare.rs @@ -68,7 +68,7 @@ impl QueryParser { }; match stmt { - Some(Node::SelectStmt(stmt)) => self.select(&ast, stmt, &mut context), + Some(Node::SelectStmt(stmt)) => self.select(stmt, &mut context), Some(Node::InsertStmt(stmt)) => self.insert(stmt.into(), &mut context), Some(Node::UpdateStmt(stmt)) => self.update(stmt.into(), &mut context), Some(Node::DeleteStmt(stmt)) => self.delete(stmt.into(), &mut context), diff --git a/pgdog/src/frontend/router/parser/query/select.rs b/pgdog/src/frontend/router/parser/query/select.rs index b2b56bd84..70ac08517 100644 --- a/pgdog/src/frontend/router/parser/query/select.rs +++ b/pgdog/src/frontend/router/parser/query/select.rs @@ -1,5 +1,3 @@ -use crate::frontend::router::parser::cache::Ast; - use super::*; use pg_raw_parse::walk; use pg_raw_parse::{Node, nodes}; @@ -16,7 +14,6 @@ impl QueryParser { /// pub(super) fn select( &mut self, - _cached_ast: &Ast, stmt: &nodes::SelectStmt, context: &mut QueryParserContext, ) -> Result { diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/aggregate/engine.rs b/pgdog/src/frontend/router/parser/rewrite/statement/aggregate/engine.rs index d4b062a8d..c4205f1ed 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/aggregate/engine.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/aggregate/engine.rs @@ -4,9 +4,7 @@ use itertools::*; use pg_raw_parse::{Node, make, nodes}; use super::{AggregateHelper, HelperKind}; -use crate::frontend::router::parser::rewrite::statement::projection::{ - ProjectionRewritePlan, RewriteOutput, -}; +use crate::frontend::router::parser::rewrite::statement::projection::ProjectionRewritePlan; /// Query rewrite engine. Currently supports injecting helper aggregates for AVG and /// variance-related functions that require additional helper aggregates when run @@ -20,7 +18,7 @@ impl AggregatesRewrite { select: &mut nodes::SelectStmtMut<'a, '_>, mem: make::MemoryToken<'a>, aggregate: &Aggregate, - ) -> RewriteOutput { + ) -> ProjectionRewritePlan { let mut plan = ProjectionRewritePlan::default(); let helper_nodes = aggregate @@ -40,15 +38,13 @@ impl AggregatesRewrite { .into_iter() .map(move |spec| (target, spec)) }) - .enumerate() - .map(|(idx, (target, HelperSpec { func, kind }))| { + .map(|(target, HelperSpec { func, kind })| { let helper_alias = format!("__pgdog_{}_col{}", kind.alias_suffix(), target.column()); let node = mem.make_res_target(Some(&helper_alias), mem.empty(), func.uncast()); - plan.add_aggregate_helper(AggregateHelper { + plan.aggregate_helpers.push(AggregateHelper { target_column: target.column(), - projected_column: select.target_list().len() + idx, distinct: target.is_distinct(), kind, alias: helper_alias, @@ -57,14 +53,12 @@ impl AggregatesRewrite { }) .collect::>(); - if helper_nodes.is_empty() { - RewriteOutput::default() - } else { + if !helper_nodes.is_empty() { select .target_list_mut() .extend(mem, mem.make_list(&helper_nodes)); - RewriteOutput::new(plan) } + plan } fn build_sum_of_squares_func<'a>( @@ -166,10 +160,10 @@ mod tests { use crate::frontend::router::parser::aggregate::Aggregate; use pg_raw_parse::{Node, Owned, make, nodes}; - fn rewrite(sql: &str) -> (Owned, RewriteOutput) { + fn rewrite(sql: &str) -> (Owned, ProjectionRewritePlan) { let ast = pg_raw_parse::parse(sql).unwrap(); - let mut output = None; + let mut plan = None; let select = make::owned(|mem| { let Node::SelectStmt(stmt) = ast.stmts().next().unwrap() else { unreachable!("not a select"); @@ -177,32 +171,31 @@ mod tests { let mut stmt = mem.make_unique(stmt); let aggregate = Aggregate::parse(&stmt, &Default::default()); - output = Some(AggregatesRewrite::rewrite_select( + plan = Some(AggregatesRewrite::rewrite_select( &mut stmt.as_mut(), mem, &aggregate, )); stmt }); - (select, output.unwrap()) + (select, plan.unwrap()) } #[test] fn rewrite_engine_noop() { - let (ast, output) = rewrite("SELECT COUNT(price) FROM menu"); - assert!(output.plan.is_noop()); + let (ast, plan) = rewrite("SELECT COUNT(price) FROM menu"); + assert!(plan.is_noop()); assert_eq!(ast.target_list().len(), 1); } #[test] fn rewrite_engine_adds_helper() { - let (ast, output) = rewrite("SELECT AVG(price) FROM menu"); - assert!(!output.plan.is_noop()); - assert_eq!(output.plan.drop_columns().collect::>(), &[1]); - assert_eq!(output.plan.aggregate_helpers().len(), 1); - let helper = &output.plan.aggregate_helpers()[0]; + let (ast, plan) = rewrite("SELECT AVG(price) FROM menu"); + assert!(!plan.is_noop()); + assert_eq!(plan.aggregate_helpers.len(), 1); + let helper = &plan.aggregate_helpers[0]; assert_eq!(helper.target_column, 0); - assert_eq!(helper.projected_column, 1); + assert_eq!(helper.alias, "__pgdog_count_col0"); assert!(!helper.distinct); assert!(matches!(helper.kind, HelperKind::Count)); @@ -218,12 +211,11 @@ mod tests { #[test] fn rewrite_engine_handles_mismatched_pair() { - let (ast, output) = rewrite("SELECT COUNT(price::numeric), AVG(price) FROM menu"); - assert_eq!(output.plan.drop_columns().collect::>(), &[2]); - assert_eq!(output.plan.aggregate_helpers().len(), 1); - let helper = &output.plan.aggregate_helpers()[0]; + let (ast, plan) = rewrite("SELECT COUNT(price::numeric), AVG(price) FROM menu"); + assert_eq!(plan.aggregate_helpers.len(), 1); + let helper = &plan.aggregate_helpers[0]; assert_eq!(helper.target_column, 1); - assert_eq!(helper.projected_column, 2); + assert_eq!(helper.alias, "__pgdog_count_col1"); assert!(!helper.distinct); assert!(matches!(helper.kind, HelperKind::Count)); @@ -241,18 +233,17 @@ mod tests { #[test] fn rewrite_engine_multiple_avg_helpers() { - let (ast, output) = rewrite("SELECT AVG(price), AVG(discount) FROM menu"); - assert_eq!(output.plan.drop_columns().collect::>(), &[2, 3]); - assert_eq!(output.plan.aggregate_helpers().len(), 2); + let (ast, plan) = rewrite("SELECT AVG(price), AVG(discount) FROM menu"); + assert_eq!(plan.aggregate_helpers.len(), 2); - let helper_price = &output.plan.aggregate_helpers()[0]; + let helper_price = &plan.aggregate_helpers[0]; assert_eq!(helper_price.target_column, 0); - assert_eq!(helper_price.projected_column, 2); + assert_eq!(helper_price.alias, "__pgdog_count_col0"); assert!(matches!(helper_price.kind, HelperKind::Count)); - let helper_discount = &output.plan.aggregate_helpers()[1]; + let helper_discount = &plan.aggregate_helpers[1]; assert_eq!(helper_discount.target_column, 1); - assert_eq!(helper_discount.projected_column, 3); + assert_eq!(helper_discount.alias, "__pgdog_count_col1"); assert!(matches!(helper_discount.kind, HelperKind::Count)); let aggregate = Aggregate::parse(&ast, &Default::default()); @@ -269,14 +260,12 @@ mod tests { #[test] fn rewrite_engine_stddev_helpers() { - let (ast, output) = rewrite("SELECT STDDEV(price) FROM menu"); - assert!(!output.plan.is_noop()); - assert_eq!(output.plan.drop_columns().collect::>(), &[1, 2, 3]); - assert_eq!(output.plan.aggregate_helpers().len(), 3); - - let kinds: Vec = output - .plan - .aggregate_helpers() + let (ast, plan) = rewrite("SELECT STDDEV(price) FROM menu"); + assert!(!plan.is_noop()); + assert_eq!(plan.aggregate_helpers.len(), 3); + + let kinds: Vec = plan + .aggregate_helpers .iter() .map(|helper| { assert_eq!(helper.target_column, 0); diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/aggregate/mod.rs b/pgdog/src/frontend/router/parser/rewrite/statement/aggregate/mod.rs index 07da4bb48..5709b4589 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/aggregate/mod.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/aggregate/mod.rs @@ -4,14 +4,19 @@ pub(crate) use super::projection::AggregateHelper; pub(crate) use engine::AggregatesRewrite; +/// Type of aggregate function added to the result set. #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub(crate) enum HelperKind { + /// `COUNT(*)` or `COUNT(column)`. Count, + /// `SUM(column)`. Sum, + /// `SUM(POWER(column, 2))`. SumSquares, } impl HelperKind { + /// Suffix used in the projected helper's internal alias. pub(crate) fn alias_suffix(self) -> &'static str { match self { Self::Count => "count", diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/order_by.rs b/pgdog/src/frontend/router/parser/rewrite/statement/order_by.rs index a7007970d..7c0599d97 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/order_by.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/order_by.rs @@ -4,6 +4,7 @@ use crate::frontend::router::parser::Column; use super::projection::{OrderByHelper, OrderBySource, ProjectionRewritePlan}; +/// Compare complete column references so `a.price` and `b.price` remain distinct. fn same_column(left: &nodes::ColumnRef, right: &nodes::ColumnRef) -> bool { left.fields() .into_iter() @@ -11,6 +12,9 @@ fn same_column(left: &nodes::ColumnRef, right: &nodes::ColumnRef) -> bool { .eq(right.fields().into_iter().map(Node::as_str)) } +/// An unqualified `*` covers every column. A qualified star only covers an +/// equally-qualified ORDER BY reference; mixed qualification is handled +/// conservatively by projecting a helper and resolving its alias at runtime. fn star_covers(star: &nodes::ColumnRef, column: &nodes::ColumnRef) -> bool { let star_fields = star.fields(); let column_fields = column.fields(); @@ -30,6 +34,7 @@ fn star_covers(star: &nodes::ColumnRef, column: &nodes::ColumnRef) -> bool { .map(Node::as_str)))) } +/// Whether the SELECT list already returns the value needed for this sort. fn projects_column(select: &nodes::SelectStmtMut<'_, '_>, column: &nodes::ColumnRef) -> bool { let fields = column.fields(); let unqualified = fields.len() == 1; @@ -46,6 +51,8 @@ fn projects_column(select: &nodes::SelectStmtMut<'_, '_>, column: &nodes::Column }) } +/// Project missing sort values so PgDog can merge shard results. Helpers use +/// aliases rather than AST target positions because `*` expands only in Postgres. pub(super) fn rewrite_select<'a>( select: &mut nodes::SelectStmtMut<'a, '_>, mem: make::MemoryToken<'a>, @@ -91,17 +98,16 @@ pub(super) fn rewrite_select<'a>( continue; } - let projected_column = select.target_list().len() + helpers.len(); let alias = format!("__pgdog_order_col{current_sort_position}"); helpers.push(mem.make_res_target( Some(&alias), mem.empty(), mem.make_unique(node).uncast(), )); - plan.add_order_by_helper(OrderByHelper { + plan.order_by_helpers.push(OrderByHelper { sort_position: current_sort_position, source: source.expect("only columns and vector expressions need helpers"), - projected_column, + alias, }); } @@ -142,8 +148,8 @@ mod tests { let (sql, plan) = rewrite("SELECT id FROM products ORDER BY price"); assert!(sql.contains("price AS __pgdog_order_col0")); - assert_eq!(plan.order_by_helpers().len(), 1); - assert_eq!(plan.order_by_helpers()[0].projected_column, 1); + assert_eq!(plan.order_by_helpers.len(), 1); + assert_eq!(plan.order_by_helpers[0].alias, "__pgdog_order_col0"); } #[test] @@ -175,7 +181,15 @@ mod tests { let (sql, plan) = rewrite("SELECT a.* FROM a JOIN b ON a.id = b.a_id ORDER BY b.score"); assert!(sql.contains("b.score AS __pgdog_order_col0")); - assert_eq!(plan.order_by_helpers().len(), 1); + assert_eq!(plan.order_by_helpers.len(), 1); + } + + #[test] + fn projects_unqualified_sort_for_qualified_star() { + let (sql, plan) = rewrite("SELECT p.* FROM products p ORDER BY price"); + + assert!(sql.contains("price AS __pgdog_order_col0")); + assert_eq!(plan.order_by_helpers[0].alias, "__pgdog_order_col0"); } #[test] @@ -184,8 +198,8 @@ mod tests { rewrite("SELECT a.price FROM a JOIN b ON a.id = b.a_id ORDER BY a.price, b.price"); assert!(sql.contains("b.price AS __pgdog_order_col1")); - assert_eq!(plan.order_by_helpers().len(), 1); - assert_eq!(plan.order_by_helpers()[0].sort_position, 1); + assert_eq!(plan.order_by_helpers.len(), 1); + assert_eq!(plan.order_by_helpers[0].sort_position, 1); } #[test] @@ -194,6 +208,6 @@ mod tests { assert!(sql.contains("embedding <-> $1")); assert!(sql.contains("AS __pgdog_order_col0")); - assert_eq!(plan.order_by_helpers().len(), 1); + assert_eq!(plan.order_by_helpers.len(), 1); } } diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs b/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs index 5c207f0f5..fa13063e7 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs @@ -11,14 +11,14 @@ use super::order_by; use crate::backend::schema::Schema; use crate::frontend::router::parser::{Aggregate, OrderBy}; use crate::frontend::{ClientRequest, PreparedStatements}; -use crate::net::ProtocolMessage; +use crate::net::{ProtocolMessage, RowDescription}; use pg_raw_parse::{Node, StmtList, make}; +use std::collections::BTreeSet; use std::sync::Arc; #[derive(Debug, Clone, PartialEq)] pub(crate) struct AggregateHelper { pub(crate) target_column: usize, - pub(crate) projected_column: usize, pub(crate) distinct: bool, pub(crate) kind: HelperKind, pub(crate) alias: String, @@ -28,7 +28,7 @@ pub(crate) struct AggregateHelper { pub(crate) struct OrderByHelper { pub(crate) sort_position: usize, pub(crate) source: OrderBySource, - pub(crate) projected_column: usize, + pub(crate) alias: String, } #[derive(Debug, Clone, PartialEq)] @@ -52,8 +52,8 @@ impl OrderByHelper { #[derive(Debug, Clone, Default, PartialEq)] pub(crate) struct ProjectionRewritePlan { - aggregate_helpers: Vec, - order_by_helpers: Vec, + pub(crate) aggregate_helpers: Vec, + pub(crate) order_by_helpers: Vec, } impl ProjectionRewritePlan { @@ -61,37 +61,18 @@ impl ProjectionRewritePlan { self.aggregate_helpers.is_empty() && self.order_by_helpers.is_empty() } - pub(crate) fn drop_columns(&self) -> impl Iterator + '_ { + pub(crate) fn drop_columns(&self, row_description: &RowDescription) -> BTreeSet { self.aggregate_helpers .iter() - .map(|helper| helper.projected_column) + .map(|helper| helper.alias.as_str()) .chain( self.order_by_helpers .iter() - .map(|helper| helper.projected_column), + .map(|helper| helper.alias.as_str()), ) + .filter_map(|alias| row_description.field_index(alias)) + .collect() } - - pub(crate) fn aggregate_helpers(&self) -> &[AggregateHelper] { - &self.aggregate_helpers - } - - pub(crate) fn order_by_helpers(&self) -> &[OrderByHelper] { - &self.order_by_helpers - } - - pub(crate) fn add_aggregate_helper(&mut self, helper: AggregateHelper) { - self.aggregate_helpers.push(helper); - } - - pub(crate) fn add_order_by_helper(&mut self, helper: OrderByHelper) { - self.order_by_helpers.push(helper); - } -} - -#[derive(Debug, Default, Clone)] -pub(crate) struct RewriteOutput { - pub(crate) plan: ProjectionRewritePlan, } #[derive(Debug, Clone)] @@ -100,12 +81,6 @@ pub(crate) struct PostRouteRewrite { plan: ProjectionRewritePlan, } -impl RewriteOutput { - pub(crate) fn new(plan: ProjectionRewritePlan) -> Self { - Self { plan } - } -} - pub(crate) fn finalize_after_route( request: &mut ClientRequest, schema: &Schema, @@ -168,9 +143,9 @@ pub(crate) fn finalize_after_route( if !rewrite.plan.is_noop() && let Some(route) = request.route.as_mut() { - route.set_projection_rewrite_plan(rewrite.plan.clone()); + route.projection_rewrite_plan = rewrite.plan.clone(); let mut order_by = route.order_by().to_vec(); - for helper in rewrite.plan.order_by_helpers() { + for helper in &rewrite.plan.order_by_helpers { // Prefer the structural position. The source fallback handles a // bind-dependent vector sort omitted from this execution's route. let position = order_by @@ -182,9 +157,9 @@ pub(crate) fn finalize_after_route( continue; }; *sort = if sort.asc() { - OrderBy::Asc(helper.projected_column + 1) + OrderBy::AscColumn(helper.alias.clone()) } else { - OrderBy::Desc(helper.projected_column + 1) + OrderBy::DescColumn(helper.alias.clone()) }; } route.set_order_by(order_by); @@ -211,7 +186,7 @@ fn build( let rewritten = make::owned(|mem| { let mut select = mem.make_unique(select); if !aggregate.is_empty() { - plan = AggregatesRewrite::rewrite_select(&mut select.as_mut(), mem, &aggregate).plan; + plan = AggregatesRewrite::rewrite_select(&mut select.as_mut(), mem, &aggregate); } order_by::rewrite_select(&mut select.as_mut(), mem, &mut plan); if rewrite_offset { @@ -234,22 +209,26 @@ mod tests { #[test] fn projection_plan_tracks_helpers() { let mut plan = ProjectionRewritePlan::default(); - plan.add_aggregate_helper(AggregateHelper { + plan.aggregate_helpers.push(AggregateHelper { target_column: 0, - projected_column: 1, distinct: false, kind: HelperKind::Count, alias: "__pgdog_count_col0".into(), }); - plan.add_order_by_helper(OrderByHelper { + plan.order_by_helpers.push(OrderByHelper { sort_position: 0, source: OrderBySource::Column("created_at".into()), - projected_column: 2, + alias: "__pgdog_order_col0".into(), }); assert!(!plan.is_noop()); - assert_eq!(plan.drop_columns().collect::>(), [1, 2]); - assert_eq!(plan.aggregate_helpers().len(), 1); - assert_eq!(plan.order_by_helpers().len(), 1); + let row_description = RowDescription::new(&[ + crate::net::Field::double("avg"), + crate::net::Field::bigint("__pgdog_count_col0"), + crate::net::Field::timestamp("__pgdog_order_col0"), + ]); + assert_eq!(plan.drop_columns(&row_description), BTreeSet::from([1, 2])); + assert_eq!(plan.aggregate_helpers.len(), 1); + assert_eq!(plan.order_by_helpers.len(), 1); } } diff --git a/pgdog/src/frontend/router/parser/route.rs b/pgdog/src/frontend/router/parser/route.rs index 8f6c6b0f5..382fa94cf 100644 --- a/pgdog/src/frontend/router/parser/route.rs +++ b/pgdog/src/frontend/router/parser/route.rs @@ -111,7 +111,7 @@ pub(crate) struct Route { /// `DISTINCT` clause, if set. distinct: Option, /// Temporary columns projected for cross-shard result processing. - projection_rewrite_plan: ProjectionRewritePlan, + pub(crate) projection_rewrite_plan: ProjectionRewritePlan, /// Our query explain plan. We attach /// this to the `EXPLAIN` output. explain: Option, @@ -404,14 +404,6 @@ impl Route { self.is_cross_shard() && self.is_write() } - pub(crate) fn projection_rewrite_plan(&self) -> &ProjectionRewritePlan { - &self.projection_rewrite_plan - } - - pub(crate) fn set_projection_rewrite_plan(&mut self, plan: ProjectionRewritePlan) { - self.projection_rewrite_plan = plan; - } - pub(super) fn with_temp_table_change(mut self, temp_table: Option) -> Self { self.temp_table_change = temp_table; self From aec5f304f3e12ee757bbd91c8be2515e51787834 Mon Sep 17 00:00:00 2001 From: Nupur Agrawal Date: Tue, 22 Sep 2026 17:59:14 +0530 Subject: [PATCH 13/14] more fixes --- .../pool/connection/multi_shard/test.rs | 1 + .../query_engine/test/rewrite_projection.rs | 53 ++++ .../parser/rewrite/statement/order_by.rs | 267 +++++++++++++++--- .../parser/rewrite/statement/projection.rs | 22 ++ 4 files changed, 301 insertions(+), 42 deletions(-) diff --git a/pgdog/src/backend/pool/connection/multi_shard/test.rs b/pgdog/src/backend/pool/connection/multi_shard/test.rs index e2a73b901..52b5df9d2 100644 --- a/pgdog/src/backend/pool/connection/multi_shard/test.rs +++ b/pgdog/src/backend/pool/connection/multi_shard/test.rs @@ -70,6 +70,7 @@ fn test_order_by_helper_after_star_expansion_is_dropped_after_sorting() { sort_position: 0, source: OrderBySource::Column("price".into()), alias: "__pgdog_order_col0".into(), + injected: true, }); let mut route = Route::select( ShardWithPriority::new_default_unset(Shard::All), diff --git a/pgdog/src/frontend/client/query_engine/test/rewrite_projection.rs b/pgdog/src/frontend/client/query_engine/test/rewrite_projection.rs index 533efb683..0dcdec5cb 100644 --- a/pgdog/src/frontend/client/query_engine/test/rewrite_projection.rs +++ b/pgdog/src/frontend/client/query_engine/test/rewrite_projection.rs @@ -298,6 +298,59 @@ fn cached_projection_does_not_depend_on_first_route_order() { ); } +#[test] +fn aliased_projected_sort_column_remaps_route() { + let sql = "SELECT price AS item_price FROM products ORDER BY price"; + let mut request = ClientRequest::from(vec![ProtocolMessage::Query(Query::new(sql))]); + request.ast = Some(Ast::new_record(sql).unwrap()); + request.route = Some(Route::select( + ShardWithPriority::new_table(Shard::All), + vec![OrderBy::AscColumn("price".into())], + Default::default(), + Limit::default(), + None, + )); + + projection::finalize_after_route(&mut request, &Schema::default(), None).unwrap(); + + let query = match &request.messages[0] { + ProtocolMessage::Query(query) => query, + _ => panic!("expected Query"), + }; + assert!(!query.query().contains("__pgdog_order_col")); + assert_eq!( + request.route().order_by(), + &[OrderBy::AscColumn("item_price".into())] + ); + assert!(!request.route().projection_rewrite_plan.order_by_helpers[0].injected); +} + +#[test] +fn duplicate_sort_column_names_use_injected_helper() { + let sql = "SELECT a.price, b.price FROM a JOIN b ON a.id = b.a_id ORDER BY b.price"; + let mut request = ClientRequest::from(vec![ProtocolMessage::Query(Query::new(sql))]); + request.ast = Some(Ast::new_record(sql).unwrap()); + request.route = Some(Route::select( + ShardWithPriority::new_table(Shard::All), + vec![OrderBy::AscColumn("price".into())], + Default::default(), + Limit::default(), + None, + )); + + projection::finalize_after_route(&mut request, &Schema::default(), None).unwrap(); + + let query = match &request.messages[0] { + ProtocolMessage::Query(query) => query, + _ => panic!("expected Query"), + }; + assert!(query.query().contains("b.price AS __pgdog_order_col0")); + assert_eq!( + request.route().order_by(), + &[OrderBy::AscColumn("__pgdog_order_col0".into())] + ); +} + #[test] fn helper_replaces_the_matching_duplicate_order_by_position() { let sql = "SELECT a.price FROM a JOIN b ON a.id = b.a_id ORDER BY a.price, b.price"; diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/order_by.rs b/pgdog/src/frontend/router/parser/rewrite/statement/order_by.rs index 7c0599d97..25e3f687f 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/order_by.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/order_by.rs @@ -4,6 +4,29 @@ use crate::frontend::router::parser::Column; use super::projection::{OrderByHelper, OrderBySource, ProjectionRewritePlan}; +impl OrderBySource { + fn name(&self) -> &str { + match self { + Self::Column(name) | Self::Vector(name) => name, + } + } +} + +fn column_name(column: &nodes::ColumnRef) -> Option<&str> { + column + .fields() + .into_iter() + .next_back() + .and_then(Node::as_str) +} + +fn is_star(column: &nodes::ColumnRef) -> bool { + matches!( + column.fields().into_iter().next_back(), + Some(Node::A_Star(_)) + ) +} + /// Compare complete column references so `a.price` and `b.price` remain distinct. fn same_column(left: &nodes::ColumnRef, right: &nodes::ColumnRef) -> bool { left.fields() @@ -34,25 +57,106 @@ fn star_covers(star: &nodes::ColumnRef, column: &nodes::ColumnRef) -> bool { .map(Node::as_str)))) } -/// Whether the SELECT list already returns the value needed for this sort. -fn projects_column(select: &nodes::SelectStmtMut<'_, '_>, column: &nodes::ColumnRef) -> bool { - let fields = column.fields(); - let unqualified = fields.len() == 1; - let name = fields.into_iter().next_back().and_then(Node::as_str); - - select.target_list().iter().any(|target| { - (unqualified && target.name() == name) - || matches!(target.val(), Node::ColumnRef(projected) - if star_covers(projected, column) - || same_column(projected, column) - || (unqualified - && projected.fields().into_iter().next_back().and_then(Node::as_str) - == name)) - }) +fn target_output_name(target: &nodes::ResTarget) -> Option<&str> { + if let Some(alias) = target.name() { + return Some(alias); + } + match target.val() { + Node::ColumnRef(column) if !is_star(column) => column_name(column), + _ => None, + } +} + +fn single_relation(select: &nodes::SelectStmtMut<'_, '_>) -> bool { + let from = select.from_clause(); + from.len() == 1 && matches!(from.first(), Some(Node::RangeVar(_))) } -/// Project missing sort values so PgDog can merge shard results. Helpers use -/// aliases rather than AST target positions because `*` expands only in Postgres. +fn explicit_output_name( + target: &nodes::ResTarget, + column: &nodes::ColumnRef, + name: &str, + unqualified: bool, +) -> Option { + if unqualified && target.name() == Some(name) { + return Some(name.to_owned()); + } + let Node::ColumnRef(projected) = target.val() else { + return None; + }; + if is_star(projected) { + return None; + } + if same_column(projected, column) || (unqualified && column_name(projected) == Some(name)) { + Some(target.name().unwrap_or(name).to_owned()) + } else { + None + } +} + +/// Cross-shard sort looks up a RowDescription name. Reuse a projected column +/// only when that name is unique; stars and duplicate names get a helper. +fn unique_output_name( + select: &nodes::SelectStmtMut<'_, '_>, + column: &nodes::ColumnRef, +) -> Option { + let name = column_name(column)?; + let unqualified = column.fields().len() == 1; + + let mut star_match = false; + let mut output = None; + for target in select.target_list() { + if let Some(found) = explicit_output_name(target, column, name, unqualified) { + output = Some(found); + star_match = false; + break; + } + if matches!(target.val(), Node::ColumnRef(projected) if star_covers(projected, column)) { + star_match = true; + } + } + let output = output.or_else(|| star_match.then(|| name.to_owned()))?; + + let mut sources = 0; + let mut unqualified_star = false; + for target in select.target_list() { + if let Node::ColumnRef(projected) = target.val() + && is_star(projected) + { + sources += 1; + if projected.fields().len() == 1 { + unqualified_star = true; + } + } else if target_output_name(target) == Some(output.as_str()) { + sources += 1; + } + } + + if sources == 1 && (!unqualified_star || single_relation(select)) { + Some(output) + } else { + None + } +} + +fn push_helper( + plan: &mut ProjectionRewritePlan, + sort_position: usize, + source: OrderBySource, + alias: String, + injected: bool, +) { + plan.order_by_helpers.push(OrderByHelper { + sort_position, + source, + alias, + injected, + }); +} + +/// Project missing or ambiguous sort values so PgDog can merge shard results. +/// Helpers use aliases rather than AST target positions because `*` expands +/// only in Postgres. pub(super) fn rewrite_select<'a>( select: &mut nodes::SelectStmtMut<'a, '_>, mem: make::MemoryToken<'a>, @@ -63,12 +167,9 @@ pub(super) fn rewrite_select<'a>( for sort in select.sort_clause() { let node = sort.node(); let source = match node { - Node::ColumnRef(column) => column - .fields() - .into_iter() - .next_back() - .and_then(Node::as_str) - .map(|name| OrderBySource::Column(name.to_owned())), + Node::ColumnRef(column) => { + column_name(column).map(|name| OrderBySource::Column(name.to_owned())) + } Node::A_Expr(expr) if expr.name().iter().next().and_then(Node::as_str) == Some("<->") => { @@ -89,26 +190,49 @@ pub(super) fn rewrite_select<'a>( let current_sort_position = sort_position; sort_position += 1; - let needs_helper = match node { - Node::ColumnRef(column) => !projects_column(select, column), - Node::A_Expr(_) => source.is_some(), - _ => false, - }; - if !needs_helper { - continue; + match node { + Node::ColumnRef(column) => match unique_output_name(select, column) { + Some(name) if source.as_ref().is_some_and(|source| source.name() == name) => {} + Some(name) => push_helper( + plan, + current_sort_position, + source.expect("column sorts always have a source"), + name, + false, + ), + None => { + let alias = format!("__pgdog_order_col{current_sort_position}"); + helpers.push(mem.make_res_target( + Some(&alias), + mem.empty(), + mem.make_unique(node).uncast(), + )); + push_helper( + plan, + current_sort_position, + source.expect("column sorts always have a source"), + alias, + true, + ); + } + }, + Node::A_Expr(_) if source.is_some() => { + let alias = format!("__pgdog_order_col{current_sort_position}"); + helpers.push(mem.make_res_target( + Some(&alias), + mem.empty(), + mem.make_unique(node).uncast(), + )); + push_helper( + plan, + current_sort_position, + source.expect("vector sorts always have a source"), + alias, + true, + ); + } + _ => {} } - - let alias = format!("__pgdog_order_col{current_sort_position}"); - helpers.push(mem.make_res_target( - Some(&alias), - mem.empty(), - mem.make_unique(node).uncast(), - )); - plan.order_by_helpers.push(OrderByHelper { - sort_position: current_sort_position, - source: source.expect("only columns and vector expressions need helpers"), - alias, - }); } if !helpers.is_empty() { @@ -150,6 +274,7 @@ mod tests { assert!(sql.contains("price AS __pgdog_order_col0")); assert_eq!(plan.order_by_helpers.len(), 1); assert_eq!(plan.order_by_helpers[0].alias, "__pgdog_order_col0"); + assert!(plan.order_by_helpers[0].injected); } #[test] @@ -160,6 +285,63 @@ mod tests { assert!(plan.is_noop()); } + #[test] + fn remaps_aliased_projected_sort_column() { + let (sql, plan) = rewrite("SELECT price AS item_price FROM products ORDER BY price"); + + assert!(!sql.contains("__pgdog_order_col")); + assert_eq!(plan.order_by_helpers.len(), 1); + assert_eq!(plan.order_by_helpers[0].alias, "item_price"); + assert!(!plan.order_by_helpers[0].injected); + } + + #[test] + fn skips_sort_by_output_alias() { + let (sql, plan) = rewrite("SELECT price AS item_price FROM products ORDER BY item_price"); + + assert!(!sql.contains("__pgdog_order_col")); + assert!(plan.is_noop()); + } + + #[test] + fn injects_helper_for_duplicate_output_names() { + let (sql, plan) = + rewrite("SELECT a.price, b.price FROM a JOIN b ON a.id = b.a_id ORDER BY b.price"); + + assert!(sql.contains("b.price AS __pgdog_order_col0")); + assert_eq!(plan.order_by_helpers.len(), 1); + assert!(plan.order_by_helpers[0].injected); + } + + #[test] + fn remaps_unique_alias_among_same_named_columns() { + let (sql, plan) = rewrite( + "SELECT a.price AS a_price, b.price AS b_price FROM a JOIN b ON a.id = b.a_id ORDER BY b.price", + ); + + assert!(!sql.contains("__pgdog_order_col")); + assert_eq!(plan.order_by_helpers.len(), 1); + assert_eq!(plan.order_by_helpers[0].alias, "b_price"); + assert!(!plan.order_by_helpers[0].injected); + } + + #[test] + fn injects_helper_when_star_can_collide() { + let (sql, plan) = + rewrite("SELECT a.*, b.price FROM a JOIN b ON a.id = b.a_id ORDER BY b.price"); + + assert!(sql.contains("b.price AS __pgdog_order_col0")); + assert!(plan.order_by_helpers[0].injected); + } + + #[test] + fn injects_helper_for_unqualified_star_join() { + let (sql, plan) = rewrite("SELECT * FROM a JOIN b ON a.id = b.a_id ORDER BY b.price"); + + assert!(sql.contains("b.price AS __pgdog_order_col0")); + assert!(plan.order_by_helpers[0].injected); + } + #[test] fn skips_star_select() { let (sql, plan) = rewrite("SELECT * FROM products ORDER BY id"); @@ -209,5 +391,6 @@ mod tests { assert!(sql.contains("embedding <-> $1")); assert!(sql.contains("AS __pgdog_order_col0")); assert_eq!(plan.order_by_helpers.len(), 1); + assert!(plan.order_by_helpers[0].injected); } } diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs b/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs index fa13063e7..772246586 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs @@ -29,6 +29,9 @@ pub(crate) struct OrderByHelper { pub(crate) sort_position: usize, pub(crate) source: OrderBySource, pub(crate) alias: String, + /// False when the SELECT list already has this unique name and we only + /// remap the route. Those columns must stay in the client result. + pub(crate) injected: bool, } #[derive(Debug, Clone, PartialEq)] @@ -68,6 +71,7 @@ impl ProjectionRewritePlan { .chain( self.order_by_helpers .iter() + .filter(|helper| helper.injected) .map(|helper| helper.alias.as_str()), ) .filter_map(|alias| row_description.field_index(alias)) @@ -219,6 +223,7 @@ mod tests { sort_position: 0, source: OrderBySource::Column("created_at".into()), alias: "__pgdog_order_col0".into(), + injected: true, }); assert!(!plan.is_noop()); @@ -231,4 +236,21 @@ mod tests { assert_eq!(plan.aggregate_helpers.len(), 1); assert_eq!(plan.order_by_helpers.len(), 1); } + + #[test] + fn remapped_order_by_alias_is_not_dropped() { + let mut plan = ProjectionRewritePlan::default(); + plan.order_by_helpers.push(OrderByHelper { + sort_position: 0, + source: OrderBySource::Column("price".into()), + alias: "item_price".into(), + injected: false, + }); + + let row_description = RowDescription::new(&[ + crate::net::Field::bigint("id"), + crate::net::Field::numeric("item_price"), + ]); + assert!(plan.drop_columns(&row_description).is_empty()); + } } From c93158aeda437805a5465b3008c188c5e9e0ef61 Mon Sep 17 00:00:00 2001 From: Nupur Agrawal Date: Tue, 22 Sep 2026 19:35:15 +0530 Subject: [PATCH 14/14] clippy --- pgdog/src/backend/pool/connection/aggregate.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pgdog/src/backend/pool/connection/aggregate.rs b/pgdog/src/backend/pool/connection/aggregate.rs index 0c908a5b3..7efddb306 100644 --- a/pgdog/src/backend/pool/connection/aggregate.rs +++ b/pgdog/src/backend/pool/connection/aggregate.rs @@ -571,7 +571,7 @@ mod test { ); rows.push_back(row); } - let plan = AggregateRewritePlan::default(); + let plan = ProjectionRewritePlan::default(); let mut result = Aggregates::new(&rows, &decoder, &aggregate, &plan) .expect("count aggregate") .aggregate()