diff --git a/pgdog/src/backend/server.rs b/pgdog/src/backend/server.rs index 3bf9d603f..4f44f998d 100644 --- a/pgdog/src/backend/server.rs +++ b/pgdog/src/backend/server.rs @@ -605,10 +605,10 @@ impl Server { // Compare client and server params. let mut executed = if !params.identical(&self.client_params) { // Construct client parameter SET queries. - let tracked = params.tracked(); + let tracked = params.tracked_and_different(&self.client_params); // Construct RESET queries to reset any current params // to their default values. - let mut queries = self.client_params.reset_queries(); + let mut queries = self.client_params.reset_queries(params); // Combine both to create a new, fresh session state // on this connection. @@ -624,7 +624,7 @@ impl Server { } // Update params on this connection. - self.client_params = tracked; + self.client_params = params.tracked(); queries.len() } else { @@ -2313,7 +2313,7 @@ pub mod test { let changed = server .link_client(FrontendPid::new(), ¶ms, None) .await?; - assert_eq!(changed, 2); // RESET, SET. + assert_eq!(changed, 1); // SET only; the parameter already exists. let changed = server .link_client(FrontendPid::new(), ¶ms, None) diff --git a/pgdog/src/net/parameter.rs b/pgdog/src/net/parameter.rs index 88bcc76af..a105b181c 100644 --- a/pgdog/src/net/parameter.rs +++ b/pgdog/src/net/parameter.rs @@ -349,11 +349,42 @@ impl Parameters { if entries > 0 { hasher.finish() } else { 0 } } - pub fn tracked(&self) -> Parameters { - let params = self - .params + /// Iterate over parameters that we track with SET queries. + pub(crate) fn tracked_iter(&self) -> impl Iterator { + self.params .iter() .filter(|(k, _)| !UNTRACKED_PARAMS.contains(k)) + } + + /// Filter our parameters that we would track with SET queries. + pub(crate) fn tracked(&self) -> Parameters { + let params = self + .tracked_iter() + .map(|(k, v)| (k.clone(), v.clone())) + .collect::>(); + + let hash = Self::compute_hash(¶ms); + + Self { + params, + hash, + ..Default::default() + } + } + + /// Calculate the parameters we need to update on the server, + /// excluding the parameters we do not track because they have no effect, e.g., "pgdog"."role", + /// or which cannot be changed, e.g., "user". + /// + /// # Arguments + /// + /// - `other`: Parameters stored on the server. + /// + pub(crate) fn tracked_and_different(&self, other: &Self) -> Parameters { + let params = self + .tracked_iter() + // Ignore parameters that have identical values, they don't need to be updated. + .filter(|(k, v)| other.get(k).map(|other| other != *v).unwrap_or(true)) .map(|(k, v)| (k.clone(), v.clone())) .collect::>(); @@ -406,9 +437,18 @@ impl Parameters { } } - pub fn reset_queries(&self) -> Vec { + /// Create a list of `RESET` queries that will reset parameters + /// back to their default value. + /// + /// This will ignore all parameters that are about to be SET + /// by incoming parameters. It will only reset parameters + /// that are currently set on the server and which do not + /// 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() } @@ -529,6 +569,30 @@ mod test { assert!(Parameters::default().identical(&Parameters::default())); } + #[test] + fn test_tracked_and_different() { + let mut client = Parameters::default(); + client.insert("application_name", "client"); + client.insert("statement_timeout", "1001ms"); + client.insert("search_path", "public"); + + let mut server = Parameters::default(); + server.insert("application_name", "server"); + server.insert("statement_timeout", "1001ms"); + + let different = client.tracked_and_different(&server); + + assert_eq!( + different.get("application_name"), + Some(&ParameterValue::String("client".into())) + ); + assert_eq!( + different.get("search_path"), + Some(&ParameterValue::String("public".into())) + ); + assert!(!different.contains_key("statement_timeout")); + } + #[test] fn test_insert_transaction_non_local() { let mut params = Parameters::default(); @@ -819,7 +883,7 @@ mod test { assert_eq!(timeout.first().unwrap(), "5s"); // Get reset queries before resetting (reset_queries uses current params) - let reset_queries = params.reset_queries(); + let reset_queries = params.reset_queries(&Parameters::default()); assert_eq!(reset_queries.len(), 2); // Execute reset queries on server @@ -859,7 +923,7 @@ mod test { assert_eq!(timeout.first().unwrap(), "5s"); // Get reset queries and execute on server - let reset_queries = params.reset_queries(); + let reset_queries = params.reset_queries(&Parameters::default()); for query in reset_queries { server.execute(query).await.unwrap(); }