diff --git a/.schema/pgdog.schema.json b/.schema/pgdog.schema.json index a8609bdc1..f67f0d134 100644 --- a/.schema/pgdog.schema.json +++ b/.schema/pgdog.schema.json @@ -224,6 +224,7 @@ "$ref": "#/$defs/Rewrite", "default": { "enabled": false, + "omni_non_deterministic_functions": "ignore", "primary_key": "ignore", "shard_key": "error", "split_inserts": "error" @@ -1939,6 +1940,11 @@ "type": "boolean", "default": false }, + "omni_non_deterministic_functions": { + "description": "Behavior when an `INSERT` is headed to an omnisharded table using a function (such as date-time functions)\nthat will not be consistent when performing the functions separately on each shard.\nThus, it re-writes all such functions before performing the `INSERT` with constant values to maintain consistency.\n\nExample: `NOW()` is re-written to `2026-09-15 18:14:09.123456-05` (or whatever the current time is)\nbefore performing the individual `INSERT` operations.\n\nThis applies to both `DEFAULT` table schema and functions called within a VALUES list of an `INSERT`.\n\n`ignore` allows the `INSERT` without modification.\n\n_Default:_ `ignore`\n\n", + "$ref": "#/$defs/RewriteMode", + "default": "ignore" + }, "primary_key": { "description": "Behavior for `INSERT` missing a `BIGINT` primary key: `error` rejects, `rewrite` auto-injects `pgdog.unique_id()`, `ignore` allows without modification.\n\n_Default:_ `ignore`\n\n", "$ref": "#/$defs/RewriteMode", diff --git a/Cargo.lock b/Cargo.lock index 6ce4fcc9f..b1e957662 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1018,6 +1018,16 @@ dependencies = [ "windows-link", ] +[[package]] +name = "chrono-tz" +version = "0.10.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a6139a8597ed92cf816dfb33f5dd6cf0bb93a6adc938f11039f371bc5bcd26c3" +dependencies = [ + "chrono", + "phf 0.12.1", +] + [[package]] name = "clang-sys" version = "1.8.1" @@ -2437,6 +2447,7 @@ version = "0.1.0" dependencies = [ "bytes", "chrono", + "chrono-tz", "futures-util", "libc", "native-tls", @@ -3115,6 +3126,7 @@ dependencies = [ "bytes", "cc", "chrono", + "chrono-tz", "clap", "crc32c", "csv-core", @@ -3301,16 +3313,34 @@ dependencies = [ "serde_json", ] +[[package]] +name = "phf" +version = "0.12.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "913273894cec178f401a31ec4b656318d95473527be05c0752cc41cdc32be8b7" +dependencies = [ + "phf_shared 0.12.1", +] + [[package]] name = "phf" version = "0.13.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c1562dc717473dbaa4c1f85a36410e03c047b2e7df7f45ee938fbef64ae7fadf" dependencies = [ - "phf_shared", + "phf_shared 0.13.1", "serde", ] +[[package]] +name = "phf_shared" +version = "0.12.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "06005508882fb681fd97892ecff4b7fd0fee13ef1aa569f8695dae7ab9099981" +dependencies = [ + "siphasher", +] + [[package]] name = "phf_shared" version = "0.13.1" @@ -4946,7 +4976,7 @@ dependencies = [ "log", "parking_lot", "percent-encoding", - "phf", + "phf 0.13.1", "pin-project-lite", "postgres-protocol", "postgres-types", diff --git a/integration/pgdog.toml b/integration/pgdog.toml index 41fc3e8e0..d0574aed0 100644 --- a/integration/pgdog.toml +++ b/integration/pgdog.toml @@ -44,6 +44,7 @@ enabled = false shard_key = "ignore" split_inserts = "error" # primary_key = "rewrite" +omni_non_deterministic_functions = "rewrite" # ------------------------------------------------------------------------------ # ----- Database :: pgdog ------------------------------------------------------ diff --git a/integration/python/test_asyncpg.py b/integration/python/test_asyncpg.py index 46ee80896..428016ed0 100644 --- a/integration/python/test_asyncpg.py +++ b/integration/python/test_asyncpg.py @@ -689,3 +689,45 @@ async def test_pgdog_role_selection(): pass assert got_err + + +# Test to make sure everything works with unnamed prepared statements when we need to re-write +# Bind for `TimeFunction` +@pytest.mark.asyncio +async def test_omni_time_function_unnamed_statement(): + """`statement_cache_size=0` makes asyncpg use unnamed prepared statements. + """ + conn = await asyncpg.connect( + user="pgdog", + password="pgdog", + database="pgdog_sharded", + host="127.0.0.1", + port=6432, + statement_cache_size=0, + ) + row_id = random.randrange(1 << 32, 1 << 63) + + try: + result = await conn.execute( + "INSERT INTO sharded_omni (id, value, created_at) VALUES ($1, $2, now())", + row_id, + "unnamed", + ) + assert result == "INSERT 0 1" + + # Every shard must have gotten the same timestamp. + created_at = [ + await conn.fetchval( + f"/* pgdog_shard: {shard} */ SELECT created_at" + " FROM sharded_omni WHERE id = $1", + row_id, + ) + for shard in (0, 1) + ] + assert created_at[0] is not None + assert created_at[0] == created_at[1] + finally: + await conn.execute("DELETE FROM sharded_omni WHERE id = $1", row_id) + await conn.close() + + no_out_of_sync() diff --git a/integration/rust/Cargo.toml b/integration/rust/Cargo.toml index 386134ff7..ff2ee4347 100644 --- a/integration/rust/Cargo.toml +++ b/integration/rust/Cargo.toml @@ -25,4 +25,5 @@ libc = "0.2" rand = "0.9" bytes.workspace = true rust_decimal = { version = "1.42.0", features = ["macros"] } +chrono-tz = "0.10.4" pgdog-stats = { path = "../../pgdog-stats" } diff --git a/integration/rust/tests/integration/mod.rs b/integration/rust/tests/integration/mod.rs index 082df0d21..0d7de7da6 100644 --- a/integration/rust/tests/integration/mod.rs +++ b/integration/rust/tests/integration/mod.rs @@ -23,6 +23,7 @@ pub mod max; pub mod multi_set; pub mod notify; pub mod offset; +pub mod omni_timestamps; 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_timestamps.rs new file mode 100644 index 000000000..77a12f32e --- /dev/null +++ b/integration/rust/tests/integration/omni_timestamps.rs @@ -0,0 +1,706 @@ +use std::ops::Sub; + +use crate::setup::admin_sqlx; +use crate::setup::connection_sqlx_direct_db; +use crate::setup::connections_sqlx; +use chrono::DateTime; +use chrono::Duration; +use chrono::FixedOffset; +use chrono::NaiveDate; +use chrono::NaiveDateTime; +use chrono::NaiveTime; +use chrono::Utc; +use chrono_tz::Tz; +use sqlx::PgTransaction; +use sqlx::Pool; +use sqlx::Postgres; +use sqlx::Transaction; +use sqlx::postgres::PgRow; +use sqlx::postgres::types::PgTimeTz; +use sqlx::{Executor, Row}; + +// TODO: Test changing schema for a column while this is cached +// TODO: Test for other caching issues +// 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 + +/// 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. +/// +/// Tests against NYC timezone. +#[tokio::test] +async fn omni_timestamp_rewrite_local_time() { + let schema = "test_time_text text, test_time_regular time DEFAULT LOCALTIME"; + let insertion_col = "test_time_text"; + let func_call = "LOCALTIME(3)"; + + reusable_func_test(schema, insertion_col, func_call, async |pg_row, dog_row| { + assert_text_format(pg_row, dog_row, insertion_col, |s| { + NaiveTime::parse_from_str(s, "%H:%M:%S%.f") + }) + .await; + assert_equality( + pg_row, + dog_row, + "test_time_regular", + Duration::seconds(5), + |t: NaiveTime| t, + ) + .await; + }) + .await; +} + +/// CURRENT_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 & timetz col. +#[tokio::test] +async fn omni_timestamp_rewrite_current_time() { + let schema = "test_time_text text, test_time_regular timetz DEFAULT CURRENT_TIME"; + let insertion_col = "test_time_text"; + let func_call = "CURRENT_TIME(3)"; + + reusable_func_test(schema, insertion_col, func_call, async |pg_row, dog_row| { + assert_text_format(pg_row, dog_row, insertion_col, |s| { + NaiveTime::parse_from_str(s, "%H:%M:%S%.f%#z") + }) + .await; + assert_equality( + pg_row, + dog_row, + "test_time_regular", + Duration::seconds(5), + |t: PgTimeTz| t.time - t.offset, + ) + .await; + }) + .await; +} + +/// CURRENT_DATE testing +/// - Case 1: `test_date_text` has no DEFAULT and it uses text col. +/// - Case 2: `test_date_regular` has DEFAULT and it uses date col. +#[tokio::test] +async fn omni_timestamp_rewrite_current_date() { + let schema = "test_date_text text, test_date_regular date DEFAULT CURRENT_DATE"; + let insertion_col = "test_date_text"; + let func_call = "CURRENT_DATE"; + + reusable_func_test(schema, insertion_col, func_call, async |pg_row, dog_row| { + assert_text_format(pg_row, dog_row, insertion_col, |s| { + NaiveDate::parse_from_str(s, "%Y-%m-%d") + }) + .await; + + assert_equality( + pg_row, + dog_row, + "test_date_regular", + Duration::hours(25), + |d: NaiveDate| d, + ) + .await; + }) + .await; +} + +/// LOCALTIMESTAMP testing +/// - Case 1: `test_timestamp_text` has no DEFAULT w/ precision arg & text col. +/// - Case 2: `test_timestamp_regular` has DEFAULT w/ no precision arg & timestamp col. +#[tokio::test] +async fn omni_timestamp_rewrite_local_timestamp() { + let schema = + "test_timestamp_text text, test_timestamp_regular timestamp DEFAULT LOCALTIMESTAMP"; + let insertion_col = "test_timestamp_text"; + let func_call = "LOCALTIMESTAMP(2)"; + + reusable_func_test(schema, insertion_col, func_call, async |pg_row, dog_row| { + assert_text_format(pg_row, dog_row, insertion_col, |s| { + NaiveDateTime::parse_from_str(s, "%Y-%m-%d %H:%M:%S%.f") + }) + .await; + assert_equality( + pg_row, + dog_row, + "test_timestamp_regular", + Duration::seconds(5), + |t: NaiveDateTime| t, + ) + .await; + }) + .await; +} + +/// clock_timestamp() testing +/// - Case 1: `test_clock_text` has no DEFAULT & text col. +/// - Case 2: `test_clock_regular` has DEFAULT & timestamptz col. +#[tokio::test] +async fn omni_timestamp_rewrite_clock_timestamp() { + let schema = "test_clock_text text, test_clock_regular timestamptz DEFAULT clock_timestamp()"; + let insertion_col = "test_clock_text"; + let func_call = "clock_timestamp()"; + + reusable_func_test(schema, insertion_col, func_call, async |pg_row, dog_row| { + assert_text_format(pg_row, dog_row, insertion_col, |s| { + DateTime::parse_from_str(s, "%Y-%m-%d %H:%M:%S%.f%#z") + }) + .await; + assert_equality( + pg_row, + dog_row, + "test_clock_regular", + Duration::seconds(5), + |t: DateTime| t, + ) + .await; + }) + .await; +} + +/// statement_timestamp() testing +/// - Case 1: `test_statement_text` has no DEFAULT & text col. +/// - Case 2: `test_statement_regular` has DEFAULT & timestamptz col. +#[tokio::test] +async fn omni_timestamp_rewrite_statement_timestamp() { + let schema = "test_statement_text text, test_statement_regular timestamptz DEFAULT statement_timestamp()"; + let insertion_col = "test_statement_text"; + let func_call = "statement_timestamp()"; + + reusable_func_test(schema, insertion_col, func_call, async |pg_row, dog_row| { + assert_text_format(pg_row, dog_row, insertion_col, |s| { + DateTime::parse_from_str(s, "%Y-%m-%d %H:%M:%S%.f%#z") + }) + .await; + assert_equality( + pg_row, + dog_row, + "test_statement_regular", + Duration::seconds(5), + |t: DateTime| t, + ) + .await; + }) + .await; +} + +/// statement_timestamp() = same for every value in one statement (explicit and DEFAULT, across rows) +/// However it changes between statements in the same transaction. +#[tokio::test] +async fn omni_timestamp_rewrite_statement_timestamp_consistency() { + let conn = connections_sqlx().await; + let conn = conn.get(1).unwrap(); + conn.execute("DROP TABLE IF EXISTS dummy_omni_table") + .await + .unwrap(); + conn.execute( + "CREATE TABLE dummy_omni_table(id BIGSERIAL PRIMARY KEY, explicit timestamptz, by_default timestamptz DEFAULT statement_timestamp())", + ) + .await + .unwrap(); + admin_sqlx().await.execute("RELOAD").await.unwrap(); + + let mut transaction = conn.begin().await.unwrap(); + + let first_rows = sqlx::query( + "INSERT INTO dummy_omni_table(id, explicit) VALUES ($1, statement_timestamp()), ($2, statement_timestamp()) RETURNING *", + ) + .bind(1) + .bind(2) + .fetch_all(&mut *transaction) + .await + .unwrap(); + + let second_row = sqlx::query( + "INSERT INTO dummy_omni_table(id, explicit) VALUES ($1, statement_timestamp()) RETURNING *", + ) + .bind(3) + .fetch_one(&mut *transaction) + .await + .unwrap(); + + transaction.rollback().await.unwrap(); + conn.execute("DROP TABLE dummy_omni_table").await.unwrap(); + + let times = |row: &PgRow| { + ( + row.get::, _>("explicit"), + row.get::, _>("by_default"), + ) + }; + + let (row_1_explicit, row_1_default) = times(&first_rows[0]); + let (row_2_explicit, row_2_default) = times(&first_rows[1]); + let (row_3_explicit, _) = times(&second_row); + + assert_eq!(row_1_explicit, row_2_explicit); + assert_eq!(row_1_explicit, row_1_default); + assert_eq!(row_1_default, row_2_default); + assert_ne!(row_1_explicit, row_3_explicit); +} + +/// timeofday() testing +/// - Case 1: `test_timeofday_text` has no DEFAULT & text col. +#[tokio::test] +async fn omni_timestamp_rewrite_time_of_day() { + const TIME_OF_DAY_FORMAT: &str = "%a %b %d %H:%M:%S%.f %Y %Z"; + + let schema = "test_timeofday_text text"; + let insertion_col = "test_timeofday_text"; + let func_call = "timeofday()"; + + reusable_func_test(schema, insertion_col, func_call, async |pg_row, dog_row| { + assert_text_format(pg_row, dog_row, insertion_col, |s| { + NaiveDateTime::parse_from_str(s, TIME_OF_DAY_FORMAT) + }) + .await; + assert_equality( + pg_row, + dog_row, + insertion_col, + Duration::seconds(5), + |s: String| NaiveDateTime::parse_from_str(&s, TIME_OF_DAY_FORMAT).unwrap(), + ) + .await; + }) + .await; +} + +/// Asserts, after normalization, that the value for `col_name` for `pg_row` and `dog_row` are within +/// the bound of `acceptable_diff` +async fn assert_equality( + pg_row: &PgRow, + dog_row: &PgRow, + col_name: &str, + acceptable_diff: Duration, + normalize: impl Fn(T) -> U, +) where + T: for<'r> sqlx::Decode<'r, sqlx::Postgres> + sqlx::Type, + U: Sub, +{ + let (pg_value, dog_value) = ( + normalize(pg_row.get::(col_name)), + normalize(dog_row.get::(col_name)), + ); + + assert!((pg_value - dog_value).abs() < acceptable_diff); +} + +/// Asserts that the parsed value for `col_name` from `pg_row` and `dog_row` both work; +/// proving that both (Postgres and PgDog) work the same, and that both are correctly formatted. +async fn assert_text_format( + pg_row: &PgRow, + dog_row: &PgRow, + col_name: &str, + parse: impl Fn(&str) -> chrono::ParseResult, +) { + let (pg_text, dog_text) = ( + pg_row.get::<&str, &str>(col_name), + dog_row.get::<&str, &str>(col_name), + ); + + assert!(parse(pg_text).is_ok()); + assert!(parse(dog_text).is_ok()); +} + +/// Inserts 3 rows +/// - Row 1: Simple +/// - Row 2: Extended +/// - Row 3: Prepare / Execute +/// +/// Uses RETURNING * on each, and returns all the PgRows for analysis. +async fn test_simple_extended_and_prepare( + conn: &Pool, + cols: &str, + vals: &str, +) -> Vec { + // TODO: Could be useful to test BOTH UTC and NYC. + let mut transaction = conn.begin().await.unwrap(); + + // Allows us to test things like local time (instead of everything being UTC) + transaction + .execute("SET TIME ZONE 'America/New_York'") + .await + .unwrap(); + + let row1 = sqlx::raw_sql( + format!("INSERT INTO dummy_omni_table(id, {cols}) VALUES (1, {vals}) RETURNING *").as_str(), + ) + .fetch_one(&mut *transaction) + .await + .unwrap(); + + // By putting the Bind parameter first, it forces us to use the Binary text format, testing vs the already tested String. + let row2 = sqlx::query( + format!("INSERT INTO dummy_omni_table(id, {cols}) VALUES ($1, {vals}) RETURNING *") + .as_str(), + ) + .bind(2) + .fetch_one(&mut *transaction) + .await + .unwrap(); + + sqlx::raw_sql( + format!( + "PREPARE stmt AS INSERT INTO dummy_omni_table(id, {cols}) VALUES ($1, {vals}) RETURNING *" + ) + .as_str(), + ) + .execute(&mut *transaction) + .await + .unwrap(); + + let row3 = sqlx::raw_sql("EXECUTE stmt(3)") + .fetch_one(&mut *transaction) + .await + .unwrap(); + + transaction.rollback().await.unwrap(); + + vec![row1, row2, row3] +} + +async fn reusable_func_test( + schema: &str, + insertion_col: &str, + func_call: &str, + validate: impl AsyncFn(&PgRow, &PgRow), +) { + let conn = connections_sqlx().await; + let conn = conn.get(1).unwrap(); + conn.execute("DROP TABLE IF EXISTS dummy_omni_table") + .await + .unwrap(); + conn.execute( + format!("CREATE TABLE IF NOT EXISTS dummy_omni_table(id BIGSERIAL PRIMARY KEY, {schema})") + .as_str(), + ) + .await + .unwrap(); + admin_sqlx().await.execute("RELOAD").await.unwrap(); + + let pg_rows = test_simple_extended_and_prepare( + &connection_sqlx_direct_db("shard_0").await, + insertion_col, + func_call, + ) + .await; + + let dog_rows = test_simple_extended_and_prepare( + connections_sqlx().await.get(1).unwrap(), + insertion_col, + func_call, + ) + .await; + + for (pg_row, dog_row) in pg_rows.iter().zip(&dog_rows) { + validate(pg_row, dog_row).await; + } + + conn.execute("DROP TABLE dummy_omni_table").await.unwrap(); +} + +/// NOTE: The tests below assert that everything is intercepted and handled; therefore, I didn't try to mimic that in the above tests, +/// given that they share a re-usable abstraction (would be redundant) +/// +/// Re-usable harness for other tests (simple protocol, extended protocol, prepare/execute) to equally test if +/// different INSERT methods work correctly. +/// +/// - Before running, it resets everything (re-create table on each shard) +/// - Tests with different timezones to ensure `timestamp` vs `timestamptz` works as intended; +/// more specifically, INSERT timezone is different from SELECT timezone. +/// - Tests DEFAULT schema (both implicit and explicit) +/// - Tests functions within the VALUES list. +/// - Tests multiple VALUES lists (...), (....) +/// - Ensures time consistency across databases for the omnisharded column. +/// - Ensures consistency across columns within the same row (if multiple time funcs) +/// +/// Does all of this within transactions, and after conclusion, performs a rollback. +async fn run_test(perform_insert: F) +where + F: AsyncFn(&mut PgTransaction), +{ + let conn = connections_sqlx().await; + let conn = conn.get(1).unwrap(); + conn.execute("DROP TABLE IF EXISTS public.test_omni_ts") + .await + .unwrap(); + conn.execute("CREATE TABLE IF NOT EXISTS public.test_omni_ts(id BIGSERIAL PRIMARY KEY, created_at TIMESTAMP, created_at_tz TIMESTAMPTZ, created_at_default TIMESTAMP DEFAULT CURRENT_TIMESTAMP, created_at_tz_default TIMESTAMPTZ DEFAULT TRANSACTION_TIMESTAMP())").await.unwrap(); + admin_sqlx().await.execute("RELOAD").await.unwrap(); + + for (insertion_tz, fetch_tz) in [ + ("America/Los_Angeles", "America/New_York"), + ("America/New_York", "America/Los_Angeles"), + // TODO: could also check equal + ] { + let mut sesh = conn.begin().await.unwrap(); + + sesh.execute(format!("SET TIME ZONE '{insertion_tz}'").as_str()) + .await + .unwrap(); + + perform_insert(&mut sesh).await; + check_shards(&mut sesh, fetch_tz, insertion_tz).await; + + sesh.rollback().await.unwrap(); + } + + conn.execute("DROP TABLE public.test_omni_ts") + .await + .unwrap(); +} + +/// Simple protocol case (see `run_test` for details) +#[tokio::test] +async fn omni_timestamp_rewrite_simple_protocol() { + run_test(async |sesh| { + // Routes to both (shard 0, shard 1) because it's omnisharded. + sqlx::raw_sql( + "INSERT INTO test_omni_ts(id, created_at, created_at_tz) VALUES (1, now(), now()), (2, now(), now())", + ) + .execute(&mut **sesh) + .await + .unwrap(); + + sqlx::raw_sql( + "INSERT INTO test_omni_ts(id, created_at, created_at_tz) VALUES (3, now(), now())", + ) + .execute(&mut **sesh) + .await + .unwrap(); + + sqlx::raw_sql( + "INSERT INTO test_omni_ts(id, created_at, created_at_tz, created_at_default, created_at_tz_default) VALUES (4, now(), now(), DEFAULT, DEFAULT)", + ) + .execute(&mut **sesh) + .await + .unwrap(); + }).await; +} + +/// Extended protocol case (see `run_test` for details) +#[tokio::test] +async fn omni_timestamp_rewrite_extended_protocol() { + run_test(async |sesh| { + // Routes to both (shard 0, shard 1) because it's omnisharded. + sqlx::query( + "INSERT INTO test_omni_ts(id, created_at, created_at_tz) VALUES ($1, NOW(), now()), ($2, TRANSACTION_TIMESTAMP(), CURRENT_TIMESTAMP)", + ) + .bind(1).bind(2) + .execute(&mut **sesh) + .await + .unwrap(); + + sqlx::query( + "INSERT INTO test_omni_ts(id, created_at, created_at_tz) VALUES ($1, now(), now())", + ) + .bind(3) + .execute(&mut **sesh) + .await + .unwrap(); + + sqlx::query( + "INSERT INTO test_omni_ts(id, created_at, created_at_tz, created_at_default, created_at_tz_default) VALUES ($1, now(), now(), DEFAULT, DEFAULT)", + ) + .bind(4) + .execute(&mut **sesh) + .await + .unwrap(); + }).await; +} + +/// Prepare/execute case (see `run_test` for details) +#[tokio::test] +async fn omni_timestamp_rewrite_prepare_execute() { + run_test(async |sesh| { + // Routes to both (shard 0, shard 1) because it's omnisharded. + sqlx::raw_sql( + "PREPARE stmt AS INSERT INTO test_omni_ts(id, created_at, created_at_tz) VALUES ($1, now(), now()), ($2, now(), now())", + ) + .execute(&mut **sesh) + .await + .unwrap(); + + sqlx::raw_sql( + "PREPARE stmt2 AS INSERT INTO test_omni_ts(id, created_at, created_at_tz) VALUES ($1, now(), now())" + ).execute(&mut **sesh).await.unwrap(); + + sqlx::raw_sql( + "PREPARE stmt3 AS INSERT INTO test_omni_ts(id, created_at, created_at_tz, created_at_default, created_at_tz_default) VALUES ($1, now(), now(), DEFAULT, DEFAULT)" + ).execute(&mut **sesh).await.unwrap(); + + sqlx::raw_sql("EXECUTE stmt(1, 2)") + .execute(&mut **sesh) + .await + .unwrap(); + + sqlx::raw_sql("EXECUTE stmt2(3)") + .execute(&mut **sesh) + .await + .unwrap(); + + + sqlx::raw_sql("EXECUTE stmt3(4)") + .execute(&mut **sesh) + .await + .unwrap(); + }).await; +} + +/// TODO: Doc comment. +async fn check_shards(sesh: &mut Transaction<'_, Postgres>, fetch_tz: &str, insertion_tz: &str) { + let now = Utc::now(); + let (shard_0_rows, shard_1_rows) = fetch_rows_with_tz(sesh, fetch_tz).await; + + for (shard_0_row, shard_1_row) in shard_0_rows.iter().zip(&shard_1_rows) { + assert_timestamp_col_validity( + shard_0_row, + shard_1_row, + ColumnType::ByDefault, + &now, + insertion_tz, + ) + .await; + assert_timestamp_tz_col_validity(shard_0_row, shard_1_row, ColumnType::ByDefault, &now) + .await; + + assert_timestamp_col_validity( + shard_0_row, + shard_1_row, + ColumnType::Regular, + &now, + insertion_tz, + ) + .await; + assert_timestamp_tz_col_validity(shard_0_row, shard_1_row, ColumnType::Regular, &now).await; + + // Same shard time equality in same column? + { + let created_at = shard_0_row.get::, &str>("created_at_tz"); + let created_at_default = + shard_0_row.get::, &str>("created_at_tz_default"); + + assert_eq!(created_at, created_at_default); + } + } + + { + // Same INSERT: VALUES (...), (...) + // now() should be the same. + let shard_0_row_1 = shard_0_rows.first().unwrap(); + let shard_0_row_2 = shard_0_rows.get(1).unwrap(); + + let created_at_first_insert = shard_0_row_1.get::, &str>("created_at_tz"); + let created_at_second_insert = shard_0_row_2.get::, &str>("created_at_tz"); + + assert_eq!(created_at_first_insert, created_at_second_insert); + } + + { + // Added in same transaction. Separate INSERTs. + // now() should be the same. + let shard_0_row_1 = shard_0_rows.first().unwrap(); + let shard_0_row_3 = shard_0_rows.get(2).unwrap(); + + let created_at_first_insert = shard_0_row_1.get::, &str>("created_at_tz"); + let created_at_second_insert = shard_0_row_3.get::, &str>("created_at_tz"); + + assert_eq!(created_at_first_insert, created_at_second_insert); + } +} + +enum ColumnType { + /// Col has DEFAULT in table schema. + ByDefault, + /// Function called and specified explicitly in VALUES list for the INSERT + Regular, +} + +/// These will be the same (based on UTC) regardless of being inserted/fetched in different timezones. +async fn assert_timestamp_tz_col_validity( + shard_0_row: &PgRow, + shard_1_row: &PgRow, + col_type: ColumnType, + now: &DateTime, +) { + let created_at_tz_col = match col_type { + ColumnType::ByDefault => "created_at_tz_default", + ColumnType::Regular => "created_at_tz", + }; + + let first = shard_0_row.get::, &str>(created_at_tz_col); + let second = shard_1_row.get::, &str>(created_at_tz_col); + + // Cross-shard equality? + assert_eq!(first, second); + + // Equal to UTC? + assert!(*now > first); + assert!((*now - first) < Duration::seconds(5)); +} + +/// `fetch_time_zone` will differ from `insert_time_zone` with timestamp column (created_at) +/// This is because offset information is stripped when inserted in Postgres. +async fn assert_timestamp_col_validity( + shard_0_row: &PgRow, + shard_1_row: &PgRow, + col_type: ColumnType, + now: &DateTime, + insertion_tz: &str, +) { + let created_at_col = match col_type { + ColumnType::ByDefault => "created_at_default", + ColumnType::Regular => "created_at", + }; + + let first = shard_0_row.get::(created_at_col); + let second = shard_1_row.get::(created_at_col); + + // Cross-shard equality? + assert_eq!(first, second); + + // Local time in LA and NY. + let (los_angeles_time, new_york_time) = ( + now.with_timezone(&Tz::America__Los_Angeles).naive_local(), + now.with_timezone(&Tz::America__New_York).naive_local(), + ); + + let (la_time_diff, ny_time_diff) = ( + (los_angeles_time - first).abs(), + (new_york_time - first).abs(), + ); + + if insertion_tz.eq("America/Los_Angeles") { + // insert LA tz, fetch NY tz + assert!(ny_time_diff > Duration::hours(2) && ny_time_diff < Duration::hours(4)); + assert!(la_time_diff < Duration::seconds(5)); + } else { + // insert NY tz, fetch LA tz + assert!(la_time_diff > Duration::hours(2) && la_time_diff < Duration::hours(4)); + assert!(ny_time_diff < Duration::seconds(5)); + } +} + +/// TODO: Docs +async fn fetch_rows_with_tz( + sesh: &mut Transaction<'_, Postgres>, + fetch_tz: &str, +) -> (Vec, Vec) { + sesh.execute(format!("SET TIME ZONE '{fetch_tz}'").as_str()) + .await + .unwrap(); + + // Force to route to the individual shards to ensure no divergence. + ( + sesh.fetch_all( + "/* pgdog_shard: 0 */ SELECT * FROM public.test_omni_ts WHERE id IN (1, 2, 3, 4)", + ) + .await + .unwrap(), + sesh.fetch_all( + "/* pgdog_shard: 1 */ SELECT * FROM public.test_omni_ts WHERE id IN (1, 2, 3, 4)", + ) + .await + .unwrap(), + ) +} diff --git a/pgdog-config/src/rewrite.rs b/pgdog-config/src/rewrite.rs index 0cea2afd2..0988d252f 100644 --- a/pgdog-config/src/rewrite.rs +++ b/pgdog-config/src/rewrite.rs @@ -92,6 +92,23 @@ pub struct Rewrite { /// #[serde(default = "Rewrite::default_primary_key")] pub primary_key: RewriteMode, + + /// Behavior when an `INSERT` is headed to an omnisharded table using a function (such as date-time functions) + /// that will not be consistent when performing the functions separately on each shard. + /// Thus, it re-writes all such functions before performing the `INSERT` with constant values to maintain consistency. + /// + /// Example: `NOW()` is re-written to `2026-09-15 18:14:09.123456-05` (or whatever the current time is) + /// before performing the individual `INSERT` operations. + /// + /// This applies to both `DEFAULT` table schema and functions called within a VALUES list of an `INSERT`. + /// + /// `ignore` allows the `INSERT` without modification. + /// + /// _Default:_ `ignore` + /// + /// + #[serde(default = "Rewrite::default_omni_non_deterministic_functions")] + pub omni_non_deterministic_functions: RewriteMode, } impl Default for Rewrite { @@ -101,6 +118,7 @@ impl Default for Rewrite { shard_key: Self::default_shard_key(), split_inserts: Self::default_split_inserts(), primary_key: Self::default_primary_key(), + omni_non_deterministic_functions: Self::default_omni_non_deterministic_functions(), } } } @@ -117,4 +135,8 @@ impl Rewrite { const fn default_primary_key() -> RewriteMode { RewriteMode::Ignore } + + const fn default_omni_non_deterministic_functions() -> RewriteMode { + RewriteMode::Ignore + } } diff --git a/pgdog/Cargo.toml b/pgdog/Cargo.toml index 9e53e5de0..572b57073 100644 --- a/pgdog/Cargo.toml +++ b/pgdog/Cargo.toml @@ -54,6 +54,7 @@ url = "2" rmp-serde = "1" rust_decimal = { version = "1.36", features = ["db-postgres", "macros", "maths"] } chrono = "0.4" +chrono-tz = "0.10.4" hyper = { version = "1", features = ["full"] } http-body-util = "0.1" hyper-util = { version = "0.1", features = ["full"] } diff --git a/pgdog/src/admin/set.rs b/pgdog/src/admin/set.rs index 64b148315..ad92d7cc5 100644 --- a/pgdog/src/admin/set.rs +++ b/pgdog/src/admin/set.rs @@ -135,6 +135,13 @@ impl Command for Set { .map_err(|_| Error::Syntax)?; } + "rewrite_omni_non_deterministic_functions" => { + config.config.rewrite.omni_non_deterministic_functions = self + .value + .parse::() + .map_err(|_| Error::Syntax)?; + } + "rewrite_enabled" => { config.config.rewrite.enabled = Self::from_json(&self.value)?; } diff --git a/pgdog/src/backend/pool/cluster.rs b/pgdog/src/backend/pool/cluster.rs index d7fcc9b8d..bb512efdc 100644 --- a/pgdog/src/backend/pool/cluster.rs +++ b/pgdog/src/backend/pool/cluster.rs @@ -19,7 +19,10 @@ use crate::{ ConnectionRecovery, MultiTenant, PoolerMode, ReadWriteSplit, ReadWriteStrategy, User, }, frontend::{ClientRequest, RegexParser, router::round_robin}, - net::{bind::Parameter as BindParameter, messages::DataRow, messages::FrontendPid}, + net::{ + bind::Parameter as BindParameter, messages::DataRow, messages::FrontendPid, + parameter::ParameterValue, + }, }; use super::{ @@ -443,6 +446,7 @@ impl Cluster { cluster.rewrite.enabled = false; cluster.rewrite.shard_key = RewriteMode::Ignore; cluster.rewrite.split_inserts = RewriteMode::Ignore; + cluster.rewrite.omni_non_deterministic_functions = RewriteMode::Ignore; cluster } @@ -489,6 +493,14 @@ impl Cluster { &self.shards } + /// The database's default `TimeZone` + pub(crate) fn default_timezone(&self) -> Option<&ParameterValue> { + self.shards + .iter() + .flat_map(|shard| shard.pool_iter()) + .find_map(|pool| pool.cached_params()?.get("TimeZone")) + } + pub(crate) fn passwords(&self) -> &[PasswordKind] { &self.passwords } diff --git a/pgdog/src/backend/pool/connection/mirror/mod.rs b/pgdog/src/backend/pool/connection/mirror/mod.rs index bc5292e0e..5f71475d6 100644 --- a/pgdog/src/backend/pool/connection/mirror/mod.rs +++ b/pgdog/src/backend/pool/connection/mirror/mod.rs @@ -11,9 +11,9 @@ use tracing::{debug, error, warn}; use crate::backend::Cluster; use crate::config::{ConfigAndUsers, config}; -use crate::frontend::client::TransactionType; use crate::frontend::client::query_engine::{QueryEngine, QueryEngineContext}; use crate::frontend::client::timeouts::Timeouts; +use crate::frontend::client::transaction_type::Transaction; use crate::frontend::{ClientComms, PreparedStatements}; use crate::net::{FrontendPid, Parameter, Parameters, Stream}; use crate::tasks; @@ -47,7 +47,7 @@ pub(crate) struct Mirror { /// Stream that absorbs all data. pub(crate) stream: Stream, /// Transaction state. - pub(crate) transaction: Option, + pub(crate) transaction: Option, } impl Mirror { diff --git a/pgdog/src/backend/pool/monitor.rs b/pgdog/src/backend/pool/monitor.rs index 0bd7d89f8..8b1a13607 100644 --- a/pgdog/src/backend/pool/monitor.rs +++ b/pgdog/src/backend/pool/monitor.rs @@ -464,6 +464,7 @@ impl Monitor { conn.set_credentials_generation(guard.credentials_generation()); } conn.apply_lifetime_jitter(max_age, max_age_jitter); + pool.cache_params(conn.params()); return Ok(conn); } diff --git a/pgdog/src/backend/pool/pool_impl.rs b/pgdog/src/backend/pool/pool_impl.rs index a3a800f4c..6891c33f6 100644 --- a/pgdog/src/backend/pool/pool_impl.rs +++ b/pgdog/src/backend/pool/pool_impl.rs @@ -196,6 +196,18 @@ impl Pool { } } + /// Server parameters + pub(crate) fn cached_params(&self) -> Option<&Parameters> { + self.inner.params.get() + } + + /// Record the server parameters of a newly created connection + pub(super) fn cache_params(&self, params: &Parameters) { + if self.inner.params.get().is_none() { + let _ = self.inner.params.set(params.clone()); + } + } + /// Get server parameters, fetch them if necessary. pub(crate) async fn params(&self, request: &Request) -> Result<&Parameters, Error> { if let Some(params) = self.inner.params.get() { diff --git a/pgdog/src/backend/prepared_statements.rs b/pgdog/src/backend/prepared_statements.rs index 14a93270d..080bb0c5b 100644 --- a/pgdog/src/backend/prepared_statements.rs +++ b/pgdog/src/backend/prepared_statements.rs @@ -104,10 +104,14 @@ pub(crate) struct PreparedStatements { parses: VecDeque, // Describes being executed now on the connection. describes: VecDeque, + // Statement names of every statement Describe sent (to match each ParameterDescription to its statement) + parameter_describes: VecDeque, config: PreparedStatementsConfig, memory_used: usize, oids: Arc, server_state: State, + // Client parameter count of the unnamed statement being rewritten + anonymous_client_params: Option, } #[cfg(test)] @@ -126,10 +130,12 @@ impl PreparedStatements { state: ProtocolState::default(), parses: VecDeque::new(), describes: VecDeque::new(), + parameter_describes: VecDeque::new(), config: PreparedStatementsConfig::default(), memory_used: 0, oids, server_state: State::Idle, + anonymous_client_params: None, } } @@ -139,6 +145,11 @@ impl PreparedStatements { self.config = config; } + /// Number of parameters the client wrote in the unnamed statement this request rewrites. + pub(crate) fn set_anonymous_client_params(&mut self, params: Option) { + self.anonymous_client_params = params; + } + pub(super) fn set_server_state(&mut self, state: State) { self.server_state = state; } @@ -205,6 +216,11 @@ impl PreparedStatements { } } ProtocolMessage::Describe(describe) => { + if describe.is_statement() { + self.parameter_describes + .push_back(describe.statement().to_string()); + } + if !describe.anonymous() { let message = self.check_prepared(describe.statement())?; @@ -402,6 +418,7 @@ impl PreparedStatements { // These prepared statements have not been prepared, even if they // are syntactically valid. self.describes.clear(); + self.parameter_describes.clear(); self.parses.clear(); } @@ -448,7 +465,8 @@ impl PreparedStatements { } 't' => { - self.rewrite_parameter_description_data_types(message)?; + let statement = self.parameter_describes.pop_front(); + self.rewrite_parameter_description(message, statement.as_deref())?; } _ => (), @@ -655,17 +673,37 @@ impl PreparedStatements { } } - fn rewrite_parameter_description_data_types(&self, message: &mut Message) -> Result<(), Error> { - let Some(mappings) = self.oids.get() else { - return Ok(()); + /// Rewrite the ParameterDescription for `statement` to hide any parameters the rewrite engine added. + /// Asyncpg (& possibly others) check their argument count against this before sending Bind + fn rewrite_parameter_description( + &self, + message: &mut Message, + statement: Option<&str>, + ) -> Result<(), Error> { + let mappings = self + .oids + .get() + .map(|mappings| &mappings.shard_to_canonical) + .filter(|mappings| !mappings.is_empty()); + let client_params = match statement.filter(|name| !name.is_empty()) { + Some(name) => self.global_cache.read().client_params(name), + None => self.anonymous_client_params, }; - let mappings = &mappings.shard_to_canonical; - if mappings.is_empty() { + + if mappings.is_none() && client_params.is_none() { return Ok(()); } let mut parameter_description = ParameterDescription::from_bytes(message.payload())?; - parameter_description.rewrite_data_types(mappings); + if let Some(mappings) = mappings { + parameter_description.rewrite_data_types(mappings); + } + + // Note: This relies on the invariant that the first X parameters are all client-provided + // params, while the ones we re-write are appended to the end. + if let Some(client_params) = client_params { + parameter_description.truncate(client_params as usize); + } message.replace_payload(parameter_description.to_bytes()); Ok(()) } @@ -914,6 +952,48 @@ pub(crate) mod test { rewritten_name } + /// Describe `name` -> forward a ParameterDescription + fn describe_parameters(ps: &mut PreparedStatements, name: &str, oids: Vec) -> Vec { + let result = ps + .handle(&ProtocolMessage::Describe(Describe::new_statement(name))) + .unwrap(); + assert!(matches!(result, HandleResult::Prepend(_))); + + let mut parse_complete = Message::new(ParseComplete.to_bytes()); + ps.forward(&mut parse_complete).unwrap(); + + let mut message = Message::new(ParameterDescription::new(oids).to_bytes()); + ps.forward(&mut message).unwrap(); + + ParameterDescription::from_bytes(message.payload()) + .unwrap() + .params() + .to_vec() + } + + #[test] + fn parameter_description_hides_rewrite_engine_params() { + let mut ps = new_extended(); + let name = insert_global("param_desc_rewritten", "SELECT $1 AS param_desc_rewritten"); + FrontendPreparedStatements::global().write().rewrite( + &Parse::named(&name, "SELECT $1 AS param_desc_rewritten, $2::text"), + 1, + ); + + assert_eq!(describe_parameters(&mut ps, &name, vec![23, 25]), vec![23]); + } + + #[test] + fn parameter_description_untouched_without_rewrite() { + let mut ps = new_extended(); + let name = insert_global("param_desc_plain", "SELECT $1, $2 AS param_desc_plain"); + + assert_eq!( + describe_parameters(&mut ps, &name, vec![23, 25]), + vec![23, 25] + ); + } + #[test] fn ensure_prepared_completes_after_backend_responses() { let mut ps = new_extended(); diff --git a/pgdog/src/backend/replication/logical/subscriber/context.rs b/pgdog/src/backend/replication/logical/subscriber/context.rs index 53b571ab3..a55ae9541 100644 --- a/pgdog/src/backend/replication/logical/subscriber/context.rs +++ b/pgdog/src/backend/replication/logical/subscriber/context.rs @@ -5,7 +5,7 @@ use crate::{ backend::Cluster, frontend::{ BufferedQuery, ClientRequest, Command, PreparedStatements, Router, RouterContext, - client::Sticky, + client::{QueryTimestamps, Sticky}, router::{ parser::{AstContext, Cache, Shard}, sharding::lookup, @@ -50,7 +50,7 @@ impl StreamContext { let parse = stmt.clone(); let mut request = ClientRequest::from(vec![parse.clone().into(), bind.clone().into()]); - let ast_context = AstContext::from_cluster(cluster, &PARAMS); + let ast_context = AstContext::from_cluster(cluster, &PARAMS, QueryTimestamps::now()); let ast = Cache::get().query( &BufferedQuery::Prepared(parse), &ast_context, diff --git a/pgdog/src/backend/server.rs b/pgdog/src/backend/server.rs index 132ad3441..6fe2abb40 100644 --- a/pgdog/src/backend/server.rs +++ b/pgdog/src/backend/server.rs @@ -470,6 +470,9 @@ impl Server { self.in_transaction = true; } + self.prepared_statements + .set_anonymous_client_params(client_request.anonymous_client_params); + for message in client_request.messages.iter() { self.send_one(message).await?; } @@ -2243,7 +2246,14 @@ pub(crate) mod test { let mut prep = PreparedStatements::new(); let name = "test"; let query = Bytes::from("SELECT 1::bigint".to_owned()); - let prepare = prep.insert_prepare(name, query.clone(), None, &RewritePlan::default(), None); + let prepare = prep.insert_prepare( + name, + query.clone(), + None, + &RewritePlan::default(), + None, + vec![], + ); assert_eq!(prepare.name(), "__pgdog_1"); server diff --git a/pgdog/src/frontend/client/mod.rs b/pgdog/src/frontend/client/mod.rs index 9d35b194c..96eb23e9a 100644 --- a/pgdog/src/frontend/client/mod.rs +++ b/pgdog/src/frontend/client/mod.rs @@ -5,8 +5,9 @@ use std::net::SocketAddr; use std::sync::Arc; -use std::time::{Duration, Instant}; +use std::time::Duration; +use chrono::{DateTime, Utc}; use pgdog_config::users::PasswordKind; use timeouts::Timeouts; use tokio::{select, spawn}; @@ -41,7 +42,7 @@ pub(crate) mod transaction_type; use query_engine::QueryEngine; pub(crate) use sticky::Sticky; -pub(crate) use transaction_type::TransactionType; +pub(crate) use transaction_type::{QueryTimestamps, Transaction, TransactionType}; /// PostgreSQL client. /// @@ -75,7 +76,7 @@ pub(crate) struct Client { // Client prepared statements cache. prepared_statements: PreparedStatements, // Client transaction state. - transaction: Option, + transaction: Option, // Current timeouts to use for client/server communication. // These change based on client state, e.g. if client is running query, // the `query_timeout` is active, and if the client is idle, the `client_idle_timeout` is. @@ -96,6 +97,8 @@ pub(crate) struct Client { query_log_stdout: bool, /// Maximum query message size before a warning is logged. query_size_limit: Option, + /// When we received the first message of the current request. + statement_start: DateTime, } /// Inputs to the per-user client certificate check. @@ -435,6 +438,7 @@ impl Client { database: database.to_string(), query_log_stdout: false, query_size_limit: None, + statement_start: Utc::now(), })) } @@ -475,6 +479,7 @@ impl Client { database: "pgdog".to_string(), query_log_stdout: false, query_size_limit: None, + statement_start: Utc::now(), } } @@ -610,7 +615,8 @@ impl Client { QueryEngineResult::Split { requests, extended } => { let mut requests = requests.into_iter(); if extended { - self.transaction.get_or_insert(TransactionType::Implicit); + self.transaction + .get_or_insert(Transaction::new(TransactionType::Implicit)); } while let Some(mut request) = requests.next() { @@ -647,9 +653,6 @@ impl Client { ) -> Result { self.client_request.clear(); - // Only start timer once we receive the first message. - let mut timer = None; - // Check config once per request. let config = config::config(); // Configure prepared statements cache. @@ -660,6 +663,7 @@ impl Client { self.stream_buffer .set_size_limit_block(config.config.general.frontend_query_size_limit_block()); + let mut has_set_time: bool = false; while !self.client_request.is_complete() { let idle_timeout = self .timeouts @@ -695,8 +699,9 @@ impl Client { } }; - if timer.is_none() { - timer = Some(Instant::now()); + if !has_set_time { + has_set_time = true; + self.statement_start = Utc::now(); } // Terminate (B & F). @@ -708,10 +713,11 @@ impl Client { } } + let elapsed_time = Utc::now() - self.statement_start; if !enabled!(LogLevel::TRACE) { debug!( "request buffered [{:.4}ms] {:?}", - timer.unwrap().elapsed().as_secs_f64() * 1000.0, + elapsed_time.as_seconds_f64() * 1000.0, self.client_request .messages .iter() @@ -721,7 +727,7 @@ impl Client { } else { trace!( "request buffered [{:.4}ms]\n{:#?}", - timer.unwrap().elapsed().as_secs_f64() * 1000.0, + elapsed_time.as_seconds_f64() * 1000.0, self.client_request, ); } diff --git a/pgdog/src/frontend/client/query_engine/context.rs b/pgdog/src/frontend/client/query_engine/context.rs index 7f9af0cb1..262c2d012 100644 --- a/pgdog/src/frontend/client/query_engine/context.rs +++ b/pgdog/src/frontend/client/query_engine/context.rs @@ -2,10 +2,15 @@ use crate::{ backend::pool::{connection::mirror::Mirror, stats::MemoryStats}, frontend::{ Client, ClientRequest, PreparedStatements, - client::{Sticky, TransactionType, timeouts::Timeouts}, + client::{ + Sticky, + timeouts::Timeouts, + transaction_type::{QueryTimestamps, Transaction}, + }, }, net::{FrontendPid, Parameters, Stream}, }; +use chrono::{DateTime, Utc}; use super::split::Pipeline; @@ -26,7 +31,7 @@ pub(crate) struct QueryEngineContext<'a> { /// Client's socket to send responses to. pub(super) stream: &'a mut Stream, /// Client in transaction? - pub(super) transaction: Option, + pub(super) transaction: Option, /// Timeouts pub(super) timeouts: Timeouts, /// Cross shard queries are disabled. @@ -43,6 +48,8 @@ pub(crate) struct QueryEngineContext<'a> { pub(super) query_log_stdout: bool, /// Maximum query message size before a warning is logged. pub(super) query_size_limit: Option, + /// When we received the first message of the request. + pub(super) statement_start: DateTime, } impl<'a> QueryEngineContext<'a> { @@ -66,6 +73,7 @@ impl<'a> QueryEngineContext<'a> { sticky: client.sticky, query_log_stdout: client.query_log_stdout, query_size_limit: client.query_size_limit, + statement_start: client.statement_start, } } @@ -96,13 +104,19 @@ impl<'a> QueryEngineContext<'a> { sticky: Sticky::new(), query_log_stdout: false, query_size_limit: None, + statement_start: Utc::now(), } } - pub(crate) fn transaction(&self) -> Option { + pub(crate) fn transaction(&self) -> Option { self.transaction } + /// Request itself can start a transaction, so this is computed "on demand" + pub(crate) fn timestamps(&self) -> QueryTimestamps { + QueryTimestamps::new(self.transaction.as_ref(), self.statement_start) + } + pub(crate) fn in_transaction(&self) -> bool { self.transaction.is_some() } diff --git a/pgdog/src/frontend/client/query_engine/discard.rs b/pgdog/src/frontend/client/query_engine/discard.rs index c2f9252af..f8a7adc91 100644 --- a/pgdog/src/frontend/client/query_engine/discard.rs +++ b/pgdog/src/frontend/client/query_engine/discard.rs @@ -1,4 +1,6 @@ -use crate::frontend::{client::TransactionType, router::parameter_hints::PGDOG_PIN}; +use crate::frontend::{ + client::Transaction, client::TransactionType, router::parameter_hints::PGDOG_PIN, +}; use crate::net::{CommandComplete, Protocol, ReadyForQuery}; use super::*; @@ -15,12 +17,17 @@ impl QueryEngine { match target { DiscardTarget::All if context.in_transaction() => { - context.transaction = Some(match context.transaction { - Some(TransactionType::ReadOnly | TransactionType::ErrorReadOnly) => { - TransactionType::ErrorReadOnly - } - _ => TransactionType::ErrorReadWrite, - }); + context.transaction = Some(Transaction::new( + match context + .transaction + .map(|transaction| transaction.transaction_type()) + { + Some(TransactionType::ReadOnly | TransactionType::ErrorReadOnly) => { + TransactionType::ErrorReadOnly + } + _ => TransactionType::ErrorReadWrite, + }, + )); self.error_response(context, ErrorResponse::discard_all_in_transaction()) .await?; return Ok(()); diff --git a/pgdog/src/frontend/client/query_engine/end_transaction.rs b/pgdog/src/frontend/client/query_engine/end_transaction.rs index 6e134f945..a6215416d 100644 --- a/pgdog/src/frontend/client/query_engine/end_transaction.rs +++ b/pgdog/src/frontend/client/query_engine/end_transaction.rs @@ -141,7 +141,7 @@ impl QueryEngine { mod tests { use super::*; use crate::config::load_test; - use crate::frontend::client::TransactionType; + use crate::frontend::client::{Transaction, TransactionType}; use crate::net::Stream; #[tokio::test] @@ -151,7 +151,7 @@ mod tests { // Create a test client with DevNull stream (doesn't require real I/O) let mut client = crate::frontend::Client::new_test(Stream::dev_null(), Parameters::default()); - client.transaction = Some(TransactionType::ReadWrite); + client.transaction = Some(Transaction::new(TransactionType::ReadWrite)); // Create a default query engine (avoids backend connection) let mut engine = QueryEngine::from_client(&client).unwrap(); @@ -161,9 +161,14 @@ mod tests { assert!(result.is_ok(), "end_transaction should succeed"); assert_eq!( - context.transaction, None, + context + .transaction + .map(|transaction| transaction.transaction_type()), + None, "Transaction state should be None, but is {:?}", - context.transaction + context + .transaction + .map(|transaction| transaction.transaction_type()) ); } } diff --git a/pgdog/src/frontend/client/query_engine/query.rs b/pgdog/src/frontend/client/query_engine/query.rs index 838726bf0..0e23a1e34 100644 --- a/pgdog/src/frontend/client/query_engine/query.rs +++ b/pgdog/src/frontend/client/query_engine/query.rs @@ -2,7 +2,7 @@ use tracing::{info, trace}; use crate::{ frontend::{ - client::TransactionType, + client::{TransactionType, transaction_type::Transaction}, router::parser::{explain_trace::ExplainTrace, rewrite::statement::plan::RewriteResult}, }, net::{ @@ -196,10 +196,12 @@ impl QueryEngine { match state { TransactionState::Error => { - let error_state = match context.transaction { - Some(TransactionType::ReadOnly) => Some(TransactionType::ErrorReadOnly), + let error_state = match context.transaction.map(|t| t.transaction_type()) { + Some(TransactionType::ReadOnly) => { + Some(Transaction::new(TransactionType::ErrorReadOnly)) + } Some(TransactionType::ReadWrite | TransactionType::Implicit) => { - Some(TransactionType::ErrorReadWrite) + Some(Transaction::new(TransactionType::ErrorReadWrite)) } _ => None, }; @@ -221,20 +223,22 @@ impl QueryEngine { self.end_two_pc(false).await?; two_pc_auto = true; } - match context.transaction { + match context.transaction.map(|t| t.transaction_type()) { // Query parser is disabled, so the server is responsible for telling us // we started a transaction. None => { - context.transaction = Some(TransactionType::ReadWrite); + context.transaction = + Some(Transaction::new(TransactionType::ReadWrite)); } // Restore transaction state after rollback to savepoint. Some(TransactionType::ErrorReadOnly) => { - context.transaction = Some(TransactionType::ReadOnly); + context.transaction = Some(Transaction::new(TransactionType::ReadOnly)); } Some(TransactionType::ErrorReadWrite) => { - context.transaction = Some(TransactionType::ReadWrite); + context.transaction = + Some(Transaction::new(TransactionType::ReadWrite)); } _ => (), diff --git a/pgdog/src/frontend/client/query_engine/result.rs b/pgdog/src/frontend/client/query_engine/result.rs index 1aa4d4be7..63c918930 100644 --- a/pgdog/src/frontend/client/query_engine/result.rs +++ b/pgdog/src/frontend/client/query_engine/result.rs @@ -1,9 +1,9 @@ -use crate::frontend::{ClientRequest, client::TransactionType}; +use crate::frontend::{ClientRequest, client::Transaction}; /// Query engine execution result. pub(crate) enum QueryEngineResult { /// Query engine is done executing the request. - Done(Option), + Done(Option), /// Query engine requests the request to be resubmitted /// as a series of separate requests. Split { diff --git a/pgdog/src/frontend/client/query_engine/rewrite.rs b/pgdog/src/frontend/client/query_engine/rewrite.rs index 5e82c0f1c..f07d37128 100644 --- a/pgdog/src/frontend/client/query_engine/rewrite.rs +++ b/pgdog/src/frontend/client/query_engine/rewrite.rs @@ -37,10 +37,17 @@ impl QueryEngine { let query = context.client_request.query()?; if let Some(query) = query { let cluster = self.backend.cluster()?; - let ast_ctx = AstContext::from_cluster(cluster, context.params); + let ast_ctx = AstContext::from_cluster(cluster, context.params, context.timestamps()); let ast = Cache::get().query(&query, &ast_ctx, context.prepared_statements)?; - let rewrite_result = ast.rewrite_plan.apply(context.client_request).await?; + let rewrite_result = ast + .rewrite_plan + .apply( + context.client_request, + ast_ctx.timezone, + ast_ctx.query_timestamps, + ) + .await?; context.client_request.ast = Some(ast); Ok(Some(rewrite_result)) } else { diff --git a/pgdog/src/frontend/client/query_engine/start_transaction.rs b/pgdog/src/frontend/client/query_engine/start_transaction.rs index 70ea4d02e..dee219501 100644 --- a/pgdog/src/frontend/client/query_engine/start_transaction.rs +++ b/pgdog/src/frontend/client/query_engine/start_transaction.rs @@ -1,5 +1,5 @@ use crate::{ - frontend::client::TransactionType, + frontend::client::{TransactionType, transaction_type::Transaction}, net::{ BindComplete, CommandComplete, NoData, NoticeResponse, ParameterDescription, ParseComplete, Protocol, ProtocolMessage, ReadyForQuery, @@ -17,7 +17,7 @@ impl QueryEngine { transaction_type: TransactionType, extended: bool, ) -> Result<(), Error> { - context.transaction = Some(transaction_type); + context.transaction = Some(Transaction::new(transaction_type)); if self.backend.connected() { self.execute(context, None).await?; diff --git a/pgdog/src/frontend/client/query_engine/test/extended_transaction.rs b/pgdog/src/frontend/client/query_engine/test/extended_transaction.rs index e89384f0c..1a8ee83ee 100644 --- a/pgdog/src/frontend/client/query_engine/test/extended_transaction.rs +++ b/pgdog/src/frontend/client/query_engine/test/extended_transaction.rs @@ -1,6 +1,6 @@ use crate::{ expect_message, - frontend::client::TransactionType, + frontend::client::{Transaction, TransactionType}, net::{ BindComplete, CommandComplete, NoData, NoticeResponse, ParameterDescription, ParseComplete, ReadyForQuery, @@ -84,7 +84,7 @@ async fn begin_multiple_describes() { #[tokio::test] async fn commit_statement_describe() { let mut client = TestClient::new_replicas(Parameters::default()).await; - client.client.transaction = Some(TransactionType::ReadWrite); + client.client.transaction = Some(Transaction::new(TransactionType::ReadWrite)); client.client.client_request = ClientRequest::from(vec![ Parse::named("c", "COMMIT").into(), Bind::new_statement("c").into(), diff --git a/pgdog/src/frontend/client/transaction_type.rs b/pgdog/src/frontend/client/transaction_type.rs index 829dc2798..41d6d25ed 100644 --- a/pgdog/src/frontend/client/transaction_type.rs +++ b/pgdog/src/frontend/client/transaction_type.rs @@ -1,3 +1,75 @@ +use chrono::{DateTime, Utc}; + +/// TODO: Not sure if this Transaction refactor is the best way to store the Transaction start time +/// However, if that field just goes separately on the Client, +/// there's more verbosity everywhere we update the Transaction (and potential for bugs) +#[derive(Debug, Clone, Copy)] +pub(crate) struct Transaction { + transaction_type: TransactionType, + start_time: DateTime, +} + +impl Transaction { + pub(crate) fn new(transaction_type: TransactionType) -> Self { + Self { + transaction_type, + start_time: Utc::now(), + } + } + + pub(crate) fn transaction_type(&self) -> TransactionType { + self.transaction_type + } + + pub(crate) fn start_time(&self) -> DateTime { + self.start_time + } + + // pub(crate) fn read_only(&self) -> bool { + // self.transaction_type.read_only() + // } + + pub(crate) fn write(&self) -> bool { + self.transaction_type.write() + } + + pub(crate) fn error(&self) -> bool { + self.transaction_type.error() + } +} + +/// Reference times used to rewrite time functions +/// (e.g. now(), statement_timestamp()) consistently across shards. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) struct QueryTimestamps { + /// Start of the transaction. + /// If not in transaction, this is same as statement start. + pub(crate) transaction_start: DateTime, + /// When we received the first message of the client's request. + pub(crate) statement_start: DateTime, +} + +impl Default for QueryTimestamps { + fn default() -> Self { + QueryTimestamps::now() + } +} + +impl QueryTimestamps { + pub(crate) fn new(transaction: Option<&Transaction>, statement_start: DateTime) -> Self { + Self { + transaction_start: transaction + .map(|t| t.start_time()) + .unwrap_or(statement_start), + statement_start, + } + } + + pub(crate) fn now() -> Self { + Self::new(None, Utc::now()) + } +} + #[derive(Debug, Clone, Copy, PartialEq, Eq, Default)] pub(crate) enum TransactionType { ReadOnly, diff --git a/pgdog/src/frontend/client_request.rs b/pgdog/src/frontend/client_request.rs index 155060b05..88ed647a1 100644 --- a/pgdog/src/frontend/client_request.rs +++ b/pgdog/src/frontend/client_request.rs @@ -33,6 +33,8 @@ pub(crate) struct ClientRequest { pub(crate) ast: Option, /// Last Parse we received. pub(crate) last_parse: Option, + /// How many parameters the client wrote in the unnamed prepared statement + pub(crate) anonymous_client_params: Option, } impl MemoryUsage for ClientRequest { @@ -58,6 +60,7 @@ impl ClientRequest { route: None, ast: None, last_parse: None, + anonymous_client_params: None, } } @@ -215,6 +218,7 @@ impl ClientRequest { route: self.route.clone(), ast: self.ast.clone(), last_parse: None, + anonymous_client_params: self.anonymous_client_params, } } @@ -389,6 +393,7 @@ impl From> for ClientRequest { route: None, ast: None, last_parse: None, + anonymous_client_params: None, } } } diff --git a/pgdog/src/frontend/prepared_statements/global_cache.rs b/pgdog/src/frontend/prepared_statements/global_cache.rs index 4908b10db..c5396c842 100644 --- a/pgdog/src/frontend/prepared_statements/global_cache.rs +++ b/pgdog/src/frontend/prepared_statements/global_cache.rs @@ -1,5 +1,5 @@ use crate::{ - frontend::RewritePlan, + frontend::{RewritePlan, router::parser::rewrite::statement::plan::GeneratedParam}, net::{ Prepare, messages::{Parse, RowDescription}, @@ -70,6 +70,7 @@ impl GlobalCache { stmt: StatementType::Parse { parse, rewrite: None, + client_params: None, }, cache_key: cache_key.clone(), row_description: None, @@ -94,6 +95,7 @@ impl GlobalCache { // to use `RewritePlan` for `offset_plan` too (which isn't possible; see comment below) rewrite_plan: &RewritePlan, offset_plan: Option, + generated_params: Vec, ) -> (bool, Prepare) { let cache_key = CacheKey::Simple { query: original_query.clone(), @@ -114,14 +116,15 @@ impl GlobalCache { }; let statement = Statement { - stmt: StatementType::Prepare { + stmt: StatementType::Prepare(PreparedPlan { prepare: prepare.clone(), unique_ids: rewrite_plan.unique_ids, // The reason this isn't using [`rewrite_plan.offset`] is that in `rewrite_single_prepared`, // for `PrepareStmt`, we don't set `offset` on`RewritePlan` yet. We only attach `offset` // to the plan for `ExecuteStmt`, and we need access to `OffsetPlan` for both here. offset_plan, - }, + generated_params, + }), row_description: None, cache_key: cache_key.clone(), }; @@ -131,12 +134,19 @@ impl GlobalCache { } /// Rewrite prepared statement in the global cache. - pub(crate) fn rewrite(&mut self, parse: &Parse) { + /// `client_params` indicates how many Bind parameters the original statement has. + pub(crate) fn rewrite(&mut self, parse: &Parse, client_params: u16) { if let Some(stmt) = self.names.get_mut(parse.name()) { - stmt.set_rewrite(parse); + stmt.set_rewrite(parse, client_params); } } + /// Number of parameters the client's original statement has + /// (if we re-write, we must catch and not send back the extra cols ParameterDescriptions) + pub(crate) fn client_params(&self, name: &str) -> Option { + self.names.get(name).and_then(|stmt| stmt.client_params()) + } + /// Client sent a Describe for a prepared statement and received a RowDescription. /// We record the RowDescription for later use by the results decoder. pub(crate) fn insert_row_description(&mut self, name: &str, row_description: RowDescription) { @@ -158,17 +168,12 @@ impl GlobalCache { /// Get the [`Prepare`] message for a globally unique prepare statement name. pub(crate) fn prepare(&self, name: &str) -> Option { - self.prepare_and_unique_ids(name) - .map(|(prepare, _, _)| prepare) + self.prepared_plan(name).map(|plan| plan.prepare) } - pub(crate) fn prepare_and_unique_ids( - &self, - name: &str, - ) -> Option<(Prepare, u16, Option)> { - self.names - .get(name) - .and_then(|p| p.prepare_and_unique_ids()) + /// Fetch the `PreparedPlan` for a globally unique prepare statement name. + pub(crate) fn prepared_plan(&self, name: &str) -> Option { + self.names.get(name).and_then(|p| p.prepared_plan()) } /// Get the rewritten Parse statement. @@ -392,8 +397,9 @@ mod test { let query = Bytes::from("PREPARE __pgdog_template_name AS SELECT $1"); let parse = Parse::named("client_stmt", "SELECT $1"); - let (_, first) = cache.insert_prepare(query.clone(), None, &RewritePlan::default(), None); - let (_, second) = cache.insert_prepare(query, None, &RewritePlan::default(), None); + let (_, first) = + cache.insert_prepare(query.clone(), None, &RewritePlan::default(), None, vec![]); + let (_, second) = cache.insert_prepare(query, None, &RewritePlan::default(), None, vec![]); assert_eq!(first, second); assert_eq!(cache.len(), 1); diff --git a/pgdog/src/frontend/prepared_statements/mod.rs b/pgdog/src/frontend/prepared_statements/mod.rs index 3e8b40654..11a66af20 100644 --- a/pgdog/src/frontend/prepared_statements/mod.rs +++ b/pgdog/src/frontend/prepared_statements/mod.rs @@ -8,7 +8,10 @@ use parking_lot::RwLock; use crate::{ config::PreparedStatementsLevel, - frontend::{RewritePlan, router::parser::rewrite::statement::offset::OffsetPlan}, + frontend::{ + RewritePlan, + router::parser::rewrite::statement::{offset::OffsetPlan, plan::GeneratedParam}, + }, net::{Parse, Prepare, ProtocolMessage}, }; @@ -29,7 +32,7 @@ pub(crate) use global_cache::GlobalCache; // Maintenance tasks are spawned in main.rs. pub(crate) use maintenance::*; pub(crate) use rewrite::Rewrite; -pub(crate) use statement::{Statement, StatementType}; +pub(crate) use statement::{PreparedPlan, Statement, StatementType}; static CACHE: Lazy = Lazy::new(PreparedStatements::default); @@ -122,6 +125,7 @@ impl PreparedStatements { rewrite_plan: &RewritePlan, // Needs to be separate from `RewritePlan`. See comment in `global_cache.rs`. offset_plan: Option, + generated_params: Vec, ) -> Prepare { let (_new, prepare) = { self.global.write().insert_prepare( @@ -129,6 +133,7 @@ impl PreparedStatements { rewritten_query, rewrite_plan, offset_plan, + generated_params, ) }; @@ -143,14 +148,11 @@ impl PreparedStatements { self.local.get(name) } - /// Get a globally unique [`Prepare`] message using the client name as key. - pub(crate) fn prepare_and_unique_ids( - &self, - name: &str, - ) -> Option<(Prepare, u16, Option)> { + /// Get a globally unique `PreparedPlan` using the client name as key. + pub(crate) fn prepared_plan(&self, name: &str) -> Option { self.local .get(name) - .and_then(|name| self.global.read().prepare_and_unique_ids(name)) + .and_then(|name| self.global.read().prepared_plan(name)) } /// Number of prepared statements in the client's cache. diff --git a/pgdog/src/frontend/prepared_statements/statement.rs b/pgdog/src/frontend/prepared_statements/statement.rs index bfc8fc402..264c83839 100644 --- a/pgdog/src/frontend/prepared_statements/statement.rs +++ b/pgdog/src/frontend/prepared_statements/statement.rs @@ -1,5 +1,6 @@ use crate::{ - frontend::router::parser::rewrite::statement::offset::OffsetPlan, net::Prepare, + frontend::router::parser::rewrite::statement::{offset::OffsetPlan, plan::GeneratedParam}, + net::Prepare, stats::memory::MemoryUsage, }; @@ -12,34 +13,41 @@ pub(crate) struct Statement { pub(super) cache_key: CacheKey, } +#[derive(Debug, Clone)] +pub(crate) struct PreparedPlan { + pub(crate) prepare: Prepare, + + /// The number of calls to `pgdog.unique_id` which were previously + /// rewritten. If this value is greater than zero, it is expected + /// that the query in the [`Parse`] message referenced by + /// [`Self::prepare`] was previously rewritten to replace those calls + /// with bind parameter placeholder numbered after all others + pub(crate) unique_ids: u16, + + /// Used to keep track of LIMIT + OFFSET queries (stemming from Prepare), + /// where we have to re-write `A_Const` nodes with `ParamRefs`, so that we can dynamically + /// modify limit/offset values before execution if it ends up being cross-shard. + pub(crate) offset_plan: Option, + + pub(crate) generated_params: Vec, +} + #[derive(Debug, Clone)] pub(crate) enum StatementType { Parse { parse: Parse, rewrite: Option, + client_params: Option, }, - Prepare { - prepare: Prepare, - /// The number of calls to `pgdog.unique_id` which were previously - /// rewritten. If this value is greater than zero, it is expected - /// that the query in the [`Parse`] message referenced by - /// [`Self::prepare`] was previously rewritten to replace those calls - /// with bind parameter placeholder numbered after all others - unique_ids: u16, - - /// Used to keep track of LIMIT + OFFSET queries (stemming from Prepare), - /// where we have to re-write `A_Const` nodes with `ParamRefs`, so that we can dynamically - /// modify limit/offset values before execution if it ends up being cross-shard. - offset_plan: Option, - }, + Prepare(PreparedPlan), } impl MemoryUsage for StatementType { fn memory_usage(&self) -> usize { match self { - Self::Prepare { prepare, .. } => prepare.len(), - Self::Parse { parse, rewrite } => { + Self::Prepare(plan) => plan.prepare.len(), + Self::Parse { parse, rewrite, .. } => { parse.len() + rewrite .as_ref() @@ -71,13 +79,9 @@ impl Statement { } } - pub(super) fn prepare_and_unique_ids(&self) -> Option<(Prepare, u16, Option)> { + pub(super) fn prepared_plan(&self) -> Option { match &self.stmt { - StatementType::Prepare { - prepare, - unique_ids, - offset_plan, - } => Some((prepare.clone(), *unique_ids, offset_plan.clone())), + StatementType::Prepare(plan) => Some(plan.clone()), _ => None, } } @@ -93,12 +97,22 @@ impl Statement { &self.cache_key } - pub(super) fn set_rewrite(&mut self, parse: &Parse) { + pub(crate) fn client_params(&self) -> Option { + match self.stmt { + StatementType::Parse { client_params, .. } => client_params, + _ => None, + } + } + + pub(super) fn set_rewrite(&mut self, parse: &Parse, params: u16) { if let StatementType::Parse { - ref mut rewrite, .. + ref mut rewrite, + ref mut client_params, + .. } = self.stmt { - *rewrite = Some(parse.clone()) + *rewrite = Some(parse.clone()); + *client_params = Some(params); } } } @@ -111,7 +125,7 @@ mod test { pub(crate) fn query(&self) -> &str { match self.stmt { StatementType::Parse { ref parse, .. } => parse.query(), - StatementType::Prepare { ref prepare, .. } => prepare.query(), + StatementType::Prepare(ref plan) => plan.prepare.query(), } } } diff --git a/pgdog/src/frontend/router/context.rs b/pgdog/src/frontend/router/context.rs index 475b75ade..121a2f564 100644 --- a/pgdog/src/frontend/router/context.rs +++ b/pgdog/src/frontend/router/context.rs @@ -1,9 +1,10 @@ use super::{Error, ParameterHints}; +use crate::frontend::client::transaction_type::Transaction; use crate::{ backend::{Cluster, Schema}, frontend::{ BufferedQuery, ClientRequest, - client::{Sticky, TransactionType}, + client::Sticky, router::{Ast, parser::StatementParameters, sharding::ResolvedLookups}, }, net::Parameters, @@ -20,7 +21,7 @@ pub(crate) struct RouterContext<'a> { /// Client parameters, e.g. search_path. pub(super) parameter_hints: ParameterHints<'a>, /// Client inside transaction, - pub(super) transaction: Option, + pub(super) transaction: Option, /// Currently executing COPY statement. pub(super) copy_mode: bool, /// Do we have an executable buffer? @@ -46,7 +47,7 @@ impl<'a> RouterContext<'a> { buffer: &'a ClientRequest, cluster: &'a Cluster, params: &'a Parameters, - transaction: Option, + transaction: Option, sticky: Sticky, ) -> Result { let query = buffer.query()?; @@ -83,7 +84,7 @@ impl<'a> RouterContext<'a> { self.transaction.is_some() } - pub(crate) fn transaction(&self) -> &Option { + pub(crate) fn transaction(&self) -> &Option { &self.transaction } } diff --git a/pgdog/src/frontend/router/parser/cache/ast.rs b/pgdog/src/frontend/router/parser/cache/ast.rs index 3a94302eb..a63348a92 100644 --- a/pgdog/src/frontend/router/parser/cache/ast.rs +++ b/pgdog/src/frontend/router/parser/cache/ast.rs @@ -10,13 +10,11 @@ use tracing::warn; use super::super::{Error, Route, StatementRewrite, StatementRewriteContext}; use super::Stats; -use crate::backend::schema::Schema; +use crate::config::Role; use crate::frontend::PreparedStatements; use crate::frontend::router::parser::cache::AstQuery; use crate::frontend::router::parser::rewrite::statement::RewritePlan; use crate::frontend::router::sharding::ShardOrLookup; -use crate::net::parameter::ParameterValue; -use crate::{backend::ShardingSchema, config::Role}; /// Abstract syntax tree (query) cache entry, /// with statistics. @@ -70,11 +68,8 @@ impl Ast { /// Parse statement and run the rewrite engine, if necessary. pub(super) fn new( query: &AstQuery, - schema: &ShardingSchema, - db_schema: &Schema, + ctx: &super::AstContext<'_>, prepared_statements: &mut PreparedStatements, - user: &str, - search_path: Option<&ParameterValue>, ) -> Result { let now = Instant::now(); @@ -86,10 +81,12 @@ impl Ast { extended: query.original_query.extended(), prepared: query.original_query.prepared(), prepared_statements, - schema, - db_schema, - user, - search_path, + schema: &ctx.sharding_schema, + db_schema: &ctx.db_schema, + user: ctx.user, + search_path: ctx.search_path, + timezone: ctx.timezone, + query_timestamps: ctx.query_timestamps, }); let mut rewrite_plan = Default::default(); let ast = make::try_owned(|mem| { @@ -105,13 +102,13 @@ impl Ast { let mut stats = Stats::new(); stats.parse_time += elapsed; - if let Some(threshold) = schema.log_min_duration_parse + if let Some(threshold) = ctx.sharding_schema.log_min_duration_parse && elapsed >= threshold { warn!( "[slow_query_parse] parse_time_in_ms={}ms truncated_query=\"{}\"", elapsed.as_millis(), - query.truncated_query(schema.log_query_sample_length), + query.truncated_query(ctx.sharding_schema.log_query_sample_length), ); } @@ -129,22 +126,6 @@ impl Ast { }) } - /// Parse statement using AstContext for schema and user information. - pub(super) fn with_context( - query: &AstQuery, - ctx: &super::AstContext<'_>, - prepared_statements: &mut PreparedStatements, - ) -> Result { - Self::new( - query, - &ctx.sharding_schema, - &ctx.db_schema, - prepared_statements, - ctx.user, - ctx.search_path, - ) - } - /// Record new AST entry, without rewriting or comment-routing. pub(crate) fn new_record(query: &str) -> Result { let ast = pg_raw_parse::parse(query)?; diff --git a/pgdog/src/frontend/router/parser/cache/cache_impl.rs b/pgdog/src/frontend/router/parser/cache/cache_impl.rs index 9990e98e8..365ba0308 100644 --- a/pgdog/src/frontend/router/parser/cache/cache_impl.rs +++ b/pgdog/src/frontend/router/parser/cache/cache_impl.rs @@ -127,7 +127,7 @@ impl Cache { } // Parse query without holding lock. - let mut entry = Ast::with_context( + let mut entry = Ast::new( &AstQuery { original_query: query, query_without_comment: query_and_comment.query, @@ -169,7 +169,7 @@ impl Cache { ) -> Result { let query_and_comment = parse_edge_comment(query.query(), &ctx.sharding_schema)?; - let mut entry = Ast::with_context( + let mut entry = Ast::new( &AstQuery { original_query: query, query_without_comment: query_and_comment.query, diff --git a/pgdog/src/frontend/router/parser/cache/context.rs b/pgdog/src/frontend/router/parser/cache/context.rs index 774f70219..e4e8f46f5 100644 --- a/pgdog/src/frontend/router/parser/cache/context.rs +++ b/pgdog/src/frontend/router/parser/cache/context.rs @@ -4,6 +4,7 @@ use crate::backend::ShardingSchema; use crate::backend::pool::Cluster; use crate::backend::schema::Schema; use crate::frontend::BufferedQuery; +use crate::frontend::client::QueryTimestamps; use crate::net::Parameters; use crate::net::parameter::ParameterValue; @@ -23,16 +24,28 @@ pub(crate) struct AstContext<'a> { pub(crate) user: &'a str, /// Search path for table lookups. pub(crate) search_path: Option<&'a ParameterValue>, + /// Allows `timestamp` types to use the Client's local time when excecuting a `TimeFunction` + pub(crate) timezone: Option<&'a ParameterValue>, + /// Statement, and transaction DateTime relevant to the current Query (if not being cached) + pub(crate) query_timestamps: QueryTimestamps, } impl<'a> AstContext<'a> { /// Create AstContext from a Cluster and Parameters. - pub(crate) fn from_cluster(cluster: &'a Cluster, params: &'a Parameters) -> Self { + pub(crate) fn from_cluster( + cluster: &'a Cluster, + params: &'a Parameters, + query_timestamps: QueryTimestamps, + ) -> Self { Self { sharding_schema: cluster.sharding_schema(), db_schema: cluster.schema(), user: cluster.user(), search_path: params.get("search_path"), + timezone: params + .get("timezone") + .or_else(|| cluster.default_timezone()), + query_timestamps, } } } diff --git a/pgdog/src/frontend/router/parser/context.rs b/pgdog/src/frontend/router/parser/context.rs index 76a16774b..cf026b7a3 100644 --- a/pgdog/src/frontend/router/parser/context.rs +++ b/pgdog/src/frontend/router/parser/context.rs @@ -99,7 +99,9 @@ impl<'a> QueryParserContext<'a> { pub(super) fn write_override(&self) -> bool { let role = self.router_context.parameter_hints.compute_role(); let txn_write = matches!( - self.router_context.transaction(), + self.router_context + .transaction() + .map(|t| t.transaction_type()), Some(TransactionType::ReadWrite | TransactionType::Implicit) ) && self.rw_conservative(); // prefer_primary defaults reads to the primary; an explicit replica hint opts out. diff --git a/pgdog/src/frontend/router/parser/query/explain.rs b/pgdog/src/frontend/router/parser/query/explain.rs index 12411618a..f011fe2d9 100644 --- a/pgdog/src/frontend/router/parser/query/explain.rs +++ b/pgdog/src/frontend/router/parser/query/explain.rs @@ -51,7 +51,7 @@ mod tests { use crate::backend::Cluster; use crate::config::{self, config}; - use crate::frontend::client::Sticky; + use crate::frontend::client::{QueryTimestamps, Sticky}; use crate::frontend::router::parser::{AstContext, Cache}; use crate::frontend::{BufferedQuery, ClientRequest, PreparedStatements, RouterContext}; use crate::net::{ @@ -76,7 +76,7 @@ mod tests { let cluster = Cluster::new_test(&config()); let mut stmts = PreparedStatements::default(); let params = Parameters::default(); - let ast_ctx = AstContext::from_cluster(&cluster, ¶ms); + let ast_ctx = AstContext::from_cluster(&cluster, ¶ms, QueryTimestamps::default()); let buffered = BufferedQuery::Query(Query::new(sql)); let ast = Cache::get().query(&buffered, &ast_ctx, &mut stmts).unwrap(); @@ -108,7 +108,7 @@ mod tests { let cluster = Cluster::new_test(&config()); let mut stmts = PreparedStatements::default(); let params = Parameters::default(); - let ast_ctx = AstContext::from_cluster(&cluster, ¶ms); + let ast_ctx = AstContext::from_cluster(&cluster, ¶ms, QueryTimestamps::default()); let buffered = BufferedQuery::Prepared(Parse::new_anonymous(sql)); let ast = Cache::get().query(&buffered, &ast_ctx, &mut stmts).unwrap(); diff --git a/pgdog/src/frontend/router/parser/query/show.rs b/pgdog/src/frontend/router/parser/query/show.rs index c4050fbcd..440d64886 100644 --- a/pgdog/src/frontend/router/parser/query/show.rs +++ b/pgdog/src/frontend/router/parser/query/show.rs @@ -32,7 +32,7 @@ impl QueryParser { mod test_show { use crate::backend::Cluster; use crate::config::config; - use crate::frontend::client::Sticky; + use crate::frontend::client::{QueryTimestamps, Sticky}; use crate::frontend::router::QueryParser; use crate::frontend::router::parser::{AstContext, Cache, Shard}; use crate::frontend::{BufferedQuery, ClientRequest, PreparedStatements, RouterContext}; @@ -44,7 +44,7 @@ mod test_show { let c = Cluster::new_test(&config()); let mut parser = QueryParser::default(); let params = Parameters::default(); - let ctx = AstContext::from_cluster(&c, ¶ms); + let ctx = AstContext::from_cluster(&c, ¶ms, QueryTimestamps::default()); // First call let query = "SHOW TRANSACTION ISOLATION LEVEL"; diff --git a/pgdog/src/frontend/router/parser/query/test/mod.rs b/pgdog/src/frontend/router/parser/query/test/mod.rs index d16050f5a..4892c005d 100644 --- a/pgdog/src/frontend/router/parser/query/test/mod.rs +++ b/pgdog/src/frontend/router/parser/query/test/mod.rs @@ -2,6 +2,7 @@ use crate::{ config::config, + frontend::client::QueryTimestamps, net::{ Close, Format, Parameters, Sync, messages::{bind::Parameter, parse::Parse}, @@ -15,7 +16,7 @@ use crate::config::ReadWriteStrategy; use crate::frontend::router::parser::{AstContext, Cache}; use crate::frontend::{ BufferedQuery, ClientRequest, PreparedStatements, RouterContext, - client::{Sticky, TransactionType}, + client::{Sticky, Transaction, TransactionType}, }; use crate::net::messages::Query; @@ -48,7 +49,7 @@ fn parse_query(query: &str) -> Command { let cluster = Cluster::new_test(&config()); let buffered = BufferedQuery::Query(Query::new(query)); let params = Parameters::default(); - let ctx = AstContext::from_cluster(&cluster, ¶ms); + let ctx = AstContext::from_cluster(&cluster, ¶ms, QueryTimestamps::default()); let ast = Cache::get() .query(&buffered, &ctx, &mut PreparedStatements::default()) .unwrap(); @@ -70,14 +71,18 @@ macro_rules! command { let cluster = Cluster::new_test(&crate::config::config()); let buffered = BufferedQuery::Query(Query::new($query)); let params = Parameters::default(); - let ctx = crate::frontend::router::parser::AstContext::from_cluster(&cluster, ¶ms); + let ctx = crate::frontend::router::parser::AstContext::from_cluster( + &cluster, + ¶ms, + QueryTimestamps::default(), + ); let ast = crate::frontend::router::parser::Cache::get() .query(&buffered, &ctx, &mut PreparedStatements::default()) .unwrap(); let mut client_request = ClientRequest::from(vec![Query::new(query).into()]); client_request.ast = Some(ast); let transaction = if $in_transaction { - Some(TransactionType::ReadWrite) + Some(Transaction::new(TransactionType::ReadWrite)) } else { None }; @@ -122,7 +127,11 @@ macro_rules! query_parser { let mut prep_stmts = PreparedStatements::default(); let params = Parameters::default(); - let ctx = crate::frontend::router::parser::AstContext::from_cluster(&cluster, ¶ms); + let ctx = crate::frontend::router::parser::AstContext::from_cluster( + &cluster, + ¶ms, + QueryTimestamps::default(), + ); let mut ast = crate::frontend::router::parser::Cache::get() .query(&buffered_query, &ctx, &mut prep_stmts) @@ -131,7 +140,7 @@ macro_rules! query_parser { client_request.ast = Some(ast); let maybe_transaction = if $in_transaction { - Some(TransactionType::ReadWrite) + Some(Transaction::new(TransactionType::ReadWrite)) } else { None }; @@ -176,8 +185,11 @@ macro_rules! parse { let cluster = Cluster::new_test(&crate::config::config()); let buffered = BufferedQuery::Prepared(Parse::new_anonymous($query)); let client_params = Parameters::default(); - let ctx = - crate::frontend::router::parser::AstContext::from_cluster(&cluster, &client_params); + let ctx = crate::frontend::router::parser::AstContext::from_cluster( + &cluster, + &client_params, + QueryTimestamps::default(), + ); let ast = crate::frontend::router::parser::Cache::get() .query(&buffered, &ctx, &mut PreparedStatements::default()) .unwrap(); @@ -462,13 +474,13 @@ fn test_set() { let mut prep_stmts = PreparedStatements::default(); let buffered_query = BufferedQuery::Query(Query::new(query_str)); let params = Parameters::default(); - let ctx = AstContext::from_cluster(&cluster, ¶ms); + let ctx = AstContext::from_cluster(&cluster, ¶ms, QueryTimestamps::default()); let ast = Cache::get() .query(&buffered_query, &ctx, &mut prep_stmts) .unwrap(); let mut buffer: ClientRequest = vec![Query::new(query_str).into()].into(); buffer.ast = Some(ast); - let transaction = Some(TransactionType::ReadWrite); + let transaction = Some(Transaction::new(TransactionType::ReadWrite)); let router_context = RouterContext::new(&buffer, &cluster, ¶ms, transaction, Sticky::new()).unwrap(); let mut context = QueryParserContext::new(router_context).unwrap(); @@ -607,13 +619,13 @@ WHERE t2.account = ( "; let buffered_query = BufferedQuery::Query(Query::new(query_str)); let params = Parameters::default(); - let ctx = AstContext::from_cluster(&cluster, ¶ms); + let ctx = AstContext::from_cluster(&cluster, ¶ms, QueryTimestamps::now()); let ast = Cache::get() .query(&buffered_query, &ctx, &mut prep_stmts) .unwrap(); let mut buffer: ClientRequest = vec![Query::new(query_str).into()].into(); buffer.ast = Some(ast); - let transaction = Some(TransactionType::ReadWrite); + let transaction = Some(Transaction::new(TransactionType::ReadWrite)); let router_context = RouterContext::new(&buffer, &cluster, ¶ms, transaction, Sticky::new()).unwrap(); let mut context = QueryParserContext::new(router_context).unwrap(); diff --git a/pgdog/src/frontend/router/parser/query/test/setup.rs b/pgdog/src/frontend/router/parser/query/test/setup.rs index e41a2de2c..cc51713ae 100644 --- a/pgdog/src/frontend/router/parser/query/test/setup.rs +++ b/pgdog/src/frontend/router/parser/query/test/setup.rs @@ -7,7 +7,7 @@ use crate::{ config::{self, ReadWriteStrategy, config}, frontend::{ ClientRequest, Command, PreparedStatements, RouterContext, - client::{Sticky, TransactionType}, + client::{QueryTimestamps, Sticky, TransactionType, transaction_type::Transaction}, router::{ QueryParser, parser::{AstContext, Cache, Error}, @@ -22,7 +22,7 @@ pub(super) use crate::net::*; pub(crate) struct QueryParserTest { cluster: Cluster, params: Parameters, - transaction: Option, + transaction: Option, sticky: Sticky, prepared: PreparedStatements, pub(crate) parser: QueryParser, @@ -92,7 +92,7 @@ impl QueryParserTest { /// Set whether we're in a transaction. pub(crate) fn in_transaction(mut self, in_tx: bool) -> Self { self.transaction = if in_tx { - Some(TransactionType::ReadWrite) + Some(Transaction::new(TransactionType::ReadWrite)) } else { None }; @@ -102,7 +102,7 @@ impl QueryParserTest { /// Set the exact transaction state, e.g. `TransactionType::ReadOnly` for /// a `BEGIN READ ONLY` transaction. pub(crate) fn with_transaction(mut self, transaction: TransactionType) -> Self { - self.transaction = Some(transaction); + self.transaction = Some(Transaction::new(transaction)); self } @@ -215,7 +215,11 @@ impl QueryParserTest { if use_parser { // Some requests (like Close) don't have a query if let Ok(Some(buffered_query)) = request.query() { - let ctx = AstContext::from_cluster(&self.cluster, &self.params); + let ctx = AstContext::from_cluster( + &self.cluster, + &self.params, + QueryTimestamps::default(), + ); // The engine surfaces cache-time errors (e.g. a comment // directive that fails to resolve) as client errors. let ast = Cache::get().query(&buffered_query, &ctx, &mut self.prepared)?; diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/auto_id.rs b/pgdog/src/frontend/router/parser/rewrite/statement/auto_id.rs index 347fd3147..3f693d863 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/auto_id.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/auto_id.rs @@ -110,7 +110,7 @@ impl StatementRewrite<'_> { } /// Get the table from an INSERT statement. - fn get_insert_table<'a>(&self, insert: &'a nodes::InsertStmt) -> (Table<'a>, bool) { + pub(crate) fn get_insert_table<'a>(&self, insert: &'a nodes::InsertStmt) -> (Table<'a>, bool) { let relation = insert.relation().expect("INSERT always has table"); let is_sharded = StatementParser::new(insert.into(), None, self.schema, None).is_sharded( self.db_schema, @@ -232,6 +232,8 @@ mod split_tests; mod tests { use super::super::nextval::SequenceCall; use super::super::plan::GeneratedId; + use crate::frontend::client::QueryTimestamps; + use crate::frontend::router::parser::rewrite::statement::plan::GeneratedParam; use crate::frontend::router::sharding::ShardedTable; use indexmap::IndexMap; use pgdog_config::{Rewrite, SystemCatalogsBehavior}; @@ -551,6 +553,8 @@ mod tests { db_schema, user: "", search_path: None, + timezone: None, + query_timestamps: QueryTimestamps::default(), }); let mut plan = Default::default(); let ast = make::try_owned(|mem| { @@ -693,16 +697,20 @@ mod tests { assert_eq!(plan.auto_id_injected, injected); assert_eq!(plan.unique_ids, 0); assert_eq!( - plan.generated_ids, + plan.generated_params, vec![ - ( - 1, - GeneratedId::Sequence(SequenceCall::Nextval(sequence.to_owned())) - ), - ( - 2, - GeneratedId::Sequence(SequenceCall::Nextval(sequence.to_owned())) - ), + GeneratedParam { + param_num: 1, + generated_id: GeneratedId::Sequence(SequenceCall::Nextval( + sequence.to_owned() + )) + }, + GeneratedParam { + param_num: 2, + generated_id: GeneratedId::Sequence(SequenceCall::Nextval( + sequence.to_owned() + )) + }, ] ); } @@ -722,7 +730,7 @@ mod tests { assert_eq!(sql, original); assert_eq!(plan.auto_id_injected, 0); - assert!(plan.generated_ids.is_empty()); + assert!(plan.generated_params.is_empty()); } } } diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/auto_id/split_tests.rs b/pgdog/src/frontend/router/parser/rewrite/statement/auto_id/split_tests.rs index 2dc025256..a38301426 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/auto_id/split_tests.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/auto_id/split_tests.rs @@ -3,6 +3,7 @@ use super::super::plan::RewriteResult; use super::tests::make_schema_with_bigint_pk; use super::*; use crate::backend::ShardingSchema; +use crate::frontend::client::QueryTimestamps; use crate::frontend::router::parser::StatementRewriteContext; use crate::frontend::{ClientRequest, PreparedStatements}; use crate::net::messages::bind::{Format, Parameter}; @@ -30,6 +31,8 @@ fn split_plan(sql: &str, extended: bool, prepared: bool) -> RewritePlan { db_schema: &db_schema, user: "", search_path: None, + timezone: None, + query_timestamps: QueryTimestamps::default(), }); let mut plan = RewritePlan::default(); make::owned(|mem| { @@ -128,7 +131,11 @@ async fn test_nextval_auto_id_extended_splits_keep_generated_parameters() { let mut prepare_request = ClientRequest::from(vec![ProtocolMessage::Parse(parse.clone())]); let result = plan - .apply(&mut prepare_request) + .apply( + &mut prepare_request, + None, + crate::frontend::client::QueryTimestamps::now(), + ) .await .expect("prepare succeeds"); assert!(matches!(result, RewriteResult::InPlace { .. })); @@ -140,11 +147,16 @@ async fn test_nextval_auto_id_extended_splits_keep_generated_parameters() { &[format], ); let mut value = 200; - plan.apply_generated_ids(&mut bind, async |call: &SequenceCall| { - assert_eq!(call, &SequenceCall::Nextval("users_id_seq".into())); - value += 1; - Ok(value) - }) + plan.apply_generated_ids( + &mut bind, + None, + crate::frontend::client::QueryTimestamps::now(), + async |call: &SequenceCall| { + assert_eq!(call, &SequenceCall::Nextval("users_id_seq".into())); + value += 1; + Ok(value) + }, + ) .await .expect("sequence values appended"); let request = if separate_parse { @@ -188,7 +200,7 @@ async fn test_nextval_auto_id_extended_splits_keep_generated_parameters() { .as_ref() .expect("AST") .rewrite_plan - .generated_ids + .generated_params .is_empty() ); } diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/error.rs b/pgdog/src/frontend/router/parser/rewrite/statement/error.rs index 0e7a02ad4..8c576c0f4 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/error.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/error.rs @@ -46,4 +46,12 @@ pub(crate) enum Error { #[error("missing or invalid parameters in Execute")] IncorrectExecuteParameters, + + #[error( + "TimeZone {0} is not supported for time functions on omnisharded tables, only IANA names like UTC or America/New_York are (other formats are future work)" + )] + UnsupportedTimeZone(String), + + #[error("could not determine the session TimeZone for time functions on omnisharded tables")] + UnknownTimeZone, } diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/insert.rs b/pgdog/src/frontend/router/parser/rewrite/statement/insert.rs index 1e4e40f22..fd4f33274 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/insert.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/insert.rs @@ -258,6 +258,7 @@ mod tests { use crate::backend::ShardingSchema; use crate::backend::schema::Schema; use crate::frontend::PreparedStatements; + use crate::frontend::client::QueryTimestamps; use crate::frontend::router::parser::StatementRewriteContext; use crate::net::messages::bind::{Format, Parameter}; @@ -294,6 +295,8 @@ mod tests { db_schema: &db_schema, user: "", search_path: None, + timezone: None, + query_timestamps: QueryTimestamps::default(), }); let mut plan = RewritePlan::default(); rewriter.split_insert(insert, &mut plan).unwrap(); diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs b/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs index 813f0be50..f49ab68e6 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/mod.rs @@ -1,10 +1,12 @@ //! Statement rewriter. -use crate::backend::ShardingSchema; use crate::backend::schema::Schema; +use crate::config::config; use crate::frontend::PreparedStatements; use crate::frontend::router::parser::AstContext; +use crate::frontend::router::parser::rewrite::statement::plan::GeneratedParam; use crate::net::parameter::ParameterValue; +use crate::{backend::ShardingSchema, frontend::client::QueryTimestamps}; use pg_raw_parse::{Node, NodeMut, make, nodes, transform, walk}; pub(crate) mod aggregate; @@ -15,11 +17,14 @@ pub(crate) mod nextval; 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; pub(crate) use error::Error; pub(crate) use insert::InsertSplit; +use pgdog_config::RewriteMode; +//use pgdog_config::RewriteMode; use plan::GeneratedId; pub(crate) use plan::RewritePlan; pub(crate) use simple_prepared::PrepareExecute; @@ -43,6 +48,10 @@ pub(crate) struct StatementRewriteContext<'a> { pub(crate) user: &'a str, /// Search path for table lookups. pub(crate) search_path: Option<&'a ParameterValue>, + /// Timezone for now() time generation for TIMEZONE columns. + pub(crate) timezone: Option<&'a ParameterValue>, + /// Statement, and transaction DateTime relevant to the current Query (if not being cached) + pub(crate) query_timestamps: QueryTimestamps, } #[derive(Debug)] @@ -66,6 +75,10 @@ pub(crate) struct StatementRewrite<'a> { user: &'a str, /// Search path for table lookups. search_path: Option<&'a ParameterValue>, + /// Timezone for now() time generation for TIMEZONE columns. + timezone: Option<&'a ParameterValue>, + /// Statement, and transaction DateTime relevant to the current Query (if not being cached) + query_timestamps: QueryTimestamps, } impl<'a> StatementRewrite<'a> { @@ -83,6 +96,8 @@ impl<'a> StatementRewrite<'a> { db_schema: ctx.db_schema, user: ctx.user, search_path: ctx.search_path, + timezone: ctx.timezone, + query_timestamps: ctx.query_timestamps, } } @@ -93,6 +108,8 @@ impl<'a> StatementRewrite<'a> { db_schema: self.db_schema.clone(), user: self.user, search_path: self.search_path, + timezone: self.timezone, + query_timestamps: self.query_timestamps, } } @@ -133,7 +150,9 @@ impl<'a> StatementRewrite<'a> { // This must run BEFORE the unique_id rewriter so the injected // function calls get processed. match stmt.stmt_mut() { - NodeMut::InsertStmt(insert) => self.inject_auto_id(insert, mem, &mut plan)?, + NodeMut::InsertStmt(insert) => { + self.inject_auto_id(insert, mem, &mut plan)?; + } NodeMut::PrepareStmt(mut prepare) => { if let NodeMut::InsertStmt(insert) = prepare.query_mut() { self.inject_auto_id(insert, mem, &mut plan)?; @@ -152,8 +171,10 @@ impl<'a> StatementRewrite<'a> { Ok(Some(replacement)) => { plan.unique_ids += 1; if self.extended { - plan.generated_ids - .push(((next_param - 1) as u16, GeneratedId::UniqueId)); + plan.generated_params.push(GeneratedParam { + param_num: (next_param - 1) as u16, + generated_id: GeneratedId::UniqueId, + }); } self.rewritten = true; node.replace(replacement); @@ -185,8 +206,38 @@ impl<'a> StatementRewrite<'a> { self.limit_offset(&select, &mut plan); } + let timestamp_rewrite = !matches!( + config().config.rewrite.omni_non_deterministic_functions, + RewriteMode::Ignore + ); + + if timestamp_rewrite { + match stmt.stmt_mut() { + NodeMut::InsertStmt(_) => { + self.rewrite_timestamp_functions( + stmt.stmt_mut(), + mem, + &mut next_param, + &mut plan, + )?; + } + NodeMut::PrepareStmt(mut prepare) => { + if matches!(prepare.query_mut(), NodeMut::InsertStmt(_)) { + self.rewrite_timestamp_functions( + prepare.query_mut(), + mem, + &mut next_param, + &mut plan, + )?; + } + } + _ => {} + } + } + // Handle top-level PREPARE/EXECUTE statements. - let prepared_result = self.rewrite_simple_prepared(stmt.stmt_mut(), mem, &mut plan)?; + let prepared_result = + self.rewrite_simple_prepared(stmt.stmt_mut(), mem, &mut plan, timestamp_rewrite)?; if prepared_result.rewritten { self.rewritten = true; plan.prepare_rewrites = prepared_result.rewrites; diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/nextval.rs b/pgdog/src/frontend/router/parser/rewrite/statement/nextval.rs index 7a359cf62..25c8a62cc 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/nextval.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/nextval.rs @@ -3,6 +3,7 @@ use std::collections::HashMap; use pg_raw_parse::{ConstValue, Node, make, transform, walk}; use crate::frontend::router::parser::rewrite::ee; +use crate::frontend::router::parser::rewrite::statement::plan::GeneratedParam; use super::plan::GeneratedId; use super::{Error, RewritePlan, StatementRewrite}; @@ -118,8 +119,10 @@ impl StatementRewrite<'_> { let sequence = sequence_call(node)?; let param = *next_param; *next_param += 1; - plan.generated_ids - .push((param as u16, GeneratedId::Sequence(sequence))); + plan.generated_params.push(GeneratedParam { + generated_id: GeneratedId::Sequence(sequence), + param_num: param as u16, + }); // Retain simple-protocol SQL even when a sequence call is the only rewrite. self.rewritten = true; self.extended.then(|| { @@ -204,7 +207,9 @@ mod tests { use crate::backend::{ShardingSchema, schema::Schema}; use crate::frontend::ClientRequest; use crate::frontend::PreparedStatements; + use crate::frontend::client::QueryTimestamps; use crate::frontend::router::parser::StatementRewriteContext; + use crate::frontend::router::parser::rewrite::statement::plan::GeneratedParam; use crate::net::messages::bind::{Format, Parameter}; use crate::net::{Bind, Parse, ProtocolMessage, Query}; use pgdog_config::Rewrite; @@ -229,6 +234,8 @@ mod tests { db_schema: &db_schema, user: "test", search_path: None, + timezone: None, + query_timestamps: QueryTimestamps::default(), }); let mut plan = RewritePlan::default(); let ast = make::owned(|mem| { @@ -255,25 +262,32 @@ mod tests { assert_eq!(plan.params, 1); assert_eq!(plan.unique_ids, 1); assert_eq!( - plan.generated_ids, + plan.generated_params, vec![ - ( - 2, - GeneratedId::Sequence(SequenceCall::Nextval("sequence.name".to_owned())) - ), - (3, GeneratedId::UniqueId), - ( - 4, - GeneratedId::Sequence(SequenceCall::Currval("other.seq".to_owned())) - ), - ( - 5, - GeneratedId::Sequence(SequenceCall::Setval { + GeneratedParam { + param_num: 2, + generated_id: GeneratedId::Sequence(SequenceCall::Nextval( + "sequence.name".to_owned() + )) + }, + GeneratedParam { + param_num: 3, + generated_id: GeneratedId::UniqueId + }, + GeneratedParam { + param_num: 4, + generated_id: GeneratedId::Sequence(SequenceCall::Currval( + "other.seq".to_owned() + )) + }, + GeneratedParam { + param_num: 5, + generated_id: GeneratedId::Sequence(SequenceCall::Setval { name: "sequence.name".to_owned(), value: 42, is_called: false, }) - ), + } ] ); assert_eq!(plan.stmt.as_deref(), Some(sql.as_str())); @@ -286,16 +300,20 @@ mod tests { let (sql, plan) = rewrite(original, false); assert_eq!(sql, original); assert_eq!( - plan.generated_ids, + plan.generated_params, vec![ - ( - 1, - GeneratedId::Sequence(SequenceCall::Nextval("sequence.name".to_owned())) - ), - ( - 2, - GeneratedId::Sequence(SequenceCall::Nextval("sequence.name".to_owned())) - ), + GeneratedParam { + param_num: 1, + generated_id: GeneratedId::Sequence(SequenceCall::Nextval( + "sequence.name".to_owned() + )) + }, + GeneratedParam { + param_num: 2, + generated_id: GeneratedId::Sequence(SequenceCall::Nextval( + "sequence.name".to_owned() + )) + }, ] ); assert_eq!(plan.unique_ids, 0); @@ -313,18 +331,20 @@ mod tests { ); assert_eq!(sql, "INSERT INTO t (id) VALUES ($1::bigint), ($2::bigint)"); assert_eq!( - plan.generated_ids, + plan.generated_params, vec![ - ( - 1, - GeneratedId::Sequence(SequenceCall::Nextval( + GeneratedParam { + param_num: 1, + generated_id: GeneratedId::Sequence(SequenceCall::Nextval( "\"My Schema\".\"My Sequence\"".to_owned() )) - ), - ( - 2, - GeneratedId::Sequence(SequenceCall::Nextval("other.seq".to_owned())) - ), + }, + GeneratedParam { + param_num: 2, + generated_id: GeneratedId::Sequence(SequenceCall::Nextval( + "other.seq".to_owned() + )) + }, ] ); } @@ -463,14 +483,19 @@ mod tests { let mut bind = Bind::new_params_codes("stmt", &original_params, &codes); let mut calls = Vec::new(); let mut value = -2i64; - plan.apply_generated_ids(&mut bind, async |call: &SequenceCall| { - let SequenceCall::Nextval(name) = call else { - panic!("expected nextval"); - }; - calls.push(name.to_owned()); - value += 1; - Ok(value) - }) + plan.apply_generated_ids( + &mut bind, + None, + crate::frontend::client::QueryTimestamps::now(), + async |call: &SequenceCall| { + let SequenceCall::Nextval(name) = call else { + panic!("expected nextval"); + }; + calls.push(name.to_owned()); + value += 1; + Ok(value) + }, + ) .await .expect("sequence values appended"); @@ -532,9 +557,14 @@ mod tests { }; for expected in [2, 4] { let mut bind = Bind::default(); - plan.apply_generated_ids(&mut bind, &mut nextval) - .await - .expect("values"); + plan.apply_generated_ids( + &mut bind, + None, + crate::frontend::client::QueryTimestamps::now(), + &mut nextval, + ) + .await + .expect("values"); assert_eq!( bind.parameter(1) .expect("format") @@ -556,7 +586,11 @@ mod tests { let (_, plan) = rewrite(&format!("SELECT pgdog.{call}"), true); let mut request = ClientRequest::from(vec![ProtocolMessage::Bind(Bind::default())]); let error = plan - .apply(&mut request) + .apply( + &mut request, + None, + crate::frontend::client::QueryTimestamps::now(), + ) .await .expect_err("EE hook rejects sequence"); assert!(matches!(error, Error::Enterprise(ee::Error::EERequired))); @@ -577,7 +611,11 @@ mod tests { Parse::new_anonymous(&original), )]); extended_plan - .apply(&mut request) + .apply( + &mut request, + None, + crate::frontend::client::QueryTimestamps::now(), + ) .await .expect("prepare does not fetch"); @@ -585,7 +623,11 @@ mod tests { let mut request = ClientRequest::from(vec![ProtocolMessage::Query(Query::new(&original))]); let error = simple_plan - .apply(&mut request) + .apply( + &mut request, + None, + crate::frontend::client::QueryTimestamps::now(), + ) .await .expect_err("simple query calls the EE hook"); assert!(matches!(error, Error::Enterprise(ee::Error::EERequired))); @@ -631,8 +673,11 @@ mod tests { let original = format!("SELECT {call}"); let (sql, plan) = rewrite(&original, extended); assert_eq!( - plan.generated_ids, - [(1, GeneratedId::Sequence(expected.clone()))], + plan.generated_params, + [GeneratedParam { + param_num: 1, + generated_id: GeneratedId::Sequence(expected.clone()), + }], "{call}" ); let canonical = original.replace( @@ -711,9 +756,14 @@ mod tests { for first in [1, 44] { if extended { let mut bind = Bind::default(); - plan.apply_generated_ids(&mut bind, &mut execute) - .await - .expect("values appended"); + plan.apply_generated_ids( + &mut bind, + None, + crate::frontend::client::QueryTimestamps::now(), + &mut execute, + ) + .await + .expect("values appended"); assert_eq!(bind.params_raw().len(), 6); for (index, value) in [first, first, -42, -42, 42, 43].into_iter().enumerate() { assert_eq!( @@ -783,7 +833,7 @@ mod tests { ] { for extended in [false, true] { let (_, plan) = rewrite(&format!("SELECT {call}"), extended); - assert!(plan.generated_ids.is_empty(), "{call}"); + assert!(plan.generated_params.is_empty(), "{call}"); assert!(plan.is_empty(), "{call}"); } } @@ -834,7 +884,7 @@ mod tests { "{call}" ); let (_, plan) = rewrite(&format!("SELECT {call}"), true); - assert!(plan.generated_ids.is_empty(), "{call}"); + assert!(plan.generated_params.is_empty(), "{call}"); assert!(plan.is_empty(), "{call}"); } } diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/offset.rs b/pgdog/src/frontend/router/parser/rewrite/statement/offset.rs index 8ce84b0d2..166885593 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/offset.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/offset.rs @@ -328,6 +328,8 @@ mod tests { db_schema: &db_schema, user: "test", search_path: None, + timezone: None, + query_timestamps: QueryTimestamps::default(), }); let mut plan = RewritePlan::default(); rewrite.limit_offset( diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs b/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs index 94bac8bce..9f5bbb7ed 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/plan.rs @@ -1,8 +1,3 @@ -use crate::frontend::{ClientRequest, PreparedStatements}; -use crate::net::messages::bind::{Format, Parameter}; -use crate::net::{Bind, Parse, ProtocolMessage, Query}; -use crate::unique_id::UniqueId; - use super::super::ee; use super::insert::{build_resolved_split_requests, build_split_requests}; use super::nextval::SequenceCall; @@ -10,11 +5,26 @@ use super::offset::OffsetPlan; 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::{ClientRequest, PreparedStatements}; +use crate::net::messages::bind::{Format, Parameter}; +use crate::net::{Bind, Parse, ProtocolMessage, Query, parameter::ParameterValue}; +use crate::unique_id::UniqueId; +/// TODO: Docs. +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct GeneratedParam { + pub(crate) generated_id: GeneratedId, + pub(crate) param_num: u16, +} + +/// TODO: Document that this is also stored in PreparedStatement cache. #[derive(Debug, Clone, PartialEq, Eq)] pub(crate) enum GeneratedId { UniqueId, Sequence(SequenceCall), + ProxyTime(TimeFunction), } /// Statement rewrite plan. @@ -37,7 +47,8 @@ pub(crate) struct RewritePlan { /// One-based parameter indexes and ID sources in allocation order. /// Simple protocol records sequence calls here without using the indexes. - pub(crate) generated_ids: Vec<(u16, GeneratedId)>, + /// TODO: Document that this is also stored in PreparedStatement cache. + pub(crate) generated_params: Vec, /// Rewritten SQL statement. pub(crate) stmt: Option, @@ -87,58 +98,98 @@ impl RewritePlan { pub(crate) fn is_empty(&self) -> bool { self.unique_ids == 0 && self.auto_id_injected == 0 - && self.generated_ids.is_empty() + && self.generated_params.is_empty() && self.stmt.is_none() && self.prepare_rewrites.is_empty() && self.insert_split.is_empty() && self.aggregates.is_noop() && self.sharding_key_update.is_none() && self.offset.is_none() + // TODO: Check here. } /// Append generated unique IDs and sequence values to a Bind message. - async fn apply_bind(&self, bind: &mut Bind) -> Result<(), Error> { - self.apply_generated_ids(bind, SequenceCall::execute).await + async fn apply_bind( + &self, + bind: &mut Bind, + timezone: Option<&ParameterValue>, + timestamps: QueryTimestamps, + ) -> Result<(), Error> { + self.apply_generated_ids(bind, timezone, timestamps, SequenceCall::execute) + .await } /// Append values in the same order their placeholders were allocated. + /// + /// `timezone` = client's setting (or the database default when None) pub(super) async fn apply_generated_ids( &self, bind: &mut Bind, + timezone: Option<&ParameterValue>, + timestamps: QueryTimestamps, mut execute: impl AsyncFnMut(&SequenceCall) -> Result, ) -> Result<(), Error> { let format = bind.default_param_format(); - for (_, source) in &self.generated_ids { - let id = match source { - GeneratedId::UniqueId => UniqueId::generator()?.next_id(), - GeneratedId::Sequence(call) => execute(call).await?, - }; - let param = match format { - Format::Binary => Parameter::new(&id.to_be_bytes()), - Format::Text => Parameter::new(itoa::Buffer::new().format(id).as_bytes()), + for generated_param in &self.generated_params { + let source = &generated_param.generated_id; + let num = generated_param.param_num; + assert_eq!(bind.params_raw().len() + 1, num as usize); + + let param = match source { + GeneratedId::UniqueId => { + Self::convert_int_to_param(UniqueId::generator()?.next_id(), format) + } + GeneratedId::Sequence(call) => { + Self::convert_int_to_param(execute(call).await?, format) + } + GeneratedId::ProxyTime(time) => { + let (text, binary) = time.formatted_time(×tamps, timezone)?; + match format { + Format::Binary => Parameter::new(binary.as_slice()), + Format::Text => Parameter::new(text.as_bytes()), + } + } }; + bind.push_param(param, format); } Ok(()) } + fn convert_int_to_param(id: i64, format: Format) -> Parameter { + match format { + Format::Binary => Parameter::new(&id.to_be_bytes()), + Format::Text => Parameter::new(itoa::Buffer::new().format(id).as_bytes()), + } + } + /// Apply the rewrite plan to a Parse message by updating the SQL. - fn apply_parse(&self, parse: &mut Parse) { + /// + /// Returns the client's parameter count for an unnamed statement + fn apply_parse(&self, parse: &mut Parse) -> Option { if let Some(ref stmt) = self.stmt { + let client_params = self.params.max(parse.num_data_types()); + parse.set_query(stmt); if !parse.anonymous() { - PreparedStatements::global().write().rewrite(parse); + PreparedStatements::global() + .write() + .rewrite(parse, client_params); + } else { + return Some(client_params); } } + + None } /// Apply the rewrite plan to a Query message by updating the SQL. async fn apply_query(&self, query: &mut Query) -> Result<(), Error> { if self - .generated_ids + .generated_params .iter() - .any(|(_, source)| matches!(source, GeneratedId::Sequence(_))) + .any(|source| matches!(source.generated_id, GeneratedId::Sequence(_))) { if let Some(stmt) = self.rewrite_sequence_simple().await? { query.set_query(&stmt); @@ -151,7 +202,12 @@ impl RewritePlan { } /// Apply the rewrite plan to a ClientRequest. - pub(crate) async fn apply(&self, request: &mut ClientRequest) -> Result { + pub(crate) async fn apply( + &self, + request: &mut ClientRequest, + timezone: Option<&ParameterValue>, + timestamps: QueryTimestamps, + ) -> Result { // Prepend any required Prepare messages for EXECUTE statements. if !self.prepare_rewrites.is_empty() { self.prepare_rewrites @@ -169,15 +225,21 @@ impl RewritePlan { }); } + let mut anonymous_client_params = None; + for message in request.messages.iter_mut() { match message { - ProtocolMessage::Parse(parse) => self.apply_parse(parse), + ProtocolMessage::Parse(parse) => { + anonymous_client_params = self.apply_parse(parse); + } ProtocolMessage::Query(query) => self.apply_query(query).await?, - ProtocolMessage::Bind(bind) => self.apply_bind(bind).await?, + ProtocolMessage::Bind(bind) => self.apply_bind(bind, timezone, timestamps).await?, _ => {} } } + request.anonymous_client_params = anonymous_client_params; + self.apply_after_messages(request) } @@ -191,9 +253,9 @@ impl RewritePlan { // those since insert split will return the same row(s) as multi-tuple insert. if !self.insert_split.is_empty() && request.is_executable() { if self - .generated_ids + .generated_params .iter() - .any(|(_, source)| matches!(source, GeneratedId::Sequence(_))) + .any(|source| matches!(source.generated_id, GeneratedId::Sequence(_))) && let Some(query) = request.messages.iter().find_map(|message| match message { ProtocolMessage::Query(query) => Some(query), _ => None, @@ -246,7 +308,9 @@ mod tests { let _guard = set_env_var("NODE_ID", "pgdog-1"); let plan = RewritePlan::default(); let mut bind = Bind::default(); - plan.apply_bind(&mut bind).await.unwrap(); + plan.apply_bind(&mut bind, None, QueryTimestamps::now()) + .await + .unwrap(); assert_eq!(bind.params_raw().len(), 0); } @@ -255,11 +319,16 @@ mod tests { let _guard = set_env_var("NODE_ID", "pgdog-1"); let plan = RewritePlan { unique_ids: 1, - generated_ids: vec![(1, GeneratedId::UniqueId)], + generated_params: vec![GeneratedParam { + generated_id: GeneratedId::UniqueId, + param_num: 1, + }], ..Default::default() }; let mut bind = Bind::default(); - plan.apply_bind(&mut bind).await.unwrap(); + plan.apply_bind(&mut bind, None, QueryTimestamps::now()) + .await + .unwrap(); assert_eq!(bind.params_raw().len(), 1); // Default format is Text, so data should be a string @@ -277,13 +346,18 @@ mod tests { let plan = RewritePlan { params: 1, unique_ids: 1, - generated_ids: vec![(2, GeneratedId::UniqueId)], + generated_params: vec![GeneratedParam { + param_num: 2, + generated_id: GeneratedId::UniqueId, + }], ..Default::default() }; // Create bind with uniform binary format (1 code applies to all) let mut bind = Bind::new_params_codes("test", &[Parameter::new(b"existing")], &[Format::Binary]); - plan.apply_bind(&mut bind).await.unwrap(); + plan.apply_bind(&mut bind, None, QueryTimestamps::now()) + .await + .unwrap(); assert_eq!(bind.params_raw().len(), 2); // Should use binary format: 8 bytes big-endian @@ -303,7 +377,10 @@ mod tests { let plan = RewritePlan { params: 2, unique_ids: 1, - generated_ids: vec![(3, GeneratedId::UniqueId)], + generated_params: vec![GeneratedParam { + param_num: 3, + generated_id: GeneratedId::UniqueId, + }], ..Default::default() }; // Create bind with one-to-one format codes @@ -312,7 +389,9 @@ mod tests { &[Parameter::new(b"a"), Parameter::new(b"b")], &[Format::Binary, Format::Binary], ); - plan.apply_bind(&mut bind).await.unwrap(); + plan.apply_bind(&mut bind, None, QueryTimestamps::now()) + .await + .unwrap(); assert_eq!(bind.params_raw().len(), 3); // New param should be text (default for one-to-one) @@ -330,15 +409,26 @@ mod tests { let _guard = set_env_var("NODE_ID", "pgdog-1"); let plan = RewritePlan { unique_ids: 3, - generated_ids: vec![ - (1, GeneratedId::UniqueId), - (2, GeneratedId::UniqueId), - (3, GeneratedId::UniqueId), + generated_params: vec![ + GeneratedParam { + param_num: 1, + generated_id: GeneratedId::UniqueId, + }, + GeneratedParam { + param_num: 2, + generated_id: GeneratedId::UniqueId, + }, + GeneratedParam { + param_num: 3, + generated_id: GeneratedId::UniqueId, + }, ], ..Default::default() }; let mut bind = Bind::default(); - plan.apply_bind(&mut bind).await.unwrap(); + plan.apply_bind(&mut bind, None, QueryTimestamps::now()) + .await + .unwrap(); assert_eq!(bind.params_raw().len(), 3); let mut ids = HashSet::new(); @@ -356,14 +446,25 @@ mod tests { let plan = RewritePlan { params: 2, unique_ids: 2, - generated_ids: vec![(3, GeneratedId::UniqueId), (4, GeneratedId::UniqueId)], + generated_params: vec![ + GeneratedParam { + param_num: 3, + generated_id: GeneratedId::UniqueId, + }, + GeneratedParam { + param_num: 4, + generated_id: GeneratedId::UniqueId, + }, + ], ..Default::default() }; let mut bind = Bind::new_params( "test", &[Parameter::new(b"existing1"), Parameter::new(b"existing2")], ); - plan.apply_bind(&mut bind).await.unwrap(); + plan.apply_bind(&mut bind, None, QueryTimestamps::now()) + .await + .unwrap(); assert_eq!(bind.params_raw().len(), 4); assert_eq!(bind.params_raw()[0].data.as_ref(), b"existing1"); 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 495f74a71..2d7ee82db 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/simple_prepared.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/simple_prepared.rs @@ -7,10 +7,16 @@ use pg_raw_parse::{ use crate::{ frontend::{ - PreparedStatements, - router::parser::{Limit, rewrite::statement::offset::OffsetPlan}, + prepared_statements::PreparedPlan, + router::parser::{ + Limit, + rewrite::statement::{ + offset::OffsetPlan, + plan::{GeneratedId, GeneratedParam}, + }, + }, }, - net::{PREPARE_TEMPLATE_NAME, Prepare}, + net::{PREPARE_TEMPLATE_NAME, Prepare, parameter::ParameterValue}, unique_id::UniqueId, }; @@ -63,6 +69,7 @@ impl StatementRewrite<'_> { node: NodeMut<'a, '_>, mem: MemoryToken<'a>, plan: &mut RewritePlan, + timestamp_rewrite: bool, ) -> Result { let mut result = SimplePreparedResult::default(); @@ -70,7 +77,7 @@ impl StatementRewrite<'_> { return Ok(result); } - match rewrite_single_prepared(node, mem, self.prepared_statements, plan)? { + match self.rewrite_single_prepared(node, mem, plan, timestamp_rewrite)? { SimplePreparedRewrite::Prepared { prepare } => { result.rewrites.push(PrepareExecute::Prepare(prepare)); result.rewritten = true; @@ -84,75 +91,120 @@ impl StatementRewrite<'_> { Ok(result) } -} -/// Rewrites a single `PREPARE` or `EXECUTE` node. -fn rewrite_single_prepared<'a>( - node: NodeMut<'a, '_>, - mem: MemoryToken<'a>, - prepared_statements: &mut PreparedStatements, - plan: &mut RewritePlan, -) -> Result { - match node { - NodeMut::PrepareStmt(mut stmt) => { - let client_name = stmt.name().expect("prepare must have a name").to_owned(); - - // Create a globally unique key using the query text - // with a hardcoded name. - stmt.set_name(Some(mem.copy_string(PREPARE_TEMPLATE_NAME))); - - let original_query = Bytes::from(pg_raw_parse::deparse(&*stmt)?.as_str().to_owned()); - - // Is the query a SELECT? Do we have both LIMIT and OFFSET in the SELECT? - let offset_plan: Option = create_offset_plan(mem, &mut stmt); - - let new_query = offset_plan - .as_ref() - .map(|_| { - pg_raw_parse::deparse(&*stmt) - .map(|deparse_result| Bytes::from(deparse_result.as_str().to_owned())) - }) - .transpose()?; - - let prepare = prepared_statements.insert_prepare( - &client_name, - original_query, - new_query, - plan, - offset_plan, - ); + /// Rewrites a single `PREPARE` or `EXECUTE` node. + fn rewrite_single_prepared<'a>( + &mut self, + node: NodeMut<'a, '_>, + mem: MemoryToken<'a>, + plan: &mut RewritePlan, + timestamp_rewrite: bool, + ) -> Result { + match node { + NodeMut::PrepareStmt(mut stmt) => { + let client_name = stmt.name().expect("prepare must have a name").to_owned(); + + // Create a globally unique key using the query text + // with a hardcoded name. + stmt.set_name(Some(mem.copy_string(PREPARE_TEMPLATE_NAME))); + + let original_query = + Bytes::from(pg_raw_parse::deparse(&*stmt)?.as_str().to_owned()); + + // Is the query a SELECT? Do we have both LIMIT and OFFSET in the SELECT? + let offset_plan: Option = create_offset_plan(mem, &mut stmt); + + let new_query = offset_plan + .as_ref() + .map(|_| { + pg_raw_parse::deparse(&*stmt) + .map(|deparse_result| Bytes::from(deparse_result.as_str().to_owned())) + }) + .transpose()?; + + let generated_params = plan.generated_params.clone(); + let prepare = self.prepared_statements.insert_prepare( + &client_name, + original_query, + new_query, + plan, + offset_plan, + generated_params, + ); - stmt.set_name(Some(mem.copy_string(prepare.name()))); + stmt.set_name(Some(mem.copy_string(prepare.name()))); - Ok(SimplePreparedRewrite::Prepared { prepare }) - } + Ok(SimplePreparedRewrite::Prepared { prepare }) + } - NodeMut::ExecuteStmt(mut stmt) => { - let stmt_name = stmt.name().expect("EXECUTE always has name"); + NodeMut::ExecuteStmt(mut stmt) => { + let stmt_name = stmt.name().expect("EXECUTE always has name"); + + if let Some(PreparedPlan { + prepare, + unique_ids, + offset_plan, + generated_params, + }) = self.prepared_statements.prepared_plan(stmt_name) + { + if let Some(mut offset_plan) = offset_plan { + // Note: This needs to be ordered before the offset_val/limit_val adjustment. + insert_offset_params(&mut stmt, mem, &offset_plan); + update_offset_plan_fields(&mut offset_plan, &mut stmt)?; + + plan.offset = Some(offset_plan); + } + + // TODO: Should we be setting this on Plan? Pros? Cons? + // TODO: Double check that this only runs on omnisharded (as well as Bind/Execute, etc) + if timestamp_rewrite { + plan.generated_params = generated_params; + self.insert_generated_ids( + &mut stmt, + mem, + &plan.generated_params, + self.timezone, + )?; + } + + // Rewrite EXECUTE statement to match the rewrite + // we did on the PREPARE statement. + insert_unique_ids(&mut stmt, mem, unique_ids)?; + + stmt.set_name(Some(mem.copy_string(prepare.name()))); + Ok(SimplePreparedRewrite::Executed { prepare }) + } else { + Err(Error::ExecuteMissingPrepare(stmt_name.to_owned())) + } + } - if let Some((prepare, unique_ids, offset_plan)) = - prepared_statements.prepare_and_unique_ids(stmt_name) - { - if let Some(mut offset_plan) = offset_plan { - // Note: This needs to be ordered before the offset_val/limit_val adjustment. - insert_offset_params(&mut stmt, mem, &offset_plan); - update_offset_plan_fields(&mut offset_plan, &mut stmt)?; + _ => Ok(SimplePreparedRewrite::None), + } + } - plan.offset = Some(offset_plan); + fn insert_generated_ids<'a>( + &self, + stmt: &mut ExecuteStmtMut<'a, '_>, + mem: MemoryToken<'a>, + generated_params: &Vec, + timezone: Option<&ParameterValue>, + ) -> Result<(), Error> { + for param in generated_params { + let (text, _) = match ¶m.generated_id { + GeneratedId::ProxyTime(time) => { + time.formatted_time(&self.query_timestamps, timezone)? } + // TODO: It seems very straightforward to support the rest (if we want to support them for PREPARE) + _ => continue, + }; - // Rewrite EXECUTE statement to match the rewrite - // we did on the PREPARE statement. - insert_unique_ids(&mut stmt, mem, unique_ids)?; - - stmt.set_name(Some(mem.copy_string(prepare.name()))); - Ok(SimplePreparedRewrite::Executed { prepare }) - } else { - Err(Error::ExecuteMissingPrepare(stmt_name.to_owned())) - } + stmt.params_mut().push( + mem, + mem.make_a_const(ConstValue::String(text.as_str())).uncast(), + ); } - _ => Ok(SimplePreparedRewrite::None), + Ok(()) } } @@ -379,6 +431,8 @@ mod tests { use crate::backend::ShardingSchema; use crate::backend::schema::Schema; use crate::config::PreparedStatementsLevel; + use crate::frontend::PreparedStatements; + use crate::frontend::client::QueryTimestamps; use crate::test_utils::set_env_var; use pg_raw_parse::Node; use pgdog_config::Rewrite; @@ -418,6 +472,8 @@ mod tests { db_schema: &self.db_schema, user: "", search_path: None, + timezone: None, + query_timestamps: QueryTimestamps::default(), }); let mut plan = Default::default(); let ast = pg_raw_parse::make::try_owned(|mem| { @@ -541,9 +597,9 @@ mod tests { assert_eq!(ctx.ps.global.read().len(), 1); // Verify the OffsetPlan is correct from the PreparedStatement name used. - let (fetched_prepare, _, offset_plan) = - ctx.ps.prepare_and_unique_ids("test_stmt").unwrap(); - let offset_plan = offset_plan.unwrap(); + let fetched = ctx.ps.prepared_plan("test_stmt").unwrap(); + let fetched_prepare = fetched.prepare; + let offset_plan = fetched.offset_plan.unwrap(); assert_eq!( fetched_prepare.query, "PREPARE __pgdog_template_name AS SELECT * FROM sharded LIMIT $1 OFFSET $2" @@ -600,9 +656,9 @@ mod tests { assert_eq!(ctx.ps.global.read().len(), 2); // Verify the OffsetPlan is correct using the PreparedStatement name used. - let (fetched_prepare, _, offset_plan) = - ctx.ps.prepare_and_unique_ids("test_stmt2").unwrap(); - let offset_plan = offset_plan.unwrap(); + let fetched = ctx.ps.prepared_plan("test_stmt2").unwrap(); + let fetched_prepare = fetched.prepare; + let offset_plan = fetched.offset_plan.unwrap(); assert_eq!( fetched_prepare.query, "PREPARE __pgdog_template_name AS SELECT * FROM sharded LIMIT $1 OFFSET $2" diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/timestamp.rs b/pgdog/src/frontend/router/parser/rewrite/statement/timestamp.rs new file mode 100644 index 000000000..356af71d2 --- /dev/null +++ b/pgdog/src/frontend/router/parser/rewrite/statement/timestamp.rs @@ -0,0 +1,651 @@ +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() + } + } +} diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/unique_id.rs b/pgdog/src/frontend/router/parser/rewrite/statement/unique_id.rs index 7e6db6b23..517b9050b 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/unique_id.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/unique_id.rs @@ -58,12 +58,12 @@ mod tests { use pgdog_config::Rewrite; use super::*; - use crate::backend::ShardingSchema; use crate::backend::schema::Schema; use crate::frontend::PreparedStatements; use crate::frontend::router::parser::StatementRewriteContext; use crate::frontend::router::parser::rewrite::statement::RewritePlan; use crate::test_utils::set_env_var; + use crate::{backend::ShardingSchema, frontend::client::QueryTimestamps}; use pg_raw_parse::{Owned, nodes}; fn default_schema() -> ShardingSchema { @@ -288,6 +288,8 @@ mod tests { db_schema: &db_schema, user: "", search_path: None, + timezone: None, + query_timestamps: QueryTimestamps::default(), }); let mut plan = Default::default(); let ast = make::owned(|mem| { diff --git a/pgdog/src/frontend/router/parser/rewrite/statement/update.rs b/pgdog/src/frontend/router/parser/rewrite/statement/update.rs index bc8e76b27..5900f2b5a 100644 --- a/pgdog/src/frontend/router/parser/rewrite/statement/update.rs +++ b/pgdog/src/frontend/router/parser/rewrite/statement/update.rs @@ -468,6 +468,8 @@ mod test { prepared_statements: &mut stmts, user: "", search_path: None, + timezone: None, + query_timestamps: QueryTimestamps::default(), }; let mut plan = RewritePlan::default(); StatementRewrite::new(ctx).sharding_key_update( diff --git a/pgdog/src/net/messages/parameter_description.rs b/pgdog/src/net/messages/parameter_description.rs index f329e2c38..0d9f6e7a4 100644 --- a/pgdog/src/net/messages/parameter_description.rs +++ b/pgdog/src/net/messages/parameter_description.rs @@ -44,6 +44,11 @@ impl ParameterDescription { Self { params: Vec::new() } } + /// Keep only the first `len` parameters. + pub(crate) fn truncate(&mut self, len: usize) { + self.params.truncate(len); + } + pub(crate) fn rewrite_data_types(&mut self, mapping: &HashMap) { for param in &mut self.params { if let Some(&canonical) = mapping.get(&(*param as u32)) { @@ -58,6 +63,10 @@ mod test { use super::*; impl ParameterDescription { + pub(crate) fn new(params: Vec) -> Self { + Self { params } + } + /// Type OIDs of the parameters, in order. pub(crate) fn params(&self) -> &[i32] { &self.params diff --git a/pgdog/src/net/messages/parse.rs b/pgdog/src/net/messages/parse.rs index aec7ad895..4c26441ef 100644 --- a/pgdog/src/net/messages/parse.rs +++ b/pgdog/src/net/messages/parse.rs @@ -112,6 +112,14 @@ impl Parse { self.data_types.clone() } + /// Number of parameter data types the client declared. + pub(crate) fn num_data_types(&self) -> u16 { + self.data_types + .get(..2) + .map(|count| u16::from_be_bytes([count[0], count[1]])) + .unwrap_or_default() + } + /// Update the SQL for this prepared statement. pub(crate) fn set_query(&mut self, query: &str) { self.query = c_string_bytes(query);