Skip to content

Commit 3866939

Browse files
committed
Allow custom responses for malformed request headers
1 parent 702f690 commit 3866939

2 files changed

Lines changed: 265 additions & 2 deletions

File tree

pingora-proxy/src/lib.rs

Lines changed: 228 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -352,8 +352,8 @@ where
352352
"Fail to proxy: {e}, downstream session type: {}",
353353
downstream_session.session_type()
354354
);
355-
downstream_session
356-
.respond_error(400)
355+
self.inner
356+
.request_error_filter(&mut downstream_session, &e)
357357
.await
358358
.unwrap_or_else(|e| {
359359
error!("failed to send error response to downstream: {e}");
@@ -1891,6 +1891,7 @@ mod tests {
18911891
read_buf: Vec<u8>,
18921892
read_pos: usize,
18931893
write_buf: Arc<Mutex<Vec<u8>>>,
1894+
shutdown: Arc<AtomicBool>,
18941895
}
18951896

18961897
impl StaticVirtualSocket {
@@ -1899,6 +1900,7 @@ mod tests {
18991900
read_buf: read_buf.to_vec(),
19001901
read_pos: 0,
19011902
write_buf,
1903+
shutdown: Arc::new(AtomicBool::new(false)),
19021904
}
19031905
}
19041906
}
@@ -1934,6 +1936,7 @@ mod tests {
19341936
}
19351937

19361938
fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
1939+
self.shutdown.store(true, Ordering::Relaxed);
19371940
Poll::Ready(Ok(()))
19381941
}
19391942
}
@@ -1977,6 +1980,229 @@ mod tests {
19771980
}
19781981
}
19791982

