From f5d74c0a9bdfecc02abaa5df01632968c23cbad1 Mon Sep 17 00:00:00 2001 From: jkaczman Date: Thu, 1 Oct 2026 11:38:51 -0400 Subject: [PATCH 1/2] feat: SELECT non-deterministic transaction-time rewrites --- Cargo.lock | 1 + integration/rust/Cargo.toml | 1 + .../integration/non_deterministic_funcs.rs | 105 ++++++- .../router/parser/rewrite/statement/mod.rs | 8 +- .../statement/non_deterministic_funcs/mod.rs | 288 +++++++----------- .../non_deterministic_funcs/rewrite_insert.rs | 161 ++++++++++ .../rewrite_transaction_time.rs | 108 +++++++ .../{time.rs => time_function.rs} | 16 +- .../{uuid.rs => uuid_function.rs} | 0 9 files changed, 510 insertions(+), 178 deletions(-) create mode 100644 pgdog/src/frontend/router/parser/rewrite/statement/non_deterministic_funcs/rewrite_insert.rs create mode 100644 pgdog/src/frontend/router/parser/rewrite/statement/non_deterministic_funcs/rewrite_transaction_time.rs rename pgdog/src/frontend/router/parser/rewrite/statement/non_deterministic_funcs/{time.rs => time_function.rs} (95%) rename pgdog/src/frontend/router/parser/rewrite/statement/non_deterministic_funcs/{uuid.rs => uuid_function.rs} (100%) diff --git a/Cargo.lock b/Cargo.lock index 317d07819..3b3d1b55c 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2449,6 +2449,7 @@ dependencies = [ "chrono", "chrono-tz", "futures-util", + "itertools 0.15.0", "libc", "native-tls", "ordered-float", diff --git a/integration/rust/Cargo.toml b/integration/rust/Cargo.toml index 126b2e3b2..dbeef7a89 100644 --- a/integration/rust/Cargo.toml +++ b/integration/rust/Cargo.toml @@ -27,3 +27,4 @@ bytes.workspace = true rust_decimal = { version = "1.42.0", features = ["macros"] } chrono-tz = "0.10.4" pgdog-stats = { path = "../../pgdog-stats" } +itertools = "0.15" diff --git a/integration/rust/tests/integration/non_deterministic_funcs.rs b/integration/rust/tests/integration/non_deterministic_funcs.rs index 26f428130..67e15d233 100644 --- a/integration/rust/tests/integration/non_deterministic_funcs.rs +++ b/integration/rust/tests/integration/non_deterministic_funcs.rs @@ -31,7 +31,7 @@ use sqlx::{Executor, Row}; /// This tests a case where we're performing two INSERTs which resolve to different Shards. This uses 2 connections. /// Previously, the now() values would be different in each. #[tokio::test] -async fn two_conns_transaction_time_reuse() { +async fn two_conns_transaction_time_reuse_insert() { let single_sharded_list_pool = PgPoolOptions::new() .max_connections(1) .connect("postgres://pgdog:pgdog@127.0.0.1:6432/single_sharded_list?application_name=sqlx") @@ -65,6 +65,109 @@ async fn two_conns_transaction_time_reuse() { transaction.rollback().await.unwrap(); } +/// Ensure that SELECT now() and CURRENT_TIMESTAMP return the correct type, are consistent +/// within a transaction across multiple queries, and return within the correct column name +/// despite being re-written. +#[tokio::test] +async fn transaction_time_select_equality() { + run_and_assert_eq::>("SELECT now();", "now").await; + run_and_assert_eq::>("SELECT CURRENT_TIMESTAMP;", "current_timestamp").await; +} + +/// Helper method to run a query twice, in one transaction, in both simple and extended protocol, +/// and assert that the response is equal (and of the correct type) +async fn run_and_assert_eq(query: &str, col: &str) +where + T: PartialEq + + std::fmt::Debug + + sqlx::Type + + for<'a> sqlx::Decode<'a, sqlx::Postgres>, +{ + let conn = connections_sqlx().await; + let conn = conn.get(1).unwrap(); + let mut transaction = conn.begin().await.unwrap(); + + for simple_protocol in [false, true] { + let mut rows = vec![]; + for _ in 0..2 { + rows.push(if simple_protocol { + sqlx::raw_sql(query) + .fetch_one(&mut *transaction) + .await + .unwrap() + } else { + sqlx::query(query) + .fetch_one(&mut *transaction) + .await + .unwrap() + }); + } + + assert_eq!(rows[0].get::(col), rows[1].get::(col)); + } +} + +/// Ensure that now() functionality isn't broken for use within a query +/// Tests things like comparisons and addition with an interval in a SELECT query +#[tokio::test] +async fn transaction_time_usage_in_select_query() { + let conn = connections_sqlx().await; + let conn = conn.get(1).unwrap(); + + let mut transaction = conn.begin().await.unwrap(); + + // 'created_at' col has default now() + sqlx::query("INSERT INTO sharded (id) VALUES (1)") + .execute(&mut *transaction) + .await + .unwrap(); + + for use_simple_protocol in [false, true] { + // Test comparisons + let query_tries = [ + ("SELECT * FROM sharded WHERE created_at = now()", 1), + ("SELECT * FROM sharded WHERE created_at < now()", 0), + ("SELECT * FROM sharded WHERE created_at > now()", 0), + ]; + for (query, expected_rows) in query_tries { + let rows_fetched = if use_simple_protocol { + sqlx::raw_sql(query) + .fetch_all(&mut *transaction) + .await + .unwrap() + } else { + sqlx::query(query) + .fetch_all(&mut *transaction) + .await + .unwrap() + }; + assert_eq!(rows_fetched.len(), expected_rows); + } + + // Test addition + let query = "SELECT now() + INTERVAL '1 day' AS this_time_tomorrow;"; + + let time_plus_1_day: DateTime = if use_simple_protocol { + let row: PgRow = sqlx::raw_sql(query) + .fetch_one(&mut *transaction) + .await + .unwrap(); + row.get::<_, &str>("this_time_tomorrow") + } else { + sqlx::query_scalar(query) + .fetch_one(&mut *transaction) + .await + .unwrap() + }; + + let now: DateTime = Utc::now(); + assert!( + time_plus_1_day > now + Duration::hours(23) + && time_plus_1_day < now + Duration::hours(25) + ); + } +} + /// Test that an INSERT into an omnisharded table which uses UUID functions is /// re-written to a constant to be consistent across all shards. /// diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs b/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs index 5e0ba0579..311bba5b8 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs @@ -212,11 +212,15 @@ impl<'a> StatementRewrite<'a> { if nd_function_rewrite { match stmt.stmt_mut() { - NodeMut::InsertStmt(_) => { + // TODO: we could also support UPDATE / etc + NodeMut::InsertStmt(_) | NodeMut::SelectStmt(_) => { self.rewrite_nd_functions(stmt.stmt_mut(), mem, &mut next_param, &mut plan)?; } NodeMut::PrepareStmt(mut prepare) => { - if matches!(prepare.query_mut(), NodeMut::InsertStmt(_)) { + if matches!( + prepare.query_mut(), + NodeMut::InsertStmt(_) | NodeMut::SelectStmt(_) + ) { self.rewrite_nd_functions( prepare.query_mut(), mem, diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/non_deterministic_funcs/mod.rs b/pgdog/src/frontend/router/parser/rewrite/statement/non_deterministic_funcs/mod.rs index ebd6c5f9d..ae4604cc5 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/non_deterministic_funcs/mod.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/non_deterministic_funcs/mod.rs @@ -1,11 +1,8 @@ -use std::ops::Deref; - use pg_raw_parse::{ ConstValue, Node, NodeMut, list::NodeList, make::{MemoryToken, Unique}, raw::SQLValueFunctionOp, - transform::{TransformClosure, transform_node}, }; use pgdog_stats::{Column, Relation}; @@ -26,6 +23,8 @@ use crate::{ }; /// No need to expose these outside. +mod rewrite_insert; +mod rewrite_transaction_time; mod time; mod uuid; @@ -88,12 +87,19 @@ enum NDFunctionType { UUIDFunction(UUIDFunctionType), } +/// We cover two cases in this code: (1) transaction time functions (SELECT, all INSERT) and (2) omnisharded INSERT +#[derive(PartialEq)] +enum RewriteCase { + TransactionTimeFunction, + OmnishardedInsert, +} + impl NDFunctionType { /// Convert `SQLValueFunctionOp` (e.g. current_date, current_time... non ()) to `NDFunctionType` fn from_sql_value_function( op: SQLValueFunctionOp::Type, typmod: i32, - is_sharded: bool, + rewrite_case: &RewriteCase, ) -> Option { let nd_func_type = TimeFunctionType::from_sql_value_function(op, typmod) .map(NDFunctionType::TimeFunction) @@ -102,9 +108,11 @@ impl NDFunctionType { .map(NDFunctionType::UUIDFunction) }); + // Verify the `SQLValueFunction` is transaction time based, if we're + // only re-writing transaction time functions. if let Some(nd_func) = nd_func_type - && is_sharded - && !nd_func.apply_rewrite_on_sharded_tables() + && *rewrite_case == RewriteCase::TransactionTimeFunction + && !nd_func.is_transaction_time_function() { return None; } @@ -136,13 +144,22 @@ impl NDFunctionType { /// Why not re-write everything? This implementation isn't perfect; some things aren't implemented (e.g. interval arg for uuidv7). /// Best not to do anything we don't have to; otherwise, it could unnecessarily break something for someone. /// (maybe they're reliant on the config setting for something else) - fn apply_rewrite_on_sharded_tables(&self) -> bool { + fn is_transaction_time_function(&self) -> bool { match self { - Self::TimeFunction(tf) => tf.apply_rewrite_on_sharded_tables(), + Self::TimeFunction(tf) => tf.is_transaction_time_function(), Self::UUIDFunction(_) => false, } } + /// Used to determine what type the Postgres function outputs by default, so we can cast the text into it + /// in cases like SELECT where we don't have a corresponding column type to go by. + fn output_type_as_str(&self) -> &'static str { + match self { + Self::TimeFunction(tf) => tf.default_output_type().into_postgres_str(), + Self::UUIDFunction(_) => "uuid", + } + } + /// Convert `TimeFunction` in VALUES list /// /// Parse both `FuncCall`s and `SQLValueFunction`s here. @@ -151,7 +168,7 @@ impl NDFunctionType { fn from_node( node: Node, column_relation: Option<&Column>, - is_sharded: bool, + rewrite_case: &RewriteCase, ) -> Result, Error> { match node { Node::FuncCall(func) => { @@ -163,12 +180,12 @@ impl NDFunctionType { return Ok(None); }; - Self::from_func_call(func_name, Some(func.args()), is_sharded) + Self::from_func_call(func_name, Some(func.args()), rewrite_case) } Node::SQLValueFunction(func) => Ok(Self::from_sql_value_function( func.op, func.typmod, - is_sharded, + rewrite_case, )), // If DEFAULT is in a VALUES list; fetch the column based on index. @@ -177,7 +194,7 @@ impl NDFunctionType { return Self::from_func_call( relation.column_default.as_str(), None, - is_sharded, + rewrite_case, ); } @@ -198,7 +215,7 @@ impl NDFunctionType { fn from_func_call( func_name: &str, args: Option<&NodeList>, - is_sharded: bool, + case: &RewriteCase, ) -> Result, Error> { // Normalize the function name. Postgres does this. let func_name = func_name.to_lowercase(); @@ -207,7 +224,12 @@ impl NDFunctionType { .iter() .chain(UUIDFunctionType::ALL_VARIANTS.iter()) { - if is_sharded && !variant.apply_rewrite_on_sharded_tables() { + // If we're re-writing for a transaction time function, and + // the current variant we're looking at isn't one, then + // we should skip. + if *case == RewriteCase::TransactionTimeFunction + && !variant.is_transaction_time_function() + { continue; } @@ -287,33 +309,62 @@ impl StatementRewrite<'_> { next_param: &mut i32, plan: &mut RewritePlan, ) -> Result<(), Error> { - let mut parser = StatementParser::new(stmt.as_ref(), None, self.schema, None); + if matches!(stmt.as_ref(), Node::InsertStmt(_)) { + let mut parser = StatementParser::new(stmt.as_ref(), None, self.schema, None); + + let mut nd_rewrite = NDRewrite { + rewrite: self, + plan, + next_param, + mem, + statement_type: StatementType::Insert, + }; - // We allow `ShardedTable`s on a case-by-case basis (see `apply_rewrite_on_sharded_tables`) - let is_sharded = parser.is_sharded(self.db_schema, self.user, self.search_path); + // We allow `ShardedTable`s on a case-by-case basis (see `apply_rewrite_on_sharded_tables`) + let is_sharded = parser.is_sharded( + nd_rewrite.rewrite.db_schema, + nd_rewrite.rewrite.user, + nd_rewrite.rewrite.search_path, + ); - let Some((relation, cols, not_covered_cols)) = self.find_not_used_cols(&mut stmt, mem) - else { - return Ok(()); - }; + let Some((relation, cols, not_covered_cols)) = + nd_rewrite.rewrite.find_not_used_cols(&mut stmt, mem) + else { + return Ok(()); + }; - let mut nd_rewrite = NDRewrite { - rewrite: self, - plan, - next_param, - mem, - relation, - cols, - is_sharded, - }; + let insert_context = InsertContext { + relation, + cols, + rewrite_case: if is_sharded { + RewriteCase::TransactionTimeFunction + } else { + RewriteCase::OmnishardedInsert + }, + }; + + // 1. iterates through Schema to find DEFAULT columns + // 2. adds the column to target list & all the values lists (ParamRef or String) + nd_rewrite.handle_adding_defaults(&mut stmt, ¬_covered_cols, &insert_context)?; + + // Replaces all non-deterministic function calls (ParamRef or String) + nd_rewrite.transform_func_calls_in_insert(stmt, &insert_context)?; + } else if let NodeMut::SelectStmt(select_stmt) = stmt { + // Still rewrite other kind of statements that use ND functions that + // are reliant on transaction start time - // 1. iterates through Schema to find DEFAULT columns - // 2. adds the column to target list & all the values lists (ParamRef or String) - nd_rewrite.handle_adding_defaults(&mut stmt, ¬_covered_cols)?; + // TODO: Eventually cover Update, Delete, etc - // Replaces all non-deterministic function calls (ParamRef or String) - nd_rewrite.transform_func_calls(stmt)?; + let mut nd_rewrite = NDRewrite { + rewrite: self, + plan, + next_param, + mem, + statement_type: StatementType::Select, + }; + nd_rewrite.transform_func_calls_transaction_time(select_stmt)?; + } Ok(()) } @@ -364,147 +415,24 @@ struct NDRewrite<'mem, 'a, 's> { /// could do like a .next_param() method on `RewritePlan` next_param: &'a mut i32, mem: MemoryToken<'mem>, + statement_type: StatementType, +} + +enum StatementType { + Insert, + Select, +} + +/// Extra context necessary for re-writing INSERT statements, where we need +/// information to determine things like if it's sharded, or what kind of column each column is +struct InsertContext<'mem> { relation: Relation, cols: Unique<'mem, &'mem NodeList>, - is_sharded: bool, + rewrite_case: RewriteCase, } impl<'mem, 'a, 's> NDRewrite<'mem, 'a, 's> { /// Replaces all non-deterministic function calls (ParamRef or String) - /// Used by all to handle re-writes. - fn transform_func_calls(&mut self, stmt: NodeMut<'mem, '_>) -> Result<(), Error> { - // If any Error is caught during transform_node, update this, and it'll be returned when the transform is done. - // Have this workaround because it's within a closure. - let mut err: Option = None; - transform_node( - stmt, - &mut TransformClosure::new(|node| match &*node { - // TODO: Is this guaranteed to be a VALUES list? - // What if it's something unrelated in the statement? - NodeMut::NodeList(list_of_values) => { - // VALUES (...), (...) where (...) is what we're inspecting (one NodeList) - - let mut cloned_values = self.mem.make_unique(list_of_values.deref()); - let mut changed = false; - - // The reason this is iterating over the NodeList instead of individual - // FuncCalls is that we must know where we are within a VALUES, as that - // allows us to know the present column's data type (for potential later coersion) - for (i, value) in list_of_values.iter().enumerate() { - let col_relation = if self.cols.is_empty() { - self.relation.columns.get_index(i).map(|(_, column)| column) - } else { - match self.cols.get(i) { - Some(Node::ResTarget(target)) => target - .name() - .and_then(|name| self.relation.columns.get(name)), - _ => None, - } - }; - - match NDFunctionType::from_node(value, col_relation, self.is_sharded) { - Ok(Some(nd_function_type)) => { - let Some(col_relation) = col_relation else { - continue; - }; - - let nd_function = NDFunction { - nd_function_type, - column_type: col_relation.data_type.clone(), - }; - - let node = self.make_node(&nd_function); - match node { - Ok(node) => { - // Replace the specific node within the list. - cloned_values.as_mut().set(i, node); - changed = true; - } - Err(e) => { - err.get_or_insert(e); - break; - } - } - } - - Ok(None) => continue, - Err(e) => { - err.get_or_insert(e); - break; - } - } - } - - // Replaces the entire VALUES list at once with the one we cloned and re-wrote. - if changed { - node.replace(cloned_values.uncast()); - - // Do not continue to traverse. - return None; - } - - Some(node) - } - _ => Some(node), - }), - ); - - err.map(Err).unwrap_or(Ok(())) - } - - /// Iterates through Schema to find DEFAULT columns - /// Adds the column to target list & all the values lists (ParamRef or String) - fn handle_adding_defaults( - &mut self, - mut stmt: &mut NodeMut<'mem, '_>, - not_covered_cols: &Vec, - ) -> Result<(), Error> { - let NodeMut::InsertStmt(insert_stmt) = &mut stmt else { - return Ok(()); - }; - - for col in not_covered_cols { - let col_relation = self.relation.columns.get(col.as_str()).unwrap(); - - let nd_function_type = NDFunctionType::from_func_call( - &col_relation.column_default, - None, - self.is_sharded, - )?; - - let Some(nd_function_type) = nd_function_type else { - continue; - }; - - let nd_function = NDFunction { - nd_function_type, - column_type: col_relation.data_type.clone(), - }; - - // Add to the list of cols in the INSERT. - insert_stmt.cols_mut().push( - self.mem, - self.mem - .make_res_target(Some(col), self.mem.empty(), self.mem.none()) - .uncast(), - ); - - let NodeMut::SelectStmt(select_stmt) = &mut insert_stmt.select_stmt_mut() else { - return Ok(()); - }; - - // Have to add the now() to every single select VALUES list now. - // VALUES (...), (....) - for values_list in select_stmt.values_lists_mut() { - let mut node_list_mut = values_list.expect_node_list(); - - node_list_mut.push(self.mem, self.make_node(&nd_function)?); - } - } - - Ok(()) - } - /// If simple protocol, make an A_Const node with the String constant of the formatted function output. /// If extended or prepare, make a ParamRef, so that we can cache it and put in the formatted output later. fn make_node(&mut self, nd_function: &NDFunction) -> Result>, Error> { @@ -518,9 +446,25 @@ impl<'mem, 'a, 's> NDRewrite<'mem, 'a, 's> { // The statement is discarded when the error is returned (thus, value doesn't matter) Err(err) => return Err(err), }; - self.mem + + let constant_text_node = self + .mem .make_a_const(ConstValue::String(text.as_str())) - .uncast() + .uncast(); + + match self.statement_type { + StatementType::Insert => constant_text_node, + StatementType::Select => self + .mem + .make_type_cast( + constant_text_node, + self.mem.make_list(&[self + .mem + .make_string(Some(nd_function.col_type_to_type_cast_alias()))]), + ) + .uncast(), + } + .uncast() } else { let param_ref = self.mem.make_param_ref(*self.next_param); *self.next_param += 1; diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/non_deterministic_funcs/rewrite_insert.rs b/pgdog/src/frontend/router/parser/rewrite/statement/non_deterministic_funcs/rewrite_insert.rs new file mode 100644 index 000000000..29a6a1346 --- /dev/null +++ b/pgdog/src/frontend/router/parser/rewrite/statement/non_deterministic_funcs/rewrite_insert.rs @@ -0,0 +1,161 @@ +use std::ops::Deref; + +use pg_raw_parse::{ + Node, NodeMut, + transform::{TransformClosure, transform_node}, +}; + +use crate::frontend::router::parser::rewrite::statement::{ + Error, + non_deterministic_funcs::{InsertContext, NDFunction, NDFunctionType, NDRewrite}, +}; + +impl<'mem, 'a, 's> NDRewrite<'mem, 'a, 's> { + /// Replaces all non-deterministic function calls (ParamRef or String) + /// Used by all to handle re-writes. + pub(super) fn transform_func_calls_in_insert( + &mut self, + stmt: NodeMut<'mem, '_>, + insert_context: &InsertContext, + ) -> Result<(), Error> { + // If any Error is caught during transform_node, update this, and it'll be returned when the transform is done. + // Have this workaround because it's within a closure. + let mut err: Option = None; + transform_node( + stmt, + &mut TransformClosure::new(|node| match &*node { + // TODO: Is this guaranteed to be a VALUES list? + // What if it's something unrelated in the statement? + NodeMut::NodeList(list_of_values) => { + // VALUES (...), (...) where (...) is what we're inspecting (one NodeList) + + let mut cloned_values = self.mem.make_unique(list_of_values.deref()); + let mut changed = false; + + // The reason this is iterating over the NodeList instead of individual + // FuncCalls is that we must know where we are within a VALUES, as that + // allows us to know the present column's data type (for potential later coersion) + for (i, value) in list_of_values.iter().enumerate() { + let col_relation = if insert_context.cols.is_empty() { + insert_context + .relation + .columns + .get_index(i) + .map(|(_, column)| column) + } else { + match insert_context.cols.get(i) { + Some(Node::ResTarget(target)) => target + .name() + .and_then(|name| insert_context.relation.columns.get(name)), + _ => None, + } + }; + + match NDFunctionType::from_node( + value, + col_relation, + &insert_context.rewrite_case, + ) { + Ok(Some(nd_function_type)) => { + let Some(col_relation) = col_relation else { + continue; + }; + + let nd_function = NDFunction { + nd_function_type, + column_type: col_relation.data_type.clone(), + }; + + let node = self.make_node(&nd_function); + match node { + Ok(node) => { + // Replace the specific node within the list. + cloned_values.as_mut().set(i, node); + changed = true; + } + Err(e) => { + err.get_or_insert(e); + break; + } + } + } + + Ok(None) => continue, + Err(e) => { + err.get_or_insert(e); + break; + } + } + } + + // Replaces the entire VALUES list at once with the one we cloned and re-wrote. + if changed { + node.replace(cloned_values.uncast()); + + // Do not continue to traverse. + return None; + } + + Some(node) + } + _ => Some(node), + }), + ); + + err.map(Err).unwrap_or(Ok(())) + } + + /// Iterates through Schema to find DEFAULT columns + /// Adds the column to target list & all the values lists (ParamRef or String) + pub(super) fn handle_adding_defaults( + &mut self, + mut stmt: &mut NodeMut<'mem, '_>, + not_covered_cols: &Vec, + insert_context: &InsertContext<'mem>, + ) -> Result<(), Error> { + let NodeMut::InsertStmt(insert_stmt) = &mut stmt else { + return Ok(()); + }; + + for col in not_covered_cols { + let col_relation = insert_context.relation.columns.get(col.as_str()).unwrap(); + + let nd_function_type = NDFunctionType::from_func_call( + &col_relation.column_default, + None, + &insert_context.rewrite_case, + )?; + + let Some(nd_function_type) = nd_function_type else { + continue; + }; + + let nd_function = NDFunction { + nd_function_type, + column_type: col_relation.data_type.clone(), + }; + + // Add to the list of cols in the INSERT. + insert_stmt.cols_mut().push( + self.mem, + self.mem + .make_res_target(Some(col), self.mem.empty(), self.mem.none()) + .uncast(), + ); + + let NodeMut::SelectStmt(select_stmt) = &mut insert_stmt.select_stmt_mut() else { + return Ok(()); + }; + + // Have to add the now() to every single select VALUES list now. + // VALUES (...), (....) + for values_list in select_stmt.values_lists_mut() { + let mut node_list_mut = values_list.expect_node_list(); + + node_list_mut.push(self.mem, self.make_node(&nd_function)?); + } + } + + Ok(()) + } +} diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/non_deterministic_funcs/rewrite_transaction_time.rs b/pgdog/src/frontend/router/parser/rewrite/statement/non_deterministic_funcs/rewrite_transaction_time.rs new file mode 100644 index 000000000..7317af051 --- /dev/null +++ b/pgdog/src/frontend/router/parser/rewrite/statement/non_deterministic_funcs/rewrite_transaction_time.rs @@ -0,0 +1,108 @@ +use std::ops::Deref; + +use pg_raw_parse::{ + Node, + nodes::SelectStmtMut, + transform::{self, Transform}, +}; + +use crate::frontend::router::parser::rewrite::statement::{ + Error, + non_deterministic_funcs::{NDFunction, NDFunctionType, NDRewrite, RewriteCase}, +}; + +impl<'mem, 'a, 's> NDRewrite<'mem, 'a, 's> { + pub(super) fn transform_func_calls_transaction_time<'mutref>( + &mut self, + stmt: SelectStmtMut<'mem, 'mutref>, + ) -> Result<(), Error> { + let mut transform_select = ReplaceTransactionTimeSelect { + nd_rewrite: self, + outer_error: None, + }; + + transform_select.transform_select_stmt(stmt); + transform_select.outer_error.map(Err).unwrap_or(Ok(())) + } +} +/// transform_node doesn't let us work with ResTargets, so we have to implement transform ourselves +/// see +struct ReplaceTransactionTimeSelect<'mutr, 'mem, 'a, 's> { + /// Ability to reference `self` within the Transform impl. + nd_rewrite: &'mutr mut NDRewrite<'mem, 'a, 's>, + /// Replaced with an Error if we come across one, so we can return an Error from this function + /// to the client. + outer_error: Option, +} + +impl<'mutr, 'mem, 'a, 's> Transform<'mem> for ReplaceTransactionTimeSelect<'mutr, 'mem, 'a, 's> { + /// Case, basic: SELECT now() + /// + /// This means now is a ResTarget, and Postgres will output the timestamptz w/ **a now column** + /// If we replace all FunctionCall nodes with Strings, a ?col? will be returned as Postgres doesn't know + /// that the client called now(). + /// + /// This is handled by naming the ResTarget below. + fn transform_res_target<'mutref>( + &mut self, + mut node: pg_raw_parse::nodes::ResTargetMut<'mem, 'mutref>, + ) { + if node.name().is_none() + && matches!(node.val(), Node::FuncCall(_) | Node::SQLValueFunction(_)) + && let Some(nd_function_type) = match NDFunctionType::from_node( + node.val(), + None, + &RewriteCase::TransactionTimeFunction, + ) { + Ok(nd_function_type) => nd_function_type, + Err(e) => { + self.outer_error.get_or_insert(e); + return; + } + } + { + node.set_name(Some( + self.nd_rewrite.mem.copy_string(nd_function_type.name()), + )); + } + + // Continue to traverse. + transform::transform_res_target(node, self); + } + + /// This covers both `FuncCall` and `SQLValueFunction` replacement. + fn transform_node<'mutref>( + &mut self, + node: pg_raw_parse::transform::Assignable<'mem, 'mutref>, + ) { + if let Some(nd_function_type) = match NDFunctionType::from_node( + node.deref().as_ref(), + None, + &RewriteCase::TransactionTimeFunction, + ) { + Ok(nd_function_type) => nd_function_type, + Err(e) => { + self.outer_error.get_or_insert(e); + return; + } + } { + let nd_func = NDFunction { + nd_function_type, + column_type: nd_function_type.output_type_as_str().to_string(), + }; + + node.replace(match self.nd_rewrite.make_node(&nd_func) { + Ok(node) => node, + Err(e) => { + self.outer_error.get_or_insert(e); + return; + } + }); + + return; + } + + // Continue to traverse. + transform::transform_node(node.into_inner(), self); + } +} diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/non_deterministic_funcs/time.rs b/pgdog/src/frontend/router/parser/rewrite/statement/non_deterministic_funcs/time_function.rs similarity index 95% rename from pgdog/src/frontend/router/parser/rewrite/statement/non_deterministic_funcs/time.rs rename to pgdog/src/frontend/router/parser/rewrite/statement/non_deterministic_funcs/time_function.rs index 6800615af..8be8e5a92 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/non_deterministic_funcs/time.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/non_deterministic_funcs/time_function.rs @@ -15,7 +15,7 @@ use pg_raw_parse::raw::SQLValueFunctionOp; /// Represents what Postgres type the `TimeFunction` would normally output. #[derive(PartialEq)] -enum TimeFunctionOutput { +pub(super) enum TimeFunctionOutput { Date, TimeWithTimeZone, TimestampWithTimeZone, @@ -62,6 +62,16 @@ impl TimeFunctionOutput { } } + pub(super) fn into_postgres_str(self) -> &'static str { + match self { + Self::Date => "date", + Self::TimeWithTimeZone => "timetz", + Self::TextFormattedTimestampWithTimeZone | Self::TimestampWithTimeZone => "timestamptz", + Self::Timestamp => "timestamp", + Self::Time => "time", + } + } + /// Postgres trims trailing zeros from fractional seconds /// It also drops the dot when there's none fn fractional_seconds(nanoseconds: u32) -> String { @@ -220,12 +230,12 @@ impl TimeFunctionType { /// If we're considering re-writing a `TimeFunction`, the decision as to whether or not we should /// rewrite rests solely on the corresponding `TimeReference` being `TransactionStart`. Otherwise, /// there's no point; Postgres can achieve the same functionality without our assistance. - pub(super) fn apply_rewrite_on_sharded_tables(&self) -> bool { + pub(super) fn is_transaction_time_function(&self) -> bool { self.time_reference() == TimeReference::TransactionStart } /// Represents what Postgres type the `TimeFunction` would normally output. - fn default_output_type(self) -> TimeFunctionOutput { + pub(super) fn default_output_type(self) -> TimeFunctionOutput { match self { Self::CurrentTimestamp(_) | Self::ClockTimestamp diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/non_deterministic_funcs/uuid.rs b/pgdog/src/frontend/router/parser/rewrite/statement/non_deterministic_funcs/uuid_function.rs similarity index 100% rename from pgdog/src/frontend/router/parser/rewrite/statement/non_deterministic_funcs/uuid.rs rename to pgdog/src/frontend/router/parser/rewrite/statement/non_deterministic_funcs/uuid_function.rs From df2730469c593ff80a78a2c7ff7459d4095666ab Mon Sep 17 00:00:00 2001 From: jkaczman Date: Thu, 1 Oct 2026 11:56:23 -0400 Subject: [PATCH 2/2] Fix merge conflicts --- .../router/parser/rewrite/statement/mod.rs | 9 ++------- .../statement/non_deterministic_funcs/mod.rs | 16 ++++++++-------- 2 files changed, 10 insertions(+), 15 deletions(-) diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs b/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs index 6ba86e63f..912fb7fe8 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs @@ -195,19 +195,14 @@ impl<'a> StatementRewrite<'a> { match stmt.stmt_mut() { // TODO: we could also support UPDATE / etc NodeMut::InsertStmt(_) | NodeMut::SelectStmt(_) => { - self.rewrite_nd_functions(stmt.stmt_mut(), mem, &mut next_param, &mut plan)?; + self.rewrite_nd_functions(stmt.stmt_mut(), mem, &mut plan.bind_params)?; } NodeMut::PrepareStmt(mut prepare) => { if matches!( prepare.query_mut(), NodeMut::InsertStmt(_) | NodeMut::SelectStmt(_) ) { - self.rewrite_nd_functions( - prepare.query_mut(), - mem, - &mut next_param, - &mut plan, - )?; + self.rewrite_nd_functions(prepare.query_mut(), mem, &mut plan.bind_params)?; } } _ => {} diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/non_deterministic_funcs/mod.rs b/pgdog/src/frontend/router/parser/rewrite/statement/non_deterministic_funcs/mod.rs index 3e62a7ee4..d09d0d7ec 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/non_deterministic_funcs/mod.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/non_deterministic_funcs/mod.rs @@ -14,7 +14,9 @@ use crate::{ StatementParser, StatementRewrite, Table, rewrite::statement::{ Error, - non_deterministic_funcs::{time::TimeFunctionType, uuid::UUIDFunctionType}, + non_deterministic_funcs::{ + time_function::TimeFunctionType, uuid_function::UUIDFunctionType, + }, plan::BindParam, }, }, @@ -25,8 +27,8 @@ use crate::{ /// No need to expose these outside. mod rewrite_insert; mod rewrite_transaction_time; -mod time; -mod uuid; +mod time_function; +mod uuid_function; /// A non-deterministic function that we must re-write when writing to an omnisharded table (or now() for sharded), /// so that we can maintain consistency instead of generating a different value (from executing the function) @@ -308,12 +310,11 @@ impl StatementRewrite<'_> { bind_params: &mut BindParams, ) -> Result<(), Error> { if matches!(stmt.as_ref(), Node::InsertStmt(_)) { - let mut parser = StatementParser::new(stmt.as_ref(), None, self.schema, None); + let mut parser = StatementParser::new(stmt.as_ref(), None, self.schema); let mut nd_rewrite = NDRewrite { rewrite: self, - plan, - next_param, + bind_params, mem, statement_type: StatementType::Insert, }; @@ -355,8 +356,7 @@ impl StatementRewrite<'_> { let mut nd_rewrite = NDRewrite { rewrite: self, - plan, - next_param, + bind_params, mem, statement_type: StatementType::Select, };