From ef7f79d2a07305d3fadfd91fd9c9547518bb8e12 Mon Sep 17 00:00:00 2001 From: jkaczman Date: Wed, 16 Sep 2026 07:29:54 -0400 Subject: [PATCH 1/2] Early draft / refactor supporting UUID func re-writes for omni writes --- Cargo.lock | 4 + integration/rust/Cargo.toml | 2 +- integration/rust/tests/integration/mod.rs | 2 +- ...mps.rs => omni_non_deterministic_funcs.rs} | 70 ++++ .../router/parser/rewrite/statement/mod.rs | 2 +- .../mod.rs} | 369 ++++-------------- .../statement/non_deterministic_funcs/time.rs | 283 ++++++++++++++ .../statement/non_deterministic_funcs/uuid.rs | 51 +++ .../router/parser/rewrite/statement/plan.rs | 10 +- .../rewrite/statement/simple_prepared.rs | 4 +- 10 files changed, 499 insertions(+), 298 deletions(-) rename integration/rust/tests/integration/{omni_timestamps.rs => omni_non_deterministic_funcs.rs} (90%) rename pgdog/src/frontend/router/parser/rewrite/statement/{timestamp.rs => non_deterministic_funcs/mod.rs} (51%) create mode 100644 pgdog/src/frontend/router/parser/rewrite/statement/non_deterministic_funcs/time.rs create mode 100644 pgdog/src/frontend/router/parser/rewrite/statement/non_deterministic_funcs/uuid.rs diff --git a/Cargo.lock b/Cargo.lock index 5204b3aff..43cabcb68 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -4533,6 +4533,7 @@ dependencies = [ "tokio-stream", "tracing", "url", + "uuid", ] [[package]] @@ -4615,6 +4616,7 @@ dependencies = [ "stringprep", "thiserror", "tracing", + "uuid", "whoami 1.6.1", ] @@ -4656,6 +4658,7 @@ dependencies = [ "stringprep", "thiserror", "tracing", + "uuid", "whoami 1.6.1", ] @@ -4682,6 +4685,7 @@ dependencies = [ "thiserror", "tracing", "url", + "uuid", ] [[package]] diff --git a/integration/rust/Cargo.toml b/integration/rust/Cargo.toml index ff2ee4347..126b2e3b2 100644 --- a/integration/rust/Cargo.toml +++ b/integration/rust/Cargo.toml @@ -10,7 +10,7 @@ test = true tokio-postgres = {version = "0.7.13", features = ["with-uuid-1"]} postgres-native-tls = "0.5" native-tls = "0.2" -sqlx = { version = "0.8.6", features = ["postgres", "runtime-tokio", "tls-native-tls", "bigdecimal", "chrono", "json", "rust_decimal"]} +sqlx = { version = "0.8.6", features = ["postgres", "runtime-tokio", "tls-native-tls", "bigdecimal", "chrono", "json", "rust_decimal", "uuid"]} tokio = { version = "1", features = ["full"]} futures-util.workspace = true uuid.workspace = true diff --git a/integration/rust/tests/integration/mod.rs b/integration/rust/tests/integration/mod.rs index 75b623780..59baf2780 100644 --- a/integration/rust/tests/integration/mod.rs +++ b/integration/rust/tests/integration/mod.rs @@ -24,7 +24,7 @@ pub mod max; pub mod multi_set; pub mod notify; pub mod offset; -pub mod omni_timestamps; +pub mod omni_non_deterministic_funcs; pub mod partial_req; pub mod per_stmt_routing; pub mod prepared; diff --git a/integration/rust/tests/integration/omni_timestamps.rs b/integration/rust/tests/integration/omni_non_deterministic_funcs.rs similarity index 90% rename from integration/rust/tests/integration/omni_timestamps.rs rename to integration/rust/tests/integration/omni_non_deterministic_funcs.rs index 77a12f32e..d3c209661 100644 --- a/integration/rust/tests/integration/omni_timestamps.rs +++ b/integration/rust/tests/integration/omni_non_deterministic_funcs.rs @@ -24,6 +24,76 @@ use sqlx::{Executor, Row}; // TODO: Test to make sure this doesn't affect harded tables (it doesn't; but doesn't hurt to assert that) // TODO: Assert what happens if we don't explicitly set timezone +// +#[tokio::test] +async fn omni_uuid_rewrite() { + let sharded_conn = connections_sqlx().await; + let sharded_conn = sharded_conn.get(1).unwrap(); + + // TODO: Check that UUID v7 shift interval is parsed correctly (and works without errors) + + // Create a test table through PgDog, and then reload, so that the Schema is loaded. + { + sharded_conn + .execute("DROP TABLE IF EXISTS test_omni_uuid") + .await + .unwrap(); + sharded_conn + .execute( + "CREATE TABLE IF NOT EXISTS test_omni_uuid( + id BIGSERIAL PRIMARY KEY, + uuid4 uuid, + uuid7 uuid, + uuid4_default uuid DEFAULT gen_random_uuid(), + uuid7_default uuid DEFAULT uuidv7(), + uuid7_default_explicit uuid DEFAULT uuidv7())", + ) + .await + .unwrap(); + + admin_sqlx().await.execute("RELOAD").await.unwrap(); + } + + // Other functions below test that general rewrites work in all protocols. + // It would be redundant to test that here, as they all share the same re-usable structure. + { + let mut transaction = sharded_conn.begin().await.unwrap(); + + transaction + .execute( + "INSERT INTO test_omni_uuid(id, uuid4, uuid7, uuid7_default_explicit) + VALUES(1, uuidv4(), uuidv7(), DEFAULT)", + ) + .await + .unwrap(); + + let (shard_0_row, shard_1_row) = ( + transaction + .fetch_one("/* pgdog_shard: 0 */ SELECT * FROM test_omni_uuid") + .await + .unwrap(), + transaction + .fetch_one("/* pgdog_shard: 1 */ SELECT * FROM test_omni_uuid") + .await + .unwrap(), + ); + + assert_eq!(shard_0_row.columns().len(), shard_1_row.columns().len()); + + // Iterate through all columns ensuring the UUIDs are consistent across shards for each case. + for col_num in 1..shard_0_row.columns().len() { + let (shard_0_uuid, shard_1_uuid) = ( + shard_0_row.get::(col_num), + shard_1_row.get::(col_num), + ); + + assert_eq!(shard_0_uuid, shard_1_uuid); + } + + transaction.rollback().await.unwrap(); + } +} + /// LOCAL_TIME testing /// - Case 1: `test_time_text` has no DEFAULT w/ precision arg & text col. /// - Case 2: `test_time_regular` has DEFAULT w/ no precision arg & time col. diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs b/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs index f49ab68e6..40acccb3a 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs @@ -14,10 +14,10 @@ pub(crate) mod auto_id; pub(crate) mod error; pub(crate) mod insert; pub(crate) mod nextval; +pub(crate) mod non_deterministic_funcs; pub(crate) mod offset; pub(crate) mod plan; pub(crate) mod simple_prepared; -pub(crate) mod timestamp; pub(crate) mod unique_id; pub(crate) mod update; diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/timestamp.rs b/pgdog/src/frontend/router/parser/rewrite/statement/non_deterministic_funcs/mod.rs similarity index 51% rename from pgdog/src/frontend/router/parser/rewrite/statement/timestamp.rs rename to pgdog/src/frontend/router/parser/rewrite/statement/non_deterministic_funcs/mod.rs index 356af71d2..fe9b33bdd 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/timestamp.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/non_deterministic_funcs/mod.rs @@ -1,8 +1,5 @@ -use std::fmt; use std::ops::Deref; -use chrono::{DateTime, Offset, SubsecRound, TimeZone, Timelike, Utc}; -use chrono_tz::Tz; use pg_raw_parse::{ ConstValue, Node, NodeMut, list::NodeList, @@ -21,6 +18,7 @@ use crate::{ StatementParser, StatementRewrite, Table, rewrite::statement::{ Error, + non_deterministic_funcs::{time::TimeFunctionType, uuid::UUIDFunctionType}, plan::{GeneratedId, GeneratedParam}, }, }, @@ -28,50 +26,38 @@ use crate::{ net::parameter::ParameterValue, }; -/// A "parsed" time function the Client specified; either from database schema or manual commands. -/// Column type represents the data type attached, so that we can correctly assemble the String. -#[derive(Debug, Clone, PartialEq, Eq)] -pub(crate) struct TimeFunction { - time_function_type: TimeFunctionType, +/// No need to expose these outside. +mod time; +mod uuid; + +/// A non-deterministic function that we must re-write when writing to an omnisharded table, +/// so that we can maintain consistency instead of generating a different value (from executing the function) +/// on each shard. This re-writes function calls to a constant. +#[derive(Debug, PartialEq, Eq, Clone)] +pub(crate) struct NDFunction { + nd_function_type: NDFunctionType, + /// TODO: What happens if the database schema is changed? This should be invalidated if cached. column_type: String, } -impl TimeFunction { - /// Based on the internal values (col type, arguments passed, ...), - /// and the reference time (e.g., transaction time), generate - /// the String and binary equivalent to be put in the final String. - pub(crate) fn formatted_time( +impl NDFunction { + /// Generate a Postgres-ready String and binary equivalent for the non-deterministic function. + /// Binary format is always text-based as we explicitly inner-cast the ParamRefs we use with ::text + /// (and have an outer re-cast into the actual column data type) + pub(crate) fn write_as_constant( &self, timestamps: &QueryTimestamps, timezone_param: Option<&ParameterValue>, ) -> Result<(String, Vec), Error> { - let timestamp = self.column_type.eq("timestamp without time zone"); - - let reference_time = match self.time_function_type.time_reference() { - TimeReference::Current => Utc::now(), - TimeReference::TransactionStart => timestamps.transaction_start, - TimeReference::StatementStart => timestamps.statement_start, - }; - - let mut time_output: TimeFunctionOutput = self.time_function_type.default_output_type(); - - // Column expects 'timestamp', function outputs 'timestamptz', need to convert. - if time_output == TimeFunctionOutput::TimestampWithTimeZone && timestamp { - time_output = TimeFunctionOutput::Timestamp - } - - let precision = self.time_function_type.precision(); - - let tz = match session_time_zone(timezone_param) { - Ok(tz) => tz, - Err(_) if time_output == TimeFunctionOutput::TimestampWithTimeZone => Tz::UTC, - Err(err) => return Err(err), - }; + let formatted_string = match self.nd_function_type { + NDFunctionType::TimeFunction(tf) => { + tf.formatted_time(&self.column_type, timestamps, timezone_param) + } + NDFunctionType::UUIDFunction(uuid) => uuid.format(), + }?; - let formatted_string = time_output.format(&reference_time, &tz, precision); let binary = formatted_string.as_bytes().to_vec(); - Ok((formatted_string, binary)) } @@ -90,190 +76,44 @@ impl TimeFunction { "timestamp with time zone" => "timestamptz", "timestamp without time zone" => "timestamp", "time without time zone" => "time", - string => string, - } - } -} - -/// The session's `TimeZone` as an IANA name, e.g. `UTC` or `America/New_York`. -/// -/// Offsets like `+00:00` or `<-08>+08` are future work. -fn session_time_zone(timezone: Option<&ParameterValue>) -> Result { - let timezone = timezone.ok_or(Error::UnknownTimeZone)?; - timezone - .as_str() - .and_then(|value| value.parse::().ok()) - .ok_or_else(|| Error::UnsupportedTimeZone(timezone.to_string())) -} - -/// Postgres trims trailing zeros from fractional seconds -/// It also drops the dot when there's none -fn fractional_seconds(nanoseconds: u32) -> String { - let microseconds = nanoseconds / 1_000; - - if microseconds == 0 { - return String::new(); - } - - format!(".{microseconds:06}") - .trim_end_matches('0') - .to_string() -} - -/// Postgres prints UTC offsets as +HH -/// Adds :MM and :SS when they're non-zero. -fn utc_offset(local_minus_utc: i32) -> String { - let sign = if local_minus_utc < 0 { '-' } else { '+' }; - let total_seconds = local_minus_utc.unsigned_abs(); - let (hours, minutes, seconds) = ( - total_seconds / 3600, - total_seconds / 60 % 60, - total_seconds % 60, - ); - - match (minutes, seconds) { - (0, 0) => format!("{sign}{hours:02}"), - (_, 0) => format!("{sign}{hours:02}:{minutes:02}"), - _ => format!("{sign}{hours:02}:{minutes:02}:{seconds:02}"), - } -} - -/// Represents the kind of `TimeFunction` that we're re-writing. -/// If an Option argument is present and Some(..), the Client specified precision. -/// -#[derive(Debug, Copy, Clone, PartialEq, Eq)] -pub(crate) enum TimeFunctionType { - CurrentDate, - CurrentTime(Option), - CurrentTimestamp(Option), - ClockTimestamp, - LocalTime(Option), - LocalTimestamp(Option), - Now, - StatementTimestamp, - TimeOfDay, - TransactionTimestamp, -} - -/// Represents what Postgres type the `TimeFunction` would normally output. -#[derive(PartialEq)] -enum TimeFunctionOutput { - Date, - TimeWithTimeZone, - TimestampWithTimeZone, - Time, - Timestamp, - - /// Specially formatted (e.g. EST instead of -05) as it's intended for a text col. - TextFormattedTimestampWithTimeZone, -} - -impl TimeFunctionOutput { - /// Formats `utc_time` how Postgres outputs for this time func. Timezone taken into account (`tz`). - /// Seconds with fractions rounded to `precision` (which are capped by Postgres at 6) - fn format(&self, utc_time: &DateTime, tz: &Z, precision: u8) -> String - where - Z: TimeZone, - Z::Offset: fmt::Display, - { - let local_time = utc_time.with_timezone(tz); - let rounded = local_time - .clone() - .round_subsecs(u16::from(precision.min(6))); - - let date = rounded.format("%Y-%m-%d"); - let time = format!( - "{}{}", - rounded.format("%H:%M:%S"), - fractional_seconds(rounded.nanosecond()) - ); - let offset = utc_offset(rounded.offset().fix().local_minus_utc()); - - match self { - Self::Date => local_time.format("%Y-%m-%d").to_string(), - Self::Time => time, - Self::TimeWithTimeZone => format!("{time}{offset}"), - Self::Timestamp => format!("{date} {time}"), - Self::TimestampWithTimeZone => format!("{date} {time}{offset}"), - Self::TextFormattedTimestampWithTimeZone => format!( - "{}.{:06} {}", - local_time.format("%a %b %d %H:%M:%S"), - local_time.nanosecond() / 1_000, - local_time.format("%Y %Z"), - ), + other => other, } } } -enum TimeReference { - /// Changes with statement execution - Current, - - TransactionStart, - - /// "returns the start time of the current statement (more specifically, - /// the time of receipt of the latest command message from the client)." - StatementStart, +/// Simple wrapper around `TimeFunctionType` and `UUIDFunctionType` to pass functions through +/// depending on the type of non-deterministic function that was parsed. +#[derive(Debug, Clone, PartialEq, Eq, Copy)] +enum NDFunctionType { + TimeFunction(TimeFunctionType), + UUIDFunction(UUIDFunctionType), } -impl TimeFunctionType { - /// For easy iteration over all enum variants - /// "CurrentTimestamp" is purposefully ordered before "CurrentTime" (+ LocalTimestamp/LocalTime) to prevent - /// partial match bugs with .starts_with (we can't match on NAME() as not all require ()) - const ALL_VARIANTS: [Self; 10] = [ - Self::CurrentDate, - Self::CurrentTimestamp(None), - Self::CurrentTime(None), - Self::ClockTimestamp, - Self::LocalTimestamp(None), - Self::LocalTime(None), - Self::Now, - Self::StatementTimestamp, - Self::TimeOfDay, - Self::TransactionTimestamp, - ]; - - /// The precision of partial seconds that should be displayed in the `TimeFunction`'s output. - fn precision(self) -> u8 { - match self { - Self::CurrentTime(precision) - | Self::LocalTime(precision) - | Self::LocalTimestamp(precision) - | Self::CurrentTimestamp(precision) => precision, - _ => None, - } - .unwrap_or(6) +impl NDFunctionType { + /// Convert `SQLValueFunctionOp` (e.g. current_date, current_time... non ()) to `NDFunctionType` + fn from_sql_value_function(op: SQLValueFunctionOp::Type, typmod: i32) -> Option { + TimeFunctionType::from_sql_value_function(op, typmod) + .map(NDFunctionType::TimeFunction) + .or_else(|| { + UUIDFunctionType::from_sql_value_function(op, typmod) + .map(NDFunctionType::UUIDFunction) + }) } - /// What point of time (current, transaction start, statement start) should we base the - /// `TimeFunction`'s output on? - fn time_reference(self) -> TimeReference { + /// Postgres-formatted function name for this function type. + /// Used for pattern matching to determine what kind of function (if any) is present. + fn name(&self) -> &str { match self { - Self::ClockTimestamp | Self::TimeOfDay => TimeReference::Current, - Self::CurrentDate - | Self::CurrentTime(_) - | Self::CurrentTimestamp(_) - | Self::LocalTime(_) - | Self::LocalTimestamp(_) - | Self::TransactionTimestamp - | Self::Now => TimeReference::TransactionStart, - Self::StatementTimestamp => TimeReference::StatementStart, + Self::TimeFunction(tf) => tf.name(), + Self::UUIDFunction(uuid) => uuid.name(), } } - /// Represents what Postgres type the `TimeFunction` would normally output. - fn default_output_type(self) -> TimeFunctionOutput { + /// If the type has a parameter (for precision), return the same type with that parameter. + fn with_param(&self, param: u8) -> Self { match self { - Self::CurrentTimestamp(_) - | Self::ClockTimestamp - | Self::Now - | Self::StatementTimestamp - | Self::TransactionTimestamp => TimeFunctionOutput::TimestampWithTimeZone, - Self::CurrentTime(_) => TimeFunctionOutput::TimeWithTimeZone, - Self::CurrentDate => TimeFunctionOutput::Date, - Self::LocalTime(_) => TimeFunctionOutput::Time, - Self::LocalTimestamp(_) => TimeFunctionOutput::Timestamp, - Self::TimeOfDay => TimeFunctionOutput::TextFormattedTimestampWithTimeZone, + Self::TimeFunction(tf) => Self::TimeFunction(tf.with_param(param)), + Self::UUIDFunction(uuid) => Self::UUIDFunction(uuid.with_param(param)), } } @@ -292,61 +132,18 @@ impl TimeFunctionType { } Node::SQLValueFunction(func) => Self::from_sql_value_function(func.op, func.typmod), - Node::SetToDefault(_) => { - // If DEFAULT is in a VALUES list; fetch the column based on index. - if let Some(column) = column_relation - && let Ok(time_function_type) = - column.column_default.parse::() - { - return Some(time_function_type); - } + // If DEFAULT is in a VALUES list; fetch the column based on index. + Node::SetToDefault(_) => column_relation + .map(|column| column.column_default.parse::()) + .and_then(Result::ok), - None - } _ => None, } } - - /// Convert `SQLValueFunctionOp` (e.g. current_date, current_time... non ()) to `TimeFunctionType` - fn from_sql_value_function(op: SQLValueFunctionOp::Type, typmod: i32) -> Option { - use SQLValueFunctionOp::*; - - let precision = u8::try_from(typmod).ok(); - - Some(match op { - SVFOP_CURRENT_DATE => Self::CurrentDate, - SVFOP_CURRENT_TIME => Self::CurrentTime(None), - SVFOP_CURRENT_TIME_N => Self::CurrentTime(precision), - SVFOP_CURRENT_TIMESTAMP => Self::CurrentTimestamp(None), - SVFOP_CURRENT_TIMESTAMP_N => Self::CurrentTimestamp(precision), - SVFOP_LOCALTIME => Self::LocalTime(None), - SVFOP_LOCALTIME_N => Self::LocalTime(precision), - SVFOP_LOCALTIMESTAMP => Self::LocalTimestamp(None), - SVFOP_LOCALTIMESTAMP_N => Self::LocalTimestamp(precision), - // Others: CURRENT_USER, CURRENT_SCHEMA... not relevant here - _ => return None, - }) - } - - /// Postgres formatted String to match against Client-provided names in query. - fn name(self) -> &'static str { - match self { - Self::CurrentDate => "current_date", - Self::CurrentTime(_) => "current_time", - Self::CurrentTimestamp(_) => "current_timestamp", - Self::ClockTimestamp => "clock_timestamp", - Self::LocalTime(_) => "localtime", - Self::LocalTimestamp(_) => "localtimestamp", - Self::Now => "now", - Self::StatementTimestamp => "statement_timestamp", - Self::TimeOfDay => "timeofday", - Self::TransactionTimestamp => "transaction_timestamp", - } - } } -/// Client `now()`.parse() -> TimeFunctionType::Now() -impl FromStr for TimeFunctionType { +/// Client `now()`.parse() -> NDFunction::TimeFunction(TimeFunctionType::Now()) +impl FromStr for NDFunctionType { type Err = Option; /// TODO: Doc comment @@ -356,7 +153,10 @@ impl FromStr for TimeFunctionType { fn from_str(s: &str) -> Result { let s = s.to_lowercase(); - for variant in Self::ALL_VARIANTS { + for variant in TimeFunctionType::ALL_VARIANTS + .iter() + .chain(UUIDFunctionType::ALL_VARIANTS.iter()) + { let variant_name = &variant.name(); if s.starts_with(variant_name) { // TODO: I think this can be written better @@ -370,15 +170,9 @@ impl FromStr for TimeFunctionType { Err(_) => continue, }; - return Ok(match variant { - Self::CurrentTime(_) => Self::CurrentTime(Some(after_to_int)), - Self::CurrentTimestamp(_) => Self::CurrentTimestamp(Some(after_to_int)), - Self::LocalTime(_) => Self::LocalTime(Some(after_to_int)), - Self::LocalTimestamp(_) => Self::LocalTimestamp(Some(after_to_int)), - _ => continue, - }); + return Ok(variant.with_param(after_to_int)); } else if after.eq("()") || after.is_empty() { - return Ok(variant); + return Ok(*variant); } else { continue; } @@ -414,7 +208,7 @@ impl StatementRewrite<'_> { return Ok(()); }; - let mut timestamp_rewrite = TimestampRewrite { + let mut nd_rewrite = NDRewrite { rewrite: self, plan, next_param, @@ -426,12 +220,12 @@ impl StatementRewrite<'_> { // 1. iterates through Schema to find DEFAULT columns // 2. adds the column to target list & all the values lists (ParamRef or String) - timestamp_rewrite.handle_adding_defaults(&mut stmt, ¬_covered_cols); + nd_rewrite.handle_adding_defaults(&mut stmt, ¬_covered_cols); // Replaces all time function calls (ParamRef or String) - timestamp_rewrite.transform_func_calls(stmt); + nd_rewrite.transform_func_calls(stmt); - timestamp_rewrite.error.map_or(Ok(()), Err) + nd_rewrite.error.map_or(Ok(()), Err) } /// Fetch the table Relation, so that we can get the relevant Schema for each column. @@ -474,7 +268,7 @@ impl StatementRewrite<'_> { } } -struct TimestampRewrite<'mem, 'a, 's> { +struct NDRewrite<'mem, 'a, 's> { rewrite: &'a mut StatementRewrite<'s>, plan: &'a mut RewritePlan, /// TODO: Replace `next_param` with plan.param directly @@ -488,7 +282,7 @@ struct TimestampRewrite<'mem, 'a, 's> { error: Option, } -impl<'mem, 'a, 's> TimestampRewrite<'mem, 'a, 's> { +impl<'mem, 'a, 's> NDRewrite<'mem, 'a, 's> { /// Replaces all time function calls (ParamRef or String) /// Used by all to handle re-writes. fn transform_func_calls(&mut self, stmt: NodeMut<'mem, '_>) { @@ -518,22 +312,20 @@ impl<'mem, 'a, 's> TimestampRewrite<'mem, 'a, 's> { } }; - if let Some(time_function_type) = - TimeFunctionType::from_node(value, col_relation) + if let Some(nd_function_type) = + NDFunctionType::from_node(value, col_relation) { let Some(col_relation) = col_relation else { continue; }; - let time_function = TimeFunction { - time_function_type, + let nd_function = NDFunction { + nd_function_type, column_type: col_relation.data_type.clone(), }; // Replace the specific node within the list. - cloned_values - .as_mut() - .set(i, self.make_node(&time_function)); + cloned_values.as_mut().set(i, self.make_node(&nd_function)); changed = true; } } @@ -566,13 +358,12 @@ impl<'mem, 'a, 's> TimestampRewrite<'mem, 'a, 's> { for col in not_covered_cols { let col_relation = self.relation.columns.get(col.as_str()).unwrap(); - let Ok(time_function_type) = col_relation.column_default.parse::() - else { + let Ok(nd_function_type) = col_relation.column_default.parse::() else { continue; }; - let time_function = TimeFunction { - time_function_type, + let nd_function = NDFunction { + nd_function_type, column_type: col_relation.data_type.clone(), }; @@ -593,19 +384,19 @@ impl<'mem, 'a, 's> TimestampRewrite<'mem, 'a, 's> { 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(&time_function)); + node_list_mut.push(self.mem, self.make_node(&nd_function)); } } } /// If simple protocol, make an A_Const node with the String constant of the formatted time. /// If extended or prepare, make a ParamRef, so that we can cache it and put in the formatted time later. - fn make_node(&mut self, time_function: &TimeFunction) -> Unique<'mem, Node<'mem>> { + fn make_node(&mut self, nd_function: &NDFunction) -> Unique<'mem, Node<'mem>> { self.rewrite.rewritten = true; if !self.rewrite.extended && !self.rewrite.prepared { - let text = match time_function - .formatted_time(&self.rewrite.query_timestamps, self.rewrite.timezone) + let text = match nd_function + .write_as_constant(&self.rewrite.query_timestamps, self.rewrite.timezone) { Ok((text, _)) => text, // The statement is discarded when the error is returned (thus, value doesn't matter) @@ -624,7 +415,7 @@ impl<'mem, 'a, 's> TimestampRewrite<'mem, 'a, 's> { // TODO: add a method to plan() for this... self.plan.generated_params.push(GeneratedParam { param_num: (*self.next_param - 1) as u16, - generated_id: GeneratedId::ProxyTime(time_function.clone()), + generated_id: GeneratedId::NDFunction(nd_function.clone()), }); // Example: CAST($1::pg_catalog.text AS timetz) @@ -643,7 +434,7 @@ impl<'mem, 'a, 's> TimestampRewrite<'mem, 'a, 's> { .uncast(), self.mem.make_list(&[self .mem - .make_string(Some(time_function.col_type_to_type_cast_alias()))]), + .make_string(Some(nd_function.col_type_to_type_cast_alias()))]), ) .uncast() } 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.rs new file mode 100644 index 000000000..3210a6443 --- /dev/null +++ b/pgdog/src/frontend/router/parser/rewrite/statement/non_deterministic_funcs/time.rs @@ -0,0 +1,283 @@ +use std::fmt; + +use chrono::{DateTime, Offset, SubsecRound, TimeZone, Timelike, Utc}; +use chrono_tz::Tz; + +use crate::{ + frontend::{ + client::QueryTimestamps, + router::parser::rewrite::statement::{Error, non_deterministic_funcs::NDFunctionType}, + }, + net::parameter::ParameterValue, +}; + +use pg_raw_parse::raw::SQLValueFunctionOp; + +/// Represents what Postgres type the `TimeFunction` would normally output. +#[derive(PartialEq)] +enum TimeFunctionOutput { + Date, + TimeWithTimeZone, + TimestampWithTimeZone, + Time, + Timestamp, + + /// Specially formatted (e.g. EST instead of -05) as it's intended for a text col. + TextFormattedTimestampWithTimeZone, +} + +impl TimeFunctionOutput { + /// Formats `utc_time` how Postgres outputs for this time func. Timezone taken into account (`tz`). + /// Seconds with fractions rounded to `precision` (which are capped by Postgres at 6) + fn format(&self, utc_time: &DateTime, tz: &Z, precision: u8) -> String + where + Z: TimeZone, + Z::Offset: fmt::Display, + { + let local_time = utc_time.with_timezone(tz); + let rounded = local_time + .clone() + .round_subsecs(u16::from(precision.min(6))); + + let date = rounded.format("%Y-%m-%d"); + let time = format!( + "{}{}", + rounded.format("%H:%M:%S"), + Self::fractional_seconds(rounded.nanosecond()) + ); + let offset = Self::utc_offset(rounded.offset().fix().local_minus_utc()); + + match self { + Self::Date => local_time.format("%Y-%m-%d").to_string(), + Self::Time => time, + Self::TimeWithTimeZone => format!("{time}{offset}"), + Self::Timestamp => format!("{date} {time}"), + Self::TimestampWithTimeZone => format!("{date} {time}{offset}"), + Self::TextFormattedTimestampWithTimeZone => format!( + "{}.{:06} {}", + local_time.format("%a %b %d %H:%M:%S"), + local_time.nanosecond() / 1_000, + local_time.format("%Y %Z"), + ), + } + } + + /// Postgres trims trailing zeros from fractional seconds + /// It also drops the dot when there's none + fn fractional_seconds(nanoseconds: u32) -> String { + let microseconds = nanoseconds / 1_000; + + if microseconds == 0 { + return String::new(); + } + + format!(".{microseconds:06}") + .trim_end_matches('0') + .to_string() + } + + /// Postgres prints UTC offsets as +HH + /// Adds :MM and :SS when they're non-zero. + fn utc_offset(local_minus_utc: i32) -> String { + let sign = if local_minus_utc < 0 { '-' } else { '+' }; + let total_seconds = local_minus_utc.unsigned_abs(); + let (hours, minutes, seconds) = ( + total_seconds / 3600, + total_seconds / 60 % 60, + total_seconds % 60, + ); + + match (minutes, seconds) { + (0, 0) => format!("{sign}{hours:02}"), + (_, 0) => format!("{sign}{hours:02}:{minutes:02}"), + _ => format!("{sign}{hours:02}:{minutes:02}:{seconds:02}"), + } + } +} + +enum TimeReference { + /// Changes with statement execution + Current, + + TransactionStart, + + /// "returns the start time of the current statement (more specifically, + /// the time of receipt of the latest command message from the client)." + StatementStart, +} + +/// Represents the kind of `TimeFunction` that we're re-writing. +/// If an Option argument is present and Some(..), the Client specified precision. +/// +#[derive(Debug, Copy, Clone, PartialEq, Eq)] +pub(super) enum TimeFunctionType { + CurrentDate, + CurrentTime(Option), + CurrentTimestamp(Option), + ClockTimestamp, + LocalTime(Option), + LocalTimestamp(Option), + Now, + StatementTimestamp, + TimeOfDay, + TransactionTimestamp, +} + +impl TimeFunctionType { + /// For easy iteration over all enum variants + /// "CurrentTimestamp" is purposefully ordered before "CurrentTime" (+ LocalTimestamp/LocalTime) to prevent + /// partial match bugs with .starts_with (we can't match on NAME() as not all require ()) + pub(super) const ALL_VARIANTS: [NDFunctionType; 10] = [ + NDFunctionType::TimeFunction(Self::CurrentDate), + NDFunctionType::TimeFunction(Self::CurrentTimestamp(None)), + NDFunctionType::TimeFunction(Self::CurrentTime(None)), + NDFunctionType::TimeFunction(Self::ClockTimestamp), + NDFunctionType::TimeFunction(Self::LocalTimestamp(None)), + NDFunctionType::TimeFunction(Self::LocalTime(None)), + NDFunctionType::TimeFunction(Self::Now), + NDFunctionType::TimeFunction(Self::StatementTimestamp), + NDFunctionType::TimeFunction(Self::TimeOfDay), + NDFunctionType::TimeFunction(Self::TransactionTimestamp), + ]; + + /// If the type has a parameter (for precision), return the same type with that parameter. + pub(super) fn with_param(self, precision: u8) -> Self { + match self { + Self::CurrentTime(_) => Self::CurrentTime(Some(precision)), + Self::CurrentTimestamp(_) => Self::CurrentTimestamp(Some(precision)), + Self::LocalTime(_) => Self::LocalTime(Some(precision)), + Self::LocalTimestamp(_) => Self::LocalTimestamp(Some(precision)), + _ => self, + } + } + + /// Based on the internal values (col type, arguments passed, ...), + /// and the reference time (e.g., transaction time), generate + /// the String and binary equivalent to be put in the final String. + pub(super) fn formatted_time( + &self, + column_type: &str, + timestamps: &QueryTimestamps, + timezone_param: Option<&ParameterValue>, + ) -> Result { + let timestamp = column_type.eq("timestamp without time zone"); + + let reference_time = match self.time_reference() { + TimeReference::Current => Utc::now(), + TimeReference::TransactionStart => timestamps.transaction_start, + TimeReference::StatementStart => timestamps.statement_start, + }; + + let mut time_output: TimeFunctionOutput = self.default_output_type(); + + // Column expects 'timestamp', function outputs 'timestamptz', need to convert. + if time_output == TimeFunctionOutput::TimestampWithTimeZone && timestamp { + time_output = TimeFunctionOutput::Timestamp + } + + let precision = self.precision(); + + let tz = match Self::session_time_zone(timezone_param) { + Ok(tz) => tz, + Err(_) if time_output == TimeFunctionOutput::TimestampWithTimeZone => Tz::UTC, + Err(err) => return Err(err), + }; + + Ok(time_output.format(&reference_time, &tz, precision)) + } + + /// The precision of partial seconds that should be displayed in the `TimeFunction`'s output. + fn precision(self) -> u8 { + match self { + Self::CurrentTime(precision) + | Self::LocalTime(precision) + | Self::LocalTimestamp(precision) + | Self::CurrentTimestamp(precision) => precision, + _ => None, + } + .unwrap_or(6) + } + + /// What point of time (current, transaction start, statement start) should we base the + /// `TimeFunction`'s output on? + fn time_reference(self) -> TimeReference { + match self { + Self::ClockTimestamp | Self::TimeOfDay => TimeReference::Current, + Self::CurrentDate + | Self::CurrentTime(_) + | Self::CurrentTimestamp(_) + | Self::LocalTime(_) + | Self::LocalTimestamp(_) + | Self::TransactionTimestamp + | Self::Now => TimeReference::TransactionStart, + Self::StatementTimestamp => TimeReference::StatementStart, + } + } + + /// Represents what Postgres type the `TimeFunction` would normally output. + fn default_output_type(self) -> TimeFunctionOutput { + match self { + Self::CurrentTimestamp(_) + | Self::ClockTimestamp + | Self::Now + | Self::StatementTimestamp + | Self::TransactionTimestamp => TimeFunctionOutput::TimestampWithTimeZone, + Self::CurrentTime(_) => TimeFunctionOutput::TimeWithTimeZone, + Self::CurrentDate => TimeFunctionOutput::Date, + Self::LocalTime(_) => TimeFunctionOutput::Time, + Self::LocalTimestamp(_) => TimeFunctionOutput::Timestamp, + Self::TimeOfDay => TimeFunctionOutput::TextFormattedTimestampWithTimeZone, + } + } + + /// The session's `TimeZone` as an IANA name, e.g. `UTC` or `America/New_York`. + /// + /// TODO: Offsets like `+00:00` or `<-08>+08` are future work. + fn session_time_zone(timezone: Option<&ParameterValue>) -> Result { + let timezone = timezone.ok_or(Error::UnknownTimeZone)?; + timezone + .as_str() + .and_then(|value| value.parse::().ok()) + .ok_or_else(|| Error::UnsupportedTimeZone(timezone.to_string())) + } + + /// Postgres formatted String to match against Client-provided names in query. + pub(super) fn name(self) -> &'static str { + match self { + Self::CurrentDate => "current_date", + Self::CurrentTime(_) => "current_time", + Self::CurrentTimestamp(_) => "current_timestamp", + Self::ClockTimestamp => "clock_timestamp", + Self::LocalTime(_) => "localtime", + Self::LocalTimestamp(_) => "localtimestamp", + Self::Now => "now", + Self::StatementTimestamp => "statement_timestamp", + Self::TimeOfDay => "timeofday", + Self::TransactionTimestamp => "transaction_timestamp", + } + } + + /// Convert `SQLValueFunctionOp` (e.g. current_date, current_time... non ()) to `TimeFunctionType` + pub(super) fn from_sql_value_function( + op: SQLValueFunctionOp::Type, + typmod: i32, + ) -> Option { + use SQLValueFunctionOp::*; + + let precision = u8::try_from(typmod).ok(); + + Some(match op { + SVFOP_CURRENT_DATE => Self::CurrentDate, + SVFOP_CURRENT_TIME => Self::CurrentTime(None), + SVFOP_CURRENT_TIME_N => Self::CurrentTime(precision), + SVFOP_CURRENT_TIMESTAMP => Self::CurrentTimestamp(None), + SVFOP_CURRENT_TIMESTAMP_N => Self::CurrentTimestamp(precision), + SVFOP_LOCALTIME => Self::LocalTime(None), + SVFOP_LOCALTIME_N => Self::LocalTime(precision), + SVFOP_LOCALTIMESTAMP => Self::LocalTimestamp(None), + SVFOP_LOCALTIMESTAMP_N => Self::LocalTimestamp(precision), + // Others: CURRENT_USER, CURRENT_SCHEMA... not relevant here + _ => return None, + }) + } +} 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.rs new file mode 100644 index 000000000..5d509396d --- /dev/null +++ b/pgdog/src/frontend/router/parser/rewrite/statement/non_deterministic_funcs/uuid.rs @@ -0,0 +1,51 @@ +use pg_raw_parse::raw::SQLValueFunctionOp; + +use crate::frontend::router::parser::rewrite::statement::{ + Error, non_deterministic_funcs::NDFunctionType, +}; + +/// TODO: Docs. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(super) enum UUIDFunctionType { + Uuidv4, + Uuidv7, // TODO: Accept a smallint param. + GenRandomUuid, +} + +impl UUIDFunctionType { + /// For easy iteration over all enum variants for pattern matching. + pub(super) const ALL_VARIANTS: [NDFunctionType; 3] = [ + NDFunctionType::UUIDFunction(Self::Uuidv4), + NDFunctionType::UUIDFunction(Self::Uuidv7), + NDFunctionType::UUIDFunction(Self::GenRandomUuid), + ]; + + /// Convert `SQLValueFunctionOp` (e.g. current_date, current_time... non ()) to `UUIDFunctionType` + /// There are no such cases for UUID functions. + pub(super) fn from_sql_value_function( + _op: SQLValueFunctionOp::Type, + _typmod: i32, + ) -> Option { + None + } + + /// Postgres formatted String to match against Client-provided names in query. + pub(super) fn name(self) -> &'static str { + match self { + Self::Uuidv4 => "uuidv4", + Self::Uuidv7 => "uuidv7", + Self::GenRandomUuid => "gen_random_uuid", + } + } + + /// If the type has a parameter (for precision), return the same type with that parameter. + /// TODO: Handle this for UUIDv7. + pub(super) fn with_param(self, _precision: u8) -> Self { + self + } + + pub(super) fn format(self) -> Result { + // TODO: Generate a random UUIDv4 / UUIDv7 based on the `UUIDFunctionType` + todo!() + } +} diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs b/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs index 9f5bbb7ed..629e60e87 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs @@ -6,7 +6,7 @@ use super::{ Error, InsertSplit, PrepareExecute, ShardingKeyUpdate, aggregate::AggregateRewritePlan, }; use crate::frontend::client::QueryTimestamps; -use crate::frontend::router::parser::rewrite::statement::timestamp::TimeFunction; +use crate::frontend::router::parser::rewrite::statement::non_deterministic_funcs::NDFunction; use crate::frontend::{ClientRequest, PreparedStatements}; use crate::net::messages::bind::{Format, Parameter}; use crate::net::{Bind, Parse, ProtocolMessage, Query, parameter::ParameterValue}; @@ -24,7 +24,9 @@ pub(crate) struct GeneratedParam { pub(crate) enum GeneratedId { UniqueId, Sequence(SequenceCall), - ProxyTime(TimeFunction), + /// This represents a function (such as date/time, UUID) that was re-written to a constant + /// to be consistent across shards for omni writes. + NDFunction(NDFunction), } /// Statement rewrite plan. @@ -142,8 +144,8 @@ impl RewritePlan { GeneratedId::Sequence(call) => { Self::convert_int_to_param(execute(call).await?, format) } - GeneratedId::ProxyTime(time) => { - let (text, binary) = time.formatted_time(×tamps, timezone)?; + GeneratedId::NDFunction(nd_func) => { + let (text, binary) = nd_func.write_as_constant(×tamps, timezone)?; match format { Format::Binary => Parameter::new(binary.as_slice()), Format::Text => Parameter::new(text.as_bytes()), 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 2d7ee82db..92c98453c 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/simple_prepared.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/simple_prepared.rs @@ -191,8 +191,8 @@ impl StatementRewrite<'_> { ) -> Result<(), Error> { for param in generated_params { let (text, _) = match ¶m.generated_id { - GeneratedId::ProxyTime(time) => { - time.formatted_time(&self.query_timestamps, timezone)? + GeneratedId::NDFunction(nd_func) => { + nd_func.write_as_constant(&self.query_timestamps, timezone)? } // TODO: It seems very straightforward to support the rest (if we want to support them for PREPARE) _ => continue, From 1d6511c026370844ad60374424928e3e0d315623 Mon Sep 17 00:00:00 2001 From: jkaczman Date: Wed, 16 Sep 2026 15:06:31 -0400 Subject: [PATCH 2/2] Finish support for basic uuidv4(), uuidv7(), gen_random_uuid(). --- Cargo.toml | 2 +- .../omni_non_deterministic_funcs.rs | 22 +- .../router/parser/rewrite/statement/error.rs | 5 + .../statement/non_deterministic_funcs/mod.rs | 212 ++++++++++++------ .../statement/non_deterministic_funcs/uuid.rs | 32 ++- 5 files changed, 191 insertions(+), 82 deletions(-) diff --git a/Cargo.toml b/Cargo.toml index 41fff8e2a..c4d70ad10 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -31,7 +31,7 @@ indexmap = { version = "2.14", features = ["serde"] } serde_with = { version = "3.22", features = ["macros", "schemars_1"] } reqwest = { version = "0.13", features = ["json"] } bytes = "1" -uuid = { version = "1", features = ["v4", "serde"] } +uuid = { version = "1", features = ["v4", "v7", "serde"] } parking_lot = "0.12" futures-util = "0.3.32" diff --git a/integration/rust/tests/integration/omni_non_deterministic_funcs.rs b/integration/rust/tests/integration/omni_non_deterministic_funcs.rs index d3c209661..670d2f053 100644 --- a/integration/rust/tests/integration/omni_non_deterministic_funcs.rs +++ b/integration/rust/tests/integration/omni_non_deterministic_funcs.rs @@ -24,7 +24,14 @@ use sqlx::{Executor, Row}; // TODO: Test to make sure this doesn't affect harded tables (it doesn't; but doesn't hurt to assert that) // TODO: Assert what happens if we don't explicitly set timezone -// +/// Test that an INSERT into an omnisharded table which uses UUID functions is +/// re-written to a constant to be consistent across all shards. +/// +/// - Verifies DEFAULT works (in schema, in VALUES) +/// - Verifies functions in VALUES work. +/// - Tests uuidv4(), uuidv7(), gen_random_uuid() +/// +/// It also asserts that uuidv7(interval) does **NOT** work right now, as we don't have an easy way to parse intervals. #[tokio::test] async fn omni_uuid_rewrite() { let sharded_conn = connections_sqlx().await; @@ -90,6 +97,19 @@ async fn omni_uuid_rewrite() { assert_eq!(shard_0_uuid, shard_1_uuid); } + // Specify an INTERVAL as an argument within uuidv7. This should fail with an Error. + // It would require us to parse Postgres intervals (possible, but not supported yet) + let err = transaction + .execute( + "INSERT INTO test_omni_uuid(id, uuid4, uuid7, uuid7_default_explicit) + VALUES(2, uuidv4(), uuidv7(INTERVAL '-2 weeks'), DEFAULT)", + ) + .await + .err() + .unwrap(); + + assert!(err.to_string().contains("parser: rewrite: could not determine how to parse the argument passed in uuidv7; it is likely not supported yet")); + transaction.rollback().await.unwrap(); } } diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/error.rs b/pgdog/src/frontend/router/parser/rewrite/statement/error.rs index 8c576c0f4..2948adc15 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/error.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/error.rs @@ -54,4 +54,9 @@ pub(crate) enum Error { #[error("could not determine the session TimeZone for time functions on omnisharded tables")] UnknownTimeZone, + + #[error( + "could not determine how to parse the argument passed in {0}; it is likely not supported yet" + )] + UnsupportedArgument(String), } 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 fe9b33bdd..3df2d3400 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 @@ -8,7 +8,6 @@ use pg_raw_parse::{ transform::{TransformClosure, transform_node}, }; use pgdog_stats::{Column, Relation}; -use std::str::FromStr; use crate::{ frontend::{ @@ -113,7 +112,7 @@ impl NDFunctionType { fn with_param(&self, param: u8) -> Self { match self { Self::TimeFunction(tf) => Self::TimeFunction(tf.with_param(param)), - Self::UUIDFunction(uuid) => Self::UUIDFunction(uuid.with_param(param)), + Self::UUIDFunction(uuid) => Self::UUIDFunction(uuid.with_param()), } } @@ -122,64 +121,111 @@ impl NDFunctionType { /// Parse both `FuncCall`s and `SQLValueFunction`s here. /// `now()` = `FuncCall`, /// `CURRENT_TIMESTAMP`, `LOCALTIME` = `SQLValueFunction`, - fn from_node(node: Node, column_relation: Option<&Column>) -> Option { + fn from_node(node: Node, column_relation: Option<&Column>) -> Result, Error> { match node { Node::FuncCall(func) => { - let Node::String(str) = func.funcname().first()? else { - return None; + let Some(Node::String(str)) = func.funcname().first() else { + return Ok(None); }; - str.sval()?.parse().ok() + + let Some(func_name) = str.sval() else { + return Ok(None); + }; + + Self::from_func_call(func_name, Some(func.args())) } - Node::SQLValueFunction(func) => Self::from_sql_value_function(func.op, func.typmod), + Node::SQLValueFunction(func) => Ok(Self::from_sql_value_function(func.op, func.typmod)), // If DEFAULT is in a VALUES list; fetch the column based on index. - Node::SetToDefault(_) => column_relation - .map(|column| column.column_default.parse::()) - .and_then(Result::ok), + Node::SetToDefault(_) => { + if let Some(relation) = column_relation { + return Self::from_func_call(relation.column_default.as_str(), None); + } - _ => None, + Ok(None) + } + + _ => Ok(None), } } -} -/// Client `now()`.parse() -> NDFunction::TimeFunction(TimeFunctionType::Now()) -impl FromStr for NDFunctionType { - type Err = Option; - - /// TODO: Doc comment - /// Not sure it's necessary to error in this circumstance. - /// Would mean they didn't correctly call the function; - /// Postgres will error them out (unless we have a logic bug) - fn from_str(s: &str) -> Result { - let s = s.to_lowercase(); + /// Parse the function to determine if it's one we should re-write (date/time, uuid) + /// If it is, return the corresponding `NDFunctionType`. + /// + /// from_str(Client `now()`) -> NDFunction::TimeFunction(TimeFunctionType::Now()) + /// This doesn't use the FromStr trait because I wanted to return an `Option` type + /// + /// Returns an Error when the argument is not supported by us yet (e.g. intervals for uuidv7()) + fn from_func_call(func_name: &str, args: Option<&NodeList>) -> Result, Error> { + // Normalize the function name. Postgres does this. + let func_name = func_name.to_lowercase(); for variant in TimeFunctionType::ALL_VARIANTS .iter() .chain(UUIDFunctionType::ALL_VARIANTS.iter()) { let variant_name = &variant.name(); - if s.starts_with(variant_name) { - // TODO: I think this can be written better - let after = s.replace(' ', ""); - let after = &after[variant_name.len()..]; - if after.starts_with('(') && after.ends_with(')') && after.len() >= 3 { - let after = &after[1..after.len() - 1]; - - let after_to_int: u8 = match after.parse() { - Ok(integer_argument) => integer_argument, - Err(_) => continue, + if func_name.starts_with(variant_name) { + // If `args` is None, this means that it was passed by Schema (DEFAULT), where we haven't run pg_raw_parse + // TODO: pg_raw_parse could (potentially) parse this instead. + if args.is_none() { + let func_name = func_name.replace(' ', ""); + let arguments = &func_name[variant_name.len()..]; + + if arguments.starts_with('(') + && arguments.ends_with(')') + && arguments.len() >= 3 + { + let arguments = &arguments[1..arguments.len() - 1]; + + let after_to_int: u8 = match arguments.parse() { + Ok(integer_argument) => integer_argument, + Err(_) => { + // Unsupported argument in Schema + return Err(Error::UnsupportedArgument(variant_name.to_string())); + } + }; + + return Ok(Some(variant.with_param(after_to_int))); + } else if arguments.eq("()") || arguments.is_empty() { + return Ok(Some(*variant)); + } else { + continue; + } + } else if let Some(args) = args { + if args.len() >= 2 { + // Only support a singular parameter right now; + // none of our current functions require more than that. + return Err(Error::UnsupportedArgument(variant_name.to_string())); + } + + let Some(first_arg) = args.get(0) else { + // No arguments. As-is. + return Ok(Some(*variant)); }; - return Ok(variant.with_param(after_to_int)); - } else if after.eq("()") || after.is_empty() { - return Ok(*variant); - } else { - continue; + if let Node::A_Const(constant) = first_arg + && let Some(constant) = constant.val() + && let Some(constant_number) = constant.numeric_value::() + { + match constant_number.try_into() { + Ok(constant_number) => { + return Ok(Some(variant.with_param(constant_number))); + } + Err(_) => { + return Err(Error::UnsupportedArgument(variant_name.to_string())); + } + } + } else { + // Only support parsing out a numeric constant right now. + // TODO: uuidv7 param + return Err(Error::UnsupportedArgument(variant_name.to_string())); + } } } } - Err(None) + Ok(None) } } @@ -215,17 +261,16 @@ impl StatementRewrite<'_> { mem, relation, cols, - error: None, }; // 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); + nd_rewrite.handle_adding_defaults(&mut stmt, ¬_covered_cols)?; // Replaces all time function calls (ParamRef or String) - nd_rewrite.transform_func_calls(stmt); + nd_rewrite.transform_func_calls(stmt)?; - nd_rewrite.error.map_or(Ok(()), Err) + Ok(()) } /// Fetch the table Relation, so that we can get the relevant Schema for each column. @@ -277,15 +322,15 @@ struct NDRewrite<'mem, 'a, 's> { mem: MemoryToken<'mem>, relation: Relation, cols: Unique<'mem, &'mem NodeList>, - - // Error from formatting a time. - error: Option, } impl<'mem, 'a, 's> NDRewrite<'mem, 'a, 's> { /// Replaces all time function calls (ParamRef or String) /// Used by all to handle re-writes. - fn transform_func_calls(&mut self, stmt: NodeMut<'mem, '_>) { + 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 { @@ -312,21 +357,37 @@ impl<'mem, 'a, 's> NDRewrite<'mem, 'a, 's> { } }; - if let Some(nd_function_type) = - NDFunctionType::from_node(value, col_relation) - { - let Some(col_relation) = col_relation else { - continue; - }; - - let nd_function = NDFunction { - nd_function_type, - column_type: col_relation.data_type.clone(), - }; - - // Replace the specific node within the list. - cloned_values.as_mut().set(i, self.make_node(&nd_function)); - changed = true; + match NDFunctionType::from_node(value, col_relation) { + 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; + } } } @@ -343,6 +404,8 @@ impl<'mem, 'a, 's> NDRewrite<'mem, 'a, 's> { _ => Some(node), }), ); + + err.map(Err).unwrap_or(Ok(())) } /// Iterates through Schema to find DEFAULT columns @@ -351,14 +414,18 @@ impl<'mem, 'a, 's> NDRewrite<'mem, 'a, 's> { &mut self, mut stmt: &mut NodeMut<'mem, '_>, not_covered_cols: &Vec, - ) { + ) -> Result<(), Error> { let NodeMut::InsertStmt(insert_stmt) = &mut stmt else { - return; + return Ok(()); }; for col in not_covered_cols { let col_relation = self.relation.columns.get(col.as_str()).unwrap(); - let Ok(nd_function_type) = col_relation.column_default.parse::() else { + + let nd_function_type = + NDFunctionType::from_func_call(&col_relation.column_default, None)?; + + let Some(nd_function_type) = nd_function_type else { continue; }; @@ -376,7 +443,7 @@ impl<'mem, 'a, 's> NDRewrite<'mem, 'a, 's> { ); let NodeMut::SelectStmt(select_stmt) = &mut insert_stmt.select_stmt_mut() else { - return; + return Ok(()); }; // Have to add the now() to every single select VALUES list now. @@ -384,26 +451,25 @@ impl<'mem, 'a, 's> NDRewrite<'mem, 'a, 's> { 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)); + 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 time. /// If extended or prepare, make a ParamRef, so that we can cache it and put in the formatted time later. - fn make_node(&mut self, nd_function: &NDFunction) -> Unique<'mem, Node<'mem>> { + fn make_node(&mut self, nd_function: &NDFunction) -> Result>, Error> { self.rewrite.rewritten = true; - if !self.rewrite.extended && !self.rewrite.prepared { + Ok(if !self.rewrite.extended && !self.rewrite.prepared { let text = match nd_function .write_as_constant(&self.rewrite.query_timestamps, self.rewrite.timezone) { Ok((text, _)) => text, // The statement is discarded when the error is returned (thus, value doesn't matter) - Err(err) => { - self.error.get_or_insert(err); - String::new() - } + Err(err) => return Err(err), }; self.mem .make_a_const(ConstValue::String(text.as_str())) @@ -437,6 +503,6 @@ impl<'mem, 'a, 's> NDRewrite<'mem, 'a, 's> { .make_string(Some(nd_function.col_type_to_type_cast_alias()))]), ) .uncast() - } + }) } } 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.rs index 5d509396d..3685d1781 100644 --- 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.rs @@ -1,14 +1,17 @@ +use chrono::Utc; use pg_raw_parse::raw::SQLValueFunctionOp; +use uuid::{ContextV7, Timestamp, Uuid}; use crate::frontend::router::parser::rewrite::statement::{ Error, non_deterministic_funcs::NDFunctionType, }; -/// TODO: Docs. +/// Represents the kind of `UUIDFunction` that we're re-writing. +/// #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub(super) enum UUIDFunctionType { Uuidv4, - Uuidv7, // TODO: Accept a smallint param. + Uuidv7, //TODO: Postgres supports a parameter for an interval to be specified to shift the timestamp. GenRandomUuid, } @@ -38,14 +41,29 @@ impl UUIDFunctionType { } } - /// If the type has a parameter (for precision), return the same type with that parameter. - /// TODO: Handle this for UUIDv7. - pub(super) fn with_param(self, _precision: u8) -> Self { + /// If the type has a parameter, return the same type with that parameter. + /// TODO: Support intervals for uuidv7 (Param enum to generalize precision / interval) + pub(super) fn with_param(self) -> Self { self } + /// Generate a random UUIDv4 / UUIDv7 based on the `UUIDFunctionType` pub(super) fn format(self) -> Result { - // TODO: Generate a random UUIDv4 / UUIDv7 based on the `UUIDFunctionType` - todo!() + Ok(match self { + Self::Uuidv4 | Self::GenRandomUuid => Uuid::new_v4().to_string(), + Self::Uuidv7 => { + // I considered re-using `QueryTimestamps` (which stores statement and transaction times), however, + // what if we have multiple function calls within the same INSERT? That would mean generating the same + // UUIDs (as it's deterministic if the input time is the same), meaning we have to account for that + // by always generating a new time. + let current_time = Utc::now(); + Uuid::new_v7(Timestamp::from_unix( + ContextV7::new(), + current_time.timestamp_millis() as u64, + current_time.timestamp_subsec_nanos(), + )) + .to_string() + } + }) } }