mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-05 00:47:48 +08:00
refactor(workspace): enforce layered crate boundaries
This commit is contained in:
@@ -0,0 +1,399 @@
|
||||
//! Bounded request-body buffering for frontdoor adapters.
|
||||
//!
|
||||
//! The policy reserves weighted memory before reading a body and holds the
|
||||
//! reservation through the caller's normalization callback. This keeps body
|
||||
//! buffering independent from gateway business routing while preventing a
|
||||
//! burst of compressed requests from bypassing the memory budget.
|
||||
|
||||
use axum::body::{to_bytes, Body};
|
||||
use bytes::Bytes;
|
||||
use http::{header, HeaderMap, StatusCode};
|
||||
use std::error::Error as StdError;
|
||||
use std::sync::Arc;
|
||||
use std::time::{Duration, Instant};
|
||||
use tokio::sync::{OwnedSemaphorePermit, Semaphore};
|
||||
|
||||
pub const DEFAULT_BODY_BUFFER_PERMIT_BYTES: usize = 64 * 1024;
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct BodyBufferPolicy {
|
||||
max_bytes: u64,
|
||||
read_timeout: Duration,
|
||||
queue_timeout: Duration,
|
||||
budget_bytes: usize,
|
||||
permit_bytes: usize,
|
||||
budget: Arc<Semaphore>,
|
||||
}
|
||||
|
||||
impl BodyBufferPolicy {
|
||||
pub fn new(
|
||||
max_bytes: u64,
|
||||
read_timeout: Duration,
|
||||
queue_timeout: Duration,
|
||||
budget_bytes: usize,
|
||||
budget: Arc<Semaphore>,
|
||||
) -> Self {
|
||||
Self::with_permit_bytes(
|
||||
max_bytes,
|
||||
read_timeout,
|
||||
queue_timeout,
|
||||
budget_bytes,
|
||||
DEFAULT_BODY_BUFFER_PERMIT_BYTES,
|
||||
budget,
|
||||
)
|
||||
}
|
||||
|
||||
pub fn with_permit_bytes(
|
||||
max_bytes: u64,
|
||||
read_timeout: Duration,
|
||||
queue_timeout: Duration,
|
||||
budget_bytes: usize,
|
||||
permit_bytes: usize,
|
||||
budget: Arc<Semaphore>,
|
||||
) -> Self {
|
||||
Self {
|
||||
max_bytes,
|
||||
read_timeout,
|
||||
queue_timeout,
|
||||
budget_bytes,
|
||||
permit_bytes: permit_bytes.max(1),
|
||||
budget,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn max_bytes(&self) -> u64 {
|
||||
self.max_bytes
|
||||
}
|
||||
|
||||
pub fn read_timeout(&self) -> Duration {
|
||||
self.read_timeout
|
||||
}
|
||||
|
||||
pub fn queue_timeout(&self) -> Duration {
|
||||
self.queue_timeout
|
||||
}
|
||||
|
||||
pub fn budget_bytes(&self) -> usize {
|
||||
self.budget_bytes
|
||||
}
|
||||
|
||||
pub fn reservation_bytes(&self, headers: &HeaderMap) -> usize {
|
||||
reservation_bytes(headers, self.max_bytes)
|
||||
}
|
||||
|
||||
pub fn reservation_permits(&self, reservation_bytes: usize) -> u32 {
|
||||
reservation_permits(reservation_bytes, self.permit_bytes)
|
||||
}
|
||||
|
||||
pub async fn reserve(
|
||||
&self,
|
||||
headers: &HeaderMap,
|
||||
) -> Result<BodyBufferReservation, BodyBufferError> {
|
||||
if let Some(declared) = declared_content_length(headers) {
|
||||
if declared > self.max_bytes {
|
||||
return Err(BodyBufferError::TooLarge {
|
||||
limit_bytes: self.max_bytes,
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
let requested_bytes = self.reservation_bytes(headers);
|
||||
let permits = self.reservation_permits(requested_bytes);
|
||||
let timeout_ms = duration_millis(self.queue_timeout);
|
||||
let permit = match tokio::time::timeout(
|
||||
self.queue_timeout,
|
||||
Arc::clone(&self.budget).acquire_many_owned(permits),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(Ok(permit)) => permit,
|
||||
Ok(Err(_)) | Err(_) => {
|
||||
return Err(BodyBufferError::Overloaded {
|
||||
requested_bytes,
|
||||
budget_bytes: self.budget_bytes,
|
||||
timeout_ms,
|
||||
});
|
||||
}
|
||||
};
|
||||
|
||||
Ok(BodyBufferReservation {
|
||||
permit,
|
||||
max_bytes: self.max_bytes,
|
||||
read_timeout: self.read_timeout,
|
||||
requested_bytes,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub struct BodyBufferReservation {
|
||||
permit: OwnedSemaphorePermit,
|
||||
max_bytes: u64,
|
||||
read_timeout: Duration,
|
||||
requested_bytes: usize,
|
||||
}
|
||||
|
||||
impl BodyBufferReservation {
|
||||
pub fn requested_bytes(&self) -> usize {
|
||||
self.requested_bytes
|
||||
}
|
||||
|
||||
pub async fn collect(self, body: Body) -> Result<BufferedBody, BodyBufferError> {
|
||||
let Self {
|
||||
permit,
|
||||
max_bytes,
|
||||
read_timeout,
|
||||
requested_bytes,
|
||||
} = self;
|
||||
let started_at = Instant::now();
|
||||
let body_limit = usize::try_from(max_bytes).unwrap_or(usize::MAX);
|
||||
let bytes = match tokio::time::timeout(read_timeout, to_bytes(body, body_limit)).await {
|
||||
Ok(Ok(bytes)) => bytes,
|
||||
Ok(Err(error)) if collection_exceeded_limit(&error) => {
|
||||
return Err(BodyBufferError::TooLarge {
|
||||
limit_bytes: max_bytes,
|
||||
});
|
||||
}
|
||||
Ok(Err(error)) => {
|
||||
return Err(BodyBufferError::ReadFailed {
|
||||
message: error.to_string(),
|
||||
});
|
||||
}
|
||||
Err(_) => {
|
||||
return Err(BodyBufferError::Timeout {
|
||||
timeout_ms: duration_millis(read_timeout),
|
||||
});
|
||||
}
|
||||
};
|
||||
|
||||
Ok(BufferedBody {
|
||||
bytes,
|
||||
permit: Some(permit),
|
||||
requested_bytes,
|
||||
elapsed: started_at.elapsed(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub struct BufferedBody {
|
||||
bytes: Bytes,
|
||||
permit: Option<OwnedSemaphorePermit>,
|
||||
requested_bytes: usize,
|
||||
elapsed: Duration,
|
||||
}
|
||||
|
||||
impl BufferedBody {
|
||||
pub fn bytes(&self) -> &Bytes {
|
||||
&self.bytes
|
||||
}
|
||||
|
||||
pub fn requested_bytes(&self) -> usize {
|
||||
self.requested_bytes
|
||||
}
|
||||
|
||||
pub fn elapsed(&self) -> Duration {
|
||||
self.elapsed
|
||||
}
|
||||
|
||||
/// Apply normalization while retaining the memory permit until the
|
||||
/// callback completes.
|
||||
pub fn try_map<T, E>(self, map: impl FnOnce(Bytes) -> Result<T, E>) -> Result<T, E> {
|
||||
let Self { bytes, permit, .. } = self;
|
||||
let result = map(bytes);
|
||||
drop(permit);
|
||||
result
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, PartialEq, Eq)]
|
||||
pub enum BodyBufferError {
|
||||
TooLarge {
|
||||
limit_bytes: u64,
|
||||
},
|
||||
Overloaded {
|
||||
requested_bytes: usize,
|
||||
budget_bytes: usize,
|
||||
timeout_ms: u64,
|
||||
},
|
||||
Timeout {
|
||||
timeout_ms: u64,
|
||||
},
|
||||
ReadFailed {
|
||||
message: String,
|
||||
},
|
||||
}
|
||||
|
||||
impl BodyBufferError {
|
||||
pub fn http_status(&self) -> StatusCode {
|
||||
match self {
|
||||
Self::TooLarge { .. } => StatusCode::PAYLOAD_TOO_LARGE,
|
||||
Self::Overloaded { .. } => StatusCode::SERVICE_UNAVAILABLE,
|
||||
Self::Timeout { .. } => StatusCode::REQUEST_TIMEOUT,
|
||||
Self::ReadFailed { .. } => StatusCode::BAD_REQUEST,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn client_message(&self) -> String {
|
||||
match self {
|
||||
Self::TooLarge { limit_bytes } => format!("Request body exceeds {limit_bytes} bytes"),
|
||||
Self::Overloaded { .. } => {
|
||||
"Request body buffering capacity is temporarily exhausted".to_string()
|
||||
}
|
||||
Self::Timeout { .. } => {
|
||||
"Request body read timed out before the gateway could route the request".to_string()
|
||||
}
|
||||
Self::ReadFailed { .. } => "Failed to read request body".to_string(),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn reason(&self) -> &'static str {
|
||||
match self {
|
||||
Self::TooLarge { .. } => "request_body_too_large",
|
||||
Self::Overloaded { .. } => "request_body_buffer_overloaded",
|
||||
Self::Timeout { .. } => "request_body_read_timeout",
|
||||
Self::ReadFailed { .. } => "request_body_read_failed",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn declared_content_length(headers: &HeaderMap) -> Option<u64> {
|
||||
headers
|
||||
.get(header::CONTENT_LENGTH)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.and_then(|value| value.trim().parse::<u64>().ok())
|
||||
}
|
||||
|
||||
fn reservation_bytes(headers: &HeaderMap, max_bytes: u64) -> usize {
|
||||
let max_bytes = usize::try_from(max_bytes).unwrap_or(usize::MAX);
|
||||
let encoded = headers
|
||||
.get(header::CONTENT_ENCODING)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.map(str::trim)
|
||||
.is_some_and(|value| !value.is_empty() && !value.eq_ignore_ascii_case("identity"));
|
||||
if encoded {
|
||||
return max_bytes;
|
||||
}
|
||||
declared_content_length(headers)
|
||||
.map(|value| usize::try_from(value).unwrap_or(usize::MAX).min(max_bytes))
|
||||
.unwrap_or(max_bytes)
|
||||
}
|
||||
|
||||
fn reservation_permits(reservation_bytes: usize, permit_bytes: usize) -> u32 {
|
||||
let permits = reservation_bytes
|
||||
.max(1)
|
||||
.saturating_add(permit_bytes.saturating_sub(1))
|
||||
/ permit_bytes.max(1);
|
||||
u32::try_from(permits).unwrap_or(u32::MAX).max(1)
|
||||
}
|
||||
|
||||
fn duration_millis(duration: Duration) -> u64 {
|
||||
u64::try_from(duration.as_millis()).unwrap_or(u64::MAX)
|
||||
}
|
||||
|
||||
fn collection_exceeded_limit(error: &(dyn StdError + 'static)) -> bool {
|
||||
let mut current = Some(error);
|
||||
while let Some(error) = current {
|
||||
if error.to_string().contains("length limit exceeded") {
|
||||
return true;
|
||||
}
|
||||
current = error.source();
|
||||
}
|
||||
false
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{BodyBufferError, BodyBufferPolicy, DEFAULT_BODY_BUFFER_PERMIT_BYTES};
|
||||
use axum::body::{Body, Bytes};
|
||||
use futures_util::{stream, StreamExt};
|
||||
use http::{header, HeaderMap, HeaderValue};
|
||||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
use tokio::sync::Semaphore;
|
||||
|
||||
fn policy(max_bytes: u64, timeout: Duration, budget: Arc<Semaphore>) -> BodyBufferPolicy {
|
||||
BodyBufferPolicy::with_permit_bytes(
|
||||
max_bytes,
|
||||
timeout,
|
||||
timeout,
|
||||
max_bytes as usize,
|
||||
DEFAULT_BODY_BUFFER_PERMIT_BYTES,
|
||||
budget,
|
||||
)
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn rejects_declared_content_length_before_reading_body() {
|
||||
let mut headers = HeaderMap::new();
|
||||
headers.insert(header::CONTENT_LENGTH, HeaderValue::from_static("6"));
|
||||
let error = policy(5, Duration::from_secs(1), Arc::new(Semaphore::new(1)))
|
||||
.reserve(&headers)
|
||||
.await
|
||||
.expect_err("declared body should be rejected");
|
||||
assert_eq!(error, BodyBufferError::TooLarge { limit_bytes: 5 });
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn holds_weighted_permit_through_normalization_callback() {
|
||||
let budget = Arc::new(Semaphore::new(1));
|
||||
let mut headers = HeaderMap::new();
|
||||
headers.insert(header::CONTENT_LENGTH, HeaderValue::from_static("2"));
|
||||
let reservation = policy(1024, Duration::from_secs(1), Arc::clone(&budget))
|
||||
.reserve(&headers)
|
||||
.await
|
||||
.expect("reservation should succeed");
|
||||
let buffered = reservation
|
||||
.collect(Body::from(Bytes::from_static(b"{}")))
|
||||
.await
|
||||
.expect("body should collect");
|
||||
assert_eq!(budget.available_permits(), 0);
|
||||
let normalized = buffered
|
||||
.try_map(Ok::<_, ()>)
|
||||
.expect("mapping should succeed");
|
||||
assert_eq!(normalized.as_ref(), b"{}");
|
||||
assert_eq!(budget.available_permits(), 1);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn rejects_chunked_body_when_collected_bytes_exceed_limit() {
|
||||
let reservation = policy(5, Duration::from_secs(1), Arc::new(Semaphore::new(1)))
|
||||
.reserve(&HeaderMap::new())
|
||||
.await
|
||||
.expect("reservation should succeed");
|
||||
let error = reservation
|
||||
.collect(Body::from(Bytes::from_static(b"abcdef")))
|
||||
.await
|
||||
.expect_err("chunked body should remain bounded while reading");
|
||||
assert_eq!(error, BodyBufferError::TooLarge { limit_bytes: 5 });
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn times_out_slow_body_reads() {
|
||||
let stream = stream::once(async { Ok::<Bytes, std::io::Error>(Bytes::from_static(b"{")) })
|
||||
.chain(stream::pending());
|
||||
let reservation = policy(1024, Duration::from_millis(5), Arc::new(Semaphore::new(1)))
|
||||
.reserve(&HeaderMap::new())
|
||||
.await
|
||||
.expect("reservation should succeed");
|
||||
let error = reservation
|
||||
.collect(Body::from_stream(stream))
|
||||
.await
|
||||
.expect_err("slow body should time out");
|
||||
assert_eq!(error, BodyBufferError::Timeout { timeout_ms: 5 });
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn rejects_when_weighted_budget_is_exhausted() {
|
||||
let budget = Arc::new(Semaphore::new(1));
|
||||
let _held = Arc::clone(&budget)
|
||||
.acquire_owned()
|
||||
.await
|
||||
.expect("test permit should be available");
|
||||
let error = policy(1024, Duration::from_millis(5), budget)
|
||||
.reserve(&HeaderMap::new())
|
||||
.await
|
||||
.expect_err("exhausted budget should fail closed");
|
||||
assert!(matches!(error, BodyBufferError::Overloaded { .. }));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,17 @@
|
||||
pub mod body;
|
||||
pub mod middleware;
|
||||
mod request_id;
|
||||
|
||||
pub use body::{
|
||||
BodyBufferError, BodyBufferPolicy, BodyBufferReservation, BufferedBody,
|
||||
DEFAULT_BODY_BUFFER_PERMIT_BYTES,
|
||||
};
|
||||
|
||||
pub use middleware::access_log::{
|
||||
access_log_middleware, sanitize_access_log_path, should_downgrade_access_log,
|
||||
GatewayRequestAcceptedAt, RequestLogEmitted,
|
||||
};
|
||||
pub use middleware::cf_headers::{
|
||||
apply_cf_header_stripping, strip_cf_headers_middleware, CfConnectingIp,
|
||||
};
|
||||
pub use request_id::short_request_id;
|
||||
@@ -0,0 +1,682 @@
|
||||
use std::time::Instant;
|
||||
|
||||
use axum::body::Body;
|
||||
use axum::extract::Request;
|
||||
use axum::http::header::{HeaderName, HeaderValue};
|
||||
use axum::http::Method;
|
||||
use axum::middleware::Next;
|
||||
use axum::response::Response;
|
||||
use tracing::{info, trace, warn};
|
||||
|
||||
use aether_ai_formats::api::sanitize_request_path_and_query;
|
||||
|
||||
use crate::request_id::short_request_id;
|
||||
|
||||
pub const TRACE_ID_HEADER: &str = "x-trace-id";
|
||||
pub const EXECUTION_PATH_HEADER: &str = "x-aether-execution-path";
|
||||
pub const CONTROL_ROUTE_CLASS_HEADER: &str = "x-aether-control-route-class";
|
||||
pub const CONTROL_REQUEST_ID_HEADER: &str = "x-aether-control-request-id";
|
||||
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
pub struct RequestLogEmitted;
|
||||
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
pub struct GatewayRequestAcceptedAt(pub Instant);
|
||||
|
||||
fn extract_or_generate_trace_id(headers: &http::HeaderMap) -> String {
|
||||
headers
|
||||
.get(TRACE_ID_HEADER)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
.unwrap_or_else(|| uuid::Uuid::new_v4().to_string())
|
||||
}
|
||||
|
||||
fn is_usage_detail_path(path: &str) -> bool {
|
||||
let Some(detail_id) = path.strip_prefix("/api/admin/usage/") else {
|
||||
return false;
|
||||
};
|
||||
!detail_id.is_empty()
|
||||
&& !detail_id.contains('/')
|
||||
&& !matches!(detail_id, "active" | "records" | "stats" | "heatmap")
|
||||
}
|
||||
|
||||
pub fn should_downgrade_access_log(method: &Method, path: &str) -> bool {
|
||||
if method != Method::GET {
|
||||
return false;
|
||||
}
|
||||
let normalized_path = path.split('?').next().unwrap_or(path);
|
||||
matches!(
|
||||
normalized_path,
|
||||
"/api/admin/usage/active"
|
||||
| "/api/users/me/usage/active"
|
||||
| "/api/admin/usage/records"
|
||||
| "/api/admin/usage/stats"
|
||||
| "/api/admin/usage/aggregation/stats"
|
||||
| "/api/admin/usage/heatmap"
|
||||
| "/api/admin/usage/cache-affinity/interval-timeline"
|
||||
| "/api/admin/usage/cache-affinity/ttl-analysis"
|
||||
| "/api/admin/usage/cache-affinity/hit-analysis"
|
||||
| "/api/admin/users"
|
||||
| "/api/admin/monitoring/cache/stats"
|
||||
| "/api/admin/monitoring/cache/model-mapping/stats"
|
||||
| "/api/admin/monitoring/cache/config"
|
||||
| "/api/admin/monitoring/cache/redis-keys"
|
||||
| "/api/admin/monitoring/cache/affinities"
|
||||
) || is_usage_detail_path(normalized_path)
|
||||
|| normalized_path.starts_with("/api/admin/monitoring/trace/")
|
||||
}
|
||||
|
||||
pub fn sanitize_access_log_path(path: &str) -> String {
|
||||
sanitize_request_path_and_query(path, None).unwrap_or_else(|| "/".to_string())
|
||||
}
|
||||
|
||||
pub async fn access_log_middleware(mut request: Request<Body>, next: Next) -> Response {
|
||||
let started_at = Instant::now();
|
||||
request
|
||||
.extensions_mut()
|
||||
.insert(GatewayRequestAcceptedAt(started_at));
|
||||
let method = request.method().clone();
|
||||
let raw_path = request
|
||||
.uri()
|
||||
.path_and_query()
|
||||
.map(|value| value.as_str().to_string())
|
||||
.unwrap_or_else(|| "/".to_string());
|
||||
let path = sanitize_access_log_path(&raw_path);
|
||||
let trace_id = extract_or_generate_trace_id(request.headers());
|
||||
if !request.headers().contains_key(TRACE_ID_HEADER) {
|
||||
request.headers_mut().insert(
|
||||
HeaderName::from_static(TRACE_ID_HEADER),
|
||||
HeaderValue::from_str(&trace_id).expect("trace id should be a valid header value"),
|
||||
);
|
||||
}
|
||||
trace!(
|
||||
event_name = "http_request_started",
|
||||
log_type = "access",
|
||||
status = "started",
|
||||
trace_id = %trace_id,
|
||||
request_id = "-",
|
||||
method = %method,
|
||||
path = %path,
|
||||
route_class = "pending",
|
||||
execution_path = "pending",
|
||||
"gateway request started"
|
||||
);
|
||||
let mut response = next.run(request).await;
|
||||
if !response.headers().contains_key(TRACE_ID_HEADER) {
|
||||
response.headers_mut().insert(
|
||||
HeaderName::from_static(TRACE_ID_HEADER),
|
||||
HeaderValue::from_str(&trace_id).expect("trace id should be a valid header value"),
|
||||
);
|
||||
}
|
||||
if response.extensions().get::<RequestLogEmitted>().is_none() {
|
||||
let route_class = response
|
||||
.headers()
|
||||
.get(CONTROL_ROUTE_CLASS_HEADER)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.unwrap_or("local");
|
||||
let execution_path = response
|
||||
.headers()
|
||||
.get(EXECUTION_PATH_HEADER)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.unwrap_or("local_route");
|
||||
let request_id = response
|
||||
.headers()
|
||||
.get(CONTROL_REQUEST_ID_HEADER)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.filter(|value| !value.trim().is_empty())
|
||||
.unwrap_or("-");
|
||||
let request_id = short_request_id(request_id);
|
||||
let status_code = response.status().as_u16();
|
||||
let elapsed_ms = started_at.elapsed().as_millis() as u64;
|
||||
if response.status().is_server_error() {
|
||||
warn!(
|
||||
event_name = "http_request_failed",
|
||||
log_type = "access",
|
||||
status = "failed",
|
||||
status_code,
|
||||
trace_id = %trace_id,
|
||||
request_id,
|
||||
method = %method,
|
||||
path = %path,
|
||||
route_class,
|
||||
execution_path,
|
||||
elapsed_ms,
|
||||
"gateway request failed"
|
||||
);
|
||||
} else if should_downgrade_access_log(&method, &path) {
|
||||
trace!(
|
||||
event_name = "http_request_completed",
|
||||
log_type = "access",
|
||||
status = "completed",
|
||||
status_code,
|
||||
trace_id = %trace_id,
|
||||
request_id,
|
||||
method = %method,
|
||||
path = %path,
|
||||
route_class,
|
||||
execution_path,
|
||||
elapsed_ms,
|
||||
"gateway completed request"
|
||||
);
|
||||
} else {
|
||||
info!(
|
||||
event_name = "http_request_completed",
|
||||
log_type = "access",
|
||||
status = "completed",
|
||||
status_code,
|
||||
trace_id = %trace_id,
|
||||
request_id,
|
||||
method = %method,
|
||||
path = %path,
|
||||
route_class,
|
||||
execution_path,
|
||||
elapsed_ms,
|
||||
"gateway completed request"
|
||||
);
|
||||
}
|
||||
}
|
||||
response
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
access_log_middleware, sanitize_access_log_path, should_downgrade_access_log,
|
||||
CONTROL_REQUEST_ID_HEADER, CONTROL_ROUTE_CLASS_HEADER, EXECUTION_PATH_HEADER,
|
||||
TRACE_ID_HEADER,
|
||||
};
|
||||
use axum::body::Body;
|
||||
use axum::http::{Method, Request, Response, StatusCode};
|
||||
use axum::routing::get;
|
||||
use axum::Router;
|
||||
use bytes::Bytes;
|
||||
use futures_util::stream;
|
||||
use std::sync::{Arc, Mutex};
|
||||
use tower::ServiceExt;
|
||||
use tracing_subscriber::filter::LevelFilter;
|
||||
use tracing_subscriber::prelude::*;
|
||||
|
||||
#[derive(Clone, Default)]
|
||||
struct SharedBuffer(Arc<Mutex<Vec<u8>>>);
|
||||
|
||||
struct SharedBufferWriter(Arc<Mutex<Vec<u8>>>);
|
||||
|
||||
impl SharedBuffer {
|
||||
fn lines(&self) -> Vec<serde_json::Value> {
|
||||
String::from_utf8(self.0.lock().expect("buffer should lock").clone())
|
||||
.expect("buffer should contain valid utf-8")
|
||||
.lines()
|
||||
.filter(|line| !line.trim().is_empty())
|
||||
.map(|line| serde_json::from_str(line).expect("json log line should parse"))
|
||||
.collect()
|
||||
}
|
||||
}
|
||||
|
||||
impl std::io::Write for SharedBufferWriter {
|
||||
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
|
||||
self.0
|
||||
.lock()
|
||||
.expect("buffer should lock")
|
||||
.extend_from_slice(buf);
|
||||
Ok(buf.len())
|
||||
}
|
||||
|
||||
fn flush(&mut self) -> std::io::Result<()> {
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
impl<'a> tracing_subscriber::fmt::writer::MakeWriter<'a> for SharedBuffer {
|
||||
type Writer = SharedBufferWriter;
|
||||
|
||||
fn make_writer(&'a self) -> Self::Writer {
|
||||
SharedBufferWriter(Arc::clone(&self.0))
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn access_log_path_redacts_credential_query_values() {
|
||||
assert_eq!(
|
||||
sanitize_access_log_path(
|
||||
"/v1beta/models/gemini-3-flash-preview:generateContent?key=secret&alt=sse&pageSize=10&token=hidden"
|
||||
),
|
||||
"/v1beta/models/gemini-3-flash-preview:generateContent?alt=sse&pageSize=10"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "current_thread")]
|
||||
async fn access_log_emits_sanitized_path() {
|
||||
let writer = SharedBuffer::default();
|
||||
let subscriber = tracing_subscriber::registry().with(
|
||||
tracing_subscriber::fmt::layer()
|
||||
.json()
|
||||
.flatten_event(true)
|
||||
.with_current_span(false)
|
||||
.with_span_list(false)
|
||||
.with_writer(writer.clone())
|
||||
.with_filter(LevelFilter::INFO),
|
||||
);
|
||||
let dispatch = tracing::Dispatch::new(subscriber);
|
||||
let _guard = tracing::dispatcher::set_default(&dispatch);
|
||||
|
||||
let app = Router::new()
|
||||
.route(
|
||||
"/v1beta/models/gemini-3-flash-preview:generateContent",
|
||||
get(|| async { Response::new(Body::empty()) }),
|
||||
)
|
||||
.layer(axum::middleware::from_fn(access_log_middleware));
|
||||
|
||||
let _response = app
|
||||
.oneshot(
|
||||
Request::builder()
|
||||
.uri("/v1beta/models/gemini-3-flash-preview:generateContent?key=secret&alt=sse")
|
||||
.body(Body::empty())
|
||||
.expect("request should build"),
|
||||
)
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
|
||||
let logs = writer.lines();
|
||||
assert_eq!(logs.len(), 1);
|
||||
assert_eq!(
|
||||
logs[0]["path"],
|
||||
"/v1beta/models/gemini-3-flash-preview:generateContent?alt=sse"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "current_thread")]
|
||||
async fn access_log_emits_completed_events_by_default() {
|
||||
let writer = SharedBuffer::default();
|
||||
let subscriber = tracing_subscriber::registry().with(
|
||||
tracing_subscriber::fmt::layer()
|
||||
.json()
|
||||
.flatten_event(true)
|
||||
.with_current_span(false)
|
||||
.with_span_list(false)
|
||||
.with_writer(writer.clone())
|
||||
.with_filter(LevelFilter::INFO),
|
||||
);
|
||||
let dispatch = tracing::Dispatch::new(subscriber);
|
||||
let _guard = tracing::dispatcher::set_default(&dispatch);
|
||||
|
||||
let app = Router::new()
|
||||
.route(
|
||||
"/ok",
|
||||
get(|| async {
|
||||
let mut response = Response::new(Body::empty());
|
||||
response.headers_mut().insert(
|
||||
CONTROL_ROUTE_CLASS_HEADER,
|
||||
"local".parse().expect("header should parse"),
|
||||
);
|
||||
response.headers_mut().insert(
|
||||
EXECUTION_PATH_HEADER,
|
||||
"local_route".parse().expect("header should parse"),
|
||||
);
|
||||
response.headers_mut().insert(
|
||||
CONTROL_REQUEST_ID_HEADER,
|
||||
"req-123".parse().expect("header should parse"),
|
||||
);
|
||||
response
|
||||
}),
|
||||
)
|
||||
.layer(axum::middleware::from_fn(access_log_middleware));
|
||||
|
||||
let response = app
|
||||
.oneshot(
|
||||
Request::builder()
|
||||
.uri("/ok")
|
||||
.body(Body::empty())
|
||||
.expect("request should build"),
|
||||
)
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
|
||||
assert!(response.headers().contains_key(TRACE_ID_HEADER));
|
||||
|
||||
let logs = writer.lines();
|
||||
assert_eq!(logs.len(), 1);
|
||||
assert_eq!(logs[0]["event_name"], "http_request_completed");
|
||||
assert_eq!(logs[0]["status"], "completed");
|
||||
assert_eq!(logs[0]["status_code"], 200);
|
||||
assert_eq!(logs[0]["request_id"], "req-123");
|
||||
assert_eq!(logs[0]["route_class"], "local");
|
||||
assert_eq!(logs[0]["execution_path"], "local_route");
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "current_thread")]
|
||||
async fn access_log_propagates_generated_trace_id_to_downstream_handler() {
|
||||
let app = Router::new()
|
||||
.route(
|
||||
"/trace",
|
||||
get(|headers: http::HeaderMap| async move {
|
||||
Response::builder()
|
||||
.status(StatusCode::OK)
|
||||
.header(
|
||||
"x-seen-trace-id",
|
||||
headers
|
||||
.get(TRACE_ID_HEADER)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.unwrap_or("-"),
|
||||
)
|
||||
.body(Body::empty())
|
||||
.expect("response should build")
|
||||
}),
|
||||
)
|
||||
.layer(axum::middleware::from_fn(access_log_middleware));
|
||||
|
||||
let response = app
|
||||
.oneshot(
|
||||
Request::builder()
|
||||
.uri("/trace")
|
||||
.body(Body::empty())
|
||||
.expect("request should build"),
|
||||
)
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
|
||||
let response_trace_id = response
|
||||
.headers()
|
||||
.get(TRACE_ID_HEADER)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.expect("response trace id should exist")
|
||||
.to_string();
|
||||
let seen_trace_id = response
|
||||
.headers()
|
||||
.get("x-seen-trace-id")
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.expect("downstream seen trace id should exist")
|
||||
.to_string();
|
||||
|
||||
assert_eq!(seen_trace_id, response_trace_id);
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "current_thread")]
|
||||
async fn access_log_shortens_long_request_ids() {
|
||||
let writer = SharedBuffer::default();
|
||||
let subscriber = tracing_subscriber::registry().with(
|
||||
tracing_subscriber::fmt::layer()
|
||||
.json()
|
||||
.flatten_event(true)
|
||||
.with_current_span(false)
|
||||
.with_span_list(false)
|
||||
.with_writer(writer.clone())
|
||||
.with_filter(LevelFilter::INFO),
|
||||
);
|
||||
let dispatch = tracing::Dispatch::new(subscriber);
|
||||
let _guard = tracing::dispatcher::set_default(&dispatch);
|
||||
|
||||
let app = Router::new()
|
||||
.route(
|
||||
"/ok",
|
||||
get(|| async {
|
||||
let mut response = Response::new(Body::empty());
|
||||
response.headers_mut().insert(
|
||||
CONTROL_ROUTE_CLASS_HEADER,
|
||||
"local".parse().expect("header should parse"),
|
||||
);
|
||||
response.headers_mut().insert(
|
||||
EXECUTION_PATH_HEADER,
|
||||
"local_route".parse().expect("header should parse"),
|
||||
);
|
||||
response.headers_mut().insert(
|
||||
CONTROL_REQUEST_ID_HEADER,
|
||||
"d07e1e94-41b8-409f-a18a-27993ae7ecb1"
|
||||
.parse()
|
||||
.expect("header should parse"),
|
||||
);
|
||||
response
|
||||
}),
|
||||
)
|
||||
.layer(axum::middleware::from_fn(access_log_middleware));
|
||||
|
||||
let _response = app
|
||||
.oneshot(
|
||||
Request::builder()
|
||||
.uri("/ok")
|
||||
.body(Body::empty())
|
||||
.expect("request should build"),
|
||||
)
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
|
||||
let logs = writer.lines();
|
||||
assert_eq!(logs[0]["request_id"], "d07e1e94");
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "current_thread")]
|
||||
async fn access_log_emits_failed_events_by_default_for_server_errors() {
|
||||
let writer = SharedBuffer::default();
|
||||
let subscriber = tracing_subscriber::registry().with(
|
||||
tracing_subscriber::fmt::layer()
|
||||
.json()
|
||||
.flatten_event(true)
|
||||
.with_current_span(false)
|
||||
.with_span_list(false)
|
||||
.with_writer(writer.clone())
|
||||
.with_filter(LevelFilter::INFO),
|
||||
);
|
||||
let dispatch = tracing::Dispatch::new(subscriber);
|
||||
let _guard = tracing::dispatcher::set_default(&dispatch);
|
||||
|
||||
let app = Router::new()
|
||||
.route(
|
||||
"/fail",
|
||||
get(|| async {
|
||||
Response::builder()
|
||||
.status(StatusCode::BAD_GATEWAY)
|
||||
.header(CONTROL_ROUTE_CLASS_HEADER, "passthrough")
|
||||
.header(EXECUTION_PATH_HEADER, "execution_runtime_sync")
|
||||
.body(Body::empty())
|
||||
.expect("response should build")
|
||||
}),
|
||||
)
|
||||
.layer(axum::middleware::from_fn(access_log_middleware));
|
||||
|
||||
let _response = app
|
||||
.oneshot(
|
||||
Request::builder()
|
||||
.uri("/fail")
|
||||
.body(Body::empty())
|
||||
.expect("request should build"),
|
||||
)
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
|
||||
let logs = writer.lines();
|
||||
assert_eq!(logs.len(), 1);
|
||||
assert_eq!(logs[0]["event_name"], "http_request_failed");
|
||||
assert_eq!(logs[0]["status"], "failed");
|
||||
assert_eq!(logs[0]["status_code"], 502);
|
||||
assert_eq!(logs[0]["route_class"], "passthrough");
|
||||
assert_eq!(logs[0]["execution_path"], "execution_runtime_sync");
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "current_thread")]
|
||||
async fn access_log_treats_client_errors_as_completed_events() {
|
||||
let writer = SharedBuffer::default();
|
||||
let subscriber = tracing_subscriber::registry().with(
|
||||
tracing_subscriber::fmt::layer()
|
||||
.json()
|
||||
.flatten_event(true)
|
||||
.with_current_span(false)
|
||||
.with_span_list(false)
|
||||
.with_writer(writer.clone())
|
||||
.with_filter(LevelFilter::INFO),
|
||||
);
|
||||
let dispatch = tracing::Dispatch::new(subscriber);
|
||||
let _guard = tracing::dispatcher::set_default(&dispatch);
|
||||
|
||||
let app = Router::new()
|
||||
.route(
|
||||
"/missing",
|
||||
get(|| async {
|
||||
Response::builder()
|
||||
.status(StatusCode::UNAUTHORIZED)
|
||||
.header(CONTROL_ROUTE_CLASS_HEADER, "auth")
|
||||
.header(EXECUTION_PATH_HEADER, "local_auth_denied")
|
||||
.body(Body::empty())
|
||||
.expect("response should build")
|
||||
}),
|
||||
)
|
||||
.layer(axum::middleware::from_fn(access_log_middleware));
|
||||
|
||||
let _response = app
|
||||
.oneshot(
|
||||
Request::builder()
|
||||
.uri("/missing")
|
||||
.body(Body::empty())
|
||||
.expect("request should build"),
|
||||
)
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
|
||||
let logs = writer.lines();
|
||||
assert_eq!(logs.len(), 1);
|
||||
assert_eq!(logs[0]["event_name"], "http_request_completed");
|
||||
assert_eq!(logs[0]["status"], "completed");
|
||||
assert_eq!(logs[0]["status_code"], 401);
|
||||
assert_eq!(logs[0]["route_class"], "auth");
|
||||
assert_eq!(logs[0]["execution_path"], "local_auth_denied");
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "current_thread")]
|
||||
async fn access_log_emits_completed_events_for_streaming_responses() {
|
||||
let writer = SharedBuffer::default();
|
||||
let subscriber = tracing_subscriber::registry().with(
|
||||
tracing_subscriber::fmt::layer()
|
||||
.json()
|
||||
.flatten_event(true)
|
||||
.with_current_span(false)
|
||||
.with_span_list(false)
|
||||
.with_writer(writer.clone())
|
||||
.with_filter(LevelFilter::INFO),
|
||||
);
|
||||
let dispatch = tracing::Dispatch::new(subscriber);
|
||||
let _guard = tracing::dispatcher::set_default(&dispatch);
|
||||
|
||||
let app = Router::new()
|
||||
.route(
|
||||
"/stream",
|
||||
get(|| async {
|
||||
let body = Body::from_stream(stream::iter(vec![
|
||||
Ok::<Bytes, std::convert::Infallible>(Bytes::from("chunk-1")),
|
||||
Ok::<Bytes, std::convert::Infallible>(Bytes::from("chunk-2")),
|
||||
]));
|
||||
Response::builder()
|
||||
.status(StatusCode::OK)
|
||||
.header(CONTROL_ROUTE_CLASS_HEADER, "ai_public")
|
||||
.header(EXECUTION_PATH_HEADER, "execution_runtime_stream")
|
||||
.header(CONTROL_REQUEST_ID_HEADER, "req-stream")
|
||||
.body(body)
|
||||
.expect("response should build")
|
||||
}),
|
||||
)
|
||||
.layer(axum::middleware::from_fn(access_log_middleware));
|
||||
|
||||
let response = app
|
||||
.oneshot(
|
||||
Request::builder()
|
||||
.uri("/stream")
|
||||
.body(Body::empty())
|
||||
.expect("request should build"),
|
||||
)
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
|
||||
let logs = writer.lines();
|
||||
assert_eq!(logs.len(), 1);
|
||||
assert_eq!(logs[0]["event_name"], "http_request_completed");
|
||||
assert_eq!(logs[0]["status_code"], 200);
|
||||
assert_eq!(logs[0]["request_id"], "req-stream");
|
||||
assert_eq!(logs[0]["route_class"], "ai_public");
|
||||
assert_eq!(logs[0]["execution_path"], "execution_runtime_stream");
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "current_thread")]
|
||||
async fn access_log_downgrades_usage_active_polling_to_trace() {
|
||||
let writer = SharedBuffer::default();
|
||||
let subscriber = tracing_subscriber::registry().with(
|
||||
tracing_subscriber::fmt::layer()
|
||||
.json()
|
||||
.flatten_event(true)
|
||||
.with_current_span(false)
|
||||
.with_span_list(false)
|
||||
.with_writer(writer.clone())
|
||||
.with_filter(LevelFilter::TRACE),
|
||||
);
|
||||
let dispatch = tracing::Dispatch::new(subscriber);
|
||||
let _guard = tracing::dispatcher::set_default(&dispatch);
|
||||
|
||||
let app = Router::new()
|
||||
.route(
|
||||
"/api/admin/usage/active",
|
||||
get(|| async {
|
||||
Response::builder()
|
||||
.status(StatusCode::OK)
|
||||
.header(CONTROL_ROUTE_CLASS_HEADER, "admin_proxy")
|
||||
.header(EXECUTION_PATH_HEADER, "public_proxy_passthrough")
|
||||
.body(Body::empty())
|
||||
.expect("response should build")
|
||||
}),
|
||||
)
|
||||
.layer(axum::middleware::from_fn(access_log_middleware));
|
||||
|
||||
let _response = app
|
||||
.oneshot(
|
||||
Request::builder()
|
||||
.uri("/api/admin/usage/active?ids=req-1")
|
||||
.body(Body::empty())
|
||||
.expect("request should build"),
|
||||
)
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
|
||||
let logs = writer.lines();
|
||||
assert_eq!(logs.len(), 2);
|
||||
assert_eq!(logs[0]["level"], "TRACE");
|
||||
assert_eq!(logs[0]["event_name"], "http_request_started");
|
||||
assert_eq!(logs[1]["level"], "TRACE");
|
||||
assert_eq!(logs[1]["event_name"], "http_request_completed");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn access_log_marks_usage_active_paths_as_high_frequency() {
|
||||
assert!(should_downgrade_access_log(
|
||||
&Method::GET,
|
||||
"/api/admin/usage/active"
|
||||
));
|
||||
assert!(should_downgrade_access_log(
|
||||
&Method::GET,
|
||||
"/api/admin/usage/active?ids=req-1"
|
||||
));
|
||||
assert!(should_downgrade_access_log(
|
||||
&Method::GET,
|
||||
"/api/users/me/usage/active"
|
||||
));
|
||||
assert!(should_downgrade_access_log(
|
||||
&Method::GET,
|
||||
"/api/admin/usage/records?limit=20"
|
||||
));
|
||||
assert!(should_downgrade_access_log(
|
||||
&Method::GET,
|
||||
"/api/admin/usage/123e4567-e89b-12d3-a456-426614174000?include_bodies=false"
|
||||
));
|
||||
assert!(should_downgrade_access_log(
|
||||
&Method::GET,
|
||||
"/api/admin/monitoring/trace/req-123?attempted_only=false"
|
||||
));
|
||||
assert!(should_downgrade_access_log(
|
||||
&Method::GET,
|
||||
"/api/admin/monitoring/cache/stats"
|
||||
));
|
||||
assert!(!should_downgrade_access_log(
|
||||
&Method::DELETE,
|
||||
"/api/admin/monitoring/cache/affinity/provider/key/model/openai:responses"
|
||||
));
|
||||
assert!(!should_downgrade_access_log(&Method::GET, "/v1/responses"));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,161 @@
|
||||
use axum::{extract::Request, middleware::Next, response::Response, Router};
|
||||
use http::{header::HeaderName, HeaderMap};
|
||||
|
||||
const CF_EXACT_HEADERS: &[&str] = &["cdn-loop", "true-client-ip"];
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct CfConnectingIp(pub String);
|
||||
|
||||
fn should_strip_cf_header(name: &HeaderName) -> bool {
|
||||
let normalized = name.as_str();
|
||||
normalized.starts_with("cf-") || CF_EXACT_HEADERS.contains(&normalized)
|
||||
}
|
||||
|
||||
fn cf_connecting_ip(headers: &HeaderMap) -> Option<String> {
|
||||
headers
|
||||
.get("cf-connecting-ip")
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(|value| value.chars().take(45).collect())
|
||||
}
|
||||
|
||||
fn strip_cf_headers(headers: &mut HeaderMap) {
|
||||
let to_remove: Vec<_> = headers
|
||||
.keys()
|
||||
.filter(|name| should_strip_cf_header(name))
|
||||
.cloned()
|
||||
.collect();
|
||||
for name in to_remove {
|
||||
headers.remove(name);
|
||||
}
|
||||
}
|
||||
|
||||
pub fn apply_cf_header_stripping(router: Router) -> Router {
|
||||
router.layer(axum::middleware::from_fn(strip_cf_headers_middleware))
|
||||
}
|
||||
|
||||
pub async fn strip_cf_headers_middleware(mut request: Request, next: Next) -> Response {
|
||||
if let Some(client_ip) = cf_connecting_ip(request.headers()) {
|
||||
request.extensions_mut().insert(CfConnectingIp(client_ip));
|
||||
}
|
||||
strip_cf_headers(request.headers_mut());
|
||||
|
||||
let mut response = next.run(request).await;
|
||||
strip_cf_headers(response.headers_mut());
|
||||
response
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use axum::body::{to_bytes, Body};
|
||||
use axum::routing::any;
|
||||
use axum::Router;
|
||||
use http::{HeaderValue, Request, Response};
|
||||
use tower::ServiceExt;
|
||||
|
||||
use super::apply_cf_header_stripping;
|
||||
|
||||
#[tokio::test]
|
||||
async fn strips_cf_prefixed_and_exact_headers_from_request_and_response() {
|
||||
let app = apply_cf_header_stripping(Router::new().route(
|
||||
"/",
|
||||
any(|headers: http::HeaderMap| async move {
|
||||
let leaked = headers.contains_key("cf-ipcity")
|
||||
|| headers.contains_key("cf-ray")
|
||||
|| headers.contains_key("cf-connecting-ip")
|
||||
|| headers.contains_key("true-client-ip")
|
||||
|| headers.contains_key("cdn-loop");
|
||||
let mut response =
|
||||
Response::new(Body::from(if leaked { "leaked" } else { "clean" }));
|
||||
response.headers_mut().insert(
|
||||
http::header::HeaderName::from_static("cf-ipcity"),
|
||||
HeaderValue::from_static("Shanghai"),
|
||||
);
|
||||
response.headers_mut().insert(
|
||||
http::header::HeaderName::from_static("cf-cache-status"),
|
||||
HeaderValue::from_static("HIT"),
|
||||
);
|
||||
response.headers_mut().insert(
|
||||
http::header::HeaderName::from_static("true-client-ip"),
|
||||
HeaderValue::from_static("1.1.1.1"),
|
||||
);
|
||||
response.headers_mut().insert(
|
||||
http::header::HeaderName::from_static("cdn-loop"),
|
||||
HeaderValue::from_static("cloudflare"),
|
||||
);
|
||||
response
|
||||
}),
|
||||
));
|
||||
|
||||
let response = app
|
||||
.oneshot(
|
||||
Request::builder()
|
||||
.uri("/")
|
||||
.header("cf-ipcity", "Shanghai")
|
||||
.header("cf-ray", "abc123")
|
||||
.header("cf-connecting-ip", "203.0.113.10")
|
||||
.header("true-client-ip", "1.1.1.1")
|
||||
.header("cdn-loop", "cloudflare")
|
||||
.body(Body::empty())
|
||||
.expect("request should build"),
|
||||
)
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
|
||||
assert!(response.headers().get("cf-ipcity").is_none());
|
||||
assert!(response.headers().get("cf-cache-status").is_none());
|
||||
assert!(response.headers().get("cf-connecting-ip").is_none());
|
||||
assert!(response.headers().get("true-client-ip").is_none());
|
||||
assert!(response.headers().get("cdn-loop").is_none());
|
||||
|
||||
let body = to_bytes(response.into_body(), usize::MAX)
|
||||
.await
|
||||
.expect("body should be readable");
|
||||
assert_eq!(body.as_ref(), b"clean");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn preserves_non_cf_headers() {
|
||||
let app = apply_cf_header_stripping(Router::new().route(
|
||||
"/",
|
||||
any(|headers: http::HeaderMap| async move {
|
||||
let mut response = Response::new(Body::from(
|
||||
headers
|
||||
.get("x-custom-header")
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.unwrap_or_default()
|
||||
.to_string(),
|
||||
));
|
||||
response.headers_mut().insert(
|
||||
http::header::HeaderName::from_static("x-custom-response"),
|
||||
HeaderValue::from_static("kept"),
|
||||
);
|
||||
response
|
||||
}),
|
||||
));
|
||||
|
||||
let response = app
|
||||
.oneshot(
|
||||
Request::builder()
|
||||
.uri("/")
|
||||
.header("x-custom-header", "kept")
|
||||
.body(Body::empty())
|
||||
.expect("request should build"),
|
||||
)
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
|
||||
assert_eq!(
|
||||
response
|
||||
.headers()
|
||||
.get("x-custom-response")
|
||||
.and_then(|value| value.to_str().ok()),
|
||||
Some("kept")
|
||||
);
|
||||
let body = to_bytes(response.into_body(), usize::MAX)
|
||||
.await
|
||||
.expect("body should be readable");
|
||||
assert_eq!(body.as_ref(), b"kept");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,2 @@
|
||||
pub mod access_log;
|
||||
pub mod cf_headers;
|
||||
@@ -0,0 +1,74 @@
|
||||
pub fn short_request_id(value: &str) -> String {
|
||||
let trimmed = value.trim();
|
||||
if trimmed.is_empty() {
|
||||
return "-".to_string();
|
||||
}
|
||||
|
||||
if trimmed.chars().count() <= 12 {
|
||||
return trimmed.to_string();
|
||||
}
|
||||
|
||||
if looks_like_uuid(trimmed) {
|
||||
return trimmed.chars().take(8).collect();
|
||||
}
|
||||
|
||||
let prefix: String = trimmed.chars().take(6).collect();
|
||||
let suffix: String = trimmed
|
||||
.chars()
|
||||
.rev()
|
||||
.take(4)
|
||||
.collect::<String>()
|
||||
.chars()
|
||||
.rev()
|
||||
.collect();
|
||||
format!("{prefix}...{suffix}")
|
||||
}
|
||||
|
||||
fn looks_like_uuid(value: &str) -> bool {
|
||||
let bytes = value.as_bytes();
|
||||
if bytes.len() != 36 {
|
||||
return false;
|
||||
}
|
||||
|
||||
for (index, byte) in bytes.iter().enumerate() {
|
||||
let is_hyphen = matches!(index, 8 | 13 | 18 | 23);
|
||||
if is_hyphen {
|
||||
if *byte != b'-' {
|
||||
return false;
|
||||
}
|
||||
continue;
|
||||
}
|
||||
|
||||
if !byte.is_ascii_hexdigit() {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
true
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::short_request_id;
|
||||
|
||||
#[test]
|
||||
fn shortens_uuid_like_request_ids_to_prefix() {
|
||||
assert_eq!(
|
||||
short_request_id("d07e1e94-41b8-409f-a18a-27993ae7ecb1"),
|
||||
"d07e1e94"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn shortens_long_named_request_ids_to_prefix_and_suffix() {
|
||||
assert_eq!(
|
||||
short_request_id("trace-openai-cli-stream-sync-direct-123"),
|
||||
"trace-...-123"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn preserves_short_request_ids() {
|
||||
assert_eq!(short_request_id("req-123"), "req-123");
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user