Tiny fix non-PD router http header missing whitelist (#16339)
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
This commit is contained in:
co-authored by
gemini-code-assist[bot]
parent
dcacc492d0
commit
66dfb8c156
@@ -203,6 +203,8 @@ def create_app(args: argparse.Namespace) -> FastAPI:
|
|||||||
except (json.JSONDecodeError, ValueError):
|
except (json.JSONDecodeError, ValueError):
|
||||||
data = {}
|
data = {}
|
||||||
|
|
||||||
|
received_headers = {k.lower(): v for k, v in request.headers.items()}
|
||||||
|
|
||||||
now = time.time()
|
now = time.time()
|
||||||
ret = {
|
ret = {
|
||||||
"id": f"cmpl-{int(now*1000)}",
|
"id": f"cmpl-{int(now*1000)}",
|
||||||
@@ -218,6 +220,7 @@ def create_app(args: argparse.Namespace) -> FastAPI:
|
|||||||
],
|
],
|
||||||
"worker_id": worker_id,
|
"worker_id": worker_id,
|
||||||
"echo": data,
|
"echo": data,
|
||||||
|
"received_headers": received_headers,
|
||||||
}
|
}
|
||||||
return make_json_response(ret, status_code=200)
|
return make_json_response(ret, status_code=200)
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,36 @@
|
|||||||
|
import pytest
|
||||||
|
import requests
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
def test_header_forwarding_whitelist(mock_workers, router_manager):
|
||||||
|
_, urls, _ = mock_workers(n=1)
|
||||||
|
rh = router_manager.start_router(worker_urls=urls)
|
||||||
|
|
||||||
|
with requests.Session() as s:
|
||||||
|
r = s.post(
|
||||||
|
f"{rh.url}/v1/completions",
|
||||||
|
json={"model": "test", "prompt": "hi", "max_tokens": 1, "stream": False},
|
||||||
|
headers={
|
||||||
|
"Authorization": "Bearer test-token",
|
||||||
|
"X-SMG-Routing-Key": "routing-123",
|
||||||
|
"X-Request-Id": "req-456",
|
||||||
|
"X-Correlation-Id": "corr-789",
|
||||||
|
"traceparent": "00-trace-span-01",
|
||||||
|
"tracestate": "vendor=value",
|
||||||
|
"X-Custom-Header": "should-not-forward",
|
||||||
|
"Cookie": "session=abc",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
assert r.status_code == 200
|
||||||
|
h = r.json().get("received_headers", {})
|
||||||
|
|
||||||
|
assert h.get("authorization") == "Bearer test-token"
|
||||||
|
assert h.get("x-request-id") == "req-456"
|
||||||
|
assert h.get("x-correlation-id") == "corr-789"
|
||||||
|
assert h.get("traceparent") == "00-trace-span-01"
|
||||||
|
assert h.get("tracestate") == "vendor=value"
|
||||||
|
|
||||||
|
assert "x-smg-routing-key" not in h
|
||||||
|
assert "x-custom-header" not in h
|
||||||
|
assert "cookie" not in h
|
||||||
@@ -183,3 +183,54 @@ pub fn extract_auth_header(
|
|||||||
.and_then(|k| HeaderValue::from_str(&format!("Bearer {}", k)).ok())
|
.and_then(|k| HeaderValue::from_str(&format!("Bearer {}", k)).ok())
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[inline]
|
||||||
|
pub fn should_forward_request_header(name: &str) -> bool {
|
||||||
|
let lower_name = name.to_ascii_lowercase();
|
||||||
|
matches!(
|
||||||
|
lower_name.as_str(),
|
||||||
|
"authorization" | "x-request-id" | "x-correlation-id" | "traceparent" | "tracestate"
|
||||||
|
) || lower_name.starts_with("x-request-id-")
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::*;
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_should_forward_request_header_whitelist() {
|
||||||
|
assert!(should_forward_request_header("authorization"));
|
||||||
|
assert!(should_forward_request_header("Authorization"));
|
||||||
|
assert!(should_forward_request_header("AUTHORIZATION"));
|
||||||
|
assert!(should_forward_request_header("x-request-id"));
|
||||||
|
assert!(should_forward_request_header("X-Request-Id"));
|
||||||
|
assert!(should_forward_request_header("x-correlation-id"));
|
||||||
|
assert!(should_forward_request_header("X-Correlation-ID"));
|
||||||
|
assert!(should_forward_request_header("traceparent"));
|
||||||
|
assert!(should_forward_request_header("Traceparent"));
|
||||||
|
assert!(should_forward_request_header("tracestate"));
|
||||||
|
assert!(should_forward_request_header("Tracestate"));
|
||||||
|
assert!(should_forward_request_header("x-request-id-user"));
|
||||||
|
assert!(should_forward_request_header("X-Request-ID-Span"));
|
||||||
|
assert!(should_forward_request_header("x-request-id-123"));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_should_forward_request_header_blocked() {
|
||||||
|
assert!(!should_forward_request_header("content-type"));
|
||||||
|
assert!(!should_forward_request_header("Content-Type"));
|
||||||
|
assert!(!should_forward_request_header("content-length"));
|
||||||
|
assert!(!should_forward_request_header("host"));
|
||||||
|
assert!(!should_forward_request_header("Host"));
|
||||||
|
assert!(!should_forward_request_header("connection"));
|
||||||
|
assert!(!should_forward_request_header("transfer-encoding"));
|
||||||
|
assert!(!should_forward_request_header("accept"));
|
||||||
|
assert!(!should_forward_request_header("accept-encoding"));
|
||||||
|
assert!(!should_forward_request_header("user-agent"));
|
||||||
|
assert!(!should_forward_request_header("cookie"));
|
||||||
|
assert!(!should_forward_request_header("x-custom-header"));
|
||||||
|
assert!(!should_forward_request_header("x-api-key"));
|
||||||
|
assert!(!should_forward_request_header("x-smg-routing-key"));
|
||||||
|
assert!(!should_forward_request_header("X-SMG-Routing-Key"));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -1045,17 +1045,7 @@ impl PDRouter {
|
|||||||
}
|
}
|
||||||
if let Some(headers) = headers {
|
if let Some(headers) = headers {
|
||||||
for (name, value) in headers.iter() {
|
for (name, value) in headers.iter() {
|
||||||
let name_lc = name.as_str().to_ascii_lowercase();
|
if header_utils::should_forward_request_header(name.as_str()) {
|
||||||
// Whitelist important end-to-end headers, skip hop-by-hop
|
|
||||||
let forward = matches!(
|
|
||||||
name_lc.as_str(),
|
|
||||||
"authorization"
|
|
||||||
| "x-request-id"
|
|
||||||
| "x-correlation-id"
|
|
||||||
| "traceparent" // W3C Trace Context
|
|
||||||
| "tracestate" // W3C Trace Context
|
|
||||||
) || name_lc.starts_with("x-request-id-");
|
|
||||||
if forward {
|
|
||||||
if let Ok(val) = value.to_str() {
|
if let Ok(val) = value.to_str() {
|
||||||
request = request.header(name, val);
|
request = request.header(name, val);
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -3,10 +3,7 @@ use std::{sync::Arc, time::Instant};
|
|||||||
use axum::{
|
use axum::{
|
||||||
body::{to_bytes, Body},
|
body::{to_bytes, Body},
|
||||||
extract::Request,
|
extract::Request,
|
||||||
http::{
|
http::{header::CONTENT_TYPE, HeaderMap, HeaderValue, Method, StatusCode},
|
||||||
header::{CONTENT_LENGTH, CONTENT_TYPE},
|
|
||||||
HeaderMap, HeaderValue, Method, StatusCode,
|
|
||||||
},
|
|
||||||
response::{IntoResponse, Response},
|
response::{IntoResponse, Response},
|
||||||
Json,
|
Json,
|
||||||
};
|
};
|
||||||
@@ -91,7 +88,6 @@ impl Router {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Helper method to proxy GET requests to the first available worker
|
|
||||||
async fn proxy_get_request(&self, req: Request<Body>, endpoint: &str) -> Response {
|
async fn proxy_get_request(&self, req: Request<Body>, endpoint: &str) -> Response {
|
||||||
let headers = header_utils::copy_request_headers(&req);
|
let headers = header_utils::copy_request_headers(&req);
|
||||||
|
|
||||||
@@ -99,10 +95,7 @@ impl Router {
|
|||||||
Ok(worker_url) => {
|
Ok(worker_url) => {
|
||||||
let mut request_builder = self.client.get(format!("{}/{}", worker_url, endpoint));
|
let mut request_builder = self.client.get(format!("{}/{}", worker_url, endpoint));
|
||||||
for (name, value) in headers {
|
for (name, value) in headers {
|
||||||
// Use eq_ignore_ascii_case to avoid string allocation
|
if header_utils::should_forward_request_header(&name) {
|
||||||
if !name.eq_ignore_ascii_case("content-type")
|
|
||||||
&& !name.eq_ignore_ascii_case("content-length")
|
|
||||||
{
|
|
||||||
request_builder = request_builder.header(name, value);
|
request_builder = request_builder.header(name, value);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -361,14 +354,10 @@ impl Router {
|
|||||||
return error::service_unavailable("no_workers", "No available workers");
|
return error::service_unavailable("no_workers", "No available workers");
|
||||||
}
|
}
|
||||||
|
|
||||||
// Pre-filter headers once before the loop to avoid repeated lowercasing
|
|
||||||
let filtered_headers: Vec<_> = headers
|
let filtered_headers: Vec<_> = headers
|
||||||
.map(|hdrs| {
|
.map(|hdrs| {
|
||||||
hdrs.iter()
|
hdrs.iter()
|
||||||
.filter(|(name, _)| {
|
.filter(|(name, _)| header_utils::should_forward_request_header(name.as_str()))
|
||||||
!name.as_str().eq_ignore_ascii_case("content-type")
|
|
||||||
&& !name.as_str().eq_ignore_ascii_case("content-length")
|
|
||||||
})
|
|
||||||
.collect()
|
.collect()
|
||||||
})
|
})
|
||||||
.unwrap_or_default();
|
.unwrap_or_default();
|
||||||
@@ -542,11 +531,9 @@ impl Router {
|
|||||||
request_builder = request_builder.header("Authorization", auth_header);
|
request_builder = request_builder.header("Authorization", auth_header);
|
||||||
}
|
}
|
||||||
|
|
||||||
// Copy all headers from original request if provided
|
|
||||||
if let Some(headers) = headers {
|
if let Some(headers) = headers {
|
||||||
for (name, value) in headers {
|
for (name, value) in headers {
|
||||||
// Skip Content-Type and Content-Length as .json() sets them
|
if header_utils::should_forward_request_header(name.as_str()) {
|
||||||
if *name != CONTENT_TYPE && *name != CONTENT_LENGTH {
|
|
||||||
request_builder = request_builder.header(name, value);
|
request_builder = request_builder.header(name, value);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user