From a3cf5fe256f17d1e0b2cad8ef327da26678e708c Mon Sep 17 00:00:00 2001 From: zipg Date: Mon, 10 Aug 2026 10:32:42 +0800 Subject: [PATCH] fix(postgres): avoid crash on PG-compatible database connect Co-authored-by: zipg <4047349+zipg@users.noreply.github.com> --- crates/dbx-core/src/db/postgres.rs | 84 ++++++++++++++++++++++++------ 1 file changed, 67 insertions(+), 17 deletions(-) diff --git a/crates/dbx-core/src/db/postgres.rs b/crates/dbx-core/src/db/postgres.rs index d36e70af6..b0e97289b 100644 --- a/crates/dbx-core/src/db/postgres.rs +++ b/crates/dbx-core/src/db/postgres.rs @@ -1623,11 +1623,15 @@ async fn connect_with_local_timezone(url: &str, fallback_timeout: Duration, time /// (server address, server port, backend PID). The PID alone is not unique /// across different servers, so the server's own address/port disambiguate /// (`inet_server_addr()` is NULL for Unix sockets, hence the fallback). -type PostgresConnectionKey = (String, i32, i32); +type PostgresConnectionKey = (String, String, String); -const POSTGRES_CONNECTION_IDENTITY_SQL: &str = "SELECT pg_backend_pid(), \ +const POSTGRES_CONNECTION_IDENTITY_SQL: &str = "SELECT pg_backend_pid()::text, \ COALESCE(host(inet_server_addr()), 'unix'), \ - COALESCE(inet_server_port(), current_setting('port')::integer)"; + COALESCE(inet_server_port()::text, current_setting('port'))"; + +fn postgres_connection_key_from_row(row: &Row) -> Option { + Some((row.try_get::<_, String>(1).ok()?, row.try_get::<_, String>(2).ok()?, row.try_get::<_, String>(0).ok()?)) +} /// Notice buffers for live connections, keyed by connection identity. Entries /// are weak so they disappear once the pooled connection (and its driver @@ -1728,13 +1732,14 @@ where // attributed to query results on this connection and are logged // by the driver task instead. Never fail the connection over this. if let Ok(row) = client.query_one(POSTGRES_CONNECTION_IDENTITY_SQL, &[]).await { - let key: PostgresConnectionKey = (row.get(1), row.get(2), row.get(0)); - let buffer = Arc::new(Mutex::new(Vec::new())); - let mut buffers = postgres_notice_buffers().lock().unwrap_or_else(|poisoned| poisoned.into_inner()); - buffers.retain(|_, weak| weak.strong_count() > 0); - buffers.insert(key, Arc::downgrade(&buffer)); - drop(buffers); - *notice_buffer.lock().unwrap_or_else(|poisoned| poisoned.into_inner()) = Some(buffer); + if let Some(key) = postgres_connection_key_from_row(&row) { + let buffer = Arc::new(Mutex::new(Vec::new())); + let mut buffers = postgres_notice_buffers().lock().unwrap_or_else(|poisoned| poisoned.into_inner()); + buffers.retain(|_, weak| weak.strong_count() > 0); + buffers.insert(key, Arc::downgrade(&buffer)); + drop(buffers); + *notice_buffer.lock().unwrap_or_else(|poisoned| poisoned.into_inner()) = Some(buffer); + } } Ok((client, conn_task)) @@ -1770,7 +1775,7 @@ async fn resolve_postgres_client_key(client: &deadpool_postgres::Client) -> Opti .query_one(POSTGRES_CONNECTION_IDENTITY_SQL, &[]) .await .ok() - .map(|row| (row.get(1), row.get(2), row.get(0))); + .and_then(|row| postgres_connection_key_from_row(&row)); let mut keys = postgres_client_keys().lock().unwrap_or_else(|poisoned| poisoned.into_inner()); keys.retain(|_, (cached, _)| cached.strong_count() > 0); let cache_key = Arc::as_ptr(&client.statement_cache) as usize; @@ -5092,9 +5097,52 @@ mod tests { } } + #[test] + fn postgres_connection_identity_normalizes_vendor_numeric_types_to_text() { + assert!(POSTGRES_CONNECTION_IDENTITY_SQL.contains("pg_backend_pid()::text")); + assert!(POSTGRES_CONNECTION_IDENTITY_SQL.contains("inet_server_port()::text")); + } + + #[tokio::test] + #[ignore = "requires DBX_TEST_POSTGRES_URL pointing at a PostgreSQL-compatible database"] + async fn postgres_connection_identity_supports_compatible_servers() { + let url = std::env::var("DBX_TEST_POSTGRES_URL").expect("DBX_TEST_POSTGRES_URL"); + let pool = connect_with_local_timezone(&url, Duration::from_secs(10), "UTC") + .await + .expect("connect PostgreSQL-compatible database"); + let client = pool.get().await.expect("checkout PostgreSQL-compatible database"); + let key = postgres_client_key(&client).await.expect("resolve text connection identity"); + + assert!(!key.0.is_empty()); + assert!(!key.1.is_empty()); + assert!(!key.2.is_empty()); + } + + #[tokio::test] + #[ignore = "requires DBX_TEST_POSTGRES_URL pointing at a PostgreSQL database"] + async fn postgres_connection_identity_preserves_notice_capture() { + let url = std::env::var("DBX_TEST_POSTGRES_URL").expect("DBX_TEST_POSTGRES_URL"); + let pool = connect_with_local_timezone(&url, Duration::from_secs(10), "UTC") + .await + .expect("connect PostgreSQL database"); + let client = pool.get().await.expect("checkout PostgreSQL database"); + + let result = execute_query_with_max_rows_inner( + &client, + "DO $$ BEGIN RAISE NOTICE 'dbx notice identity regression'; END $$", + None, + false, + None, + ) + .await + .expect("execute statement with notice"); + + assert!(result.messages.iter().any(|message| message.message == "dbx notice identity regression")); + } + #[test] fn take_notices_for_key_returns_buffered_notices_and_empties_buffer() { - let key = ("test-host".to_string(), 9_000_001, 9_000_001); + let key = ("test-host".to_string(), "9000001".to_string(), "9000001".to_string()); let buffer = Arc::new(Mutex::new(vec![test_query_message("first"), test_query_message("second")])); postgres_notice_buffers() .lock() @@ -5118,8 +5166,8 @@ mod tests { #[test] fn take_notices_for_key_prunes_dead_buffers_and_misses_return_empty() { - let live_key = ("test-host".to_string(), 9_000_002, 9_000_002); - let dead_key = ("test-host".to_string(), 9_000_003, 9_000_003); + let live_key = ("test-host".to_string(), "9000002".to_string(), "9000002".to_string()); + let dead_key = ("test-host".to_string(), "9000003".to_string(), "9000003".to_string()); let live = Arc::new(Mutex::new(vec![test_query_message("live")])); let dead = Arc::new(Mutex::new(vec![test_query_message("dead")])); { @@ -5129,7 +5177,9 @@ mod tests { } drop(dead); - assert!(take_notices_for_key(&("test-host".to_string(), 9_000_004, 9_000_004)).is_empty()); + assert!( + take_notices_for_key(&("test-host".to_string(), "9000004".to_string(), "9000004".to_string())).is_empty() + ); assert!(take_notices_for_key(&dead_key).is_empty()); let buffers = postgres_notice_buffers().lock().unwrap_or_else(|poisoned| poisoned.into_inner()); assert!(!buffers.contains_key(&dead_key)); @@ -5145,8 +5195,8 @@ mod tests { fn take_notices_for_key_distinguishes_same_pid_on_different_servers() { // Backend PIDs collide across servers; the (address, port, pid) key // keeps notice attribution separate. - let key_a = ("server-a".to_string(), 5432, 42); - let key_b = ("server-b".to_string(), 5432, 42); + let key_a = ("server-a".to_string(), "5432".to_string(), "42".to_string()); + let key_b = ("server-b".to_string(), "5432".to_string(), "42".to_string()); let buffer_a = Arc::new(Mutex::new(vec![test_query_message("from-a")])); let buffer_b = Arc::new(Mutex::new(vec![test_query_message("from-b")])); {