diff --git a/src/webserver/oidc.rs b/src/webserver/oidc.rs index 21a40e24..a362d08b 100644 --- a/src/webserver/oidc.rs +++ b/src/webserver/oidc.rs @@ -18,6 +18,7 @@ use actix_web::{ }; use anyhow::{anyhow, Context}; use awc::Client; +use futures_util::StreamExt; use openidconnect::core::{ CoreAuthDisplay, CoreAuthPrompt, CoreErrorResponseType, CoreGenderClaim, CoreJsonWebKey, CoreJweContentEncryptionAlgorithm, CoreJwsSigningAlgorithm, CoreRevocableToken, @@ -429,21 +430,21 @@ async fn handle_request( } Ok(None) => { log::trace!("No authenticated user found"); - handle_unauthenticated_request(oidc_state, request) + handle_unauthenticated_request(oidc_state, request).await } Err(e) => { log::debug!("An auth cookie is present but could not be verified. Redirecting to OIDC provider to re-authenticate. {e:?}"); if let Some(c) = http_client { oidc_state.maybe_refresh(c, OIDC_CLIENT_MIN_REFRESH_INTERVAL); } - handle_unauthenticated_request(oidc_state, request) + handle_unauthenticated_request(oidc_state, request).await } } } -fn handle_unauthenticated_request( +async fn handle_unauthenticated_request( oidc_state: &OidcState, - request: ServiceRequest, + mut request: ServiceRequest, ) -> MiddlewareResponse { log::debug!("Handling unauthenticated request to {}", request.path()); @@ -453,6 +454,17 @@ fn handle_unauthenticated_request( log::debug!("Redirecting to OIDC provider"); + // Drain the request body to prevent broken pipes with buffering proxies + if request.method() != actix_web::http::Method::GET { + let mut payload = request.take_payload(); + while let Some(chunk) = payload.next().await { + if let Err(e) = chunk { + log::warn!("Error draining payload: {}", e); + break; + } + } + } + let initial_url = request.uri().to_string(); let redirect_count = get_redirect_count(&request); let response = build_auth_provider_redirect_response(oidc_state, &initial_url, redirect_count); diff --git a/tests/oidc/mod.rs b/tests/oidc/mod.rs index d3c1a4f2..b0848070 100644 --- a/tests/oidc/mod.rs +++ b/tests/oidc/mod.rs @@ -11,6 +11,7 @@ use serde::{Deserialize, Serialize}; use serde_json::json; use sqlpage::webserver::http::create_app; use std::collections::HashMap; +use std::path::PathBuf; use std::sync::{Arc, Mutex}; use std::time::Duration; use tokio_util::sync::{CancellationToken, DropGuard}; @@ -653,3 +654,74 @@ async fn test_slow_token_endpoint_does_not_freeze_server() { .unwrap(); assert_eq!(resp.status(), StatusCode::SEE_OTHER); } + +#[actix_web::test] +async fn test_oidc_unauthenticated_post_drains_body() { + use sqlpage::AppState; + crate::common::init_log(); + let provider = FakeOidcProvider::new().await; + + let tmp_dir_path = std::env::temp_dir().join(format!("sqlpage_test_{}", rand::random::())); + std::fs::create_dir_all(&tmp_dir_path).unwrap(); + let web_root = tmp_dir_path.display().to_string().replace('\\', "/"); + std::fs::write(tmp_dir_path.join("index.sql"), "SELECT 'debug' AS component;").unwrap(); + + let mut config = crate::common::test_config(); + config.oidc_issuer_url = Some(openidconnect::IssuerUrl::new(provider.issuer_url.clone()).unwrap()); + config.oidc_client_id = provider.client_id.clone(); + config.oidc_client_secret = Some(provider.client_secret.clone()); + config.oidc_protected_paths = vec!["/".to_string()]; + config.web_root = PathBuf::from(&web_root); + config.environment = sqlpage::app_config::DevOrProd::Production; + + let app_state = Arc::new(AppState::init(&config).await.unwrap()); + + let server_state = Arc::clone(&app_state); + let server = HttpServer::new(move || { + create_app(Data::from(Arc::clone(&server_state))) + }) + .bind("127.0.0.1:0") + .unwrap(); + let addr = server.addrs()[0]; + let server_handle = tokio::spawn(server.run()); + + let payload = vec![0u8; 1024 * 1024]; // 1MB payload + + // Use a raw TCP stream to see if we get a broken pipe when writing the payload + use tokio::io::{AsyncReadExt, AsyncWriteExt}; + let mut stream = tokio::net::TcpStream::connect(addr).await.unwrap(); + let request_header = format!( + "POST / HTTP/1.1\r\nHost: {}\r\nContent-Length: {}\r\nAccept: text/html\r\n\r\n", + addr, payload.len() + ); + stream.write_all(request_header.as_bytes()).await.unwrap(); + + let (mut reader, mut writer) = stream.into_split(); + + // Start a task to write the payload + let writer_handle = tokio::spawn(async move { + let chunk_iter = payload.chunks(64 * 1024); + for chunk in chunk_iter { + if let Err(e) = writer.write_all(chunk).await { + return Err(e); + } + tokio::task::yield_now().await; + } + Ok(()) + }); + + // Read the response + let mut buf = [0u8; 1024]; + let n = reader.read(&mut buf).await.unwrap(); + let response = String::from_utf8_lossy(&buf[..n]); + assert!(response.contains("HTTP/1.1 303 See Other")); + + // Ensure the writer finished successfully. Without the fix, this might return a Broken Pipe error. + let write_res = writer_handle.await.unwrap(); + if let Err(e) = write_res { + panic!("Failed to write payload after receiving 303: {:?}. This is the bug!", e); + } + + server_handle.abort(); +} +