oidc drain

This commit is contained in:
lovasoa
2026-03-09 21:55:41 +01:00
parent 357b5bbca2
commit 66ec006baa
2 changed files with 88 additions and 4 deletions
+16 -4
View File
@@ -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);
+72
View File
@@ -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::<u64>()));
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();
}