oidc drain
This commit is contained in:
+16
-4
@@ -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);
|
||||
|
||||
@@ -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();
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user