From 4472af995dc9f43f909c9e398afbd5f3e3db19bc Mon Sep 17 00:00:00 2001 From: jkaczman Date: Fri, 2 Oct 2026 13:35:41 -0400 Subject: [PATCH] feat: UPDATE non-deterministic transaction-time rewrites --- .../integration/non_deterministic_funcs.rs | 60 +++++++++++++++++++ .../router/parser/rewrite/statement/mod.rs | 5 +- .../statement/non_deterministic_funcs/mod.rs | 28 +++++++-- .../rewrite_transaction_time.rs | 23 ++----- 4 files changed, 89 insertions(+), 27 deletions(-) diff --git a/integration/rust/tests/integration/non_deterministic_funcs.rs b/integration/rust/tests/integration/non_deterministic_funcs.rs index 67e15d233..59a74fb17 100644 --- a/integration/rust/tests/integration/non_deterministic_funcs.rs +++ b/integration/rust/tests/integration/non_deterministic_funcs.rs @@ -65,6 +65,66 @@ async fn two_conns_transaction_time_reuse_insert() { transaction.rollback().await.unwrap(); } +/// Verify that usage of now() / CURRENT_TIMESTAMP in an UPDATE statement uses transaction time re-writes. +#[tokio::test] +async fn transaction_time_update_statement_rewrites() { + let conn = connections_sqlx().await; + let conn = conn.get(1).unwrap(); + + let row_created_at_before_time: DateTime = { + conn.execute("TRUNCATE sharded").await.unwrap(); + + let row_before_transaction_start_time: PgRow = conn + .fetch_one("INSERT INTO sharded(id) VALUES (1) RETURNING *") + .await + .unwrap(); + + conn.execute("INSERT INTO sharded(id) VALUES (2) RETURNING *") + .await + .unwrap(); + + row_before_transaction_start_time.get::<_, &str>("created_at") + }; + + let mut transaction = conn.begin().await.unwrap(); + + let transaction_start_time: DateTime = { + let transaction_start_time_row: PgRow = transaction + .fetch_one("INSERT INTO sharded(id) VALUES (3) RETURNING *") + .await + .unwrap(); + + transaction_start_time_row.get::<_, &str>("created_at") + }; + + assert!(row_created_at_before_time != transaction_start_time); + + // Update both to the transaction time NOW(). + // We previously did an INSERT for both before the transaction started, so they'll differ + // (as asserted above) + let (row_tt_extended_protocol_time, row_tt_simple_protocol_time) = { + let row_tt_extended_protocol: PgRow = transaction + .fetch_one("UPDATE sharded SET created_at = NOW() WHERE id = 1 RETURNING *") + .await + .unwrap(); + + let row_tt_simple_protocol: PgRow = sqlx::raw_sql( + "UPDATE sharded SET created_at = CURRENT_TIMESTAMP WHERE id = 2 RETURNING *", + ) + .fetch_one(&mut *transaction) + .await + .unwrap(); + + ( + row_tt_extended_protocol.get::, &str>("created_at"), + row_tt_simple_protocol.get::, &str>("created_at"), + ) + }; + + assert!(transaction_start_time == row_tt_simple_protocol_time); + assert!(transaction_start_time == row_tt_extended_protocol_time); +} + /// 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. diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs b/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs index 912fb7fe8..446395ae1 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs @@ -193,14 +193,13 @@ impl<'a> StatementRewrite<'a> { if nd_function_rewrite { match stmt.stmt_mut() { - // TODO: we could also support UPDATE / etc - NodeMut::InsertStmt(_) | NodeMut::SelectStmt(_) => { + NodeMut::InsertStmt(_) | NodeMut::SelectStmt(_) | NodeMut::UpdateStmt(_) => { 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(_) + NodeMut::InsertStmt(_) | NodeMut::SelectStmt(_) | NodeMut::UpdateStmt(_) ) { 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 d09d0d7ec..16d19c8ac 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 @@ -3,6 +3,7 @@ use pg_raw_parse::{ list::NodeList, make::{MemoryToken, Unique}, raw::SQLValueFunctionOp, + transform::Transform, }; use pgdog_stats::{Column, Relation}; @@ -15,6 +16,7 @@ use crate::{ rewrite::statement::{ Error, non_deterministic_funcs::{ + rewrite_transaction_time::ReplaceTransactionTime, time_function::TimeFunctionType, uuid_function::UUIDFunctionType, }, plan::BindParam, @@ -348,20 +350,35 @@ impl StatementRewrite<'_> { // 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 { + } else if matches!(stmt, NodeMut::SelectStmt(_) | NodeMut::UpdateStmt(_)) { // Still rewrite other kind of statements that use ND functions that // are reliant on transaction start time - // TODO: Eventually cover Update, Delete, etc + let statement_type = match stmt { + NodeMut::SelectStmt(_) => StatementType::Select, + NodeMut::UpdateStmt(_) => StatementType::Update, + _ => panic!("only select and update covered"), + }; let mut nd_rewrite = NDRewrite { rewrite: self, bind_params, mem, - statement_type: StatementType::Select, + statement_type, + }; + + let mut transform_tt = ReplaceTransactionTime { + nd_rewrite: &mut nd_rewrite, + outer_error: None, }; - nd_rewrite.transform_func_calls_transaction_time(select_stmt)?; + match stmt { + NodeMut::SelectStmt(select_stmt) => transform_tt.transform_select_stmt(select_stmt), + NodeMut::UpdateStmt(update_stmt) => transform_tt.transform_update_stmt(update_stmt), + _ => panic!("only select and update covered"), + } + + return transform_tt.outer_error.map(Err).unwrap_or(Ok(())); } Ok(()) } @@ -416,6 +433,7 @@ struct NDRewrite<'mem, 'a, 's> { enum StatementType { Insert, Select, + Update, } /// Extra context necessary for re-writing INSERT statements, where we need @@ -448,7 +466,7 @@ impl<'mem, 'a, 's> NDRewrite<'mem, 'a, 's> { .uncast(); match self.statement_type { - StatementType::Insert => constant_text_node, + StatementType::Insert | StatementType::Update => constant_text_node, StatementType::Select => self .mem .make_type_cast( 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 index 7317af051..bd26b8c6b 100644 --- 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 @@ -2,7 +2,6 @@ use std::ops::Deref; use pg_raw_parse::{ Node, - nodes::SelectStmtMut, transform::{self, Transform}, }; @@ -11,31 +10,17 @@ use crate::frontend::router::parser::rewrite::statement::{ 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> { +pub(super) struct ReplaceTransactionTime<'mutr, 'mem, 'a, 's> { /// Ability to reference `self` within the Transform impl. - nd_rewrite: &'mutr mut NDRewrite<'mem, 'a, 's>, + pub(super) 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, + pub(super) outer_error: Option, } -impl<'mutr, 'mem, 'a, 's> Transform<'mem> for ReplaceTransactionTimeSelect<'mutr, 'mem, 'a, 's> { +impl<'mutr, 'mem, 'a, 's> Transform<'mem> for ReplaceTransactionTime<'mutr, 'mem, 'a, 's> { /// Case, basic: SELECT now() /// /// This means now is a ResTarget, and Postgres will output the timestamptz w/ **a now column**