diff --git a/src/intercept.rs b/src/intercept.rs index e430c0b..446613b 100644 --- a/src/intercept.rs +++ b/src/intercept.rs @@ -24,6 +24,23 @@ fn single_text_response(column_name: &str, value: &str) -> PgWireResult PgWireResult> { + let fields = Arc::new(vec![text_field(column_name, Type::BOOL)]); + + let mut encoder = DataRowEncoder::new(Arc::clone(&fields)); + encoder.encode_field(&value)?; + let row = encoder.take_row(); + + Ok(vec![Response::Query(QueryResponse::new( + fields, + stream::iter(vec![Ok(row)]), + ))]) +} + /// If the query can be answered locally instead of forwarded to Trino, /// build the response and return it. Returns `None` for queries that /// should pass through to Trino. @@ -37,11 +54,15 @@ fn single_text_response(column_name: &str, value: &str) -> PgWireResult Option>> { let trimmed = query.trim(); if trimmed.is_empty() { @@ -95,6 +116,17 @@ pub fn try_intercept( return Some(resp); } + // Npgsql (Power BI) and pgjdbc probe the connection's encryption state + // right after connecting with `SELECT ssl FROM pg_stat_ssl WHERE pid = + // pg_backend_pid()`. Trino has no `pg_stat_ssl`, so forwarding it aborts + // the session before the client runs a single user query. Answer locally + // with the real TLS state of this connection; the `WHERE pid = ...` clause + // is moot because a backend only ever sees its own row in pg_stat_ssl. + if parsed_query.references_table("pg_stat_ssl") { + tracing::trace!(query = trimmed, "Intercept: pg_stat_ssl"); + return Some(single_bool_response("ssl", client_is_secure)); + } + if let Some(resp) = crate::catalog::handle_catalog_query(parsed_query) { tracing::trace!(query = trimmed, "Intercept: pg_catalog"); return Some(resp); @@ -188,7 +220,7 @@ mod tests { fn assert_intercepted(query: &str) { let parsed_query = ParsedQuery::new(query); assert!( - try_intercept(query, &parsed_query, "test_catalog", "test_schema").is_some(), + try_intercept(query, &parsed_query, "test_catalog", "test_schema", false).is_some(), "expected query to be intercepted: {query}" ); } @@ -196,7 +228,7 @@ mod tests { fn assert_not_intercepted(query: &str) { let parsed_query = ParsedQuery::new(query); assert!( - try_intercept(query, &parsed_query, "test_catalog", "test_schema").is_none(), + try_intercept(query, &parsed_query, "test_catalog", "test_schema", false).is_none(), "expected query to NOT be intercepted: {query}" ); } @@ -246,7 +278,7 @@ mod tests { for &(query, expected) in cases { let parsed_query = ParsedQuery::new(query); - let result = try_intercept(query, &parsed_query, "test_catalog", "test_schema") + let result = try_intercept(query, &parsed_query, "test_catalog", "test_schema", false) .unwrap_or_else(|| panic!("SHOW not intercepted: {query}")); assert!(result.is_ok(), "SHOW returned error for: {query}"); @@ -280,6 +312,45 @@ mod tests { assert_intercepted("SELECT current_setting('server_version')"); } + /// The Npgsql/pgjdbc SSL probe is answered locally instead of forwarded + /// to Trino (which has no `pg_stat_ssl` and would abort the connection). + #[test] + fn pg_stat_ssl_probe_intercepted() { + assert_intercepted("SELECT ssl FROM pg_stat_ssl WHERE pid = pg_backend_pid()"); + assert_intercepted("select ssl from pg_stat_ssl"); + } + + /// The `ssl` column reports the connection's real TLS state and is typed + /// BOOL (not VARCHAR) so the RowDescription oid matches what the driver + /// decodes. + #[test] + fn pg_stat_ssl_reports_tls_state_as_bool() { + let query = "SELECT ssl FROM pg_stat_ssl WHERE pid = pg_backend_pid()"; + let parsed_query = ParsedQuery::new(query); + + for is_secure in [true, false] { + let resp = try_intercept(query, &parsed_query, "cat", "sch", is_secure) + .expect("pg_stat_ssl must be intercepted") + .expect("intercept must not error"); + assert_eq!(resp.len(), 1); + match &resp[0] { + Response::Query(qr) => { + assert_eq!(qr.row_schema.len(), 1); + assert_eq!(qr.row_schema[0].name(), "ssl"); + assert_eq!(qr.row_schema[0].datatype(), &Type::BOOL); + } + other => panic!("expected Query response, got: {other:?}"), + } + } + } + + /// Regression: a user table literally named `pg_stat_ssl` in a literal + /// must not trip the probe intercept. + #[test] + fn pg_stat_ssl_in_literal_not_intercepted() { + assert_not_intercepted("SELECT 'pg_stat_ssl' AS sentinel"); + } + #[test] fn regular_queries_not_intercepted() { assert_not_intercepted("SELECT 1"); diff --git a/src/query_extended.rs b/src/query_extended.rs index a5e2491..ac9bf8d 100644 --- a/src/query_extended.rs +++ b/src/query_extended.rs @@ -91,6 +91,8 @@ impl ExtendedQueryHandler for GatewayExtendedQueryHandler { let query = &portal.statement.statement; tracing::debug!(query, "Extended query execute"); + let client_is_secure = client.is_secure(); + let conn_state = client .session_extensions() .get::() @@ -118,6 +120,7 @@ impl ExtendedQueryHandler for GatewayExtendedQueryHandler { &conn_state.config, Some(&conn_state.active_query_id), Some(&portal.result_column_format), + client_is_secure, ) .await?; let response = responses @@ -182,6 +185,8 @@ impl ExtendedQueryHandler for GatewayExtendedQueryHandler { return Ok(DescribeStatementResponse::no_data()); } + let client_is_secure = client.is_secure(); + let conn_state = client .session_extensions() .get::() @@ -201,6 +206,7 @@ impl ExtendedQueryHandler for GatewayExtendedQueryHandler { &conn_state.config, Some(&conn_state.active_query_id), None, + client_is_secure, ) .await?; let response = responses @@ -230,6 +236,8 @@ impl ExtendedQueryHandler for GatewayExtendedQueryHandler { let query = &portal.statement.statement; tracing::debug!(query, "Extended query describe portal"); + let client_is_secure = client.is_secure(); + let conn_state = client .session_extensions() .get::() @@ -241,6 +249,7 @@ impl ExtendedQueryHandler for GatewayExtendedQueryHandler { &conn_state.config, Some(&conn_state.active_query_id), Some(&portal.result_column_format), + client_is_secure, ) .await?; let response = responses diff --git a/src/query_pipeline.rs b/src/query_pipeline.rs index cea8b2b..e3cae9c 100644 --- a/src/query_pipeline.rs +++ b/src/query_pipeline.rs @@ -49,6 +49,7 @@ pub(crate) async fn process_query( config: &Arc, active_query_id: Option<&ActiveQueryId>, result_format: Option<&Format>, + client_is_secure: bool, ) -> PgWireResult> { tracing::trace!(query, "Pipeline: enter"); @@ -60,6 +61,7 @@ pub(crate) async fn process_query( config, active_query_id, result_format, + client_is_secure, ) .await; } @@ -67,8 +69,15 @@ pub(crate) async fn process_query( tracing::trace!(count = pieces.len(), "Pipeline: multi-statement input"); let mut out = Vec::with_capacity(pieces.len()); for stmt in &pieces { - match process_single_statement(stmt, trino_client, config, active_query_id, result_format) - .await + match process_single_statement( + stmt, + trino_client, + config, + active_query_id, + result_format, + client_is_secure, + ) + .await { Ok(mut responses) => out.append(&mut responses), // User-visible errors (e.g. a Trino syntax error on statement N @@ -107,6 +116,7 @@ async fn process_single_statement( config: &Arc, active_query_id: Option<&ActiveQueryId>, result_format: Option<&Format>, + client_is_secure: bool, ) -> PgWireResult> { // The query is parsed up to three times: once here (for routing // checks), once by the multi-statement splitter in the public @@ -122,6 +132,7 @@ async fn process_single_statement( &parsed_query, &config.trino_catalog, &config.trino_schema, + client_is_secure, ) { tracing::trace!("Pipeline: static intercept matched"); return result; diff --git a/src/query_simple.rs b/src/query_simple.rs index f404a58..2ef50be 100644 --- a/src/query_simple.rs +++ b/src/query_simple.rs @@ -26,6 +26,10 @@ impl SimpleQueryHandler for GatewayQueryHandler { { tracing::debug!(query, "Simple query received"); + // Captured before borrowing session state, so the `pg_stat_ssl` + // intercept can report this connection's real TLS state. + let client_is_secure = client.is_secure(); + let conn_state = client .session_extensions() .get::() @@ -39,6 +43,7 @@ impl SimpleQueryHandler for GatewayQueryHandler { &conn_state.config, Some(&conn_state.active_query_id), None, + client_is_secure, ) .await; match &result {