diff --git a/.schema/pgdog.schema.json b/.schema/pgdog.schema.json index a334988f3..cae9440bc 100644 --- a/.schema/pgdog.schema.json +++ b/.schema/pgdog.schema.json @@ -220,6 +220,7 @@ "enabled": false, "primary_key": "ignore", "shard_key": "error", + "simple_to_prepared": false, "split_inserts": "error" } }, @@ -1870,6 +1871,11 @@ "$ref": "#/$defs/RewriteMode", "default": "error" }, + "simple_to_prepared": { + "description": "Rewrite simple queries to prepared statements.", + "type": "boolean", + "default": false + }, "split_inserts": { "description": "Behavior for multi-row `INSERT` on sharded tables: `error` rejects, `rewrite` distributes rows to their shards, `ignore` forwards unchanged.\n\n_Default:_ `error`\n\n", "$ref": "#/$defs/RewriteMode", diff --git a/cli.sh b/cli.sh index b21c93248..f871de793 100755 --- a/cli.sh +++ b/cli.sh @@ -17,7 +17,7 @@ function admin() { # - protocol: simple|extended|prepared # function bench() { - PGPASSWORD=pgdog pgbench -h 127.0.0.1 -p 6432 -U pgdog pgdog --protocol ${1:-simple} -t 100000000 -c 10 -P 1 -f pgdog/tests/pgbouncer/pgbench-parser.sql + PGPASSWORD=pgdog pgbench -h 127.0.0.1 -p 6432 -U pgdog pgdog --protocol ${1:-simple} -t 100000000 -c 10 -P 1 -S } function bench_init() { diff --git a/integration/simple_to_prepared/pgdog.toml b/integration/simple_to_prepared/pgdog.toml new file mode 100644 index 000000000..de06af273 --- /dev/null +++ b/integration/simple_to_prepared/pgdog.toml @@ -0,0 +1,11 @@ +[general] +idle_healthcheck_delay = 1000000000 +query_parser = "on" + +[[databases]] +name = "pgdog" +host = "127.0.0.1" + + +[rewrite] +simple_to_prepared = false diff --git a/integration/simple_to_prepared/users.toml b/integration/simple_to_prepared/users.toml new file mode 100644 index 000000000..539bb1832 --- /dev/null +++ b/integration/simple_to_prepared/users.toml @@ -0,0 +1,4 @@ +[[users]] +name = "pgdog" +password = "pgdog" +database = "pgdog" diff --git a/pgdog-config/src/rewrite.rs b/pgdog-config/src/rewrite.rs index d7d0adb25..7121a563e 100644 --- a/pgdog-config/src/rewrite.rs +++ b/pgdog-config/src/rewrite.rs @@ -88,6 +88,10 @@ pub struct Rewrite { /// #[serde(default = "Rewrite::default_primary_key")] pub primary_key: RewriteMode, + + /// Rewrite simple queries to prepared statements. + #[serde(default)] + pub simple_to_prepared: bool, } impl Default for Rewrite { @@ -97,6 +101,7 @@ impl Default for Rewrite { shard_key: Self::default_shard_key(), split_inserts: Self::default_split_inserts(), primary_key: Self::default_primary_key(), + simple_to_prepared: bool::default(), } } } diff --git a/pgdog/src/frontend/client/query_engine/query.rs b/pgdog/src/frontend/client/query_engine/query.rs index cb250d22e..399dccd37 100644 --- a/pgdog/src/frontend/client/query_engine/query.rs +++ b/pgdog/src/frontend/client/query_engine/query.rs @@ -242,12 +242,20 @@ impl QueryEngine { // Do this before flushing, because flushing can take time. self.cleanup_backend(context)?; - trace!("{:#?} >>> {:?}", message, context.stream.peer_addr()); - - if flush { - context.stream.send_flush(&message).await?; - } else { - context.stream.send(&message).await?; + let forward_to_client = context + .rewrite_result + .as_ref() + .map(|rewrite| rewrite.apply_after_execution(&message).forward()) + .unwrap_or(true); + + if forward_to_client { + trace!("{:#?} >>> {:?}", message, context.stream.peer_addr()); + + if flush { + context.stream.send_flush(&message).await?; + } else { + context.stream.send(&message).await?; + } } if code == 'Z' { 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 3514e3a6c..6ae4adca9 100644 --- a/pgdog/src/frontend/client/query_engine/test/rewrite_offset.rs +++ b/pgdog/src/frontend/client/query_engine/test/rewrite_offset.rs @@ -18,7 +18,7 @@ async fn run_test(messages: Vec) -> Option { engine.parse_and_rewrite(&mut context).await.unwrap(); match context.rewrite_result { - Some(RewriteResult::InPlace { offset }) => offset, + Some(RewriteResult::InPlace { offset, .. }) => offset, other => panic!("expected InPlace, got {:?}", other), } } diff --git a/pgdog/src/frontend/prepared_statements/mod.rs b/pgdog/src/frontend/prepared_statements/mod.rs index 7419cceb8..f6b7c66b1 100644 --- a/pgdog/src/frontend/prepared_statements/mod.rs +++ b/pgdog/src/frontend/prepared_statements/mod.rs @@ -2,6 +2,7 @@ use std::{collections::HashMap, sync::Arc, time::Duration}; +use bytes::Bytes; use once_cell::sync::Lazy; use parking_lot::RwLock; use tracing::debug; @@ -34,6 +35,7 @@ pub struct PreparedStatements { // mapping the client statement name -> __pgdog__ name from global cache pub(super) local: HashMap, pub(super) level: PreparedStatementsLevel, + rewritten_simple_to_prepared: HashMap, pub(super) memory_used: usize, } @@ -43,6 +45,7 @@ impl Default for PreparedStatements { global: Arc::new(RwLock::new(GlobalCache::default())), local: HashMap::default(), level: PreparedStatementsLevel::Extended, + rewritten_simple_to_prepared: HashMap::new(), memory_used: 0, } } @@ -66,8 +69,28 @@ impl PreparedStatements { Ok(()) } + /// Manually map a local prepared statement to a global one. + /// + /// Warning: don't use this unless you understand the side-effects: + /// + /// 1. When client disconnects, this statement's global counter will be decreased by 1. + /// 2. The statement will not be removed from the global cache until the client disconnects + /// because clients are not aware of this and will never close it. + /// + pub(crate) fn insert_rewritten_simple_to_prepared(&mut self, parse: &Parse) -> String { + if let Some(name) = self.rewritten_simple_to_prepared.get(&parse.query_ref()) { + name.to_owned() + } else { + let (_new, name) = { self.global.write().insert(parse) }; + self.local.insert(name.to_owned(), name.to_owned()); + self.rewritten_simple_to_prepared + .insert(parse.query_ref(), name.clone()); + name + } + } + /// Register prepared statement with the global cache. - pub fn insert(&mut self, parse: &mut Parse) { + pub(crate) fn insert(&mut self, parse: &mut Parse) { let (_new, name) = { self.global.write().insert(parse) }; let key = parse.name(); let existed = self.local.insert(key.to_owned(), name.clone()); diff --git a/pgdog/src/frontend/router/parser/cache/ast.rs b/pgdog/src/frontend/router/parser/cache/ast.rs index e6308c0a1..58921054f 100644 --- a/pgdog/src/frontend/router/parser/cache/ast.rs +++ b/pgdog/src/frontend/router/parser/cache/ast.rs @@ -68,7 +68,7 @@ impl Deref for Ast { impl Ast { /// Parse statement and run the rewrite engine, if necessary. - pub(super) fn new( + fn new( query: &AstQuery, schema: &ShardingSchema, db_schema: &Schema, @@ -78,6 +78,7 @@ impl Ast { ) -> Result { let now = Instant::now(); let ast = pg_raw_parse::parse(query.query_without_comment).map_err(Error::Parse)?; + let multiple_statements = ast.stmts().count() > 1; // Run the rewrite unconditionally. Even when a shard comment will // route the query to a specific shard, we need to know whether the @@ -91,6 +92,7 @@ impl Ast { db_schema, user, search_path, + multiple_statements, }); let mut rewrite_plan = Default::default(); let ast = make::try_owned(|mem| { diff --git a/pgdog/src/frontend/router/parser/cache/test.rs b/pgdog/src/frontend/router/parser/cache/test.rs index fda1cb071..a6cd3acce 100644 --- a/pgdog/src/frontend/router/parser/cache/test.rs +++ b/pgdog/src/frontend/router/parser/cache/test.rs @@ -359,3 +359,28 @@ fn test_truncated_query_non_ascii_char_boundary() { let ast_query = AstQuery::from_query(&buffered); assert_eq!(ast_query.truncated_query(9), "SELECT '€"); } + +#[test] +fn rejects_rewritten_multi_statement_queries() { + let mut ctx = test_context(); + ctx.sharding_schema.rewrite.simple_to_prepared = true; + let mut prepared_statements = PreparedStatements::default(); + + let unchanged = + BufferedQuery::Query(Query::new("SELECT current_user; SELECT current_database()")); + Cache::get() + .query(&unchanged, &ctx, &mut prepared_statements) + .expect("multi-statement query without rewrites should parse"); + + let rewritten = BufferedQuery::Query(Query::new("SELECT 1; SELECT 2")); + let error = Cache::get() + .query(&rewritten, &ctx, &mut prepared_statements) + .expect_err("rewritten multi-statement query should be rejected"); + + assert!(matches!( + error, + crate::frontend::router::parser::Error::Rewrite( + crate::frontend::router::parser::rewrite::statement::Error::MultiStatementRewrite + ) + )); +} 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 c56a074bf..25993205e 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/auto_id.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/auto_id.rs @@ -500,6 +500,7 @@ mod tests { db_schema, user: "", search_path: None, + multiple_statements: false, }); let mut plan = Default::default(); let ast = make::try_owned(|mem| { diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/error.rs b/pgdog/src/frontend/router/parser/rewrite/statement/error.rs index 18e81a7b5..f5c40be9b 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/error.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/error.rs @@ -40,4 +40,10 @@ pub enum Error { #[error("prepared statement '{0}' does not exist")] ExecuteMissingPrepare(String), + + #[error("prepared statement: {0}")] + PreparedStmt(#[from] crate::frontend::prepared_statements::Error), + + #[error("cannot rewrite a multi-statement query")] + MultiStatementRewrite, } diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/insert.rs b/pgdog/src/frontend/router/parser/rewrite/statement/insert.rs index 6e127a59c..40fc21b3e 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/insert.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/insert.rs @@ -273,6 +273,7 @@ mod tests { db_schema: &db_schema, user: "", search_path: None, + multiple_statements: false, }); let mut plan = RewritePlan::default(); rewriter.split_insert(insert, &mut plan).unwrap(); diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs b/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs index 8ef1105d2..6a9a5c211 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs @@ -15,6 +15,7 @@ pub mod insert; pub mod offset; pub mod plan; pub mod simple_prepared; +pub mod simple_to_prepared; pub mod unique_id; pub mod update; @@ -22,6 +23,7 @@ pub use error::Error; pub use insert::InsertSplit; pub(crate) use plan::RewritePlan; pub use simple_prepared::SimplePreparedResult; +pub(crate) use simple_to_prepared::*; pub(crate) use update::*; /// Statement rewrite engine context. @@ -42,6 +44,8 @@ pub struct StatementRewriteContext<'a> { pub user: &'a str, /// Search path for table lookups. pub search_path: Option<&'a ParameterValue>, + /// Whether the query contains more than one SQL statement. + pub multiple_statements: bool, } #[derive(Debug)] @@ -65,6 +69,8 @@ pub struct StatementRewrite<'a> { user: &'a str, /// Search path for table lookups. search_path: Option<&'a ParameterValue>, + /// Whether the query contains more than one SQL statement. + multiple_statements: bool, } impl<'a> StatementRewrite<'a> { @@ -82,6 +88,7 @@ impl<'a> StatementRewrite<'a> { db_schema: ctx.db_schema, user: ctx.user, search_path: ctx.search_path, + multiple_statements: ctx.multiple_statements, } } @@ -104,6 +111,10 @@ impl<'a> StatementRewrite<'a> { ) -> Result { let mut plan = RewritePlan::default(); + // N.B. The simple to prepared rewriter should run first. + // All subsequent rewriters will act on the prepared statement. + self.rewrite_simple_to_prepared(stmt.stmt_mut(), mem, &mut plan)?; + match stmt.stmt() { Node::InsertStmt(_) | Node::SelectStmt(_) @@ -162,8 +173,19 @@ impl<'a> StatementRewrite<'a> { self.limit_offset(&select, &mut plan); } + if self.rewritten && self.multiple_statements { + return Err(Error::MultiStatementRewrite); + } + if self.rewritten { - plan.stmt = Some(pg_raw_parse::deparse(&*stmt)?.as_str().to_owned()); + let stmt = pg_raw_parse::deparse(&*stmt)?.as_str().to_owned(); + + // N.B. careful with ordering. This should run before insert splits, etc. + // since we want to make sure the statement is registered with the global cache. + plan.simple_to_prepared + .step_two(self.prepared_statements, &stmt)?; + + plan.stmt = Some(stmt); } if let Node::InsertStmt(insert) = stmt.stmt() { diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/offset.rs b/pgdog/src/frontend/router/parser/rewrite/statement/offset.rs index 2f9440043..a15a81ef7 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/offset.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/offset.rs @@ -248,6 +248,7 @@ mod tests { db_schema: &db_schema, user: "test", search_path: None, + multiple_statements: false, }); let mut plan = RewritePlan::default(); rewrite.limit_offset( diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs b/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs index 6b2bbb9e8..20358f1ca 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs @@ -1,11 +1,13 @@ use crate::frontend::{ClientRequest, PreparedStatements}; use crate::net::messages::bind::{Format, Parameter}; -use crate::net::{Bind, Parse, ProtocolMessage, Query}; +use crate::net::{Bind, Message, Parse, Protocol, ProtocolMessage, Query}; use crate::unique_id::UniqueId; use super::insert::build_split_requests; use super::offset::OffsetPlan; -use super::{Error, InsertSplit, ShardingKeyUpdate, aggregate::AggregateRewritePlan}; +use super::{ + Error, InsertSplit, ShardingKeyUpdate, SimpleToPreparedPlan, aggregate::AggregateRewritePlan, +}; /// Statement rewrite plan. /// @@ -46,24 +48,70 @@ pub struct RewritePlan { /// Limit/offset pagination. pub(crate) offset: Option, + + /// Simple to prepared rewrite. + pub(crate) simple_to_prepared: SimpleToPreparedPlan, } #[derive(Debug, Clone)] pub(crate) enum RewriteResult { - InPlace { offset: Option }, + InPlace { + offset: Option, + simple_to_prepared: bool, + }, InsertSplit(Vec), ShardingKeyUpdate(ShardingKeyUpdate), } +/// Action to be taken by the query engine +/// given the state of the rewrite and the message +/// received from Postgres. +#[derive(Debug, Clone, PartialEq)] +pub(crate) enum AfterExecutionAction { + /// Forward message as-is to the client. + Forward, + /// Drop the message. + Drop, +} + +impl AfterExecutionAction { + /// Forward the message to the client. + pub(crate) fn forward(&self) -> bool { + self == &Self::Forward + } +} + 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(()), } } + + /// Apply any filtering/rewriting rules to messages received from Postgres + /// given the rewrite performed by the rewrite engine. + /// + /// # Arguments + /// + /// - `message`: Message received from a Postgres server. + /// + pub(crate) fn apply_after_execution(&self, message: &Message) -> AfterExecutionAction { + if let Self::InPlace { + simple_to_prepared: true, + .. + } = self + { + if matches!(message.code(), '1' | '2' | 't' | 'n') { + return AfterExecutionAction::Drop; + } + } + + AfterExecutionAction::Forward + } } impl RewritePlan { @@ -117,6 +165,9 @@ impl RewritePlan { /// Apply the rewrite plan to a ClientRequest. pub(crate) fn apply(&self, request: &mut ClientRequest) -> Result { + // This needs to run first! + let simple_to_prepared = self.simple_to_prepared.apply(request); + // Prepend any required Prepare messages for EXECUTE statements. if !self.prepares.is_empty() { let prepends: Vec = self @@ -157,6 +208,7 @@ impl RewritePlan { Ok(RewriteResult::InPlace { offset: self.offset.clone(), + simple_to_prepared, }) } } 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 30a73288a..7cd630e6c 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/simple_prepared.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/simple_prepared.rs @@ -148,6 +148,7 @@ mod tests { db_schema: &self.db_schema, user: "", search_path: None, + multiple_statements: false, }); let mut plan = Default::default(); let ast = pg_raw_parse::make::try_owned(|mem| { diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/simple_to_prepared.rs b/pgdog/src/frontend/router/parser/rewrite/statement/simple_to_prepared.rs new file mode 100644 index 000000000..c3a0a32ff --- /dev/null +++ b/pgdog/src/frontend/router/parser/rewrite/statement/simple_to_prepared.rs @@ -0,0 +1,364 @@ +use crate::{ + frontend::ClientRequest, + net::{ + Describe, Execute, Parse, ProtocolMessage, Sync, + bind::{Bind, Parameter}, + }, +}; +use pg_raw_parse::{ + ConstValue, Node, NodeMut, + make::{MemoryToken, Unique}, + nodes, + transform::{self, Assignable, Transform}, +}; + +use super::*; + +#[derive(Default, Clone, Debug)] +pub(crate) struct SimpleToPreparedPlan { + /// Parameters using text encoding. + pub(crate) params: Vec, + + pub(crate) step_two: Option, +} + +#[derive(Default, Clone, Debug)] +pub(crate) struct SimpleToPreparedPlanStepTwo { + pub(super) bind: Bind, + pub(super) parse: Parse, +} + +impl SimpleToPreparedPlan { + /// This step runs after all other rewriters are done. + /// + /// This is to ensure we cache the prepared statement after all rewrites are complete. + /// + pub(super) fn step_two( + &mut self, + prepared_statements: &mut PreparedStatements, + stmt: &str, + ) -> Result<(), Error> { + if self.params.is_empty() { + return Ok(()); + } + + let mut parse = Parse::new_anonymous(stmt); + let name = prepared_statements.insert_rewritten_simple_to_prepared(&parse); + parse.rename(&name); + + let parse = Parse::named(&name, stmt); + let bind = Bind::new_params(&name, &self.params); + + self.step_two = Some(SimpleToPreparedPlanStepTwo { parse, bind }); + + Ok(()) + } + + /// Rewrite the request from simple protocol to prepared. + /// + /// INVARIANT: the request contains one single [`crate::net::Query`] message. + /// This is enforced by [`crate::frontend::client::Client::buffer`] and [`ClientRequest::is_complete`]. + /// + pub(crate) fn apply(&self, request: &mut ClientRequest) -> bool { + if let Some(ref step_two) = self.step_two { + request.clear(); + request.push(ProtocolMessage::Parse(step_two.parse.clone())); + request.push(ProtocolMessage::Describe(Describe::new_statement( + step_two.parse.name(), + ))); + request.push(ProtocolMessage::Bind(step_two.bind.clone())); + request.push(ProtocolMessage::Execute(Execute::new())); + request.push(ProtocolMessage::Sync(Sync)); + true + } else { + false + } + } +} + +impl StatementRewrite<'_> { + /// Rewrite a simple query protocol request to a prepared one. + /// + /// # Example + /// + /// ```sql + /// SELECT * FROM users WHERE id = 1; + /// ``` + /// + /// becomes + /// + /// ```sql + /// SELECT * FROM users WHERE id = $1; + /// ``` + /// + /// This also returns parameters extracted from the query, e.g.: + /// + /// ```no_compile + /// vec![ + /// Parameter { + /// data: b"1", + /// len: 1, + /// } + /// ] + /// ``` + /// + pub(super) fn rewrite_simple_to_prepared<'a>( + &mut self, + node: NodeMut<'a, '_>, + mem: MemoryToken<'a>, + plan: &mut RewritePlan, + ) -> Result<(), Error> { + // Only rewrite simple statements. + if self.extended || self.prepared || !self.schema.rewrite.simple_to_prepared { + return Ok(()); + } + + let simple_plan = rewrite_literals(node, mem); + + if !simple_plan.params.is_empty() { + self.rewritten = true; + plan.simple_to_prepared = simple_plan; + } + + Ok(()) + } +} + +/// Replaces constants in executable expressions while leaving constants which +/// are part of SQL syntax (such as the precision and scale in `numeric(5, 2)`) +/// untouched. +fn rewrite_literals<'a>(node: NodeMut<'a, '_>, mem: MemoryToken<'a>) -> SimpleToPreparedPlan { + if !matches!(&node, NodeMut::SelectStmt(_)) { + return SimpleToPreparedPlan::default(); + } + + let mut rewriter = LiteralRewriter { + mem, + params: Vec::new(), + next_param: 1, + }; + transform::transform_node(node, &mut rewriter); + + SimpleToPreparedPlan { + params: rewriter.params, + ..Default::default() + } +} + +struct LiteralRewriter<'mem> { + mem: MemoryToken<'mem>, + params: Vec, + next_param: i32, +} + +impl<'mem> LiteralRewriter<'mem> { + fn parameter(value: Option>) -> Option<(Parameter, &'static str)> { + match value { + // An untyped NULL gets its type from the surrounding expression. + // Casting it to an arbitrary type can make otherwise valid + // expressions fail (for example, `integer IS DISTINCT FROM NULL`). + None => None, + Some(ConstValue::Integer(value)) => Some(( + Parameter::new(itoa::Buffer::new().format(value).as_bytes()), + "int8", + )), + Some(ConstValue::Float(value)) => Some(( + Parameter::new(value.as_bytes()), + if value.parse::().is_ok() { + "int8" + } else { + "numeric" + }, + )), + Some(ConstValue::String(value)) => Some(( + Parameter::new(value.as_bytes()), + if value.parse::().is_ok() { + "uuid" + } else { + "text" + }, + )), + Some(ConstValue::BitString(value)) => Some((Parameter::new(value.as_bytes()), "bit")), + Some(ConstValue::Boolean(value)) => Some(( + Parameter::new(if value { b"true" } else { b"false" }), + "bool", + )), + Some(_) => None, + } + } + + fn replacement( + &mut self, + value: Option>, + explicit_type: bool, + ) -> Option>> { + let (parameter, parameter_type) = Self::parameter(value)?; + + self.params.push(parameter); + let parameter = self.mem.make_param_ref(self.next_param).uncast(); + let replacement = if explicit_type { + parameter + } else { + let type_name = if parameter_type == "text" { + self.mem + .make_list(&[self.mem.make_string(Some(parameter_type))]) + } else { + self.mem.make_list(&[ + self.mem.make_string(Some("pg_catalog")), + self.mem.make_string(Some(parameter_type)), + ]) + }; + self.mem.make_type_cast(parameter, type_name).uncast() + }; + self.next_param += 1; + Some(replacement) + } +} + +impl<'mem> Transform<'mem> for LiteralRewriter<'mem> { + fn transform_node<'mutref>(&mut self, node: Assignable<'mem, 'mutref>) { + let replacement = match &*node { + NodeMut::A_Const(constant) => self.replacement(constant.val(), false), + _ => None, + }; + + if let Some(replacement) = replacement { + node.replace(replacement); + } else { + transform::transform_node(node.into_inner(), self); + } + } + + fn transform_type_cast<'mutref>(&mut self, mut node: nodes::TypeCastMut<'mem, 'mutref>) { + let replacement = match node.arg() { + Node::A_Const(constant) => self.replacement(constant.val(), true), + _ => None, + }; + + if let Some(replacement) = replacement { + node.set_arg(replacement); + } else { + transform::transform_type_cast(node, self); + } + } + + // Type modifiers are represented as A_Const nodes too, but replacing them + // would produce invalid SQL (`numeric($1, $2)`). The value being cast is + // reached through TypeCast.arg and is still rewritten normally. + fn transform_type_name<'mutref>(&mut self, _node: nodes::TypeNameMut<'mem, 'mutref>) {} +} + +#[cfg(test)] +mod tests { + use super::*; + + fn rewrite(sql: &str) -> (String, Vec) { + let parsed = pg_raw_parse::parse(sql).expect("test query should parse"); + let mut params = Vec::new(); + let rewritten = pg_raw_parse::make::owned(|mem| { + let mut copy = mem.make_unique(&*parsed.into_inner()); + let mut stmt = copy + .as_mut() + .into_iter() + .next() + .expect("test query should contain a statement"); + let plan = rewrite_literals(stmt.stmt_mut(), mem); + params = plan.params; + copy + }); + let sql = pg_raw_parse::deparse_stmts(&*rewritten) + .expect("rewritten query should deparse") + .as_str() + .to_owned(); + (sql, params) + } + + #[test] + fn rewrites_constants_in_parameter_order() { + let (sql, params) = rewrite("SELECT 42, 'hello', true, NULL, 1.25"); + + assert_eq!( + sql, + "SELECT $1::bigint, $2::text, $3::boolean, NULL, $4::numeric" + ); + assert_eq!(params.len(), 4); + assert_eq!(params[0].data.as_ref(), b"42"); + assert_eq!(params[1].data.as_ref(), b"hello"); + assert_eq!(params[2].data.as_ref(), b"true"); + assert_eq!(params[3].data.as_ref(), b"1.25"); + } + + #[test] + fn leaves_untyped_nulls_in_place() { + let (sql, params) = rewrite("SELECT 1 IS DISTINCT FROM NULL, NULL::integer"); + + assert_eq!(sql, "SELECT $1::bigint IS DISTINCT FROM NULL, NULL::int"); + assert_eq!(params.len(), 1); + assert_eq!(params[0].data.as_ref(), b"1"); + } + + #[test] + fn leaves_cast_type_modifiers_in_place() { + let (sql, params) = rewrite("SELECT 5::numeric(10, 2)"); + + assert_eq!(sql, "SELECT $1::numeric(10, 2)"); + assert_eq!(params.len(), 1); + assert_eq!(params[0].data.as_ref(), b"5"); + } + + #[test] + fn rewrites_nested_expressions() { + let (sql, params) = rewrite( + "SELECT * FROM users WHERE id = 7 AND name IN ('alice', 'bob') LIMIT 10 OFFSET 2", + ); + + assert_eq!( + sql, + "SELECT * FROM users WHERE id = $1::bigint AND name IN ($2::text, $3::text) LIMIT $5::bigint OFFSET $4::bigint" + ); + let params: Vec<_> = params + .iter() + .map(|parameter| parameter.data.as_ref()) + .collect(); + assert_eq!( + params, + [ + b"7" as &[u8], + b"alice" as &[u8], + b"bob" as &[u8], + b"2" as &[u8], + b"10" as &[u8] + ] + ); + } + + #[test] + fn does_not_rewrite_non_dml_statements() { + let (sql, params) = rewrite("CREATE TABLE measurements (value numeric(10, 2) DEFAULT 5)"); + + assert_eq!( + sql, + "CREATE TABLE measurements (value numeric(10, 2) DEFAULT 5)" + ); + assert!(params.is_empty()); + + let (sql, params) = rewrite("EXPLAIN SELECT 5"); + + assert_eq!(sql, "EXPLAIN SELECT 5"); + assert!(params.is_empty()); + } + + #[test] + fn does_not_rewrite_writes() { + for statement in [ + "INSERT INTO measurements (value) VALUES (5)", + "UPDATE measurements SET value = 5", + "DELETE FROM measurements WHERE value = 5", + ] { + let (sql, params) = rewrite(statement); + + assert_eq!(sql, statement); + assert!(params.is_empty()); + } + } +} 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..c56a9b678 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/unique_id.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/unique_id.rs @@ -288,6 +288,7 @@ mod tests { db_schema: &db_schema, user: "", search_path: None, + multiple_statements: false, }); let mut plan = Default::default(); let ast = make::owned(|mem| { diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/update.rs b/pgdog/src/frontend/router/parser/rewrite/statement/update.rs index 3a1f91416..9fb03990b 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/update.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/update.rs @@ -471,6 +471,7 @@ mod test { prepared_statements: &mut stmts, user: "", search_path: None, + multiple_statements: false, }; let mut plan = RewritePlan::default(); StatementRewrite::new(ctx).sharding_key_update(