diff --git a/pgdog/src/frontend/client/query_engine/context.rs b/pgdog/src/frontend/client/query_engine/context.rs index 8dd4c8113..9c083e8e8 100644 --- a/pgdog/src/frontend/client/query_engine/context.rs +++ b/pgdog/src/frontend/client/query_engine/context.rs @@ -18,7 +18,7 @@ pub(crate) struct QueryEngineContext<'a> { /// Client session parameters. pub(super) params: &'a mut Parameters, /// Request. - pub(super) client_request: &'a mut ClientRequest, + pub(crate) client_request: &'a mut ClientRequest, /// How many requests are left to execute in an extended pipeline. pub(super) pipeline: Pipeline, /// Client's socket to send responses to. @@ -106,4 +106,15 @@ impl<'a> QueryEngineContext<'a> { pub(crate) fn in_error(&self) -> bool { self.transaction.map(|t| t.error()).unwrap_or_default() } + + /// Mark the transaction failed + pub(crate) fn set_transaction_error(&mut self) { + self.transaction = match self.transaction { + Some(TransactionType::ReadOnly) => Some(TransactionType::ErrorReadOnly), + Some(TransactionType::ReadWrite | TransactionType::Implicit) => { + Some(TransactionType::ErrorReadWrite) + } + _ => None, + }; + } } diff --git a/pgdog/src/frontend/client/query_engine/mod.rs b/pgdog/src/frontend/client/query_engine/mod.rs index 498ec830a..6011eee84 100644 --- a/pgdog/src/frontend/client/query_engine/mod.rs +++ b/pgdog/src/frontend/client/query_engine/mod.rs @@ -146,8 +146,8 @@ impl QueryEngine { } // Rewrite statement if necessary. - let rewrite_result = match self.parse_and_rewrite(context) { - Ok(rewrite_result) => rewrite_result, + let (query_planner, offset_plan) = match self.parse_and_rewrite(context).await { + Ok(result) => result, Err(e) => { self.error_response(context, ErrorResponse::syntax(e.to_string())) .await?; @@ -162,7 +162,7 @@ impl QueryEngine { } // Route transaction to the right servers. - if !self.route_query(context, rewrite_result.as_ref()).await? { + if !self.route_query(context, offset_plan.as_ref()).await? { self.update_stats(context); debug!("query has nowhere to go"); return Ok(QueryEngineResult::Done(context.transaction())); @@ -240,11 +240,11 @@ impl QueryEngine { context.params.rollback(); } - Command::Query(_) => self.execute(context, rewrite_result).await?, + Command::Query(_) => self.execute(context, query_planner).await?, Command::Listen { .. } | Command::Notify { .. } | Command::Unlisten(_) if self.backend.session_mode() => { - self.execute(context, rewrite_result).await? + self.execute(context, query_planner).await? } Command::Listen { channel, shard } => { self.listen(context, &channel.clone(), shard.clone()) @@ -268,7 +268,7 @@ impl QueryEngine { Command::ResetAll => { self.reset_all(context).await?; } - Command::Copy(_) => self.execute(context, rewrite_result).await?, + Command::Copy(_) => self.execute(context, query_planner).await?, Command::Deallocate => self.deallocate(context).await?, Command::Discard { extended } => self.discard(context, *extended).await?, Command::Split(queries) => return Ok(Self::build_simple_split(queries)), diff --git a/pgdog/src/frontend/client/query_engine/multi_step/error.rs b/pgdog/src/frontend/client/query_engine/multi_step/error.rs index 3c5819f0b..d398a2374 100644 --- a/pgdog/src/frontend/client/query_engine/multi_step/error.rs +++ b/pgdog/src/frontend/client/query_engine/multi_step/error.rs @@ -7,9 +7,18 @@ pub(crate) enum Error { #[error("{0}")] Update(#[from] UpdateError), + #[error("{0}")] + Insert(#[from] InsertError), + #[error("frontend: {0}")] Frontend(Box), + #[error("parser: {0}")] + Parser(#[from] crate::frontend::router::parser::Error), + + #[error("deparse: {0}")] + Deparse(#[from] pg_raw_parse::Error), + #[error("backend: {0}")] Backend(#[from] crate::backend::Error), @@ -26,11 +35,30 @@ pub(crate) enum Error { Net(#[from] crate::net::Error), } +impl Error { + /// Errors the client should see as an `ErrorResponse`. + /// Otherwise, it's an internal failure and propagates up which closes the connection. + pub(crate) fn into_client_error(self) -> Result { + match self { + Self::Execution(error) => Ok(*error), + err @ (Self::Update(_) | Self::Insert(_) | Self::Rewrite(_)) => { + Ok(ErrorResponse::from_err(&err)) + } + err => Err(err), + } + } +} + #[derive(Debug, Error)] pub(crate) enum UpdateError { #[error("sharding key updates are forbidden")] Disabled, + /// Parser flagged a sharding key update but the planner can't continue. + /// If we let it continue, this could cause unintended side effects. + #[error("sharding key update plan doesn't match the parsed statement")] + PlanMismatch, + #[error("sharding key update must be executed inside a transaction")] TransactionRequired, @@ -42,6 +70,26 @@ pub(crate) enum UpdateError { #[error("sharding key update would move a row referenced by an ON DELETE foreign key")] ForeignKeyOnDelete, + + #[error("sharding key update expected an UPDATE statement")] + NotAnUpdate, + + #[error("sharding key update step \"{0}\" response is missing or incomplete")] + MissingStepResponse(&'static str), +} + +#[derive(Debug, Error)] +pub(crate) enum InsertError { + #[error("multi-tuple insert requires multi-shard binding")] + MultiShardRequired, + + /// Parser flagged a multi insert but the planner can't continue. + /// If we let it continue, this could cause unintended side effects. + #[error("multi-tuple insert plan doesn't match the parsed statement")] + PlanMismatch, + + #[error("cache: {0}")] + Cache(String), } impl From for Error { diff --git a/pgdog/src/frontend/client/query_engine/multi_step/insert.rs b/pgdog/src/frontend/client/query_engine/multi_step/insert.rs deleted file mode 100644 index a36abe2c9..000000000 --- a/pgdog/src/frontend/client/query_engine/multi_step/insert.rs +++ /dev/null @@ -1,135 +0,0 @@ -use super::{CommandType, MultiServerState}; -use crate::{ - frontend::{ - ClientRequest, Command, Router, RouterContext, - client::query_engine::{QueryEngine, QueryEngineContext}, - router::{ - Route, - parser::route::{Shard, ShardWithPriority}, - }, - }, - net::Protocol, -}; - -use super::super::Error; - -#[derive(Debug)] -pub(crate) struct InsertMulti<'a> { - /// Requests split by the rewrite engine. - requests: Vec, - /// Execution state. - state: MultiServerState, - /// Query engine. - engine: &'a mut QueryEngine, -} - -impl<'a> InsertMulti<'a> { - /// Create multi-shard INSERT handler - /// from query engine and a set of routed requests. - pub(crate) fn from_engine(engine: &'a mut QueryEngine, requests: Vec) -> Self { - Self { - state: MultiServerState::new(requests.len()), - requests, - engine, - } - } - - /// If every split routes to the same `Shard::Direct(n)`, return that shard - /// number. Returns `None` when the splits span multiple shards or contain - /// any non-direct routing. - fn uniform_shard(&self) -> Option { - let first = match self.requests.first()?.route.as_ref()?.shard() { - Shard::Direct(n) => *n, - _ => return None, - }; - - self.requests - .iter() - .skip(1) - .all(|req| { - matches!( - req.route.as_ref().map(|r| r.shard()), - Some(Shard::Direct(n)) if *n == first - ) - }) - .then_some(first) - } - - /// Execute the multi-shard INSERT. - pub(crate) async fn execute( - &'a mut self, - context: &mut QueryEngineContext<'_>, - ) -> Result { - let cluster = self.engine.backend.cluster()?; - for request in self.requests.iter_mut() { - let context = RouterContext::new( - request, - cluster, - context.params, - context.transaction(), - context.sticky, - )?; - let mut router = Router::new(); - let command = router.query(context)?; - if let Command::Query(route) = command { - request.route = Some(route.clone()); - } else { - return Err(Error::NoRoute); - } - } - - // All tuples map to the same shard: send the original multi-row INSERT - // as a single statement, skipping the multi-step path entirely. - if let Some(shard_n) = self.uniform_shard() { - context.client_request.route = Some(Route::write(ShardWithPriority::new_table( - Shard::Direct(shard_n), - ))); - self.engine - .backend - .handle_client_request( - context.client_request, - &mut self.engine.router, - self.engine.streaming, - ) - .await?; - while self.engine.backend.has_more_messages() { - let message = self.engine.read_server_message().await?; - self.engine.process_server_message(context, message).await?; - } - return Ok(false); - } - - if !self.engine.backend.is_multishard() { - return Err(Error::MultiShardRequired); - } - - for request in self.requests.iter() { - self.engine - .backend - .handle_client_request(request, &mut self.engine.router, self.engine.streaming) - .await?; - - while self.engine.backend.has_more_messages() { - let message = self.engine.read_server_message().await?; - - if self.state.forward(&message)? { - self.engine.process_server_message(context, message).await?; - } - } - } - - if let Some(cc) = self.state.command_complete(CommandType::Insert) { - self.engine - .process_server_message(context, cc.message()) - .await?; - } - - if let Some(rfq) = self.state.ready_for_query(context.in_transaction()) { - self.engine - .process_server_message(context, rfq.message()) - .await?; - } - - Ok(self.state.error()) - } -} diff --git a/pgdog/src/frontend/client/query_engine/multi_step/mod.rs b/pgdog/src/frontend/client/query_engine/multi_step/mod.rs index 3bf237e74..e37d7e954 100644 --- a/pgdog/src/frontend/client/query_engine/multi_step/mod.rs +++ b/pgdog/src/frontend/client/query_engine/multi_step/mod.rs @@ -1,14 +1,10 @@ -pub(crate) mod error; +pub mod error; pub(crate) mod forward_check; -pub(crate) mod insert; -pub(crate) mod state; -pub(crate) mod update; +pub(crate) mod shared; -pub(crate) use error::{Error, UpdateError}; -pub(crate) use forward_check::*; -pub(crate) use insert::InsertMulti; -pub(crate) use state::{CommandType, MultiServerState}; -pub(crate) use update::UpdateMulti; +pub(crate) mod types; + +pub(crate) mod ops; #[cfg(test)] mod test; diff --git a/pgdog/src/frontend/client/query_engine/multi_step/ops/insert.rs b/pgdog/src/frontend/client/query_engine/multi_step/ops/insert.rs new file mode 100644 index 000000000..818e08347 --- /dev/null +++ b/pgdog/src/frontend/client/query_engine/multi_step/ops/insert.rs @@ -0,0 +1,243 @@ +use crate::frontend::BufferedQuery; +use crate::frontend::client::query_engine::multi_step::error::{Error, InsertError}; +use crate::frontend::client::query_engine::multi_step::types::{ + ForwardToClient, QueryPlanner, ResponseHistory, Step, StepRequest, +}; +use crate::frontend::router::parser::rewrite::statement::Error as RewriteError; +use crate::frontend::router::parser::{AstContext, Cache}; +use crate::net::{CommandComplete, Message, Parse, Query, ReadyForQuery}; +use crate::{ + frontend::{ + ClientRequest, + client::query_engine::{QueryEngine, QueryEngineContext}, + router::{ + Route, + parser::route::{Shard, ShardWithPriority}, + }, + }, + net::Protocol, +}; +use indexmap::IndexSet; +use itertools::Itertools; +use pg_raw_parse::{Node, NodeMut, deparse, make, nodes, walk}; + +impl QueryPlanner { + /// Try to create a `QueryPlanner` where we suspect a multi-INSERT. + pub(crate) fn plan_multi_insert( + engine: &QueryEngine, + context: &mut QueryEngineContext, + request: &ClientRequest, + ) -> Result, Error> { + let Some(ref ast) = request.ast else { + debug_assert!(false, "planner dispatched without an AST"); + return Err(InsertError::PlanMismatch.into()); + }; + + let Ok(Node::InsertStmt(insert)) = ast.ast.stmts().exactly_one() else { + debug_assert!(false, "insert_split flagged on a non-INSERT statement"); + return Err(InsertError::PlanMismatch.into()); + }; + + let steps = create_steps(insert, engine, context, request)?; + if steps.is_empty() { + debug_assert!(false, "insert_split flagged with fewer than two tuples"); + return Err(InsertError::PlanMismatch.into()); + } + if !Self::checks(engine, context, &steps)? { + return Ok(None); + } + + #[derive(Debug, Clone)] + struct MultiInsertAggregate; + + impl ForwardToClient for MultiInsertAggregate { + fn forward_to_client( + &self, + context: &QueryEngineContext, + map: ResponseHistory, + ) -> Vec { + let mut messages = map.compose_all(context.client_request); + + // Sum the per-tuple INSERT tags into one. + let rows: usize = map + .steps() + .iter() + .filter_map(|step| step.command_complete.as_ref()) + .filter_map(|cc| cc.rows().ok().flatten()) + .sum(); + messages.push(CommandComplete::new(format!("INSERT 0 {}", rows)).message()); + messages.push(ReadyForQuery::in_transaction(context.in_transaction()).message()); + + messages + } + } + + Ok(Some(QueryPlanner { + steps, + forward_to_client: Some(Box::new(MultiInsertAggregate {})), + })) + } + + fn checks( + engine: &QueryEngine, + context: &mut QueryEngineContext<'_>, + steps: &[Step], + ) -> Result { + // All tuples map to the same shard: send the original multi-row INSERT + // as a single statement, skipping the multi-step path entirely. + if let Some(shard_n) = Self::uniform_shard(context.client_request, steps) { + context.client_request.route = Some(Route::write(ShardWithPriority::new_table( + Shard::Direct(shard_n), + ))); + + return Ok(false); + } + + // TODO: I think this is an approximation of the old execute-time check as this is planning-based; + // are we connected yet? Is this suitably tested by an integration test? + if engine.backend.connected() && !engine.backend.is_multishard() { + return Err(InsertError::MultiShardRequired.into()); + } + + Ok(true) + } + + /// If every split routes to the same `Shard::Direct(n)`, return that shard + /// number. Returns `None` when the splits span multiple shards or contain + /// any non-direct routing. + fn uniform_shard(original_request: &ClientRequest, steps: &[Step]) -> Option { + let direct_shard = |step: &Step| { + let route = match &step.request { + StepRequest::Raw => original_request.route(), + StepRequest::Statement(statement) => &statement.route, + }; + + match route.shard() { + Shard::Direct(n) => Some(*n), + _ => None, + } + }; + + let first = direct_shard(steps.first()?)?; + + steps + .iter() + .skip(1) + .all(|step| direct_shard(step) == Some(first)) + .then_some(first) + } +} + +/// Split up multi-tuple INSERT statements into separate single-tuple statements +/// for individual execution. +/// +/// # Example +/// +/// ```sql +/// INSERT INTO my_table (id, value) VALUES ($1, $2), ($3, $4) +/// ``` +/// +/// becomes +/// +/// ```sql +/// INSERT INTO my_table (id, value) VALUES ($1, $2) +/// INSERT INTO my_table (id, value) VALUES ($1, $2) -- These are copied from params $3 and $4 +/// ``` +/// +pub(crate) fn create_steps( + insert: &nodes::InsertStmt, + engine: &QueryEngine, + context: &mut QueryEngineContext<'_>, + original_request: &ClientRequest, +) -> Result, Error> { + let mut steps: Vec = Vec::new(); + + let mut splits = Vec::new(); + make::try_owned(|mem| { + let mut copy = mem.make_unique(insert); + + if let Node::SelectStmt(select) = insert.select_stmt() { + for list in select.values_lists() { + let (params, select) = build_single_tuple_select(mem, list); + copy.as_mut().set_select_stmt(select.uncast()); + splits.push((params, deparse(&*copy)?.as_str().to_string())); + } + } + + Ok::<_, RewriteError>(copy) + })?; + + // There's no point of continuing if there's not >= 2 splits. + if splits.len() <= 1 { + return Ok(steps); + } + + // Now create Ast for each split (needs mutable borrow of prepared_statements) + let (extended, prepared) = context + .client_request + .query()? + .map(|query| (query.extended(), query.prepared())) + .unwrap_or_default(); + let cache = Cache::get(); + let ctx = AstContext::from_cluster(engine.backend.cluster()?, context.params); + + for (params, stmt) in splits.iter() { + let query = if extended { + BufferedQuery::Prepared(Parse::named("", stmt)) + } else { + BufferedQuery::Query(Query::new(stmt)) + }; + let ast = cache + .query(&query, &ctx, &mut *context.prepared_statements) + .map_err(|e| InsertError::Cache(e.to_string()))?; + + // If this is a named prepared statement, register the split in the global cache + // and store the assigned name for use in Bind messages. + let statement_name = if prepared { + // Name will be assigned by `insert`. + let mut parse = Parse::named("", stmt); + context.prepared_statements.insert(&mut parse); + Some(parse.name().to_owned()) + } else { + None + }; + + let request = QueryPlanner::build_request( + engine, + context, + original_request, + &ast, + stmt, + params, + statement_name.as_deref(), + )?; + + steps.push(Step { + save_key: None, + request, + }); + } + + Ok(steps) +} + +/// Build a single-tuple INSERT from the original statement with just one values_list. +/// Returns the parameter positions (0-indexed) and the SQL string. +fn build_single_tuple_select<'mem>( + mem: make::MemoryToken<'mem>, + values_list: Node<'_>, +) -> (IndexSet, make::Unique<'mem, &'mem nodes::SelectStmt>) { + let mut tuple = mem.make_unique(values_list); + + let mut params = IndexSet::new(); + walk::walk_mut(tuple.as_mut(), |node| { + if let NodeMut::ParamRef(param) = node { + params.insert(param.number as _); + param.set_number(params.get_index_of(&(param.number as u16)).unwrap() as i32 + 1) + } + }); + + let mut select = mem.make_node::(); + select.as_mut().set_values_lists(mem.make_list(&[tuple])); + (params, select) +} diff --git a/pgdog/src/frontend/client/query_engine/multi_step/ops/mod.rs b/pgdog/src/frontend/client/query_engine/multi_step/ops/mod.rs new file mode 100644 index 000000000..5f09da01f --- /dev/null +++ b/pgdog/src/frontend/client/query_engine/multi_step/ops/mod.rs @@ -0,0 +1,5 @@ +pub(crate) mod update; + +pub(crate) mod insert; + +pub(crate) mod normal; diff --git a/pgdog/src/frontend/client/query_engine/multi_step/ops/normal.rs b/pgdog/src/frontend/client/query_engine/multi_step/ops/normal.rs new file mode 100644 index 000000000..807c0d548 --- /dev/null +++ b/pgdog/src/frontend/client/query_engine/multi_step/ops/normal.rs @@ -0,0 +1,22 @@ +use crate::frontend::client::query_engine::multi_step::types::{QueryPlanner, Step, StepRequest}; + +impl QueryPlanner { + /// Fallback when we don't match to another `QueryPlannerType`; + /// Runs the `ClientRequest` as normal without any special handling. + /// Allows us to push all execution flow through the `QueryPlanner` instead of special cases. + pub(crate) fn plan_normal() -> QueryPlanner { + let solo_step = Self::construct_solo_step(); + QueryPlanner { + steps: vec![solo_step], + // Everything is forwarded by-request; don't need anything in aggregate + forward_to_client: None, + } + } + + fn construct_solo_step() -> Step { + Step { + save_key: None, + request: StepRequest::Raw, + } + } +} diff --git a/pgdog/src/frontend/client/query_engine/multi_step/ops/update.rs b/pgdog/src/frontend/client/query_engine/multi_step/ops/update.rs new file mode 100644 index 000000000..e2a99304d --- /dev/null +++ b/pgdog/src/frontend/client/query_engine/multi_step/ops/update.rs @@ -0,0 +1,502 @@ +use crate::backend::ShardingSchema; +use crate::frontend::ClientRequest; +use crate::frontend::client::query_engine::multi_step::error::{Error, UpdateError}; +use crate::frontend::client::query_engine::multi_step::types::{ + ForwardToClient, QueryPlanner, ResponseHistory, SaveKey, StatementRequest, StatementSource, + Step, StepProtocol, StepRequest, +}; +use crate::frontend::client::query_engine::{QueryEngine, QueryEngineContext}; +use crate::frontend::router::parser::rewrite::statement::Error as RewriteError; +use crate::frontend::router::parser::{Column, Table, Value}; +use crate::frontend::router::sharding::ShardedTable; +use crate::frontend::router::{Ast, Route}; +use crate::net::bind::Parameter; +use crate::net::{ + Bind, CommandComplete, DataRow, Message, Parse, Protocol, ReadyForQuery, RowDescription, +}; +use indexmap::IndexSet; +use itertools::Itertools; +use pg_raw_parse::make::{owned, try_owned}; +use pg_raw_parse::nodes::ResTarget; +use pg_raw_parse::{Node, deparse, nodes}; +use pgdog_config::RewriteMode; +use pgdog_postgres_types::Format; +use std::collections::{HashMap, HashSet}; +use std::fmt::Debug; + +impl QueryPlanner { + /// Try to create a `QueryPlanner` where we suspect a `ShardingKeyUpdate` + pub(crate) async fn plan_sharding_key_update( + request: &ClientRequest, + engine: &mut QueryEngine, + context: &mut QueryEngineContext<'_>, + schema: ShardingSchema, + ) -> Result, Error> { + // Get the AST + let Some(ref ast) = request.ast else { + debug_assert!(false, "planner dispatched without an AST"); + return Err(UpdateError::PlanMismatch.into()); + }; + + // Get the Update Statement node + let Ok(Node::UpdateStmt(original_update_stmt)) = ast.ast.stmts().exactly_one() else { + return Ok(None); + }; + + let table = original_update_stmt + .relation() + .map(Table::from) + .expect("UPDATE always has a table"); + + // Fetch the `ResTarget` for the original update statement, so that we can formulate + // the INSERT `Step` + let Some(new_target) = original_update_stmt.target_list().into_iter().find(|c| { + Column::try_from(*c).is_ok_and(|mut c| { + c.qualify(table); + schema.tables().get_table(c).is_some() + }) + }) else { + return Ok(None); + }; + + let Some(new_route) = + Self::on_same_shard(engine, new_target, original_update_stmt, request, context)? + else { + // CHECK: Are we on the same shard? If so, abort and use the original request. + // This compares two SELECTs (old + new sharding key) through router.query + return Ok(None); + }; + + if Self::has_destructive_on_delete_reference(engine, original_update_stmt, context)? { + return Err(UpdateError::ForeignKeyOnDelete.into()); + } + + // Check if we are allowed to do this operation by the config. + // Propagated as an error so execution stops + if engine.backend.cluster()?.rewrite().shard_key == RewriteMode::Error { + return Err(UpdateError::Disabled.into()); + } + + // Do this check at the last possible moment (in case transactions changed in future) + // TODO: I think this is an approximation of the old execute-time check as this is planning-based; + // are we connected yet? Is this suitably tested by an integration test? + if !context.in_transaction() + || (engine.backend.connected() && !engine.backend.is_multishard()) + { + engine.cleanup_backend(context)?; + return Err(UpdateError::TransactionRequired.into()); + } + + // ****** We've completed all our pre-execution checks. ****** + + let steps: Vec = vec![ + Self::construct_delete_step(engine, request, original_update_stmt, context)?, + Self::construct_insert_step(request, ast, original_update_stmt, new_route)?, + ]; + + // After completion of all Steps, we return this: + #[derive(Debug, Clone)] + struct ShardingKeyUpdateResponse; + impl ForwardToClient for ShardingKeyUpdateResponse { + fn forward_to_client( + &self, + context: &QueryEngineContext, + map: ResponseHistory, + ) -> Vec { + // The INSERT step's responses (RETURNING rows and protocol acks) + // We don't worry about DELETE since it's an internal step. + let mut messages = match map.get(SaveKey::ShardingKeyUpdateInsert) { + Some(insert) => ResponseHistory::compose(&[insert], context.client_request), + None => Vec::new(), + }; + + // Only allows update one row at a time + // We use 0 when the DELETE matched nothing. + let rows = map + .get(SaveKey::ShardingKeyUpdateDelete) + .map(|delete| delete.rows.len()) + .unwrap_or_default(); + + messages.push(CommandComplete::new(format!("UPDATE {}", rows)).message()); + messages.push(ReadyForQuery::in_transaction(context.in_transaction()).message()); + messages + } + } + + Ok(Some(QueryPlanner { + steps, + forward_to_client: Some(Box::new(ShardingKeyUpdateResponse {})), + })) + } + + /// Dependent on `construct_delete_step` + fn construct_insert_step( + request: &ClientRequest, + ast: &Ast, + original_update_stmt: &nodes::UpdateStmt, + new_route: Route, + ) -> Result { + let mut target_map = HashMap::new(); + let mut inlined = HashSet::new(); + for target in original_update_stmt.target_list() { + if let Some(name) = target.name() { + if let Ok(Value::Placeholder(number)) = Value::try_from(target.val()) { + target_map.insert(name.to_string(), number); + } else { + inlined.insert(name.to_string()); + } + } + } + + /// Builds the INSERT from the row the DELETE step returned + /// Avoids using schema cache *which can be stale* (in ref. to convos about this), + /// e.g. after a migration that bypassed pgdog; this is why the RowDescription is used. + #[derive(Clone, Debug)] + struct InsertStep { + /// The original UPDATE + ast: Ast, + /// Columns changed via a placeholder + target_map: HashMap, + /// Columns changed via an expression that *stays* in the SQL. + inlined: HashSet, + params: Option, + } + + impl InsertStep { + /// If `Ok(None)`: the DELETE matched nothing; there's no row to move. + fn deleted_row<'a>( + &self, + map: &'a ResponseHistory, + ) -> Result, Error> { + // TODO: Can this be a broad assert? + debug_assert!(map.get(SaveKey::ShardingKeyUpdateDelete).is_some()); + let delete = map + .get(SaveKey::ShardingKeyUpdateDelete) + .ok_or(UpdateError::MissingStepResponse("delete"))?; + + let rows = delete.rows.len(); + if rows >= 2 { + return Err(UpdateError::TooManyRows(rows).into()); + } + + let Some(data_row) = delete.rows.first() else { + return Ok(None); + }; + let row_description = delete + .row_description + .as_ref() + .ok_or(UpdateError::MissingStepResponse("delete"))?; + + Ok(Some((row_description, data_row))) + } + + fn update_stmt(&self) -> Result<&nodes::UpdateStmt, Error> { + match self.ast.ast.stmts().exactly_one() { + Ok(Node::UpdateStmt(update)) => Ok(update), + _ => Err(UpdateError::NotAnUpdate.into()), + } + } + } + + impl StatementSource for InsertStep { + fn resolve(&self, map: &ResponseHistory) -> Result, Error> { + let Some((row_description, data_row)) = self.deleted_row(map)? else { + // Nothing was deleted; skip the INSERT; emit UPDATE 0 + return Ok(None); + }; + let update = self.update_stmt()?; + + let insert = try_owned(|mem| -> Result<_, Error> { + let mut columns = Vec::new(); + let mut values = Vec::new(); + let mut placeholders = 0; + + for field in row_description.fields.iter() { + let name = field.name.as_str(); + columns.push( + mem.make_res_target(Some(name), mem.empty(), mem.none()) + .uncast(), + ); + + if self.inlined.contains(name) { + let value = update + .target_list() + .iter() + .find_map(|rt| { + if rt.name() == Some(name) { + Some(rt.val()) + } else { + None + } + }) + .expect("inlined columns come from the target list"); + values.push(mem.make_unique(value)); + } else { + // $1, $2, $3 + placeholders += 1; + values.push(mem.make_param_ref(placeholders).uncast()); + } + } + + let mut insert = mem.make_node::(); + insert + .as_mut() + .set_relation(mem.make_unique(update.relation())); + insert.as_mut().set_cols(mem.make_list(&columns)); + let mut select = mem.make_node::(); + select + .as_mut() + .set_values_lists(mem.make_list(&[mem.make_list(&values)])); + insert.as_mut().set_select_stmt(select.uncast()); + insert + .as_mut() + .set_returning_clause(mem.make_unique(update.returning_clause())); + Ok(mem.make_list(&[mem.make_raw_stmt(insert.uncast())])) + })?; + + let parse = Parse::new_anonymous(deparse(insert.first().unwrap())?.as_str()); + + let mut bind = Bind::new_statement(""); + for (idx, field) in row_description.fields.iter().enumerate() { + let name = field.name.as_str(); + if self.inlined.contains(name) { + continue; + } + + if let Some(number) = self.target_map.get(name) { + let number = *number; + let param = self + .params + .as_ref() + .and_then(|p| p.parameter(number as usize - 1).transpose()) + .ok_or(RewriteError::MissingParameter(number as u16))??; + bind.push_param(param.parameter().clone(), param.format()); + } else { + // This column wasn't changed, get the value from the select. + debug_assert!(data_row.get_raw(idx).is_some()); + let value = data_row + .get_raw(idx) + .ok_or(RewriteError::MissingColumn(idx))?; + + if value.is_null() { + bind.push_param(Parameter::new_null(), Format::Text); + } else { + bind.push_param(Parameter::new(value), Format::Text); + } + } + } + + Ok(Some((parse, bind))) + } + } + + let insert = InsertStep { + ast: ast.clone(), + target_map, + inlined, + params: request.parameters()?.cloned(), + }; + + Ok(Step { + save_key: Some(SaveKey::ShardingKeyUpdateInsert), + request: StepRequest::Statement(Box::new(StatementRequest { + source: Box::new(insert), + protocol: StepProtocol::Extended, + route: new_route, + ast: None, + })), + }) + } + + /// Static DELETE, RETURNING * + fn construct_delete_step( + engine: &QueryEngine, + client_request: &ClientRequest, + original_update_stmt: &nodes::UpdateStmt, + context: &QueryEngineContext<'_>, + ) -> Result { + let mut params = IndexSet::new(); + let delete = owned(|mem| { + let mut delete = mem.make_node::(); + delete + .as_mut() + .set_relation(mem.make_unique(original_update_stmt.relation())); + delete + .as_mut() + .set_where_clause(mem.make_unique(original_update_stmt.where_clause())); + delete.as_mut().set_returning_clause( + mem.make_returning_clause( + mem.make_list(&[mem + .make_res_target( + None, + mem.empty(), + mem.make_column_ref( + mem.make_list(&[mem.make_node::().uncast()]), + ) + .uncast(), + ) + .uncast()]), + ) + .as_option(), + ); + params = QueryPlanner::rewrite_params(delete.as_mut().into()); + mem.make_list(&[mem.make_raw_stmt(delete.uncast())]) + }); + + let stmt = deparse(delete.first().unwrap())?.as_str().to_owned(); + let ast = Ast::from_raw_stmts(delete); + + Ok(Step { + save_key: Some(SaveKey::ShardingKeyUpdateDelete), + request: Self::build_request( + engine, + context, + client_request, + &ast, + &stmt, + ¶ms, + None, + )?, + }) + } + + /// If we do a SELECT with the new sharding key as the target, + /// will it differ from a SELECT with the old sharding key in terms of + /// which `Shard` it resolves to? + fn on_same_shard( + engine: &QueryEngine, + new_target: &ResTarget, + original_update_stmt: &nodes::UpdateStmt, + original_request: &ClientRequest, + context: &QueryEngineContext, + ) -> Result, Error> { + let select_star = owned(|mem| { + let mut select_stmt = mem.make_node::(); + select_stmt.as_mut().set_target_list( + mem.make_list(&[mem.make_res_target( + None, + mem.empty(), + mem.make_column_ref( + mem.make_list(&[mem.make_node::().uncast()]), + ) + .uncast(), + )]), + ); + select_stmt.as_mut().set_from_clause( + mem.make_list(&[mem + .make_unique( + original_update_stmt + .relation() + .expect("UPDATE always has a table"), + ) + .uncast()]), + ); + select_stmt + }); + + let mut params = IndexSet::new(); + let check = owned(|mem| { + let mut select_stmt = mem.make_unique(&*select_star); + select_stmt.as_mut().set_where_clause( + mem.make_a_expr( + nodes::A_Expr_Kind::AEXPR_OP, + mem.make_list(&[mem.make_string(Some("=")).uncast()]), + mem.make_column_ref( + mem.make_list(&[mem.make_string(new_target.name()).uncast()]), + ) + .uncast(), + mem.make_unique(new_target.val()), + ) + .uncast(), + ); + params = Self::rewrite_params(select_stmt.as_mut().into()); + mem.make_list(&[mem.make_raw_stmt(select_stmt.uncast())]) + }); + + let stmt = deparse(check.first().unwrap())?; + let ast = &Ast::from_raw_stmts(check); + + let check_request = Self::build_request( + engine, + context, + original_request, + ast, + stmt.as_str(), + ¶ms, + None, + )?; + + let new_route = match check_request { + StepRequest::Statement(statement) => statement.route, + StepRequest::Raw => unreachable!("build_request always returns a routed statement"), + }; + + let same_shard = match original_request.route.as_ref() { + Some(route) => route.shard().eq(new_route.shard()), + None => { + let mut original = original_request.clone(); + Self::route(engine, &mut original, context)?; + original.route().shard().eq(new_route.shard()) + } + }; + + Ok((!same_shard).then_some(new_route)) + } + + /// Returns true if a column is referenced by a foreign key whose ON DELETE + /// action would be unsafe during a sharding-key row move. + fn has_destructive_on_delete_reference( + engine: &QueryEngine, + from_update: &nodes::UpdateStmt, + context: &QueryEngineContext<'_>, + ) -> Result { + let cluster = engine.backend.cluster()?; + let schema = cluster.schema(); + let table = Self::target_table(from_update); + + let Some(relation) = schema.table(table, cluster.user(), context.params.search_path()) + else { + return Ok(false); + }; + let Some(sharded_table) = Self::sharded_table(cluster.sharded_tables(), from_update) else { + return Ok(false); + }; + + Ok(schema.has_destructive_on_delete_reference( + relation.schema(), + &relation.name, + &sharded_table.column, + )) + } + + pub(crate) fn sharded_table<'a>( + sharded_tables: &'a [ShardedTable], + from_update: &nodes::UpdateStmt, + ) -> Option<&'a ShardedTable> { + let table = Self::target_table(from_update); + + sharded_tables.iter().find(|sharded| { + if let Some(name) = sharded.name.as_ref() + && !table.name_match(name) + { + return false; + } + + if let Some(schema) = sharded.schema.as_ref() + && let Some(table_schema) = table.schema + && table_schema != schema + { + return false; + } + + from_update + .target_list() + .iter() + .any(|rt| rt.name() == Some(&*sharded.column)) + }) + } + + pub(crate) fn target_table(from_update: &nodes::UpdateStmt) -> Table<'_> { + Table::from(from_update.relation().expect("UPDATE always has table")) + } +} diff --git a/pgdog/src/frontend/client/query_engine/multi_step/shared.rs b/pgdog/src/frontend/client/query_engine/multi_step/shared.rs new file mode 100644 index 000000000..ba9b80d48 --- /dev/null +++ b/pgdog/src/frontend/client/query_engine/multi_step/shared.rs @@ -0,0 +1,373 @@ +// TODO: This could use more unit tests +// TODO: Can maybe use lifetimes to get rid of some .clone()s + +use crate::frontend::client::query_engine::multi_step::error::{Error, UpdateError}; +use crate::frontend::client::query_engine::multi_step::forward_check::ForwardCheck; +use crate::frontend::client::query_engine::multi_step::types::{ + QueryPlanner, QueryPlannerType, ResponseHistory, StatementRequest, StatementSource, + StepProtocol, StepRequest, StepResponses, +}; +use crate::frontend::client::query_engine::{QueryEngine, QueryEngineContext}; +use crate::frontend::router::parser::rewrite::statement::Error as RewriteError; +use crate::frontend::router::{Ast, Route}; +use crate::frontend::{BufferedQuery, ClientRequest, Command, Router, RouterContext}; +use crate::net::{ + Bind, BindComplete, CommandComplete, DataRow, Describe, ErrorResponse, Execute, FromBytes, + Message, Parse, ParseComplete, Protocol, Query, ReadyForQuery, RowDescription, Sync, ToBytes, + TransactionState, +}; +use bytes::Buf; +use indexmap::IndexSet; +use pg_raw_parse::{NodeMut, walk}; + +impl StatementRequest { + /// Construct a `ClientRequest` based on statically/dynamically resolving the current `Step` + /// based on prior `Step` responses. Handles both simple and extended protocol. + pub(crate) fn assemble(&self, map: &ResponseHistory) -> Result, Error> { + let Some((parse, bind)) = self.source.resolve(map)? else { + return Ok(None); + }; + + let mut request = ClientRequest::default(); + match self.protocol { + StepProtocol::Simple => { + request.push(Query::new(parse.query()).into()); + } + StepProtocol::Extended => { + let name = parse.name().to_owned(); + request.push(parse.into()); + // So we get both T and t. + request.push(Describe::new_statement(&name).into()); + request.push(bind.into()); + request.push(Execute::new().into()); + request.push(Sync.into()); + } + } + request.route = Some(self.route.clone()); + request.ast = self.ast.clone(); + + Ok(Some(request)) + } +} + +/// Represents a statement that's fully known during planning. Needs no dynamic resolving. +#[derive(Debug, Clone)] +struct RewrittenStatement { + parse: Parse, + bind: Bind, +} + +impl StatementSource for RewrittenStatement { + fn resolve(&self, _map: &ResponseHistory) -> Result, Error> { + Ok(Some((self.parse.clone(), self.bind.clone()))) + } +} + +impl ResponseHistory { + /// Protocol acks and the row set the client expects ahead of the final `CommandComplete`, + /// composed from the given `steps` responses and filtered by what the client's request asked for. + pub(crate) fn compose( + steps: &[&StepResponses], + client_request: &ClientRequest, + ) -> Vec { + let mut check = ForwardCheck::new(client_request); + let mut messages = Vec::new(); + + if check.forward('1') { + messages.push(ParseComplete.message()); + } + if let Some(parameter_description) = steps + .iter() + .find_map(|step| step.parameter_description.clone()) + && check.forward('t') + { + messages.push(parameter_description); + } + if let Some(row_description) = steps.iter().find_map(|step| step.row_description.clone()) + && check.forward('T') + { + messages.push(row_description.message()); + } + if check.forward('2') { + messages.push(BindComplete.message()); + } + for step in steps { + for row in &step.rows { + if check.forward('D') { + messages.push(row.message()); + } + } + } + + messages + } + + /// [`Self::compose`] for every `Step` (maintaining execution order) + pub(crate) fn compose_all(&self, client_request: &ClientRequest) -> Vec { + Self::compose(&self.steps().iter().collect::>(), client_request) + } +} + +impl QueryEngine { + /// Handles execution of all `Step`s in the `QueryPlanner` one-by-one, in a serial, + /// sequential way. Handles dynamic resolving, executing returned `ClientRequest`s, + /// checking to see if we should save Responses, and checking to see if we should forward. + pub(crate) async fn run_steps( + &mut self, + context: &mut QueryEngineContext<'_>, + planner: &QueryPlanner, + ) -> Result<(), crate::frontend::Error> { + // TODO: I think there needs to be some work in regard to this; current functionality: + // + // - Aggregated plans buffer every step's responses and `ForwardToClient` + // composes the client's reply from them at the end + // + // - Normal plan streams its single Step responses. + // + // If we consider a subquery, we may want to stream the outer query's Responses without + // waiting. Sage and I discussed briefly about potential Postgres protocol violation w.r.t. + // an Error occuring after Rows are sent back. Haven't tested yet, but thought I'd mention + // here as to future work. + let mut map = ResponseHistory::default(); + let aggregate = planner.forward_to_client.is_some(); + + // Iterate serially + for step in &planner.steps { + let assembled; + let client_request = match &step.request { + StepRequest::Raw => context.client_request, + StepRequest::Statement(statement) => match statement.assemble(&map)? { + Some(request) => { + assembled = request; + &assembled + } + // The `Step` resolves to nothing + // (e.g. an INSERT whose source row doesn't exist), skip it here + None => continue, + }, + }; + + self.backend + .handle_client_request(client_request, &mut self.router, self.streaming) + .await?; + + let mut responses = StepResponses { + key: step.save_key, + ..Default::default() + }; + let mut step_error = None; + while self.backend.has_more_messages() + && !self.backend.in_copy_mode() + && !self.streaming + { + let message = self.read_server_message().await?; + if aggregate && message.code() == 'E' { + step_error = Some(ErrorResponse::try_from(message)?); + continue; + } + + // Prevent case where: + // - A step failed. + // - Trailing RFQ is held from the client, + // Its aborted-transaction state should still be applied. + // Otherwise, we think the transaction is healthy and COMMIT half-commits + if step_error.is_some() + && message.code() == 'Z' + && ReadyForQuery::from_bytes(message.to_bytes())?.state()? + == TransactionState::Error + { + context.set_transaction_error(); + continue; + } + + if aggregate { + match message.code() { + 'T' => { + responses.row_description = + Some(RowDescription::from_bytes(message.to_bytes())?) + } + 't' => responses.parameter_description = Some(message), + 'D' => responses + .rows + .push(DataRow::from_bytes(message.to_bytes())?), + 'C' => { + responses.command_complete = + Some(CommandComplete::from_bytes(message.to_bytes())?) + } + _ => (), + } + } else { + self.process_server_message(context, message).await?; + } + } + + if let Some(error) = step_error { + return Err(Error::Execution(Box::new(error)).into()); + } + + if aggregate { + map.push(responses); + } + } + + // Forward whatever we need to the client at the end (in aggregate) + if let Some(ftc) = &planner.forward_to_client { + let messages = ftc.forward_to_client(context, map); + for message_to_forward in messages { + self.process_server_message(context, message_to_forward) + .await?; + } + } + + Ok(()) + } +} + +impl QueryPlanner { + pub(crate) async fn plan_query( + request: &ClientRequest, + engine: &mut QueryEngine, + context: &mut QueryEngineContext<'_>, + mut query_planner_type: QueryPlannerType, + ) -> Result, Error> { + if !request.is_executable() { + query_planner_type = QueryPlannerType::Normal; + }; + + // Based on `QueryPlannerType`, plan out the `Steps` we should take. + match query_planner_type { + QueryPlannerType::InsertSplit => Self::plan_multi_insert(engine, context, request), + QueryPlannerType::ShardingKeyUpdate => { + let schema = engine.backend.cluster()?.sharding_schema(); + Self::plan_sharding_key_update(request, engine, context, schema).await + } + QueryPlannerType::Normal => Ok(None), + } + } + + /// Build a routed `StepRequest` from statement. + /// Use the same protocol as the original statement. + /// + /// TODO: This is a lot of parameters passed in; move / turn it into a struct? + pub(crate) fn build_request( + engine: &QueryEngine, + context: &QueryEngineContext<'_>, + original: &ClientRequest, + ast: &Ast, + stmt: &str, + params: &IndexSet, + statement_name: Option<&str>, + ) -> Result { + let query = original.query()?.ok_or(RewriteError::EmptyQuery)?; + let name = statement_name.unwrap_or_default(); + + let (protocol, parse, bind) = match query { + BufferedQuery::Query(_) => ( + StepProtocol::Simple, + Parse::named(name, stmt), + Bind::new_statement(name), + ), + BufferedQuery::Prepared(original_parse) => { + let data_types = Self::rewrite_data_types(&original_parse, params); + let bind = match original.parameters()? { + Some(bind) => Self::rewrite_bind(params, bind, name)?, + // This shouldn't really happen since we don't rewrite + // non-executable requests. + None => Bind::new_statement(name), + }; + ( + StepProtocol::Extended, + Parse::named(name, stmt).with_data_types(&data_types), + bind, + ) + } + }; + + let mut statement = StatementRequest { + source: Box::new(RewrittenStatement { parse, bind }), + protocol, + route: Route::default(), + ast: Some(ast.clone()), + }; + + // Deliberately uses an empty history so we can route the solo request. + let mut probe = statement + .assemble(&ResponseHistory::default())? + .ok_or(RewriteError::EmptyQuery)?; + Self::route(engine, &mut probe, context)?; + statement.route = probe.route.take().unwrap_or_default(); + + Ok(StepRequest::Statement(Box::new(statement))) + } + + fn rewrite_data_types(parse: &Parse, params: &IndexSet) -> Vec { + let mut bytes = parse.data_types_ref(); + let count = bytes.get_i16().max(0) as usize; + let declared: Vec = (0..count).map(|_| bytes.get_u32()).collect(); + + params + .iter() + .map(|original| { + declared + .get(*original as usize - 1) + .copied() + .unwrap_or_default() + }) + .collect() + } + + pub(crate) fn route( + query_engine: &QueryEngine, + request: &mut ClientRequest, + context: &QueryEngineContext<'_>, + ) -> Result<(), Error> { + let cluster = query_engine.backend.cluster()?; + + let context = RouterContext::new( + request, + cluster, + context.params, + context.transaction(), + context.sticky, + )?; + let mut router = Router::new(); + let command = router.query(context)?; + if let Command::Query(route) = command { + request.route = Some(route.clone()); + } else { + return Err(UpdateError::NoRoute.into()); + } + + Ok(()) + } + + /// Visit all ParamRef nodes in a ParseResult and renumber them sequentially. + /// Returns a sorted list of the original parameter numbers. + pub(crate) fn rewrite_params(node: NodeMut<'_, '_>) -> IndexSet { + let mut params = IndexSet::new(); + walk::walk_mut(node, |node| { + if let NodeMut::ParamRef(param) = node { + params.insert(param.number as _); + param.set_number(params.get_index_of(&(param.number as u16)).unwrap() as i32 + 1) + } + }); + params + } + + /// Create new Bind message for the statement from original Bind. + pub(crate) fn rewrite_bind( + params: &IndexSet, + bind: &Bind, + statement_name: &str, + ) -> Result { + let mut new = Bind::new_statement(statement_name); + for param in params { + let param = bind + .parameter(*param as usize - 1)? + .ok_or(RewriteError::MissingParameter(*param))?; + new.push_param(param.parameter().clone(), param.format()); + } + + Ok(new) + } +} diff --git a/pgdog/src/frontend/client/query_engine/multi_step/state.rs b/pgdog/src/frontend/client/query_engine/multi_step/state.rs deleted file mode 100644 index d10116840..000000000 --- a/pgdog/src/frontend/client/query_engine/multi_step/state.rs +++ /dev/null @@ -1,83 +0,0 @@ -use fnv::FnvHashMap as HashMap; - -use super::super::Error; -use crate::net::{CommandComplete, FromBytes, Message, Protocol, ReadyForQuery, ToBytes}; - -#[derive(Debug, Clone)] -pub(crate) enum CommandType { - Insert, -} - -#[derive(Debug, Clone)] -pub(crate) struct MultiServerState { - servers: usize, - rows: usize, - counters: HashMap, -} - -impl MultiServerState { - /// New multi-server execution state. - pub(crate) fn new(servers: usize) -> Self { - Self { - servers, - rows: 0, - counters: HashMap::default(), - } - } - - /// Should the message be forwarded to the client. - pub(crate) fn forward(&mut self, message: &Message) -> Result { - let code = message.code(); - let count = self.counters.entry(code).or_default(); - *count += 1; - - Ok(match code { - 'T' | '1' | '2' | '3' | 't' => *count == 1, - 'C' => { - let command_complete = CommandComplete::from_bytes(message.to_bytes())?; - self.rows += command_complete.rows()?.unwrap_or(0); - false - } - 'Z' => false, - 'n' => *count == self.servers && !self.counters.contains_key(&'D'), - 'I' => *count == self.servers && !self.counters.contains_key(&'C'), - _ => true, - }) - } - - /// Number of rows returned. - pub(crate) fn rows(&self) -> usize { - self.rows - } - - /// Error happened. - pub(crate) fn error(&self) -> bool { - self.counters.contains_key(&'E') - } - - /// Create CommandComplete (C) message. - pub(crate) fn command_complete(&self, command_type: CommandType) -> Option { - if !self.counters.contains_key(&'C') || self.error() { - return None; - } - - let name = match command_type { - CommandType::Insert => "INSERT 0", - }; - - Some(CommandComplete::new(format!("{} {}", name, self.rows()))) - } - - /// Create ReadyForQuery (C) message. - pub(crate) fn ready_for_query(&self, in_transaction: bool) -> Option { - if !self.counters.contains_key(&'Z') { - return None; - } - - if self.error() { - Some(ReadyForQuery::error()) - } else { - Some(ReadyForQuery::in_transaction(in_transaction)) - } - } -} diff --git a/pgdog/src/frontend/client/query_engine/multi_step/test/insert.rs b/pgdog/src/frontend/client/query_engine/multi_step/test/insert.rs index 7385e3bef..d6cb367d7 100644 --- a/pgdog/src/frontend/client/query_engine/multi_step/test/insert.rs +++ b/pgdog/src/frontend/client/query_engine/multi_step/test/insert.rs @@ -23,15 +23,15 @@ async fn test_same_shard_insert_uses_direct_route() { ]); let mut context = QueryEngineContext::new(&mut client.client); - let rewrite_result = client.engine.parse_and_rewrite(&mut context).unwrap(); + let (query_planner, offset_plan) = client.engine.parse_and_rewrite(&mut context).await.unwrap(); client .engine - .route_query(&mut context, rewrite_result.as_ref()) + .route_query(&mut context, offset_plan.as_ref()) .await .unwrap(); client .engine - .execute(&mut context, rewrite_result) + .execute(&mut context, query_planner) .await .unwrap(); @@ -58,15 +58,15 @@ async fn test_cross_shard_insert_uses_all_shards() { ]); let mut context = QueryEngineContext::new(&mut client.client); - let rewrite_result = client.engine.parse_and_rewrite(&mut context).unwrap(); + let (query_planner, offset_plan) = client.engine.parse_and_rewrite(&mut context).await.unwrap(); client .engine - .route_query(&mut context, rewrite_result.as_ref()) + .route_query(&mut context, offset_plan.as_ref()) .await .unwrap(); client .engine - .execute(&mut context, rewrite_result) + .execute(&mut context, query_planner) .await .unwrap(); diff --git a/pgdog/src/frontend/client/query_engine/multi_step/test/mod.rs b/pgdog/src/frontend/client/query_engine/multi_step/test/mod.rs index 6ee7f9981..af72e2eb9 100644 --- a/pgdog/src/frontend/client/query_engine/multi_step/test/mod.rs +++ b/pgdog/src/frontend/client/query_engine/multi_step/test/mod.rs @@ -7,7 +7,9 @@ use crate::{ pub(crate) mod insert; pub(crate) mod prepared; +pub(crate) mod sharding_key_update; pub(crate) mod simple; +pub(crate) mod split_insert; pub(crate) mod update; async fn truncate_table(table: &str, stream: &mut TcpStream) { diff --git a/pgdog/src/frontend/client/query_engine/multi_step/test/sharding_key_update.rs b/pgdog/src/frontend/client/query_engine/multi_step/test/sharding_key_update.rs new file mode 100644 index 000000000..0537732e6 --- /dev/null +++ b/pgdog/src/frontend/client/query_engine/multi_step/test/sharding_key_update.rs @@ -0,0 +1,642 @@ +// Note: Most of these tests previously were in parser/rewrite/statement/update.rs + +use crate::frontend::router::sharding::ShardedTable; +use indexmap::{IndexSet, indexset}; +use pg_raw_parse::{Node, nodes}; +use pgdog_config::{Rewrite, RewriteMode}; + +use crate::backend::ShardingSchema; +use crate::backend::{ShardedTables, replication::ShardedSchemas}; +use crate::frontend::ClientRequest; +use crate::frontend::client::query_engine::QueryEngineContext; +use crate::frontend::client::query_engine::multi_step::error::Error; +use crate::frontend::client::query_engine::multi_step::types::{ + QueryPlanner, ResponseHistory, SaveKey, StatementSource, StepRequest, StepResponses, +}; +use crate::frontend::client::test::TestClient; +use crate::frontend::router::parser::rewrite::statement::Error as RewriteError; +use crate::frontend::router::parser::{AstContext, Cache, Error as ParserError}; +use crate::net::messages::row_description::Field; +use crate::net::{ + Bind, DataRow, Execute, Parameters, Parse, Query, RowDescription, Sync, bind::Parameter, +}; + +fn default_schema() -> ShardingSchema { + ShardingSchema { + shards: 2, + tables: ShardedTables::new( + vec![ShardedTable { + database: "pgdog".into(), + name: Some("sharded".into()), + column: "id".into(), + ..Default::default() + }], + vec![], + false, + pgdog_config::SystemCatalogsBehavior::default(), + ), + schemas: ShardedSchemas::new(vec![]), + rewrite: Rewrite { + enabled: true, + shard_key: RewriteMode::Rewrite, + ..Default::default() + }, + ..Default::default() + } +} + +#[derive(Debug)] +struct Statement { + stmt: String, + params: IndexSet, +} + +#[derive(Debug)] +struct TargetTable { + name: String, +} + +#[derive(Debug)] +struct ShardingKeyUpdate { + query: String, + delete: Statement, + insert: Option>, +} + +impl ShardingKeyUpdate { + fn with_update(&self, f: impl FnOnce(&nodes::UpdateStmt) -> R) -> R { + let stmt = pg_raw_parse::parse(&self.query).unwrap(); + match stmt.stmts().next().unwrap() { + Node::UpdateStmt(stmt) => f(stmt), + _ => panic!("Not an update"), + } + } + + fn is_returning(&self) -> bool { + self.with_update(|update| update.returning_clause().is_some()) + } + + fn target_table(&self) -> TargetTable { + self.with_update(|update| TargetTable { + name: QueryPlanner::target_table(update).name.to_string(), + }) + } + + fn sharded_table(&self, tables: &[ShardedTable]) -> Option { + self.with_update(|update| { + QueryPlanner::sharded_table(tables, update).map(|table| table.column.to_string()) + }) + } +} + +fn placeholders(sql: &str) -> u16 { + sql.split('$') + .skip(1) + .filter_map(|rest| { + rest.chars() + .take_while(|c| c.is_ascii_digit()) + .collect::() + .parse::() + .ok() + }) + .max() + .unwrap_or_default() +} + +fn bind_params(request: &ClientRequest) -> IndexSet { + request + .parameters() + .unwrap() + .map(|bind| { + bind.params_raw() + .iter() + .map(|param| { + std::str::from_utf8(¶m.data) + .unwrap() + .parse::() + .unwrap() + }) + .collect() + }) + .unwrap_or_default() +} + +async fn run_test_with( + client: &mut TestClient, + query: &str, +) -> Result, Error> { + client.send_simple(Query::new("BEGIN")).await; + client.read_until('Z').await.unwrap(); + + let params = (1..=placeholders(query)) + .map(|number| Parameter::new(number.to_string().as_bytes())) + .collect::>(); + client.client.client_request = ClientRequest::from(vec![ + Parse::new_anonymous(query).into(), + Bind::new_params("", ¶ms).into(), + Execute::new().into(), + Sync.into(), + ]); + + let mut context = QueryEngineContext::new(&mut client.client); + + let ast = { + let cluster = client.engine.backend.cluster()?; + let ast_context = AstContext::from_cluster(cluster, context.params); + let buffered = context.client_request.query()?.unwrap(); + Cache::get().query(&buffered, &ast_context, context.prepared_statements)? + }; + context.client_request.ast = Some(ast); + + client.engine.route_query(&mut context, None).await?; + client.engine.connect_transaction(&mut context).await?; + + let schema = client.engine.backend.cluster()?.sharding_schema(); + let request = context.client_request.clone(); + let Some(planner) = + QueryPlanner::plan_sharding_key_update(&request, &mut client.engine, &mut context, schema) + .await? + else { + return Ok(None); + }; + + let StepRequest::Statement(ref statement) = planner.steps[0].request else { + unreachable!("delete should not be raw") + }; + let delete = statement + .assemble(&ResponseHistory::default())? + .expect("delete step resolves statically"); + let delete = Statement { + stmt: delete.query()?.unwrap().query().to_string(), + params: bind_params(&delete), + }; + + let insert = planner.steps.get(1).and_then(|step| match &step.request { + StepRequest::Statement(statement) => Some(statement.source.clone()), + _ => None, + }); + + Ok(Some(ShardingKeyUpdate { + query: query.to_string(), + delete, + insert, + })) +} + +async fn run_test(query: &str) -> Result, Error> { + let mut client = TestClient::new_rewrites(Parameters::default()).await; + run_test_with(&mut client, query).await +} + +#[tokio::test] +async fn test_select_basic_where_param() { + let result = run_test("UPDATE sharded SET id = $1 WHERE email = $2") + .await + .unwrap() + .unwrap(); + + // SELECT should have WHERE clause with param renumbered to $1 + assert_eq!( + result.delete.stmt, + "DELETE FROM sharded WHERE email = $1 RETURNING *" + ); + assert_eq!(result.delete.params, indexset![2]); + + let schema = default_schema(); + let tables = schema.tables.tables(); + assert_eq!(result.target_table().name, "sharded"); + assert_eq!(result.sharded_table(tables).unwrap(), "id"); +} + +#[tokio::test] +async fn test_select_multiple_where_params() { + let result = run_test("UPDATE sharded SET id = $1 WHERE email = $2 AND name = $3") + .await + .unwrap() + .unwrap(); + + assert_eq!( + result.delete.stmt, + "DELETE FROM sharded WHERE email = $1 AND name = $2 RETURNING *" + ); + assert_eq!(result.delete.params, indexset![2, 3]); + assert!(!result.is_returning()); +} + +#[tokio::test] +async fn test_select_non_sequential_params() { + // Params in WHERE are $3 and $5, should be renumbered to $1 and $2 + let result = run_test( + "UPDATE sharded SET id = $1, value = $2, other = $4 WHERE email = $3 AND name = $5", + ) + .await + .unwrap() + .unwrap(); + + assert_eq!( + result.delete.stmt, + "DELETE FROM sharded WHERE email = $1 AND name = $2 RETURNING *" + ); + assert_eq!(result.delete.params, indexset![3, 5]); +} + +#[tokio::test] +async fn test_delete_basic() { + let result = run_test("UPDATE sharded SET id = $1 WHERE email = $2") + .await + .unwrap() + .unwrap(); + + assert_eq!( + result.delete.stmt, + "DELETE FROM sharded WHERE email = $1 RETURNING *" + ); + + assert!(result.sharded_table(&[]).is_none()); + assert!( + result + .sharded_table(&[ShardedTable { + name: Some("other".into()), + column: "id".into(), + ..Default::default() + }]) + .is_none() + ); + assert!( + result + .sharded_table(&[ShardedTable { + name: Some("sharded".into()), + column: "user_id".into(), + ..Default::default() + }]) + .is_none() + ); +} + +#[tokio::test] +async fn test_no_params_in_where() { + let result = run_test("UPDATE sharded SET id = $1 WHERE email = 'test@example.com'") + .await + .unwrap() + .unwrap(); + + assert_eq!( + result.delete.stmt, + "DELETE FROM sharded WHERE email = 'test@example.com' RETURNING *" + ); + assert!(result.delete.params.is_empty()); +} + +#[tokio::test] +async fn test_where_with_in_clause() { + let result = run_test("UPDATE sharded SET id = $1 WHERE email IN ($2, $3, $4)") + .await + .unwrap() + .unwrap(); + + assert_eq!( + result.delete.stmt, + "DELETE FROM sharded WHERE email IN ($1, $2, $3) RETURNING *" + ); + assert_eq!(result.delete.params, indexset![2, 3, 4]); +} + +#[tokio::test] +async fn test_where_with_comparison_operators() { + let result = run_test("UPDATE sharded SET id = $1 WHERE count > $2 AND count < $3") + .await + .unwrap() + .unwrap(); + + assert_eq!( + result.delete.stmt, + "DELETE FROM sharded WHERE count > $1 AND count < $2 RETURNING *" + ); + assert_eq!(result.delete.params, indexset![2, 3]); +} + +#[tokio::test] +async fn test_where_with_or_condition() { + let result = run_test("UPDATE sharded SET id = $1 WHERE email = $2 OR name = $3") + .await + .unwrap() + .unwrap(); + + assert_eq!( + result.delete.stmt, + "DELETE FROM sharded WHERE email = $1 OR name = $2 RETURNING *" + ); + assert_eq!(result.delete.params, indexset![2, 3]); +} + +#[tokio::test] +async fn test_high_param_numbers() { + let result = run_test("UPDATE sharded SET id = $10 WHERE email = $20 AND name = $30") + .await + .unwrap() + .unwrap(); + + assert_eq!( + result.delete.stmt, + "DELETE FROM sharded WHERE email = $1 AND name = $2 RETURNING *" + ); + assert_eq!(result.delete.params, indexset![20, 30]); +} + +#[tokio::test] +async fn test_non_sharding_key_update_returns_none() { + // Updating a non-sharding column should return None + let result = run_test("UPDATE sharded SET email = $1 WHERE id = $2") + .await + .unwrap(); + assert!(result.is_none()); +} + +#[tokio::test] +async fn test_where_with_like() { + let result = run_test("UPDATE sharded SET id = $1 WHERE email LIKE $2") + .await + .unwrap() + .unwrap(); + + assert_eq!( + result.delete.stmt, + "DELETE FROM sharded WHERE email LIKE $1 RETURNING *" + ); + assert_eq!(result.delete.params, indexset![2]); +} + +#[tokio::test] +async fn test_where_with_is_null() { + let result = run_test("UPDATE sharded SET id = $1 WHERE email = $2 AND deleted_at IS NULL") + .await + .unwrap() + .unwrap(); + + assert_eq!( + result.delete.stmt, + "DELETE FROM sharded WHERE email = $1 AND deleted_at IS NULL RETURNING *" + ); + assert_eq!(result.delete.params, indexset![2]); +} + +#[tokio::test] +async fn test_where_with_between() { + let result = run_test("UPDATE sharded SET id = $1 WHERE created_at BETWEEN $2 AND $3") + .await + .unwrap() + .unwrap(); + + assert_eq!( + result.delete.stmt, + "DELETE FROM sharded WHERE created_at BETWEEN $1 AND $2 RETURNING *" + ); + assert_eq!(result.delete.params, indexset![2, 3]); +} + +#[tokio::test] +async fn test_same_param_used_twice() { + // Same parameter $2 used twice in WHERE clause + let result = run_test("UPDATE sharded SET id = $1 WHERE email = $2 OR name = $2") + .await + .unwrap() + .unwrap(); + + // Both occurrences should be renumbered to $1 + assert_eq!( + result.delete.stmt, + "DELETE FROM sharded WHERE email = $1 OR name = $1 RETURNING *" + ); + // Only one unique param in the mapping + assert_eq!(result.delete.params, indexset![2]); +} + +#[tokio::test] +async fn test_same_param_used_multiple_times() { + // $2 used three times + let result = run_test("UPDATE sharded SET id = $1 WHERE a = $2 AND b = $2 AND c = $2") + .await + .unwrap() + .unwrap(); + + assert_eq!( + result.delete.stmt, + "DELETE FROM sharded WHERE a = $1 AND b = $1 AND c = $1 RETURNING *" + ); + assert_eq!(result.delete.params, indexset![2]); +} + +#[tokio::test] +async fn test_mixed_repeated_and_unique_params() { + // $2 used twice, $3 used once + let result = run_test("UPDATE sharded SET id = $1 WHERE a = $2 AND b = $3 AND c = $2") + .await + .unwrap() + .unwrap(); + + assert_eq!( + result.delete.stmt, + "DELETE FROM sharded WHERE a = $1 AND b = $2 AND c = $1 RETURNING *" + ); + assert_eq!(result.delete.params, indexset![2, 3]); +} + +#[tokio::test] +async fn test_repeated_params_in_in_clause() { + // Same param repeated in IN clause (unusual but valid) + let result = run_test("UPDATE sharded SET id = $1 WHERE email IN ($2, $3, $2)") + .await + .unwrap() + .unwrap(); + + assert_eq!( + result.delete.stmt, + "DELETE FROM sharded WHERE email IN ($1, $2, $1) RETURNING *" + ); + assert_eq!(result.delete.params, indexset![2, 3]); +} + +#[tokio::test] +async fn test_sharding_key_not_changed() { + let result = run_test("UPDATE sharded SET id = $1 WHERE id = $1 AND email = $2") + .await + .unwrap(); + assert!(result.is_none()); +} + +#[tokio::test] +async fn test_unsupported_assignment() { + let result = run_test("UPDATE sharded SET id = random() WHERE id = $1").await; + std::assert_matches!( + result, + Err(Error::Parser(ParserError::Rewrite(RewriteError::UnsupportedShardingKeyUpdate(msg)))) if msg == "\"id\" = random()" + ); +} + +#[tokio::test] +async fn test_unsupported_assignment_arithmetic_add() { + let result = run_test("UPDATE sharded SET id = id + 1 WHERE id = $1").await; + std::assert_matches!( + result, + Err(Error::Parser(ParserError::Rewrite(RewriteError::UnsupportedShardingKeyUpdate(msg)))) if msg == "\"id\" = id + 1" + ); +} + +#[tokio::test] +async fn test_unsupported_assignment_arithmetic_multiply() { + let result = run_test("UPDATE sharded SET id = id * 2 WHERE id = $1").await; + std::assert_matches!( + result, + Err(Error::Parser(ParserError::Rewrite(RewriteError::UnsupportedShardingKeyUpdate(msg)))) if msg == "\"id\" = id * 2" + ); +} + +#[tokio::test] +async fn test_unsupported_assignment_arithmetic_with_param() { + let result = run_test("UPDATE sharded SET id = id + $2 WHERE id = $1").await; + std::assert_matches!( + result, + Err(Error::Parser(ParserError::Rewrite(RewriteError::UnsupportedShardingKeyUpdate(msg)))) if msg == "\"id\" = id + $2" + ); +} + +#[tokio::test] +async fn test_unsupported_assignment_now() { + let result = run_test("UPDATE sharded SET id = now() WHERE id = $1").await; + std::assert_matches!( + result, + Err(Error::Parser(ParserError::Rewrite(RewriteError::UnsupportedShardingKeyUpdate(msg)))) if msg == "\"id\" = now()" + ); +} + +#[tokio::test] +async fn test_unsupported_assignment_coalesce() { + let result = run_test("UPDATE sharded SET id = coalesce(id, 0) WHERE id = $1").await; + std::assert_matches!( + result, + Err(Error::Parser(ParserError::Rewrite(RewriteError::UnsupportedShardingKeyUpdate(msg)))) if msg == "\"id\" = COALESCE(id, 0)" + ); +} + +#[tokio::test] +async fn test_unsupported_assignment_case() { + let result = + run_test("UPDATE sharded SET id = CASE WHEN id > 0 THEN 1 ELSE 0 END WHERE id = $1").await; + std::assert_matches!( + result, + Err(Error::Parser(ParserError::Rewrite(RewriteError::UnsupportedShardingKeyUpdate(msg)))) if msg == "\"id\" = CASE WHEN id > 0 THEN 1 ELSE 0 END" + ); +} + +#[tokio::test] +async fn test_unsupported_assignment_subquery() { + let result = + run_test("UPDATE sharded SET id = (SELECT max(id) FROM sharded) WHERE id = $1").await; + std::assert_matches!( + result, + Err(Error::Parser(ParserError::Rewrite(RewriteError::UnsupportedShardingKeyUpdate(msg)))) if msg == "\"id\" = (SELECT max(id) FROM sharded)" + ); +} + +#[tokio::test] +async fn test_unsupported_assignment_column_reference() { + let result = run_test("UPDATE sharded SET id = other_column WHERE id = $1").await; + std::assert_matches!( + result, + Err(Error::Parser(ParserError::Rewrite(RewriteError::UnsupportedShardingKeyUpdate(msg)))) if msg == "\"id\" = other_column" + ); +} + +#[tokio::test] +async fn test_unsupported_assignment_concat() { + let result = run_test("UPDATE sharded SET id = id || '_suffix' WHERE id = $1").await; + std::assert_matches!( + result, + Err(Error::Parser(ParserError::Rewrite(RewriteError::UnsupportedShardingKeyUpdate(msg)))) if msg == "\"id\" = id || '_suffix'" + ); +} + +#[tokio::test] +async fn test_unsupported_assignment_negation() { + let result = run_test("UPDATE sharded SET id = -id WHERE id = $1").await; + std::assert_matches!( + result, + Err(Error::Parser(ParserError::Rewrite(RewriteError::UnsupportedShardingKeyUpdate(msg)))) if msg == "\"id\" = - id" + ); +} + +#[tokio::test] +async fn test_insert_build_request_with_expr_column() { + // Test that INSERT statement is built correctly when there are expression columns. + // The expression should appear directly in the VALUES clause. + // Use literal values (not placeholders) to avoid needing bind parameters. + let mut client = TestClient::new_rewrites(Parameters::default()).await; + let old_id = client.random_id_for_shard(0); + let new_id = client.random_id_for_shard(1); + let result = run_test_with( + &mut client, + &format!("UPDATE sharded SET id = {new_id}, value = random() WHERE id = {old_id}"), + ) + .await + .unwrap() + .unwrap(); + + // Create a mock row description matching the SELECT * result + let row_description = RowDescription::new(&[ + Field::bigint("id"), + Field::text("value"), + Field::text("other_col"), + Field::text("other_other_col"), + ]); + + // Create a mock data row with values for columns not in the UPDATE SET clause + let mut data_row = DataRow::new(); + data_row.add("1"); // id - will be overwritten by mapping + data_row.add("old_value"); // value - will be overwritten by mapping + data_row.add("other_value"); // other_col - from existing row + data_row.add("other_other_value"); // other_other_col - from existing row + + // The INSERT is built from the DELETE step's response. + let mut map = ResponseHistory::default(); + map.push(StepResponses { + key: Some(SaveKey::ShardingKeyUpdateDelete), + row_description: Some(row_description), + rows: vec![data_row], + ..Default::default() + }); + + let stmt = result + .insert + .expect("insert step should exist") + .resolve(&map) + .unwrap() + .expect("statement resolves") + .0 + .query() + .to_string(); + + // The INSERT should contain the expression random() directly in VALUES + assert!( + stmt.contains("random()"), + "INSERT statement should contain the expression: {}", + stmt + ); + // Verify it's an INSERT statement + assert!( + stmt.starts_with("INSERT INTO"), + "Should be an INSERT statement: {}", + stmt + ); + // Verify parameter numbering is correct: $1 for id, random() for email, $2 for other_col + // (not $3, which would be wrong if we used row index instead of bind param index) + let placeholders = placeholders(&stmt); + assert_eq!( + placeholders, 2, + "one placeholder per non-inlined column: {stmt}" + ); + assert!( + (1..=placeholders).all(|number| stmt.contains(&format!("${number}"))), + "Parameter numbering should be sequential without gaps: {}", + stmt + ); +} diff --git a/pgdog/src/frontend/client/query_engine/multi_step/test/split_insert.rs b/pgdog/src/frontend/client/query_engine/multi_step/test/split_insert.rs new file mode 100644 index 000000000..14d54e40b --- /dev/null +++ b/pgdog/src/frontend/client/query_engine/multi_step/test/split_insert.rs @@ -0,0 +1,319 @@ +// Note: Most of these tests previously were in parser/rewrite/statement/insert.rs + +use pg_raw_parse::Node; + +use crate::frontend::client::query_engine::multi_step::types::StepRequest; +use crate::{ + frontend::{ + ClientRequest, + client::{ + query_engine::{ + QueryEngineContext, + multi_step::{ + error::Error, + ops::insert::create_steps, + types::{QueryPlanner, ResponseHistory}, + }, + }, + test::TestClient, + }, + router::parser::rewrite::statement::Error as RewriteError, + }, + net::{ + Bind, Execute, Parameters, Parse, Sync, + messages::bind::{Format, Parameter}, + }, +}; + +#[derive(Debug)] +struct InsertSplit { + params: Vec, + stmt: String, +} + +impl InsertSplit { + fn extract_bind_params(&self, bind: &Bind) -> Result { + QueryPlanner::rewrite_bind(&self.params.iter().copied().collect(), bind, "") + } +} + +fn placeholders(sql: &str) -> u16 { + sql.split('$') + .skip(1) + .filter_map(|rest| { + rest.chars() + .take_while(|c| c.is_ascii_digit()) + .collect::() + .parse::() + .ok() + }) + .max() + .unwrap_or_default() +} + +async fn parse_and_split(sql: &str) -> Vec { + let root = pg_raw_parse::parse(sql).unwrap(); + let insert = match root.stmts().next() { + Some(Node::InsertStmt(insert)) => insert, + _ => unreachable!(), + }; + + let params = (1..=placeholders(sql)) + .map(|number| Parameter::new(number.to_string().as_bytes())) + .collect::>(); + let request = ClientRequest::from(vec![ + Parse::new_anonymous(sql).into(), + Bind::new_params("", ¶ms).into(), + Execute::new().into(), + Sync.into(), + ]); + + let mut client = TestClient::new_rewrites(Parameters::default()).await; + client.client.client_request = request.clone(); + let mut context = QueryEngineContext::new(&mut client.client); + + let steps = create_steps(insert, &client.engine, &mut context, &request).unwrap(); + + steps + .into_iter() + .map(|step| { + let StepRequest::Statement(statement) = step.request else { + unreachable!("split should not be raw") + }; + let request = statement + .assemble(&ResponseHistory::default()) + .unwrap() + .expect("split resolves statically"); + let stmt = request.query().unwrap().unwrap().query().to_string(); + let params = request + .parameters() + .unwrap() + .map(|bind| { + bind.params_raw() + .iter() + .map(|param| { + std::str::from_utf8(¶m.data) + .unwrap() + .parse::() + .unwrap() + }) + .collect() + }) + .unwrap_or_default(); + InsertSplit { params, stmt } + }) + .collect() +} + +#[tokio::test] +async fn test_split_insert_with_params() { + let splits = + parse_and_split("INSERT INTO my_table (id, value) VALUES ($1, $2), ($3, $4)").await; + + assert_eq!(splits.len(), 2); + + // First tuple uses params 0 and 1 (original $1, $2) + assert_eq!(splits[0].params.as_slice(), &[1, 2]); + assert_eq!( + splits[0].stmt, + "INSERT INTO my_table (id, value) VALUES ($1, $2)" + ); + + // Second tuple uses params 2 and 3 (original $3, $4), renumbered to $1, $2 + assert_eq!(splits[1].params.as_slice(), &[3, 4]); + assert_eq!( + splits[1].stmt, + "INSERT INTO my_table (id, value) VALUES ($1, $2)" + ); +} + +#[tokio::test] +async fn test_split_insert_single_tuple_no_split() { + let splits = parse_and_split("INSERT INTO my_table (id, value) VALUES ($1, $2)").await; + + // Single tuple should not be split + assert!(splits.is_empty()); +} + +#[tokio::test] +async fn test_split_insert_literal_values() { + let splits = + parse_and_split("INSERT INTO my_table (id, value) VALUES (1, 'a'), (2, 'b')").await; + + assert_eq!(splits.len(), 2); + + // No params for literal values + assert!(splits[0].params.is_empty()); + assert_eq!( + splits[0].stmt, + "INSERT INTO my_table (id, value) VALUES (1, 'a')" + ); + + assert!(splits[1].params.is_empty()); + assert_eq!( + splits[1].stmt, + "INSERT INTO my_table (id, value) VALUES (2, 'b')" + ); +} + +#[tokio::test] +async fn test_split_insert_mixed_params_and_literals() { + let splits = + parse_and_split("INSERT INTO my_table (id, value) VALUES ($1, 'a'), ($2, 'b')").await; + + assert_eq!(splits.len(), 2); + + assert_eq!(splits[0].params.as_slice(), &[1]); + assert_eq!( + splits[0].stmt, + "INSERT INTO my_table (id, value) VALUES ($1, 'a')" + ); + + assert_eq!(splits[1].params.as_slice(), &[2]); + assert_eq!( + splits[1].stmt, + "INSERT INTO my_table (id, value) VALUES ($1, 'b')" + ); +} + +#[tokio::test] +async fn test_extract_bind_params() { + let splits = parse_and_split("INSERT INTO t (a, b) VALUES ($1, $2), ($3, $4)").await; + let bind = Bind::new_params( + "test", + &[ + Parameter::new(b"p0"), + Parameter::new(b"p1"), + Parameter::new(b"p2"), + Parameter::new(b"p3"), + ], + ); + + // First split uses params 0 and 1 + let extracted = splits[0].extract_bind_params(&bind).unwrap(); + assert_eq!(extracted.params_raw().len(), 2); + assert_eq!(extracted.params_raw()[0].data.as_ref(), b"p0"); + assert_eq!(extracted.params_raw()[1].data.as_ref(), b"p1"); + + // Second split uses params 2 and 3 + let extracted = splits[1].extract_bind_params(&bind).unwrap(); + assert_eq!(extracted.params_raw().len(), 2); + assert_eq!(extracted.params_raw()[0].data.as_ref(), b"p2"); + assert_eq!(extracted.params_raw()[1].data.as_ref(), b"p3"); +} + +#[tokio::test] +async fn test_extract_bind_params_with_format_codes() { + let splits = parse_and_split("INSERT INTO t (a, b) VALUES ($1, $2), ($3, $4)").await; + let bind = Bind::new_params_codes( + "test", + &[ + Parameter::new(b"p0"), + Parameter::new(b"p1"), + Parameter::new(b"p2"), + Parameter::new(b"p3"), + ], + &[Format::Text, Format::Binary, Format::Text, Format::Binary], + ); + + // Second split uses params 2 and 3 (Text, Binary) + let extracted = splits[1].extract_bind_params(&bind).unwrap(); + assert_eq!(extracted.params_raw().len(), 2); + assert_eq!(extracted.params_raw()[0].data.as_ref(), b"p2"); + assert_eq!(extracted.params_raw()[1].data.as_ref(), b"p3"); + assert_eq!(extracted.format_codes_raw().len(), 2); + assert_eq!(extracted.format_codes_raw()[0], Format::Text); + assert_eq!(extracted.format_codes_raw()[1], Format::Binary); +} + +#[tokio::test] +async fn test_extract_bind_params_uniform_format() { + let splits = parse_and_split("INSERT INTO t (a) VALUES ($1), ($2)").await; + let bind = Bind::new_params_codes( + "test", + &[Parameter::new(b"p0"), Parameter::new(b"p1")], + &[Format::Binary], // Uniform format + ); + + let extracted = splits[0].extract_bind_params(&bind).unwrap(); + assert_eq!(extracted.params_raw().len(), 1); + assert_eq!(extracted.format_codes_raw().len(), 1); + assert_eq!(extracted.format_codes_raw()[0], Format::Binary); +} + +#[tokio::test] +async fn test_extract_bind_params_mixed_params_and_literals() { + let splits = parse_and_split("INSERT INTO t (a, b) VALUES ($1, 'lit1'), ($2, 'lit2')").await; + let bind = Bind::new_params( + "test", + &[ + Parameter::new(b"value_for_param1"), + Parameter::new(b"value_for_param2"), + ], + ); + + assert_eq!(splits.len(), 2); + + // First split: statement uses $1 with literal, bind extracts param 0 + assert_eq!(splits[0].stmt, "INSERT INTO t (a, b) VALUES ($1, 'lit1')"); + let extracted = splits[0].extract_bind_params(&bind).unwrap(); + assert_eq!(extracted.params_raw().len(), 1); + assert_eq!(extracted.params_raw()[0].data.as_ref(), b"value_for_param1"); + + // Second split: statement uses $1 (renumbered from $2) with literal, bind extracts param 1 + assert_eq!(splits[1].stmt, "INSERT INTO t (a, b) VALUES ($1, 'lit2')"); + let extracted = splits[1].extract_bind_params(&bind).unwrap(); + assert_eq!(extracted.params_raw().len(), 1); + assert_eq!(extracted.params_raw()[0].data.as_ref(), b"value_for_param2"); +} + +#[tokio::test] +async fn test_extract_bind_params_varying_param_counts() { + // First tuple has 2 params, second tuple has 1 param and 1 literal + let splits = parse_and_split("INSERT INTO t (a, b) VALUES ($1, $2), ($3, 'literal')").await; + let bind = Bind::new_params( + "test", + &[ + Parameter::new(b"p1"), + Parameter::new(b"p2"), + Parameter::new(b"p3"), + ], + ); + + assert_eq!(splits.len(), 2); + + // First split: uses params 0 and 1 (original $1, $2) + assert_eq!(splits[0].stmt, "INSERT INTO t (a, b) VALUES ($1, $2)"); + let extracted = splits[0].extract_bind_params(&bind).unwrap(); + assert_eq!(extracted.params_raw().len(), 2); + assert_eq!(extracted.params_raw()[0].data.as_ref(), b"p1"); + assert_eq!(extracted.params_raw()[1].data.as_ref(), b"p2"); + + // Second split: uses param 2 (original $3), renumbered to $1 + assert_eq!( + splits[1].stmt, + "INSERT INTO t (a, b) VALUES ($1, 'literal')" + ); + let extracted = splits[1].extract_bind_params(&bind).unwrap(); + assert_eq!(extracted.params_raw().len(), 1); + assert_eq!(extracted.params_raw()[0].data.as_ref(), b"p3"); +} + +#[tokio::test] +async fn test_extract_bind_params_incorrect_count() { + let splits = parse_and_split("INSERT INTO t (a, b) VALUES ($1, $2), ($3, $4)").await; + let bind = Bind::new_params( + "test", + &[ + Parameter::new(b"p1"), + Parameter::new(b"p2"), + Parameter::new(b"p3"), + ], + ); + + std::assert_matches!(splits[0].extract_bind_params(&bind), Ok(_)); + std::assert_matches!( + splits[1].extract_bind_params(&bind), + Err(Error::Rewrite(RewriteError::MissingParameter(_))) + ); +} diff --git a/pgdog/src/frontend/client/query_engine/multi_step/test/update.rs b/pgdog/src/frontend/client/query_engine/multi_step/test/update.rs index 82e8e03ea..5949b2ba3 100644 --- a/pgdog/src/frontend/client/query_engine/multi_step/test/update.rs +++ b/pgdog/src/frontend/client/query_engine/multi_step/test/update.rs @@ -4,7 +4,7 @@ use crate::{ frontend::{ ClientRequest, client::{ - query_engine::{QueryEngineContext, multi_step::UpdateMulti}, + query_engine::{QueryEngineContext, multi_step::error::Error}, test::TestClient, }, }, @@ -15,8 +15,6 @@ use crate::{ }, }; -use super::super::super::Error; - const FK_CHILD_TABLE: &str = "shard_key_update_fk_child"; const SHARDED_TABLE_DDL: &str = "CREATE TABLE IF NOT EXISTS sharded ( @@ -70,10 +68,10 @@ async fn same_shard_check(request: ClientRequest) -> Result<(), Error> { client.client().client_request.extend(request.messages); let mut context = QueryEngineContext::new(&mut client.client); - let rewrite_result = client.engine.parse_and_rewrite(&mut context)?; + let (query_planner, offset_plan) = client.engine.parse_and_rewrite(&mut context).await?; client .engine - .route_query(&mut context, rewrite_result.as_ref()) + .route_query(&mut context, offset_plan.as_ref()) .await?; assert!( @@ -86,21 +84,19 @@ async fn same_shard_check(request: ClientRequest) -> Result<(), Error> { std::assert_matches!(&*client.engine.backend, Binding::Direct(..)); let ast = context.client_request.ast.clone().expect("ast was set"); - let rewrite = ast - .rewrite_plan - .sharding_key_update - .as_ref() - .expect("sharding key update to exist"); + assert!( + ast.rewrite_plan.sharding_key_update, + "sharding key update to exist" + ); - let mut update = UpdateMulti::new(&mut client.engine, rewrite); assert!( - update.is_same_shard(&context).unwrap(), + query_planner.is_none(), "query should not trigger multi-shard update" ); // Won't error out because the query goes to the same shard // as the old shard. - update.execute(&mut context).await?; + client.engine.execute(&mut context, query_planner).await?; Ok(()) } @@ -186,7 +182,7 @@ async fn test_row_same_shard_no_transaction() { let mut context = QueryEngineContext::new(&mut client.client); - let rewrite_result = client.engine.parse_and_rewrite(&mut context).unwrap(); + let (query_planner, offset_plan) = client.engine.parse_and_rewrite(&mut context).await.unwrap(); assert!( context @@ -195,19 +191,18 @@ async fn test_row_same_shard_no_transaction() { .as_ref() .expect("ast to exist") .rewrite_plan - .sharding_key_update - .is_some(), + .sharding_key_update, "sharding key update should exist on the request" ); client .engine - .route_query(&mut context, rewrite_result.as_ref()) + .route_query(&mut context, offset_plan.as_ref()) .await .unwrap(); client .engine - .execute(&mut context, rewrite_result) + .execute(&mut context, query_planner) .await .unwrap(); @@ -729,3 +724,369 @@ async fn test_foreign_key_on_delete_sharding_key_update() { cleanup_fk_child(&mut client).await; } + +#[tokio::test] +async fn test_move_rows_insert_error_is_reported() { + let mut client = TestClient::new_rewrites(Parameters::default()).await; + ensure_sharded_table(&mut client).await; + + let shard_0_id = client.random_id_for_shard(0); + let shard_1_id = client.random_id_for_shard(1); + + // We want the INSERT step to hit a duplicate key. + for id in [shard_0_id, shard_1_id] { + client + .send_simple(Query::new(format!( + "INSERT INTO sharded (id) VALUES ({}) ON CONFLICT(id) DO NOTHING", + id + ))) + .await; + client.read_until('Z').await.unwrap(); + } + + client.send_simple(Query::new("BEGIN")).await; + client.read_until('Z').await.unwrap(); + + client + .send_simple(Query::new(format!( + "UPDATE sharded SET id = {} WHERE id = {}", + shard_1_id, shard_0_id + ))) + .await; + let error = ErrorResponse::try_from(client.read().await) + .expect("expected error from failed insert step"); + assert_eq!(error.code, "23505", "{error:?}"); + expect_message!(client.read().await, ReadyForQuery); + + // Connection still good (no transaction issues) + client.send_simple(Query::new("ROLLBACK")).await; + client.read_until('Z').await.unwrap(); +} + +/// Note: this is a case where we have to use a `RowDescription` over the cache. +#[tokio::test] +async fn test_move_rows_after_table_recreated() { + // Schema cache is stale (maybe due to a migration or something else) + let mut client = TestClient::new_rewrites(Parameters::default()) + .await + .without_schema_reload(); + + // Recreate the table with fewer columns than the schema pgdog loaded. + for ddl in [ + "DROP TABLE IF EXISTS sharded CASCADE", + "CREATE TABLE sharded (id BIGINT PRIMARY KEY, value TEXT)", + ] { + client.send(Query::new(ddl)).await; + client.try_process().await.unwrap(); + client.read_until('Z').await.unwrap(); + } + + let shard_0_id = client.random_id_for_shard(0); + let shard_1_id = client.random_id_for_shard(1); + + client + .send_simple(Query::new(format!( + "INSERT INTO sharded (id, value) VALUES ({}, 'test')", + shard_0_id + ))) + .await; + client.read_until('Z').await.unwrap(); + + client.send_simple(Query::new("BEGIN")).await; + client.read_until('Z').await.unwrap(); + + client + .send_simple(Query::new(format!( + "UPDATE sharded SET id = {} WHERE id = {}", + shard_1_id, shard_0_id + ))) + .await; + let update_reply = client.read_until('Z').await.unwrap(); + + client + .send_simple(Query::new(format!( + "SELECT id FROM sharded WHERE id = {}", + shard_1_id + ))) + .await; + let select_reply = client.read_until('Z').await.unwrap(); + + client.send_simple(Query::new("ROLLBACK")).await; + client.read_until('Z').await.unwrap(); + + // Restore the table before asserting so other tests aren't affected. + client + .send(Query::new("DROP TABLE IF EXISTS sharded")) + .await; + client.try_process().await.unwrap(); + client.read_until('Z').await.unwrap(); + ensure_sharded_table(&mut client).await; + + let cc = update_reply + .iter() + .find(|message| message.code() == 'C') + .cloned() + .unwrap_or_else(|| panic!("expected UPDATE to succeed, got: {update_reply:?}")); + assert_eq!(CommandComplete::try_from(cc).unwrap().command(), "UPDATE 1"); + assert_eq!( + select_reply + .iter() + .filter(|message| message.code() == 'D') + .count(), + 1, + "row should exist on the new shard: {select_reply:?}" + ); +} + +#[tokio::test] +async fn test_move_rows_sqlx_flow() { + let client = TestClient::new_rewrites(Parameters::default()).await; + sqlx_flow(client).await; +} + +#[tokio::test] +async fn test_move_rows_sqlx_flow_two_pc() { + let client = TestClient::new_rewrites(Parameters::default()) + .await + .with_two_pc(); + sqlx_flow(client).await; +} + +async fn sqlx_flow(mut client: TestClient) { + ensure_sharded_table(&mut client).await; + + let shard_0_id = client.random_id_for_shard(0); + let shard_1_id = client.random_id_for_shard(1); + + client + .send_simple(Query::new(format!( + "INSERT INTO sharded (id, value) VALUES ({}, 'test')", + shard_0_id + ))) + .await; + client.read_until('Z').await.unwrap(); + + client + .send(Parse::named( + "sqlx_s_1", + "UPDATE sharded SET id = $2 WHERE id = $1", + )) + .await; + client.send(Describe::new_statement("sqlx_s_1")).await; + client.send(Sync).await; + client.try_process().await.unwrap(); + client.read_until('Z').await.unwrap(); + + let bind = || { + Bind::new_params_codes_results( + "sqlx_s_1", + &[ + Parameter::new(&shard_0_id.to_be_bytes()), + Parameter::new(&shard_1_id.to_be_bytes()), + ], + &[Format::Binary, Format::Binary], + &[1], + ) + }; + + // Connection stays usable. + client.send(bind()).await; + client.send(Execute::new()).await; + client.send(Sync).await; + client.try_process().await.unwrap(); + let error = ErrorResponse::try_from(client.read().await).expect("expected error"); + assert_eq!( + error.message, + "sharding key update must be executed inside a transaction" + ); + expect_message!(client.read().await, ReadyForQuery); + + client.send_simple(Query::new("BEGIN")).await; + client.read_until('Z').await.unwrap(); + + // Same cached statement. Bind/Execute/Sync only. + client.send(bind()).await; + client.send(Execute::new()).await; + client.send(Sync).await; + client.try_process().await.unwrap(); + let reply = client.read_until('Z').await.unwrap(); + let cc = reply + .iter() + .find(|message| message.code() == 'C') + .cloned() + .unwrap_or_else(|| panic!("expected CommandComplete, got: {reply:?}")); + assert_eq!(CommandComplete::try_from(cc).unwrap().command(), "UPDATE 1"); + + for query in [ + "SELECT id FROM sharded WHERE id = {}", + "COMMIT", + "SELECT id FROM sharded WHERE id = {}", + ] { + client + .send_simple(Query::new(query.replace("{}", &shard_1_id.to_string()))) + .await; + let reply = client.read_until('Z').await.unwrap(); + if query != "COMMIT" { + assert_eq!( + reply.iter().filter(|message| message.code() == 'D').count(), + 1, + "row should be on the new shard: {reply:?}" + ); + } + } + + // Cleanup + client + .send_simple(Query::new(format!( + "DELETE FROM sharded WHERE id IN ({}, {})", + shard_0_id, shard_1_id + ))) + .await; + client.read_until('Z').await.unwrap(); +} + +#[tokio::test] +async fn test_move_rows_zero_rows() { + let mut client = TestClient::new_rewrites(Parameters::default()).await; + ensure_sharded_table(&mut client).await; + + let shard_0_id = client.random_id_for_shard(0); + let shard_1_id = client.random_id_for_shard(1); + + // No row with shard_0_id exists. + client + .send_simple(Query::new(format!( + "DELETE FROM sharded WHERE id = {}", + shard_0_id + ))) + .await; + client.read_until('Z').await.unwrap(); + + client.send_simple(Query::new("BEGIN")).await; + client.read_until('Z').await.unwrap(); + + client + .send_simple(Query::new(format!( + "UPDATE sharded SET id = {} WHERE id = {}", + shard_1_id, shard_0_id + ))) + .await; + let cc = client.read().await; + expect_message!(cc.clone(), CommandComplete); + assert_eq!(CommandComplete::try_from(cc).unwrap().command(), "UPDATE 0"); + expect_message!(client.read().await, ReadyForQuery); + + // Transaction still healthy. + client.send_simple(Query::new("COMMIT")).await; + client.read_until('Z').await.unwrap(); + + client + .send_simple(Query::new(format!( + "SELECT id FROM sharded WHERE id = {}", + shard_1_id + ))) + .await; + let reply = client.read_until('Z').await.unwrap(); + assert_eq!( + reply.iter().filter(|message| message.code() == 'D').count(), + 0, + "no row should have been created: {reply:?}" + ); +} + +#[tokio::test] +async fn test_shard_key_update_disabled_does_not_execute() { + let mut client = TestClient::new_rewrites(Parameters::default()) + .await + .with_shard_key_error(); + ensure_sharded_table(&mut client).await; + + let shard_0_id = client.random_id_for_shard(0); + let shard_1_id = client.random_id_for_shard(1); + + client + .send_simple(Query::new(format!( + "INSERT INTO sharded (id) VALUES ({}) ON CONFLICT(id) DO NOTHING", + shard_0_id + ))) + .await; + client.read_until('Z').await.unwrap(); + + client + .send_simple(Query::new(format!( + "UPDATE sharded SET id = {} WHERE id = {}", + shard_1_id, shard_0_id + ))) + .await; + let error = ErrorResponse::try_from(client.read().await).expect("expected error"); + assert_eq!( + error.message, "sharding key updates are forbidden", + "{error:?}" + ); + expect_message!(client.read().await, ReadyForQuery); + + // Row untouched. + client + .send_simple(Query::new(format!( + "SELECT id FROM sharded WHERE id = {}", + shard_0_id + ))) + .await; + let reply = client.read_until('Z').await.unwrap(); + assert_eq!( + reply.iter().filter(|message| message.code() == 'D').count(), + 1, + "row must still exist under its original id: {reply:?}" + ); +} + +#[tokio::test] +async fn test_move_rows_error_then_commit() { + let mut client = TestClient::new_rewrites(Parameters::default()).await; + ensure_sharded_table(&mut client).await; + + let shard_0_id = client.random_id_for_shard(0); + let shard_1_id = client.random_id_for_shard(1); + + // INSERT step hits a duplicate key. + for id in [shard_0_id, shard_1_id] { + client + .send_simple(Query::new(format!( + "INSERT INTO sharded (id) VALUES ({}) ON CONFLICT(id) DO NOTHING", + id + ))) + .await; + client.read_until('Z').await.unwrap(); + } + + client.send_simple(Query::new("BEGIN")).await; + client.read_until('Z').await.unwrap(); + + client + .send_simple(Query::new(format!( + "UPDATE sharded SET id = {} WHERE id = {}", + shard_1_id, shard_0_id + ))) + .await; + let error = ErrorResponse::try_from(client.read().await).expect("expected error"); + assert_eq!(error.code, "23505", "{error:?}"); + expect_message!(client.read().await, ReadyForQuery); + + // The transaction failed mid-plan (DELETE ran, INSERT errored). + // COMMIT must roll back. + client.send_simple(Query::new("COMMIT")).await; + client.read_until('Z').await.unwrap(); + + client + .send_simple(Query::new(format!( + "SELECT id FROM sharded WHERE id = {}", + shard_0_id + ))) + .await; + let reply = client.read_until('Z').await.unwrap(); + assert_eq!( + reply.iter().filter(|message| message.code() == 'D').count(), + 1, + "the failed move must not commit its DELETE: {reply:?}" + ); +} diff --git a/pgdog/src/frontend/client/query_engine/multi_step/types.rs b/pgdog/src/frontend/client/query_engine/multi_step/types.rs new file mode 100644 index 000000000..5a7238b81 --- /dev/null +++ b/pgdog/src/frontend/client/query_engine/multi_step/types.rs @@ -0,0 +1,121 @@ +use crate::frontend::client::query_engine::QueryEngineContext; +use crate::frontend::client::query_engine::multi_step::error::Error; +use crate::frontend::router::{Ast, Route}; +use crate::net::{Bind, CommandComplete, DataRow, Message, Parse, RowDescription}; +use dyn_clone::DynClone; +use std::fmt::Debug; + +/// Responses saved from a completed `Step` +#[derive(Debug, Clone, Default)] +pub(crate) struct StepResponses { + pub(crate) key: Option, + /// Need this for getting a fresh look at the table. Avoids cache problems. + pub(crate) row_description: Option, + pub(crate) parameter_description: Option, + pub(crate) rows: Vec, + pub(crate) command_complete: Option, +} + +/// What key we save under to reference later if we need to reference a certain `Step`'s +/// `StepResponses` +#[derive(Debug, Clone, Copy, PartialEq)] +pub(crate) enum SaveKey { + ShardingKeyUpdateDelete, + ShardingKeyUpdateInsert, + //TODO: Dynamic(String); when we have subqueries; we must dynamically generate names. +} + +/// Previously completed `Step` responses +#[derive(Debug, Clone, Default)] +pub(crate) struct ResponseHistory { + steps: Vec, +} + +/// We need to preserve all responses instead of just the ones we plan to look up, +/// for example, for purposes of constructing a Response based on `CommandComplete`s, +/// which is why I opted for this structure over a `HashMap` +impl ResponseHistory { + pub(crate) fn push(&mut self, responses: StepResponses) { + self.steps.push(responses); + } + + pub(crate) fn get(&self, key: SaveKey) -> Option<&StepResponses> { + self.steps.iter().find(|step| step.key == Some(key)) + } + + pub(crate) fn steps(&self) -> &[StepResponses] { + &self.steps + } +} + +/// The caller determines themselves what planning approach we take (based on parser checks) +#[derive(Debug, Clone)] +pub(crate) enum QueryPlannerType { + InsertSplit, + ShardingKeyUpdate, + /// Runs the `ClientRequest` as one step (as normal); forwards all Responses. + Normal, +} + +/// Return this to the caller after [`QueryPlanner::plan_query`] is called, for execution later on. +/// It represents everything that should be needed to fully execute the flow of a normal request +#[derive(Debug, Clone)] +pub(crate) struct QueryPlanner { + pub(crate) steps: Vec, + /// This runs at the conclusion of `steps` assuming no errors or skips + /// for how we should aggregate a Response to the Client. + pub(crate) forward_to_client: Option>, +} +#[derive(Debug, Clone)] +pub(crate) struct Step { + /// The key that the `Step` responses are saved under in `ResponseHistory` + pub(crate) save_key: Option, + /// Statically contains or dynamically constructs the `ClientRequest` + pub(crate) request: StepRequest, +} + +/// `ClientRequest` a `Step` resolves to. +/// Assembled at execution time so we can (if we'd like) dynamically resolve +/// from prior `Step` responses. +#[derive(Debug, Clone)] +pub(crate) enum StepRequest { + /// The client's own `ClientRequest` as-is. + Raw, + + // We're dynamically putting something together. + Statement(Box), +} + +/// A single statement pgdog constructed for a `Step` +#[derive(Debug, Clone)] +pub(crate) struct StatementRequest { + pub(crate) source: Box, + pub(crate) protocol: StepProtocol, + pub(crate) route: Route, + pub(crate) ast: Option, +} + +#[derive(Debug, Clone, Copy)] +pub(crate) enum StepProtocol { + Simple, + Extended, +} + +/// Produces the statement a `Step` executes; resolved against prior `Step` responses. +/// We use this in instances where we don't know the Parse/Bind upfront (due to dependency issues) +pub(crate) trait StatementSource: Debug + DynClone + Send + Sync { + fn resolve(&self, map: &ResponseHistory) -> Result, Error>; +} + +/// After the conclusion of all `steps`, look at the responses (through `map`), +/// and determine what we should pretend the Server sent back (for the Client) +/// based on looking at all of them in aggregate +// TODO: I think we discussed renaming this; forgot what was suggested. +pub(crate) trait ForwardToClient: Debug + DynClone + Send + Sync { + fn forward_to_client(&self, context: &QueryEngineContext, map: ResponseHistory) + -> Vec; +} + +// Derive Clone on all traits +dyn_clone::clone_trait_object!(ForwardToClient); +dyn_clone::clone_trait_object!(StatementSource); diff --git a/pgdog/src/frontend/client/query_engine/multi_step/update.rs b/pgdog/src/frontend/client/query_engine/multi_step/update.rs deleted file mode 100644 index 78b233c5e..000000000 --- a/pgdog/src/frontend/client/query_engine/multi_step/update.rs +++ /dev/null @@ -1,301 +0,0 @@ -use pgdog_config::RewriteMode; -use tracing::debug; - -use crate::{ - frontend::{ - ClientRequest, Command, Router, RouterContext, - client::query_engine::{QueryEngine, QueryEngineContext}, - router::parser::rewrite::statement::ShardingKeyUpdate, - }, - net::{CommandComplete, DataRow, ErrorResponse, Protocol, ReadyForQuery, RowDescription}, -}; - -use super::{Error, ForwardCheck, UpdateError}; - -#[derive(Debug, Clone, Default)] -pub(super) struct Row { - data_row: DataRow, - row_description: RowDescription, -} - -#[derive(Debug)] -pub(crate) struct UpdateMulti<'a> { - pub(super) rewrite: &'a ShardingKeyUpdate, - pub(super) engine: &'a mut QueryEngine, -} - -impl<'a> UpdateMulti<'a> { - /// Create new sharding key update handler. - pub(crate) fn new(engine: &'a mut QueryEngine, rewrite: &'a ShardingKeyUpdate) -> Self { - Self { rewrite, engine } - } - - /// Execute sharding key update, if needed. - pub(crate) async fn execute( - &mut self, - context: &mut QueryEngineContext<'_>, - ) -> Result<(), Error> { - match self.execute_internal(context).await { - Ok(()) => Ok(()), - Err(err) => { - // These are recoverable with a ROLLBACK. - if matches!(err, Error::Update(_) | Error::Execution(_)) { - self.engine - .error_response(context, ErrorResponse::from_err(&err)) - .await?; - Ok(()) - } else { - // These are bad, disconnecting the client. - Err(err) - } - } - } - } - - /// Execute sharding key update, if needed. - pub(super) async fn execute_internal( - &mut self, - context: &mut QueryEngineContext<'_>, - ) -> Result<(), Error> { - let mut check = self.rewrite.check.build_request(context.client_request)?; - self.route(&mut check, context)?; - - // The new row is on the same shard as the old row - // and we know this from the statement itself, e.g. - // - // UPDATE my_table SET shard_key = $1 WHERE shard_key = $2 - // - // This is very likely if the number of shards is low or - // you're using an ORM that puts all record columns - // into the SET clause. - // - if self.is_same_shard(context)? { - // Serve original request as-is. - debug!("[update] row is on the same shard"); - self.execute_original(context).await?; - - return Ok(()); - } - - if self.move_row(context).await?.is_none() { - // This happens, but the UPDATE's WHERE clause - // doesn't match any rows, so this whole thing is a no-op. - self.engine - .fake_command_response(context, "UPDATE 0", None) - .await?; - } - - Ok(()) - } - - /// Delete the row from the original shard and move it to the new one. - /// - /// Returns `None` if no row was returned by the DELETE query. - pub(super) async fn move_row( - &mut self, - context: &mut QueryEngineContext<'_>, - ) -> Result, Error> { - if !context.in_transaction() || !self.engine.backend.is_multishard() - // Do this check at the last possible moment. - // Just in case we change how transactions are - // routed in the future. - { - self.engine.cleanup_backend(context)?; - return Err(UpdateError::TransactionRequired.into()); - } - - if self.has_destructive_on_delete_reference(context)? { - return Err(UpdateError::ForeignKeyOnDelete.into()); - } - - let Some(row) = self.delete_and_fetch_row(context).await? else { - return Ok(None); - }; - - let mut request = self.rewrite.build_insert_request( - context.client_request, - &row.row_description, - &row.data_row, - )?; - self.route(&mut request, context)?; - - debug!("[update] executing multi-shard insert/delete"); - - // Check if we are allowed to do this operation by the config. - if self.engine.backend.cluster()?.rewrite().shard_key == RewriteMode::Error { - self.engine - .error_response(context, ErrorResponse::from_err(&UpdateError::Disabled)) - .await?; - return Ok(Some(())); - } - - self.execute_request_internal(context, &mut request, self.rewrite.is_returning()) - .await?; - - self.engine - .process_server_message(context, CommandComplete::new("UPDATE 1").message()) // We only allow to update one row at a time. - .await?; - self.engine - .process_server_message( - context, - ReadyForQuery::in_transaction(context.in_transaction()).message(), - ) - .await?; - - Ok(Some(())) - } - - fn has_destructive_on_delete_reference( - &self, - context: &QueryEngineContext<'_>, - ) -> Result { - let cluster = self.engine.backend.cluster()?; - let schema = cluster.schema(); - let table = self.rewrite.target_table(); - - let Some(relation) = schema.table(table, cluster.user(), context.params.search_path()) - else { - return Ok(false); - }; - let Some(sharded_table) = self.rewrite.sharded_table(cluster.sharded_tables()) else { - return Ok(false); - }; - - Ok(schema.has_destructive_on_delete_reference( - relation.schema(), - &relation.name, - &sharded_table.column, - )) - } - - /// Execute request and return messages to the client if forward_reply is true. - async fn execute_request_internal( - &mut self, - context: &mut QueryEngineContext<'_>, - request: &mut ClientRequest, - forward_reply: bool, - ) -> Result<(), Error> { - self.engine - .backend - .handle_client_request(request, &mut Router::default(), false) - .await?; - - let mut checker = ForwardCheck::new(context.client_request); - - while self.engine.backend.has_more_messages() { - let message = self.engine.read_server_message().await?; - let code = message.code(); - - if code == 'E' { - return Err(ErrorResponse::try_from(message)?.into()); - } - - if forward_reply && checker.forward(code) { - self.engine.process_server_message(context, message).await?; - } - } - - Ok(()) - } - - async fn execute_original( - &mut self, - context: &mut QueryEngineContext<'_>, - ) -> Result<(), Error> { - // Serve original request as-is. - self.engine - .backend - .handle_client_request( - context.client_request, - &mut self.engine.router, - self.engine.streaming, - ) - .await?; - - while self.engine.backend.has_more_messages() { - let message = self.engine.read_server_message().await?; - self.engine.process_server_message(context, message).await?; - } - - Ok(()) - } - - pub(super) async fn delete_and_fetch_row( - &mut self, - context: &mut QueryEngineContext<'_>, - ) -> Result, Error> { - let mut request = self.rewrite.delete.build_request(context.client_request)?; - self.route(&mut request, context)?; - - self.engine - .backend - .handle_client_request(&request, &mut Router::default(), false) - .await?; - - let mut row = Row::default(); - let mut rows = 0; - - while self.engine.backend.has_more_messages() { - let message = self.engine.read_server_message().await?; - match message.code() { - 'D' => { - row.data_row = DataRow::try_from(message)?; - rows += 1; - } - 'T' => row.row_description = RowDescription::try_from(message)?, - 'E' => return Err(ErrorResponse::try_from(message)?.into()), - _ => (), - } - } - - match rows { - 0 => return Ok(None), - 1 => (), - n => return Err(UpdateError::TooManyRows(n).into()), - } - - Ok(Some(row)) - } - - /// Returns true if the new sharding key resides on the same shard - /// as the old sharding key. - /// - /// This is an optimization to avoid doing a multi-shard UPDATE when - /// we don't have to. - pub(super) fn is_same_shard(&self, context: &QueryEngineContext<'_>) -> Result { - let mut check = self.rewrite.check.build_request(context.client_request)?; - self.route(&mut check, context)?; - - let new_shard = check.route().shard(); - let old_shard = context.client_request.route().shard(); - - // The sharding key isn't actually being changed - // or it maps to the same shard as before. - Ok(new_shard == old_shard) - } - - fn route( - &self, - request: &mut ClientRequest, - context: &QueryEngineContext<'_>, - ) -> Result<(), Error> { - let cluster = self.engine.backend.cluster()?; - - let context = RouterContext::new( - request, - cluster, - context.params, - context.transaction(), - context.sticky, - )?; - let mut router = Router::new(); - let command = router.query(context)?; - if let Command::Query(route) = command { - request.route = Some(route.clone()); - } else { - return Err(UpdateError::NoRoute.into()); - } - - Ok(()) - } -} diff --git a/pgdog/src/frontend/client/query_engine/query.rs b/pgdog/src/frontend/client/query_engine/query.rs index eb7e8d275..2784599e4 100644 --- a/pgdog/src/frontend/client/query_engine/query.rs +++ b/pgdog/src/frontend/client/query_engine/query.rs @@ -1,10 +1,7 @@ use tracing::{info, trace}; use crate::{ - frontend::{ - client::TransactionType, - router::parser::{explain_trace::ExplainTrace, rewrite::statement::plan::RewriteResult}, - }, + frontend::{client::TransactionType, router::parser::explain_trace::ExplainTrace}, net::{ DataRow, FromBytes, Message, Protocol, ProtocolMessage, Query, ReadyForQuery, RowDescription, ToBytes, TransactionState, @@ -13,17 +10,17 @@ use crate::{ util::safe_timeout, }; -use tracing::{debug, error}; - use super::hooks::schema::schema_changed; use super::*; +use crate::frontend::client::query_engine::multi_step::types::QueryPlanner; +use tracing::{debug, error}; impl QueryEngine { /// Handle query from client. pub(super) async fn execute( &mut self, context: &mut QueryEngineContext<'_>, - query_planner: Option, + query_planner: Option, ) -> Result<(), Error> { // Check that we're not in a transaction error state. if !self.transaction_error_check(context).await? { @@ -65,14 +62,16 @@ impl QueryEngine { } } + let planner = query_planner.unwrap_or_else(QueryPlanner::plan_normal); + let query_timeout = context.timeouts.query_timeout(&State::Active); - let result = safe_timeout( - query_timeout, - self.client_server_exchange(context, query_planner), - ) - .await; + let result = safe_timeout(query_timeout, self.run_steps(context, &planner)).await; match result { + Ok(Err(Error::Planner(err))) => match err.into_client_error() { + Ok(response) => self.error_response(context, response).await?, + Err(err) => return Err(err.into()), + }, Ok(response) => response?, Err(err) => { // Close the conn, it could be stuck executing a query @@ -85,40 +84,6 @@ impl QueryEngine { Ok(()) } - async fn client_server_exchange( - &mut self, - context: &mut QueryEngineContext<'_>, - rewrite_result: Option, - ) -> Result<(), Error> { - match rewrite_result { - Some(RewriteResult::InsertSplit(requests)) => { - Box::pin(multi_step::InsertMulti::from_engine(self, requests).execute(context)) - .await?; - } - - Some(RewriteResult::InPlace { .. }) | None => { - self.backend - .handle_client_request(context.client_request, &mut self.router, self.streaming) - .await?; - - while self.backend.has_more_messages() - && !self.backend.in_copy_mode() - && !self.streaming - { - let message = self.read_server_message().await?; - self.process_server_message(context, message).await?; - } - } - - Some(RewriteResult::ShardingKeyUpdate(sharding_key_update)) => { - Box::pin(multi_step::UpdateMulti::new(self, &sharding_key_update).execute(context)) - .await?; - } - } - - Ok(()) - } - pub(crate) async fn read_server_message(&mut self) -> Result { Ok(self.backend.read().await?) } @@ -178,14 +143,7 @@ impl QueryEngine { match state { TransactionState::Error => { - let error_state = match context.transaction { - Some(TransactionType::ReadOnly) => Some(TransactionType::ErrorReadOnly), - Some(TransactionType::ReadWrite | TransactionType::Implicit) => { - Some(TransactionType::ErrorReadWrite) - } - _ => None, - }; - context.transaction = error_state; + context.set_transaction_error(); if self.two_pc.auto() { self.end_two_pc(true).await?; // TODO: this records a 2pc transaction in client diff --git a/pgdog/src/frontend/client/query_engine/rewrite.rs b/pgdog/src/frontend/client/query_engine/rewrite.rs index 129064f66..cb051f851 100644 --- a/pgdog/src/frontend/client/query_engine/rewrite.rs +++ b/pgdog/src/frontend/client/query_engine/rewrite.rs @@ -1,5 +1,6 @@ use super::*; -use crate::frontend::router::parser::rewrite::statement::plan::RewriteResult; +use crate::frontend::client::query_engine::multi_step::types::QueryPlanner; +use crate::frontend::router::parser::rewrite::statement::offset::OffsetPlan; use crate::frontend::router::parser::{AstContext, Cache}; impl QueryEngine { @@ -20,10 +21,10 @@ impl QueryEngine { } /// Parse client request and rewrite it, if necessary. - pub(super) fn parse_and_rewrite( + pub(super) async fn parse_and_rewrite( &mut self, context: &mut QueryEngineContext<'_>, - ) -> Result, Error> { + ) -> Result<(Option, Option), Error> { let use_parser = self .backend .cluster() @@ -31,7 +32,7 @@ impl QueryEngine { .unwrap_or(false); if !use_parser { - return Ok(None); + return Ok((None, None)); } let query = context.client_request.query()?; @@ -40,11 +41,12 @@ impl QueryEngine { let ast_ctx = AstContext::from_cluster(cluster, context.params); let ast = Cache::get().query(&query, &ast_ctx, context.prepared_statements)?; - let rewrite_result = ast.rewrite_plan.apply(context.client_request)?; + let rewrite_plan = ast.rewrite_plan.clone(); context.client_request.ast = Some(ast); - Ok(Some(rewrite_result)) + let rewrite_result = rewrite_plan.apply(self, context).await?; + Ok((rewrite_result, rewrite_plan.offset)) } else { - Ok(None) + Ok((None, None)) } } } diff --git a/pgdog/src/frontend/client/query_engine/route_query.rs b/pgdog/src/frontend/client/query_engine/route_query.rs index a8f563a6f..658ab29e0 100644 --- a/pgdog/src/frontend/client/query_engine/route_query.rs +++ b/pgdog/src/frontend/client/query_engine/route_query.rs @@ -3,7 +3,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::offset::OffsetPlan; use crate::frontend::router::sharding::lookup; use crate::util::safe_timeout; @@ -66,7 +66,7 @@ impl QueryEngine { pub(super) async fn route_query( &mut self, context: &mut QueryEngineContext<'_>, - rewrite_result: Option<&RewriteResult>, + offset_plan: Option<&OffsetPlan>, ) -> Result { // Check that we can route this transaction at all. if self.backend.pooler_mode() == PoolerMode::Statement && context.client_request.is_begin() @@ -141,15 +141,21 @@ impl QueryEngine { match result { Ok(()) => { let command = self.router.command(); - context.client_request.route = Some(command.route().clone()); + // TODO: This relies on an implicit construct... + // `route` is `None` at the start of every request + // If it's `Some` at this point in the execution flow ,`QueryPlanner` pinned it already (e.g. same-shard multi-tuple INSERT collapsed). + // Maybe add a flag on `ClientRequest` to make this more clear/explicit. + if context.client_request.route.is_none() { + context.client_request.route = Some(command.route().clone()); + } trace!( "routing {:#?} to {:#?}", context.client_request.messages, command, ); // Apply post-parser rewrites, e.g. offset/limit. - if let Some(rewrite_result) = rewrite_result { - rewrite_result.apply_after_parser(context.client_request)?; + if let Some(offset_plan) = offset_plan { + offset_plan.apply_after_parser(context.client_request)?; } // Only validate shard placement for requests that actually execute diff --git a/pgdog/src/frontend/client/query_engine/test/rewrite_insert_split.rs b/pgdog/src/frontend/client/query_engine/test/rewrite_insert_split.rs index dad551131..488d10e3b 100644 --- a/pgdog/src/frontend/client/query_engine/test/rewrite_insert_split.rs +++ b/pgdog/src/frontend/client/query_engine/test/rewrite_insert_split.rs @@ -1,4 +1,4 @@ -use crate::frontend::router::parser::rewrite::statement::plan::RewriteResult; +use crate::frontend::client::query_engine::multi_step::types::{ResponseHistory, StepRequest}; use super::prelude::*; @@ -13,17 +13,27 @@ async fn run_test(messages: Vec) -> Vec { let mut context = QueryEngineContext::new(&mut client); engine.rewrite_extended(&mut context).unwrap(); - let rewrite_result = engine.parse_and_rewrite(&mut context).unwrap(); + let (planner, _offset) = engine.parse_and_rewrite(&mut context).await.unwrap(); + let planner = planner.expect("expected rewrite insert split"); assert!( - matches!(rewrite_result, Some(RewriteResult::InsertSplit(_))), + planner.forward_to_client.is_some(), "expected rewrite insert split" ); - match rewrite_result.unwrap() { - RewriteResult::InsertSplit(requests) => requests, - _ => unreachable!(), - } + planner + .steps + .into_iter() + .map(|step| { + let StepRequest::Statement(statement) = step.request else { + unreachable!("split should not be raw") + }; + statement + .assemble(&ResponseHistory::default()) + .unwrap() + .expect("split resolves statically") + }) + .collect() } #[tokio::test] @@ -60,7 +70,7 @@ async fn test_insert_split() { matches!(request[0].clone(), ProtocolMessage::Parse(parse) if parse.query() == "INSERT INTO test (id, email) VALUES ($1, $2)" && parse.anonymous()), "expected single tuple insert with no name" ); - match request[1].clone() { + match request[2].clone() { ProtocolMessage::Bind(bind) => { assert_eq!(bind.params_raw().first().unwrap().data, id); assert_eq!(bind.params_raw().get(1).unwrap().data, email); @@ -99,7 +109,7 @@ async fn test_insert_split_prepared() { matches!(request[0].clone(), ProtocolMessage::Parse(parse) if parse.query() == "INSERT INTO test (id, email) VALUES ($1, $2)" && parse.name() == "__pgdog_2"), "expected single tuple insert" ); - match request[1].clone() { + match request[2].clone() { ProtocolMessage::Bind(bind) => { assert_eq!(bind.params_raw().first().unwrap().data, id); assert_eq!(bind.params_raw().get(1).unwrap().data, email); @@ -147,7 +157,7 @@ async fn test_insert_split_not_sharded() { ]); let mut engine = QueryEngine::from_client(&client).unwrap(); let mut context = QueryEngineContext::new(&mut client); - let rewrite_result = engine.parse_and_rewrite(&mut context).unwrap(); + let (planner, _offset) = engine.parse_and_rewrite(&mut context).await.unwrap(); - assert!(rewrite_result.is_none()); + assert!(planner.is_none()); } 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 5eccc828a..bd380dedd 100644 --- a/pgdog/src/frontend/client/query_engine/test/rewrite_offset.rs +++ b/pgdog/src/frontend/client/query_engine/test/rewrite_offset.rs @@ -1,7 +1,5 @@ use crate::frontend::router::parser::Limit; -use crate::frontend::router::parser::rewrite::statement::{ - offset::OffsetPlan, plan::RewriteResult, -}; +use crate::frontend::router::parser::rewrite::statement::offset::OffsetPlan; use crate::frontend::router::parser::route::{Route, Shard, ShardWithPriority}; use super::prelude::*; @@ -15,12 +13,9 @@ async fn run_test(messages: Vec) -> Option { let mut engine = QueryEngine::from_client(&client).unwrap(); let mut context = QueryEngineContext::new(&mut client); - let rewrite_result = engine.parse_and_rewrite(&mut context).unwrap(); + let (_planner, offset) = engine.parse_and_rewrite(&mut context).await.unwrap(); - match rewrite_result { - Some(RewriteResult::InPlace { offset }) => offset, - other => panic!("expected InPlace, got {:?}", other), - } + offset } fn cross_shard_route() -> Route { @@ -102,9 +97,10 @@ async fn test_offset_limit_not_sharded() { let mut engine = QueryEngine::from_client(&client).unwrap(); let mut context = QueryEngineContext::new(&mut client); - let rewrite_result = engine.parse_and_rewrite(&mut context).unwrap(); + let (planner, offset) = engine.parse_and_rewrite(&mut context).await.unwrap(); - assert!(rewrite_result.is_none()); + assert!(planner.is_none()); + assert!(offset.is_none()); } #[tokio::test] @@ -127,7 +123,7 @@ async fn test_offset_with_unique_id_simple() { let mut engine = QueryEngine::from_client(&client).unwrap(); let mut context = QueryEngineContext::new(&mut client); - let rewrite_result = engine.parse_and_rewrite(&mut context).unwrap(); + let rewrite_result = engine.parse_and_rewrite(&mut context).await.unwrap(); // After parse_and_rewrite, the Query message should have unique_id replaced. let rewritten_sql = match &context.client_request.messages[0] { @@ -146,6 +142,7 @@ async fn test_offset_with_unique_id_simple() { // apply_after_parser with a cross-shard route. context.client_request.route = Some(cross_shard_route()); rewrite_result + .1 .as_ref() .unwrap() .apply_after_parser(context.client_request) @@ -198,7 +195,7 @@ async fn test_offset_with_unique_id_extended() { let mut engine = QueryEngine::from_client(&client).unwrap(); let mut context = QueryEngineContext::new(&mut client); - let rewrite_result = engine.parse_and_rewrite(&mut context).unwrap(); + let rewrite_result = engine.parse_and_rewrite(&mut context).await.unwrap(); // After parse_and_rewrite, Parse should have unique_id rewritten to $4::bigint. let rewritten_sql = match &context.client_request.messages[0] { @@ -213,6 +210,7 @@ async fn test_offset_with_unique_id_extended() { // apply_after_parser with cross-shard route should only rewrite Bind params. context.client_request.route = Some(cross_shard_route()); rewrite_result + .1 .as_ref() .unwrap() .apply_after_parser(context.client_request) diff --git a/pgdog/src/frontend/client/query_engine/test/rewrite_simple_prepared.rs b/pgdog/src/frontend/client/query_engine/test/rewrite_simple_prepared.rs index 408922555..627a59b36 100644 --- a/pgdog/src/frontend/client/query_engine/test/rewrite_simple_prepared.rs +++ b/pgdog/src/frontend/client/query_engine/test/rewrite_simple_prepared.rs @@ -10,7 +10,8 @@ async fn run_test(client: &mut Client, messages: &[ProtocolMessage]) -> Vec Self { + let mut config = config().deref().clone(); + config.config.rewrite.shard_key = RewriteMode::Error; + set(config).unwrap(); + reload_from_existing().unwrap(); + self + } + + pub(crate) fn with_two_pc(self) -> Self { + let mut config = config().deref().clone(); + config.config.general.two_phase_commit = true; + config.config.general.two_phase_commit_auto = Some(true); + set(config).unwrap(); + reload_from_existing().unwrap(); + self + } + + pub(crate) fn without_schema_reload(self) -> Self { + let mut config = config().deref().clone(); + config.config.general.reload_schema_on_ddl = false; + set(config).unwrap(); + reload_from_existing().unwrap(); + self + } + pub(crate) fn with_full_prepared_statements(self) -> Self { let mut config = config().deref().clone(); config.config.general.prepared_statements = pgdog_config::PreparedStatementsLevel::Full; diff --git a/pgdog/src/frontend/error.rs b/pgdog/src/frontend/error.rs index 92e6950a4..c168262dc 100644 --- a/pgdog/src/frontend/error.rs +++ b/pgdog/src/frontend/error.rs @@ -51,21 +51,15 @@ pub(crate) enum Error { #[error("rewrite: {0}")] Rewrite(#[from] crate::frontend::router::parser::rewrite::statement::Error), - #[error("query has no route")] - NoRoute, - - #[error("multi-tuple insert requires multi-shard binding")] - MultiShardRequired, - // FIXME: layer errors better so we don't have // to reach so deep into a module. #[error("{0}")] - Multi(#[from] Box), + Planner(#[from] Box), } impl From for Error { fn from(value: crate::frontend::client::query_engine::multi_step::error::Error) -> Self { - Self::Multi(Box::new(value)) + Self::Planner(Box::new(value)) } } diff --git a/pgdog/src/frontend/router/parser/cache/ast.rs b/pgdog/src/frontend/router/parser/cache/ast.rs index 3a94302eb..7339dce8e 100644 --- a/pgdog/src/frontend/router/parser/cache/ast.rs +++ b/pgdog/src/frontend/router/parser/cache/ast.rs @@ -84,7 +84,6 @@ impl Ast { // `Cache::query` can decide whether this entry is safe to cache. let mut rewriter = StatementRewrite::new(StatementRewriteContext { extended: query.original_query.extended(), - prepared: query.original_query.prepared(), prepared_statements, schema, db_schema, diff --git a/pgdog/src/frontend/router/parser/query/mod.rs b/pgdog/src/frontend/router/parser/query/mod.rs index cd505d0bd..131489254 100644 --- a/pgdog/src/frontend/router/parser/query/mod.rs +++ b/pgdog/src/frontend/router/parser/query/mod.rs @@ -453,7 +453,8 @@ impl QueryParser { if !context.router_context.executable && let Command::Query(ref query) = command && query.is_cross_shard() - && statement.rewrite_plan.insert_split.is_empty() + // TODO: Why are we checking insert_split here after checking if it's not executable? + && !statement.rewrite_plan.insert_split { context .shards_calculator diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/auto_id.rs b/pgdog/src/frontend/router/parser/rewrite/statement/auto_id.rs index d6d8595f2..ff2f4d878 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/auto_id.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/auto_id.rs @@ -504,7 +504,6 @@ mod tests { let ast = pg_raw_parse::parse(sql).unwrap(); let mut rewriter = StatementRewrite::new(StatementRewriteContext { extended: false, - prepared: false, prepared_statements: prepared, schema, db_schema, diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/error.rs b/pgdog/src/frontend/router/parser/rewrite/statement/error.rs index e29203f64..84dd98328 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/error.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/error.rs @@ -8,9 +8,6 @@ pub(crate) enum Error { #[error("parser: {0}")] Parser(#[from] pg_raw_parse::Error), - #[error("cache: {0}")] - Cache(String), - #[error("sharding key assignment unsupported: {0}")] UnsupportedShardingKeyUpdate(String), diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/insert.rs b/pgdog/src/frontend/router/parser/rewrite/statement/insert.rs index 61a59275a..71663f4b6 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/insert.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/insert.rs @@ -1,132 +1,10 @@ -use indexmap::IndexSet; -use pg_raw_parse::{Node, NodeMut, deparse, make, nodes, walk}; +use pg_raw_parse::{Node, nodes}; use pgdog_config::RewriteMode; -use crate::frontend::router::Ast; -use crate::frontend::router::parser::Cache; -use crate::frontend::{BufferedQuery, ClientRequest}; -use crate::net::{Bind, Parse, ProtocolMessage, Query}; - use super::{Error, RewritePlan, StatementRewrite}; -#[derive(Debug, Clone)] -pub(crate) struct InsertSplit { - /// Parameter positions in the original Bind message - /// that should be used to build the Bind message specific to this - /// insert statement. - params: IndexSet, - - /// The split up INSERT statement with parameters and/or values. - stmt: String, - - /// The statement AST. - ast: Ast, - - /// The global prepared statement name for this split. - /// Only set when the original statement was a named prepared statement. - statement_name: Option, -} - -impl InsertSplit { - /// Get the global prepared statement name, if this split was registered. - pub(crate) fn statement_name(&self) -> Option<&str> { - self.statement_name.as_deref() - } - - /// Build a ClientRequest from this split and the original request. - pub(crate) fn build_request(&self, request: &ClientRequest) -> Result { - let mut new_request = ClientRequest::default(); - let mut has_parse = false; - - for message in &request.messages { - let new_message = match message { - ProtocolMessage::Parse(parse) => { - has_parse = true; - let mut new_parse = parse.clone(); - new_parse.set_query(&self.stmt); - - if let Some(name) = self.statement_name() { - new_parse.rename(name); - } - - ProtocolMessage::Parse(new_parse) - } - ProtocolMessage::Query(query) => { - let mut new_query = query.clone(); - new_query.set_query(&self.stmt); - ProtocolMessage::Query(new_query) - } - ProtocolMessage::Bind(bind) => { - let new_bind = self.extract_bind_params(bind)?; - ProtocolMessage::Bind(new_bind) - } - other => other.clone(), - }; - new_request.messages.push(new_message); - new_request.ast = Some(self.ast.clone()); - } - - // When the driver prepared the statement in a separate round-trip - // (lib/pq: Parse/Describe/Sync, then Bind/Execute/Sync), the execute - // batch carries no Parse of its own and relies on `last_parse` being - // injected before the Bind. Rewrite that saved Parse to this split's - // single-tuple statement so the backend prepares the right one; - // otherwise it would still hold the original multi-tuple statement and - // reject the Bind's parameter count. - if !has_parse && let Some(parse) = &request.last_parse { - let mut split_parse = parse.clone(); - split_parse.set_query(&self.stmt); - if let Some(name) = self.statement_name() { - split_parse.rename(name); - } - new_request.last_parse = Some(split_parse); - } - - Ok(new_request) - } - - /// Extract specific parameters from a Bind message based on this split's param indices. - fn extract_bind_params(&self, bind: &Bind) -> Result { - let mut new = Bind::new_statement(self.statement_name().unwrap_or_default()); - for param in &self.params { - let param = bind - .parameter(*param as usize - 1)? - .ok_or(Error::MissingParameter(*param))?; - new.push_param(param.parameter().clone(), param.format()); - } - - Ok(new) - } -} - -/// Build separate ClientRequests for each insert split. -pub(crate) fn build_split_requests( - splits: &[InsertSplit], - request: &ClientRequest, -) -> Result, Error> { - splits - .iter() - .map(|split| split.build_request(request)) - .collect() -} - impl StatementRewrite<'_> { - /// Split up multi-tuple INSERT statements into separate single-tuple statements - /// for individual execution. - /// - /// # Example - /// - /// ```sql - /// INSERT INTO my_table (id, value) VALUES ($1, $2), ($3, $4) - /// ``` - /// - /// becomes - /// - /// ```sql - /// INSERT INTO my_table (id, value) VALUES ($1, $2) - /// INSERT INTO my_table (id, value) VALUES ($1, $2) -- These are copied from params $3 and $4 - /// ``` - /// + /// Check to see if we should try to use [`QueryPlannerType::MultiInsert`] pub(super) fn split_insert( &mut self, insert: &nodes::InsertStmt, @@ -137,346 +15,13 @@ impl StatementRewrite<'_> { return Ok(()); } - let mut splits = Vec::new(); - make::try_owned(|mem| { - let mut copy = mem.make_unique(insert); - - if let Node::SelectStmt(select) = insert.select_stmt() { - for list in select.values_lists() { - let (params, select) = self.build_single_tuple_select(mem, list); - copy.as_mut().set_select_stmt(select.uncast()); - splits.push((params, deparse(&*copy)?.as_str().to_string())); - } + if let Node::SelectStmt(select) = insert.select_stmt() { + let count = select.values_lists().len(); + if count >= 2 { + plan.insert_split = true; } - - Ok::<_, Error>(copy) - })?; - - if splits.len() <= 1 { - return Ok(()); - } - - // FIXME(sage): This is extremely duplicated with the work we do for - // multi-step updates. (#1178 for inserts) Both pieces of code have - // distinct sets of bugs. We should unify those two parts of the code - // base and make this behave consistently. - - // Now create Ast for each split (needs mutable borrow of prepared_statements) - let cache = Cache::get(); - let ctx = self.ast_context(); - for (params, stmt) in splits { - let query = if self.extended { - BufferedQuery::Prepared(Parse::named("", &stmt)) - } else { - BufferedQuery::Query(Query::new(&stmt)) - }; - let ast = cache - .query(&query, &ctx, self.prepared_statements) - .map_err(|e| Error::Cache(e.to_string()))?; - - // If this is a named prepared statement, register the split in the global cache - // and store the assigned name for use in Bind messages. - let statement_name = if self.prepared { - // Name will be assigned by `insert`. - let mut parse = Parse::named("", &stmt); - self.prepared_statements.insert(&mut parse); - Some(parse.name().to_owned()) - } else { - None - }; - - plan.insert_split.push(InsertSplit { - params, - stmt, - ast, - statement_name, - }); } Ok(()) } - - /// Build a single-tuple INSERT from the original statement with just one values_list. - /// Returns the parameter positions (0-indexed) and the SQL string. - fn build_single_tuple_select<'mem>( - &self, - mem: make::MemoryToken<'mem>, - values_list: Node<'_>, - ) -> (IndexSet, make::Unique<'mem, &'mem nodes::SelectStmt>) { - let mut tuple = mem.make_unique(values_list); - - let mut params = IndexSet::new(); - walk::walk_mut(tuple.as_mut(), |node| { - if let NodeMut::ParamRef(param) = node { - params.insert(param.number as _); - param.set_number(params.get_index_of(&(param.number as u16)).unwrap() as i32 + 1) - } - }); - - let mut select = mem.make_node::(); - select.as_mut().set_values_lists(mem.make_list(&[tuple])); - (params, select) - } -} - -#[cfg(test)] -mod tests { - use pgdog_config::Rewrite; - - use super::*; - use crate::backend::ShardingSchema; - use crate::backend::schema::Schema; - use crate::frontend::PreparedStatements; - use crate::frontend::router::parser::StatementRewriteContext; - use crate::net::messages::bind::{Format, Parameter}; - - fn default_db_schema() -> Schema { - Schema::default() - } - - fn default_schema() -> ShardingSchema { - ShardingSchema { - shards: 2, - rewrite: Rewrite { - enabled: true, - split_inserts: RewriteMode::Rewrite, - ..Default::default() - }, - ..Default::default() - } - } - - fn parse_and_split(sql: &str) -> Vec { - let root = pg_raw_parse::parse(sql).unwrap(); - let insert = match root.stmts().next() { - Some(Node::InsertStmt(insert)) => insert, - _ => unreachable!(), - }; - let mut prepared = PreparedStatements::default(); - let schema = default_schema(); - let db_schema = default_db_schema(); - let mut rewriter = StatementRewrite::new(StatementRewriteContext { - extended: false, - prepared: false, - prepared_statements: &mut prepared, - schema: &schema, - db_schema: &db_schema, - user: "", - search_path: None, - }); - let mut plan = RewritePlan::default(); - rewriter.split_insert(insert, &mut plan).unwrap(); - plan.insert_split - } - - #[test] - fn test_split_insert_with_params() { - let splits = parse_and_split("INSERT INTO my_table (id, value) VALUES ($1, $2), ($3, $4)"); - - assert_eq!(splits.len(), 2); - - // First tuple uses params 0 and 1 (original $1, $2) - assert_eq!(splits[0].params.as_slice(), &[1, 2]); - assert_eq!( - splits[0].stmt, - "INSERT INTO my_table (id, value) VALUES ($1, $2)" - ); - - // Second tuple uses params 2 and 3 (original $3, $4), renumbered to $1, $2 - assert_eq!(splits[1].params.as_slice(), &[3, 4]); - assert_eq!( - splits[1].stmt, - "INSERT INTO my_table (id, value) VALUES ($1, $2)" - ); - } - - #[test] - fn test_split_insert_single_tuple_no_split() { - let splits = parse_and_split("INSERT INTO my_table (id, value) VALUES ($1, $2)"); - - // Single tuple should not be split - assert!(splits.is_empty()); - } - - #[test] - fn test_split_insert_literal_values() { - let splits = parse_and_split("INSERT INTO my_table (id, value) VALUES (1, 'a'), (2, 'b')"); - - assert_eq!(splits.len(), 2); - - // No params for literal values - assert!(splits[0].params.is_empty()); - assert_eq!( - splits[0].stmt, - "INSERT INTO my_table (id, value) VALUES (1, 'a')" - ); - - assert!(splits[1].params.is_empty()); - assert_eq!( - splits[1].stmt, - "INSERT INTO my_table (id, value) VALUES (2, 'b')" - ); - } - - #[test] - fn test_split_insert_mixed_params_and_literals() { - let splits = - parse_and_split("INSERT INTO my_table (id, value) VALUES ($1, 'a'), ($2, 'b')"); - - assert_eq!(splits.len(), 2); - - assert_eq!(splits[0].params.as_slice(), &[1]); - assert_eq!( - splits[0].stmt, - "INSERT INTO my_table (id, value) VALUES ($1, 'a')" - ); - - assert_eq!(splits[1].params.as_slice(), &[2]); - assert_eq!( - splits[1].stmt, - "INSERT INTO my_table (id, value) VALUES ($1, 'b')" - ); - } - - #[test] - fn test_extract_bind_params() { - let splits = parse_and_split("INSERT INTO t (a, b) VALUES ($1, $2), ($3, $4)"); - let bind = Bind::new_params( - "test", - &[ - Parameter::new(b"p0"), - Parameter::new(b"p1"), - Parameter::new(b"p2"), - Parameter::new(b"p3"), - ], - ); - - // First split uses params 0 and 1 - let extracted = splits[0].extract_bind_params(&bind).unwrap(); - assert_eq!(extracted.params_raw().len(), 2); - assert_eq!(extracted.params_raw()[0].data.as_ref(), b"p0"); - assert_eq!(extracted.params_raw()[1].data.as_ref(), b"p1"); - - // Second split uses params 2 and 3 - let extracted = splits[1].extract_bind_params(&bind).unwrap(); - assert_eq!(extracted.params_raw().len(), 2); - assert_eq!(extracted.params_raw()[0].data.as_ref(), b"p2"); - assert_eq!(extracted.params_raw()[1].data.as_ref(), b"p3"); - } - - #[test] - fn test_extract_bind_params_with_format_codes() { - let splits = parse_and_split("INSERT INTO t (a, b) VALUES ($1, $2), ($3, $4)"); - let bind = Bind::new_params_codes( - "test", - &[ - Parameter::new(b"p0"), - Parameter::new(b"p1"), - Parameter::new(b"p2"), - Parameter::new(b"p3"), - ], - &[Format::Text, Format::Binary, Format::Text, Format::Binary], - ); - - // Second split uses params 2 and 3 (Text, Binary) - let extracted = splits[1].extract_bind_params(&bind).unwrap(); - assert_eq!(extracted.params_raw().len(), 2); - assert_eq!(extracted.params_raw()[0].data.as_ref(), b"p2"); - assert_eq!(extracted.params_raw()[1].data.as_ref(), b"p3"); - assert_eq!(extracted.format_codes_raw().len(), 2); - assert_eq!(extracted.format_codes_raw()[0], Format::Text); - assert_eq!(extracted.format_codes_raw()[1], Format::Binary); - } - - #[test] - fn test_extract_bind_params_uniform_format() { - let splits = parse_and_split("INSERT INTO t (a) VALUES ($1), ($2)"); - let bind = Bind::new_params_codes( - "test", - &[Parameter::new(b"p0"), Parameter::new(b"p1")], - &[Format::Binary], // Uniform format - ); - - let extracted = splits[0].extract_bind_params(&bind).unwrap(); - assert_eq!(extracted.params_raw().len(), 1); - assert_eq!(extracted.format_codes_raw().len(), 1); - assert_eq!(extracted.format_codes_raw()[0], Format::Binary); - } - - #[test] - fn test_extract_bind_params_mixed_params_and_literals() { - let splits = parse_and_split("INSERT INTO t (a, b) VALUES ($1, 'lit1'), ($2, 'lit2')"); - let bind = Bind::new_params( - "test", - &[ - Parameter::new(b"value_for_param1"), - Parameter::new(b"value_for_param2"), - ], - ); - - assert_eq!(splits.len(), 2); - - // First split: statement uses $1 with literal, bind extracts param 0 - assert_eq!(splits[0].stmt, "INSERT INTO t (a, b) VALUES ($1, 'lit1')"); - let extracted = splits[0].extract_bind_params(&bind).unwrap(); - assert_eq!(extracted.params_raw().len(), 1); - assert_eq!(extracted.params_raw()[0].data.as_ref(), b"value_for_param1"); - - // Second split: statement uses $1 (renumbered from $2) with literal, bind extracts param 1 - assert_eq!(splits[1].stmt, "INSERT INTO t (a, b) VALUES ($1, 'lit2')"); - let extracted = splits[1].extract_bind_params(&bind).unwrap(); - assert_eq!(extracted.params_raw().len(), 1); - assert_eq!(extracted.params_raw()[0].data.as_ref(), b"value_for_param2"); - } - - #[test] - fn test_extract_bind_params_varying_param_counts() { - // First tuple has 2 params, second tuple has 1 param and 1 literal - let splits = parse_and_split("INSERT INTO t (a, b) VALUES ($1, $2), ($3, 'literal')"); - let bind = Bind::new_params( - "test", - &[ - Parameter::new(b"p1"), - Parameter::new(b"p2"), - Parameter::new(b"p3"), - ], - ); - - assert_eq!(splits.len(), 2); - - // First split: uses params 0 and 1 (original $1, $2) - assert_eq!(splits[0].stmt, "INSERT INTO t (a, b) VALUES ($1, $2)"); - let extracted = splits[0].extract_bind_params(&bind).unwrap(); - assert_eq!(extracted.params_raw().len(), 2); - assert_eq!(extracted.params_raw()[0].data.as_ref(), b"p1"); - assert_eq!(extracted.params_raw()[1].data.as_ref(), b"p2"); - - // Second split: uses param 2 (original $3), renumbered to $1 - assert_eq!( - splits[1].stmt, - "INSERT INTO t (a, b) VALUES ($1, 'literal')" - ); - let extracted = splits[1].extract_bind_params(&bind).unwrap(); - assert_eq!(extracted.params_raw().len(), 1); - assert_eq!(extracted.params_raw()[0].data.as_ref(), b"p3"); - } - - #[test] - fn test_extract_bind_params_incorrect_count() { - let splits = parse_and_split("INSERT INTO t (a, b) VALUES ($1, $2), ($3, $4)"); - let bind = Bind::new_params( - "test", - &[ - Parameter::new(b"p1"), - Parameter::new(b"p2"), - Parameter::new(b"p3"), - ], - ); - - std::assert_matches!(splits[0].extract_bind_params(&bind), Ok(_)); - std::assert_matches!( - splits[1].extract_bind_params(&bind), - Err(Error::MissingParameter(_)) - ); - } } diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs b/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs index 22a8a47f5..8444e0441 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs @@ -5,7 +5,6 @@ use pg_raw_parse::{Node, NodeMut, make, nodes, transform, walk}; use crate::backend::ShardingSchema; use crate::backend::schema::Schema; use crate::frontend::PreparedStatements; -use crate::frontend::router::parser::AstContext; use crate::net::parameter::ParameterValue; pub(crate) mod aggregate; @@ -19,19 +18,14 @@ pub(crate) mod unique_id; pub(crate) mod update; pub(crate) use error::Error; -pub(crate) use insert::InsertSplit; pub(crate) use plan::RewritePlan; pub(crate) use simple_prepared::PrepareExecute; -pub(crate) use update::*; /// Statement rewrite engine context. #[derive(Debug)] pub(crate) struct StatementRewriteContext<'a> { /// The statement is using the extended protocol with placeholders. pub(crate) extended: bool, - /// The statement is named, so we need to save any derivatives into the global - /// statement cache. - pub(crate) prepared: bool, /// Reference to global prepared stmt cache. pub(crate) prepared_statements: &'a mut PreparedStatements, /// Sharding schema. @@ -52,9 +46,6 @@ pub(crate) struct StatementRewrite<'a> { /// we need to rewrite function calls with parameters /// and not actual values. extended: bool, - /// The statement is named (prepared), so we need to save - /// any derivatives into the global statement cache. - prepared: bool, /// Prepared statements cache for name mapping. prepared_statements: &'a mut PreparedStatements, /// Sharding schema for cache lookups. @@ -76,7 +67,6 @@ impl<'a> StatementRewrite<'a> { Self { rewritten: false, extended: ctx.extended, - prepared: ctx.prepared, prepared_statements: ctx.prepared_statements, schema: ctx.schema, db_schema: ctx.db_schema, @@ -85,16 +75,6 @@ impl<'a> StatementRewrite<'a> { } } - /// Create an AstContext from this rewriter's fields. - fn ast_context(&self) -> AstContext<'a> { - AstContext { - sharding_schema: self.schema.clone(), - db_schema: self.db_schema.clone(), - user: self.user, - search_path: self.search_path, - } - } - /// Maybe rewrite the statement and produce a rewrite plan /// we can apply to Bind messages. pub(crate) fn maybe_rewrite<'mem>( diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/offset.rs b/pgdog/src/frontend/router/parser/rewrite/statement/offset.rs index 3bd665e43..0fd237f0f 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/offset.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/offset.rs @@ -15,7 +15,7 @@ pub(crate) struct OffsetPlan { } impl OffsetPlan { - pub(super) fn apply_after_parser(&self, request: &mut ClientRequest) -> Result<(), Error> { + pub(crate) fn apply_after_parser(&self, request: &mut ClientRequest) -> Result<(), Error> { let route = match request.route.as_mut() { Some(route) => route, None => return Ok(()), @@ -242,7 +242,6 @@ mod tests { let mut ps = PreparedStatements::default(); let rewrite = StatementRewrite::new(StatementRewriteContext { extended: false, - prepared: false, prepared_statements: &mut ps, schema, db_schema: &db_schema, diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs b/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs index 82f3c723f..0f9159c42 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs @@ -1,13 +1,15 @@ -use crate::frontend::{ClientRequest, PreparedStatements}; +use super::offset::OffsetPlan; +use super::{Error, PrepareExecute, aggregate::AggregateRewritePlan}; +use crate::frontend::PreparedStatements; +use crate::frontend::client::query_engine; +pub(crate) use crate::frontend::client::query_engine::multi_step::types::{ + QueryPlanner, QueryPlannerType, +}; +use crate::frontend::client::query_engine::{QueryEngine, QueryEngineContext}; use crate::net::messages::bind::{Format, Parameter}; use crate::net::{Bind, Parse, ProtocolMessage, Query}; use crate::unique_id::UniqueId; - -use super::insert::build_split_requests; -use super::offset::OffsetPlan; -use super::{ - Error, InsertSplit, PrepareExecute, ShardingKeyUpdate, aggregate::AggregateRewritePlan, -}; +use std::fmt::Debug; /// Statement rewrite plan. /// @@ -34,40 +36,22 @@ pub(crate) struct RewritePlan { /// Each tuple contains (name, statement) for ProtocolMessage::Prepare. pub(crate) prepare_rewrites: Vec, - /// Splitting of multi-tuple INSERT statements into - /// multiple queries. - pub(crate) insert_split: Vec, - /// Position in the result where the count(*) or count(name) /// functions are added. pub(crate) aggregates: AggregateRewritePlan, + /// Splitting of multi-tuple INSERT statements into + /// multiple queries. + pub(crate) insert_split: bool, + /// Sharding key is being updated, we need to execute /// a multi-step plan. - pub(crate) sharding_key_update: Option, + pub(crate) sharding_key_update: bool, /// Limit/offset pagination. pub(crate) offset: Option, } -#[derive(Debug, Clone)] -pub(crate) enum RewriteResult { - InPlace { offset: Option }, - InsertSplit(Vec), - ShardingKeyUpdate(ShardingKeyUpdate), -} - -impl RewriteResult { - pub(crate) fn apply_after_parser(&self, request: &mut ClientRequest) -> Result<(), Error> { - match self { - Self::InPlace { - offset: Some(offset), - } => offset.apply_after_parser(request), - _ => Ok(()), - } - } -} - impl RewritePlan { /// True if the plan would not modify the query or its messages. /// `params` is purely informational (count of original `$N` placeholders) @@ -77,9 +61,9 @@ impl RewritePlan { && self.auto_id_injected == 0 && self.stmt.is_none() && self.prepare_rewrites.is_empty() - && self.insert_split.is_empty() + && !self.insert_split && self.aggregates.is_noop() - && self.sharding_key_update.is_none() + && !self.sharding_key_update && self.offset.is_none() } @@ -117,8 +101,13 @@ impl RewritePlan { } } - /// Apply the rewrite plan to a ClientRequest. - pub(crate) fn apply(&self, request: &mut ClientRequest) -> Result { + /// Apply the `RewritePlan` to the [`context.client_request`] to get a `QueryPlanner` back + pub(crate) async fn apply( + &self, + engine: &mut QueryEngine, + context: &mut QueryEngineContext<'_>, + ) -> Result, query_engine::multi_step::error::Error> { + let request = &mut *context.client_request; // Prepend any required Prepare messages for EXECUTE statements. if !self.prepare_rewrites.is_empty() { self.prepare_rewrites @@ -145,25 +134,20 @@ impl RewritePlan { } } - // Only rewrite executable requests. Some clients prepare the statement - // separately (e.g. go/pq with Parse, Describe, Sync). We don't need to rewrite - // those since insert split will return the same row(s) as multi-tuple insert. - if !self.insert_split.is_empty() && request.is_executable() { - let requests = build_split_requests(&self.insert_split, request)?; - return Ok(RewriteResult::InsertSplit(requests)); - } - - if let Some(sharding_key_update) = &self.sharding_key_update - && request.is_executable() - { - return Ok(RewriteResult::ShardingKeyUpdate( - sharding_key_update.clone(), - )); - } + // If the parser thinks we meet a special case, we should try planning for it. + // Else, ask the `QueryPlanner` to plan as a normal request. + let query_planner_type = { + if self.insert_split { + QueryPlannerType::InsertSplit + } else if self.sharding_key_update { + QueryPlannerType::ShardingKeyUpdate + } else { + QueryPlannerType::Normal + } + }; - Ok(RewriteResult::InPlace { - offset: self.offset.clone(), - }) + let request = context.client_request.clone(); + QueryPlanner::plan_query(&request, engine, context, query_planner_type).await } } diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/simple_prepared.rs b/pgdog/src/frontend/router/parser/rewrite/statement/simple_prepared.rs index efa7713bc..ed76cf10b 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/simple_prepared.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/simple_prepared.rs @@ -179,7 +179,6 @@ mod tests { let stmt = pg_raw_parse::parse(sql)?; let mut rewrite = StatementRewrite::new(StatementRewriteContext { extended: false, - prepared: false, prepared_statements: &mut self.ps, schema: &self.schema, db_schema: &self.db_schema, diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/unique_id.rs b/pgdog/src/frontend/router/parser/rewrite/statement/unique_id.rs index 7e6db6b23..cd7be7181 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/unique_id.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/unique_id.rs @@ -282,7 +282,6 @@ mod tests { let db_schema = default_db_schema(); let mut rewrite = StatementRewrite::new(StatementRewriteContext { extended, - prepared: false, prepared_statements: &mut ps, schema: &schema, db_schema: &db_schema, diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/update.rs b/pgdog/src/frontend/router/parser/rewrite/statement/update.rs index c1972c710..da8832579 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/update.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/update.rs @@ -1,238 +1,12 @@ -use indexmap::IndexSet; -use std::{ops::Deref, sync::Arc}; - -use pg_raw_parse::make::{owned, try_owned}; -use pg_raw_parse::{DeparseResult, Node, NodeMut, Owned, deparse, nodes, walk}; +use pg_raw_parse::make::owned; +use pg_raw_parse::{DeparseResult, Node, deparse, nodes}; use pgdog_config::RewriteMode; -use crate::{ - frontend::{ - BufferedQuery, ClientRequest, - router::{ - Ast, - parser::{Column, Table, Value}, - sharding::ShardedTable, - }, - }, - net::{ - Bind, DataRow, Describe, Execute, Flush, Format, Parse, ProtocolMessage, Query, - RowDescription, Sync, bind::Parameter, - }, -}; - use super::*; - -#[derive(Debug, Clone)] -pub(crate) struct Statement { - pub(crate) ast: Ast, - pub(crate) stmt: String, - pub(crate) params: IndexSet, -} - -impl Statement { - /// Create new Bind message for the statement from original Bind. - pub(crate) fn rewrite_bind(&self, bind: &Bind) -> Result { - let mut new = Bind::new_statement(""); // We use anonymous prepared - // statements for execution. - for param in &self.params { - let param = bind - .parameter(*param as usize - 1)? - .ok_or(Error::MissingParameter(*param))?; - new.push_param(param.parameter().clone(), param.format()); - } - - Ok(new) - } - - /// Build request from statement. - /// - /// Use the same protocol as the original statement. - /// - pub(crate) fn build_request(&self, request: &ClientRequest) -> Result { - let query = request.query()?.ok_or(Error::EmptyQuery)?; - let params = request.parameters()?; - - let mut request = ClientRequest::default(); - - match query { - BufferedQuery::Query(_) => { - request.push(Query::new(self.stmt.clone()).into()); - } - BufferedQuery::Prepared(_) => { - request.push(Parse::new_anonymous(&self.stmt).into()); - request.push(Describe::new_statement("").into()); - if let Some(params) = params { - request.push(self.rewrite_bind(params)?.into()); - request.push(Execute::new().into()); - request.push(Sync.into()); - } else { - // This shouldn't really happen since we don't rewrite - // non-executable requests. - request.push(Flush.into()); - } - } - } - - request.ast = Some(self.ast.clone()); - - Ok(request) - } -} - -#[derive(Debug, Clone)] -pub(crate) struct ShardingKeyUpdate { - inner: Arc, -} - -impl Deref for ShardingKeyUpdate { - type Target = Inner; - - fn deref(&self) -> &Self::Target { - &self.inner - } -} - -impl ShardingKeyUpdate { - pub(crate) fn sharded_table<'a>( - &self, - sharded_tables: &'a [ShardedTable], - ) -> Option<&'a ShardedTable> { - let table = self.target_table(); - - sharded_tables.iter().find(|sharded| { - if let Some(name) = sharded.name.as_ref() - && !table.name_match(name) - { - return false; - } - - if let Some(schema) = sharded.schema.as_ref() - && let Some(table_schema) = table.schema - && table_schema != schema - { - return false; - } - - self.from_update - .target_list() - .iter() - .any(|rt| rt.name() == Some(&*sharded.column)) - }) - } -} - -#[derive(Debug)] -pub(crate) struct Inner { - /// Check that the row actually moves shards. - pub(crate) check: Statement, - /// Delete old row from shard. - pub(crate) delete: Statement, - /// Update this is being constructed from - // FIXME(sage): There's no reason we need to own this, but this struct is - // ultimately a child of AstInner, where the statement is borrowed from, - // so we can't add a lifetime here. We should see if we can pass the update - // in later as needed - from_update: Owned, -} - -impl Inner { - pub(crate) fn target_table(&self) -> Table<'_> { - Table::from( - self.from_update - .relation() - .expect("UPDATE always has table"), - ) - } - - /// Build an INSERT statement built from an existing - /// UPDATE statement and a row returned by a SELECT statement. - pub(crate) fn build_insert_request( - &self, - request: &ClientRequest, - row_description: &RowDescription, - data_row: &DataRow, - ) -> Result { - let params = request.parameters()?; - let mut bind = Bind::new_statement(""); - - let insert = try_owned(|mem| -> Result<_, Error> { - let mut columns = Vec::new(); - let mut values = Vec::new(); - for (idx, field) in row_description.iter().enumerate() { - columns.push( - mem.make_res_target(Some(&*field.name), mem.empty(), mem.none()) - .uncast(), - ); - - if let Some(value) = self.from_update.target_list().iter().find_map(|rt| { - if rt.name() == Some(&*field.name) { - Some(rt.val()) - } else { - None - } - }) { - if let Ok(Value::Placeholder(number)) = Value::try_from(value) { - let param = params - .and_then(|p| p.parameter(number as usize - 1).transpose()) - .ok_or(Error::MissingParameter(number as u16))??; - bind.push_param(param.parameter().clone(), param.format()); - values.push(mem.make_param_ref(bind.params_raw().len() as _).uncast()); - } else { - values.push(mem.make_unique(value)); - } - } else { - // This column wasn't changed, get the value from the select - let value = data_row.get_raw(idx).ok_or(Error::MissingColumn(idx))?; - - if value.is_null() { - bind.push_param(Parameter::new_null(), Format::Text); - } else { - bind.push_param(Parameter::new(value), Format::Text); - } - values.push(mem.make_param_ref(bind.params_raw().len() as _).uncast()); - } - } - - let mut insert = mem.make_node::(); - insert - .as_mut() - .set_relation(mem.make_unique(self.from_update.relation())); - insert.as_mut().set_cols(mem.make_list(&columns)); - let mut select = mem.make_node::(); - select - .as_mut() - .set_values_lists(mem.make_list(&[mem.make_list(&values)])); - insert.as_mut().set_select_stmt(select.uncast()); - insert - .as_mut() - .set_returning_clause(mem.make_unique(self.from_update.returning_clause())); - Ok(mem.make_list(&[mem.make_raw_stmt(insert.uncast())])) - })?; - let stmt = deparse(insert.first().unwrap())?; - - //// Build the AST to be used with the router. - //// It's identical to the string-generated statement above. - let ast = Ast::from_raw_stmts(insert); - - let mut req = ClientRequest::from(vec![ - ProtocolMessage::from(Parse::new_anonymous(stmt.as_str())), - Describe::new_statement("").into(), // So we get both T and t, - bind.into(), - Execute::new().into(), - Sync.into(), - ]); - req.ast = Some(ast); - Ok(req) - } - - /// Do we have to return the rows to the client? - pub(crate) fn is_returning(&self) -> bool { - self.from_update.returning_clause().is_some() - } -} +use crate::frontend::router::parser::{Column, Table, Value}; impl<'a> StatementRewrite<'a> { - /// Create a plan for shardking key updates, if we suspect there is one + /// Create a plan for sharding key updates, if we suspect there is one /// in the query. pub(super) fn sharding_key_update( &mut self, @@ -243,23 +17,20 @@ impl<'a> StatementRewrite<'a> { return Ok(()); } - if let Some(value) = self.sharding_key_update_check(stmt)? { + if self.sharding_key_update_check(stmt)? { // Without a WHERE clause, this is a huge // cross-shard rewrite. if let Node::None = stmt.where_clause() { return Err(Error::WhereClauseMissing); } - plan.sharding_key_update = Some(create_stmts(stmt, value)?); + plan.sharding_key_update = true; } Ok(()) } /// Check if the sharding key could be updated. - fn sharding_key_update_check( - &'a self, - stmt: &'a nodes::UpdateStmt, - ) -> Result, Error> { + fn sharding_key_update_check(&'a self, stmt: &'a nodes::UpdateStmt) -> Result { let table = stmt .relation() .map(Table::from) @@ -271,13 +42,13 @@ impl<'a> StatementRewrite<'a> { self.schema.tables().get_table(c).is_some() }) }) else { - return Ok(None); + return Ok(false); }; // Check that it's a value assignment and not something like // id = id + 1 if Value::try_from(shard_key_assignment.val()).is_ok() { - Ok(Some(shard_key_assignment)) + Ok(true) } else { let expr = shard_key_assignment.val(); let expr = deparse_expr([expr])?; @@ -295,21 +66,10 @@ impl<'a> StatementRewrite<'a> { } } -/// Visit all ParamRef nodes in a ParseResult and renumber them sequentially. -/// Returns a sorted list of the original parameter numbers. -fn rewrite_params(node: NodeMut<'_, '_>) -> IndexSet { - let mut params = IndexSet::new(); - walk::walk_mut(node, |node| { - if let NodeMut::ParamRef(param) = node { - params.insert(param.number as _); - param.set_number(params.get_index_of(&(param.number as u16)).unwrap() as i32 + 1) - } - }); - params -} - /// Deparse an expression node by wrapping it in a SELECT statement. -fn deparse_expr<'a>(nodes: impl IntoIterator>) -> Result { +pub(crate) fn deparse_expr<'a>( + nodes: impl IntoIterator>, +) -> Result { let node = owned(|mem| { let mut select = mem.make_node::(); let res_targets = nodes @@ -324,614 +84,3 @@ fn deparse_expr<'a>(nodes: impl IntoIterator>) -> Result( - stmt: &'a nodes::UpdateStmt, - new_value: &'a nodes::ResTarget, -) -> Result { - let select_star = owned(|mem| { - let mut select_stmt = mem.make_node::(); - select_stmt.as_mut().set_target_list( - mem.make_list(&[mem.make_res_target( - None, - mem.empty(), - mem.make_column_ref(mem.make_list(&[mem.make_node::().uncast()])) - .uncast(), - )]), - ); - select_stmt.as_mut().set_from_clause( - mem.make_list(&[mem - .make_unique(stmt.relation().expect("UPDATE always has a table")) - .uncast()]), - ); - select_stmt - }); - - let mut params = IndexSet::new(); - let delete = owned(|mem| { - let mut delete = mem.make_node::(); - delete - .as_mut() - .set_relation(mem.make_unique(stmt.relation())); - delete - .as_mut() - .set_where_clause(mem.make_unique(stmt.where_clause())); - delete.as_mut().set_returning_clause( - mem.make_returning_clause( - mem.make_list(&[mem - .make_res_target( - None, - mem.empty(), - mem.make_column_ref( - mem.make_list(&[mem.make_node::().uncast()]), - ) - .uncast(), - ) - .uncast()]), - ) - .as_option(), - ); - params = rewrite_params(delete.as_mut().into()); - mem.make_list(&[mem.make_raw_stmt(delete.uncast())]) - }); - - let delete = Statement { - stmt: deparse(delete.first().unwrap())?.as_str().to_owned(), - ast: Ast::from_raw_stmts(delete), - params, - }; - - let mut params = IndexSet::new(); - let check = owned(|mem| { - let mut select_stmt = mem.make_unique(&*select_star); - select_stmt.as_mut().set_where_clause( - mem.make_a_expr( - nodes::A_Expr_Kind::AEXPR_OP, - mem.make_list(&[mem.make_string(Some("=")).uncast()]), - mem.make_column_ref(mem.make_list(&[mem.make_string(new_value.name()).uncast()])) - .uncast(), - mem.make_unique(new_value.val()), - ) - .uncast(), - ); - params = rewrite_params(select_stmt.as_mut().into()); - mem.make_list(&[mem.make_raw_stmt(select_stmt.uncast())]) - }); - - let check = Statement { - stmt: deparse(check.first().unwrap())?.as_str().to_owned(), - ast: Ast::from_raw_stmts(check), - params, - }; - - Ok(ShardingKeyUpdate { - inner: Arc::new(Inner { - delete, - check, - from_update: owned(|mem| mem.make_unique(stmt)), - }), - }) -} - -#[cfg(test)] -mod test { - use crate::frontend::router::sharding::ShardedTable; - use indexmap::indexset; - use pgdog_config::Rewrite; - - use crate::backend::schema::Schema; - use crate::backend::{ShardedTables, replication::ShardedSchemas}; - use crate::net::messages::row_description::Field; - - use super::*; - - fn default_db_schema() -> Schema { - Schema::default() - } - - fn default_schema() -> ShardingSchema { - ShardingSchema { - shards: 2, - tables: ShardedTables::new( - vec![ShardedTable { - database: "pgdog".into(), - name: Some("sharded".into()), - column: "id".into(), - ..Default::default() - }], - vec![], - false, - pgdog_config::SystemCatalogsBehavior::default(), - ), - schemas: ShardedSchemas::new(vec![]), - rewrite: Rewrite { - enabled: true, - shard_key: RewriteMode::Rewrite, - ..Default::default() - }, - ..Default::default() - } - } - - fn run_test(query: &str) -> Result, Error> { - let stmt = pg_raw_parse::parse(query)?; - let schema = default_schema(); - let db_schema = default_db_schema(); - let mut stmts = PreparedStatements::new(); - - let ctx = StatementRewriteContext { - schema: &schema, - db_schema: &db_schema, - extended: true, - prepared: false, - prepared_statements: &mut stmts, - user: "", - search_path: None, - }; - let mut plan = RewritePlan::default(); - StatementRewrite::new(ctx).sharding_key_update( - match stmt.stmts().next().unwrap() { - Node::UpdateStmt(stmt) => stmt, - _ => panic!("Not an update"), - }, - &mut plan, - )?; - Ok(plan.sharding_key_update) - } - - #[test] - fn test_select_basic_where_param() { - let result = run_test("UPDATE sharded SET id = $1 WHERE email = $2") - .unwrap() - .unwrap(); - - // SELECT should have WHERE clause with param renumbered to $1 - assert_eq!( - result.delete.stmt, - "DELETE FROM sharded WHERE email = $1 RETURNING *" - ); - assert_eq!(result.delete.params, indexset![2]); - - let schema = default_schema(); - let tables = schema.tables.tables(); - assert_eq!(result.target_table().name, "sharded"); - assert_eq!(result.sharded_table(tables).unwrap().column, "id"); - } - - #[test] - fn test_select_multiple_where_params() { - let result = run_test("UPDATE sharded SET id = $1 WHERE email = $2 AND name = $3") - .unwrap() - .unwrap(); - - assert_eq!( - result.delete.stmt, - "DELETE FROM sharded WHERE email = $1 AND name = $2 RETURNING *" - ); - assert_eq!(result.delete.params, indexset![2, 3]); - assert!(!result.is_returning()); - } - - #[test] - fn test_select_non_sequential_params() { - // Params in WHERE are $3 and $5, should be renumbered to $1 and $2 - let result = run_test( - "UPDATE sharded SET id = $1, value = $2, other = $4 WHERE email = $3 AND name = $5", - ) - .unwrap() - .unwrap(); - - assert_eq!( - result.delete.stmt, - "DELETE FROM sharded WHERE email = $1 AND name = $2 RETURNING *" - ); - assert_eq!(result.delete.params, indexset![3, 5]); - } - - #[test] - fn test_select_single_where_param() { - let result = run_test("UPDATE sharded SET id = $1 WHERE email = $2") - .unwrap() - .unwrap(); - - assert_eq!( - result.delete.stmt, - "DELETE FROM sharded WHERE email = $1 RETURNING *" - ); - assert_eq!(result.delete.params, indexset![2]); - } - - #[test] - fn test_delete_basic() { - let result = run_test("UPDATE sharded SET id = $1 WHERE email = $2") - .unwrap() - .unwrap(); - - assert_eq!( - result.delete.stmt, - "DELETE FROM sharded WHERE email = $1 RETURNING *" - ); - - assert!(result.sharded_table(&[]).is_none()); - assert!( - result - .sharded_table(&[ShardedTable { - name: Some("other".into()), - column: "id".into(), - ..Default::default() - }]) - .is_none() - ); - assert!( - result - .sharded_table(&[ShardedTable { - name: Some("sharded".into()), - column: "user_id".into(), - ..Default::default() - }]) - .is_none() - ); - } - - #[test] - fn test_delete_multiple_where_params() { - let result = run_test("UPDATE sharded SET id = $1 WHERE email = $2 AND name = $3") - .unwrap() - .unwrap(); - - assert_eq!( - result.delete.stmt, - "DELETE FROM sharded WHERE email = $1 AND name = $2 RETURNING *" - ); - } - - #[test] - fn test_no_params_in_where() { - let result = run_test("UPDATE sharded SET id = $1 WHERE email = 'test@example.com'") - .unwrap() - .unwrap(); - - assert_eq!( - result.delete.stmt, - "DELETE FROM sharded WHERE email = 'test@example.com' RETURNING *" - ); - assert!(result.delete.params.is_empty()); - } - - #[test] - fn test_where_with_in_clause() { - let result = run_test("UPDATE sharded SET id = $1 WHERE email IN ($2, $3, $4)") - .unwrap() - .unwrap(); - - assert_eq!( - result.delete.stmt, - "DELETE FROM sharded WHERE email IN ($1, $2, $3) RETURNING *" - ); - assert_eq!(result.delete.params, indexset![2, 3, 4]); - } - - #[test] - fn test_where_with_comparison_operators() { - let result = run_test("UPDATE sharded SET id = $1 WHERE count > $2 AND count < $3") - .unwrap() - .unwrap(); - - assert_eq!( - result.delete.stmt, - "DELETE FROM sharded WHERE count > $1 AND count < $2 RETURNING *" - ); - assert_eq!(result.delete.params, indexset![2, 3]); - } - - #[test] - fn test_where_with_or_condition() { - let result = run_test("UPDATE sharded SET id = $1 WHERE email = $2 OR name = $3") - .unwrap() - .unwrap(); - - assert_eq!( - result.delete.stmt, - "DELETE FROM sharded WHERE email = $1 OR name = $2 RETURNING *" - ); - assert_eq!(result.delete.params, indexset![2, 3]); - } - - #[test] - fn test_high_param_numbers() { - let result = run_test("UPDATE sharded SET id = $10 WHERE email = $20 AND name = $30") - .unwrap() - .unwrap(); - - assert_eq!( - result.delete.stmt, - "DELETE FROM sharded WHERE email = $1 AND name = $2 RETURNING *" - ); - assert_eq!(result.delete.params, indexset![20, 30]); - } - - #[test] - fn test_non_sharding_key_update_returns_none() { - // Updating a non-sharding column should return None - let result = run_test("UPDATE sharded SET email = $1 WHERE id = $2").unwrap(); - assert!(result.is_none()); - } - - #[test] - fn test_where_with_like() { - let result = run_test("UPDATE sharded SET id = $1 WHERE email LIKE $2") - .unwrap() - .unwrap(); - - assert_eq!( - result.delete.stmt, - "DELETE FROM sharded WHERE email LIKE $1 RETURNING *" - ); - assert_eq!(result.delete.params, indexset![2]); - } - - #[test] - fn test_where_with_is_null() { - let result = run_test("UPDATE sharded SET id = $1 WHERE email = $2 AND deleted_at IS NULL") - .unwrap() - .unwrap(); - - assert_eq!( - result.delete.stmt, - "DELETE FROM sharded WHERE email = $1 AND deleted_at IS NULL RETURNING *" - ); - assert_eq!(result.delete.params, indexset![2]); - } - - #[test] - fn test_where_with_between() { - let result = run_test("UPDATE sharded SET id = $1 WHERE created_at BETWEEN $2 AND $3") - .unwrap() - .unwrap(); - - assert_eq!( - result.delete.stmt, - "DELETE FROM sharded WHERE created_at BETWEEN $1 AND $2 RETURNING *" - ); - assert_eq!(result.delete.params, indexset![2, 3]); - } - - #[test] - fn test_same_param_used_twice() { - // Same parameter $2 used twice in WHERE clause - let result = run_test("UPDATE sharded SET id = $1 WHERE email = $2 OR name = $2") - .unwrap() - .unwrap(); - - // Both occurrences should be renumbered to $1 - assert_eq!( - result.delete.stmt, - "DELETE FROM sharded WHERE email = $1 OR name = $1 RETURNING *" - ); - // Only one unique param in the mapping - assert_eq!(result.delete.params, indexset![2]); - } - - #[test] - fn test_same_param_used_multiple_times() { - // $2 used three times - let result = run_test("UPDATE sharded SET id = $1 WHERE a = $2 AND b = $2 AND c = $2") - .unwrap() - .unwrap(); - - assert_eq!( - result.delete.stmt, - "DELETE FROM sharded WHERE a = $1 AND b = $1 AND c = $1 RETURNING *" - ); - assert_eq!(result.delete.params, indexset![2]); - } - - #[test] - fn test_mixed_repeated_and_unique_params() { - // $2 used twice, $3 used once - let result = run_test("UPDATE sharded SET id = $1 WHERE a = $2 AND b = $3 AND c = $2") - .unwrap() - .unwrap(); - - assert_eq!( - result.delete.stmt, - "DELETE FROM sharded WHERE a = $1 AND b = $2 AND c = $1 RETURNING *" - ); - assert_eq!(result.delete.params, indexset![2, 3]); - } - - #[test] - fn test_repeated_params_in_in_clause() { - // Same param repeated in IN clause (unusual but valid) - let result = run_test("UPDATE sharded SET id = $1 WHERE email IN ($2, $3, $2)") - .unwrap() - .unwrap(); - - assert_eq!( - result.delete.stmt, - "DELETE FROM sharded WHERE email IN ($1, $2, $1) RETURNING *" - ); - assert_eq!(result.delete.params, indexset![2, 3]); - } - - #[test] - fn test_delete_with_repeated_params() { - let result = run_test("UPDATE sharded SET id = $1 WHERE email = $2 OR name = $2") - .unwrap() - .unwrap(); - - assert_eq!( - result.delete.stmt, - "DELETE FROM sharded WHERE email = $1 OR name = $1 RETURNING *" - ); - assert_eq!(result.delete.params, indexset![2]); - } - - #[test] - fn test_sharding_key_not_changed() { - let result = run_test("UPDATE sharded SET id = $1 WHERE id = $1 AND email = $2") - .unwrap() - .unwrap(); - assert_eq!(result.check.stmt, "SELECT * FROM sharded WHERE id = $1"); - assert_eq!(result.check.params, indexset![1]); - } - - #[test] - fn test_unsupported_assignment() { - let result = run_test("UPDATE sharded SET id = random() WHERE id = $1"); - std::assert_matches!( - result, - Err(Error::UnsupportedShardingKeyUpdate(msg)) if msg == "\"id\" = random()" - ); - } - - #[test] - fn test_unsupported_assignment_arithmetic_add() { - let result = run_test("UPDATE sharded SET id = id + 1 WHERE id = $1"); - std::assert_matches!( - result, - Err(Error::UnsupportedShardingKeyUpdate(msg)) if msg == "\"id\" = id + 1" - ); - } - - #[test] - fn test_unsupported_assignment_arithmetic_multiply() { - let result = run_test("UPDATE sharded SET id = id * 2 WHERE id = $1"); - std::assert_matches!( - result, - Err(Error::UnsupportedShardingKeyUpdate(msg)) if msg == "\"id\" = id * 2" - ); - } - - #[test] - fn test_unsupported_assignment_arithmetic_with_param() { - let result = run_test("UPDATE sharded SET id = id + $2 WHERE id = $1"); - std::assert_matches!( - result, - Err(Error::UnsupportedShardingKeyUpdate(msg)) if msg == "\"id\" = id + $2" - ); - } - - #[test] - fn test_unsupported_assignment_now() { - let result = run_test("UPDATE sharded SET id = now() WHERE id = $1"); - std::assert_matches!( - result, - Err(Error::UnsupportedShardingKeyUpdate(msg)) if msg == "\"id\" = now()" - ); - } - - #[test] - fn test_unsupported_assignment_coalesce() { - let result = run_test("UPDATE sharded SET id = coalesce(id, 0) WHERE id = $1"); - std::assert_matches!( - result, - Err(Error::UnsupportedShardingKeyUpdate(msg)) if msg == "\"id\" = COALESCE(id, 0)" - ); - } - - #[test] - fn test_unsupported_assignment_case() { - let result = - run_test("UPDATE sharded SET id = CASE WHEN id > 0 THEN 1 ELSE 0 END WHERE id = $1"); - std::assert_matches!( - result, - Err(Error::UnsupportedShardingKeyUpdate(msg)) if msg == "\"id\" = CASE WHEN id > 0 THEN 1 ELSE 0 END" - ); - } - - #[test] - fn test_unsupported_assignment_subquery() { - let result = - run_test("UPDATE sharded SET id = (SELECT max(id) FROM sharded) WHERE id = $1"); - std::assert_matches!( - result, - Err(Error::UnsupportedShardingKeyUpdate(msg)) if msg == "\"id\" = (SELECT max(id) FROM sharded)" - ); - } - - #[test] - fn test_unsupported_assignment_column_reference() { - let result = run_test("UPDATE sharded SET id = other_column WHERE id = $1"); - std::assert_matches!( - result, - Err(Error::UnsupportedShardingKeyUpdate(msg)) if msg == "\"id\" = other_column" - ); - } - - #[test] - fn test_unsupported_assignment_concat() { - let result = run_test("UPDATE sharded SET id = id || '_suffix' WHERE id = $1"); - std::assert_matches!( - result, - Err(Error::UnsupportedShardingKeyUpdate(msg)) if msg == "\"id\" = id || '_suffix'" - ); - } - - #[test] - fn test_unsupported_assignment_negation() { - let result = run_test("UPDATE sharded SET id = -id WHERE id = $1"); - std::assert_matches!( - result, - Err(Error::UnsupportedShardingKeyUpdate(msg)) if msg == "\"id\" = - id" - ); - } - - #[test] - fn test_insert_build_request_with_expr_column() { - // Test that INSERT statement is built correctly when there are expression columns. - // The expression should appear directly in the VALUES clause. - // Use literal values (not placeholders) to avoid needing bind parameters. - let result = run_test("UPDATE sharded SET id = 42, email = random() WHERE id = 1") - .unwrap() - .unwrap(); - - // Create a mock row description matching the SELECT * result - let row_description = RowDescription::new(&[ - Field::bigint("id"), - Field::text("email"), - Field::text("other_col"), - Field::text("other_other_col"), - ]); - - // Create a mock data row with values for columns not in the UPDATE SET clause - let mut data_row = DataRow::new(); - data_row.add("1"); // id - will be overwritten by mapping - data_row.add("old@example.com"); // email - will be overwritten by mapping - data_row.add("other_value"); // other_col - from existing row - data_row.add("other_other_value"); // other_other_col - from existing row - - // Create a simple query request (not prepared statement) - let request = ClientRequest::from(vec![ProtocolMessage::from(Query::new( - "UPDATE sharded SET id = 42, email = random() WHERE id = 1", - ))]); - - let insert_request = result - .build_insert_request(&request, &row_description, &data_row) - .unwrap(); - - // Get the query from the request to verify the INSERT statement - let query = insert_request.query().unwrap().unwrap(); - let stmt = query.query(); - - // The INSERT should contain the expression random() directly in VALUES - assert!( - stmt.contains("random()"), - "INSERT statement should contain the expression: {}", - stmt - ); - // Verify it's an INSERT statement - assert!( - stmt.starts_with("INSERT INTO"), - "Should be an INSERT statement: {}", - stmt - ); - // Verify parameter numbering is correct: $1 for id, random() for email, $2 for other_col - // (not $3, which would be wrong if we used row index instead of bind param index) - assert!( - stmt.contains("$1") && stmt.contains("$2") && !stmt.contains("$3"), - "Parameter numbering should be sequential without gaps: {}", - stmt - ); - } -} diff --git a/pgdog/src/net/messages/parse.rs b/pgdog/src/net/messages/parse.rs index aec7ad895..eeda0728f 100644 --- a/pgdog/src/net/messages/parse.rs +++ b/pgdog/src/net/messages/parse.rs @@ -145,7 +145,6 @@ impl Parse { rewritten } - #[cfg(test)] pub(crate) fn with_data_types(&self, data_types: &[u32]) -> Self { let mut bytes = BytesMut::new(); bytes.put_u16(data_types.len() as _);