1983+
fn unread_request_session(
1984+
request: &[u8],
1985+
) -> (HttpSession, Arc<Mutex<Vec<u8>>>, Arc<AtomicBool>) {
1986+
let written = Arc::new(Mutex::new(Vec::new()));
1987+
let socket = StaticVirtualSocket::new(request, written.clone());
1988+
let shutdown = socket.shutdown.clone();
1989+
let stream = L4Stream::from(VirtualSocketStream::new(Box::new(socket)));
1990+
(HttpSession::new_http1(Box::new(stream)), written, shutdown)
1991+
}
1992+
1993+
const MALFORMED_REQUEST: &[u8] = b"GET / HTTP/1.1\r\nBad Header: value\r\n\r\n";
1994+
const INVALID_REQUEST: &[u8] =
1995+
b"GET / HTTP/1.1\r\nHost: example.com\r\nContent-Length: 1\r\nContent-Length: 2\r\n\r\n";
1996+
const CUSTOM_ERROR_BODY: &[u8] = b"bad request\n";
1997+
1998+
enum RequestErrorAction {
1999+
Respond,
2000+
Close,
2001+
Fail,
2002+
RespondThenFail,
2003+
}
2004+
2005+
struct RequestErrorProxy {
2006+
calls: AtomicUsize,
2007+
action: RequestErrorAction,
2008+
}
2009+
2010+
#[async_trait]
2011+
impl ProxyHttp for RequestErrorProxy {
2012+
type CTX = ();
2013+
2014+
fn new_ctx(&self) -> Self::CTX {
2015+
panic!("rejected requests must not create a proxy context");
2016+
}
2017+
2018+
async fn upstream_peer(
2019+
&self,
2020+
_session: &mut Session,
2021+
_ctx: &mut Self::CTX,
2022+
) -> Result<Box<HttpPeer>> {
2023+
unreachable!("rejected requests must not reach upstream selection");
2024+
}
2025+
2026+
async fn request_filter(
2027+
&self,
2028+
_session: &mut Session,
2029+
_ctx: &mut Self::CTX,
2030+
) -> Result<bool> {
2031+
unreachable!("rejected requests must not reach request filters");
2032+
}
2033+
2034+
async fn logging(&self, _session: &mut Session, _e: Option<&Error>, _ctx: &mut Self::CTX) {
2035+
unreachable!("rejected requests must not reach normal logging");
2036+
}
2037+
2038+
async fn request_error_filter(&self, session: &mut HttpSession, e: &Error) -> Result<()> {
2039+
assert_eq!(e.etype(), &InvalidHTTPHeader);
2040+
assert_eq!(e.esource(), &ErrorSource::Downstream);
2041+
self.calls.fetch_add(1, Ordering::Relaxed);
2042+
// These accessors are usable even if parsing failed before a header was created.
2043+
let _ = session.client_addr();
2044+
let _ = session.digest();
2045+
if matches!(self.action, RequestErrorAction::Close) {
2046+
return Ok(());
2047+
}
2048+
if matches!(self.action, RequestErrorAction::Fail) {
2049+
return Err(Error::new(WriteError));
2050+
}
2051+
let mut response = ResponseHeader::build(422, Some(2))?;
2052+
response.set_content_length(CUSTOM_ERROR_BODY.len())?;
2053+
response.insert_header(header::CONTENT_TYPE, "text/plain")?;
2054+
session
2055+
.write_error_response(response, Bytes::from_static(CUSTOM_ERROR_BODY))
2056+
.await?;
2057+
if matches!(self.action, RequestErrorAction::RespondThenFail) {
2058+
return Err(Error::new(WriteError));
2059+
}
2060+
Ok(())
2061+
}
2062+
}
2063+
2064+
fn request_error_proxy(action: RequestErrorAction) -> Arc<HttpProxy<RequestErrorProxy>> {
2065+
Arc::new(HttpProxy::new(
2066+
RequestErrorProxy {
2067+
calls: AtomicUsize::new(0),
2068+
action,
2069+
},
2070+
Arc::new(ServerConf::default()),
2071+
))
2072+
}
2073+
2074+
fn assert_rejection_response(
2075+
written: &Mutex<Vec<u8>>,
2076+
status: u16,
2077+
body: &[u8],
2078+
content_length: usize,
2079+
) {
2080+
let written = written.lock().unwrap();
2081+
let mut headers = [httparse::EMPTY_HEADER; 16];
2082+
let mut response = httparse::Response::new(&mut headers);
2083+
let httparse::Status::Complete(header_len) = response.parse(&written).unwrap() else {
2084+
panic!("incomplete error response");
2085+
};
2086+
assert_eq!(response.code, Some(status));
2087+
let header = |name: &str| {
2088+
response
2089+
.headers
2090+
.iter()
2091+
.find(|h| h.name.eq_ignore_ascii_case(name))
2092+
.map(|h| h.value)
2093+
};
2094+
assert_eq!(header("connection"), Some(b"close".as_slice()));
2095+
assert_eq!(
2096+
header("content-length"),
2097+
Some(content_length.to_string().as_bytes())
2098+
);
2099+
if status == 422 {
2100+
assert_eq!(header("content-type"), Some(b"text/plain".as_slice()));
2101+
assert!(header("server").is_none());
2102+
} else {
2103+
assert!(header("server").is_some());
2104+
}
2105+
// Also rules out an appended fallback response or a response to pipelined input.
2106+
assert_eq!(&written[header_len..], body);
2107+
}
2108+
2109+
#[tokio::test]
2110+
async fn request_error_filter_defaults_to_400() {
2111+
let proxy = HttpProxy::new(DefaultRetryProxy, Arc::new(ServerConf::default()));
2112+
for request in [MALFORMED_REQUEST, INVALID_REQUEST] {
2113+
let (session, written, shutdown) = unread_request_session(request);
2114+
assert!(proxy.handle_new_request(Box::new(session)).await.is_none());
2115+
assert!(shutdown.load(Ordering::Relaxed));
2116+
assert_rejection_response(&written, 400, b"", 0);
2117+
}
2118+
}
2119+
2120+
#[tokio::test]
2121+
async fn request_error_filter_customizes_rejected_requests_before_context_creation() {
2122+
for request in [MALFORMED_REQUEST, INVALID_REQUEST] {
2123+
let proxy = request_error_proxy(RequestErrorAction::Respond);
2124+
let (session, written, closed) = unread_request_session(request);
2125+
let (_shutdown_tx, shutdown) = tokio::sync::watch::channel(false);
2126+
assert!(proxy.process_new_http(session, &shutdown).await.is_none());
2127+
assert_eq!(proxy.inner.calls.load(Ordering::Relaxed), 1);
2128+
assert!(closed.load(Ordering::Relaxed));
2129+
assert_rejection_response(&written, 422, CUSTOM_ERROR_BODY, CUSTOM_ERROR_BODY.len());
2130+
}
2131+
}
2132+
2133+
#[tokio::test]
2134+
async fn request_error_filter_preserves_head_body_suppression() {
2135+
let request = b"HEAD / HTTP/1.1\r\nHost: example.com\r\nContent-Length: 1\r\nContent-Length: 2\r\n\r\n";
2136+
let proxy = request_error_proxy(RequestErrorAction::Respond);
2137+
let (session, written, shutdown) = unread_request_session(request);
2138+
assert!(proxy.handle_new_request(Box::new(session)).await.is_none());
2139+
assert_eq!(proxy.inner.calls.load(Ordering::Relaxed), 1);
2140+
assert!(shutdown.load(Ordering::Relaxed));
2141+
assert_rejection_response(&written, 422, b"", CUSTOM_ERROR_BODY.len());
2142+
}
2143+
2144+
#[tokio::test]
2145+
async fn request_error_filter_closes_without_a_fallback_response() {
2146+
for action in [RequestErrorAction::Close, RequestErrorAction::Fail] {
2147+
let proxy = request_error_proxy(action);
2148+
let (session, written, shutdown) = unread_request_session(MALFORMED_REQUEST);
2149+
assert!(proxy.handle_new_request(Box::new(session)).await.is_none());
2150+
assert_eq!(proxy.inner.calls.load(Ordering::Relaxed), 1);
2151+
assert!(shutdown.load(Ordering::Relaxed));
2152+
assert!(written.lock().unwrap().is_empty());
2153+
}
2154+
}
2155+
2156+
#[tokio::test]
2157+
async fn request_error_filter_does_not_append_a_response_after_callback_failure() {
2158+
let proxy = request_error_proxy(RequestErrorAction::RespondThenFail);
2159+
let (session, written, shutdown) = unread_request_session(MALFORMED_REQUEST);
2160+
assert!(proxy.handle_new_request(Box::new(session)).await.is_none());
2161+
assert_eq!(proxy.inner.calls.load(Ordering::Relaxed), 1);
2162+
assert!(shutdown.load(Ordering::Relaxed));
2163+
assert_rejection_response(&written, 422, CUSTOM_ERROR_BODY, CUSTOM_ERROR_BODY.len());
2164+
}
2165+
2166+
#[tokio::test]
2167+
async fn request_error_filter_does_not_process_pipelined_input_after_rejection() {
2168+
let mut request = MALFORMED_REQUEST.to_vec();
2169+
request.extend_from_slice(b"GET / HTTP/1.1\r\nHost: example.com\r\n\r\n");
2170+
let proxy = request_error_proxy(RequestErrorAction::Respond);
2171+
let (session, written, closed) = unread_request_session(&request);
2172+
let (_shutdown_tx, shutdown) = tokio::sync::watch::channel(false);
2173+
assert!(proxy.process_new_http(session, &shutdown).await.is_none());
2174+
assert_eq!(proxy.inner.calls.load(Ordering::Relaxed), 1);
2175+
assert!(closed.load(Ordering::Relaxed));
2176+
assert_rejection_response(&written, 422, CUSTOM_ERROR_BODY, CUSTOM_ERROR_BODY.len());
2177+
}
2178+
2179+
#[tokio::test]
2180+
async fn request_error_filter_skips_valid_requests_and_connection_errors() {
2181+
for request in [
2182+
b"GET / HTTP/1.1\r\nHost: example.com\r\n\r\n".as_slice(),
2183+
b"",
2184+
b"GET / HTTP/1.1\r\nHost:",
2185+
] {
2186+
let proxy = request_error_proxy(RequestErrorAction::Respond);
2187+
let (session, written, shutdown) = unread_request_session(request);
2188+
let result = proxy.handle_new_request(Box::new(session)).await;
2189+
assert_eq!(result.is_some(), request.ends_with(b"\r\n\r\n"));
2190+
assert_eq!(proxy.inner.calls.load(Ordering::Relaxed), 0);
2191+
assert!(written.lock().unwrap().is_empty());
2192+
if request.ends_with(b"Host:") {
2193+
assert!(shutdown.load(Ordering::Relaxed));
2194+
}
2195+
}
2196+
}
2197+
2198+
#[tokio::test]
2199+
async fn request_error_filter_skips_shutdown() {
2200+
let proxy = request_error_proxy(RequestErrorAction::Respond);
2201+
proxy.http_cleanup().await;
2202+
assert!(proxy.handle_new_request(pending_session()).await.is_none());
2203+
assert_eq!(proxy.inner.calls.load(Ordering::Relaxed), 0);
2204+
}
2205+
19802206
fn default_policy_would_retry_for_session(
19812207
session: &mut Session,
19822208
retry: RetryType,

pingora-proxy/src/proxy_trait.rs

Lines changed: 37 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -620,6 +620,43 @@ pub trait ProxyHttp {
620620
e
621621
}
622622

623+
/// Handle a request header rejected with [`InvalidHTTPHeader`].
624+
///
625+
/// This runs before a proxy [`Session`] or [`Self::CTX`] is created. It covers
626+
/// the malformed-header rejection in the HTTP/1 request-reading path, not all
627+
/// protocol errors. The default implementation sends a 400 response.
628+
///
629+
/// The request header may be unavailable or invalid. Request-dependent methods,
630+
/// including `req_header()` and body-reading methods, may panic and must not be
631+
/// called here. Connection metadata such as `client_addr()` and `digest()` can
632+
/// be used with `e` for logging. Normal request filters and [`Self::logging()`]
633+
/// are not called for this rejected request.
634+
///
635+
/// Override this callback to write a custom response with
636+
/// [`HttpSession::write_error_response()`], setting `Content-Length` to match
637+
/// the body, or return `Ok(())` without writing to close silently. The connection
638+
/// is closed after this callback returns, including on error. A returned error
639+
/// is logged without attempting a fallback response.
640+
///
641+
/// For example, a callback can replace the default response headers and body:
642+
///
643+
/// ```
644+
/// # use bytes::Bytes;
645+
/// # use pingora_core::protocols::http::ServerSession;
646+
/// # use pingora_error::Result;
647+
/// # use pingora_http::ResponseHeader;
648+
/// # async fn custom_error(session: &mut ServerSession) -> Result<()> {
649+
/// let body = Bytes::from_static(b"Invalid request\n");
650+
/// let mut response = ResponseHeader::build(400, Some(2))?;
651+
/// response.set_content_length(body.len())?;
652+
/// response.insert_header("Content-Type", "text/plain")?;
653+
/// session.write_error_response(response, body).await
654+
/// # }
655+
/// ```
656+
async fn request_error_filter(&self, session: &mut HttpSession, _e: &Error) -> Result<()> {
657+
session.respond_error(400).await
658+
}
659+
623660
/// This filter is called when the request encounters a fatal error.
624661
///
625662
/// Users may write an error response to the downstream if the downstream is still writable.

0 commit comments

Comments
 (0)