diff --git a/crates/pg/src/pg_server.rs b/crates/pg/src/pg_server.rs index f1de39f4efd..41b05ccddc8 100644 --- a/crates/pg/src/pg_server.rs +++ b/crates/pg/src/pg_server.rs @@ -132,12 +132,17 @@ async fn response(res: axum::response::Result, database: &str) -> Result>); + static LOGGER: Capture = Capture(Mutex::new(Vec::new())); + const DATABASE: &str = "pg-response-log-level-test"; + + impl log::Log for Capture { + fn enabled(&self, _: &log::Metadata<'_>) -> bool { + true + } + + fn log(&self, record: &log::Record<'_>) { + // Ignore records from unrelated tests running in parallel. + if record.args().to_string().contains(DATABASE) { + self.0.lock().unwrap().push(record.level()); + } + } + + fn flush(&self) {} + } + + #[tokio::test] + async fn response_distinguishes_client_and_server_failure_logs() { + log::set_logger(&LOGGER).unwrap(); + log::set_max_level(log::LevelFilter::Trace); + + for (status, message, expected_level) in [ + (StatusCode::BAD_REQUEST, "invalid SQL", log::Level::Warn), + ( + StatusCode::INTERNAL_SERVER_ERROR, + "internal server error", + log::Level::Error, + ), + ] { + LOGGER.0.lock().unwrap().clear(); + let result = response::<()>(Err((status, message).into()), DATABASE).await; + + assert!(matches!(result, Err(PgError::Sql(ref text)) if text == message)); + assert_eq!( + *LOGGER.0.lock().unwrap(), + vec![expected_level], + "unexpected log level for {status}" + ); + } + } +}