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/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/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 87% rename from integration/rust/tests/integration/omni_timestamps.rs rename to integration/rust/tests/integration/omni_non_deterministic_funcs.rs index 77a12f32e..670d2f053 100644 --- a/integration/rust/tests/integration/omni_timestamps.rs +++ b/integration/rust/tests/integration/omni_non_deterministic_funcs.rs @@ -24,6 +24,96 @@ 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; + 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); + } + + // 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(); + } +} + /// 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/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/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/non_deterministic_funcs/mod.rs b/pgdog/src/frontend/router/parser/rewrite/statement/non_deterministic_funcs/mod.rs new file mode 100644 index 000000000..3df2d3400 --- /dev/null +++ b/pgdog/src/frontend/router/parser/rewrite/statement/non_deterministic_funcs/mod.rs @@ -0,0 +1,508 @@ +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}; + +use crate::{ + frontend::{ + RewritePlan, + client::QueryTimestamps, + router::parser::{ + StatementParser, StatementRewrite, Table, + rewrite::statement::{ + Error, + non_deterministic_funcs::{time::TimeFunctionType, uuid::UUIDFunctionType}, + plan::{GeneratedId, GeneratedParam}, + }, + }, + }, + net::parameter::ParameterValue, +}; + +/// 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 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 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 binary = formatted_string.as_bytes().to_vec(); + Ok((formatted_string, binary)) + } + + /// Some data types we get aren't compatible as-is with the pg_catalog type + /// needed to specify for Bind param types / Prepare ParamRef types. This maps them to + /// the correct pg_catalog types. + /// + /// Why do we need to cast? This is because, for example, if we try to use CURRENT_TIME (timetz) with a + /// text column in a binary format Bind param, Postgres will error without an explicit cast (expected UTF-8). + /// The alternative is manually calculating the binary format, which is a lot more of a headache! :) + /// + /// TODO: There might be some missing. Could be worth emitting an Error for any unexpected/untested types. + fn col_type_to_type_cast_alias(&self) -> &str { + match self.column_type.as_str() { + "time with time zone" => "timetz", + "timestamp with time zone" => "timestamptz", + "timestamp without time zone" => "timestamp", + "time without time zone" => "time", + other => other, + } + } +} + +/// 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 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) + }) + } + + /// 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::TimeFunction(tf) => tf.name(), + Self::UUIDFunction(uuid) => uuid.name(), + } + } + + /// 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::TimeFunction(tf) => Self::TimeFunction(tf.with_param(param)), + Self::UUIDFunction(uuid) => Self::UUIDFunction(uuid.with_param()), + } + } + + /// Convert `TimeFunction` in VALUES list + /// + /// Parse both `FuncCall`s and `SQLValueFunction`s here. + /// `now()` = `FuncCall`, + /// `CURRENT_TIMESTAMP`, `LOCALTIME` = `SQLValueFunction`, + fn from_node(node: Node, column_relation: Option<&Column>) -> Result, Error> { + match node { + Node::FuncCall(func) => { + let Some(Node::String(str)) = func.funcname().first() else { + return Ok(None); + }; + + let Some(func_name) = str.sval() else { + return Ok(None); + }; + + Self::from_func_call(func_name, Some(func.args())) + } + 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(_) => { + if let Some(relation) = column_relation { + return Self::from_func_call(relation.column_default.as_str(), None); + } + + Ok(None) + } + + _ => Ok(None), + } + } + + /// 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 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)); + }; + + 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())); + } + } + } + } + + Ok(None) + } +} + +impl StatementRewrite<'_> { + /// Rewrites timestamp functions like now() into either ParamRefs or correctly formatted Strings, + /// for the purpose of maintaining consistency across databases for omni tables. + pub(super) fn rewrite_timestamp_functions<'mem, 'mutref>( + &mut self, + mut stmt: NodeMut<'mem, 'mutref>, + mem: MemoryToken<'mem>, + // TODO: Replace `next_param` with plan.param directly + next_param: &mut i32, + plan: &mut RewritePlan, + ) -> Result<(), Error> { + let mut parser = StatementParser::new(stmt.as_ref(), None, self.schema, None); + let is_sharded = parser.is_sharded(self.db_schema, self.user, self.search_path); + + // not sharded = omni + if is_sharded { + return Ok(()); + } + + // + let Some((relation, cols, not_covered_cols)) = self.find_not_used_cols(&mut stmt, mem) + else { + return Ok(()); + }; + + let mut nd_rewrite = NDRewrite { + rewrite: self, + plan, + next_param, + mem, + relation, + cols, + }; + + // 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)?; + + // Replaces all time function calls (ParamRef or String) + nd_rewrite.transform_func_calls(stmt)?; + + Ok(()) + } + + /// Fetch the table Relation, so that we can get the relevant Schema for each column. + /// Fetch the list of columns that are DEFAULT (and not already covered) + fn find_not_used_cols<'mem, 'mutref>( + &self, + stmt: &mut NodeMut<'mem, 'mutref>, + mem: MemoryToken<'mem>, + ) -> Option<(Relation, Unique<'mem, &'mem NodeList>, Vec)> { + let Node::InsertStmt(insert_stmt) = stmt.as_ref() else { + return None; + }; + + let relation = insert_stmt.relation().expect("INSERT always has table"); + let table = Table::from(relation); + + let relation = self.db_schema.table(table, self.user, self.search_path)?; + let cols = insert_stmt.cols(); + + // Find the columns that the insert does NOT cover. + let not_covered_cols: Vec = if !cols.is_empty() { + let subset: Vec<&str> = cols + .iter() + .filter_map(|col| match col { + Node::ResTarget(target) => target.name(), + _ => None, + }) + .collect(); + + relation + .column_names() + .filter(|name| !subset.contains(name)) + .map(|name| name.to_string()) + .collect() + } else { + vec![] + }; + + Some((relation.clone(), mem.make_unique(cols), not_covered_cols)) + } +} + +struct NDRewrite<'mem, 'a, 's> { + rewrite: &'a mut StatementRewrite<'s>, + plan: &'a mut RewritePlan, + /// TODO: Replace `next_param` with plan.param directly + /// could do like a .next_param() method on `RewritePlan` + next_param: &'a mut i32, + mem: MemoryToken<'mem>, + relation: Relation, + cols: Unique<'mem, &'mem NodeList>, +} + +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, '_>) -> 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) { + 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)?; + + 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 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) -> Result>, Error> { + self.rewrite.rewritten = true; + + 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) => return Err(err), + }; + self.mem + .make_a_const(ConstValue::String(text.as_str())) + .uncast() + } else { + let param_ref = self.mem.make_param_ref(*self.next_param); + *self.next_param += 1; + + // 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::NDFunction(nd_function.clone()), + }); + + // Example: CAST($1::pg_catalog.text AS timetz) + // This is 30x less code at the expense of query verbosity; + // I talk about why in doc comment on `col_type_to_type_cast_alias` + self.mem + .make_type_cast( + self.mem + .make_type_cast( + param_ref.uncast(), + self.mem.make_list(&[ + self.mem.make_string(Some("pg_catalog")), + self.mem.make_string(Some("text")), + ]), + ) + .uncast(), + self.mem.make_list(&[self + .mem + .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..3685d1781 --- /dev/null +++ b/pgdog/src/frontend/router/parser/rewrite/statement/non_deterministic_funcs/uuid.rs @@ -0,0 +1,69 @@ +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, +}; + +/// Represents the kind of `UUIDFunction` that we're re-writing. +/// +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(super) enum UUIDFunctionType { + Uuidv4, + Uuidv7, //TODO: Postgres supports a parameter for an interval to be specified to shift the timestamp. + 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, 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 { + 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() + } + }) + } +} 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, diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/timestamp.rs b/pgdog/src/frontend/router/parser/rewrite/statement/timestamp.rs deleted file mode 100644 index 356af71d2..000000000 --- a/pgdog/src/frontend/router/parser/rewrite/statement/timestamp.rs +++ /dev/null @@ -1,651 +0,0 @@ -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, - make::{MemoryToken, Unique}, - raw::SQLValueFunctionOp, - transform::{TransformClosure, transform_node}, -}; -use pgdog_stats::{Column, Relation}; -use std::str::FromStr; - -use crate::{ - frontend::{ - RewritePlan, - client::QueryTimestamps, - router::parser::{ - StatementParser, StatementRewrite, Table, - rewrite::statement::{ - Error, - plan::{GeneratedId, GeneratedParam}, - }, - }, - }, - 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, - /// 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( - &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 = time_output.format(&reference_time, &tz, precision); - let binary = formatted_string.as_bytes().to_vec(); - - Ok((formatted_string, binary)) - } - - /// Some data types we get aren't compatible as-is with the pg_catalog type - /// needed to specify for Bind param types / Prepare ParamRef types. This maps them to - /// the correct pg_catalog types. - /// - /// Why do we need to cast? This is because, for example, if we try to use CURRENT_TIME (timetz) with a - /// text column in a binary format Bind param, Postgres will error without an explicit cast (expected UTF-8). - /// The alternative is manually calculating the binary format, which is a lot more of a headache! :) - /// - /// TODO: There might be some missing. Could be worth emitting an Error for any unexpected/untested types. - fn col_type_to_type_cast_alias(&self) -> &str { - match self.column_type.as_str() { - "time with time zone" => "timetz", - "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"), - ), - } - } -} - -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, -} - -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) - } - - /// 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, - } - } - - /// Convert `TimeFunction` in VALUES list - /// - /// Parse both `FuncCall`s and `SQLValueFunction`s here. - /// `now()` = `FuncCall`, - /// `CURRENT_TIMESTAMP`, `LOCALTIME` = `SQLValueFunction`, - fn from_node(node: Node, column_relation: Option<&Column>) -> Option { - match node { - Node::FuncCall(func) => { - let Node::String(str) = func.funcname().first()? else { - return None; - }; - str.sval()?.parse().ok() - } - 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); - } - - 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 { - 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(); - - for variant in Self::ALL_VARIANTS { - 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, - }; - - 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, - }); - } else if after.eq("()") || after.is_empty() { - return Ok(variant); - } else { - continue; - } - } - } - - Err(None) - } -} - -impl StatementRewrite<'_> { - /// Rewrites timestamp functions like now() into either ParamRefs or correctly formatted Strings, - /// for the purpose of maintaining consistency across databases for omni tables. - pub(super) fn rewrite_timestamp_functions<'mem, 'mutref>( - &mut self, - mut stmt: NodeMut<'mem, 'mutref>, - mem: MemoryToken<'mem>, - // TODO: Replace `next_param` with plan.param directly - next_param: &mut i32, - plan: &mut RewritePlan, - ) -> Result<(), Error> { - let mut parser = StatementParser::new(stmt.as_ref(), None, self.schema, None); - let is_sharded = parser.is_sharded(self.db_schema, self.user, self.search_path); - - // not sharded = omni - if is_sharded { - return Ok(()); - } - - // - let Some((relation, cols, not_covered_cols)) = self.find_not_used_cols(&mut stmt, mem) - else { - return Ok(()); - }; - - let mut timestamp_rewrite = TimestampRewrite { - rewrite: self, - plan, - next_param, - 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) - timestamp_rewrite.handle_adding_defaults(&mut stmt, ¬_covered_cols); - - // Replaces all time function calls (ParamRef or String) - timestamp_rewrite.transform_func_calls(stmt); - - timestamp_rewrite.error.map_or(Ok(()), Err) - } - - /// Fetch the table Relation, so that we can get the relevant Schema for each column. - /// Fetch the list of columns that are DEFAULT (and not already covered) - fn find_not_used_cols<'mem, 'mutref>( - &self, - stmt: &mut NodeMut<'mem, 'mutref>, - mem: MemoryToken<'mem>, - ) -> Option<(Relation, Unique<'mem, &'mem NodeList>, Vec)> { - let Node::InsertStmt(insert_stmt) = stmt.as_ref() else { - return None; - }; - - let relation = insert_stmt.relation().expect("INSERT always has table"); - let table = Table::from(relation); - - let relation = self.db_schema.table(table, self.user, self.search_path)?; - let cols = insert_stmt.cols(); - - // Find the columns that the insert does NOT cover. - let not_covered_cols: Vec = if !cols.is_empty() { - let subset: Vec<&str> = cols - .iter() - .filter_map(|col| match col { - Node::ResTarget(target) => target.name(), - _ => None, - }) - .collect(); - - relation - .column_names() - .filter(|name| !subset.contains(name)) - .map(|name| name.to_string()) - .collect() - } else { - vec![] - }; - - Some((relation.clone(), mem.make_unique(cols), not_covered_cols)) - } -} - -struct TimestampRewrite<'mem, 'a, 's> { - rewrite: &'a mut StatementRewrite<'s>, - plan: &'a mut RewritePlan, - /// TODO: Replace `next_param` with plan.param directly - /// could do like a .next_param() method on `RewritePlan` - next_param: &'a mut i32, - mem: MemoryToken<'mem>, - relation: Relation, - cols: Unique<'mem, &'mem NodeList>, - - // Error from formatting a time. - error: Option, -} - -impl<'mem, 'a, 's> TimestampRewrite<'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, '_>) { - 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, - } - }; - - if let Some(time_function_type) = - TimeFunctionType::from_node(value, col_relation) - { - let Some(col_relation) = col_relation else { - continue; - }; - - let time_function = TimeFunction { - time_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)); - changed = true; - } - } - - // 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), - }), - ); - } - - /// 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, - ) { - let NodeMut::InsertStmt(insert_stmt) = &mut stmt else { - return; - }; - - 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 { - continue; - }; - - let time_function = TimeFunction { - time_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; - }; - - // 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(&time_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>> { - 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) - { - 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() - } - }; - self.mem - .make_a_const(ConstValue::String(text.as_str())) - .uncast() - } else { - let param_ref = self.mem.make_param_ref(*self.next_param); - *self.next_param += 1; - - // 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()), - }); - - // Example: CAST($1::pg_catalog.text AS timetz) - // This is 30x less code at the expense of query verbosity; - // I talk about why in doc comment on `col_type_to_type_cast_alias` - self.mem - .make_type_cast( - self.mem - .make_type_cast( - param_ref.uncast(), - self.mem.make_list(&[ - self.mem.make_string(Some("pg_catalog")), - self.mem.make_string(Some("text")), - ]), - ) - .uncast(), - self.mem.make_list(&[self - .mem - .make_string(Some(time_function.col_type_to_type_cast_alias()))]), - ) - .uncast() - } - } -}