diff --git a/pgdog/src/backend/pool/connection/aggregate.rs b/pgdog/src/backend/pool/connection/aggregate.rs index 41570f7be..7efddb306 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() @@ -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() @@ -609,7 +609,7 @@ mod test { &rows, &decoder, &aggregate, - &AggregateRewritePlan::default(), + &ProjectionRewritePlan::default(), ) .unwrap() .aggregate() @@ -654,7 +654,7 @@ mod test { &rows, &decoder, &aggregate, - &AggregateRewritePlan::default(), + &ProjectionRewritePlan::default(), ) .unwrap() .aggregate() @@ -701,7 +701,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..f523753f9 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_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 rows.iter_mut() { + for row in self.buffer.iter_mut() { row.drop_columns(&drop); } } @@ -285,7 +284,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 +311,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 a12f4609f..11b5a456f 100644 --- a/pgdog/src/backend/pool/connection/multi_shard/mod.rs +++ b/pgdog/src/backend/pool/connection/multi_shard/mod.rs @@ -259,15 +259,19 @@ 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(), &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 + .drop_helper_columns(&self.route.projection_rewrite_plan, &self.decoder); self.buffer.distinct(self.route.distinct(), &self.decoder); self.buffer.limit(self.route.limit()); } @@ -310,11 +314,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.aggregate_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 8cd70f231..52b5df9d2 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, OrderBySource, ProjectionRewritePlan}, + }, net::{BindComplete, DataRow, Field, Format}, }; @@ -60,6 +63,80 @@ fn test_inconsistent_data_rows() { } } +#[test] +fn test_order_by_helper_after_star_expansion_is_dropped_after_sorting() { + let mut plan = ProjectionRewritePlan::default(); + plan.order_by_helpers.push(OrderByHelper { + 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), + vec![OrderBy::AscColumn("__pgdog_order_col0".into())], + Default::default(), + Default::default(), + None, + ); + route.projection_rewrite_plan = plan; + let mut multi_shard = MultiShard::new(vec![0, 1], &route); + + 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()) + .unwrap() + .is_none() + ); + let client_description = multi_shard + .handle_server_message(row_description.message()) + .unwrap() + .unwrap(); + let client_description = RowDescription::from_bytes(client_description.to_bytes()).unwrap(); + assert_eq!( + 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("first") + .add("2026-01-01 00:00:00") + .add(20_i64); + let mut second = DataRow::new(); + 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(); + + 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(), 3); + 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/backend/prepared_statements.rs b/pgdog/src/backend/prepared_statements.rs index e13e45393..b10d7ca9e 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, @@ -59,6 +59,7 @@ pub(super) struct Prepare { /// Some if statement was prepared previously, but has expired since close: Option, parse: ProtocolMessage, + describe: Option, } impl Prepare { @@ -71,8 +72,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(); + } } } @@ -87,6 +95,10 @@ pub(super) enum HandleResult { rewrite: ProtocolMessage, }, PrependProtocolMessage(ProtocolMessage), + PrependProtocolMessageRewrite { + prepend: ProtocolMessage, + rewrite: ProtocolMessage, + }, } /// Server-specific prepared statements. @@ -186,6 +198,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(); @@ -201,12 +214,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 { @@ -534,9 +560,30 @@ impl PreparedStatements { // it still holds. close: expired.then(|| ProtocolMessage::Close(Close::named(name))), parse: ProtocolMessage::Parse(parse), + describe: None, })) } + /// 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 + .global_cache + .read() + .cross_shard_variant_needs_row_description(name) + { + return None; + } + + self.describes.push_back(name.to_owned()); + self.parameter_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) @@ -713,8 +760,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; @@ -798,6 +846,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"); @@ -981,6 +1098,31 @@ pub(crate) mod test { assert_eq!(describe_parameters(&mut ps, &name, vec![23, 25]), vec![23]); } + #[test] + fn parameter_description_hides_rewrite_params_for_cross_shard_variant() { + let mut ps = new_extended(); + let base = insert_global("param_desc_variant", "SELECT $1 AS param_desc_variant"); + let variant = { + let global = FrontendPreparedStatements::global(); + let mut cache = global.write(); + cache.rewrite( + &Parse::named(&base, "SELECT $1 AS param_desc_variant, $2::text"), + 1, + ); + cache + .cross_shard_variant( + &base, + "SELECT $1 AS param_desc_variant, $2::text, 1 AS __pgdog_order_col0", + ) + .unwrap() + }; + + assert_eq!( + describe_parameters(&mut ps, &variant, vec![23, 25]), + vec![23] + ); + } + #[test] fn parameter_description_untouched_without_rewrite() { let mut ps = new_extended(); diff --git a/pgdog/src/backend/server.rs b/pgdog/src/backend/server.rs index 08fa087ac..6e654d792 100644 --- a/pgdog/src/backend/server.rs +++ b/pgdog/src/backend/server.rs @@ -508,6 +508,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?; @@ -523,7 +527,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 a8f563a6f..5b4b519df 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,14 @@ 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), + )?; + 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 399e650d8..7c3aca9c7 100644 --- a/pgdog/src/frontend/client/query_engine/test/mod.rs +++ b/pgdog/src/frontend/client/query_engine/test/mod.rs @@ -33,6 +33,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..65617cde5 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,17 @@ async fn test_offset_with_unique_id_simple() { "should have bigint cast: {rewritten_sql}" ); - // apply_after_parser 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,16 +165,16 @@ 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"), - "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 15"), - "LIMIT should be 10+5=15: {final_sql}" + final_sql.contains("LIMIT 10::bigint + 5::bigint"), + "LIMIT should request limit+offset rows: {final_sql}" ); assert!( !final_sql.contains("OFFSET"), @@ -210,30 +216,108 @@ 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. 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). 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::bigint + $3::bigint", + "SQL must push down limit+offset" ); - // Bind params: $1=hello unchanged, $2=limit rewritten to 15, $3=offset rewritten to 0. 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"); } } + +#[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 new file mode 100644 index 000000000..0dcdec5cb --- /dev/null +++ b/pgdog/src/frontend/client/query_engine/test/rewrite_projection.rs @@ -0,0 +1,582 @@ +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}; +use crate::frontend::{ + PreparedStatements, + router::parser::{Limit, OrderBy}, +}; +use pgdog_vector::Vector; + +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::AscColumn("__pgdog_order_col0".into())] + ); + assert_eq!( + context + .client_request + .route() + .projection_rewrite_plan + .order_by_helpers + .len(), + 1 + ); +} + +#[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::AscColumn("__pgdog_order_col1".into())] + ); + + 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::AscColumn("__pgdog_order_col0".into()), + OrderBy::AscColumn("__pgdog_order_col1".into()) + ] + ); +} + +#[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"; + 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::AscColumn("__pgdog_order_col1".into()) + ] + ); +} + +#[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::bigint + 5::bigint")); + 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.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(), + &Limit { + limit: Some(10), + offset: Some(5), + } + ); +} + +#[tokio::test] +async fn split_anonymous_prepare_rewrites_each_execution_once() { + 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()); + + { + 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 + .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_eq!( + context + .client_request + .last_parse + .as_ref() + .unwrap() + .query() + .matches("__pgdog_count_col0") + .count(), + 1 + ); + 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 9c523f6c2..f60573c44 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: @@ -30,6 +34,7 @@ use super::*; pub(crate) struct GlobalCache { statements: HashMap, names: HashMap, + cross_shard_variants: HashMap, unused: HashSet, counter: Counter, } @@ -38,12 +43,20 @@ 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() } } impl GlobalCache { + 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) + } + /// Record a Parse message with the global cache and return a globally unique /// name PgDog is using for that statement. /// @@ -140,16 +153,54 @@ 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); + } + let client_params = self.client_params(name); + let variant_name = cross_shard_variant_name(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, + client_params, + }, + row_description: None, + cache_key, + }, + ); + + Some(variant_name) + } + /// Number of parameters the client's original statement has /// (if we re-write, we must catch and not send back the extra cols ParameterDescriptions) pub(crate) fn client_params(&self, name: &str) -> Option { - self.names.get(name).and_then(|stmt| stmt.client_params()) + self.cross_shard_variants + .get(name) + .or_else(|| self.names.get(name)) + .and_then(|stmt| stmt.client_params()) } /// 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); @@ -180,8 +231,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())) } @@ -199,7 +251,16 @@ 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()) + } + + 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. @@ -271,6 +332,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(&cross_shard_variant_name(name)); } } @@ -307,6 +370,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 @@ -348,6 +412,81 @@ 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.existing_cross_shard_variant_name(&base).as_deref(), + Some(variant.as_str()) + ); + 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 cross_shard_variant_preserves_client_parameter_count() { + let mut cache = GlobalCache::default(); + let (_, base) = cache.insert(&Parse::named("client", "SELECT $1")); + cache.rewrite(&Parse::named(&base, "SELECT $1, $2::bigint"), 1); + + let variant = cache + .cross_shard_variant(&base, "SELECT $1, $2::bigint") + .unwrap(); + + assert_eq!(cache.client_params(&base), Some(1)); + assert_eq!(cache.client_params(&variant), Some(1)); + } + #[test] fn test_prep_stmt_cache_close() { let mut cache = GlobalCache::default(); diff --git a/pgdog/src/frontend/prepared_statements/mod.rs b/pgdog/src/frontend/prepared_statements/mod.rs index 2862476db..5749054e2 100644 --- a/pgdog/src/frontend/prepared_statements/mod.rs +++ b/pgdog/src/frontend/prepared_statements/mod.rs @@ -70,6 +70,17 @@ impl PreparedStatements { Self::new().global.clone() } + 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); + } + + 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/cache/ast.rs b/pgdog/src/frontend/router/parser/cache/ast.rs index afc969555..2d45a7efd 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::config::Role; 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; /// Abstract syntax tree (query) cache entry, @@ -40,6 +42,9 @@ pub(crate) struct AstInner { pub(crate) stats: Mutex, /// Rewrite plan. pub(crate) rewrite_plan: RewritePlan, + /// 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, } @@ -51,6 +56,7 @@ impl AstInner { ast, stats: Mutex::new(Stats::new()), rewrite_plan: RewritePlan::default(), + post_route_rewrite: OnceCell::new(), query_without_comment: "".into(), } } @@ -121,6 +127,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/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 bb847ff16..ce97cb62b 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 057610804..7614f3ecf 100644 --- a/pgdog/src/frontend/router/parser/query/select.rs +++ b/pgdog/src/frontend/router/parser/query/select.rs @@ -1,7 +1,5 @@ -use crate::frontend::router::parser::cache::Ast; -use crate::frontend::router::parser::statement::AdvisoryLockId; - use super::*; +use crate::frontend::router::parser::statement::AdvisoryLockId; use pg_raw_parse::walk; use pg_raw_parse::{Node, nodes}; use pgdog_config::system_catalogs; @@ -17,7 +15,6 @@ impl QueryParser { /// pub(super) fn select( &mut self, - cached_ast: &Ast, stmt: &nodes::SelectStmt, context: &mut QueryParserContext, ) -> Result { @@ -315,7 +312,7 @@ impl QueryParser { } } - let mut query = Route::select( + let query = Route::select( context.shards_calculator.shard().clone(), order_by, aggregates, @@ -323,11 +320,6 @@ 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()); - } - Ok(Command::Query( query .with_read(!writes) 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..c4205f1ed 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,8 @@ 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; /// Query rewrite engine. Currently supports injecting helper aggregates for AVG and /// variance-related functions that require additional helper aggregates when run @@ -17,8 +18,8 @@ impl AggregatesRewrite { select: &mut nodes::SelectStmtMut<'a, '_>, mem: make::MemoryToken<'a>, aggregate: &Aggregate, - ) -> RewriteOutput { - let mut plan = AggregateRewritePlan::new(); + ) -> ProjectionRewritePlan { + let mut plan = ProjectionRewritePlan::default(); let helper_nodes = aggregate .targets() @@ -37,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_helper(HelperMapping { + plan.aggregate_helpers.push(AggregateHelper { target_column: target.column(), - helper_column: select.target_list().len() + idx, distinct: target.is_distinct(), kind, alias: helper_alias, @@ -54,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>( @@ -163,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"); @@ -174,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.helpers().len(), 1); - let helper = &output.plan.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.helper_column, 1); + assert_eq!(helper.alias, "__pgdog_count_col0"); assert!(!helper.distinct); assert!(matches!(helper.kind, HelperKind::Count)); @@ -215,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.helpers().len(), 1); - let helper = &output.plan.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.helper_column, 2); + assert_eq!(helper.alias, "__pgdog_count_col1"); assert!(!helper.distinct); assert!(matches!(helper.kind, HelperKind::Count)); @@ -238,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.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.helpers()[0]; + let helper_price = &plan.aggregate_helpers[0]; assert_eq!(helper_price.target_column, 0); - assert_eq!(helper_price.helper_column, 2); + assert_eq!(helper_price.alias, "__pgdog_count_col0"); assert!(matches!(helper_price.kind, HelperKind::Count)); - let helper_discount = &output.plan.helpers()[1]; + let helper_discount = &plan.aggregate_helpers[1]; assert_eq!(helper_discount.target_column, 1); - assert_eq!(helper_discount.helper_column, 3); + assert_eq!(helper_discount.alias, "__pgdog_count_col1"); assert!(matches!(helper_discount.kind, HelperKind::Count)); let aggregate = Aggregate::parse(&ast, &Default::default()); @@ -266,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.helpers().len(), 3); - - let kinds: Vec = output - .plan - .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 0ae8e04a5..5709b4589 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/aggregate/mod.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/aggregate/mod.rs @@ -1,39 +1,27 @@ 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; pub(crate) use engine::AggregatesRewrite; -pub(crate) use plan::{AggregateRewritePlan, HelperKind, HelperMapping, RewriteOutput}; -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(()); - } +/// 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, +} - let output = AggregatesRewrite::rewrite_select(select, mem, &aggregate); - if output.plan.is_noop() { - return Ok(()); +impl HelperKind { + /// Suffix used in the projected helper's internal alias. + pub(crate) fn alias_suffix(self) -> &'static str { + match self { + Self::Count => "count", + Self::Sum => "sum", + Self::SumSquares => "sumsq", } - - plan.aggregates = 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 8e9fc3ee5..f2b05e468 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs @@ -15,7 +15,9 @@ pub(crate) mod insert; pub(crate) mod nextval; pub(crate) mod non_deterministic_funcs; pub(crate) mod offset; +pub(crate) 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; @@ -187,8 +189,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 166885593..8d35bea45 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,38 @@ 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; - }; +/// `$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() +} - 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 - })) +/// 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>, +) { + 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()]), + limit, + offset, + ); + select.set_limit_count(combined.uncast()); + select.set_limit_offset(mem.none()); } impl StatementRewrite<'_> { @@ -267,7 +242,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 +286,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(); @@ -398,7 +368,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), @@ -412,15 +382,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)); @@ -428,7 +396,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, @@ -444,11 +412,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"); } @@ -459,7 +427,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), @@ -474,7 +442,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(), @@ -484,7 +452,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), @@ -499,12 +467,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"); } @@ -513,7 +479,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/order_by.rs b/pgdog/src/frontend/router/parser/rewrite/statement/order_by.rs new file mode 100644 index 000000000..25e3f687f --- /dev/null +++ b/pgdog/src/frontend/router/parser/rewrite/statement/order_by.rs @@ -0,0 +1,396 @@ +use pg_raw_parse::{ConstValue, Node, make, nodes}; + +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() + .into_iter() + .map(Node::as_str) + .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(); + 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 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(_))) +} + +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>, + plan: &mut ProjectionRewritePlan, +) { + let mut helpers = Vec::new(); + let mut sort_position = 0; + for sort in select.sort_clause() { + let node = sort.node(); + let source = match node { + 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("<->") => + { + [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 = source.is_some() + || matches!(node, Node::A_Const(constant) + if matches!(constant.val(), Some(ConstValue::Integer(_)))); + if !supported { + continue; + } + + let current_sort_position = sort_position; + sort_position += 1; + + 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, + ); + } + _ => {} + } + } + + 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) -> (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, &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"); + + 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] + fn skips_already_projected_sort_column() { + let (sql, plan) = rewrite("SELECT id, price FROM products ORDER BY price"); + + assert!(!sql.contains("__pgdog_order_col")); + 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"); + + 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"); + + 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"); + + assert!(sql.contains("b.score AS __pgdog_order_col0")); + 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] + 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 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("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/plan.rs b/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs index 116b2d3eb..b41d0d271 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs @@ -2,9 +2,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, aggregate::AggregateRewritePlan, -}; +use super::{Error, InsertSplit, PrepareExecute, ShardingKeyUpdate}; use crate::frontend::client::QueryTimestamps; use crate::frontend::router::parser::rewrite::statement::non_deterministic_funcs::NDFunction; use crate::frontend::{ClientRequest, PreparedStatements}; @@ -63,10 +61,6 @@ 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, - /// Sharding key is being updated, we need to execute /// a multi-step plan. pub(crate) sharding_key_update: Option, @@ -83,11 +77,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(()), } } @@ -104,10 +105,8 @@ impl RewritePlan { && self.stmt.is_none() && self.prepare_rewrites.is_empty() && self.insert_split.is_empty() - && self.aggregates.is_noop() && self.sharding_key_update.is_none() && self.offset.is_none() - // TODO: Check here. } /// Append generated unique IDs and sequence values to a Bind message. 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..772246586 --- /dev/null +++ b/pgdog/src/frontend/router/parser/rewrite/statement/projection.rs @@ -0,0 +1,256 @@ +//! 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}; +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, 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) distinct: bool, + pub(crate) kind: HelperKind, + pub(crate) alias: String, +} + +#[derive(Debug, Clone, PartialEq)] +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)] +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 { + pub(crate) aggregate_helpers: Vec, + pub(crate) 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, row_description: &RowDescription) -> BTreeSet { + self.aggregate_helpers + .iter() + .map(|helper| helper.alias.as_str()) + .chain( + self.order_by_helpers + .iter() + .filter(|helper| helper.injected) + .map(|helper| helper.alias.as_str()), + ) + .filter_map(|alias| row_description.field_index(alias)) + .collect() + } +} + +#[derive(Debug, Clone)] +pub(crate) struct PostRouteRewrite { + sql: Arc, + plan: ProjectionRewritePlan, +} + +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 Some(rewrite) = ast + .post_route_rewrite + .get_or_try_init(|| build(&ast.ast, schema, 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::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); + } + } + _ => {} + } + } + // 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() + { + parse.set_query(&rewrite.sql); + } + if !rewrite.plan.is_noop() + && let Some(route) = request.route.as_mut() + { + route.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)) + .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() { + OrderBy::AscColumn(helper.alias.clone()) + } else { + OrderBy::DescColumn(helper.alias.clone()) + }; + } + route.set_order_by(order_by); + } + + Ok(()) +} + +fn build( + ast: &StmtList, + schema: &Schema, + 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() && select.sort_clause().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); + } + order_by::rewrite_select(&mut select.as_mut(), mem, &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 })) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn projection_plan_tracks_helpers() { + let mut plan = ProjectionRewritePlan::default(); + plan.aggregate_helpers.push(AggregateHelper { + target_column: 0, + distinct: false, + kind: HelperKind::Count, + alias: "__pgdog_count_col0".into(), + }); + plan.order_by_helpers.push(OrderByHelper { + sort_position: 0, + source: OrderBySource::Column("created_at".into()), + alias: "__pgdog_order_col0".into(), + injected: true, + }); + + assert!(!plan.is_noop()); + 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); + } + + #[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()); + } +} diff --git a/pgdog/src/frontend/router/parser/route.rs b/pgdog/src/frontend/router/parser/route.rs index 2961dfa5a..28ed1006a 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. + pub(crate) projection_rewrite_plan: ProjectionRewritePlan, /// Our query explain plan. We attach /// this to the `EXPLAIN` output. explain: Option, @@ -255,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 } @@ -400,14 +402,6 @@ impl Route { self.is_cross_shard() && self.is_write() } - pub(crate) fn aggregate_rewrite_plan(&self) -> &AggregateRewritePlan { - &self.rewrite_plan - } - - pub(crate) fn set_rewrite_plan(&mut self, plan: AggregateRewritePlan) { - self.rewrite_plan = plan; - } - pub(super) fn with_temp_table_change(mut self, temp_table: Option) -> Self { self.temp_table_change = temp_table; self diff --git a/pgdog/src/net/messages/bind.rs b/pgdog/src/net/messages/bind.rs index 34f9803a7..df8555aea 100644 --- a/pgdog/src/net/messages/bind.rs +++ b/pgdog/src/net/messages/bind.rs @@ -301,18 +301,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 { @@ -553,20 +541,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;