feat(security): harden gateway request and runtime controls

This commit is contained in:
elky
2026-07-12 14:10:54 +08:00
parent bc1da3bf3f
commit 7f61bb43c7
19 changed files with 975 additions and 340 deletions
@@ -309,11 +309,14 @@ async fn parse_request_json<T>(request: Request) -> Result<T, ExecutionRuntimeAp
where
T: serde::de::DeserializeOwned,
{
let body = to_bytes(request.into_body(), usize::MAX)
.await
.map_err(|err| {
ExecutionRuntimeAppError(ExecutionRuntimeServerError::RequestRead(err.to_string()))
})?;
let body = to_bytes(
request.into_body(),
usize::try_from(crate::headers::max_request_body_bytes()).unwrap_or(usize::MAX),
)
.await
.map_err(|err| {
ExecutionRuntimeAppError(ExecutionRuntimeServerError::RequestRead(err.to_string()))
})?;
serde_json::from_slice(&body).map_err(|err| {
ExecutionRuntimeAppError(ExecutionRuntimeServerError::InvalidRequestJson(err))
})
@@ -1534,7 +1534,12 @@ async fn openai_image_sync_json_heartbeat_final_bytes(
result: Result<Option<Response<Body>>, GatewayError>,
) -> Vec<u8> {
match result {
Ok(Some(response)) => match to_bytes(response.into_body(), usize::MAX).await {
Ok(Some(response)) => match to_bytes(
response.into_body(),
crate::headers::max_internal_buffered_body_bytes(),
)
.await
{
Ok(bytes) if !bytes.is_empty() => bytes.to_vec(),
Ok(_) => openai_image_sync_json_heartbeat_error_body("empty sync image response"),
Err(err) => openai_image_sync_json_heartbeat_error_body(&err.to_string()),
@@ -896,7 +896,7 @@ async fn standard_text_sync_heartbeat_response_body_bytes(
) -> Vec<u8> {
let status_code = response.status().as_u16();
let (parts, body) = response.into_parts();
match to_bytes(body, usize::MAX).await {
match to_bytes(body, crate::headers::max_internal_buffered_body_bytes()).await {
Ok(bytes) => {
let body = match standard_text_sync_heartbeat_restore_response_body(
redaction_slot,
@@ -1197,7 +1197,12 @@ async fn openai_image_sync_heartbeat_final_bytes(
async fn openai_image_sync_heartbeat_response_body_bytes(response: Response<Body>) -> Vec<u8> {
let status_code = response.status().as_u16();
match to_bytes(response.into_body(), usize::MAX).await {
match to_bytes(
response.into_body(),
crate::headers::max_internal_buffered_body_bytes(),
)
.await
{
Ok(bytes) if status_code < 400 && !bytes.is_empty() => bytes.to_vec(),
Ok(bytes) if status_code >= 400 => {
openai_image_sync_heartbeat_error_body_from_response(status_code, bytes.as_ref())
@@ -652,7 +652,9 @@ pub(super) async fn extract_admin_provider_oauth_batch_error_detail(
response: Response<Body>,
) -> String {
let status = response.status();
let raw_body = to_bytes(response.into_body(), usize::MAX).await.ok();
let raw_body = to_bytes(response.into_body(), crate::MAX_ERROR_BODY_BYTES)
.await
.ok();
if let Some(raw_body) = raw_body {
if let Ok(value) = serde_json::from_slice::<serde_json::Value>(&raw_body) {
if let Some(detail) = value.get("detail").and_then(serde_json::Value::as_str) {
@@ -1704,9 +1704,12 @@ async fn provider_query_finalize_kiro_result(
let Some(outcome) = maybe_build_sync_finalize_outcome(trace_id, &decision, &payload)? else {
return Ok(None);
};
let bytes = to_bytes(outcome.response.into_body(), usize::MAX)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let bytes = to_bytes(
outcome.response.into_body(),
crate::headers::max_internal_buffered_body_bytes(),
)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
serde_json::from_slice::<Value>(&bytes)
.map(Some)
.map_err(|err| GatewayError::Internal(err.to_string()))
@@ -1926,9 +1929,12 @@ async fn provider_query_finalize_windsurf_result(
let Some(outcome) = maybe_build_sync_finalize_outcome(trace_id, &decision, &payload)? else {
return Ok(None);
};
let bytes = to_bytes(outcome.response.into_body(), usize::MAX)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let bytes = to_bytes(
outcome.response.into_body(),
crate::headers::max_internal_buffered_body_bytes(),
)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
serde_json::from_slice::<Value>(&bytes)
.map(Some)
.map_err(|err| GatewayError::Internal(err.to_string()))
@@ -2044,9 +2050,12 @@ async fn provider_query_finalize_openai_image_result(
let Some(outcome) = maybe_build_sync_finalize_outcome(trace_id, &decision, &payload)? else {
return Ok(None);
};
let bytes = to_bytes(outcome.response.into_body(), usize::MAX)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let bytes = to_bytes(
outcome.response.into_body(),
crate::headers::max_internal_buffered_body_bytes(),
)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
serde_json::from_slice::<Value>(&bytes)
.map(Some)
.map_err(|err| GatewayError::Internal(err.to_string()))
@@ -2412,9 +2421,12 @@ async fn provider_query_finalize_antigravity_result(
let Some(outcome) = maybe_build_sync_finalize_outcome(trace_id, &decision, &payload)? else {
return Ok(None);
};
let bytes = to_bytes(outcome.response.into_body(), usize::MAX)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let bytes = to_bytes(
outcome.response.into_body(),
crate::headers::max_internal_buffered_body_bytes(),
)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
serde_json::from_slice::<Value>(&bytes)
.map(Some)
.map_err(|err| GatewayError::Internal(err.to_string()))
@@ -3794,9 +3806,12 @@ pub(crate) async fn build_admin_provider_query_test_model_local_response(
if !response.status().is_success() {
return Ok(response);
}
let body = to_bytes(response.into_body(), usize::MAX)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let body = to_bytes(
response.into_body(),
crate::headers::max_internal_buffered_body_bytes(),
)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let parsed: Value =
serde_json::from_slice(&body).map_err(|err| GatewayError::Internal(err.to_string()))?;
@@ -0,0 +1,310 @@
use crate::api::response::build_local_http_error_response;
use crate::control::GatewayPublicRequestContext;
use crate::headers::RequestBodyNormalizationError;
use crate::{AppState, GatewayError};
use axum::body::{to_bytes, Body, Bytes};
use axum::http::{self, Response};
use std::error::Error as StdError;
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::sync::Semaphore;
use tracing::{info, warn};
const REQUEST_BODY_READ_TIMEOUT_DETAIL: &str =
"Request body read timed out before the gateway could route the request";
const REQUEST_BODY_READ_FAILED_DETAIL: &str = "Failed to read request body";
#[derive(Debug, Clone)]
pub(super) struct RequestBodyBufferPolicy {
max_bytes: u64,
read_timeout: Duration,
queue_timeout: Duration,
budget_bytes: usize,
budget: Arc<Semaphore>,
}
impl RequestBodyBufferPolicy {
pub(super) fn from_state(state: &AppState) -> Self {
Self {
max_bytes: crate::headers::max_request_body_bytes(),
read_timeout: state.frontdoor_runtime_guards.request_body_read_timeout,
queue_timeout: state.frontdoor_runtime_guards.internal_gate_queue_budget,
budget_bytes: state
.frontdoor_runtime_guards
.request_body_buffer_budget_bytes,
budget: Arc::clone(&state.request_body_buffer_budget),
}
}
#[cfg(test)]
pub(super) fn for_tests(max_bytes: u64, read_timeout: Duration) -> Self {
let budget_bytes = usize::try_from(max_bytes)
.unwrap_or(usize::MAX)
.max(crate::state::REQUEST_BODY_BUFFER_PERMIT_BYTES);
Self {
max_bytes,
read_timeout,
queue_timeout: read_timeout,
budget_bytes,
budget: Arc::new(Semaphore::new(
budget_bytes.saturating_add(crate::state::REQUEST_BODY_BUFFER_PERMIT_BYTES - 1)
/ crate::state::REQUEST_BODY_BUFFER_PERMIT_BYTES,
)),
}
}
#[cfg(test)]
pub(super) fn for_tests_with_budget(
max_bytes: u64,
read_timeout: Duration,
queue_timeout: Duration,
budget_bytes: usize,
budget: Arc<Semaphore>,
) -> Self {
Self {
max_bytes,
read_timeout,
queue_timeout,
budget_bytes,
budget,
}
}
}
#[derive(Debug)]
pub(super) enum RequestBodyBufferError {
Normalization(RequestBodyNormalizationError),
TooLarge {
limit_bytes: u64,
},
Overloaded {
requested_bytes: usize,
budget_bytes: usize,
timeout_ms: u64,
},
Timeout {
timeout_ms: u64,
},
ReadFailed {
message: String,
},
}
impl RequestBodyBufferError {
pub(super) fn http_status(&self) -> http::StatusCode {
match self {
Self::Normalization(error) => error.http_status(),
Self::TooLarge { .. } => http::StatusCode::PAYLOAD_TOO_LARGE,
Self::Overloaded { .. } => http::StatusCode::SERVICE_UNAVAILABLE,
Self::Timeout { .. } => http::StatusCode::REQUEST_TIMEOUT,
Self::ReadFailed { .. } => http::StatusCode::BAD_REQUEST,
}
}
fn client_message(&self) -> String {
match self {
Self::Normalization(error) => error.client_message(),
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_TIMEOUT_DETAIL.to_string(),
Self::ReadFailed { .. } => REQUEST_BODY_READ_FAILED_DETAIL.to_string(),
}
}
fn reason(&self) -> &'static str {
match self {
Self::Normalization(error) => match error {
RequestBodyNormalizationError::UnsupportedContentEncoding(_) => {
"unsupported_content_encoding"
}
RequestBodyNormalizationError::DecodeFailed { .. } => "decode_failed",
RequestBodyNormalizationError::DecompressedBodyTooLarge { .. } => {
"decompressed_body_too_large"
}
RequestBodyNormalizationError::RequestBodyTooLarge { .. } => {
"request_body_too_large"
}
},
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 request_body_buffer_reservation_bytes(headers: &http::HeaderMap, max_bytes: u64) -> usize {
let max_bytes = usize::try_from(max_bytes).unwrap_or(usize::MAX);
let encoded = headers
.get(http::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;
}
headers
.get(http::header::CONTENT_LENGTH)
.and_then(|value| value.to_str().ok())
.and_then(|value| value.trim().parse::<usize>().ok())
.map(|value| value.min(max_bytes))
.unwrap_or(max_bytes)
}
fn request_body_buffer_reservation_permits(reservation_bytes: usize) -> u32 {
let permits = reservation_bytes
.max(1)
.saturating_add(crate::state::REQUEST_BODY_BUFFER_PERMIT_BYTES - 1)
/ crate::state::REQUEST_BODY_BUFFER_PERMIT_BYTES;
u32::try_from(permits).unwrap_or(u32::MAX).max(1)
}
pub(super) async fn buffer_and_normalize_request_body(
request_body: &mut Option<Body>,
headers: &mut http::HeaderMap,
body_owner_expectation: &'static str,
trace_id: &str,
method: &http::Method,
path_and_query: &str,
phase: &'static str,
policy: RequestBodyBufferPolicy,
) -> Result<Bytes, RequestBodyBufferError> {
if let Err(err) =
crate::headers::check_request_content_length_with_limit(headers, policy.max_bytes)
{
return Err(RequestBodyBufferError::Normalization(err));
}
let reservation_bytes = request_body_buffer_reservation_bytes(headers, policy.max_bytes);
let reservation_permits = request_body_buffer_reservation_permits(reservation_bytes);
let queue_timeout_ms = policy.queue_timeout.as_millis() as u64;
let _budget_permit = match tokio::time::timeout(
policy.queue_timeout,
Arc::clone(&policy.budget).acquire_many_owned(reservation_permits),
)
.await
{
Ok(Ok(permit)) => permit,
Ok(Err(_)) | Err(_) => {
return Err(RequestBodyBufferError::Overloaded {
requested_bytes: reservation_bytes,
budget_bytes: policy.budget_bytes,
timeout_ms: queue_timeout_ms,
});
}
};
let read_started_at = Instant::now();
let timeout_ms = policy.read_timeout.as_millis() as u64;
info!(
event_name = "frontdoor_request_body_buffer_started",
log_type = "event",
trace_id,
method = %method,
path = %path_and_query,
phase,
max_body_bytes = policy.max_bytes,
reserved_body_bytes = reservation_bytes,
body_buffer_budget_bytes = policy.budget_bytes,
timeout_ms,
"gateway started buffering request body"
);
let body_limit = usize::try_from(policy.max_bytes).unwrap_or(usize::MAX);
let body = match tokio::time::timeout(
policy.read_timeout,
to_bytes(
request_body.take().expect(body_owner_expectation),
body_limit,
),
)
.await
{
Ok(Ok(body)) => body,
Ok(Err(err)) if request_body_collection_exceeded_limit(&err) => {
return Err(RequestBodyBufferError::TooLarge {
limit_bytes: policy.max_bytes,
});
}
Ok(Err(err)) => {
return Err(RequestBodyBufferError::ReadFailed {
message: err.to_string(),
});
}
Err(_) => {
return Err(RequestBodyBufferError::Timeout { timeout_ms });
}
};
let normalized = crate::headers::normalize_request_body_headers_and_bytes_with_limit(
headers,
body,
policy.max_bytes,
)
.map_err(RequestBodyBufferError::Normalization)?;
info!(
event_name = "frontdoor_request_body_buffer_completed",
log_type = "event",
trace_id,
method = %method,
path = %path_and_query,
phase,
body_bytes = normalized.len(),
elapsed_ms = read_started_at.elapsed().as_millis() as u64,
"gateway completed buffering request body"
);
Ok(normalized)
}
fn request_body_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
}
pub(super) fn build_request_body_buffer_error_response(
trace_id: &str,
request_context: &GatewayPublicRequestContext,
error: &RequestBodyBufferError,
) -> Result<Response<Body>, GatewayError> {
warn!(
event_name = "frontdoor_request_body_buffer_failed",
log_type = "ops",
trace_id,
method = %request_context.request_method,
path = %request_context.request_path_and_query(),
status_code = error.http_status().as_u16(),
reason = error.reason(),
detail = %error.client_message(),
read_error = match error {
RequestBodyBufferError::ReadFailed { message } => message.as_str(),
_ => "",
},
buffer_requested_bytes = match error {
RequestBodyBufferError::Overloaded { requested_bytes, .. } => *requested_bytes,
_ => 0,
},
buffer_budget_bytes = match error {
RequestBodyBufferError::Overloaded { budget_bytes, .. } => *budget_bytes,
_ => 0,
},
buffer_queue_timeout_ms = match error {
RequestBodyBufferError::Overloaded { timeout_ms, .. } => *timeout_ms,
_ => 0,
},
"gateway rejected request body before local execution planning"
);
build_local_http_error_response(
trace_id,
request_context.control_decision.as_ref(),
error.http_status(),
error.client_message().as_str(),
)
}
+109 -244
View File
@@ -1,5 +1,10 @@
mod body_buffer;
mod local;
use self::body_buffer::{
buffer_and_normalize_request_body, build_request_body_buffer_error_response,
RequestBodyBufferError, RequestBodyBufferPolicy,
};
use self::local::{
maybe_build_local_admin_proxy_response, maybe_build_local_internal_proxy_response,
};
@@ -52,7 +57,7 @@ use crate::handlers::shared::{
};
use crate::headers::{
effective_client_ip, extract_or_generate_trace_id, request_origin_from_headers_and_remote_addr,
should_skip_request_header, RequestBodyNormalizationError,
should_skip_request_header,
};
use crate::router::RequestAdmissionError;
use crate::scheduler::candidate::{
@@ -70,11 +75,7 @@ use axum::extract::{ConnectInfo, Request, State};
use axum::http::{self, header::HeaderName, header::HeaderValue, Response};
use futures_util::StreamExt;
use sha2::{Digest, Sha256};
use std::{
collections::BTreeMap,
error::Error as StdError,
time::{Duration, Instant},
};
use std::{collections::BTreeMap, time::Instant};
use tracing::{debug, info, warn};
const OPENAI_CHAT_LOCAL_EXECUTION_RUNTIME_MISS_DETAIL: &str =
@@ -98,201 +99,11 @@ const LOCAL_EXECUTION_LOOP_DETECTED_DETAIL: &str =
"Gateway detected an execution runtime request loop back into the local frontdoor";
const AUTH_API_KEY_CONCURRENCY_LIMIT_REACHED_DETAIL: &str =
"当前调用方 API Key 并发请求数已达上限,请稍后重试";
const REQUEST_BODY_READ_TIMEOUT_DETAIL: &str =
"Request body read timed out before the gateway could route the request";
const REQUEST_BODY_READ_FAILED_DETAIL: &str = "Failed to read request body";
const LOCAL_EXECUTION_PLANNING_TIMEOUT_DETAIL: &str =
"当前 AI 请求在本地执行规划阶段超时,请稍后重试";
const EXECUTION_PATH_TUNNEL_AFFINITY_FORWARD: &str = "tunnel_affinity_forward";
const MANAGEMENT_TOKEN_PREFIX: &str = "ae-";
const LEGACY_MANAGEMENT_TOKEN_PREFIX: &str = "ae_";
#[derive(Debug, Clone, Copy)]
struct RequestBodyBufferPolicy {
max_bytes: u64,
read_timeout: Duration,
}
impl RequestBodyBufferPolicy {
fn from_state(state: &AppState) -> Self {
Self {
max_bytes: crate::headers::max_request_body_bytes(),
read_timeout: state.frontdoor_runtime_guards.request_body_read_timeout,
}
}
#[cfg(test)]
fn for_tests(max_bytes: u64, read_timeout: Duration) -> Self {
Self {
max_bytes,
read_timeout,
}
}
}
#[derive(Debug)]
enum RequestBodyBufferError {
Normalization(RequestBodyNormalizationError),
TooLarge { limit_bytes: u64 },
Timeout { timeout_ms: u64 },
ReadFailed { message: String },
}
impl RequestBodyBufferError {
fn http_status(&self) -> http::StatusCode {
match self {
Self::Normalization(error) => error.http_status(),
Self::TooLarge { .. } => http::StatusCode::PAYLOAD_TOO_LARGE,
Self::Timeout { .. } => http::StatusCode::REQUEST_TIMEOUT,
Self::ReadFailed { .. } => http::StatusCode::BAD_REQUEST,
}
}
fn client_message(&self) -> String {
match self {
Self::Normalization(error) => error.client_message(),
Self::TooLarge { limit_bytes } => format!("Request body exceeds {limit_bytes} bytes"),
Self::Timeout { .. } => REQUEST_BODY_READ_TIMEOUT_DETAIL.to_string(),
Self::ReadFailed { .. } => REQUEST_BODY_READ_FAILED_DETAIL.to_string(),
}
}
fn reason(&self) -> &'static str {
match self {
Self::Normalization(error) => match error {
RequestBodyNormalizationError::UnsupportedContentEncoding(_) => {
"unsupported_content_encoding"
}
RequestBodyNormalizationError::DecodeFailed { .. } => "decode_failed",
RequestBodyNormalizationError::DecompressedBodyTooLarge { .. } => {
"decompressed_body_too_large"
}
RequestBodyNormalizationError::RequestBodyTooLarge { .. } => {
"request_body_too_large"
}
},
Self::TooLarge { .. } => "request_body_too_large",
Self::Timeout { .. } => "request_body_read_timeout",
Self::ReadFailed { .. } => "request_body_read_failed",
}
}
}
async fn buffer_and_normalize_request_body(
request_body: &mut Option<Body>,
headers: &mut http::HeaderMap,
body_owner_expectation: &'static str,
trace_id: &str,
method: &http::Method,
path_and_query: &str,
phase: &'static str,
policy: RequestBodyBufferPolicy,
) -> Result<Bytes, RequestBodyBufferError> {
if let Err(err) =
crate::headers::check_request_content_length_with_limit(headers, policy.max_bytes)
{
return Err(RequestBodyBufferError::Normalization(err));
}
let read_started_at = Instant::now();
let timeout_ms = policy.read_timeout.as_millis() as u64;
info!(
event_name = "frontdoor_request_body_buffer_started",
log_type = "event",
trace_id,
method = %method,
path = %path_and_query,
phase,
max_body_bytes = policy.max_bytes,
timeout_ms,
"gateway started buffering request body"
);
let body_limit = usize::try_from(policy.max_bytes).unwrap_or(usize::MAX);
let body = match tokio::time::timeout(
policy.read_timeout,
to_bytes(
request_body.take().expect(body_owner_expectation),
body_limit,
),
)
.await
{
Ok(Ok(body)) => body,
Ok(Err(err)) if request_body_collection_exceeded_limit(&err) => {
return Err(RequestBodyBufferError::TooLarge {
limit_bytes: policy.max_bytes,
});
}
Ok(Err(err)) => {
return Err(RequestBodyBufferError::ReadFailed {
message: err.to_string(),
});
}
Err(_) => {
return Err(RequestBodyBufferError::Timeout { timeout_ms });
}
};
let normalized = crate::headers::normalize_request_body_headers_and_bytes_with_limit(
headers,
body,
policy.max_bytes,
)
.map_err(RequestBodyBufferError::Normalization)?;
info!(
event_name = "frontdoor_request_body_buffer_completed",
log_type = "event",
trace_id,
method = %method,
path = %path_and_query,
phase,
body_bytes = normalized.len(),
elapsed_ms = read_started_at.elapsed().as_millis() as u64,
"gateway completed request body buffering"
);
Ok(normalized)
}
fn request_body_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
}
fn build_request_body_buffer_error_response(
trace_id: &str,
request_context: &GatewayPublicRequestContext,
error: &RequestBodyBufferError,
) -> Result<Response<Body>, GatewayError> {
warn!(
event_name = "frontdoor_request_body_buffer_failed",
log_type = "ops",
trace_id,
method = %request_context.request_method,
path = %request_context.request_path_and_query(),
status_code = error.http_status().as_u16(),
reason = error.reason(),
detail = %error.client_message(),
read_error = match error {
RequestBodyBufferError::ReadFailed { message } => message.as_str(),
_ => "",
},
"gateway rejected request body before local execution planning"
);
build_local_http_error_response(
trace_id,
request_context.control_decision.as_ref(),
error.http_status(),
error.client_message().as_str(),
)
}
fn finalize_request_body_buffer_rejection(
state: &AppState,
request_context: &GatewayPublicRequestContext,
@@ -303,12 +114,17 @@ fn finalize_request_body_buffer_rejection(
error: &RequestBodyBufferError,
) -> Result<Response<Body>, GatewayError> {
let response = build_request_body_buffer_error_response(trace_id, request_context, error)?;
let execution_path = if matches!(error, RequestBodyBufferError::Overloaded { .. }) {
EXECUTION_PATH_LOCAL_OVERLOADED
} else {
EXECUTION_PATH_LOCAL_INVALID_REQUEST
};
Ok(finalize_gateway_response_with_context(
state,
response,
remote_addr,
request_context,
EXECUTION_PATH_LOCAL_INVALID_REQUEST,
execution_path,
started_at,
request_permit,
))
@@ -776,9 +592,17 @@ async fn restore_redacted_sync_execution_response(
return Ok(Response::from_parts(parts, body));
};
let mut headers = collect_response_headers(&parts.headers);
let body_bytes = to_bytes(body, usize::MAX)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let response_body_limit = crate::headers::max_redacted_sync_response_body_bytes();
let body_bytes = to_bytes(
body,
usize::try_from(response_body_limit).unwrap_or(usize::MAX),
)
.await
.map_err(|err| {
GatewayError::Internal(format!(
"failed to buffer redacted sync response within {response_body_limit} bytes: {err}"
))
})?;
let restored =
crate::privacy::restore_sync_response_body(&mut headers, body_bytes.as_ref(), &session)?;
replace_response_headers(&mut parts.headers, &headers)?;
@@ -1041,49 +865,6 @@ async fn proxy_request_inner(
let started_at = Instant::now();
let client_ip = effective_client_ip(request.headers(), &remote_addr);
let trace_id = extract_or_generate_trace_id(request.headers());
match state.admin_security_ip_blacklisted(client_ip).await {
Ok(true) => {
warn!(
event_name = "frontdoor_ip_blacklist_rejected",
log_type = "event",
trace_id = %trace_id,
client_ip = %client_ip,
path = %request.uri().path(),
"gateway rejected blacklisted client IP"
);
let response = build_local_http_error_response(
&trace_id,
None,
http::StatusCode::FORBIDDEN,
"当前 IP 已被禁止访问",
)?;
return Ok(finalize_gateway_response(
&state,
response,
&trace_id,
&remote_addr,
request.method(),
request
.uri()
.path_and_query()
.map(|value| value.as_str())
.unwrap_or("/"),
None,
EXECUTION_PATH_LOCAL_AUTH_DENIED,
&started_at,
None,
));
}
Ok(false) => {}
Err(err) => warn!(
event_name = "frontdoor_ip_blacklist_check_failed",
log_type = "ops",
trace_id = %trace_id,
client_ip = %client_ip,
error = ?err,
"gateway failed open after IP blacklist check error"
),
}
let accepted_at = request
.extensions()
.get::<crate::middleware::GatewayRequestAcceptedAt>()
@@ -1161,6 +942,49 @@ async fn proxy_request_inner(
};
let request_admission_ms = started_at.elapsed().as_millis() as u64;
observe_gateway_stage_ms("frontdoor_admission", request_admission_ms);
match state.admin_security_ip_blacklisted(client_ip).await {
Ok(true) => {
warn!(
event_name = "frontdoor_ip_blacklist_rejected",
log_type = "event",
trace_id = %trace_id,
client_ip = %client_ip,
path = %request.uri().path(),
"gateway rejected blacklisted client IP"
);
let response = build_local_http_error_response(
&trace_id,
None,
http::StatusCode::FORBIDDEN,
"当前 IP 已被禁止访问",
)?;
return Ok(finalize_gateway_response(
&state,
response,
&trace_id,
&remote_addr,
request.method(),
request
.uri()
.path_and_query()
.map(|value| value.as_str())
.unwrap_or("/"),
None,
EXECUTION_PATH_LOCAL_AUTH_DENIED,
&started_at,
request_permit.take(),
));
}
Ok(false) => {}
Err(err) => warn!(
event_name = "frontdoor_ip_blacklist_check_failed",
log_type = "ops",
trace_id = %trace_id,
client_ip = %client_ip,
error = ?err,
"gateway failed open after IP blacklist check error"
),
}
let (mut parts, body) = request.into_parts();
let redaction_slot = crate::privacy::RedactionSessionSlot::default();
parts.extensions.insert(redaction_slot.clone());
@@ -2432,6 +2256,7 @@ fn local_execution_runtime_miss_route_detail(
#[cfg(test)]
mod tests {
use std::sync::Arc;
use std::time::Duration;
use super::{
@@ -2442,8 +2267,9 @@ mod tests {
RequestBodyBufferPolicy,
};
use axum::body::{to_bytes, Body, Bytes};
use axum::http::{header, HeaderMap, Method, Response};
use axum::http::{header, HeaderMap, HeaderValue, Method, Response};
use serde_json::json;
use tokio::sync::Semaphore;
#[test]
fn api_key_remote_ip_allows_unrestricted_keys() {
@@ -2622,6 +2448,45 @@ mod tests {
));
}
#[tokio::test]
async fn request_body_buffer_rejects_when_weighted_budget_is_exhausted() {
let budget = Arc::new(Semaphore::new(1));
let _held = Arc::clone(&budget)
.acquire_owned()
.await
.expect("test budget should be open");
let mut body = Some(Body::from(Bytes::from_static(b"{}")));
let mut headers = HeaderMap::new();
headers.insert(header::CONTENT_LENGTH, HeaderValue::from_static("2"));
let err = buffer_and_normalize_request_body(
&mut body,
&mut headers,
"test owns body",
"trace-body-budget",
&Method::POST,
"/v1/responses",
"test",
RequestBodyBufferPolicy::for_tests_with_budget(
1024,
Duration::from_secs(1),
Duration::from_millis(5),
crate::state::REQUEST_BODY_BUFFER_PERMIT_BYTES,
budget,
),
)
.await
.expect_err("exhausted body budget should reject quickly");
assert!(matches!(
err,
RequestBodyBufferError::Overloaded {
requested_bytes: 2,
..
}
));
}
#[test]
fn runtime_miss_detail_returns_model_specific_stream_message_when_candidates_are_unavailable() {
let decision = GatewayControlDecision::synthetic(
@@ -464,7 +464,7 @@ async fn complete_oauth_login(
.cloned()
.collect::<Vec<_>>();
let body = login_response.into_body();
let body = match to_bytes(body, usize::MAX).await {
let body = match to_bytes(body, crate::headers::max_internal_buffered_body_bytes()).await {
Ok(value) => value,
Err(_) => return redirect_oauth_error(Some(frontend_callback_url), "provider_unavailable"),
};
+30
View File
@@ -15,6 +15,10 @@ use uuid::Uuid;
const DEFAULT_MAX_REQUEST_BODY_MB: u64 = 64;
const MAX_REQUEST_BODY_MB_ENV: &str = "AETHER_MAX_REQUEST_BODY_MB";
const DEFAULT_MAX_REDACTED_SYNC_RESPONSE_BODY_MB: u64 = 64;
const MAX_REDACTED_SYNC_RESPONSE_BODY_MB_ENV: &str = "AETHER_MAX_REDACTED_SYNC_RESPONSE_BODY_MB";
const DEFAULT_MAX_INTERNAL_BUFFERED_BODY_MB: u64 = 128;
const MAX_INTERNAL_BUFFERED_BODY_MB_ENV: &str = "AETHER_MAX_INTERNAL_BUFFERED_BODY_MB";
const TRUSTED_PROXY_CIDRS_ENV: &str = "AETHER_TRUSTED_PROXY_CIDRS";
/// Upper bound applied to a request body after Content-Encoding decoding, and to
@@ -29,6 +33,24 @@ static MAX_REQUEST_BODY_BYTES: LazyLock<u64> = LazyLock::new(|| {
.saturating_mul(1024 * 1024)
});
static MAX_REDACTED_SYNC_RESPONSE_BODY_BYTES: LazyLock<u64> = LazyLock::new(|| {
std::env::var(MAX_REDACTED_SYNC_RESPONSE_BODY_MB_ENV)
.ok()
.and_then(|value| value.parse::<u64>().ok())
.filter(|value| *value > 0)
.unwrap_or(DEFAULT_MAX_REDACTED_SYNC_RESPONSE_BODY_MB)
.saturating_mul(1024 * 1024)
});
static MAX_INTERNAL_BUFFERED_BODY_BYTES: LazyLock<u64> = LazyLock::new(|| {
std::env::var(MAX_INTERNAL_BUFFERED_BODY_MB_ENV)
.ok()
.and_then(|value| value.parse::<u64>().ok())
.filter(|value| *value > 0)
.unwrap_or(DEFAULT_MAX_INTERNAL_BUFFERED_BODY_MB)
.saturating_mul(1024 * 1024)
});
static TRUSTED_PROXY_CIDRS: LazyLock<Vec<String>> = LazyLock::new(|| {
std::env::var(TRUSTED_PROXY_CIDRS_ENV)
.unwrap_or_else(|_| "127.0.0.0/8,::1/128".to_string())
@@ -43,6 +65,14 @@ pub(crate) fn max_request_body_bytes() -> u64 {
*MAX_REQUEST_BODY_BYTES
}
pub(crate) fn max_redacted_sync_response_body_bytes() -> u64 {
*MAX_REDACTED_SYNC_RESPONSE_BODY_BYTES
}
pub(crate) fn max_internal_buffered_body_bytes() -> usize {
usize::try_from(*MAX_INTERNAL_BUFFERED_BODY_BYTES).unwrap_or(usize::MAX)
}
pub(crate) fn extract_or_generate_trace_id(headers: &http::HeaderMap) -> String {
header_value_str(headers, TRACE_ID_HEADER).unwrap_or_else(|| Uuid::new_v4().to_string())
}
+115 -23
View File
@@ -261,6 +261,11 @@ const MAX_GATEWAY_LISTENER_SHARDS: usize = 64;
const DEFAULT_GATEWAY_HTTP2_MAX_CONCURRENT_STREAMS: u32 = 16_384;
const MIN_GATEWAY_HTTP2_MAX_CONCURRENT_STREAMS: u32 = 200;
const MAX_GATEWAY_HTTP2_MAX_CONCURRENT_STREAMS: u32 = 1_000_000;
const AUTO_GATEWAY_REQUESTS_PER_CPU: usize = 1_024;
const MIN_AUTO_GATEWAY_REQUEST_CONCURRENCY: usize = 512;
const MAX_AUTO_GATEWAY_REQUEST_CONCURRENCY: usize = 65_536;
const AUTO_GATEWAY_REQUEST_FD_DIVISOR: usize = 2;
const AUTO_GATEWAY_REQUEST_FD_RESERVE: usize = 256;
fn env_var_trimmed(name: &str) -> Option<String> {
std::env::var(name)
.ok()
@@ -281,6 +286,55 @@ fn available_parallelism_usize() -> usize {
.max(1)
}
fn automatic_gateway_request_concurrency_for_parallelism(parallelism: usize) -> usize {
automatic_gateway_request_concurrency_for_capacity(parallelism, None)
}
fn automatic_gateway_request_concurrency_for_capacity(
parallelism: usize,
fd_soft_limit: Option<usize>,
) -> usize {
let cpu_limit = parallelism
.max(1)
.saturating_mul(AUTO_GATEWAY_REQUESTS_PER_CPU)
.clamp(
MIN_AUTO_GATEWAY_REQUEST_CONCURRENCY,
MAX_AUTO_GATEWAY_REQUEST_CONCURRENCY,
);
let fd_limit = fd_soft_limit
.map(|limit| {
limit
.saturating_sub(AUTO_GATEWAY_REQUEST_FD_RESERVE)
.checked_div(AUTO_GATEWAY_REQUEST_FD_DIVISOR)
.unwrap_or(1)
.max(1)
})
.unwrap_or(MAX_AUTO_GATEWAY_REQUEST_CONCURRENCY);
cpu_limit.min(fd_limit).max(1)
}
fn automatic_gateway_request_concurrency() -> usize {
automatic_gateway_request_concurrency_for_capacity(
available_parallelism_usize(),
soft_fd_limit(),
)
}
fn soft_fd_limit() -> Option<usize> {
#[cfg(unix)]
{
let mut limit = libc::rlimit {
rlim_cur: 0,
rlim_max: 0,
};
let result = unsafe { libc::getrlimit(libc::RLIMIT_NOFILE, &mut limit) };
if result == 0 {
return usize::try_from(limit.rlim_cur).ok();
}
}
None
}
fn usage_queue_request_concurrency_hint(
max_in_flight_requests: Option<usize>,
distributed_request_limit: Option<usize>,
@@ -1655,13 +1709,23 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
runtime_redis_url.as_deref(),
runtime_backend,
)?;
let request_concurrency_limit = args
.max_in_flight_requests
.filter(|limit| *limit > 0)
.unwrap_or_else(automatic_gateway_request_concurrency);
let usage_queue_request_concurrency_hint = usage_queue_request_concurrency_hint(
args.max_in_flight_requests,
Some(request_concurrency_limit),
args.distributed_request_limit,
);
let usage_queue_request_concurrency_hint_source =
if args.max_in_flight_requests.is_some() || args.distributed_request_limit.is_some() {
"explicit"
} else {
"auto"
};
let usage_queue_workers = args.usage.effective_queue_workers(
args.node_role,
args.max_in_flight_requests,
Some(request_concurrency_limit),
args.distributed_request_limit,
sql_database_config.as_ref(),
);
@@ -1721,11 +1785,7 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
.unwrap_or_default(),
usage_queue_request_concurrency_hint =
usage_queue_request_concurrency_hint.unwrap_or_default(),
usage_queue_request_concurrency_hint_source = if usage_queue_request_concurrency_hint.is_some() {
"explicit"
} else {
"none"
},
usage_queue_request_concurrency_hint_source,
frontdoor_mode = "compatibility_frontdoor",
log_format = ?args.logging.log_format,
log_destination = args.logging.log_destination.as_str(),
@@ -1762,13 +1822,13 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
.unwrap_or_default(),
usage_queue_request_concurrency_hint =
usage_queue_request_concurrency_hint.unwrap_or_default(),
usage_queue_request_concurrency_hint_source =
if usage_queue_request_concurrency_hint.is_some() {
"explicit"
} else {
"none"
},
max_in_flight_requests = args.max_in_flight_requests.unwrap_or_default(),
usage_queue_request_concurrency_hint_source,
max_in_flight_requests = request_concurrency_limit,
max_in_flight_requests_source = if args.max_in_flight_requests.is_some() {
"explicit"
} else {
"auto"
},
distributed_request_limit = args.distributed_request_limit.unwrap_or_default(),
distributed_request_redis_configured = args
.distributed_request_redis_url
@@ -1822,9 +1882,7 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
{
state = state.with_video_task_store_path(path)?;
}
if let Some(limit) = args.max_in_flight_requests.filter(|limit| *limit > 0) {
state = state.with_request_concurrency_limit(limit);
}
state = state.with_request_concurrency_limit(request_concurrency_limit);
if let Some(limit) = args.distributed_request_limit.filter(|limit| *limit > 0) {
let distributed_gate = state
.runtime_state()
@@ -2315,12 +2373,14 @@ fn pending_backfills_error(
#[cfg(test)]
mod tests {
use super::{
automatic_sql_pool_config, automatic_sql_pool_config_for_parallelism,
automatic_usage_queue_workers_for_parallelism, ensure_database_backfills_are_current,
ensure_database_schema_is_current, pending_backfills_error, pending_schema_error,
resolve_healthcheck_url, Args, DatabaseDriverArg, DeploymentTopologyArg, GatewayDataArgs,
GatewayFrontdoorArgs, GatewayLogDestinationArg, GatewayLogFormatArg, GatewayLogRotationArg,
GatewayLoggingArgs, GatewayRateLimitArgs, GatewayUsageArgs, NodeRoleArg, RuntimeBackendArg,
automatic_gateway_request_concurrency_for_capacity,
automatic_gateway_request_concurrency_for_parallelism, automatic_sql_pool_config,
automatic_sql_pool_config_for_parallelism, automatic_usage_queue_workers_for_parallelism,
ensure_database_backfills_are_current, ensure_database_schema_is_current,
pending_backfills_error, pending_schema_error, resolve_healthcheck_url, Args,
DatabaseDriverArg, DeploymentTopologyArg, GatewayDataArgs, GatewayFrontdoorArgs,
GatewayLogDestinationArg, GatewayLogFormatArg, GatewayLogRotationArg, GatewayLoggingArgs,
GatewayRateLimitArgs, GatewayUsageArgs, NodeRoleArg, RuntimeBackendArg,
VideoTaskTruthSourceArg, DEFAULT_GATEWAY_HTTP2_MAX_CONCURRENT_STREAMS,
DEFAULT_GATEWAY_LISTENER_SHARDS, DEFAULT_GATEWAY_LISTEN_BACKLOG,
MAX_GATEWAY_HTTP2_MAX_CONCURRENT_STREAMS, MAX_GATEWAY_LISTENER_SHARDS,
@@ -2509,6 +2569,38 @@ mod tests {
);
}
#[test]
fn auto_gateway_request_concurrency_scales_and_clamps() {
assert_eq!(
automatic_gateway_request_concurrency_for_parallelism(1),
1_024
);
assert_eq!(
automatic_gateway_request_concurrency_for_parallelism(4),
4_096
);
assert_eq!(
automatic_gateway_request_concurrency_for_parallelism(64),
65_536
);
}
#[test]
fn auto_gateway_request_concurrency_respects_fd_budget() {
assert_eq!(
automatic_gateway_request_concurrency_for_capacity(64, Some(16_384)),
8_064
);
assert_eq!(
automatic_gateway_request_concurrency_for_capacity(64, Some(1_024)),
384
);
assert_eq!(
automatic_gateway_request_concurrency_for_capacity(64, None),
65_536
);
}
#[test]
fn explicit_migrate_runtime_config_enables_data_logs() {
let mut args = test_args();
+36 -1
View File
@@ -10,7 +10,7 @@ use aether_data_contracts::repository::quota::StoredProviderQuotaSnapshot;
use aether_runtime::ConcurrencyGate;
use aether_runtime_state::{RuntimeSemaphore, RuntimeState};
use dashmap::DashMap;
use tokio::sync::Mutex as TokioMutex;
use tokio::sync::{Mutex as TokioMutex, Semaphore};
use super::super::async_task::{VideoTaskPollerConfig, VideoTaskService};
use super::super::cache::{
@@ -36,6 +36,11 @@ const DEFAULT_REQUEST_BODY_READ_TIMEOUT_MS: u64 = 120_000;
const MIN_REQUEST_BODY_READ_TIMEOUT_MS: u64 = 1_000;
const MAX_REQUEST_BODY_READ_TIMEOUT_MS: u64 = 600_000;
const REQUEST_BODY_READ_TIMEOUT_MS_ENV: &str = "AETHER_GATEWAY_REQUEST_BODY_READ_TIMEOUT_MS";
const DEFAULT_REQUEST_BODY_BUFFER_BUDGET_MB: usize = 256;
const MIN_REQUEST_BODY_BUFFER_BUDGET_MB: usize = 64;
const MAX_REQUEST_BODY_BUFFER_BUDGET_MB: usize = 16 * 1024;
const REQUEST_BODY_BUFFER_BUDGET_MB_ENV: &str = "AETHER_GATEWAY_REQUEST_BODY_BUFFER_BUDGET_MB";
pub(crate) const REQUEST_BODY_BUFFER_PERMIT_BYTES: usize = 64 * 1024;
const DEFAULT_LOCAL_EXECUTION_PLANNING_TIMEOUT_MS: u64 = 30_000;
const MIN_LOCAL_EXECUTION_PLANNING_TIMEOUT_MS: u64 = 500;
@@ -90,6 +95,8 @@ impl std::fmt::Debug for TestExecutionRuntimeSyncOverride {
#[derive(Debug, Clone)]
pub(crate) struct FrontdoorRuntimeGuardConfig {
pub(crate) request_body_read_timeout: Duration,
pub(crate) request_body_buffer_budget_bytes: usize,
pub(crate) request_body_buffer_budget_permits: usize,
pub(crate) local_execution_planning_timeout: Duration,
pub(crate) internal_gate_queue_budget: Duration,
pub(crate) auth_capacity_cache_ttl: Duration,
@@ -108,6 +115,8 @@ impl FrontdoorRuntimeGuardConfig {
MIN_REQUEST_BODY_READ_TIMEOUT_MS,
MAX_REQUEST_BODY_READ_TIMEOUT_MS,
),
request_body_buffer_budget_bytes: request_body_buffer_budget_bytes_from_env(),
request_body_buffer_budget_permits: request_body_buffer_budget_permits_from_env(),
local_execution_planning_timeout: env_duration_ms(
LOCAL_EXECUTION_PLANNING_TIMEOUT_MS_ENV,
DEFAULT_LOCAL_EXECUTION_PLANNING_TIMEOUT_MS,
@@ -140,6 +149,8 @@ impl FrontdoorRuntimeGuardConfig {
) -> Self {
Self {
request_body_read_timeout,
request_body_buffer_budget_bytes: DEFAULT_REQUEST_BODY_BUFFER_BUDGET_MB * 1024 * 1024,
request_body_buffer_budget_permits: DEFAULT_REQUEST_BODY_BUFFER_BUDGET_MB * 16,
local_execution_planning_timeout,
internal_gate_queue_budget: Duration::from_millis(
DEFAULT_INTERNAL_GATE_QUEUE_BUDGET_MS,
@@ -153,6 +164,27 @@ impl FrontdoorRuntimeGuardConfig {
}
}
fn request_body_buffer_budget_mb_from_env() -> usize {
std::env::var(REQUEST_BODY_BUFFER_BUDGET_MB_ENV)
.ok()
.and_then(|value| value.trim().parse::<usize>().ok())
.filter(|value| *value > 0)
.unwrap_or(DEFAULT_REQUEST_BODY_BUFFER_BUDGET_MB)
.clamp(
MIN_REQUEST_BODY_BUFFER_BUDGET_MB,
MAX_REQUEST_BODY_BUFFER_BUDGET_MB,
)
}
fn request_body_buffer_budget_bytes_from_env() -> usize {
request_body_buffer_budget_mb_from_env().saturating_mul(1024 * 1024)
}
fn request_body_buffer_budget_permits_from_env() -> usize {
request_body_buffer_budget_bytes_from_env().saturating_add(REQUEST_BODY_BUFFER_PERMIT_BYTES - 1)
/ REQUEST_BODY_BUFFER_PERMIT_BYTES
}
fn env_duration_ms(key: &str, default_ms: u64, min_ms: u64, max_ms: u64) -> Duration {
let ms = std::env::var(key)
.ok()
@@ -337,6 +369,7 @@ pub struct AppState {
pub(crate) video_tasks: Arc<VideoTaskService>,
pub(crate) video_task_poller: Option<VideoTaskPollerConfig>,
pub(crate) frontdoor_runtime_guards: Arc<FrontdoorRuntimeGuardConfig>,
pub(crate) request_body_buffer_budget: Arc<Semaphore>,
pub(crate) request_gate: Option<Arc<ConcurrencyGate>>,
pub(crate) auth_snapshot_load_gate: Option<Arc<ConcurrencyGate>>,
pub(crate) candidate_planning_gate: Option<Arc<ConcurrencyGate>>,
@@ -346,6 +379,8 @@ pub struct AppState {
pub(crate) client: reqwest::Client,
pub(crate) auth_context_cache: Arc<AuthContextCache>,
pub(crate) auth_snapshot_cache: Arc<AuthSnapshotCache>,
pub(crate) admin_security_blacklist_cache: Arc<ValueCache<String, bool>>,
pub(crate) admin_security_whitelist_cache: Arc<ValueCache<String, Vec<String>>>,
pub(crate) user_model_capability_settings_cache: Arc<JsonValueCache<String>>,
pub(crate) user_feature_settings_cache: Arc<JsonValueCache<String>>,
pub(crate) auth_api_key_force_capabilities_cache:
+56
View File
@@ -266,6 +266,9 @@ impl AppState {
)),
video_task_poller: None,
frontdoor_runtime_guards: Arc::clone(&frontdoor_runtime_guards),
request_body_buffer_budget: Arc::new(tokio::sync::Semaphore::new(
frontdoor_runtime_guards.request_body_buffer_budget_permits,
)),
request_gate: None,
auth_snapshot_load_gate: frontdoor_runtime_guards
.auth_snapshot_load_gate_limit
@@ -286,6 +289,8 @@ impl AppState {
client,
auth_context_cache: Arc::new(AuthContextCache::default()),
auth_snapshot_cache: Arc::new(AuthSnapshotCache::default()),
admin_security_blacklist_cache: Arc::new(ValueCache::default()),
admin_security_whitelist_cache: Arc::new(ValueCache::default()),
user_model_capability_settings_cache: Arc::new(JsonValueCache::default()),
user_feature_settings_cache: Arc::new(JsonValueCache::default()),
auth_api_key_force_capabilities_cache: Arc::new(JsonValueCache::default()),
@@ -480,6 +485,8 @@ impl AppState {
pub fn with_runtime_state(mut self, runtime_state: Arc<RuntimeState>) -> Self {
self.runtime_state = runtime_state;
self.admin_security_blacklist_cache.clear();
self.admin_security_whitelist_cache.clear();
self.data = Arc::new(
(*self.data)
.clone()
@@ -1062,6 +1069,38 @@ impl AppState {
pub(crate) async fn metric_samples(&self) -> Vec<MetricSample> {
let mut samples = vec![service_up_sample("aether-gateway")];
let request_body_buffer_budget_bytes = self
.frontdoor_runtime_guards
.request_body_buffer_budget_bytes;
let request_body_buffer_available_bytes = self
.request_body_buffer_budget
.available_permits()
.saturating_mul(super::REQUEST_BODY_BUFFER_PERMIT_BYTES)
.min(request_body_buffer_budget_bytes);
samples.extend([
MetricSample::new(
"request_body_buffer_budget_bytes",
"Configured weighted request body buffering budget in bytes.",
MetricKind::Gauge,
u64::try_from(request_body_buffer_budget_bytes).unwrap_or(u64::MAX),
),
MetricSample::new(
"request_body_buffer_available_bytes",
"Currently available weighted request body buffering budget in bytes.",
MetricKind::Gauge,
u64::try_from(request_body_buffer_available_bytes).unwrap_or(u64::MAX),
),
MetricSample::new(
"request_body_buffer_in_use_bytes",
"Currently reserved weighted request body buffering budget in bytes.",
MetricKind::Gauge,
u64::try_from(
request_body_buffer_budget_bytes
.saturating_sub(request_body_buffer_available_bytes),
)
.unwrap_or(u64::MAX),
),
]);
if let Some(snapshot) = self.request_concurrency_snapshot() {
samples.extend(snapshot.to_metric_samples("gateway_requests"));
}
@@ -2935,6 +2974,23 @@ mod tests {
}));
}
#[tokio::test]
async fn metric_samples_include_request_body_buffer_budget_usage() {
let state = AppState::new().expect("app state should build");
let _permit = Arc::clone(&state.request_body_buffer_budget)
.acquire_many_owned(2)
.await
.expect("request body budget should be open");
let samples = state.metric_samples().await;
assert!(samples.iter().any(|sample| {
sample.name == "request_body_buffer_in_use_bytes"
&& sample.value
== u64::try_from(2 * crate::state::REQUEST_BODY_BUFFER_PERMIT_BYTES)
.unwrap_or(u64::MAX)
}));
}
#[test]
fn scheduler_affinity_epoch_blocks_stale_rewarm_after_invalidation() {
let state = AppState::new().expect("app state should build");
+1 -1
View File
@@ -29,7 +29,7 @@ pub(crate) use self::admin_types::{
pub use self::app::AppState;
pub(crate) use self::app::{
upstream_target_gate_auto_limit, upstream_target_gate_limit_from_env,
FrontdoorRuntimeGuardConfig,
FrontdoorRuntimeGuardConfig, REQUEST_BODY_BUFFER_PERMIT_BYTES,
};
pub(crate) use self::cache::{
CachedProviderTransportSnapshot, AUTH_API_KEY_LAST_USED_MAX_ENTRIES,
@@ -1,31 +1,68 @@
use crate::state::AdminSecurityBlacklistEntry;
use crate::{AppState, GatewayError};
use std::net::IpAddr;
use std::sync::LazyLock;
use std::time::Duration;
const ADMIN_SECURITY_BLACKLIST_PREFIX: &str = "ip:blacklist:";
const ADMIN_SECURITY_WHITELIST_KEY: &str = "ip:whitelist";
const ADMIN_SECURITY_CACHE_TTL_MS_ENV: &str = "AETHER_GATEWAY_SECURITY_CACHE_TTL_MS";
const DEFAULT_ADMIN_SECURITY_CACHE_TTL_MS: u64 = 1_000;
const MAX_ADMIN_SECURITY_CACHE_TTL_MS: u64 = 30_000;
const ADMIN_SECURITY_WHITELIST_CACHE_KEY: &str = "rules";
static ADMIN_SECURITY_CACHE_TTL: LazyLock<Duration> = LazyLock::new(|| {
let ttl_ms = std::env::var(ADMIN_SECURITY_CACHE_TTL_MS_ENV)
.ok()
.and_then(|value| value.trim().parse::<u64>().ok())
.unwrap_or(DEFAULT_ADMIN_SECURITY_CACHE_TTL_MS)
.min(MAX_ADMIN_SECURITY_CACHE_TTL_MS);
Duration::from_millis(ttl_ms)
});
fn admin_security_cache_ttl() -> Duration {
*ADMIN_SECURITY_CACHE_TTL
}
impl AppState {
pub(crate) async fn admin_security_ip_blacklisted(
&self,
ip_address: IpAddr,
) -> Result<bool, GatewayError> {
const ADMIN_SECURITY_BLACKLIST_PREFIX: &str = "ip:blacklist:";
self.runtime_state
.kv_exists(&format!("{ADMIN_SECURITY_BLACKLIST_PREFIX}{ip_address}"))
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
let cache_key = ip_address.to_string();
let runtime_key = format!("{ADMIN_SECURITY_BLACKLIST_PREFIX}{cache_key}");
Ok(self
.admin_security_blacklist_cache
.get_or_load_once(cache_key, admin_security_cache_ttl(), || async {
self.runtime_state
.kv_exists(&runtime_key)
.await
.map(Some)
.map_err(|err| GatewayError::Internal(err.to_string()))
})
.await?
.unwrap_or(false))
}
pub(crate) async fn admin_security_ip_whitelisted(
&self,
ip_address: IpAddr,
) -> Result<bool, GatewayError> {
const ADMIN_SECURITY_WHITELIST_KEY: &str = "ip:whitelist";
let rules = self
.runtime_state
.set_members(ADMIN_SECURITY_WHITELIST_KEY)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
.admin_security_whitelist_cache
.get_or_load_once(
ADMIN_SECURITY_WHITELIST_CACHE_KEY.to_string(),
admin_security_cache_ttl(),
|| async {
self.runtime_state
.set_members(ADMIN_SECURITY_WHITELIST_KEY)
.await
.map(Some)
.map_err(|err| GatewayError::Internal(err.to_string()))
},
)
.await?
.unwrap_or_default();
Ok(rules
.iter()
.any(|rule| crate::handlers::shared::ip_rule_pattern_matches(rule.trim(), ip_address)))
@@ -37,8 +74,6 @@ impl AppState {
reason: &str,
ttl_seconds: Option<u64>,
) -> Result<bool, GatewayError> {
const ADMIN_SECURITY_BLACKLIST_PREFIX: &str = "ip:blacklist:";
let key = format!("{ADMIN_SECURITY_BLACKLIST_PREFIX}{ip_address}");
self.runtime_state
.kv_set(
@@ -47,28 +82,40 @@ impl AppState {
ttl_seconds.map(std::time::Duration::from_secs),
)
.await
.map(|_| true)
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if let Ok(ip_address) = ip_address.parse::<IpAddr>() {
self.admin_security_blacklist_cache.insert(
ip_address.to_string(),
Some(true),
admin_security_cache_ttl(),
);
}
Ok(true)
}
pub(crate) async fn remove_admin_security_blacklist(
&self,
ip_address: &str,
) -> Result<bool, GatewayError> {
const ADMIN_SECURITY_BLACKLIST_PREFIX: &str = "ip:blacklist:";
let key = format!("{ADMIN_SECURITY_BLACKLIST_PREFIX}{ip_address}");
self.runtime_state
let removed = self
.runtime_state
.kv_delete(&key)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if let Ok(ip_address) = ip_address.parse::<IpAddr>() {
self.admin_security_blacklist_cache.insert(
ip_address.to_string(),
Some(false),
admin_security_cache_ttl(),
);
}
Ok(removed)
}
pub(crate) async fn admin_security_blacklist_stats(
&self,
) -> Result<(bool, usize, Option<String>), GatewayError> {
const ADMIN_SECURITY_BLACKLIST_PREFIX: &str = "ip:blacklist:";
let total = self
.runtime_state
.scan_keys(&format!("{ADMIN_SECURITY_BLACKLIST_PREFIX}*"), 100)
@@ -81,8 +128,6 @@ impl AppState {
pub(crate) async fn list_admin_security_blacklist(
&self,
) -> Result<Vec<AdminSecurityBlacklistEntry>, GatewayError> {
const ADMIN_SECURITY_BLACKLIST_PREFIX: &str = "ip:blacklist:";
let keys = self
.runtime_state
.scan_keys(&format!("{ADMIN_SECURITY_BLACKLIST_PREFIX}*"), 100)
@@ -123,30 +168,28 @@ impl AppState {
&self,
ip_address: &str,
) -> Result<bool, GatewayError> {
const ADMIN_SECURITY_WHITELIST_KEY: &str = "ip:whitelist";
self.runtime_state
.set_add(ADMIN_SECURITY_WHITELIST_KEY, ip_address)
.await
.map(|_| true)
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
self.admin_security_whitelist_cache.clear();
Ok(true)
}
pub(crate) async fn remove_admin_security_whitelist(
&self,
ip_address: &str,
) -> Result<bool, GatewayError> {
const ADMIN_SECURITY_WHITELIST_KEY: &str = "ip:whitelist";
self.runtime_state
let removed = self
.runtime_state
.set_remove(ADMIN_SECURITY_WHITELIST_KEY, ip_address)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
self.admin_security_whitelist_cache.clear();
Ok(removed)
}
pub(crate) async fn list_admin_security_whitelist(&self) -> Result<Vec<String>, GatewayError> {
const ADMIN_SECURITY_WHITELIST_KEY: &str = "ip:whitelist";
self.runtime_state
.set_members(ADMIN_SECURITY_WHITELIST_KEY)
.await
@@ -11,6 +11,59 @@ fn production_workspace_source(path: &Path) -> String {
.to_string()
}
#[test]
fn gateway_production_body_collection_stays_bounded() {
let root = Path::new(env!("CARGO_MANIFEST_DIR")).join("src");
let mut files = Vec::new();
collect_rust_files(&root, &mut files);
let forbidden = [
"to_bytes(body, usize::MAX)",
"to_bytes(request.into_body(), usize::MAX)",
"to_bytes(response.into_body(), usize::MAX)",
"into_body(),\n usize::MAX",
];
let violations = files
.into_iter()
.filter(|path| {
!path
.components()
.any(|component| component.as_os_str() == "tests")
})
.filter_map(|path| {
let source = production_workspace_source(&path);
let hits = forbidden
.iter()
.filter(|pattern| source.contains(**pattern))
.copied()
.collect::<Vec<_>>();
if hits.is_empty() {
None
} else {
Some(format!("{} -> {}", path.display(), hits.join(", ")))
}
})
.collect::<Vec<_>>();
assert!(
violations.is_empty(),
"production body collection must use explicit limits:\n{}",
violations.join("\n")
);
}
#[test]
fn tunnel_node_status_delivery_stays_bounded() {
let source = read_workspace_file("apps/aether-gateway/src/tunnel/embedded/hub.rs");
assert!(
!source.contains("unbounded_channel::<NodeStatusEvent>"),
"tunnel node status delivery must not use an unbounded channel"
);
assert!(
source.contains("bounded_queue::<NodeStatusEvent>"),
"tunnel node status delivery must use the tracked bounded queue"
);
}
#[test]
fn gateway_small_runtime_shims_stay_deleted() {
for path in [
@@ -85,6 +85,60 @@ async fn admin_security_whitelist_matches_cidr() {
.expect("whitelist check should succeed"));
}
#[tokio::test]
async fn admin_security_blacklist_cache_tracks_local_mutations() {
let state = AppState::new().expect("gateway should build");
let ip_address = "203.0.113.9".parse().expect("valid ip");
assert!(!state
.admin_security_ip_blacklisted(ip_address)
.await
.expect("initial blacklist check should succeed"));
state
.add_admin_security_blacklist("203.0.113.9", "manual", None)
.await
.expect("blacklist add should succeed");
assert!(state
.admin_security_ip_blacklisted(ip_address)
.await
.expect("cached blacklist check should succeed"));
state
.remove_admin_security_blacklist("203.0.113.9")
.await
.expect("blacklist remove should succeed");
assert!(!state
.admin_security_ip_blacklisted(ip_address)
.await
.expect("updated blacklist check should succeed"));
}
#[tokio::test]
async fn admin_security_whitelist_cache_invalidates_after_mutation() {
let state = AppState::new().expect("gateway should build");
let ip_address = "203.0.113.10".parse().expect("valid ip");
assert!(!state
.admin_security_ip_whitelisted(ip_address)
.await
.expect("initial whitelist check should succeed"));
state
.add_admin_security_whitelist("203.0.113.0/24")
.await
.expect("whitelist add should succeed");
assert!(state
.admin_security_ip_whitelisted(ip_address)
.await
.expect("updated whitelist check should succeed"));
state
.remove_admin_security_whitelist("203.0.113.0/24")
.await
.expect("whitelist remove should succeed");
assert!(!state
.admin_security_ip_whitelisted(ip_address)
.await
.expect("removed whitelist check should succeed"));
}
async fn send_admin_security_request(
gateway: Router,
method: reqwest::Method,
+62 -11
View File
@@ -4,7 +4,8 @@ use std::sync::{Arc, LazyLock};
use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH};
use aether_runtime::{
BoundedQueueSender, MetricKind, MetricLabel, MetricSample, QueueSendError, QueueSnapshot,
bounded_queue, BoundedQueueSender, MetricKind, MetricLabel, MetricSample, QueueSendError,
QueueSnapshot,
};
use axum::extract::ws::Message;
use bytes::Bytes;
@@ -23,6 +24,7 @@ const SOFT_AVOID_STREAM_PRESSURE_PERCENT: u64 = 85;
const OUTBOUND_BACKPRESSURE_TIMEOUT: Duration = Duration::from_secs(5);
const DEFAULT_STREAM_INITIAL_WINDOW_BYTES: u32 = 4 * 1024 * 1024;
const DEFAULT_DRAIN_DEADLINE_MS: u64 = 30_000;
const DEFAULT_NODE_STATUS_QUEUE_CAPACITY: usize = 1_024;
const CONNECTION_WARMUP: Duration = Duration::from_secs(1);
static STREAM_INITIAL_WINDOW_BYTES: LazyLock<u32> = LazyLock::new(|| {
@@ -41,6 +43,14 @@ static DRAIN_DEADLINE_MS: LazyLock<u64> = LazyLock::new(|| {
.unwrap_or(DEFAULT_DRAIN_DEADLINE_MS)
});
static NODE_STATUS_QUEUE_CAPACITY: LazyLock<usize> = LazyLock::new(|| {
std::env::var("AETHER_TUNNEL_NODE_STATUS_QUEUE_CAPACITY")
.ok()
.and_then(|value| value.parse::<usize>().ok())
.filter(|value| *value > 0)
.unwrap_or(DEFAULT_NODE_STATUS_QUEUE_CAPACITY)
});
static STREAM_MIN_WINDOW_UPDATE_BYTES: LazyLock<u32> = LazyLock::new(|| {
STREAM_INITIAL_WINDOW_BYTES
.saturating_div(4)
@@ -655,7 +665,7 @@ pub struct HubRouter {
next_conn_id: AtomicU64,
next_local_stream_id: AtomicU64,
control_plane: ControlPlaneClient,
node_status_tx: mpsc::UnboundedSender<NodeStatusEvent>,
node_status_tx: BoundedQueueSender<NodeStatusEvent>,
soft_avoid_selection_total: AtomicU64,
selection_retry_total: AtomicU64,
selection_unavailable_total: AtomicU64,
@@ -676,7 +686,8 @@ struct NodeStatusEvent {
impl HubRouter {
pub fn new(control_plane: ControlPlaneClient) -> Arc<Self> {
let (node_status_tx, mut node_status_rx) = mpsc::unbounded_channel::<NodeStatusEvent>();
let (node_status_tx, mut node_status_rx) =
bounded_queue::<NodeStatusEvent>(*NODE_STATUS_QUEUE_CAPACITY);
let worker_control_plane = control_plane.clone();
if let Ok(handle) = tokio::runtime::Handle::try_current() {
handle.spawn(async move {
@@ -797,14 +808,31 @@ impl HubRouter {
conn_count,
observed_at_unix_secs: current_unix_secs(),
};
if let Err(error) = self.node_status_tx.send(event) {
warn!(
node_id = %error.0.node_id,
connected = error.0.connected,
conn_count = error.0.conn_count,
observed_at_unix_secs = error.0.observed_at_unix_secs,
"node status worker unavailable"
);
match self.node_status_tx.try_send(event) {
Ok(()) => {}
Err(QueueSendError::Full(event)) => {
let rejected_total = self.node_status_tx.snapshot().rejected_full_total;
if rejected_total.is_power_of_two() {
warn!(
node_id = %event.node_id,
connected = event.connected,
conn_count = event.conn_count,
observed_at_unix_secs = event.observed_at_unix_secs,
queue_capacity = self.node_status_tx.capacity(),
rejected_total,
"node status queue full; dropping stale status event"
);
}
}
Err(QueueSendError::Closed(event)) => {
warn!(
node_id = %event.node_id,
connected = event.connected,
conn_count = event.conn_count,
observed_at_unix_secs = event.observed_at_unix_secs,
"node status worker unavailable"
);
}
}
}
@@ -1592,6 +1620,7 @@ impl HubRouter {
}
pub fn stats(&self) -> HubStats {
let node_status_queue = self.node_status_tx.snapshot();
let proxy_conns = self
.proxy_conns_by_id
.iter()
@@ -1716,6 +1745,12 @@ impl HubRouter {
scheduler_selected_conn_total: self
.scheduler_selected_conn_total
.load(Ordering::Relaxed),
node_status_queue_capacity: node_status_queue.capacity,
node_status_queue_depth: node_status_queue.depth,
node_status_queue_high_watermark: node_status_queue.high_watermark,
node_status_queue_enqueued_total: node_status_queue.enqueued_total,
node_status_queue_rejected_full_total: node_status_queue.rejected_full_total,
node_status_queue_rejected_closed_total: node_status_queue.rejected_closed_total,
}
}
}
@@ -1828,6 +1863,12 @@ pub struct HubStats {
pub selection_retry_total: u64,
pub selection_unavailable_total: u64,
pub scheduler_selected_conn_total: u64,
pub node_status_queue_capacity: usize,
pub node_status_queue_depth: usize,
pub node_status_queue_high_watermark: usize,
pub node_status_queue_enqueued_total: u64,
pub node_status_queue_rejected_full_total: u64,
pub node_status_queue_rejected_closed_total: u64,
}
impl HubStats {
@@ -2039,6 +2080,16 @@ impl HubStats {
);
}
let node_status_queue = QueueSnapshot {
capacity: self.node_status_queue_capacity,
depth: self.node_status_queue_depth,
high_watermark: self.node_status_queue_high_watermark,
enqueued_total: self.node_status_queue_enqueued_total,
rejected_full_total: self.node_status_queue_rejected_full_total,
rejected_closed_total: self.node_status_queue_rejected_closed_total,
};
samples.extend(node_status_queue.to_metric_samples("tunnel_node_status"));
samples
}
}