diff --git a/pgdog/src/backend/pool/cleanup.rs b/pgdog/src/backend/pool/cleanup.rs index df077a4ca..57e5bd343 100644 --- a/pgdog/src/backend/pool/cleanup.rs +++ b/pgdog/src/backend/pool/cleanup.rs @@ -13,9 +13,12 @@ static PREPARED: Lazy> = Lazy::new(|| vec![Query::new("DEALLOCATE ALL /// static DIRTY: Lazy> = Lazy::new(|| { vec![ - Query::new("RESET ALL"), // Reset all parameters. + // RESET ALL deliberately leaves role and session_authorization alone. + // Resetting session authorization also clears the active role. + Query::new("SET SESSION AUTHORIZATION DEFAULT"), + Query::new("RESET ALL"), // Reset all other parameters. Query::new("SELECT pg_advisory_unlock_all()"), // Remove all advisory locks. - Query::new("DISCARD TEMP"), // Drop all temporary tables. + Query::new("DISCARD TEMP"), // Drop all temporary tables. ] }); @@ -136,3 +139,28 @@ impl Cleanup { self.deallocate } } + +#[cfg(test)] +mod test { + use super::*; + + #[test] + fn dirty_cleanup_resets_session_identity_before_other_state() { + let cleanup = Cleanup::parameters(); + let queries = cleanup + .queries() + .iter() + .map(Query::query) + .collect::>(); + + assert_eq!( + queries, + [ + "SET SESSION AUTHORIZATION DEFAULT", + "RESET ALL", + "SELECT pg_advisory_unlock_all()", + "DISCARD TEMP", + ] + ); + } +} diff --git a/pgdog/src/backend/pool/connection/binding.rs b/pgdog/src/backend/pool/connection/binding.rs index 4c32e8868..edc52963f 100644 --- a/pgdog/src/backend/pool/connection/binding.rs +++ b/pgdog/src/backend/pool/connection/binding.rs @@ -432,9 +432,7 @@ impl Binding { let mut max = 0; for result in results { let synced = result?; - if max < synced { - max = synced; - } + max = max.max(synced); } Ok(max) } @@ -454,6 +452,18 @@ impl Binding { } } + pub(crate) fn sync_client_params(&mut self, params: &Parameters) { + match self { + Binding::Direct(server, ..) => server.sync_client_params(params), + Binding::MultiShard(servers, _) => { + for server in servers { + server.sync_client_params(params); + } + } + _ => (), + } + } + pub(crate) fn changed_params(&mut self) -> Parameters { match self { Binding::Direct(server, ..) => server.changed_params().clone(), diff --git a/pgdog/src/backend/pool/guard.rs b/pgdog/src/backend/pool/guard.rs index 26765be12..46f81b226 100644 --- a/pgdog/src/backend/pool/guard.rs +++ b/pgdog/src/backend/pool/guard.rs @@ -113,9 +113,27 @@ mod test { }, server::test::test_server, }, - net::{Describe, Flush, Parse, Protocol, ProtocolMessage, Query, Sync}, + net::{DataRow, Describe, Flush, Parse, Protocol, ProtocolMessage, Query, Sync}, }; + async fn identity(guard: &mut Guard) -> (String, String) { + let messages = guard + .execute("SELECT session_user, current_user") + .await + .unwrap(); + let row: DataRow = messages + .into_iter() + .find(|message| message.code() == 'D') + .expect("identity query should return one row") + .try_into() + .unwrap(); + + ( + row.get_text(0).expect("session_user"), + row.get_text(1).expect("current_user"), + ) + } + #[tokio::test] async fn test_cleanup_dirty() { crate::logger(); @@ -176,6 +194,49 @@ mod test { drop(guard); } + #[tokio::test] + async fn test_cleanup_dirty_resets_session_identity() { + crate::logger(); + let pool = pool(); + let mut guard = pool.get(&Request::default()).await.unwrap(); + let server_id = guard.id(); + let role = format!("pgdog_cleanup_role_{}", std::process::id()); + + guard + .execute_checked(format!("CREATE ROLE {role}")) + .await + .unwrap(); + guard + .execute_checked(format!("SET ROLE {role}")) + .await + .unwrap(); + assert_eq!(identity(&mut guard).await, ("pgdog".into(), role.clone())); + + guard.mark_dirty(true); + drop(guard); + + let mut guard = pool.get(&Request::default()).await.unwrap(); + assert_eq!(guard.id(), server_id); + assert_eq!(identity(&mut guard).await, ("pgdog".into(), "pgdog".into())); + + guard + .execute_checked(format!("SET SESSION AUTHORIZATION {role}")) + .await + .unwrap(); + assert_eq!(identity(&mut guard).await, (role.clone(), role.clone())); + + guard.mark_dirty(true); + drop(guard); + + let mut guard = pool.get(&Request::default()).await.unwrap(); + assert_eq!(guard.id(), server_id); + assert_eq!(identity(&mut guard).await, ("pgdog".into(), "pgdog".into())); + guard + .execute_checked(format!("DROP ROLE {role}")) + .await + .unwrap(); + } + #[tokio::test] async fn test_cleanup_prepared_statements() { crate::logger(); diff --git a/pgdog/src/backend/server.rs b/pgdog/src/backend/server.rs index 6e654d792..c2acdd7a7 100644 --- a/pgdog/src/backend/server.rs +++ b/pgdog/src/backend/server.rs @@ -684,6 +684,7 @@ impl Server { "DISCARD ALL" => { self.prepared_statements.clear(); self.client_params.clear(); + self.client_params.clear_session_identity(); } "RESET" => self.client_params.clear(), // Someone reset params, we're gonna need to re-sync. _ => (), @@ -729,7 +730,7 @@ impl Server { // Construct client parameter SET queries. let tracked = params.tracked_and_different(&self.client_params); // Construct RESET queries to reset any current params - // to their default values. + // to their default values. Session identity is included here. let mut queries = self.client_params.reset_queries(params); // Combine both to create a new, fresh session state @@ -1136,6 +1137,16 @@ impl Server { pub(crate) fn reset_params(&mut self) { self.client_params.clear(); + self.client_params.clear_session_identity(); + } + + /// Record the client's current parameters on this backend. + /// + /// Used after a connected `SET` so the next checkout can diff against + /// the identity and GUCs this connection actually has. + pub(crate) fn sync_client_params(&mut self, params: &Parameters) { + self.client_params = params.tracked(); + self.client_params.copy_in_transaction(params); } pub(crate) fn reset_re_synced(&mut self) { @@ -2470,6 +2481,7 @@ pub(crate) mod test { .link_client(FrontendPid::new(), ¶ms, None) .await?; assert_eq!(changed, 1); + assert!(!server.dirty()); let changed = server .link_client(FrontendPid::new(), ¶ms, None) @@ -2494,6 +2506,52 @@ pub(crate) mod test { Ok(()) } + #[tokio::test] + async fn test_link_client_reconciles_session_identity() -> Result<(), Box> + { + let mut params = Parameters::default(); + params.insert_identity("role", &"pgdog".into(), false, false); + + let mut server = test_server().await; + assert!(!server.dirty()); + + let changed = server + .link_client(FrontendPid::new(), ¶ms, None) + .await?; + + assert_eq!(changed, 2); + assert!(!server.dirty()); + + let changed = server + .link_client(FrontendPid::new(), ¶ms, None) + .await?; + assert_eq!(changed, 0, "identity stays on the server snapshot"); + + let peer = Parameters::default(); + let changed = server.link_client(FrontendPid::new(), &peer, None).await?; + assert_eq!(changed, 1, "next client resets identity from the snapshot"); + + Ok(()) + } + + #[tokio::test] + async fn test_link_client_reconciles_transaction_identity() + -> Result<(), Box> { + let mut params = Parameters::default(); + params.insert_identity("role", &"pgdog".into(), true, false); + + let mut server = test_server().await; + let changed = server + .link_client(FrontendPid::new(), ¶ms, Some("BEGIN")) + .await?; + + assert_eq!(changed, 2); + assert!(!server.dirty()); + server.rollback().await?; + + Ok(()) + } + #[tokio::test] async fn test_copy_protocol() { let mut server = test_server().await; diff --git a/pgdog/src/frontend/client/query_engine/mod.rs b/pgdog/src/frontend/client/query_engine/mod.rs index 55f7e33a2..27573fb87 100644 --- a/pgdog/src/frontend/client/query_engine/mod.rs +++ b/pgdog/src/frontend/client/query_engine/mod.rs @@ -75,6 +75,7 @@ pub(crate) struct QueryEngine { // or disconnect. manual_lock: bool, temp_tables: TempTables, + last_server_error: bool, } impl QueryEngine { @@ -103,6 +104,7 @@ impl QueryEngine { advisory_locks: AdvisoryLocks::default(), manual_lock: false, temp_tables: Default::default(), + last_server_error: false, }) } @@ -140,6 +142,7 @@ impl QueryEngine { return Ok(result); } + self.last_server_error = false; self.stats.received(client_request.total_message_len()); self.set_state(State::Active); // Client is active. diff --git a/pgdog/src/frontend/client/query_engine/query.rs b/pgdog/src/frontend/client/query_engine/query.rs index d86ecc8cd..dccc48cc7 100644 --- a/pgdog/src/frontend/client/query_engine/query.rs +++ b/pgdog/src/frontend/client/query_engine/query.rs @@ -195,6 +195,7 @@ impl QueryEngine { } if code == 'E' { + self.last_server_error = true; if let Some(state) = self.pending_explain.as_mut() { state.annotated = true; } diff --git a/pgdog/src/frontend/client/query_engine/set.rs b/pgdog/src/frontend/client/query_engine/set.rs index e046f470e..71e646c6f 100644 --- a/pgdog/src/frontend/client/query_engine/set.rs +++ b/pgdog/src/frontend/client/query_engine/set.rs @@ -34,7 +34,14 @@ impl QueryEngine { let is_pin = param.name == PGDOG_PIN; if let Some(value) = param.value.clone() { - if context.in_transaction() { + if context.params.insert_identity( + ¶m.name, + &value, + context.in_transaction(), + param.local, + ) { + continue; + } else if context.in_transaction() { context .params .insert_transaction(¶m.name, value, param.local); @@ -51,7 +58,13 @@ impl QueryEngine { } } else { fake_command = "RESET"; - context.params.reset(¶m.name); + if !context.params.reset_identity( + ¶m.name, + context.in_transaction(), + param.local, + ) { + context.params.reset(¶m.name); + } if is_pin { self.manual_lock = false; } @@ -64,6 +77,9 @@ impl QueryEngine { if self.backend.connected() { self.execute(context, client_request, None).await?; + if !self.last_server_error { + self.backend.sync_client_params(context.params); + } } else { let fake_response = set_config .then(|| params.iter().map(|p| p.value.as_ref())) @@ -119,7 +135,9 @@ impl QueryEngine { if context.in_transaction() || self.backend.connected() { context.params.reset_all(); } else { - context.params.restore_startup(context.startup_params); + context + .params + .restore_startup_parameters(context.startup_params); self.comms.update_params(context.params); } diff --git a/pgdog/src/frontend/client/query_engine/test/mod.rs b/pgdog/src/frontend/client/query_engine/test/mod.rs index 7c3aca9c7..5956f1d5b 100644 --- a/pgdog/src/frontend/client/query_engine/test/mod.rs +++ b/pgdog/src/frontend/client/query_engine/test/mod.rs @@ -36,6 +36,7 @@ mod rewrite_offset; mod rewrite_projection; mod rewrite_simple_prepared; mod schema_changed; +mod session_identity; mod set; mod set_schema_sharding; mod sharded; diff --git a/pgdog/src/frontend/client/query_engine/test/schema_changed.rs b/pgdog/src/frontend/client/query_engine/test/schema_changed.rs index e77cb1254..3aaf50ceb 100644 --- a/pgdog/src/frontend/client/query_engine/test/schema_changed.rs +++ b/pgdog/src/frontend/client/query_engine/test/schema_changed.rs @@ -14,6 +14,19 @@ fn get_pool_ids(engine: &mut QueryEngine) -> Vec { /// Helper to run DDL in a transaction and verify schema_changed is set and reload occurs. async fn assert_ddl_sets_schema_changed(ddl: &str, expected_cmd: &str, cleanup: Option<&str>) { + Box::pin(assert_ddl_sets_schema_changed_inner( + ddl, + expected_cmd, + cleanup, + )) + .await +} + +async fn assert_ddl_sets_schema_changed_inner( + ddl: &str, + expected_cmd: &str, + cleanup: Option<&str>, +) { let mut test_client = TestClient::new_sharded(Parameters::default()).await; // Capture pool IDs before DDL diff --git a/pgdog/src/frontend/client/query_engine/test/session_identity.rs b/pgdog/src/frontend/client/query_engine/test/session_identity.rs new file mode 100644 index 000000000..be28b40c8 --- /dev/null +++ b/pgdog/src/frontend/client/query_engine/test/session_identity.rs @@ -0,0 +1,292 @@ +use crate::{ + backend::databases::reload_from_existing, + config::{config, load_test, set}, + expect_message, + net::{CommandComplete, DataRow, ErrorResponse, ReadyForQuery, RowDescription}, +}; + +use super::prelude::*; + +fn load_single_connection_test_pool() { + load_test(); + + let mut config = (*config()).clone(); + config.config.general.default_pool_size = 1; + config.config.general.min_pool_size = 0; + set(config).unwrap(); + reload_from_existing().unwrap(); +} + +async fn identity(client: &mut TestClient) -> (i64, String, String) { + client + .send_simple(Query::new( + "SELECT pg_backend_pid(), session_user, current_user", + )) + .await; + expect_message!(client.read().await, RowDescription); + let row = expect_message!(client.read().await, DataRow); + expect_message!(client.read().await, CommandComplete); + expect_message!(client.read().await, ReadyForQuery); + + ( + row.get_int(0, true).expect("backend pid"), + row.get_text(1).expect("session_user"), + row.get_text(2).expect("current_user"), + ) +} + +async fn create_role(client: &mut TestClient, suffix: &str) -> String { + let role = format!("pgdog_cleanup_role_{}_{suffix}", std::process::id()); + client + .send_simple(Query::new(format!("CREATE ROLE {role}"))) + .await; + client.read_until('Z').await.unwrap(); + role +} + +async fn drop_role(client: &mut TestClient, role: &str) { + client + .send_simple(Query::new(format!("DROP ROLE {role}"))) + .await; + client.read_until('Z').await.unwrap(); +} + +fn assert_default_identity(identity: &(i64, String, String)) { + assert_eq!( + (identity.1.as_str(), identity.2.as_str()), + ("pgdog", "pgdog") + ); +} + +#[tokio::test] +async fn test_session_lock_cleanup_resets_role_before_backend_reuse() { + load_single_connection_test_pool(); + let mut source = TestClient::new(Parameters::default()).await; + let role = create_role(&mut source, "advisory").await; + + source + .send_simple(Query::new("SELECT pg_advisory_lock(707519)")) + .await; + source.read_until('Z').await.unwrap(); + source + .send_simple(Query::new(format!("SET ROLE {role}"))) + .await; + source.read_until('Z').await.unwrap(); + + let before = identity(&mut source).await; + assert_eq!( + (before.1.as_str(), before.2.as_str()), + ("pgdog", role.as_str()) + ); + drop(source.leak_pool()); + + let mut peer = TestClient::new(Parameters::default()).await; + let after = identity(&mut peer).await; + assert_eq!(after.0, before.0, "physical backend should be reused"); + assert_default_identity(&after); + drop_role(&mut peer, &role).await; +} + +#[tokio::test] +async fn test_deferred_role_reconciled_before_backend_reuse() { + load_single_connection_test_pool(); + let mut source = TestClient::new(Parameters::default()).await; + let role = create_role(&mut source, "deferred").await; + + source + .send_simple(Query::new(format!("SET ROLE {role}"))) + .await; + source.read_until('Z').await.unwrap(); + assert!( + !source.backend_connected(), + "SET ROLE should be deferred until a query checks out a backend" + ); + + let before = identity(&mut source).await; + assert_eq!( + (before.1.as_str(), before.2.as_str()), + ("pgdog", role.as_str()) + ); + assert!(!source.backend_connected()); + + source.send_simple(Query::new("RESET ALL")).await; + source.read_until('Z').await.unwrap(); + assert_eq!( + source.client().params.session_identity(false).role(), + Some(role.as_str()), + "RESET ALL must preserve the frontend's logical role" + ); + let after_reset_all = identity(&mut source).await; + assert_eq!( + after_reset_all, before, + "PostgreSQL RESET ALL must preserve SET ROLE" + ); + assert!(!source.backend_connected()); + drop(source.leak_pool()); + + let mut peer = TestClient::new(Parameters::default()).await; + let after = identity(&mut peer).await; + assert_eq!(after.0, before.0, "physical backend should be reused"); + assert_default_identity(&after); + drop_role(&mut peer, &role).await; +} + +#[tokio::test] +async fn test_session_authorization_reapplied_and_reconciled_before_backend_reuse() { + load_single_connection_test_pool(); + let mut source = TestClient::new(Parameters::default()).await; + let role = create_role(&mut source, "authorization").await; + + source + .send_simple(Query::new(format!("SET SESSION AUTHORIZATION {role}"))) + .await; + source.read_until('Z').await.unwrap(); + assert!(!source.backend_connected()); + + let before = identity(&mut source).await; + assert_eq!( + (before.1.as_str(), before.2.as_str()), + (role.as_str(), role.as_str()) + ); + assert!(!source.backend_connected()); + + let reapplied = identity(&mut source).await; + assert_eq!(reapplied, before); + assert!(!source.backend_connected()); + drop(source.leak_pool()); + + let mut peer = TestClient::new(Parameters::default()).await; + let after = identity(&mut peer).await; + assert_eq!(after.0, before.0, "physical backend should be reused"); + assert_default_identity(&after); + drop_role(&mut peer, &role).await; +} + +#[tokio::test] +async fn test_deferred_transaction_role_reconciled_before_backend_reuse() { + load_single_connection_test_pool(); + let mut source = TestClient::new(Parameters::default()).await; + let role = create_role(&mut source, "transaction").await; + + source.send_simple(Query::new("BEGIN")).await; + source.read_until('Z').await.unwrap(); + source + .send_simple(Query::new(format!("SET ROLE {role}"))) + .await; + source.read_until('Z').await.unwrap(); + assert!( + !source.backend_connected(), + "BEGIN and SET ROLE should remain deferred until the first query" + ); + + let before = identity(&mut source).await; + assert_eq!( + (before.1.as_str(), before.2.as_str()), + ("pgdog", role.as_str()) + ); + assert!(source.backend_connected()); + + source.send_simple(Query::new("COMMIT")).await; + source.read_until('Z').await.unwrap(); + assert!(!source.backend_connected()); + + let reapplied = identity(&mut source).await; + assert_eq!( + (reapplied.1.as_str(), reapplied.2.as_str()), + ("pgdog", role.as_str()) + ); + assert!(!source.backend_connected()); + drop(source.leak_pool()); + + let mut peer = TestClient::new(Parameters::default()).await; + let after = identity(&mut peer).await; + assert_eq!(after.0, before.0, "physical backend should be reused"); + assert_default_identity(&after); + drop_role(&mut peer, &role).await; +} + +#[tokio::test] +async fn test_transaction_role_rollback_restores_identity() { + load_single_connection_test_pool(); + let mut client = TestClient::new(Parameters::default()).await; + let role = create_role(&mut client, "rollback").await; + + client.send_simple(Query::new("BEGIN")).await; + client.read_until('Z').await.unwrap(); + let original = identity(&mut client).await; + + client + .send_simple(Query::new(format!("SET ROLE {role}"))) + .await; + client.read_until('Z').await.unwrap(); + let changed = identity(&mut client).await; + assert_eq!(changed.0, original.0); + assert_eq!( + (changed.1.as_str(), changed.2.as_str()), + ("pgdog", role.as_str()) + ); + + client.send_simple(Query::new("ROLLBACK")).await; + client.read_until('Z').await.unwrap(); + let restored = identity(&mut client).await; + assert_eq!(restored.0, original.0); + assert_default_identity(&restored); + drop_role(&mut client, &role).await; +} + +#[tokio::test] +async fn test_rejected_connected_role_restores_identity_on_rollback() { + load_single_connection_test_pool(); + let mut client = TestClient::new(Parameters::default()).await; + + client.send_simple(Query::new("BEGIN")).await; + client.read_until('Z').await.unwrap(); + let before = identity(&mut client).await; + + client + .send_simple(Query::new("SET ROLE pgdog_role_does_not_exist")) + .await; + expect_message!(client.read().await, ErrorResponse); + expect_message!(client.read().await, ReadyForQuery); + + client.send_simple(Query::new("ROLLBACK")).await; + client.read_until('Z').await.unwrap(); + let after = identity(&mut client).await; + assert_eq!(after.0, before.0); + assert_default_identity(&after); +} + +#[tokio::test] +async fn test_connected_transaction_role_reconciled_before_backend_reuse() { + load_single_connection_test_pool(); + let mut source = TestClient::new(Parameters::default()).await; + let role = create_role(&mut source, "connected").await; + + source.send_simple(Query::new("BEGIN")).await; + source.read_until('Z').await.unwrap(); + + let connected = identity(&mut source).await; + assert!(source.backend_connected()); + + source + .send_simple(Query::new(format!("SET ROLE {role}"))) + .await; + source.read_until('Z').await.unwrap(); + let changed = identity(&mut source).await; + assert_eq!(changed.0, connected.0); + assert_eq!( + (changed.1.as_str(), changed.2.as_str()), + ("pgdog", role.as_str()) + ); + + source.send_simple(Query::new("COMMIT")).await; + source.read_until('Z').await.unwrap(); + assert!(!source.backend_connected()); + drop(source.leak_pool()); + + let mut peer = TestClient::new(Parameters::default()).await; + let after = identity(&mut peer).await; + assert_eq!(after.0, connected.0, "physical backend should be reused"); + assert_default_identity(&after); + drop_role(&mut peer, &role).await; +} diff --git a/pgdog/src/frontend/router/parser/query/set.rs b/pgdog/src/frontend/router/parser/query/set.rs index 923da72a7..83637d1cf 100644 --- a/pgdog/src/frontend/router/parser/query/set.rs +++ b/pgdog/src/frontend/router/parser/query/set.rs @@ -23,7 +23,7 @@ impl QueryParser { .with_read(context.read_only), )) } else { - let param = Self::parse_set_param(stmt)?; + let param = Self::parse_set_param(stmt, context.query()?.query())?; Ok(Command::Set { params: vec![param], route: Route::write(context.shards_calculator.shard()), @@ -33,25 +33,38 @@ impl QueryParser { } /// Parse a single SET statement into a SetParam - fn parse_set_param(stmt: &nodes::VariableSetStmt) -> Result { - let value = if stmt.kind == VAR_SET_VALUE { + fn parse_set_param(stmt: &nodes::VariableSetStmt, query: &str) -> Result { + let mut value = if stmt.kind == VAR_SET_VALUE { Some(Self::parse_set_values(stmt)?) } else if stmt.kind == VAR_RESET || stmt.kind == VAR_SET_DEFAULT { None } else { panic!("parse_set_param called on invalid kind {}", stmt.kind); }; + let name = stmt.name().expect("SET always has name"); + + // PostgreSQL's raw parse tree normalizes both NONE and a quoted role + // named "none" to the same string. Preserve the keyword form by + // checking the original token at its parser-provided location. + if name == "role" + && value.as_ref().and_then(ParameterValue::as_str) == Some("none") + && let Some(Node::A_Const(constant)) = stmt.args().first() + && let Some(first) = query.as_bytes().get(constant.location as usize) + && !matches!(first, b'\'' | b'"') + { + value = None; + } match value { value @ Some(_) => Ok(SetParam { - name: stmt.name().expect("SET always has name").to_string(), + name: name.to_string(), value, local: stmt.is_local, }), None => Ok(SetParam { - name: stmt.name().expect("SET always has name").to_string(), + name: name.to_string(), value: None, - local: false, + local: stmt.is_local, }), } } @@ -70,12 +83,13 @@ impl QueryParser { context: &QueryParserContext, ) -> Result, Error> { let mut has_other = false; + let query = context.query()?.query(); let params = stmts .into_iter() .filter_map(|stmt| match stmt.stmt() { Node::VariableSetStmt(stmt) if stmt.kind != VAR_SET_MULTI => { - Some(Self::parse_set_param(stmt)) + Some(Self::parse_set_param(stmt, query)) } _ => { has_other = true; diff --git a/pgdog/src/frontend/router/parser/query/test/test_set.rs b/pgdog/src/frontend/router/parser/query/test/test_set.rs index 8917f2993..eb936423c 100644 --- a/pgdog/src/frontend/router/parser/query/test/test_set.rs +++ b/pgdog/src/frontend/router/parser/query/test/test_set.rs @@ -79,6 +79,80 @@ fn test_set_comment() { ); } +#[test] +fn test_set_role_is_tracked_as_parameter() { + let mut test = QueryParserTest::new(); + + let command = test.execute(vec![Query::new("SET ROLE other_user").into()]); + + match command { + Command::Set { + params, set_config, .. + } => { + assert_eq!(params.len(), 1); + assert_eq!(params[0].name, "role"); + assert_eq!( + params[0].value, + Some(ParameterValue::String("other_user".into())) + ); + assert!(!params[0].local); + assert!(!set_config); + } + _ => panic!("expected Command::Set, got {command:#?}"), + } +} + +#[test] +fn test_session_identity_commands_are_parameters() { + let mut test = QueryParserTest::new(); + + for (query, name, value) in [ + ( + "SET SESSION AUTHORIZATION other_user", + "session_authorization", + Some("other_user"), + ), + ( + "SET SESSION AUTHORIZATION DEFAULT", + "session_authorization", + None, + ), + ("RESET SESSION AUTHORIZATION", "session_authorization", None), + ("SET ROLE NONE", "role", None), + (r#"SET ROLE "none""#, "role", Some("none")), + ("RESET ROLE", "role", None), + ] { + let command = test.execute(vec![Query::new(query).into()]); + let Command::Set { params, .. } = command else { + panic!("expected Command::Set for {query}, got {command:#?}"); + }; + assert_eq!(params.len(), 1); + assert_eq!(params[0].name, name); + assert_eq!( + params[0].value.as_ref().and_then(ParameterValue::as_str), + value + ); + } +} + +#[test] +fn test_local_identity_reset_preserves_scope() { + let mut test = QueryParserTest::new(); + + for query in [ + "SET LOCAL ROLE NONE", + "SET LOCAL SESSION AUTHORIZATION DEFAULT", + ] { + let command = test.execute(vec![Query::new(query).into()]); + let Command::Set { params, .. } = command else { + panic!("expected Command::Set for {query}, got {command:#?}"); + }; + assert_eq!(params.len(), 1); + assert_eq!(params[0].value, None); + assert!(params[0].local); + } +} + #[test] fn test_set_config_null_value() { let mut test = QueryParserTest::new(); diff --git a/pgdog/src/net/parameter.rs b/pgdog/src/net/parameter.rs index fc808f47e..6b782f25a 100644 --- a/pgdog/src/net/parameter.rs +++ b/pgdog/src/net/parameter.rs @@ -32,6 +32,9 @@ static UNTRACKED_PARAMS: Lazy> = Lazy::new(|| { String::from("server_version"), String::from("server_encoding"), String::from("integer_datetimes"), + // Client SETs are tracked separately in SessionIdentity; don't replay + // the server-reported ParameterStatus as an ordinary GUC. + String::from("role"), String::from("session_authorization"), String::from("in_hot_standby"), String::from("pgdog.role"), @@ -161,6 +164,90 @@ impl ParameterValue { } } +const ROLE: &str = "role"; +const SESSION_AUTHORIZATION: &str = "session_authorization"; + +/// PostgreSQL session identity overrides. +/// +/// `None` means the authenticated connection user. +#[derive(Default, Debug, Clone, PartialEq, Eq)] +pub(crate) struct SessionIdentity { + session_authorization: Option, + role: Option, +} + +impl SessionIdentity { + fn parameter(name: &str) -> bool { + matches!(name, ROLE | SESSION_AUTHORIZATION) + } + + fn set(&mut self, name: &str, value: String) { + match name { + ROLE => self.role = Some(value), + SESSION_AUTHORIZATION => { + self.session_authorization = Some(value); + self.role = None; + } + _ => unreachable!("not a session identity parameter"), + } + } + + fn reset(&mut self, name: &str) { + match name { + ROLE => self.role = None, + SESSION_AUTHORIZATION => { + self.role = None; + self.session_authorization = None; + } + _ => unreachable!("not a session identity parameter"), + } + } + + fn set_parameter(&mut self, name: &str, value: &str) { + self.set(name, value.to_owned()); + } + + #[cfg(test)] + pub(crate) fn role(&self) -> Option<&str> { + self.role.as_deref() + } + + /// Queries that transform this identity into `target`. + fn reconcile(&self, target: &Self, local: bool) -> Vec { + if self == target { + return vec![]; + } + + let set = if local { "SET LOCAL" } else { "SET" }; + let mut queries = vec![Query::new(format!("{set} SESSION AUTHORIZATION DEFAULT"))]; + + if let Some(session_authorization) = &target.session_authorization { + queries.push(Query::new(format!( + "{set} SESSION AUTHORIZATION {}", + ParameterValue::String(session_authorization.clone()) + ))); + } + + if let Some(role) = &target.role { + queries.push(Query::new(format!( + "{set} ROLE {}", + ParameterValue::String(role.clone()) + ))); + } + + queries + } +} + +impl MemoryUsage for SessionIdentity { + fn memory_usage(&self) -> usize { + self.session_authorization + .as_ref() + .map_or(0, MemoryUsage::memory_usage) + + self.role.as_ref().map_or(0, MemoryUsage::memory_usage) + } +} + /// List of parameters. #[derive(Default, Debug, Clone, PartialEq)] pub(crate) struct Parameters { @@ -175,10 +262,16 @@ pub(crate) struct Parameters { /// what but we need to intercept them for databases that have cross shard /// queries disabled. transaction_local_params: BTreeMap, - /// Hash of `params` to avoid syncing params between clients and servers - /// when they are the same. /// Reset params. Stored here to support ROLLBACK. reset_params: BTreeMap, + /// Session identity committed outside a transaction. + identity: SessionIdentity, + /// Session identity changed by SET inside a transaction. + transaction_identity: Option>, + /// Session identity changed by SET LOCAL inside a transaction. + transaction_local_identity: Option>, + /// Hash of `params` to avoid syncing params between clients and servers + /// when they are the same. hash: u64, } @@ -196,7 +289,18 @@ impl Display for Parameters { impl MemoryUsage for Parameters { fn memory_usage(&self) -> usize { - self.params.memory_usage() + self.hash.memory_usage() + self.params.memory_usage() + + self.identity.memory_usage() + + self.transaction_identity.as_ref().map_or(0, |identity| { + std::mem::size_of::() + identity.memory_usage() + }) + + self + .transaction_local_identity + .as_ref() + .map_or(0, |identity| { + std::mem::size_of::() + identity.memory_usage() + }) + + self.hash.memory_usage() } } @@ -221,6 +325,13 @@ impl Parameters { self.hash = Self::compute_hash(&self.params); } + /// Restore the authenticated PostgreSQL identity. + pub(crate) fn clear_session_identity(&mut self) { + self.identity = SessionIdentity::default(); + self.transaction_identity = None; + self.transaction_local_identity = None; + } + /// Get parameter. pub(crate) fn get(&self, name: &str) -> Option<&ParameterValue> { if let Some(param) = self.transaction_local_params.get(name) { @@ -232,6 +343,104 @@ impl Parameters { } } + /// Store a PostgreSQL session identity parameter outside or inside a transaction. + /// + /// Returns true when `name` is a session identity parameter. + pub(crate) fn insert_identity( + &mut self, + name: &str, + value: &ParameterValue, + transaction: bool, + local: bool, + ) -> bool { + let name = name.to_lowercase(); + if !SessionIdentity::parameter(&name) { + return false; + } + if local && !transaction { + return true; + } + + let value = value + .as_str() + .expect("session identity has exactly one string value"); + + if transaction { + let base = self + .transaction_identity + .as_deref() + .unwrap_or(&self.identity) + .clone(); + if local { + self.transaction_local_identity + .get_or_insert_with(|| Box::new(base)) + .set_parameter(&name, value); + } else { + self.transaction_identity + .get_or_insert_with(|| Box::new(base)) + .set_parameter(&name, value); + if let Some(identity) = self.transaction_local_identity.as_mut() { + identity.set_parameter(&name, value); + } + } + } else { + self.identity.set_parameter(&name, value); + } + + true + } + + /// Reset a PostgreSQL session identity parameter. + /// + /// Returns true when `name` is a session identity parameter. + pub(crate) fn reset_identity(&mut self, name: &str, transaction: bool, local: bool) -> bool { + let name = name.to_lowercase(); + if !SessionIdentity::parameter(&name) { + return false; + } + if local && !transaction { + return true; + } + + if transaction { + let base = self + .transaction_identity + .as_deref() + .unwrap_or(&self.identity) + .clone(); + if local { + self.transaction_local_identity + .get_or_insert_with(|| Box::new(base)) + .reset(&name); + } else { + self.transaction_identity + .get_or_insert_with(|| Box::new(base)) + .reset(&name); + if let Some(identity) = self.transaction_local_identity.as_mut() { + identity.reset(&name); + } + } + } else { + self.identity.reset(&name); + } + + true + } + + /// Current session identity, including transaction overrides when requested. + #[cfg(test)] + pub(crate) fn session_identity(&self, transaction: bool) -> SessionIdentity { + if !transaction { + return self.identity.clone(); + } + + self.transaction_local_identity + .as_deref() + .or(self.transaction_identity.as_deref()) + .unwrap_or(&self.identity) + .clone() + } + /// Insert a parameter, but only for the duration of the transaction. pub(crate) fn insert_transaction( &mut self, @@ -264,6 +473,14 @@ impl Parameters { /// Restore parameters to the values supplied in the startup message, /// dropping everything changed since with `SET`. pub(crate) fn restore_startup(&mut self, startup: &Parameters) { + self.restore_startup_parameters(startup); + self.identity.clone_from(&startup.identity); + self.transaction_identity = None; + self.transaction_local_identity = None; + } + + /// Restore ordinary startup parameters while preserving session identity. + pub(crate) fn restore_startup_parameters(&mut self, startup: &Parameters) { self.params.clone_from(&startup.params); self.reset_params.clear(); self.hash = Self::compute_hash(&self.params); @@ -278,7 +495,7 @@ impl Parameters { keys.dedup(); for key in keys { - if !UNTRACKED_PARAMS.contains(&key) { + if !UNTRACKED_PARAMS.contains(&key) && !SessionIdentity::parameter(&key) { self.reset(&key); } } @@ -290,11 +507,17 @@ impl Parameters { "saved {} in-transaction params", self.transaction_params.len() ); - let changed = !self.transaction_params.is_empty() || !self.reset_params.is_empty(); + let changed = !self.transaction_params.is_empty() + || !self.reset_params.is_empty() + || self.transaction_identity.is_some(); + if let Some(identity) = self.transaction_identity.take() { + self.identity = *identity; + } self.params .extend(std::mem::take(&mut self.transaction_params)); self.transaction_local_params.clear(); + self.transaction_local_identity = None; self.reset_params.clear(); if changed { @@ -308,12 +531,12 @@ impl Parameters { pub(crate) fn rollback(&mut self) { self.transaction_params.clear(); self.transaction_local_params.clear(); + self.transaction_identity = None; + self.transaction_local_identity = None; - let mut reset = false; - for (name, value) in std::mem::take(&mut self.reset_params) { - self.params.insert(name, value); - reset = true; - } + let reset_params = std::mem::take(&mut self.reset_params); + let reset = !reset_params.is_empty(); + self.params.extend(reset_params); if reset { self.hash = Self::compute_hash(&self.params); @@ -344,7 +567,7 @@ impl Parameters { .filter(|(k, _)| !UNTRACKED_PARAMS.contains(k)) } - /// Filter our parameters that we would track with SET queries. + /// Filter parameters tracked on a pooled server. pub(crate) fn tracked(&self) -> Parameters { let params = self .tracked_iter() @@ -356,6 +579,7 @@ impl Parameters { Self { params, hash, + identity: self.identity.clone(), ..Default::default() } } @@ -388,7 +612,7 @@ impl Parameters { /// Merge params from self into other, generating the queries /// needed to sync that state on the server. pub(crate) fn identical(&self, other: &Self) -> bool { - self.hash == other.hash + self.hash == other.hash && self.identity == other.identity } /// Generate SET queries to change server state. @@ -404,15 +628,27 @@ impl Parameters { } if transaction_only { - let mut sets = self - .transaction_params - .iter() - .map(|(key, value)| query(key, value, false)) - .collect::>(); + let mut current = self.identity.clone(); + let mut sets = vec![]; + if let Some(identity) = self.transaction_identity.as_deref() { + sets.extend(current.reconcile(identity, false)); + current.clone_from(identity); + } + if let Some(identity) = self.transaction_local_identity.as_deref() { + sets.extend(current.reconcile(identity, true)); + } + + sets.extend( + self.transaction_params + .iter() + .filter(|(key, _)| !SessionIdentity::parameter(key)) + .map(|(key, value)| query(key, value, false)), + ); sets.extend( self.transaction_local_params .iter() + .filter(|(key, _)| !SessionIdentity::parameter(key)) .map(|(key, value)| query(key, value, true)), ); @@ -420,6 +656,7 @@ impl Parameters { } else { self.params .iter() + .filter(|(key, _)| !SessionIdentity::parameter(key)) .map(|(key, value)| query(key, value, false)) .collect() } @@ -434,11 +671,16 @@ impl Parameters { /// have a value on the incoming client. /// pub(crate) fn reset_queries(&self, other: &Self) -> Vec { - self.params - .keys() - .filter(|name| !other.contains_key(*name)) - .map(|name| Query::new(format!(r#"RESET "{}""#, name))) - .collect() + let mut queries = self.identity.reconcile(&other.identity, false); + queries.extend( + self.params + .keys() + .filter(|name| { + !other.contains_key(*name) && !SessionIdentity::parameter(name.as_str()) + }) + .map(|name| Query::new(format!(r#"RESET "{}""#, name))), + ); + queries } /// Get parameter value or returned an error. @@ -456,6 +698,10 @@ impl Parameters { /// Copy params set inside the transaction. pub(crate) fn copy_in_transaction(&mut self, other: &Self) { + self.transaction_identity + .clone_from(&other.transaction_identity); + self.transaction_local_identity + .clone_from(&other.transaction_local_identity); self.transaction_params.extend( other .transaction_params @@ -503,6 +749,9 @@ impl From> for Parameters { transaction_params: BTreeMap::new(), transaction_local_params: BTreeMap::new(), reset_params: BTreeMap::new(), + identity: SessionIdentity::default(), + transaction_identity: None, + transaction_local_identity: None, } } } @@ -527,7 +776,7 @@ mod test { use crate::net::ToBytes; use crate::net::parameter::ParameterValue; - use super::Parameters; + use super::{Parameters, SessionIdentity}; #[test] fn test_identical() { @@ -548,6 +797,235 @@ mod test { assert!(Parameters::default().identical(&Parameters::default())); } + #[test] + fn test_session_identity_reconcile_order() { + let current = SessionIdentity { + session_authorization: Some("old_auth".into()), + role: Some("old_role".into()), + }; + let target = SessionIdentity { + session_authorization: Some("new_auth".into()), + role: Some("new_role".into()), + }; + + let queries = current + .reconcile(&target, false) + .into_iter() + .map(|query| query.query().to_owned()) + .collect::>(); + + assert_eq!( + queries, + [ + "SET SESSION AUTHORIZATION DEFAULT", + r#"SET SESSION AUTHORIZATION "new_auth""#, + r#"SET ROLE "new_role""#, + ] + ); + + let queries = target + .reconcile(&SessionIdentity::default(), true) + .into_iter() + .map(|query| query.query().to_owned()) + .collect::>(); + assert_eq!(queries, ["SET LOCAL SESSION AUTHORIZATION DEFAULT"]); + } + + #[test] + fn test_session_identity_commit_and_rollback() { + let mut params = Parameters::default(); + params.insert_identity("role", &"base_role".into(), false, false); + params.insert_identity("role", &"transaction_role".into(), true, false); + + assert_eq!( + params.session_identity(true).role.as_deref(), + Some("transaction_role") + ); + params.rollback(); + assert_eq!( + params.session_identity(false).role.as_deref(), + Some("base_role") + ); + + params.insert_identity("role", &"committed_role".into(), true, false); + params.insert_identity("role", &"local_role".into(), true, true); + assert_eq!( + params.session_identity(true).role.as_deref(), + Some("local_role") + ); + params.commit(); + assert_eq!( + params.session_identity(false).role.as_deref(), + Some("committed_role") + ); + } + + #[test] + fn test_session_identity_respects_set_local_order() { + let mut params = Parameters::default(); + params.insert_identity("role", &"local_first".into(), true, true); + params.insert_identity("role", &"session_last".into(), true, false); + assert_eq!( + params.session_identity(true).role.as_deref(), + Some("session_last") + ); + params.commit(); + assert_eq!( + params.session_identity(false).role.as_deref(), + Some("session_last") + ); + + params.insert_identity("role", &"session_first".into(), true, false); + params.insert_identity("role", &"local_last".into(), true, true); + assert_eq!( + params.session_identity(true).role.as_deref(), + Some("local_last") + ); + params.commit(); + assert_eq!( + params.session_identity(false).role.as_deref(), + Some("session_first") + ); + + params.reset_identity("role", true, true); + assert_eq!(params.session_identity(true).role, None); + params.commit(); + assert_eq!( + params.session_identity(false).role.as_deref(), + Some("session_first") + ); + } + + #[test] + fn test_set_local_identity_outside_transaction_is_not_persisted() { + let mut params = Parameters::default(); + params.insert_identity("role", &"ignored".into(), false, true); + assert_eq!(params.session_identity(false), SessionIdentity::default()); + } + + #[test] + fn test_reset_session_authorization_is_transactional_and_clears_role() { + let mut params = Parameters::default(); + params.insert_identity( + "session_authorization", + &"delegated_user".into(), + false, + false, + ); + params.insert_identity("role", &"reporting".into(), false, false); + + params.reset_identity("session_authorization", true, false); + assert_eq!(params.session_identity(true), SessionIdentity::default()); + params.rollback(); + assert_eq!( + params.session_identity(false), + SessionIdentity { + session_authorization: Some("delegated_user".into()), + role: Some("reporting".into()), + } + ); + + params.reset_identity("session_authorization", true, false); + params.commit(); + assert_eq!(params.session_identity(false), SessionIdentity::default()); + } + + #[test] + fn test_reset_all_preserves_session_identity() { + let mut params = Parameters::default(); + params.insert("work_mem", "1MB"); + params.insert_identity("role", &"reporting".into(), false, false); + + params.reset_all(); + + assert!(params.get("work_mem").is_none()); + assert_eq!( + params.session_identity(false).role.as_deref(), + Some("reporting") + ); + } + + #[test] + fn test_server_snapshot_keeps_session_identity() { + let mut params = Parameters::default(); + params.insert_identity("role", &"reporting".into(), false, false); + + let snapshot = params.tracked(); + + assert_eq!( + snapshot.session_identity(false).role.as_deref(), + Some("reporting") + ); + } + + #[test] + fn test_reset_queries_include_session_identity() { + let mut server = Parameters::default(); + server.insert_identity("role", &"reporting".into(), false, false); + server.insert("work_mem", "1MB"); + + let client = Parameters::default(); + let queries = server + .reset_queries(&client) + .into_iter() + .map(|query| query.query().to_owned()) + .collect::>(); + + assert_eq!( + queries, + ["SET SESSION AUTHORIZATION DEFAULT", r#"RESET "work_mem""#,] + ); + } + + #[test] + fn test_set_queries_include_transaction_identity() { + let mut params = Parameters::default(); + params.insert_identity("role", &"reporting".into(), true, false); + params.insert_transaction("work_mem", "1MB", false); + + let queries = params + .set_queries(true) + .into_iter() + .map(|query| query.query().to_owned()) + .collect::>(); + + assert_eq!( + queries, + [ + "SET SESSION AUTHORIZATION DEFAULT", + r#"SET ROLE "reporting""#, + r#"SET "work_mem" TO "1MB""#, + ] + ); + } + + #[test] + fn test_restore_startup_parameters_preserves_session_identity() { + let mut startup = Parameters::default(); + startup.insert("search_path", "public"); + + let mut params = Parameters::default(); + params.insert("search_path", "private"); + params.insert_identity("role", &"reporting".into(), false, false); + params.restore_startup_parameters(&startup); + + assert_eq!(params.get("search_path"), startup.get("search_path")); + assert_eq!( + params.session_identity(false).role.as_deref(), + Some("reporting") + ); + } + + #[test] + fn test_role_named_none_is_distinct_from_reset() { + let mut params = Parameters::default(); + params.insert_identity("role", &"none".into(), false, false); + + assert_eq!(params.session_identity(false).role.as_deref(), Some("none")); + params.reset_identity("role", false, false); + assert_eq!(params.session_identity(false).role, None); + } + #[test] fn test_tracked_and_different() { let mut client = Parameters::default();