diff --git a/integration/rust/tests/integration/limit.rs b/integration/rust/tests/integration/limit.rs index 3bbcba93a..ddf084278 100644 --- a/integration/rust/tests/integration/limit.rs +++ b/integration/rust/tests/integration/limit.rs @@ -27,6 +27,37 @@ async fn limit_across_shards() -> Result<(), Box> { ) .await?; + let rows = sharded + .fetch_all("SELECT id FROM limit_test ORDER BY value") + .await?; + assert_eq!( + rows.iter() + .map(|row| { + assert_eq!(row.len(), 1); + row.get::("id") + }) + .collect::>(), + vec![1, 2, 3, 4, 5, 1, 2, 3, 4, 5] + ); + + let rows = sharded + .fetch_all("/* pgdog_shard: 0 */ SELECT id FROM limit_test ORDER BY value") + .await?; + assert_eq!( + rows.iter() + .map(|row| { + assert_eq!(row.len(), 1); + row.get::("id") + }) + .collect::>(), + vec![1, 2, 3, 4, 5] + ); + + let row = sharded + .fetch_one("/* pgdog_shard: 0 */ SELECT stddev(value) FROM limit_test") + .await?; + assert_eq!(row.len(), 1); + // LIMIT 5 let rows = sharded .fetch_all("SELECT value FROM limit_test ORDER BY value LIMIT 5") diff --git a/pgdog/src/backend/pool/connection/aggregate.rs b/pgdog/src/backend/pool/connection/aggregate.rs index c06bf594a..7b6544922 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, @@ -237,7 +237,7 @@ pub(super) struct Aggregates<'a> { mappings: HashMap>, decoder: &'a Decoder, aggregate: &'a Aggregate, - helper_columns: HashMap, + projected_columns: HashMap, } impl<'a> Aggregates<'a> { @@ -245,29 +245,29 @@ 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(); + let mut projected_columns: HashMap = HashMap::new(); for target in aggregate.targets() { let key = target.column(); match target.function() { AggregateFunction::Count => { - helper_columns.entry(key).or_default().count = Some(target.column()); + projected_columns.entry(key).or_default().count = Some(target.column()); } AggregateFunction::Sum => { - helper_columns.entry(key).or_default().sum = Some(target.column()); + projected_columns.entry(key).or_default().sum = Some(target.column()); } _ => {} } } - for helper in plan.helpers() { + for helper in plan.aggregate_helpers() { let Some(index) = decoder.row_description().field_index(&helper.alias) else { continue; }; - let entry = helper_columns.entry(helper.target_column).or_default(); + let entry = projected_columns.entry(helper.target_column).or_default(); match helper.kind { HelperKind::Count => entry.count = Some(index), HelperKind::Sum => entry.sum = Some(index), @@ -278,14 +278,14 @@ impl<'a> Aggregates<'a> { let helpers_present = aggregate.targets().iter().all(|target| { let key = target.column(); match target.function() { - AggregateFunction::Avg => helper_columns + AggregateFunction::Avg => projected_columns .get(&key) .and_then(|columns| columns.count) .is_some(), AggregateFunction::StddevPop | AggregateFunction::StddevSamp | AggregateFunction::VarPop - | AggregateFunction::VarSamp => helper_columns + | AggregateFunction::VarSamp => projected_columns .get(&key) .map(|columns| { columns.count.is_some() && columns.sum.is_some() && columns.sumsq.is_some() @@ -305,7 +305,7 @@ impl<'a> Aggregates<'a> { decoder, mappings: HashMap::new(), aggregate, - helper_columns, + projected_columns, }) } else { None @@ -333,7 +333,7 @@ impl<'a> Aggregates<'a> { Entry::Occupied(o) => o.into_mut(), Entry::Vacant(v) => { let accumulators = - Accumulator::from_aggregate(self.aggregate, &self.helper_columns)?; + Accumulator::from_aggregate(self.aggregate, &self.projected_columns)?; // Gather all col vals corresponding to list of passthrough indices. let mut passthrough = Vec::new(); @@ -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..983efed57 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,10 +140,10 @@ 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() { + 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: &AggregateRewritePlan) { + pub(super) fn drop_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); } } @@ -240,6 +239,7 @@ impl Buffer { #[cfg(test)] mod test { use super::*; + use crate::frontend::router::parser::rewrite::statement::projection::OrderByHelper; use crate::net::{Datum, Field, Format, RowDescription}; use bytes::Bytes; @@ -273,6 +273,37 @@ mod test { assert_eq!(i, 26); } + #[test] + fn test_sort_by_hidden_column_before_dropping_it() { + let mut buf = Buffer::default(); + let rd = RowDescription::new(&[Field::bigint("id"), Field::text("__pgdog_order_by_0")]); + let decoder = Decoder::from(rd); + let mut plan = ProjectionRewritePlan::default(); + plan.add_order_by_helper(OrderByHelper { + sort_position: 0, + projected_column: 1, + }); + + for (id, name) in [(1_i64, "z"), (2, "a"), (3, "m")] { + let mut row = DataRow::new(); + row.add(id).add(name); + buf.add(row.message()).unwrap(); + } + + buf.sort(&[OrderBy::Asc(2)], &decoder); + buf.drop_columns(&plan); + buf.mark_full(); + + let ids = std::iter::from_fn(|| buf.take()) + .map(|message| { + let row = DataRow::from_bytes(message.to_bytes()).unwrap(); + assert_eq!(row.len(), 1); + row.get::(0, Format::Text).unwrap() + }) + .collect::>(); + assert_eq!(ids, [2, 3, 1]); + } + #[test] fn test_aggregate_buffer() { let mut buf = Buffer::default(); @@ -285,7 +316,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 +343,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..bc4bc4e28 100644 --- a/pgdog/src/backend/pool/connection/multi_shard/mod.rs +++ b/pgdog/src/backend/pool/connection/multi_shard/mod.rs @@ -264,13 +264,15 @@ impl MultiShard { .aggregate( self.route.aggregate(), &self.decoder, - self.route.aggregate_rewrite_plan(), + self.route.projection_rewrite_plan(), ) .map_err(Error::from)?; self.buffer.sort(self.route.order_by(), &self.decoder); self.buffer.distinct(self.route.distinct(), &self.decoder); self.buffer.limit(self.route.limit()); + self.buffer + .drop_columns(self.route.projection_rewrite_plan()); } if has_rows { @@ -311,7 +313,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 { @@ -342,7 +344,7 @@ impl MultiShard { ) } - fn handle_data_row(&mut self, message: Message) -> Result, Error> { + fn handle_data_row(&mut self, mut message: Message) -> Result, Error> { let mut forward = None; if self.shards > 1 { @@ -366,9 +368,11 @@ impl MultiShard { { if self.route.is_omnisharded() { if self.request_state.first_backend_data == message.source().backend_id() { + self.drop_columns(&mut message)?; forward = Some(message); } } else { + self.drop_columns(&mut message)?; forward = Some(message); } } else { @@ -378,6 +382,18 @@ impl MultiShard { Ok(forward) } + fn drop_columns(&self, message: &mut Message) -> Result<(), Error> { + let plan = self.route.projection_rewrite_plan(); + if plan.is_noop() { + return Ok(()); + } + + let mut row = DataRow::from_bytes(message.to_bytes())?; + row.drop_columns(&plan.drop_columns().collect()); + message.replace_payload(row.to_bytes()); + Ok(()) + } + fn handle_passthrough(message: Message, counter: &mut usize, shards: usize) -> Option { *counter += 1; diff --git a/pgdog/src/frontend/client/query_engine/query.rs b/pgdog/src/frontend/client/query_engine/query.rs index aa88e8f99..399c64df4 100644 --- a/pgdog/src/frontend/client/query_engine/query.rs +++ b/pgdog/src/frontend/client/query_engine/query.rs @@ -3,7 +3,10 @@ use tracing::{info, trace}; use crate::{ frontend::{ client::TransactionType, - router::parser::{explain_trace::ExplainTrace, rewrite::statement::plan::RewriteResult}, + router::parser::{ + explain_trace::ExplainTrace, + rewrite::statement::{plan::RewriteResult, projection::ProjectionRewritePlan}, + }, }, net::{ DataRow, FromBytes, Message, Protocol, ProtocolMessage, Query, ReadyForQuery, @@ -13,6 +16,7 @@ use crate::{ util::safe_timeout, }; +use std::collections::BTreeSet; use tracing::{debug, error}; use super::hooks::schema::schema_changed; @@ -146,6 +150,12 @@ impl QueryEngine { context: &mut QueryEngineContext<'_>, mut message: Message, ) -> Result<(), Error> { + if !self.backend.is_multishard() { + drop_projected_columns( + &mut message, + context.client_request.route().projection_rewrite_plan(), + )?; + } self.streaming = message.streaming(); let code = message.code(); @@ -515,6 +525,30 @@ impl QueryEngine { } } +fn drop_projected_columns( + message: &mut Message, + plan: &ProjectionRewritePlan, +) -> Result<(), Error> { + if plan.is_noop() { + return Ok(()); + } + + let drop = plan.drop_columns().collect::>(); + let payload = match message.code() { + 'D' => { + let mut row = DataRow::from_bytes(message.to_bytes())?; + row.drop_columns(&drop); + row.to_bytes() + } + 'T' => RowDescription::from_bytes(message.to_bytes())? + .drop_columns(drop) + .to_bytes(), + _ => return Ok(()), + }; + message.replace_payload(payload); + Ok(()) +} + #[derive(Debug, Default, Clone)] pub(super) struct ExplainResponseState { lines: Vec, @@ -547,3 +581,36 @@ impl ExplainResponseState { self.supported && !self.annotated } } + +#[cfg(test)] +mod tests { + use super::*; + use crate::{ + frontend::router::parser::rewrite::statement::projection::OrderByHelper, + net::{Field, Format}, + }; + + #[test] + fn hidden_columns_are_removed_at_client_output_boundary() { + let mut plan = ProjectionRewritePlan::default(); + plan.add_order_by_helper(OrderByHelper { + sort_position: 0, + projected_column: 1, + }); + + let description = + RowDescription::new(&[Field::bigint("id"), Field::text("__pgdog_order_by_0")]); + let mut description = description.message(); + drop_projected_columns(&mut description, &plan).unwrap(); + let description = RowDescription::from_bytes(description.to_bytes()).unwrap(); + assert_eq!(description.fields.len(), 1); + + let mut row = DataRow::new(); + row.add(42_i64).add("alice"); + let mut row = row.message(); + drop_projected_columns(&mut row, &plan).unwrap(); + let row = DataRow::from_bytes(row.to_bytes()).unwrap(); + assert_eq!(row.len(), 1); + assert_eq!(row.get::(0, Format::Text), Some(42)); + } +} diff --git a/pgdog/src/frontend/router/parser/query/select.rs b/pgdog/src/frontend/router/parser/query/select.rs index 5dceda733..9c0a09a0e 100644 --- a/pgdog/src/frontend/router/parser/query/select.rs +++ b/pgdog/src/frontend/router/parser/query/select.rs @@ -66,13 +66,13 @@ impl QueryParser { // Early return for any direct-to-shard queries. if context.shards_calculator.shard().is_direct() { - return Ok(Command::Query( - Route::read(context.shards_calculator.shard().clone()) - .with_read(!writes) - .with_mutates(mutates) - .with_omnisharded(omnisharded) - .with_advisory_locks(advisory_locks), - )); + let mut route = Route::read(context.shards_calculator.shard().clone()) + .with_read(!writes) + .with_mutates(mutates) + .with_omnisharded(omnisharded) + .with_advisory_locks(advisory_locks); + route.set_projection_rewrite_plan(cached_ast.rewrite_plan.projection.clone()); + return Ok(Command::Query(route)); } let mut shards = HashSet::new(); @@ -135,16 +135,33 @@ impl QueryParser { .push(ShardWithPriority::new_rr_no_table(shard)); } - return Ok(Command::Query( - Route::read(context.shards_calculator.shard().clone()) - .with_read(!writes) - .with_mutates(mutates) - .with_omnisharded(omnisharded) - .with_advisory_locks(advisory_locks), - )); + let mut route = Route::read(context.shards_calculator.shard().clone()) + .with_read(!writes) + .with_mutates(mutates) + .with_omnisharded(omnisharded) + .with_advisory_locks(advisory_locks); + route.set_projection_rewrite_plan(cached_ast.rewrite_plan.projection.clone()); + return Ok(Command::Query(route)); } - let order_by = Self::select_sort(stmt, context.router_context.bind); + let mut order_by = Self::select_sort(stmt, context.router_context.bind); + for helper in cached_ast.rewrite_plan.projection.order_by_helpers() { + let Some((_, column)) = order_by + .iter_mut() + .find(|(position, _)| *position == helper.sort_position) + else { + continue; + }; + *column = if column.asc() { + OrderBy::Asc(helper.projected_column + 1) + } else { + OrderBy::Desc(helper.projected_column + 1) + }; + } + let order_by = order_by + .into_iter() + .map(|(_, order_by)| order_by) + .collect::>(); let from_clause_table_name = stmt.from_clause().first().and_then(|node| match node { Node::RangeVar(r) => Some(r.relname().expect("RangeVar always has relname")), _ => None, @@ -277,10 +294,7 @@ impl QueryParser { distinct, ); - // 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( query @@ -301,17 +315,18 @@ impl QueryParser { fn select_sort( stmt: &nodes::SelectStmt, params: Option>, - ) -> Vec { + ) -> Vec<(usize, OrderBy)> { stmt.sort_clause() .into_iter() - .filter_map(|sort_by| { + .enumerate() + .filter_map(|(position, sort_by)| { use pg_raw_parse::{ ConstValue, raw::{A_Expr_Kind::*, SortByDir::*}, }; let asc = matches!(sort_by.sortby_dir, SORTBY_DEFAULT | SORTBY_ASC); - match sort_by.node() { + let order_by = match sort_by.node() { Node::A_Const(c) if let Some(ConstValue::Integer(i)) = c.val() => { if asc { Some(OrderBy::Asc(i as _)) @@ -363,7 +378,8 @@ impl QueryParser { } _ => None, - } + }; + order_by.map(|order_by| (position, order_by)) }) .collect() } diff --git a/pgdog/src/frontend/router/parser/query/test/test_select.rs b/pgdog/src/frontend/router/parser/query/test/test_select.rs index 6bfd0b8d0..1465a4641 100644 --- a/pgdog/src/frontend/router/parser/query/test/test_select.rs +++ b/pgdog/src/frontend/router/parser/query/test/test_select.rs @@ -20,6 +20,46 @@ fn test_order_by_vector_simple() { assert!(order_by.asc()); } +#[test] +fn test_order_by_non_projected_column_uses_rewrite_helper() { + let mut test = QueryParserTest::new(); + + let command = test.execute(vec![ + Query::new("SELECT id FROM sharded ORDER BY value").into(), + ]); + + let route = command.route(); + assert_eq!( + route.order_by().first().and_then(|order| order.index()), + Some(1) + ); + assert_eq!( + route + .projection_rewrite_plan() + .drop_columns() + .collect::>(), + [1] + ); +} + +#[test] +fn test_order_by_helper_tracks_original_clause_position() { + let mut test = QueryParserTest::new(); + + let command = test.execute(vec![ + Query::new("SELECT id FROM sharded ORDER BY lower(value), value").into(), + ]); + + assert_eq!( + command + .route() + .order_by() + .first() + .and_then(|order| order.index()), + Some(1) + ); +} + #[test] fn test_order_by_vector_with_params() { let mut test = QueryParserTest::new(); 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..886a5582e 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::new(); 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..3911b026c 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,12 @@ mod engine; -mod plan; 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 super::projection::{AggregateHelper, HelperKind}; pub(crate) use engine::AggregatesRewrite; -pub(crate) use plan::{AggregateRewritePlan, HelperKind, HelperMapping, RewriteOutput}; impl StatementRewrite<'_> { /// Add missing COUNT(*) and other helps when using aggregates. @@ -32,7 +31,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 87c600ca9..507c6447c 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs @@ -14,7 +14,9 @@ pub(crate) mod error; pub(crate) mod insert; pub(crate) mod nextval; pub(crate) mod offset; +mod order_by; pub(crate) mod plan; +pub(crate) mod projection; pub(crate) mod simple_prepared; pub(crate) mod unique_id; pub(crate) mod update; @@ -183,6 +185,7 @@ impl<'a> StatementRewrite<'a> { if let NodeMut::SelectStmt(mut select) = stmt.stmt_mut() { self.rewrite_aggregates(&mut select, mem, &mut plan, self.db_schema)?; + self.rewrite_order_by(&mut select, mem, &mut plan); self.limit_offset(&select, &mut plan); } 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..0ceed0a98 --- /dev/null +++ b/pgdog/src/frontend/router/parser/rewrite/statement/order_by.rs @@ -0,0 +1,190 @@ +use pg_raw_parse::{Node, make, nodes}; + +use super::{RewritePlan, StatementRewrite, projection::OrderByHelper}; +use crate::frontend::router::parser::Column; + +impl StatementRewrite<'_> { + /// Project simple ORDER BY columns that are absent from the result so the + /// cross-shard merge can sort by their values. + pub(super) fn rewrite_order_by<'a>( + &mut self, + select: &mut nodes::SelectStmtMut<'a, '_>, + mem: make::MemoryToken<'a>, + plan: &mut RewritePlan, + ) { + if self.schema.shards == 1 || matches!(select.distinct_clause().first(), Some(Node::None)) { + return; + } + + let original_target_len = select.target_list().len(); + let helpers = select + .sort_clause() + .iter() + .enumerate() + .filter_map(|(order_by, sort)| { + let Node::ColumnRef(column_ref) = sort.node() else { + return None; + }; + let column = Column::try_from(column_ref).ok()?; + + if column_is_projected(select, column) { + return None; + } + + Some((order_by, sort.node())) + }) + .enumerate() + .map(|(helper_offset, (order_by, node))| { + let projected_column = original_target_len + helper_offset; + let alias = format!("__pgdog_order_by_{order_by}"); + let target = mem.make_res_target(Some(&alias), mem.empty(), mem.make_unique(node)); + + ( + target, + OrderByHelper { + sort_position: order_by, + projected_column, + }, + ) + }) + .collect::>(); + + if helpers.is_empty() { + return; + } + + let (targets, mappings): (Vec<_>, Vec<_>) = helpers.into_iter().unzip(); + select + .target_list_mut() + .extend(mem, mem.make_list(&targets)); + + for helper in mappings { + plan.projection.add_order_by_helper(helper); + } + self.rewritten = true; + } +} + +fn column_is_projected(select: &nodes::SelectStmt, order_by: Column<'_>) -> bool { + select.target_list().iter().any(|target| { + if target.name() == Some(order_by.name) { + return true; + } + + let Node::ColumnRef(projected_ref) = target.val() else { + return false; + }; + let Ok(projected) = Column::try_from(projected_ref) else { + return column_ref_is_matching_star(target.val(), order_by); + }; + + projected.name == order_by.name + && (order_by.table.is_none() + || (projected.table == order_by.table && projected.schema == order_by.schema)) + }) +} + +fn column_ref_is_matching_star(node: Node<'_>, order_by: Column<'_>) -> bool { + let Node::ColumnRef(column_ref) = node else { + return false; + }; + let fields = column_ref.fields(); + if !matches!(fields.iter().next_back(), Some(Node::A_Star(_))) { + return false; + } + + let qualifiers = fields + .iter() + .take(fields.len().saturating_sub(1)) + .filter_map(Node::as_str) + .collect::>(); + + match qualifiers.as_slice() { + [] => true, + [table] => order_by.table == Some(*table), + [schema, table] => order_by.schema == Some(*schema) && order_by.table == Some(*table), + _ => false, + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::backend::{ShardingSchema, schema::Schema}; + use crate::frontend::PreparedStatements; + use crate::frontend::router::parser::rewrite::statement::StatementRewriteContext; + + fn rewrite(sql: &str) -> (String, RewritePlan) { + let ast = pg_raw_parse::parse(sql).unwrap(); + let schema = ShardingSchema { + shards: 2, + ..Default::default() + }; + let db_schema = Schema::default(); + let mut prepared = PreparedStatements::default(); + let mut rewriter = StatementRewrite::new(StatementRewriteContext { + extended: false, + prepared: false, + prepared_statements: &mut prepared, + schema: &schema, + db_schema: &db_schema, + user: "postgres", + search_path: None, + }); + let mut plan = RewritePlan::default(); + let ast = make::owned(|mem| { + let mut ast = mem.make_unique(&*ast.into_inner()); + plan = rewriter + .maybe_rewrite(ast.as_mut().into_iter().next().unwrap(), mem) + .unwrap(); + ast + }); + + (pg_raw_parse::deparse_stmts(&*ast).unwrap(), plan) + } + + #[test] + fn adds_non_projected_order_by_column() { + let (sql, plan) = rewrite("SELECT id FROM users ORDER BY name"); + + assert_eq!( + sql, + "SELECT id, name AS __pgdog_order_by_0 FROM users ORDER BY name" + ); + assert_eq!(plan.projection.drop_columns().collect::>(), [1]); + assert_eq!(plan.projection.order_by_helpers().len(), 1); + assert_eq!(plan.projection.order_by_helpers()[0].sort_position, 0); + assert_eq!(plan.projection.order_by_helpers()[0].projected_column, 1); + } + + #[test] + fn leaves_projected_order_by_columns_unchanged() { + for sql in [ + "SELECT id, name FROM users ORDER BY name", + "SELECT id AS name FROM users ORDER BY name", + "SELECT * FROM users ORDER BY name", + "SELECT users.* FROM users ORDER BY users.name", + ] { + let (_, plan) = rewrite(sql); + assert!(plan.projection.order_by_helpers().is_empty(), "{sql}"); + } + } + + #[test] + fn tracks_multiple_helpers_by_order_position() { + let (_, plan) = rewrite("SELECT id FROM users ORDER BY id, name DESC, email"); + let helpers = plan.projection.order_by_helpers(); + + assert_eq!(helpers.len(), 2); + assert_eq!(helpers[0].sort_position, 1); + assert_eq!(helpers[0].projected_column, 1); + assert_eq!(helpers[1].sort_position, 2); + assert_eq!(helpers[1].projected_column, 2); + } + + #[test] + fn plain_distinct_keeps_postgres_validation() { + let (_, plan) = rewrite("SELECT DISTINCT id FROM users ORDER BY name"); + assert!(plan.projection.order_by_helpers().is_empty()); + } +} diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs b/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs index 94bac8bce..afbe33808 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 added for cross-shard aggregation and ordering. + 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() } @@ -120,6 +119,10 @@ impl RewritePlan { bind.push_param(param, format); } + for _ in self.projection.drop_columns() { + bind.push_result_format(Format::Text); + } + Ok(()) } @@ -224,6 +227,7 @@ impl RewritePlan { #[cfg(test)] mod tests { use super::*; + use crate::frontend::router::parser::rewrite::statement::projection::OrderByHelper; use crate::test_utils::set_env_var; use std::collections::HashSet; @@ -250,6 +254,27 @@ mod tests { assert_eq!(bind.params_raw().len(), 0); } + #[tokio::test] + async fn test_apply_bind_extends_per_column_result_formats() { + let mut projection = ProjectionRewritePlan::default(); + projection.add_order_by_helper(OrderByHelper { + sort_position: 0, + projected_column: 2, + }); + let plan = RewritePlan { + projection, + ..Default::default() + }; + let mut bind = Bind::new_params_codes_results("test", &[], &[], &[1, 0]); + + plan.apply_bind(&mut bind).await.unwrap(); + + assert_eq!( + bind.result_formats().collect::>(), + [Format::Binary, Format::Text, Format::Text] + ); + } + #[tokio::test] async fn test_apply_bind_text_format() { let _guard = set_env_var("NODE_ID", "pgdog-1"); 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..b9cc575dd --- /dev/null +++ b/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs @@ -0,0 +1,139 @@ +/// 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 AggregateHelper { + pub(crate) target_column: usize, + pub(crate) projected_column: usize, + pub(crate) distinct: bool, + pub(crate) kind: HelperKind, + pub(crate) alias: String, +} + +/// Column temporarily projected so PgDog can globally order shard results. +#[derive(Debug, Clone, PartialEq)] +pub(crate) struct OrderByHelper { + /// Position of the expression in the ORDER BY clause. + pub(crate) sort_position: usize, + /// Position of the temporary expression in the backend result. + pub(crate) projected_column: usize, +} + +/// Plan for temporary columns added to a query's projection. +#[derive(Debug, Clone, Default, PartialEq)] +pub(crate) struct ProjectionRewritePlan { + aggregate_helpers: Vec, + order_by_helpers: Vec, +} + +impl ProjectionRewritePlan { + /// Create a no-op projection rewrite plan. + pub(crate) fn new() -> Self { + Self { + aggregate_helpers: Vec::new(), + order_by_helpers: Vec::new(), + } + } + + /// Whether the projection and its result require no changes. + pub(crate) fn is_noop(&self) -> bool { + self.aggregate_helpers.is_empty() && self.order_by_helpers.is_empty() + } + + /// Temporary result columns to remove before forwarding to the client. + 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 rewrite_plan_noop() { + let plan = ProjectionRewritePlan::new(); + assert!(plan.is_noop()); + assert!(plan.drop_columns().count() == 0); + assert!(plan.aggregate_helpers().is_empty()); + } + + #[test] + fn rewrite_plan_helpers() { + let mut plan = ProjectionRewritePlan::new(); + plan.add_aggregate_helper(AggregateHelper { + target_column: 0, + projected_column: 1, + distinct: false, + kind: HelperKind::Count, + alias: "__pgdog_count_expr7_col0".into(), + }); + 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!(!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/route.rs b/pgdog/src/frontend/router/parser/route.rs index 94bb9ab24..c8e9a8d88 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,9 @@ 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, + /// Rewrites that add temporary result columns for cross-shard aggregation + /// and ordering. + projection_rewrite: ProjectionRewritePlan, /// Our query explain plan. We attach /// this to the `EXPLAIN` output. explain: Option, @@ -402,12 +401,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 } - 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; } pub(super) fn with_temp_table_change(mut self, temp_table: Option) -> Self { diff --git a/pgdog/src/net/messages/bind.rs b/pgdog/src/net/messages/bind.rs index 330fa9aab..3ca322be5 100644 --- a/pgdog/src/net/messages/bind.rs +++ b/pgdog/src/net/messages/bind.rs @@ -1,5 +1,6 @@ //! Bind (F) message. use crate::net::c_string_buf_len; +use bytes::BytesMut; use super::Error; use super::FromDataType; @@ -214,6 +215,23 @@ impl Bind { }) } + /// Add a result format for a column appended by query rewriting. + pub(crate) fn push_result_format(&mut self, format: Format) { + // Zero formats means all text and one format applies to every result + // column, so only an explicit per-column list needs extending. + if self.results.len() <= 2 { + return; + } + + let mut results = BytesMut::from(&self.results[..]); + results.put_i16(match format { + Format::Text => 0, + Format::Binary => 1, + }); + self.results = results.freeze(); + self.original = None; + } + pub(crate) fn new_statement(name: &str) -> Self { Self { statement: c_string_bytes(name),