Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
60 changes: 60 additions & 0 deletions integration/rust/tests/integration/non_deterministic_funcs.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<Utc> = {
conn.execute("TRUNCATE sharded").await.unwrap();

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Don't need to address this in this PR, but we should start writing tests in such a way so we can run them in parallel. Easiest way to do this is to create dedicated fixtures for each test; our config would allow this since we shard by column, not by table and as long as the table has a sharding key, it'll work.


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<Utc> = {
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::<DateTime<Utc>, &str>("created_at"),
row_tt_simple_protocol.get::<DateTime<Utc>, &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.
Expand Down
5 changes: 2 additions & 3 deletions pgdog/src/frontend/router/parser/rewrite/statement/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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)?;
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ use pg_raw_parse::{
list::NodeList,
make::{MemoryToken, Unique},
raw::SQLValueFunctionOp,
transform::Transform,
};
use pgdog_stats::{Column, Relation};

Expand All @@ -15,6 +16,7 @@ use crate::{
rewrite::statement::{
Error,
non_deterministic_funcs::{
rewrite_transaction_time::ReplaceTransactionTime,
time_function::TimeFunctionType, uuid_function::UUIDFunctionType,
},
plan::BindParam,
Expand Down Expand Up @@ -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(())
}
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,6 @@ use std::ops::Deref;

use pg_raw_parse::{
Node,
nodes::SelectStmtMut,
transform::{self, Transform},
};

Expand All @@ -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 <https://github.com/pgdogdev/pg_raw_parse/blob/f63e7f49d85612e4507081e52fc8f40349b70580/src/transform.rs#L107>
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<Error>,
pub(super) outer_error: Option<Error>,
}

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**
Expand Down
Loading