From e05097ecb50eda37b45ecf7a7049e07805c36a1f Mon Sep 17 00:00:00 2001 From: Lev Kokotov Date: Tue, 11 Aug 2026 15:05:19 -0700 Subject: [PATCH 1/9] feat: rewrite simple to prepared --- .../parser/rewrite/statement/auto_id.rs | 26 +- .../router/parser/rewrite/statement/error.rs | 3 + .../router/parser/rewrite/statement/mod.rs | 21 +- .../router/parser/rewrite/statement/plan.rs | 43 +-- .../rewrite/statement/simple_prepared.rs | 4 +- .../rewrite/statement/simple_to_prepared.rs | 275 ++++++++++++++++++ .../parser/rewrite/statement/unique_id.rs | 38 +-- 7 files changed, 353 insertions(+), 57 deletions(-) create mode 100644 pgdog/src/frontend/router/parser/rewrite/statement/simple_to_prepared.rs 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..495aca476 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/auto_id.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/auto_id.rs @@ -69,7 +69,7 @@ impl StatementRewrite<'_> { let replaced = self.replace_set_to_default_at_positions(&mut node, mem, &present_pk_positions); if replaced > 0 { - plan.auto_id_injected += replaced as u16; + plan.num_auto_id_injected += replaced as u16; self.rewritten = true; } } @@ -85,7 +85,7 @@ impl StatementRewrite<'_> { if rewrite { for column in missing_columns { self.inject_column_with_unique_id(&mut node, mem, column); - plan.auto_id_injected += 1; + plan.num_auto_id_injected += 1; } self.rewritten = true; } @@ -311,8 +311,8 @@ mod tests { ) .unwrap(); - assert_eq!(plan.auto_id_injected, 1); - assert_eq!(plan.unique_ids, 1); // confirms unique_id was processed + assert_eq!(plan.num_auto_id_injected, 1); + assert_eq!(plan.num_unique_ids, 1); // confirms unique_id was processed assert!(sql.contains("id")); // pgdog.unique_id() should be replaced with actual bigint value assert!(!sql.contains("pgdog.unique_id")); @@ -347,7 +347,7 @@ mod tests { ) .unwrap(); - assert_eq!(plan.auto_id_injected, 0); + assert_eq!(plan.num_auto_id_injected, 0); assert!(!sql.contains("id,")); } @@ -361,7 +361,7 @@ mod tests { ) .unwrap(); - assert_eq!(plan.auto_id_injected, 0); + assert_eq!(plan.num_auto_id_injected, 0); assert!(!sql.contains("pgdog.unique_id")); } @@ -375,7 +375,7 @@ mod tests { ) .unwrap(); - assert_eq!(plan.auto_id_injected, 0); + assert_eq!(plan.num_auto_id_injected, 0); assert!(!sql.contains("pgdog.unique_id")); } @@ -389,7 +389,7 @@ mod tests { ) .unwrap(); - assert_eq!(plan.auto_id_injected, 0); + assert_eq!(plan.num_auto_id_injected, 0); } #[test] @@ -402,7 +402,7 @@ mod tests { ) .unwrap(); - assert_eq!(plan.auto_id_injected, 1); + assert_eq!(plan.num_auto_id_injected, 1); assert!(sql.contains("id")); } @@ -431,7 +431,7 @@ mod tests { // DEFAULT should be replaced with unique_id assert!(!sql.to_uppercase().contains("DEFAULT")); assert!(sql.contains("::bigint")); // value is cast to bigint - assert_eq!(plan.unique_ids, 1); + assert_eq!(plan.num_unique_ids, 1); } #[test] @@ -446,7 +446,7 @@ mod tests { // Both DEFAULT values should be replaced assert!(!sql.to_uppercase().contains("DEFAULT")); - assert_eq!(plan.unique_ids, 2); + assert_eq!(plan.num_unique_ids, 2); } #[test] @@ -523,7 +523,7 @@ mod tests { .unwrap(); // users is sharded, so RewriteOmni should NOT inject auto id - assert_eq!(plan.auto_id_injected, 0); + assert_eq!(plan.num_auto_id_injected, 0); assert!(!sql.contains("::bigint")); } @@ -547,7 +547,7 @@ mod tests { .unwrap(); // users is NOT sharded, so RewriteOmni should inject auto id - assert_eq!(plan.auto_id_injected, 1); + assert_eq!(plan.num_auto_id_injected, 1); assert!(sql.contains("::bigint")); } } diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/error.rs b/pgdog/src/frontend/router/parser/rewrite/statement/error.rs index 18e81a7b5..4ee13f5ed 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/error.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/error.rs @@ -40,4 +40,7 @@ pub enum Error { #[error("prepared statement '{0}' does not exist")] ExecuteMissingPrepare(String), + + #[error("prepared statement: {0}")] + PreparedStmt(#[from] crate::frontend::prepared_statements::Error), } diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs b/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs index 8ef1105d2..6d8c14adf 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. @@ -104,13 +106,17 @@ 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(_) | Node::UpdateStmt(_) | Node::DeleteStmt(_) => walk::walk(stmt.stmt(), |node| { if let Node::ParamRef(param) = node { - plan.params = plan.params.max(param.number as u16) + plan.num_params = plan.num_params.max(param.number as u16) } }), Node::PrepareStmt(_) | Node::ExecuteStmt(_) | Node::ExplainStmt(_) => {} @@ -133,14 +139,14 @@ impl<'a> StatementRewrite<'a> { } // Track the next parameter number to use - let mut next_param = plan.params as i32 + 1; + let mut next_param = plan.num_params as i32 + 1; let mut err = None; transform::transform_node( stmt.stmt_mut(), &mut transform::TransformClosure::new(|node| { match Self::rewrite_unique_id(node.as_ref(), mem, self.extended, &mut next_param) { Ok(Some(replacement)) => { - plan.unique_ids += 1; + plan.num_unique_ids += 1; self.rewritten = true; node.replace(replacement); None @@ -163,7 +169,14 @@ impl<'a> StatementRewrite<'a> { } 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(&mut self.prepared_statements, &stmt)?; + + plan.rewritten_stmt = Some(stmt); } if let Node::InsertStmt(insert) = stmt.stmt() { diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs b/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs index 6b2bbb9e8..1e46953d5 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs @@ -5,7 +5,9 @@ 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. /// @@ -17,16 +19,16 @@ pub struct RewritePlan { /// the original statement. This is calculated first, /// and $params+n parameters are added to the statement to /// substitute values we are rewriting. - pub(crate) params: u16, + pub(crate) num_params: u16, /// Number of unique IDs to append to the Bind message. - pub(crate) unique_ids: u16, + pub(crate) num_unique_ids: u16, /// Number of auto-injected primary key columns with pgdog.unique_id(). - pub(crate) auto_id_injected: u16, + pub(crate) num_auto_id_injected: u16, /// Rewritten SQL statement. - pub(crate) stmt: Option, + pub(crate) rewritten_stmt: Option, /// Prepared statements to prepend to the client request. /// Each tuple contains (name, statement) for ProtocolMessage::Prepare. @@ -46,6 +48,9 @@ pub struct RewritePlan { /// Limit/offset pagination. pub(crate) offset: Option, + + /// Simple to prepared rewrite. + pub(crate) simple_to_prepared: SimpleToPreparedPlan, } #[derive(Debug, Clone)] @@ -71,9 +76,9 @@ impl RewritePlan { /// `params` is purely informational (count of original `$N` placeholders) /// and doesn't count as a rewrite. pub(crate) fn is_empty(&self) -> bool { - self.unique_ids == 0 - && self.auto_id_injected == 0 - && self.stmt.is_none() + self.num_unique_ids == 0 + && self.num_auto_id_injected == 0 + && self.rewritten_stmt.is_none() && self.prepares.is_empty() && self.insert_split.is_empty() && self.aggregates.is_noop() @@ -85,7 +90,7 @@ impl RewritePlan { pub(crate) fn apply_bind(&self, bind: &mut Bind) -> Result<(), Error> { let format = bind.default_param_format(); - for _ in 0..self.unique_ids { + for _ in 0..self.num_unique_ids { let generator = UniqueId::generator()?; let id = generator.next_id(); let param = match format { @@ -100,7 +105,7 @@ impl RewritePlan { /// Apply the rewrite plan to a Parse message by updating the SQL. pub(crate) fn apply_parse(&self, parse: &mut Parse) { - if let Some(ref stmt) = self.stmt { + if let Some(ref stmt) = self.rewritten_stmt { parse.set_query(stmt); if !parse.anonymous() { PreparedStatements::global().write().rewrite(parse); @@ -110,7 +115,7 @@ impl RewritePlan { /// Apply the rewrite plan to a Query message by updating the SQL. pub(crate) fn apply_query(&self, query: &mut Query) { - if let Some(ref stmt) = self.stmt { + if let Some(ref stmt) = self.rewritten_stmt { query.set_query(stmt); } } @@ -180,7 +185,7 @@ mod tests { fn test_apply_bind_text_format() { let _guard = set_env_var("NODE_ID", "pgdog-1"); let plan = RewritePlan { - unique_ids: 1, + num_unique_ids: 1, ..Default::default() }; let mut bind = Bind::default(); @@ -200,8 +205,8 @@ mod tests { fn test_apply_bind_binary_format_uniform() { let _guard = set_env_var("NODE_ID", "pgdog-1"); let plan = RewritePlan { - params: 1, - unique_ids: 1, + num_params: 1, + num_unique_ids: 1, ..Default::default() }; // Create bind with uniform binary format (1 code applies to all) @@ -225,8 +230,8 @@ mod tests { fn test_apply_bind_binary_format_one_to_one() { let _guard = set_env_var("NODE_ID", "pgdog-1"); let plan = RewritePlan { - params: 2, - unique_ids: 1, + num_params: 2, + num_unique_ids: 1, ..Default::default() }; // Create bind with one-to-one format codes @@ -252,7 +257,7 @@ mod tests { fn test_apply_bind_multiple_unique_ids() { let _guard = set_env_var("NODE_ID", "pgdog-1"); let plan = RewritePlan { - unique_ids: 3, + num_unique_ids: 3, ..Default::default() }; let mut bind = Bind::default(); @@ -272,8 +277,8 @@ mod tests { fn test_apply_bind_appends_to_existing_params() { let _guard = set_env_var("NODE_ID", "pgdog-1"); let plan = RewritePlan { - params: 2, - unique_ids: 2, + num_params: 2, + num_unique_ids: 2, ..Default::default() }; let mut bind = Bind::new_params( 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..3d1c6d2b5 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/simple_prepared.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/simple_prepared.rs @@ -174,7 +174,7 @@ mod tests { "original name should be replaced: {sql}" ); assert!(plan.prepares.is_empty()); - assert!(plan.stmt.is_some()); + assert!(plan.rewritten_stmt.is_some()); } #[test] @@ -225,6 +225,6 @@ mod tests { assert_eq!(sql, "SELECT 1, 2, 3"); assert!(plan.prepares.is_empty()); - assert!(plan.stmt.is_none()); + assert!(plan.rewritten_stmt.is_none()); } } 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..f586a1cdf --- /dev/null +++ b/pgdog/src/frontend/router/parser/rewrite/statement/simple_to_prepared.rs @@ -0,0 +1,275 @@ +use crate::{ + frontend::ClientRequest, + net::{ + Parse, ProtocolMessage, + bind::{Bind, Parameter}, + }, + util::random_string, +}; +use pg_raw_parse::{ + ConstValue, NodeMut, + make::MemoryToken, + 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: SimpleToPreparedPlanStepTwo, +} + +#[derive(Default, Clone, Debug)] +pub(crate) struct SimpleToPreparedPlanStepTwo { + pub(crate) bind: Bind, + pub(crate) parse: Parse, +} + +impl SimpleToPreparedPlan { + 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) = PreparedStatements::global().write().insert(&parse); + parse.rename(&name); + + let parse = Parse::named(&name, stmt); + let bind = Bind::new_params(&name, &self.params); + + // This will ensure the global counter for this prepared statement + // is correctly decreased when this client disconnects. + // + // We are using the global name for the local cache because + // the client doesn't know it's using prepared statements and will never + // manually close it. + prepared_statements.insert_local_mapping(parse.name(), parse.name()); + + self.step_two = SimpleToPreparedPlanStepTwo { parse, bind }; + + Ok(()) + } + + pub(crate) fn apply(&self, request: &mut ClientRequest) {} +} + +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 { + 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(_) + | NodeMut::InsertStmt(_) + | NodeMut::UpdateStmt(_) + | NodeMut::DeleteStmt(_) + ) { + 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 LiteralRewriter<'_> { + fn parameter(value: Option>) -> Option { + match value { + None => Some(Parameter::new_null()), + Some(ConstValue::Integer(value)) => { + Some(Parameter::new(itoa::Buffer::new().format(value).as_bytes())) + } + Some(ConstValue::Float(value)) + | Some(ConstValue::String(value)) + | Some(ConstValue::BitString(value)) => Some(Parameter::new(value.as_bytes())), + Some(ConstValue::Boolean(value)) => { + Some(Parameter::new(if value { b"true" } else { b"false" })) + } + Some(_) => None, + } + } +} + +impl<'mem> Transform<'mem> for LiteralRewriter<'mem> { + fn transform_node<'mutref>(&mut self, node: Assignable<'mem, 'mutref>) { + let parameter = match &*node { + NodeMut::A_Const(constant) => Self::parameter(constant.val()), + _ => None, + }; + + if let Some(parameter) = parameter { + self.params.push(parameter); + node.replace(self.mem.make_param_ref(self.next_param).uncast()); + self.next_param += 1; + } else { + transform::transform_node(node.into_inner(), 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, $2, $3, $4, $5"); + assert_eq!(params.len(), 5); + 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].len, -1); + assert_eq!(params[4].data.as_ref(), b"1.25"); + } + + #[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 AND name IN ($2, $3) LIMIT $5 OFFSET $4" + ); + 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()); + } +} 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..356832179 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/unique_id.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/unique_id.rs @@ -126,8 +126,8 @@ mod tests { let (sql, plan) = run_test("SELECT pgdog.unique_id()", true); assert_eq!(sql, "SELECT $1::bigint"); - assert_eq!(plan.params, 0); - assert_eq!(plan.unique_ids, 1); + assert_eq!(plan.num_params, 0); + assert_eq!(plan.num_unique_ids, 1); } #[test] @@ -135,8 +135,8 @@ mod tests { let (sql, plan) = run_test("SELECT pgdog.unique_id(), $1, $2", true); assert_eq!(sql, "SELECT $3::bigint, $1, $2"); - assert_eq!(plan.params, 2); - assert_eq!(plan.unique_ids, 1); + assert_eq!(plan.num_params, 2); + assert_eq!(plan.num_unique_ids, 1); } #[test] @@ -144,8 +144,8 @@ mod tests { let (sql, plan) = run_test("SELECT pgdog.unique_id(), pgdog.unique_id()", true); assert_eq!(sql, "SELECT $1::bigint, $2::bigint"); - assert_eq!(plan.params, 0); - assert_eq!(plan.unique_ids, 2); + assert_eq!(plan.num_params, 0); + assert_eq!(plan.num_unique_ids, 2); } #[test] @@ -157,8 +157,8 @@ mod tests { !sql.contains("pgdog.unique_id"), "Function should be replaced: {sql}" ); - assert_eq!(plan.params, 0); - assert_eq!(plan.unique_ids, 1); + assert_eq!(plan.num_params, 0); + assert_eq!(plan.num_unique_ids, 1); } #[test] @@ -171,7 +171,7 @@ mod tests { !sql.contains("pgdog.unique_id"), "Functions should be replaced: {sql}" ); - assert_eq!(plan.unique_ids, 2); + assert_eq!(plan.num_unique_ids, 2); } #[test] @@ -179,7 +179,7 @@ mod tests { let (sql, plan) = run_test("SELECT 1, 2, 3", true); assert_eq!(sql, "SELECT 1, 2, 3"); - assert_eq!(plan.unique_ids, 0); + assert_eq!(plan.num_unique_ids, 0); } #[test] @@ -190,7 +190,7 @@ mod tests { ); assert_eq!(sql, "INSERT INTO t (id, name) VALUES ($1::bigint, 'test')"); - assert_eq!(plan.unique_ids, 1); + assert_eq!(plan.num_unique_ids, 1); } #[test] @@ -201,7 +201,7 @@ mod tests { ); assert_eq!(sql, "INSERT INTO t (id) VALUES ($1::bigint), ($2::bigint)"); - assert_eq!(plan.unique_ids, 2); + assert_eq!(plan.num_unique_ids, 2); } #[test] @@ -209,7 +209,7 @@ mod tests { let (sql, plan) = run_test("INSERT INTO t (id) SELECT pgdog.unique_id() FROM s", true); assert_eq!(sql, "INSERT INTO t (id) SELECT $1::bigint FROM s"); - assert_eq!(plan.unique_ids, 1); + assert_eq!(plan.num_unique_ids, 1); } #[test] @@ -220,7 +220,7 @@ mod tests { ); assert_eq!(sql, "UPDATE t SET id = $1::bigint WHERE name = 'test'"); - assert_eq!(plan.unique_ids, 1); + assert_eq!(plan.num_unique_ids, 1); } #[test] @@ -231,7 +231,7 @@ mod tests { ); assert_eq!(sql, "UPDATE t SET name = 'new' WHERE id = $1::bigint"); - assert_eq!(plan.unique_ids, 1); + assert_eq!(plan.num_unique_ids, 1); } #[test] @@ -239,7 +239,7 @@ mod tests { let (sql, plan) = run_test("DELETE FROM t WHERE id = pgdog.unique_id()", true); assert_eq!(sql, "DELETE FROM t WHERE id = $1::bigint"); - assert_eq!(plan.unique_ids, 1); + assert_eq!(plan.num_unique_ids, 1); } #[test] @@ -253,7 +253,7 @@ mod tests { sql, "INSERT INTO t (id) VALUES ($1::bigint) RETURNING $2::bigint" ); - assert_eq!(plan.unique_ids, 2); + assert_eq!(plan.num_unique_ids, 2); } #[test] @@ -264,7 +264,7 @@ mod tests { ); assert_eq!(sql, "EXPLAIN INSERT INTO t (id) SELECT $1::bigint FROM s"); - assert_eq!(plan.unique_ids, 1); + assert_eq!(plan.num_unique_ids, 1); } #[test] @@ -272,7 +272,7 @@ mod tests { let (sql, plan) = run_test("EXPLAIN SELECT pgdog.unique_id()", true); assert_eq!(sql, "EXPLAIN SELECT $1::bigint"); - assert_eq!(plan.unique_ids, 1); + assert_eq!(plan.num_unique_ids, 1); } fn run_test(sql: &str, extended: bool) -> (String, RewritePlan) { From 5a1cfaf58ca61c78c84d25768a8c9ba8739ced67 Mon Sep 17 00:00:00 2001 From: Lev Kokotov Date: Tue, 11 Aug 2026 15:17:29 -0700 Subject: [PATCH 2/9] sae --- pgdog/src/backend/prepared_statements.rs | 4 + pgdog/src/backend/server.rs | 6 + .../frontend/client/query_engine/request.rs | 138 ++++++++++++++++++ pgdog/src/frontend/client_request.rs | 6 + pgdog/src/frontend/prepared_statements/mod.rs | 14 +- .../router/parser/rewrite/statement/plan.rs | 3 + .../rewrite/statement/simple_to_prepared.rs | 23 ++- 7 files changed, 188 insertions(+), 6 deletions(-) create mode 100644 pgdog/src/frontend/client/query_engine/request.rs diff --git a/pgdog/src/backend/prepared_statements.rs b/pgdog/src/backend/prepared_statements.rs index dde70b9d4..7b35fd36c 100644 --- a/pgdog/src/backend/prepared_statements.rs +++ b/pgdog/src/backend/prepared_statements.rs @@ -437,6 +437,10 @@ impl PreparedStatements { &self.state } + pub(crate) fn ignore_message(&mut self, code: impl Into) { + self.state.add_ignore(code); + } + /// Get mutable reference to protocol state. pub fn state_mut(&mut self) -> &mut ProtocolState { &mut self.state diff --git a/pgdog/src/backend/server.rs b/pgdog/src/backend/server.rs index cf3132398..203aae073 100644 --- a/pgdog/src/backend/server.rs +++ b/pgdog/src/backend/server.rs @@ -476,6 +476,12 @@ impl Server { } self.flush().await?; + if client_request.simple_to_prepared_rewrite { + self.prepared_statements.ignore_message('1'); + self.prepared_statements.ignore_message('2'); + self.prepared_statements.ignore_message('t'); + } + // The whole request is now in server's hands. // We can recover the connection from this point on. self.sending_request = false; diff --git a/pgdog/src/frontend/client/query_engine/request.rs b/pgdog/src/frontend/client/query_engine/request.rs new file mode 100644 index 000000000..6a868437b --- /dev/null +++ b/pgdog/src/frontend/client/query_engine/request.rs @@ -0,0 +1,138 @@ +use std::ops::{Deref, DerefMut}; + +use crate::frontend::{ + ClientRequest, + router::{Ast, Route}, +}; + +#[derive(Debug, Clone, thiserror::Error)] +pub(crate) enum Error { + #[error("only initial requests can be parsed")] + InitialToParsed, + + #[error("request was not parsed")] + NoAst, + + #[error("requested was not routed yet")] + NoRoute, +} + +pub(crate) enum EngineRequest<'a> { + Initial(&'a mut ClientRequest), + Parsed(ParsedEngineRequest<'a>), + Routed(RoutedEngineRequest<'a>), + Transitioning, +} + +impl EngineRequest<'_> { + pub(crate) fn into_parsed(&mut self, ast: Option) -> Result<(), Error> { + let request = match std::mem::replace(self, Self::Transitioning) { + Self::Initial(request) => request, + request => { + *self = request; + return Err(Error::InitialToParsed); + } + }; + + *self = Self::Parsed(ParsedEngineRequest { ast, request }); + Ok(()) + } + + pub(crate) fn into_routed(&mut self, route: Route) -> Result<(), Error> { + let request = match std::mem::replace(self, Self::Transitioning) { + Self::Parsed(request) => request, + request => { + *self = request; + return Err(Error::NoAst); + } + }; + + *self = Self::Routed(RoutedEngineRequest { route, request }); + Ok(()) + } + + pub(crate) fn ast(&self) -> Result<&Option, Error> { + match self { + Self::Initial(_) => Err(Error::NoAst), + Self::Parsed(parsed) => Ok(&parsed.ast), + Self::Routed(routed) => Ok(&routed.request.ast), + Self::Transitioning => unreachable!(), + } + } + + pub(crate) fn route(&self) -> Result<&Route, Error> { + match self { + Self::Routed(routed) => Ok(&routed.route), + _ => Err(Error::NoRoute), + } + } + + pub(crate) fn route_mut(&mut self) -> Result<&mut Route, Error> { + match self { + Self::Routed(routed) => Ok(&mut routed.route), + _ => Err(Error::NoRoute), + } + } +} + +impl Deref for EngineRequest<'_> { + type Target = ClientRequest; + + fn deref(&self) -> &Self::Target { + match self { + Self::Initial(req) => req, + Self::Parsed(parsed) => parsed.deref(), + Self::Routed(routed) => routed.deref(), + Self::Transitioning => unreachable!(), + } + } +} + +impl DerefMut for EngineRequest<'_> { + fn deref_mut(&mut self) -> &mut Self::Target { + match self { + Self::Initial(req) => req, + Self::Parsed(parsed) => parsed.deref_mut(), + Self::Routed(routed) => routed.deref_mut(), + Self::Transitioning => unreachable!(), + } + } +} + +pub(crate) struct ParsedEngineRequest<'a> { + pub(crate) ast: Option, + pub(crate) request: &'a mut ClientRequest, +} + +pub(crate) struct RoutedEngineRequest<'a> { + pub(crate) route: Route, + pub(crate) request: ParsedEngineRequest<'a>, +} + +impl Deref for ParsedEngineRequest<'_> { + type Target = ClientRequest; + + fn deref(&self) -> &Self::Target { + self.request + } +} + +impl DerefMut for ParsedEngineRequest<'_> { + fn deref_mut(&mut self) -> &mut Self::Target { + self.request + } +} + +impl Deref for RoutedEngineRequest<'_> { + type Target = ClientRequest; + + fn deref(&self) -> &Self::Target { + self.request.request + } +} + +impl DerefMut for RoutedEngineRequest<'_> { + fn deref_mut(&mut self) -> &mut Self::Target { + self.request.request + } +} diff --git a/pgdog/src/frontend/client_request.rs b/pgdog/src/frontend/client_request.rs index 901b67503..d11ebd1bc 100644 --- a/pgdog/src/frontend/client_request.rs +++ b/pgdog/src/frontend/client_request.rs @@ -33,6 +33,8 @@ pub struct ClientRequest { pub ast: Option, /// Last Parse we received. pub last_parse: Option, + /// Simple to prepared rewrite requires us to drop some messages from the server. + pub(crate) simple_to_prepared_rewrite: bool, } impl MemoryUsage for ClientRequest { @@ -58,6 +60,7 @@ impl ClientRequest { route: None, ast: None, last_parse: None, + simple_to_prepared_rewrite: false, } } @@ -92,6 +95,7 @@ impl ClientRequest { self.messages.clear(); self.route = None; self.ast = None; + self.simple_to_prepared_rewrite = false; } /// We received a complete request and we are ready to @@ -215,6 +219,7 @@ impl ClientRequest { route: self.route.clone(), ast: self.ast.clone(), last_parse: None, + simple_to_prepared_rewrite: self.simple_to_prepared_rewrite, } } @@ -393,6 +398,7 @@ impl From> for ClientRequest { route: None, ast: None, last_parse: None, + simple_to_prepared_rewrite: false, } } } diff --git a/pgdog/src/frontend/prepared_statements/mod.rs b/pgdog/src/frontend/prepared_statements/mod.rs index 7419cceb8..784b65bb9 100644 --- a/pgdog/src/frontend/prepared_statements/mod.rs +++ b/pgdog/src/frontend/prepared_statements/mod.rs @@ -66,8 +66,20 @@ 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_local_mapping(&mut self, local: &str, global: &str) { + self.local.insert(local.to_owned(), global.to_owned()); + } + /// 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/rewrite/statement/plan.rs b/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs index 1e46953d5..9ad928de5 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs @@ -122,6 +122,9 @@ impl RewritePlan { /// Apply the rewrite plan to a ClientRequest. pub(crate) fn apply(&self, request: &mut ClientRequest) -> Result { + // This needs to run first! + self.simple_to_prepared.apply(request); + // Prepend any required Prepare messages for EXECUTE statements. if !self.prepares.is_empty() { let prepends: Vec = self 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 index f586a1cdf..19ebf2e69 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/simple_to_prepared.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/simple_to_prepared.rs @@ -1,10 +1,9 @@ use crate::{ frontend::ClientRequest, net::{ - Parse, ProtocolMessage, + Describe, Execute, Parse, ProtocolMessage, Sync, bind::{Bind, Parameter}, }, - util::random_string, }; use pg_raw_parse::{ ConstValue, NodeMut, @@ -25,8 +24,8 @@ pub(crate) struct SimpleToPreparedPlan { #[derive(Default, Clone, Debug)] pub(crate) struct SimpleToPreparedPlanStepTwo { - pub(crate) bind: Bind, - pub(crate) parse: Parse, + pub(super) bind: Bind, + pub(super) parse: Parse, } impl SimpleToPreparedPlan { @@ -60,7 +59,21 @@ impl SimpleToPreparedPlan { Ok(()) } - pub(crate) fn apply(&self, request: &mut ClientRequest) {} + pub(crate) fn apply(&self, request: &mut ClientRequest) { + if self.params.is_empty() { + return; + } + + request.simple_to_prepared_rewrite = true; + request.clear(); + request.push(ProtocolMessage::Parse(self.step_two.parse.clone())); + request.push(ProtocolMessage::Describe(Describe::new_statement( + self.step_two.parse.name(), + ))); + request.push(ProtocolMessage::Bind(self.step_two.bind.clone())); + request.push(ProtocolMessage::Execute(Execute::new())); + request.push(ProtocolMessage::Sync(Sync)); + } } impl StatementRewrite<'_> { From 05d117d3124dfce4ac4391857b8a19e172d6c6ec Mon Sep 17 00:00:00 2001 From: Lev Kokotov Date: Tue, 11 Aug 2026 16:40:46 -0700 Subject: [PATCH 3/9] save --- pgdog-config/src/rewrite.rs | 5 + pgdog/src/frontend/router/parser/cache/ast.rs | 2 +- .../router/parser/rewrite/statement/mod.rs | 2 +- .../rewrite/statement/simple_to_prepared.rs | 139 +++++++++++++----- 4 files changed, 106 insertions(+), 42 deletions(-) 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/router/parser/cache/ast.rs b/pgdog/src/frontend/router/parser/cache/ast.rs index e6308c0a1..6fb57d29b 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, diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs b/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs index 6d8c14adf..10bdf5471 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs @@ -174,7 +174,7 @@ impl<'a> StatementRewrite<'a> { // 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(&mut self.prepared_statements, &stmt)?; + .step_two(self.prepared_statements, &stmt)?; plan.rewritten_stmt = Some(stmt); } 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 index 19ebf2e69..e2495c471 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/simple_to_prepared.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/simple_to_prepared.rs @@ -6,8 +6,8 @@ use crate::{ }, }; use pg_raw_parse::{ - ConstValue, NodeMut, - make::MemoryToken, + ConstValue, Node, NodeMut, + make::{MemoryToken, Unique}, nodes, transform::{self, Assignable, Transform}, }; @@ -19,7 +19,7 @@ pub(crate) struct SimpleToPreparedPlan { /// Parameters using text encoding. pub(crate) params: Vec, - pub(crate) step_two: SimpleToPreparedPlanStepTwo, + pub(crate) step_two: Option, } #[derive(Default, Clone, Debug)] @@ -54,25 +54,23 @@ impl SimpleToPreparedPlan { // manually close it. prepared_statements.insert_local_mapping(parse.name(), parse.name()); - self.step_two = SimpleToPreparedPlanStepTwo { parse, bind }; + self.step_two = Some(SimpleToPreparedPlanStepTwo { parse, bind }); Ok(()) } pub(crate) fn apply(&self, request: &mut ClientRequest) { - if self.params.is_empty() { - return; + if let Some(ref step_two) = self.step_two { + request.simple_to_prepared_rewrite = true; + 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)); } - - request.simple_to_prepared_rewrite = true; - request.clear(); - request.push(ProtocolMessage::Parse(self.step_two.parse.clone())); - request.push(ProtocolMessage::Describe(Describe::new_statement( - self.step_two.parse.name(), - ))); - request.push(ProtocolMessage::Bind(self.step_two.bind.clone())); - request.push(ProtocolMessage::Execute(Execute::new())); - request.push(ProtocolMessage::Sync(Sync)); } } @@ -109,7 +107,7 @@ impl StatementRewrite<'_> { plan: &mut RewritePlan, ) -> Result<(), Error> { // Only rewrite simple statements. - if self.extended || self.prepared { + if self.extended || self.prepared || !self.schema.rewrite.simple_to_prepared { return Ok(()); } @@ -157,40 +155,90 @@ struct LiteralRewriter<'mem> { next_param: i32, } -impl LiteralRewriter<'_> { - fn parameter(value: Option>) -> Option { +impl<'mem> LiteralRewriter<'mem> { + fn parameter(value: Option>) -> Option<(Parameter, &'static str)> { match value { - None => Some(Parameter::new_null()), - Some(ConstValue::Integer(value)) => { - Some(Parameter::new(itoa::Buffer::new().format(value).as_bytes())) - } - Some(ConstValue::Float(value)) - | Some(ConstValue::String(value)) - | Some(ConstValue::BitString(value)) => Some(Parameter::new(value.as_bytes())), - Some(ConstValue::Boolean(value)) => { - Some(Parameter::new(if value { b"true" } else { b"false" })) - } + // 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()), "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 parameter = match &*node { - NodeMut::A_Const(constant) => Self::parameter(constant.val()), + let replacement = match &*node { + NodeMut::A_Const(constant) => self.replacement(constant.val(), false), _ => None, }; - if let Some(parameter) = parameter { - self.params.push(parameter); - node.replace(self.mem.make_param_ref(self.next_param).uncast()); - self.next_param += 1; + 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. @@ -226,13 +274,24 @@ mod tests { fn rewrites_constants_in_parameter_order() { let (sql, params) = rewrite("SELECT 42, 'hello', true, NULL, 1.25"); - assert_eq!(sql, "SELECT $1, $2, $3, $4, $5"); - assert_eq!(params.len(), 5); + 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].len, -1); - assert_eq!(params[4].data.as_ref(), b"1.25"); + 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] @@ -252,7 +311,7 @@ mod tests { assert_eq!( sql, - "SELECT * FROM users WHERE id = $1 AND name IN ($2, $3) LIMIT $5 OFFSET $4" + "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() From 36b5fede9cee1889d41d32ab80acdabbaf82a489 Mon Sep 17 00:00:00 2001 From: Lev Kokotov Date: Wed, 12 Aug 2026 13:09:08 -0700 Subject: [PATCH 4/9] context --- pgdog/src/frontend/router/parser/cache/ast.rs | 2 ++ .../src/frontend/router/parser/cache/test.rs | 25 +++++++++++++++++++ .../parser/rewrite/statement/auto_id.rs | 1 + .../router/parser/rewrite/statement/error.rs | 3 +++ .../router/parser/rewrite/statement/insert.rs | 1 + .../router/parser/rewrite/statement/mod.rs | 9 +++++++ .../router/parser/rewrite/statement/offset.rs | 1 + .../rewrite/statement/simple_prepared.rs | 1 + .../rewrite/statement/simple_to_prepared.rs | 11 ++++++++ .../parser/rewrite/statement/unique_id.rs | 1 + .../router/parser/rewrite/statement/update.rs | 1 + 11 files changed, 56 insertions(+) diff --git a/pgdog/src/frontend/router/parser/cache/ast.rs b/pgdog/src/frontend/router/parser/cache/ast.rs index 6fb57d29b..58921054f 100644 --- a/pgdog/src/frontend/router/parser/cache/ast.rs +++ b/pgdog/src/frontend/router/parser/cache/ast.rs @@ -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 495aca476..65e2b4483 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 4ee13f5ed..f5c40be9b 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/error.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/error.rs @@ -43,4 +43,7 @@ pub enum Error { #[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 10bdf5471..262d57e11 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs @@ -44,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)] @@ -67,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> { @@ -84,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, } } @@ -168,6 +173,10 @@ impl<'a> StatementRewrite<'a> { self.limit_offset(&select, &mut plan); } + if self.rewritten && self.multiple_statements { + return Err(Error::MultiStatementRewrite); + } + if self.rewritten { let stmt = pg_raw_parse::deparse(&*stmt)?.as_str().to_owned(); 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/simple_prepared.rs b/pgdog/src/frontend/router/parser/rewrite/statement/simple_prepared.rs index 3d1c6d2b5..c10b57a11 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 index e2495c471..1bb57e1a3 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/simple_to_prepared.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/simple_to_prepared.rs @@ -29,6 +29,10 @@ pub(crate) struct SimpleToPreparedPlanStepTwo { } 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, @@ -40,6 +44,8 @@ impl SimpleToPreparedPlan { let mut parse = Parse::new_anonymous(stmt); + // This is what QueryEngine::rewrite_extended does, + // with a small change, see insert_local_mapping below. let (_, name) = PreparedStatements::global().write().insert(&parse); parse.rename(&name); @@ -59,6 +65,11 @@ impl SimpleToPreparedPlan { 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) { if let Some(ref step_two) = self.step_two { request.simple_to_prepared_rewrite = true; 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 356832179..38536050e 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( From 5d1205cc4ca167705c7bd1c5b547db65df6fa85b Mon Sep 17 00:00:00 2001 From: Lev Kokotov Date: Wed, 12 Aug 2026 13:09:22 -0700 Subject: [PATCH 5/9] save --- .schema/pgdog.schema.json | 6 ++++++ 1 file changed, 6 insertions(+) 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", From 427d9f705547248e213991ca38ae32ea761e1840 Mon Sep 17 00:00:00 2001 From: Lev Kokotov Date: Wed, 12 Aug 2026 14:28:52 -0700 Subject: [PATCH 6/9] hmm --- cli.sh | 2 +- integration/simple_to_prepared/pgdog.toml | 11 +++++++++++ integration/simple_to_prepared/users.toml | 4 ++++ pgdog/src/backend/prepared_statements.rs | 4 ---- pgdog/src/backend/server.rs | 6 ------ pgdog/src/frontend/client/query_engine/query.rs | 15 ++++++++++----- pgdog/src/frontend/prepared_statements/mod.rs | 13 +++++++++++-- .../rewrite/statement/simple_to_prepared.rs | 15 ++------------- 8 files changed, 39 insertions(+), 31 deletions(-) create mode 100644 integration/simple_to_prepared/pgdog.toml create mode 100644 integration/simple_to_prepared/users.toml 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/src/backend/prepared_statements.rs b/pgdog/src/backend/prepared_statements.rs index 7b35fd36c..dde70b9d4 100644 --- a/pgdog/src/backend/prepared_statements.rs +++ b/pgdog/src/backend/prepared_statements.rs @@ -437,10 +437,6 @@ impl PreparedStatements { &self.state } - pub(crate) fn ignore_message(&mut self, code: impl Into) { - self.state.add_ignore(code); - } - /// Get mutable reference to protocol state. pub fn state_mut(&mut self) -> &mut ProtocolState { &mut self.state diff --git a/pgdog/src/backend/server.rs b/pgdog/src/backend/server.rs index 203aae073..cf3132398 100644 --- a/pgdog/src/backend/server.rs +++ b/pgdog/src/backend/server.rs @@ -476,12 +476,6 @@ impl Server { } self.flush().await?; - if client_request.simple_to_prepared_rewrite { - self.prepared_statements.ignore_message('1'); - self.prepared_statements.ignore_message('2'); - self.prepared_statements.ignore_message('t'); - } - // The whole request is now in server's hands. // We can recover the connection from this point on. self.sending_request = false; diff --git a/pgdog/src/frontend/client/query_engine/query.rs b/pgdog/src/frontend/client/query_engine/query.rs index cb250d22e..7560648fa 100644 --- a/pgdog/src/frontend/client/query_engine/query.rs +++ b/pgdog/src/frontend/client/query_engine/query.rs @@ -242,12 +242,17 @@ impl QueryEngine { // Do this before flushing, because flushing can take time. self.cleanup_backend(context)?; - trace!("{:#?} >>> {:?}", message, context.stream.peer_addr()); + let drop_simple_to_prepared = context.client_request.simple_to_prepared_rewrite + && matches!(message.code(), '1' | '2' | 't'); - if flush { - context.stream.send_flush(&message).await?; - } else { - context.stream.send(&message).await?; + if !drop_simple_to_prepared { + 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/prepared_statements/mod.rs b/pgdog/src/frontend/prepared_statements/mod.rs index 784b65bb9..ceb2b15cf 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, } } @@ -74,8 +77,14 @@ impl PreparedStatements { /// 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_local_mapping(&mut self, local: &str, global: &str) { - self.local.insert(local.to_owned(), global.to_owned()); + 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()); + name + } } /// Register prepared statement with the global cache. 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 index 1bb57e1a3..eda83ec73 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/simple_to_prepared.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/simple_to_prepared.rs @@ -43,23 +43,12 @@ impl SimpleToPreparedPlan { } let mut parse = Parse::new_anonymous(stmt); - - // This is what QueryEngine::rewrite_extended does, - // with a small change, see insert_local_mapping below. - let (_, name) = PreparedStatements::global().write().insert(&parse); + 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); - // This will ensure the global counter for this prepared statement - // is correctly decreased when this client disconnects. - // - // We are using the global name for the local cache because - // the client doesn't know it's using prepared statements and will never - // manually close it. - prepared_statements.insert_local_mapping(parse.name(), parse.name()); - self.step_two = Some(SimpleToPreparedPlanStepTwo { parse, bind }); Ok(()) @@ -72,7 +61,6 @@ impl SimpleToPreparedPlan { /// pub(crate) fn apply(&self, request: &mut ClientRequest) { if let Some(ref step_two) = self.step_two { - request.simple_to_prepared_rewrite = true; request.clear(); request.push(ProtocolMessage::Parse(step_two.parse.clone())); request.push(ProtocolMessage::Describe(Describe::new_statement( @@ -81,6 +69,7 @@ impl SimpleToPreparedPlan { request.push(ProtocolMessage::Bind(step_two.bind.clone())); request.push(ProtocolMessage::Execute(Execute::new())); request.push(ProtocolMessage::Sync(Sync)); + request.simple_to_prepared_rewrite = true; } } } From 2251702a02cfcafa75587ff44df8ea37b811ec83 Mon Sep 17 00:00:00 2001 From: Lev Kokotov Date: Wed, 12 Aug 2026 14:47:01 -0700 Subject: [PATCH 7/9] remove rewrite context from request --- .../src/frontend/client/query_engine/query.rs | 9 +- .../frontend/client/query_engine/request.rs | 138 ------------------ .../query_engine/test/rewrite_offset.rs | 2 +- pgdog/src/frontend/client_request.rs | 6 - pgdog/src/frontend/prepared_statements/mod.rs | 2 + .../router/parser/rewrite/statement/plan.rs | 50 ++++++- .../rewrite/statement/simple_to_prepared.rs | 15 +- 7 files changed, 68 insertions(+), 154 deletions(-) delete mode 100644 pgdog/src/frontend/client/query_engine/request.rs diff --git a/pgdog/src/frontend/client/query_engine/query.rs b/pgdog/src/frontend/client/query_engine/query.rs index 7560648fa..399dccd37 100644 --- a/pgdog/src/frontend/client/query_engine/query.rs +++ b/pgdog/src/frontend/client/query_engine/query.rs @@ -242,10 +242,13 @@ impl QueryEngine { // Do this before flushing, because flushing can take time. self.cleanup_backend(context)?; - let drop_simple_to_prepared = context.client_request.simple_to_prepared_rewrite - && matches!(message.code(), '1' | '2' | 't'); + let forward_to_client = context + .rewrite_result + .as_ref() + .map(|rewrite| rewrite.apply_after_execution(&message).forward()) + .unwrap_or(true); - if !drop_simple_to_prepared { + if forward_to_client { trace!("{:#?} >>> {:?}", message, context.stream.peer_addr()); if flush { diff --git a/pgdog/src/frontend/client/query_engine/request.rs b/pgdog/src/frontend/client/query_engine/request.rs deleted file mode 100644 index 6a868437b..000000000 --- a/pgdog/src/frontend/client/query_engine/request.rs +++ /dev/null @@ -1,138 +0,0 @@ -use std::ops::{Deref, DerefMut}; - -use crate::frontend::{ - ClientRequest, - router::{Ast, Route}, -}; - -#[derive(Debug, Clone, thiserror::Error)] -pub(crate) enum Error { - #[error("only initial requests can be parsed")] - InitialToParsed, - - #[error("request was not parsed")] - NoAst, - - #[error("requested was not routed yet")] - NoRoute, -} - -pub(crate) enum EngineRequest<'a> { - Initial(&'a mut ClientRequest), - Parsed(ParsedEngineRequest<'a>), - Routed(RoutedEngineRequest<'a>), - Transitioning, -} - -impl EngineRequest<'_> { - pub(crate) fn into_parsed(&mut self, ast: Option) -> Result<(), Error> { - let request = match std::mem::replace(self, Self::Transitioning) { - Self::Initial(request) => request, - request => { - *self = request; - return Err(Error::InitialToParsed); - } - }; - - *self = Self::Parsed(ParsedEngineRequest { ast, request }); - Ok(()) - } - - pub(crate) fn into_routed(&mut self, route: Route) -> Result<(), Error> { - let request = match std::mem::replace(self, Self::Transitioning) { - Self::Parsed(request) => request, - request => { - *self = request; - return Err(Error::NoAst); - } - }; - - *self = Self::Routed(RoutedEngineRequest { route, request }); - Ok(()) - } - - pub(crate) fn ast(&self) -> Result<&Option, Error> { - match self { - Self::Initial(_) => Err(Error::NoAst), - Self::Parsed(parsed) => Ok(&parsed.ast), - Self::Routed(routed) => Ok(&routed.request.ast), - Self::Transitioning => unreachable!(), - } - } - - pub(crate) fn route(&self) -> Result<&Route, Error> { - match self { - Self::Routed(routed) => Ok(&routed.route), - _ => Err(Error::NoRoute), - } - } - - pub(crate) fn route_mut(&mut self) -> Result<&mut Route, Error> { - match self { - Self::Routed(routed) => Ok(&mut routed.route), - _ => Err(Error::NoRoute), - } - } -} - -impl Deref for EngineRequest<'_> { - type Target = ClientRequest; - - fn deref(&self) -> &Self::Target { - match self { - Self::Initial(req) => req, - Self::Parsed(parsed) => parsed.deref(), - Self::Routed(routed) => routed.deref(), - Self::Transitioning => unreachable!(), - } - } -} - -impl DerefMut for EngineRequest<'_> { - fn deref_mut(&mut self) -> &mut Self::Target { - match self { - Self::Initial(req) => req, - Self::Parsed(parsed) => parsed.deref_mut(), - Self::Routed(routed) => routed.deref_mut(), - Self::Transitioning => unreachable!(), - } - } -} - -pub(crate) struct ParsedEngineRequest<'a> { - pub(crate) ast: Option, - pub(crate) request: &'a mut ClientRequest, -} - -pub(crate) struct RoutedEngineRequest<'a> { - pub(crate) route: Route, - pub(crate) request: ParsedEngineRequest<'a>, -} - -impl Deref for ParsedEngineRequest<'_> { - type Target = ClientRequest; - - fn deref(&self) -> &Self::Target { - self.request - } -} - -impl DerefMut for ParsedEngineRequest<'_> { - fn deref_mut(&mut self) -> &mut Self::Target { - self.request - } -} - -impl Deref for RoutedEngineRequest<'_> { - type Target = ClientRequest; - - fn deref(&self) -> &Self::Target { - self.request.request - } -} - -impl DerefMut for RoutedEngineRequest<'_> { - fn deref_mut(&mut self) -> &mut Self::Target { - self.request.request - } -} diff --git a/pgdog/src/frontend/client/query_engine/test/rewrite_offset.rs b/pgdog/src/frontend/client/query_engine/test/rewrite_offset.rs index 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/client_request.rs b/pgdog/src/frontend/client_request.rs index d11ebd1bc..901b67503 100644 --- a/pgdog/src/frontend/client_request.rs +++ b/pgdog/src/frontend/client_request.rs @@ -33,8 +33,6 @@ pub struct ClientRequest { pub ast: Option, /// Last Parse we received. pub last_parse: Option, - /// Simple to prepared rewrite requires us to drop some messages from the server. - pub(crate) simple_to_prepared_rewrite: bool, } impl MemoryUsage for ClientRequest { @@ -60,7 +58,6 @@ impl ClientRequest { route: None, ast: None, last_parse: None, - simple_to_prepared_rewrite: false, } } @@ -95,7 +92,6 @@ impl ClientRequest { self.messages.clear(); self.route = None; self.ast = None; - self.simple_to_prepared_rewrite = false; } /// We received a complete request and we are ready to @@ -219,7 +215,6 @@ impl ClientRequest { route: self.route.clone(), ast: self.ast.clone(), last_parse: None, - simple_to_prepared_rewrite: self.simple_to_prepared_rewrite, } } @@ -398,7 +393,6 @@ impl From> for ClientRequest { route: None, ast: None, last_parse: None, - simple_to_prepared_rewrite: false, } } } diff --git a/pgdog/src/frontend/prepared_statements/mod.rs b/pgdog/src/frontend/prepared_statements/mod.rs index ceb2b15cf..f6b7c66b1 100644 --- a/pgdog/src/frontend/prepared_statements/mod.rs +++ b/pgdog/src/frontend/prepared_statements/mod.rs @@ -83,6 +83,8 @@ impl PreparedStatements { } 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 } } diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs b/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs index 9ad928de5..e0daff440 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs @@ -1,6 +1,6 @@ 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; @@ -55,20 +55,63 @@ pub struct RewritePlan { #[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 { @@ -123,7 +166,7 @@ impl RewritePlan { /// Apply the rewrite plan to a ClientRequest. pub(crate) fn apply(&self, request: &mut ClientRequest) -> Result { // This needs to run first! - self.simple_to_prepared.apply(request); + let simple_to_prepared = self.simple_to_prepared.apply(request); // Prepend any required Prepare messages for EXECUTE statements. if !self.prepares.is_empty() { @@ -165,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_to_prepared.rs b/pgdog/src/frontend/router/parser/rewrite/statement/simple_to_prepared.rs index eda83ec73..e719a4f37 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/simple_to_prepared.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/simple_to_prepared.rs @@ -59,7 +59,7 @@ impl SimpleToPreparedPlan { /// 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) { + 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())); @@ -69,7 +69,9 @@ impl SimpleToPreparedPlan { request.push(ProtocolMessage::Bind(step_two.bind.clone())); request.push(ProtocolMessage::Execute(Execute::new())); request.push(ProtocolMessage::Sync(Sync)); - request.simple_to_prepared_rewrite = true; + true + } else { + false } } } @@ -174,7 +176,14 @@ impl<'mem> LiteralRewriter<'mem> { "numeric" }, )), - Some(ConstValue::String(value)) => Some((Parameter::new(value.as_bytes()), "text")), + 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" }), From 4a286e68694dffdfeec844eb6718429ca0198297 Mon Sep 17 00:00:00 2001 From: Lev Kokotov Date: Wed, 12 Aug 2026 14:51:41 -0700 Subject: [PATCH 8/9] remove diff --- .../parser/rewrite/statement/auto_id.rs | 26 ++++++------- .../router/parser/rewrite/statement/mod.rs | 8 ++-- .../router/parser/rewrite/statement/plan.rs | 36 +++++++++--------- .../rewrite/statement/simple_prepared.rs | 4 +- .../parser/rewrite/statement/unique_id.rs | 38 +++++++++---------- 5 files changed, 56 insertions(+), 56 deletions(-) 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 65e2b4483..25993205e 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/auto_id.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/auto_id.rs @@ -69,7 +69,7 @@ impl StatementRewrite<'_> { let replaced = self.replace_set_to_default_at_positions(&mut node, mem, &present_pk_positions); if replaced > 0 { - plan.num_auto_id_injected += replaced as u16; + plan.auto_id_injected += replaced as u16; self.rewritten = true; } } @@ -85,7 +85,7 @@ impl StatementRewrite<'_> { if rewrite { for column in missing_columns { self.inject_column_with_unique_id(&mut node, mem, column); - plan.num_auto_id_injected += 1; + plan.auto_id_injected += 1; } self.rewritten = true; } @@ -311,8 +311,8 @@ mod tests { ) .unwrap(); - assert_eq!(plan.num_auto_id_injected, 1); - assert_eq!(plan.num_unique_ids, 1); // confirms unique_id was processed + assert_eq!(plan.auto_id_injected, 1); + assert_eq!(plan.unique_ids, 1); // confirms unique_id was processed assert!(sql.contains("id")); // pgdog.unique_id() should be replaced with actual bigint value assert!(!sql.contains("pgdog.unique_id")); @@ -347,7 +347,7 @@ mod tests { ) .unwrap(); - assert_eq!(plan.num_auto_id_injected, 0); + assert_eq!(plan.auto_id_injected, 0); assert!(!sql.contains("id,")); } @@ -361,7 +361,7 @@ mod tests { ) .unwrap(); - assert_eq!(plan.num_auto_id_injected, 0); + assert_eq!(plan.auto_id_injected, 0); assert!(!sql.contains("pgdog.unique_id")); } @@ -375,7 +375,7 @@ mod tests { ) .unwrap(); - assert_eq!(plan.num_auto_id_injected, 0); + assert_eq!(plan.auto_id_injected, 0); assert!(!sql.contains("pgdog.unique_id")); } @@ -389,7 +389,7 @@ mod tests { ) .unwrap(); - assert_eq!(plan.num_auto_id_injected, 0); + assert_eq!(plan.auto_id_injected, 0); } #[test] @@ -402,7 +402,7 @@ mod tests { ) .unwrap(); - assert_eq!(plan.num_auto_id_injected, 1); + assert_eq!(plan.auto_id_injected, 1); assert!(sql.contains("id")); } @@ -431,7 +431,7 @@ mod tests { // DEFAULT should be replaced with unique_id assert!(!sql.to_uppercase().contains("DEFAULT")); assert!(sql.contains("::bigint")); // value is cast to bigint - assert_eq!(plan.num_unique_ids, 1); + assert_eq!(plan.unique_ids, 1); } #[test] @@ -446,7 +446,7 @@ mod tests { // Both DEFAULT values should be replaced assert!(!sql.to_uppercase().contains("DEFAULT")); - assert_eq!(plan.num_unique_ids, 2); + assert_eq!(plan.unique_ids, 2); } #[test] @@ -524,7 +524,7 @@ mod tests { .unwrap(); // users is sharded, so RewriteOmni should NOT inject auto id - assert_eq!(plan.num_auto_id_injected, 0); + assert_eq!(plan.auto_id_injected, 0); assert!(!sql.contains("::bigint")); } @@ -548,7 +548,7 @@ mod tests { .unwrap(); // users is NOT sharded, so RewriteOmni should inject auto id - assert_eq!(plan.num_auto_id_injected, 1); + assert_eq!(plan.auto_id_injected, 1); assert!(sql.contains("::bigint")); } } diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs b/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs index 262d57e11..6a9a5c211 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs @@ -121,7 +121,7 @@ impl<'a> StatementRewrite<'a> { | Node::UpdateStmt(_) | Node::DeleteStmt(_) => walk::walk(stmt.stmt(), |node| { if let Node::ParamRef(param) = node { - plan.num_params = plan.num_params.max(param.number as u16) + plan.params = plan.params.max(param.number as u16) } }), Node::PrepareStmt(_) | Node::ExecuteStmt(_) | Node::ExplainStmt(_) => {} @@ -144,14 +144,14 @@ impl<'a> StatementRewrite<'a> { } // Track the next parameter number to use - let mut next_param = plan.num_params as i32 + 1; + let mut next_param = plan.params as i32 + 1; let mut err = None; transform::transform_node( stmt.stmt_mut(), &mut transform::TransformClosure::new(|node| { match Self::rewrite_unique_id(node.as_ref(), mem, self.extended, &mut next_param) { Ok(Some(replacement)) => { - plan.num_unique_ids += 1; + plan.unique_ids += 1; self.rewritten = true; node.replace(replacement); None @@ -185,7 +185,7 @@ impl<'a> StatementRewrite<'a> { plan.simple_to_prepared .step_two(self.prepared_statements, &stmt)?; - plan.rewritten_stmt = Some(stmt); + plan.stmt = Some(stmt); } if let Node::InsertStmt(insert) = stmt.stmt() { diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs b/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs index e0daff440..20358f1ca 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs @@ -19,16 +19,16 @@ pub struct RewritePlan { /// the original statement. This is calculated first, /// and $params+n parameters are added to the statement to /// substitute values we are rewriting. - pub(crate) num_params: u16, + pub(crate) params: u16, /// Number of unique IDs to append to the Bind message. - pub(crate) num_unique_ids: u16, + pub(crate) unique_ids: u16, /// Number of auto-injected primary key columns with pgdog.unique_id(). - pub(crate) num_auto_id_injected: u16, + pub(crate) auto_id_injected: u16, /// Rewritten SQL statement. - pub(crate) rewritten_stmt: Option, + pub(crate) stmt: Option, /// Prepared statements to prepend to the client request. /// Each tuple contains (name, statement) for ProtocolMessage::Prepare. @@ -119,9 +119,9 @@ impl RewritePlan { /// `params` is purely informational (count of original `$N` placeholders) /// and doesn't count as a rewrite. pub(crate) fn is_empty(&self) -> bool { - self.num_unique_ids == 0 - && self.num_auto_id_injected == 0 - && self.rewritten_stmt.is_none() + self.unique_ids == 0 + && self.auto_id_injected == 0 + && self.stmt.is_none() && self.prepares.is_empty() && self.insert_split.is_empty() && self.aggregates.is_noop() @@ -133,7 +133,7 @@ impl RewritePlan { pub(crate) fn apply_bind(&self, bind: &mut Bind) -> Result<(), Error> { let format = bind.default_param_format(); - for _ in 0..self.num_unique_ids { + for _ in 0..self.unique_ids { let generator = UniqueId::generator()?; let id = generator.next_id(); let param = match format { @@ -148,7 +148,7 @@ impl RewritePlan { /// Apply the rewrite plan to a Parse message by updating the SQL. pub(crate) fn apply_parse(&self, parse: &mut Parse) { - if let Some(ref stmt) = self.rewritten_stmt { + if let Some(ref stmt) = self.stmt { parse.set_query(stmt); if !parse.anonymous() { PreparedStatements::global().write().rewrite(parse); @@ -158,7 +158,7 @@ impl RewritePlan { /// Apply the rewrite plan to a Query message by updating the SQL. pub(crate) fn apply_query(&self, query: &mut Query) { - if let Some(ref stmt) = self.rewritten_stmt { + if let Some(ref stmt) = self.stmt { query.set_query(stmt); } } @@ -232,7 +232,7 @@ mod tests { fn test_apply_bind_text_format() { let _guard = set_env_var("NODE_ID", "pgdog-1"); let plan = RewritePlan { - num_unique_ids: 1, + unique_ids: 1, ..Default::default() }; let mut bind = Bind::default(); @@ -252,8 +252,8 @@ mod tests { fn test_apply_bind_binary_format_uniform() { let _guard = set_env_var("NODE_ID", "pgdog-1"); let plan = RewritePlan { - num_params: 1, - num_unique_ids: 1, + params: 1, + unique_ids: 1, ..Default::default() }; // Create bind with uniform binary format (1 code applies to all) @@ -277,8 +277,8 @@ mod tests { fn test_apply_bind_binary_format_one_to_one() { let _guard = set_env_var("NODE_ID", "pgdog-1"); let plan = RewritePlan { - num_params: 2, - num_unique_ids: 1, + params: 2, + unique_ids: 1, ..Default::default() }; // Create bind with one-to-one format codes @@ -304,7 +304,7 @@ mod tests { fn test_apply_bind_multiple_unique_ids() { let _guard = set_env_var("NODE_ID", "pgdog-1"); let plan = RewritePlan { - num_unique_ids: 3, + unique_ids: 3, ..Default::default() }; let mut bind = Bind::default(); @@ -324,8 +324,8 @@ mod tests { fn test_apply_bind_appends_to_existing_params() { let _guard = set_env_var("NODE_ID", "pgdog-1"); let plan = RewritePlan { - num_params: 2, - num_unique_ids: 2, + params: 2, + unique_ids: 2, ..Default::default() }; let mut bind = Bind::new_params( 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 c10b57a11..7cd630e6c 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/simple_prepared.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/simple_prepared.rs @@ -175,7 +175,7 @@ mod tests { "original name should be replaced: {sql}" ); assert!(plan.prepares.is_empty()); - assert!(plan.rewritten_stmt.is_some()); + assert!(plan.stmt.is_some()); } #[test] @@ -226,6 +226,6 @@ mod tests { assert_eq!(sql, "SELECT 1, 2, 3"); assert!(plan.prepares.is_empty()); - assert!(plan.rewritten_stmt.is_none()); + assert!(plan.stmt.is_none()); } } 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 38536050e..c56a9b678 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/unique_id.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/unique_id.rs @@ -126,8 +126,8 @@ mod tests { let (sql, plan) = run_test("SELECT pgdog.unique_id()", true); assert_eq!(sql, "SELECT $1::bigint"); - assert_eq!(plan.num_params, 0); - assert_eq!(plan.num_unique_ids, 1); + assert_eq!(plan.params, 0); + assert_eq!(plan.unique_ids, 1); } #[test] @@ -135,8 +135,8 @@ mod tests { let (sql, plan) = run_test("SELECT pgdog.unique_id(), $1, $2", true); assert_eq!(sql, "SELECT $3::bigint, $1, $2"); - assert_eq!(plan.num_params, 2); - assert_eq!(plan.num_unique_ids, 1); + assert_eq!(plan.params, 2); + assert_eq!(plan.unique_ids, 1); } #[test] @@ -144,8 +144,8 @@ mod tests { let (sql, plan) = run_test("SELECT pgdog.unique_id(), pgdog.unique_id()", true); assert_eq!(sql, "SELECT $1::bigint, $2::bigint"); - assert_eq!(plan.num_params, 0); - assert_eq!(plan.num_unique_ids, 2); + assert_eq!(plan.params, 0); + assert_eq!(plan.unique_ids, 2); } #[test] @@ -157,8 +157,8 @@ mod tests { !sql.contains("pgdog.unique_id"), "Function should be replaced: {sql}" ); - assert_eq!(plan.num_params, 0); - assert_eq!(plan.num_unique_ids, 1); + assert_eq!(plan.params, 0); + assert_eq!(plan.unique_ids, 1); } #[test] @@ -171,7 +171,7 @@ mod tests { !sql.contains("pgdog.unique_id"), "Functions should be replaced: {sql}" ); - assert_eq!(plan.num_unique_ids, 2); + assert_eq!(plan.unique_ids, 2); } #[test] @@ -179,7 +179,7 @@ mod tests { let (sql, plan) = run_test("SELECT 1, 2, 3", true); assert_eq!(sql, "SELECT 1, 2, 3"); - assert_eq!(plan.num_unique_ids, 0); + assert_eq!(plan.unique_ids, 0); } #[test] @@ -190,7 +190,7 @@ mod tests { ); assert_eq!(sql, "INSERT INTO t (id, name) VALUES ($1::bigint, 'test')"); - assert_eq!(plan.num_unique_ids, 1); + assert_eq!(plan.unique_ids, 1); } #[test] @@ -201,7 +201,7 @@ mod tests { ); assert_eq!(sql, "INSERT INTO t (id) VALUES ($1::bigint), ($2::bigint)"); - assert_eq!(plan.num_unique_ids, 2); + assert_eq!(plan.unique_ids, 2); } #[test] @@ -209,7 +209,7 @@ mod tests { let (sql, plan) = run_test("INSERT INTO t (id) SELECT pgdog.unique_id() FROM s", true); assert_eq!(sql, "INSERT INTO t (id) SELECT $1::bigint FROM s"); - assert_eq!(plan.num_unique_ids, 1); + assert_eq!(plan.unique_ids, 1); } #[test] @@ -220,7 +220,7 @@ mod tests { ); assert_eq!(sql, "UPDATE t SET id = $1::bigint WHERE name = 'test'"); - assert_eq!(plan.num_unique_ids, 1); + assert_eq!(plan.unique_ids, 1); } #[test] @@ -231,7 +231,7 @@ mod tests { ); assert_eq!(sql, "UPDATE t SET name = 'new' WHERE id = $1::bigint"); - assert_eq!(plan.num_unique_ids, 1); + assert_eq!(plan.unique_ids, 1); } #[test] @@ -239,7 +239,7 @@ mod tests { let (sql, plan) = run_test("DELETE FROM t WHERE id = pgdog.unique_id()", true); assert_eq!(sql, "DELETE FROM t WHERE id = $1::bigint"); - assert_eq!(plan.num_unique_ids, 1); + assert_eq!(plan.unique_ids, 1); } #[test] @@ -253,7 +253,7 @@ mod tests { sql, "INSERT INTO t (id) VALUES ($1::bigint) RETURNING $2::bigint" ); - assert_eq!(plan.num_unique_ids, 2); + assert_eq!(plan.unique_ids, 2); } #[test] @@ -264,7 +264,7 @@ mod tests { ); assert_eq!(sql, "EXPLAIN INSERT INTO t (id) SELECT $1::bigint FROM s"); - assert_eq!(plan.num_unique_ids, 1); + assert_eq!(plan.unique_ids, 1); } #[test] @@ -272,7 +272,7 @@ mod tests { let (sql, plan) = run_test("EXPLAIN SELECT pgdog.unique_id()", true); assert_eq!(sql, "EXPLAIN SELECT $1::bigint"); - assert_eq!(plan.num_unique_ids, 1); + assert_eq!(plan.unique_ids, 1); } fn run_test(sql: &str, extended: bool) -> (String, RewritePlan) { From fb1ec75c3e3a01e4b12b7ebf858e3bc0ee11e8a1 Mon Sep 17 00:00:00 2001 From: Lev Kokotov Date: Wed, 12 Aug 2026 15:15:02 -0700 Subject: [PATCH 9/9] only rewrite selects --- .../rewrite/statement/simple_to_prepared.rs | 22 +++++++++++++------ 1 file changed, 15 insertions(+), 7 deletions(-) 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 index e719a4f37..c3a0a32ff 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/simple_to_prepared.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/simple_to_prepared.rs @@ -128,13 +128,7 @@ impl StatementRewrite<'_> { /// 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(_) - | NodeMut::InsertStmt(_) - | NodeMut::UpdateStmt(_) - | NodeMut::DeleteStmt(_) - ) { + if !matches!(&node, NodeMut::SelectStmt(_)) { return SimpleToPreparedPlan::default(); } @@ -353,4 +347,18 @@ mod tests { 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()); + } + } }