fix: harden concurrency limits and high-RPM runtime paths

Bound request, stream, queue, and shutdown resource lifetimes. Reduce scheduler and Redis hot-path work and isolate database maintenance. Include regression coverage, load probes, and concurrency audit results.
This commit is contained in:
elky
2026-09-10 08:14:58 +08:00
parent 361952ada9
commit ecc16673eb
149 changed files with 27963 additions and 1926 deletions
@@ -71,8 +71,13 @@ struct Args {
distributed_request_command_timeout_ms: u64,
}
fn main() -> Result<(), Box<dyn std::error::Error>> {
let _log_shutdown = aether_runtime::LogShutdownGuard::new();
run()
}
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
async fn run() -> Result<(), Box<dyn std::error::Error>> {
let _ = rustls::crypto::ring::default_provider().install_default();
init_service_runtime(ServiceRuntimeConfig::new(
@@ -88,8 +88,13 @@ struct Args {
distributed_request_command_timeout_ms: u64,
}
fn main() -> Result<(), Box<dyn std::error::Error>> {
let _log_shutdown = aether_runtime::LogShutdownGuard::new();
run()
}
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
async fn run() -> Result<(), Box<dyn std::error::Error>> {
init_service_runtime(ServiceRuntimeConfig::new(
"aether-tunnel-standalone",
"aether_gateway=info",
@@ -996,7 +996,7 @@ impl<'a> RequestedModelAttemptPageCursor<'a> {
);
if page_is_exact_auth_api_key_concurrency_limited(&page) {
if self.wait_for_auth_api_key_concurrency_retry().await {
if self.wait_for_auth_api_key_concurrency_retry().await? {
continue;
}
self.persist_final_auth_api_key_concurrency_skips(page.skipped_candidates)
@@ -1087,20 +1087,23 @@ impl<'a> RequestedModelAttemptPageCursor<'a> {
}
}
async fn wait_for_auth_api_key_concurrency_retry(&mut self) -> bool {
async fn wait_for_auth_api_key_concurrency_retry(&mut self) -> Result<bool, GatewayError> {
let now = Instant::now();
let deadline = *self
.auth_api_key_concurrency_wait_deadline
.get_or_insert(now + AUTH_API_KEY_CONCURRENCY_WAIT_BUDGET);
if now >= deadline {
return false;
if !crate::scheduler::candidate::wait_for_auth_api_key_concurrency_retry(
self.state.app(),
Some(&self.auth_snapshot),
deadline,
AUTH_API_KEY_CONCURRENCY_RETRY_DELAY,
)
.await?
{
return Ok(false);
}
let sleep_duration =
AUTH_API_KEY_CONCURRENCY_RETRY_DELAY.min(deadline.saturating_duration_since(now));
tokio::time::sleep(sleep_duration).await;
self.page_cursor.restart_scan();
true
Ok(true)
}
async fn persist_final_auth_api_key_concurrency_skips(
@@ -2289,6 +2292,103 @@ mod tests {
candidate
}
#[tokio::test]
async fn auth_concurrency_wait_paged_scan_retries_once_at_original_deadline() {
let now = current_unix_ms();
let active = serde_json::from_value(json!({
"id": "active-candidate",
"request_id": "active-request",
"api_key_id": "api-key-1",
"candidate_index": 0,
"retry_index": 0,
"status": "pending",
"is_cached": false,
"created_at_unix_ms": now,
"started_at_unix_ms": now
}))
.expect("active candidate should build");
let repository = Arc::new(InMemoryRequestCandidateRepository::seed([active]));
let app = AppState::new()
.expect("state should build")
.with_data_state_for_tests(
GatewayDataState::with_request_candidate_repository_for_tests(repository),
);
let mut auth_snapshot = sample_auth_snapshot();
auth_snapshot.api_key_concurrent_limit = Some(1);
let page_cursor = LocalCandidatePreselectionPageCursor::new(
PlannerAppState::new(&app),
&crate::system_features::ModelDirectivePolicySnapshot::default(),
"openai:chat",
"gpt-5",
None,
true,
None,
&auth_snapshot,
None,
None,
None,
false,
LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModel,
true,
Some("trace-auth-wait"),
)
.await;
let mut cursor = RequestedModelAttemptPageCursor {
state: PlannerAppState::new(&app),
trace_id: "trace-auth-wait".to_string(),
client_api_format: "openai:chat".to_string(),
requested_model: "gpt-5".to_string(),
auth_snapshot,
client_session_affinity: None,
required_capabilities: None,
routing_policy: None,
sticky_session_token: None,
request_auth_channel: None,
skipped_user_id: "user-1".to_string(),
skipped_api_key_id: "api-key-1".to_string(),
skipped_required_capabilities: None,
skipped_error_context: "test auth wait",
record_runtime_miss_diagnostic: false,
resolution_mode: LocalCandidateResolutionMode::Standard,
decorate_skipped_candidate: Arc::new(identity_skipped_candidate),
page_cursor,
pending_items: VecDeque::new(),
skipped_provider_ids: BTreeSet::new(),
skipped_endpoint_ids: BTreeSet::new(),
skipped_credential_ids: BTreeSet::new(),
candidate_count: 0,
next_candidate_index: 0,
remembered_affinity: false,
scheduler_cache_affinity_enabled: false,
auth_api_key_concurrency_wait_deadline: None,
deferred_error: None,
};
let started = Instant::now();
let mut scan_restarts = 0;
while cursor
.wait_for_auth_api_key_concurrency_retry()
.await
.expect("auth wait should succeed")
{
scan_restarts += 1;
}
assert_eq!(
scan_restarts, 1,
"blocked polls must not restart page scans"
);
assert!(started.elapsed() >= AUTH_API_KEY_CONCURRENCY_WAIT_BUDGET);
let original_deadline = cursor.auth_api_key_concurrency_wait_deadline;
assert!(!cursor
.wait_for_auth_api_key_concurrency_retry()
.await
.expect("expired auth wait should succeed"));
assert_eq!(
cursor.auth_api_key_concurrency_wait_deadline,
original_deadline
);
}
#[tokio::test]
async fn pool_group_keys_are_not_persisted_as_available_before_attempt() {
let repository = Arc::new(InMemoryRequestCandidateRepository::default());
@@ -1,9 +1,7 @@
use aether_scheduler_core::{ClientSessionAffinity, SchedulerMinimalCandidateSelectionCandidate};
use std::time::Duration;
use tokio::time::Instant;
use super::{GatewayAuthApiKeySnapshot, PlannerAppState};
use crate::clock::current_unix_secs;
use crate::constants::{
API_KEY_CONCURRENCY_WAIT_POLL_INTERVAL_MS, API_KEY_CONCURRENCY_WAIT_TIMEOUT_MS,
};
@@ -97,11 +95,13 @@ impl<'a> PlannerAppState<'a> {
),
GatewayError,
> {
let wait_timeout = Duration::from_millis(API_KEY_CONCURRENCY_WAIT_TIMEOUT_MS);
let wait_interval = Duration::from_millis(API_KEY_CONCURRENCY_WAIT_POLL_INTERVAL_MS.max(1));
let wait_deadline = Instant::now() + wait_timeout;
let mut attempt_now_unix_secs = now_unix_secs;
loop {
crate::scheduler::candidate::select_with_auth_concurrency_wait(
self.app(),
auth_snapshot,
now_unix_secs,
Duration::from_millis(API_KEY_CONCURRENCY_WAIT_TIMEOUT_MS),
Duration::from_millis(API_KEY_CONCURRENCY_WAIT_POLL_INTERVAL_MS),
|attempt_now_unix_secs| async move {
let result = crate::scheduler::candidate::list_selectable_candidates_with_skip_reasons_for_request_operation(
self.app().data.as_ref(),
self.app(),
@@ -118,21 +118,13 @@ impl<'a> PlannerAppState<'a> {
)
.await?;
if !crate::scheduler::candidate::is_exact_all_skipped_by_auth_limit(
let auth_limit_blocked = crate::scheduler::candidate::is_exact_all_skipped_by_auth_limit(
&result.0, &result.1,
) {
return Ok(result);
}
let now = Instant::now();
if now >= wait_deadline {
return Ok(result);
}
let remaining = wait_deadline.duration_since(now);
tokio::time::sleep(wait_interval.min(remaining)).await;
attempt_now_unix_secs = current_unix_secs();
}
);
Ok((result, auth_limit_blocked))
},
)
.await
}
#[allow(clippy::too_many_arguments)]
@@ -178,13 +170,14 @@ impl<'a> PlannerAppState<'a> {
now_unix_secs: u64,
ordering_config: SchedulerOrderingConfig,
) -> Result<Vec<SchedulerMinimalCandidateSelectionCandidate>, GatewayError> {
let wait_timeout = Duration::from_millis(API_KEY_CONCURRENCY_WAIT_TIMEOUT_MS);
let wait_interval = Duration::from_millis(API_KEY_CONCURRENCY_WAIT_POLL_INTERVAL_MS.max(1));
let wait_deadline = Instant::now() + wait_timeout;
let mut attempt_now_unix_secs = now_unix_secs;
loop {
let (result, auth_limit_blocked) = crate::scheduler::candidate::list_selectable_candidates_for_required_capability_without_requested_model_with_auth_limit_signal(
crate::scheduler::candidate::select_with_auth_concurrency_wait(
self.app(),
auth_snapshot,
now_unix_secs,
Duration::from_millis(API_KEY_CONCURRENCY_WAIT_TIMEOUT_MS),
Duration::from_millis(API_KEY_CONCURRENCY_WAIT_POLL_INTERVAL_MS),
|attempt_now_unix_secs| {
crate::scheduler::candidate::list_selectable_candidates_for_required_capability_without_requested_model_with_auth_limit_signal(
self.app().data.as_ref(),
self.app(),
candidate_api_format,
@@ -195,20 +188,8 @@ impl<'a> PlannerAppState<'a> {
attempt_now_unix_secs,
ordering_config,
)
.await?;
if !auth_limit_blocked {
return Ok(result);
}
let now = Instant::now();
if now >= wait_deadline {
return Ok(result);
}
let remaining = wait_deadline.duration_since(now);
tokio::time::sleep(wait_interval.min(remaining)).await;
attempt_now_unix_secs = current_unix_secs();
}
},
)
.await
}
}
@@ -81,6 +81,19 @@ impl GatewayDataState {
}
}
pub(crate) async fn list_recent_runtime_request_candidates(
&self,
limit: usize,
) -> Result<Vec<StoredRequestCandidate>, DataLayerError> {
match &self.request_candidate_reader {
Some(repository) => repository
.list_recent_runtime(limit)
.await
.map(sanitize_request_candidate_rows),
None => Ok(Vec::new()),
}
}
pub(crate) async fn list_finalized_request_candidates_by_endpoint_ids_since(
&self,
endpoint_ids: &[String],
+25 -2
View File
@@ -19,9 +19,14 @@ use aether_data::repository::management_tokens::{
};
use aether_data::repository::proxy_nodes::{InMemoryProxyNodeRepository, StoredProxyNode};
use aether_data::repository::proxy_nodes::{ProxyNodeReadRepository, ProxyNodeWriteRepository};
use aether_data::repository::routing_profiles::InMemoryRoutingGroupRepository;
use aether_data::repository::users::{
InMemoryUserReadRepository, StoredUserAuthRecord, UserReadRepository,
};
use aether_data_contracts::repository::routing_profiles::{
StoredRoutingGroup, StoredRoutingGroupBinding, StoredRoutingGroupVersion,
};
use aether_routing_core::RoutingGroupConfig;
use sha2::{Digest, Sha256};
use super::{GatewayDataConfig, GatewayDataState};
@@ -142,6 +147,24 @@ impl GatewayDataState {
provider_catalog_repository;
let usage_reader: Arc<dyn UsageReadRepository> = usage_repository.clone();
let usage_writer: Arc<dyn UsageWriteRepository> = usage_repository;
let routing_groups = Arc::new(InMemoryRoutingGroupRepository::seed(
[StoredRoutingGroup {
id: "system-default".to_string(),
name: "system-default".to_string(),
description: Some("pressure harness routing strategy".to_string()),
enabled: true,
is_system_default: true,
sort_order: 0,
config_json: serde_json::to_value(RoutingGroupConfig::default())
.expect("default routing config should serialize"),
version: 1,
created_at: 1,
updated_at: 1,
published_at: Some(1),
}],
std::iter::empty::<StoredRoutingGroupBinding>(),
std::iter::empty::<StoredRoutingGroupVersion>(),
));
Self {
config: GatewayDataConfig::disabled().with_encryption_key(encryption_key),
@@ -174,8 +197,8 @@ impl GatewayDataState {
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
routing_group_reader: Some(routing_groups.clone()),
routing_group_writer: Some(routing_groups),
usage_reader: Some(usage_reader),
usage_writer: Some(usage_writer),
user_reader: None,
@@ -35,7 +35,6 @@ use crate::ai_serving::{
SkippedLocalExecutionCandidate,
};
use crate::clock::current_unix_ms;
use crate::handlers::shared::provider_pool::read_admin_provider_pool_runtime_state;
use crate::handlers::shared::provider_pool::{
admin_provider_pool_cache_affinity_enabled, admin_provider_pool_config_from_config_value,
};
@@ -44,6 +43,9 @@ use crate::handlers::shared::provider_pool::{
read_admin_provider_pool_key_cooldown_reason, AdminProviderPoolConfig,
AdminProviderPoolRuntimeState, AdminProviderPoolSchedulingPreset,
};
use crate::handlers::shared::provider_pool::{
read_provider_pool_scheduling_runtime_state, read_provider_pool_sticky_bound_key_id,
};
use crate::handlers::shared::{parse_catalog_auth_config_json, provider_key_health_summary};
use crate::maintenance::spawn_pool_quota_probe_replenish_for_request;
use crate::orchestration::LocalExecutionCandidateMetadata;
@@ -141,7 +143,7 @@ async fn schedule_pool_page_candidates(
AdminProviderPoolRuntimeState::default()
} else {
let runtime_started_at = std::time::Instant::now();
let runtime = read_admin_provider_pool_runtime_state(
let runtime = read_provider_pool_scheduling_runtime_state(
state.app().runtime_state.as_ref(),
provider_id.as_str(),
&key_ids,
@@ -982,15 +984,13 @@ impl<'a> PoolKeyCursor<'a> {
if !admin_provider_pool_cache_affinity_enabled(&pool_config) {
return None;
}
let runtime = read_admin_provider_pool_runtime_state(
let sticky_key_id = read_provider_pool_sticky_bound_key_id(
self.state.app().runtime_state.as_ref(),
self.group.candidate.provider_id.as_str(),
&[],
&pool_config,
self.sticky_session_token.as_deref(),
)
.await;
let sticky_key_id = runtime.sticky_bound_key_id?;
.await?;
if self
.routing_overlay
.as_ref()
@@ -2028,7 +2028,9 @@ mod tests {
};
use crate::data::GatewayDataState;
use crate::handlers::shared::provider_pool::{
admin_provider_pool_cache_affinity_enabled, record_admin_provider_pool_error,
admin_provider_pool_cache_affinity_enabled, admin_provider_pool_config_from_config_value,
read_admin_provider_pool_runtime_state, read_provider_pool_scheduling_runtime_state,
record_admin_provider_pool_error, record_admin_provider_pool_success,
AdminProviderPoolRuntimeState,
};
use crate::orchestration::LocalExecutionCandidateMetadata;
@@ -2060,6 +2062,111 @@ mod tests {
use std::collections::{BTreeMap, BTreeSet, VecDeque};
use std::sync::Arc;
#[tokio::test]
async fn scheduling_runtime_preserves_pool_ranking_and_cost_rejections() {
let runtime = aether_runtime_state::RuntimeState::memory(
aether_runtime_state::MemoryRuntimeStateConfig::default(),
);
let writer_config = admin_provider_pool_config_from_config_value(Some(&json!({
"pool_advanced": {
"cost_limit_per_key_tokens": 100,
"scheduling_presets": [
{"preset": "cache_affinity", "enabled": true},
{"preset": "latency_first", "enabled": true}
]
}
})))
.expect("writer pool config");
for (key_id, cost, latency) in [("key-a", 100, 10), ("key-b", 20, 100)] {
record_admin_provider_pool_success(
&runtime,
"provider-pool",
key_id,
&writer_config,
Some(key_id),
cost,
Some(latency),
)
.await;
}
let key_ids = vec!["key-a".to_string(), "key-b".to_string()];
for (preset, cost_limit) in [
("cache_affinity", None),
("priority_first", None),
("latency_first", None),
("cost_first", None),
("quota_balanced", None),
("latency_first", Some(100)),
] {
let provider_config = json!({
"pool_advanced": {
"cost_limit_per_key_tokens": cost_limit,
"scheduling_presets": [{"preset": preset, "enabled": true}]
}
});
let pool_config = admin_provider_pool_config_from_config_value(Some(&provider_config))
.expect("reader pool config");
let admin = read_admin_provider_pool_runtime_state(
&runtime,
"provider-pool",
&key_ids,
&pool_config,
Some("key-a"),
)
.await;
let scheduling = read_provider_pool_scheduling_runtime_state(
&runtime,
"provider-pool",
&key_ids,
&pool_config,
Some("key-a"),
)
.await;
let run = |snapshot| {
let candidates = key_ids
.iter()
.map(|key_id| {
sample_eligible_candidate(
"provider-pool",
"endpoint-1",
key_id,
10,
Some(provider_config.clone()),
)
})
.collect();
let (scheduled, skipped) = apply_local_execution_pool_scheduler_with_runtime_map(
candidates,
&BTreeMap::from([("provider-pool".to_string(), snapshot)]),
&BTreeMap::new(),
);
(
scheduled
.into_iter()
.map(|item| item.candidate.key_id)
.collect::<Vec<_>>(),
skipped
.into_iter()
.map(|item| (item.candidate.key_id, item.skip_reason))
.collect::<Vec<_>>(),
)
};
let expected = run(admin);
let actual = run(scheduling);
assert_eq!(
actual, expected,
"preset: {preset}, cost limit: {cost_limit:?}"
);
if cost_limit.is_some() {
assert_eq!(actual.0, vec!["key-b"]);
assert_eq!(
actual.1,
vec![("key-a".to_string(), "pool_cost_limit_reached")]
);
}
}
}
#[test]
fn pool_scheduler_groups_interleaved_candidates_and_reorders_internal_keys() {
let pool_first = sample_eligible_candidate(
@@ -234,7 +234,9 @@ impl Drop for AttemptCancellationGuard {
);
return;
};
let usage_producer = state.usage_runtime.track_producer();
handle.spawn(async move {
let _usage_producer = usage_producer;
settle_cancelled_attempt(state, armed, error_type, error_message).await;
});
}
@@ -674,8 +674,10 @@ impl ExecutionAttemptLifecycle {
let billing_void = settlement.billing.is_void();
let usage_runtime = Arc::clone(&state.usage_runtime);
let usage_data = Arc::clone(state.usage_lifecycle_data_state());
let usage_producer = usage_runtime.track_producer();
self.stage_guard
.await_detachable_stage(self.trace_id.as_str(), "usage_terminal", async move {
let _usage_producer = usage_producer;
usage_runtime
.record_stream_terminal(
usage_data.as_ref(),
@@ -20,6 +20,7 @@ mod response_header_rules;
mod server;
pub(crate) mod stream;
mod stream_pump;
mod stream_read_timeout;
pub(crate) mod submission;
pub(crate) mod sync;
pub(crate) mod transport;
@@ -0,0 +1,251 @@
use std::ops::Deref;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::{Arc, LazyLock};
const DEFAULT_STREAM_CAPTURE_MEMORY_BUDGET_BYTES: usize = 128 * 1024 * 1024;
const STREAM_CAPTURE_MEMORY_BUDGET_ENV: &str = "AETHER_GATEWAY_STREAM_CAPTURE_MEMORY_BUDGET_BYTES";
static STREAM_CAPTURE_BUDGET: LazyLock<Arc<StreamCaptureBudget>> = LazyLock::new(|| {
StreamCaptureBudget::new(
std::env::var(STREAM_CAPTURE_MEMORY_BUDGET_ENV)
.ok()
.and_then(|value| value.trim().parse().ok())
.unwrap_or(DEFAULT_STREAM_CAPTURE_MEMORY_BUDGET_BYTES),
)
});
#[derive(Debug)]
pub(super) struct StreamCaptureBudget {
available: AtomicUsize,
}
impl StreamCaptureBudget {
pub(super) fn new(bytes: usize) -> Arc<Self> {
Arc::new(Self {
available: AtomicUsize::new(bytes),
})
}
fn reserve_up_to(&self, wanted: usize, minimum: usize) -> usize {
let mut available = self.available.load(Ordering::Relaxed);
loop {
let reserved = wanted.min(available);
if reserved < minimum {
return 0;
}
match self.available.compare_exchange_weak(
available,
available - reserved,
Ordering::Relaxed,
Ordering::Relaxed,
) {
Ok(_) => return reserved,
Err(current) => available = current,
}
}
}
fn release(&self, bytes: usize) {
self.available.fetch_add(bytes, Ordering::Relaxed);
}
}
/// Only retained diagnostic bytes belong here. Protocol and billing observers
/// must consume the original chunks independently of capture admission.
#[derive(Debug)]
pub(super) struct StreamBodyCapture {
bytes: Vec<u8>,
budget: Arc<StreamCaptureBudget>,
}
impl Default for StreamBodyCapture {
fn default() -> Self {
Self::with_budget(Arc::clone(&STREAM_CAPTURE_BUDGET))
}
}
impl StreamBodyCapture {
pub(super) fn with_budget(budget: Arc<StreamCaptureBudget>) -> Self {
Self {
bytes: Vec::new(),
budget,
}
}
pub(super) fn append(&mut self, chunk: &[u8], limit: usize, truncated: &mut bool) {
if chunk.is_empty() || *truncated {
return;
}
let wanted_len = self.bytes.len().saturating_add(chunk.len()).min(limit);
if wanted_len > self.bytes.capacity() {
// Keep the old allocation charged until its replacement has been
// allocated and copied, including their overlap during growth.
let wanted_capacity = wanted_len
.max(self.bytes.capacity().saturating_mul(2))
.min(limit);
let reserved = self
.budget
.reserve_up_to(wanted_capacity, self.bytes.capacity().saturating_add(1));
if reserved > 0 {
let mut replacement = Vec::new();
if replacement.try_reserve_exact(reserved).is_ok() {
let extra = replacement.capacity().saturating_sub(reserved);
if extra == 0 || self.budget.reserve_up_to(extra, extra) == extra {
replacement.extend_from_slice(&self.bytes);
let old = std::mem::replace(&mut self.bytes, replacement);
let old_capacity = old.capacity();
drop(old);
self.budget.release(old_capacity);
} else {
drop(replacement);
self.budget.release(reserved);
}
} else {
self.budget.release(reserved);
}
}
}
let keep = wanted_len
.min(self.bytes.capacity())
.saturating_sub(self.bytes.len());
self.bytes.extend_from_slice(&chunk[..keep]);
// Once bytes are omitted, never append a later suffix to this prefix.
*truncated = keep < chunk.len();
}
}
impl Deref for StreamBodyCapture {
type Target = [u8];
fn deref(&self) -> &Self::Target {
&self.bytes
}
}
impl Drop for StreamBodyCapture {
fn drop(&mut self) {
let bytes = std::mem::take(&mut self.bytes);
let capacity = bytes.capacity();
drop(bytes);
self.budget.release(capacity);
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn stream_capture_budget_is_shared_and_released_on_drop() {
let budget = StreamCaptureBudget::new(12);
let mut provider = StreamBodyCapture::with_budget(Arc::clone(&budget));
let mut client = StreamBodyCapture::with_budget(Arc::clone(&budget));
let mut provider_truncated = false;
let mut client_truncated = false;
provider.append(b"12345678", 64, &mut provider_truncated);
client.append(b"abcdefgh", 64, &mut client_truncated);
assert_eq!(&*provider, b"12345678");
assert_eq!(&*client, b"abcd");
assert!(!provider_truncated);
assert!(client_truncated);
assert_eq!(budget.available.load(Ordering::Relaxed), 0);
drop(provider);
client.append(b"later", 64, &mut client_truncated);
assert_eq!(&*client, b"abcd");
drop(client);
assert_eq!(budget.available.load(Ordering::Relaxed), 12);
}
#[test]
fn stream_capture_budget_charges_capacity_and_reallocation_overlap() {
let budget = StreamCaptureBudget::new(16);
let mut capture = StreamBodyCapture::with_budget(Arc::clone(&budget));
let mut truncated = false;
capture.append(b"1234", 64, &mut truncated);
capture.append(b"5", 64, &mut truncated);
assert_eq!(capture.bytes.capacity(), 8);
assert_eq!(budget.available.load(Ordering::Relaxed), 8);
capture.append(b"6789", 64, &mut truncated);
assert_eq!(&*capture, b"12345678");
assert!(truncated);
assert_eq!(budget.available.load(Ordering::Relaxed), 8);
drop(capture);
assert_eq!(budget.available.load(Ordering::Relaxed), 16);
}
#[test]
fn stream_capture_budget_zero_disables_capture_without_allocating() {
let budget = StreamCaptureBudget::new(0);
let mut capture = StreamBodyCapture::with_budget(budget);
let mut truncated = false;
capture.append(b"data", 64, &mut truncated);
assert!(capture.is_empty());
assert_eq!(capture.bytes.capacity(), 0);
assert!(truncated);
}
#[test]
fn stream_capture_budget_exhaustion_uses_existing_spare_capacity() {
let budget = StreamCaptureBudget::new(14);
let mut capture = StreamBodyCapture::with_budget(budget);
let mut truncated = false;
capture.append(b"1234", 64, &mut truncated);
capture.append(b"5", 64, &mut truncated);
assert_eq!(capture.bytes.capacity(), 8);
capture.append(b"6789", 64, &mut truncated);
assert_eq!(&*capture, b"12345678");
assert_eq!(capture.bytes.capacity(), 8);
assert!(truncated);
}
#[test]
fn stream_capture_budget_local_limit_keeps_a_contiguous_prefix() {
let budget = StreamCaptureBudget::new(128);
let mut capture = StreamBodyCapture::with_budget(Arc::clone(&budget));
let mut truncated = false;
capture.append(b"abcdef", 3, &mut truncated);
assert_eq!(&*capture, b"abc");
assert!(truncated);
assert_eq!(budget.available.load(Ordering::Relaxed), 125);
}
#[test]
fn stream_capture_budget_concurrent_growth_and_drop_never_exceeds_capacity() {
const LIMIT: usize = 256;
const THREADS: usize = 8;
let budget = StreamCaptureBudget::new(LIMIT);
let barrier = std::sync::Barrier::new(THREADS);
let held = AtomicUsize::new(0);
std::thread::scope(|scope| {
for index in 0..THREADS {
let budget = &budget;
let barrier = &barrier;
let held = &held;
scope.spawn(move || {
for _ in 0..32 {
let mut capture = StreamBodyCapture::with_budget(Arc::clone(budget));
let mut truncated = false;
barrier.wait();
capture.append(&[1; 16], LIMIT, &mut truncated);
capture.append(&[2; 48], LIMIT, &mut truncated);
held.fetch_add(capture.bytes.capacity(), Ordering::SeqCst);
barrier.wait();
if index == 0 {
let retained = held.load(Ordering::SeqCst);
assert!(retained <= LIMIT);
assert_eq!(retained + budget.available.load(Ordering::Relaxed), LIMIT,);
}
barrier.wait();
drop(capture);
barrier.wait();
if index == 0 {
assert_eq!(budget.available.load(Ordering::Relaxed), LIMIT);
held.store(0, Ordering::SeqCst);
}
barrier.wait();
}
});
}
});
}
}
File diff suppressed because it is too large Load Diff
@@ -1,7 +1,10 @@
mod capture_budget;
mod commit_policy;
mod error;
mod execution;
mod usage_fallback;
pub(crate) use execution::{
execute_execution_runtime_stream, execute_execution_runtime_stream_with_retry_scope,
ClientVisibleStreamCompletionTracker,
};
File diff suppressed because it is too large Load Diff
@@ -12,7 +12,6 @@ use async_stream::stream;
use axum::body::Bytes;
use base64::Engine as _;
use futures_util::{Stream, StreamExt};
use http_body_util::BodyExt;
use serde_json::Value;
use tracing::warn;
@@ -21,9 +20,14 @@ use crate::ai_serving::api::{
normalize_provider_private_report_context, StreamingStandardTerminalObserver,
};
use crate::execution_runtime::ndjson::encode_stream_frame_ndjson;
use crate::execution_runtime::stream::ClientVisibleStreamCompletionTracker;
use crate::execution_runtime::stream_read_timeout::{
await_stream_idle_read, stream_idle_timeout_message,
};
use crate::execution_runtime::transport::{
append_upstream_response_body_chunk, decode_response_body_bytes,
stream_first_byte_timeout_message, DirectUpstreamResponse,
direct_upstream_response_byte_stream, stream_first_byte_timeout_message,
DirectUpstreamResponse,
};
use crate::execution_runtime::DirectUpstreamStreamExecution;
use crate::GatewayError;
@@ -31,6 +35,15 @@ use crate::GatewayError;
const STREAM_USAGE_OBSERVER_MAX_LINE_BYTES: usize = 1024 * 1024;
const UPSTREAM_STREAM_READ_ERROR_MESSAGE: &str = "Upstream response stream failed";
fn upstream_stream_error_category(response: &DirectUpstreamResponse) -> &'static str {
match response {
DirectUpstreamResponse::Reqwest(_) => "reqwest_body_read_failed",
DirectUpstreamResponse::HyperH2c(_) => "hyper_body_read_failed",
DirectUpstreamResponse::BrowserWreq(_) => "browser_body_read_failed",
DirectUpstreamResponse::LocalTunnel(_) => "tunnel_body_read_failed",
}
}
pub(crate) fn build_direct_execution_frame_stream(
execution: DirectUpstreamStreamExecution,
) -> impl Stream<Item = Result<Bytes, IoError>> + Send + 'static {
@@ -49,9 +62,11 @@ pub(crate) fn build_direct_execution_frame_stream(
started_at,
response_observation,
stream_first_byte_timeout,
stream_idle_timeout,
upstream_target_permit,
} = execution;
let _upstream_target_permit = upstream_target_permit;
let upstream_error_category = upstream_stream_error_category(&response);
let mut observer_context = stream_summary_report_context;
if observer_context
@@ -74,6 +89,7 @@ pub(crate) fn build_direct_execution_frame_stream(
let mut private_stream_normalizer =
maybe_build_provider_private_stream_normalizer(Some(&observer_context));
let mut stream_terminal_observer = StreamingStandardTerminalObserver::default();
let mut stream_completion = ClientVisibleStreamCompletionTracker::default();
let mut observer_buffered = Vec::new();
if should_buffer_non_stream_response(
@@ -87,6 +103,7 @@ pub(crate) fn build_direct_execution_frame_stream(
response,
started_at,
stream_first_byte_timeout,
stream_idle_timeout,
)
.await
{
@@ -166,6 +183,7 @@ pub(crate) fn build_direct_execution_frame_stream(
ttfb_ms,
upstream_bytes,
first_byte_timeout,
idle_timeout,
}) => {
match encode_headers_frame(
status_code,
@@ -180,6 +198,8 @@ pub(crate) fn build_direct_execution_frame_stream(
}
let error_frame = if let Some(timeout) = first_byte_timeout {
encode_first_byte_timeout_frame(timeout)
} else if let Some(timeout) = idle_timeout {
encode_idle_timeout_frame(timeout)
} else {
encode_error_frame(message)
};
@@ -228,7 +248,9 @@ pub(crate) fn build_direct_execution_frame_stream(
let mut prefetched_body_failed = false;
for item in prefetched_body {
match item {
Ok(chunk) if chunk.is_empty() => continue,
Ok(chunk) => {
stream_completion.observe_chunk(&chunk);
if ttfb_ms.is_none() {
ttfb_ms = Some(started_at.elapsed().as_millis() as u64);
}
@@ -280,328 +302,97 @@ pub(crate) fn build_direct_execution_frame_stream(
}
}
if !prefetched_body_failed {
match response {
DirectUpstreamResponse::Reqwest(response) => {
let mut bytes_stream = response.bytes_stream();
loop {
let item = if ttfb_ms.is_none() {
match await_stream_first_byte(
bytes_stream.next(),
started_at,
stream_first_byte_timeout,
)
.await
{
Ok(item) => item,
Err(timeout) => {
match encode_first_byte_timeout_frame(timeout) {
Ok(frame) => yield Ok(frame),
Err(err) => {
yield Err(err);
return;
}
}
break;
}
let mut bytes_stream = direct_upstream_response_byte_stream(VecDeque::new(), response);
loop {
let item = if ttfb_ms.is_none() {
match await_stream_first_byte(
bytes_stream.next(), started_at, stream_first_byte_timeout,
).await {
Ok(item) => item,
Err(timeout) => {
match encode_first_byte_timeout_frame(timeout) {
Ok(frame) => yield Ok(frame),
Err(err) => yield Err(err),
}
} else {
bytes_stream.next().await
};
let Some(item) = item else {
break;
};
match item {
Ok(chunk) => {
if ttfb_ms.is_none() {
ttfb_ms = Some(started_at.elapsed().as_millis() as u64);
}
if !first_chunk_telemetry_emitted {
match encode_telemetry_frame(ttfb_ms, ttfb_ms, upstream_bytes) {
Ok(frame) => yield Ok(frame),
Err(err) => {
yield Err(err);
return;
}
}
first_chunk_telemetry_emitted = true;
}
upstream_bytes += chunk.len() as u64;
observe_stream_chunk(
&mut stream_terminal_observer,
&normalized_observer_context,
private_stream_normalizer.as_mut(),
&mut observer_buffered,
chunk.as_ref(),
);
match encode_data_frame(&chunk) {
Ok(frame) => yield Ok(frame),
Err(err) => {
yield Err(err);
return;
}
}
}
Err(_err) => {
let message = UPSTREAM_STREAM_READ_ERROR_MESSAGE.to_string();
warn!(
event_name = "stream_pump_body_read_error",
log_type = "ops",
status_code,
upstream_bytes,
error_category = "reqwest_body_read_failed",
"upstream body stream read error"
);
match encode_error_frame(message) {
Ok(frame) => yield Ok(frame),
Err(encode_err) => {
yield Err(encode_err);
return;
}
}
break;
}
}
}
}
DirectUpstreamResponse::HyperH2c(response) => {
let mut bytes_stream = response.into_body().into_data_stream();
loop {
let item = if ttfb_ms.is_none() {
match await_stream_first_byte(
bytes_stream.next(),
started_at,
stream_first_byte_timeout,
)
.await
} else {
match await_stream_idle_read(bytes_stream.next(), stream_idle_timeout).await {
Ok(item) => item,
Err(timeout) => {
drop(bytes_stream);
if stream_completion.successful_completion()
|| (!stream_completion.observed_terminal()
&& stream_terminal_observer.latest_summary().is_some_and(|summary| {
summary.observed_finish && summary.parser_error.is_none()
&& summary.finish_reason.as_deref() != Some("error")
}))
{
Ok(item) => item,
Err(timeout) => {
match encode_first_byte_timeout_frame(timeout) {
Ok(frame) => yield Ok(frame),
Err(err) => {
yield Err(err);
return;
}
}
break;
}
}
} else {
bytes_stream.next().await
};
let Some(item) = item else {
break;
};
match item {
Ok(chunk) => {
if ttfb_ms.is_none() {
ttfb_ms = Some(started_at.elapsed().as_millis() as u64);
}
if !first_chunk_telemetry_emitted {
match encode_telemetry_frame(ttfb_ms, ttfb_ms, upstream_bytes) {
Ok(frame) => yield Ok(frame),
Err(err) => {
yield Err(err);
return;
}
}
first_chunk_telemetry_emitted = true;
}
upstream_bytes += chunk.len() as u64;
observe_stream_chunk(
&mut stream_terminal_observer,
&normalized_observer_context,
private_stream_normalizer.as_mut(),
&mut observer_buffered,
chunk.as_ref(),
);
match encode_data_frame(&chunk) {
Ok(frame) => yield Ok(frame),
Err(err) => {
yield Err(err);
return;
}
}
}
Err(_err) => {
let message = UPSTREAM_STREAM_READ_ERROR_MESSAGE.to_string();
warn!(
event_name = "stream_pump_body_read_error",
log_type = "ops",
status_code,
upstream_bytes,
error_category = "hyper_body_read_failed",
"upstream body stream read error"
);
match encode_error_frame(message) {
Ok(frame) => yield Ok(frame),
Err(encode_err) => {
yield Err(encode_err);
return;
}
}
break;
}
if stream_terminal_observer.latest_summary().is_some_and(|summary| {
summary.observed_finish && summary.parser_error.is_some()
}) {
// The terminal summary carries the original provider failure.
break;
}
match encode_idle_timeout_frame(timeout) {
Ok(frame) => yield Ok(frame),
Err(err) => yield Err(err),
}
break;
}
}
}
DirectUpstreamResponse::BrowserWreq(response) => {
let mut bytes_stream = response.bytes_stream();
loop {
let item = if ttfb_ms.is_none() {
match await_stream_first_byte(
bytes_stream.next(),
started_at,
stream_first_byte_timeout,
)
.await
{
Ok(item) => item,
Err(timeout) => {
match encode_first_byte_timeout_frame(timeout) {
Ok(frame) => yield Ok(frame),
Err(err) => {
yield Err(err);
return;
}
}
break;
}
}
} else {
bytes_stream.next().await
};
let Some(item) = item else {
break;
};
match item {
Ok(chunk) => {
if ttfb_ms.is_none() {
ttfb_ms = Some(started_at.elapsed().as_millis() as u64);
}
if !first_chunk_telemetry_emitted {
match encode_telemetry_frame(ttfb_ms, ttfb_ms, upstream_bytes) {
Ok(frame) => yield Ok(frame),
Err(err) => {
yield Err(err);
return;
}
}
first_chunk_telemetry_emitted = true;
}
upstream_bytes += chunk.len() as u64;
observe_stream_chunk(
&mut stream_terminal_observer,
&normalized_observer_context,
private_stream_normalizer.as_mut(),
&mut observer_buffered,
chunk.as_ref(),
);
match encode_data_frame(&chunk) {
Ok(frame) => yield Ok(frame),
Err(err) => {
yield Err(err);
return;
}
}
}
Err(_err) => {
let message = UPSTREAM_STREAM_READ_ERROR_MESSAGE.to_string();
warn!(
event_name = "stream_pump_body_read_error",
log_type = "ops",
status_code,
upstream_bytes,
error_category = "browser_body_read_failed",
"upstream body stream read error"
);
match encode_error_frame(message) {
Ok(frame) => yield Ok(frame),
Err(encode_err) => {
yield Err(encode_err);
return;
}
}
break;
}
};
let Some(item) = item else { break };
match item {
Ok(chunk) => {
stream_completion.observe_chunk(&chunk);
if ttfb_ms.is_none() {
ttfb_ms = Some(started_at.elapsed().as_millis() as u64);
}
}
}
DirectUpstreamResponse::LocalTunnel(mut response) => loop {
let item = if ttfb_ms.is_none() {
match await_stream_first_byte(
response.next_chunk(),
started_at,
stream_first_byte_timeout,
)
.await
{
Ok(item) => item,
Err(timeout) => {
match encode_first_byte_timeout_frame(timeout) {
Ok(frame) => yield Ok(frame),
Err(err) => {
yield Err(err);
return;
}
}
break;
}
}
} else {
response.next_chunk().await
};
match item {
Ok(Some(chunk)) => {
if ttfb_ms.is_none() {
ttfb_ms = Some(started_at.elapsed().as_millis() as u64);
}
if !first_chunk_telemetry_emitted {
match encode_telemetry_frame(ttfb_ms, ttfb_ms, upstream_bytes) {
Ok(frame) => yield Ok(frame),
Err(err) => {
yield Err(err);
return;
}
}
first_chunk_telemetry_emitted = true;
}
upstream_bytes += chunk.len() as u64;
observe_stream_chunk(
&mut stream_terminal_observer,
&normalized_observer_context,
private_stream_normalizer.as_mut(),
&mut observer_buffered,
chunk.as_ref(),
);
match encode_data_frame(&chunk) {
if !first_chunk_telemetry_emitted {
match encode_telemetry_frame(ttfb_ms, ttfb_ms, upstream_bytes) {
Ok(frame) => yield Ok(frame),
Err(err) => {
yield Err(err);
return;
}
}
first_chunk_telemetry_emitted = true;
}
Ok(None) => break,
Err(_message) => {
warn!(
event_name = "stream_pump_body_read_error",
log_type = "ops",
status_code,
upstream_bytes,
error_category = "tunnel_body_read_failed",
"upstream body stream read error"
);
match encode_error_frame(UPSTREAM_STREAM_READ_ERROR_MESSAGE.to_string()) {
Ok(frame) => yield Ok(frame),
Err(encode_err) => {
yield Err(encode_err);
return;
}
upstream_bytes += chunk.len() as u64;
observe_stream_chunk(
&mut stream_terminal_observer,
&normalized_observer_context,
private_stream_normalizer.as_mut(),
&mut observer_buffered,
chunk.as_ref(),
);
match encode_data_frame(&chunk) {
Ok(frame) => yield Ok(frame),
Err(err) => {
yield Err(err);
return;
}
break;
}
}
Err(_) => {
warn!(
event_name = "stream_pump_body_read_error",
log_type = "ops",
status_code,
upstream_bytes,
error_category = upstream_error_category,
"upstream body stream read error"
);
match encode_error_frame(UPSTREAM_STREAM_READ_ERROR_MESSAGE.to_string()) {
Ok(frame) => yield Ok(frame),
Err(err) => yield Err(err),
}
break;
}
}
}
}
@@ -704,6 +495,22 @@ fn encode_first_byte_timeout_frame(timeout: Duration) -> Result<Bytes, IoError>
})
}
fn encode_idle_timeout_frame(timeout: Duration) -> Result<Bytes, IoError> {
encode_stream_frame_ndjson(&StreamFrame {
frame_type: StreamFrameType::Error,
payload: StreamFramePayload::Error {
error: ExecutionError {
kind: ExecutionErrorKind::ReadTimeout,
phase: ExecutionPhase::StreamRead,
message: stream_idle_timeout_message(timeout),
upstream_status: None,
retryable: true,
failover_recommended: true,
},
},
})
}
async fn await_stream_first_byte<T, F>(
future: F,
started_at: Instant,
@@ -737,6 +544,7 @@ struct BufferedUpstreamBodyError {
ttfb_ms: Option<u64>,
upstream_bytes: u64,
first_byte_timeout: Option<Duration>,
idle_timeout: Option<Duration>,
}
fn append_buffered_upstream_body_chunk(
@@ -752,6 +560,7 @@ fn append_buffered_upstream_body_chunk(
ttfb_ms,
upstream_bytes: *upstream_bytes,
first_byte_timeout: None,
idle_timeout: None,
}
})
}
@@ -816,16 +625,52 @@ fn should_buffer_non_stream_response(
}
async fn buffer_non_sse_upstream_body(
mut prefetched_body: VecDeque<Result<Bytes, String>>,
prefetched_body: VecDeque<Result<Bytes, String>>,
response: DirectUpstreamResponse,
started_at: Instant,
stream_first_byte_timeout: Option<Duration>,
stream_idle_timeout: Option<Duration>,
) -> Result<BufferedUpstreamBody, BufferedUpstreamBodyError> {
let mut body_bytes = Vec::new();
let mut upstream_bytes = 0u64;
let mut ttfb_ms = None;
while let Some(item) = prefetched_body.pop_front() {
let upstream_error_category = upstream_stream_error_category(&response);
let mut bytes_stream = direct_upstream_response_byte_stream(prefetched_body, response);
loop {
let item = if ttfb_ms.is_none() {
match await_stream_first_byte(
bytes_stream.next(),
started_at,
stream_first_byte_timeout,
)
.await
{
Ok(item) => item,
Err(timeout) => {
return Err(BufferedUpstreamBodyError {
message: stream_first_byte_timeout_message(timeout),
ttfb_ms,
upstream_bytes,
first_byte_timeout: Some(timeout),
idle_timeout: None,
})
}
}
} else {
match await_stream_idle_read(bytes_stream.next(), stream_idle_timeout).await {
Ok(item) => item,
Err(timeout) => {
return Err(BufferedUpstreamBodyError {
message: stream_idle_timeout_message(timeout),
ttfb_ms,
upstream_bytes,
first_byte_timeout: None,
idle_timeout: Some(timeout),
})
}
}
};
let Some(item) = item else { break };
match item {
Ok(chunk) => {
if ttfb_ms.is_none() {
@@ -838,246 +683,24 @@ async fn buffer_non_sse_upstream_body(
&mut upstream_bytes,
)?;
}
Err(_message) => {
Err(_) => {
warn!(
event_name = "stream_pump_body_read_error",
log_type = "ops",
upstream_bytes,
error_category = upstream_error_category,
"upstream body stream read error"
);
return Err(BufferedUpstreamBodyError {
message: UPSTREAM_STREAM_READ_ERROR_MESSAGE.to_string(),
ttfb_ms,
upstream_bytes,
first_byte_timeout: None,
idle_timeout: None,
});
}
}
}
match response {
DirectUpstreamResponse::Reqwest(response) => {
let mut bytes_stream = response.bytes_stream();
loop {
let item = if ttfb_ms.is_none() {
match await_stream_first_byte(
bytes_stream.next(),
started_at,
stream_first_byte_timeout,
)
.await
{
Ok(item) => item,
Err(timeout) => {
return Err(BufferedUpstreamBodyError {
message: stream_first_byte_timeout_message(timeout),
ttfb_ms,
upstream_bytes,
first_byte_timeout: Some(timeout),
});
}
}
} else {
bytes_stream.next().await
};
let Some(item) = item else {
break;
};
match item {
Ok(chunk) => {
if ttfb_ms.is_none() {
ttfb_ms = Some(started_at.elapsed().as_millis() as u64);
}
append_buffered_upstream_body_chunk(
&mut body_bytes,
&chunk,
ttfb_ms,
&mut upstream_bytes,
)?;
}
Err(_err) => {
let message = UPSTREAM_STREAM_READ_ERROR_MESSAGE.to_string();
warn!(
event_name = "stream_pump_body_read_error",
log_type = "ops",
upstream_bytes,
error_category = "reqwest_body_read_failed",
"upstream body stream read error"
);
return Err(BufferedUpstreamBodyError {
message,
ttfb_ms,
upstream_bytes,
first_byte_timeout: None,
});
}
}
}
}
DirectUpstreamResponse::HyperH2c(response) => {
let mut bytes_stream = response.into_body().into_data_stream();
loop {
let item = if ttfb_ms.is_none() {
match await_stream_first_byte(
bytes_stream.next(),
started_at,
stream_first_byte_timeout,
)
.await
{
Ok(item) => item,
Err(timeout) => {
return Err(BufferedUpstreamBodyError {
message: stream_first_byte_timeout_message(timeout),
ttfb_ms,
upstream_bytes,
first_byte_timeout: Some(timeout),
});
}
}
} else {
bytes_stream.next().await
};
let Some(item) = item else {
break;
};
match item {
Ok(chunk) => {
if ttfb_ms.is_none() {
ttfb_ms = Some(started_at.elapsed().as_millis() as u64);
}
append_buffered_upstream_body_chunk(
&mut body_bytes,
&chunk,
ttfb_ms,
&mut upstream_bytes,
)?;
}
Err(_err) => {
let message = UPSTREAM_STREAM_READ_ERROR_MESSAGE.to_string();
warn!(
event_name = "stream_pump_body_read_error",
log_type = "ops",
upstream_bytes,
error_category = "hyper_body_read_failed",
"upstream body stream read error"
);
return Err(BufferedUpstreamBodyError {
message,
ttfb_ms,
upstream_bytes,
first_byte_timeout: None,
});
}
}
}
}
DirectUpstreamResponse::BrowserWreq(response) => {
let mut bytes_stream = response.bytes_stream();
loop {
let item = if ttfb_ms.is_none() {
match await_stream_first_byte(
bytes_stream.next(),
started_at,
stream_first_byte_timeout,
)
.await
{
Ok(item) => item,
Err(timeout) => {
return Err(BufferedUpstreamBodyError {
message: stream_first_byte_timeout_message(timeout),
ttfb_ms,
upstream_bytes,
first_byte_timeout: Some(timeout),
});
}
}
} else {
bytes_stream.next().await
};
let Some(item) = item else {
break;
};
match item {
Ok(chunk) => {
if ttfb_ms.is_none() {
ttfb_ms = Some(started_at.elapsed().as_millis() as u64);
}
append_buffered_upstream_body_chunk(
&mut body_bytes,
&chunk,
ttfb_ms,
&mut upstream_bytes,
)?;
}
Err(_err) => {
let message = UPSTREAM_STREAM_READ_ERROR_MESSAGE.to_string();
warn!(
event_name = "stream_pump_body_read_error",
log_type = "ops",
upstream_bytes,
error_category = "browser_body_read_failed",
"upstream body stream read error"
);
return Err(BufferedUpstreamBodyError {
message,
ttfb_ms,
upstream_bytes,
first_byte_timeout: None,
});
}
}
}
}
DirectUpstreamResponse::LocalTunnel(mut response) => loop {
let item = if ttfb_ms.is_none() {
match await_stream_first_byte(
response.next_chunk(),
started_at,
stream_first_byte_timeout,
)
.await
{
Ok(item) => item,
Err(timeout) => {
return Err(BufferedUpstreamBodyError {
message: stream_first_byte_timeout_message(timeout),
ttfb_ms,
upstream_bytes,
first_byte_timeout: Some(timeout),
});
}
}
} else {
response.next_chunk().await
};
match item {
Ok(Some(chunk)) => {
if ttfb_ms.is_none() {
ttfb_ms = Some(started_at.elapsed().as_millis() as u64);
}
append_buffered_upstream_body_chunk(
&mut body_bytes,
&chunk,
ttfb_ms,
&mut upstream_bytes,
)?;
}
Ok(None) => break,
Err(_message) => {
warn!(
event_name = "stream_pump_body_read_error",
log_type = "ops",
upstream_bytes,
error_category = "tunnel_body_read_failed",
"upstream body stream read error"
);
return Err(BufferedUpstreamBodyError {
message: UPSTREAM_STREAM_READ_ERROR_MESSAGE.to_string(),
ttfb_ms,
upstream_bytes,
first_byte_timeout: None,
});
}
}
},
}
Ok(BufferedUpstreamBody {
body_bytes,
ttfb_ms,
@@ -1605,6 +1228,89 @@ mod tests {
assert_eq!(error.get("failover_recommended"), Some(&Value::Bool(true)));
}
#[tokio::test]
async fn direct_execution_frame_stream_enforces_idle_timeout_after_first_byte() {
for (content_type, first_chunk, expect_timeout, provider_format) in [
("text/event-stream", "data: hello\n\n", true, "openai:chat"),
("application/json", "{\"message\":", true, "openai:chat"),
("text/event-stream", "data: [DONE]\n\n", false, "openai:chat"),
("text/event-stream", "event: response.completed\ndata: {\"type\":\"response.completed\",\"response\":{\"status\":\"completed\"}}\n\n", false, "openai:responses"),
("text/event-stream", "event: response.incomplete\ndata: {\"type\":\"response.incomplete\",\"response\":{}}\n\n", true, "openai:responses"),
] {
let listener = crate::test_support::bind_loopback_listener().await.unwrap();
let addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let mut request = [0_u8; 4096];
socket.read(&mut request).await.unwrap();
let response = if content_type == "application/json" {
format!("HTTP/1.1 200 OK\r\ncontent-type: {content_type}\r\ncontent-length: 1024\r\n\r\n{first_chunk}")
} else {
format!(
"HTTP/1.1 200 OK\r\ncontent-type: {content_type}\r\ntransfer-encoding: chunked\r\n\r\n{:x}\r\n{first_chunk}\r\n",
first_chunk.len(),
)
};
socket.write_all(response.as_bytes()).await.unwrap();
socket.flush().await.unwrap();
tokio::time::sleep(Duration::from_secs(5)).await;
});
let execution = DirectSyncExecutionRuntime::new()
.execute_stream(&ExecutionPlan {
request_id: "req-stream-idle-timeout".into(),
candidate_id: Some("cand-stream-idle-timeout".into()),
provider_name: Some("openai".into()),
provider_id: "prov-1".into(),
endpoint_id: "ep-1".into(),
key_id: "key-1".into(),
method: "POST".into(),
url: format!("http://{addr}/chat"),
headers: BTreeMap::from([("content-type".into(), "application/json".into())]),
content_type: Some("application/json".into()),
content_encoding: None,
body: RequestBody::from_json(serde_json::json!({"stream": true})),
stream: true,
client_api_format: "openai:chat".into(),
provider_api_format: provider_format.into(),
model_name: Some("gpt-5".into()),
proxy: None,
transport_profile: None,
timeouts: Some(ExecutionTimeouts {
first_byte_ms: Some(1_000),
read_ms: Some(10),
..ExecutionTimeouts::default()
}),
})
.await
.expect("stream response headers");
let frames = tokio::time::timeout(
Duration::from_secs(1),
build_direct_execution_frame_stream(execution).collect::<Vec<_>>(),
)
.await;
server.abort();
let frames = frames
.expect("idle timeout must terminate both SSE and buffered JSON")
.into_iter()
.map(|line| serde_json::from_slice::<Value>(&line.unwrap()).unwrap())
.collect::<Vec<_>>();
let errors = frames
.iter()
.filter(|frame| frame["type"] == "error")
.collect::<Vec<_>>();
assert_eq!(errors.len(), usize::from(expect_timeout));
if expect_timeout {
assert_eq!(errors[0]["payload"]["error"]["kind"], "read_timeout");
assert_eq!(errors[0]["payload"]["error"]["phase"], "stream_read");
}
assert!(frames.iter().any(|frame| frame["type"] == "eof"));
assert_eq!(
frames.iter().any(|frame| frame["type"] == "data"),
content_type == "text/event-stream"
);
}
}
#[tokio::test]
async fn direct_execution_frame_stream_emits_telemetry_before_first_data_frame() {
let listener = crate::test_support::bind_loopback_listener()
@@ -0,0 +1,187 @@
use std::future::Future;
use std::time::Duration;
use aether_contracts::ExecutionPlan;
use axum::body::Bytes;
use futures_util::{Stream, StreamExt};
const STREAM_IDLE_TIMEOUT_MS_ENV: &str = "AETHER_GATEWAY_UPSTREAM_STREAM_IDLE_TIMEOUT_MS";
const DEFAULT_STREAM_IDLE_TIMEOUT_MS: u64 = 300_000;
pub(crate) fn resolve_stream_idle_timeout(plan: &ExecutionPlan) -> Option<Duration> {
if !plan.stream {
return None;
}
let configured = std::env::var(STREAM_IDLE_TIMEOUT_MS_ENV).ok();
stream_idle_timeout_from_config(
plan.timeouts.as_ref().and_then(|timeouts| timeouts.read_ms),
configured.as_deref(),
)
}
fn stream_idle_timeout_from_config(
read_ms: Option<u64>,
configured: Option<&str>,
) -> Option<Duration> {
let timeout_ms = read_ms
.or_else(|| configured.and_then(|value| value.trim().parse::<u64>().ok()))
.unwrap_or(DEFAULT_STREAM_IDLE_TIMEOUT_MS);
// Zero explicitly disables the idle limit for providers with long silent reasoning phases.
(timeout_ms > 0).then(|| Duration::from_millis(timeout_ms))
}
pub(crate) fn stream_idle_timeout_message(timeout: Duration) -> String {
format!(
"provider stream idle read timeout after {} ms",
timeout.as_millis()
)
}
pub(crate) async fn await_stream_idle_read<T>(
future: impl Future<Output = T>,
timeout: Option<Duration>,
) -> Result<T, Duration> {
match timeout {
Some(timeout) => tokio::time::timeout(timeout, future)
.await
.map_err(|_| timeout),
None => Ok(future.await),
}
}
pub(crate) fn skip_empty_upstream_chunks<E: Send + 'static>(
upstream: impl Stream<Item = Result<Bytes, E>> + Send + 'static,
) -> impl Stream<Item = Result<Bytes, E>> + Send {
async_stream::stream! {
tokio::pin!(upstream);
while let Some(item) = upstream.next().await {
match item {
Ok(chunk) if chunk.is_empty() => {
// Empty frames are not progress; yield so an always-ready source cannot
// monopolize the executor or prevent its enclosing timeout from firing.
tokio::task::yield_now().await;
}
item => yield item,
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::{
atomic::{AtomicBool, Ordering},
Arc,
};
#[test]
fn idle_timeout_configuration_preserves_provider_override_and_explicit_disable() {
assert_eq!(
stream_idle_timeout_from_config(None, None),
Some(Duration::from_secs(300))
);
assert_eq!(
stream_idle_timeout_from_config(None, Some(" invalid ")),
Some(Duration::from_secs(300))
);
assert_eq!(
stream_idle_timeout_from_config(None, Some(" 600000 ")),
Some(Duration::from_secs(600))
);
assert_eq!(
stream_idle_timeout_from_config(Some(120_000), Some("600000")),
Some(Duration::from_secs(120))
);
assert_eq!(
stream_idle_timeout_from_config(Some(0), Some("600000")),
None
);
assert_eq!(stream_idle_timeout_from_config(None, Some("0")), None);
}
#[tokio::test]
async fn idle_timeout_cancels_the_pending_upstream_read() {
struct DropMarker(Arc<AtomicBool>);
impl Drop for DropMarker {
fn drop(&mut self) {
self.0.store(true, Ordering::SeqCst);
}
}
let dropped = Arc::new(AtomicBool::new(false));
let marker = DropMarker(Arc::clone(&dropped));
let outcome = await_stream_idle_read(
async move {
let _marker = marker;
std::future::pending::<()>().await;
},
Some(Duration::from_millis(5)),
)
.await;
assert_eq!(outcome, Err(Duration::from_millis(5)));
assert!(dropped.load(Ordering::SeqCst));
}
#[tokio::test]
async fn idle_timeout_allows_progressing_stream_to_outlive_one_timeout() {
for _ in 0..3 {
assert_eq!(
await_stream_idle_read(
async {
tokio::time::sleep(Duration::from_millis(10)).await;
1
},
Some(Duration::from_millis(25))
)
.await,
Ok(1)
);
}
assert_eq!(
await_stream_idle_read(
async {
tokio::time::sleep(Duration::from_millis(30)).await;
2
},
None
)
.await,
Ok(2)
);
}
#[tokio::test]
async fn empty_upstream_chunks_do_not_reset_idle_timeout() {
let upstream = futures_util::stream::repeat(Ok::<_, ()>(Bytes::new()));
let filtered = skip_empty_upstream_chunks(upstream);
tokio::pin!(filtered);
let outcome = tokio::time::timeout(
Duration::from_secs(1),
await_stream_idle_read(filtered.next(), Some(Duration::from_millis(5))),
)
.await
.expect("empty ready chunks must yield to the idle timer");
assert_eq!(outcome, Err(Duration::from_millis(5)));
}
#[tokio::test]
async fn downstream_keepalive_ticks_do_not_reset_pending_upstream_idle_timeout() {
let read = await_stream_idle_read(
std::future::pending::<()>(),
Some(Duration::from_millis(30)),
);
tokio::pin!(read);
let mut keepalive = tokio::time::interval(Duration::from_millis(2));
let mut ticks = 0;
loop {
tokio::select! {
result = &mut read => {
assert_eq!(result, Err(Duration::from_millis(30)));
assert!(ticks > 0);
break;
}
_ = keepalive.tick() => { ticks += 1; }
}
}
}
}
@@ -270,7 +270,9 @@ impl Drop for SyncAttemptTerminalGuard {
let candidate_started_unix_ms = self.candidate_started_unix_ms;
let candidate_started_at = self.candidate_started_at;
if let Ok(handle) = tokio::runtime::Handle::try_current() {
let usage_producer = state.usage_runtime.track_producer();
handle.spawn(async move {
let _usage_producer = usage_producer;
record_sync_attempt_forced_terminal_state(
state,
plan,
@@ -36,7 +36,7 @@ use hyper::body::Incoming as HyperIncomingBody;
use hyper::client::conn::http2::SendRequest as HyperH2cSendRequest;
use hyper_util::client::legacy::connect::HttpConnector;
use hyper_util::client::legacy::Client as HyperLegacyClient;
use hyper_util::rt::{TokioExecutor, TokioIo};
use hyper_util::rt::{TokioExecutor, TokioIo, TokioTimer};
use reqwest::header::{HeaderMap, HeaderName, HeaderValue};
use reqwest::redirect::Policy;
use serde::Serialize;
@@ -50,6 +50,7 @@ use tokio::sync::OnceCell as TokioOnceCell;
use crate::ai_serving::api::extract_provider_private_stream_error_body;
#[cfg(test)]
use crate::execution_runtime::remote_compat::execute_sync_plan_via_remote_execution_runtime;
use crate::execution_runtime::stream_read_timeout::resolve_stream_idle_timeout;
use crate::execution_runtime::windsurf::maybe_execute_windsurf_sync;
use crate::frontdoor_loop_guard::{
configured_gateway_frontdoor_base_url, gateway_frontdoor_self_loop_guard_error,
@@ -109,7 +110,7 @@ const DIRECT_REQWEST_PREWARM_SYNC_CLIENTS_ENV: &str =
"AETHER_GATEWAY_DIRECT_REQWEST_PREWARM_SYNC_CLIENTS";
const DEFAULT_H2_TARGET_STREAMS_PER_CLIENT: usize = 8;
const DEFAULT_HTTP1_TARGET_STREAMS_PER_CLIENT: usize = 512;
const DEFAULT_DIRECT_H2C_POOL_MAX_IDLE_PER_HOST: usize = 512;
const DEFAULT_DIRECT_H2C_POOL_MAX_IDLE_PER_HOST: usize = 32;
const DEFAULT_DIRECT_H2C_TARGET_STREAMS_PER_CLIENT: usize = 128;
const DEFAULT_DIRECT_H2C_SENDER_SELECT_WINDOW: usize = 4;
const MAX_DIRECT_H2C_DRIVER_RUNTIME_THREADS: usize = 16;
@@ -382,6 +383,7 @@ static DIRECT_H2C_SENDER_CACHE: LazyLock<
static DIRECT_H2C_POOL_MAX_IDLE_PER_HOST: LazyLock<usize> = LazyLock::new(|| {
env_positive_usize(DIRECT_H2C_POOL_MAX_IDLE_PER_HOST_ENV)
.unwrap_or(DEFAULT_DIRECT_H2C_POOL_MAX_IDLE_PER_HOST)
.min(1024)
});
static DIRECT_H2C_SENDER_SELECT_WINDOW: LazyLock<usize> = LazyLock::new(|| {
@@ -1198,6 +1200,42 @@ pub(crate) enum DirectUpstreamResponse {
LocalTunnel(tunnel::DirectRelayResponse),
}
pub(crate) fn direct_upstream_response_byte_stream(
prefetched_body: VecDeque<Result<Bytes, String>>,
response: DirectUpstreamResponse,
) -> futures_util::stream::BoxStream<'static, Result<Bytes, String>> {
let response_stream = match response {
DirectUpstreamResponse::Reqwest(response) => response
.bytes_stream()
.map(|item| item.map_err(|err| format_upstream_request_error(&err)))
.boxed(),
DirectUpstreamResponse::HyperH2c(response) => response
.into_body()
.into_data_stream()
.map(|item| item.map_err(|err| format_hyper_error_chain(&err)))
.boxed(),
DirectUpstreamResponse::BrowserWreq(response) => response
.bytes_stream()
.map(|item| item.map_err(|err| format_wreq_upstream_request_error(&err)))
.boxed(),
DirectUpstreamResponse::LocalTunnel(mut response) => async_stream::stream! {
loop {
match response.next_chunk().await {
Ok(Some(chunk)) => yield Ok(chunk),
Ok(None) => break,
Err(err) => {
yield Err(err);
break;
}
}
}
}
.boxed(),
};
let upstream = futures_util::stream::iter(prefetched_body).chain(response_stream);
crate::execution_runtime::stream_read_timeout::skip_empty_upstream_chunks(upstream).boxed()
}
pub(crate) struct DirectUpstreamStreamExecution {
pub(crate) request_id: String,
pub(crate) candidate_id: Option<String>,
@@ -1214,6 +1252,7 @@ pub(crate) struct DirectUpstreamStreamExecution {
pub(crate) started_at: Instant,
pub(crate) response_observation: ExecutionResponseObservation,
pub(crate) stream_first_byte_timeout: Option<Duration>,
pub(crate) stream_idle_timeout: Option<Duration>,
pub(crate) upstream_target_permit: Option<UpstreamTargetAdmissionPermit>,
}
@@ -1347,6 +1386,7 @@ impl DirectSyncExecutionRuntime {
request_order_id,
},
stream_first_byte_timeout: resolve_stream_first_byte_timeout(plan),
stream_idle_timeout: resolve_stream_idle_timeout(plan),
upstream_target_permit: None,
})
}
@@ -1494,6 +1534,7 @@ pub(crate) async fn execute_stream_plan_via_local_tunnel(
request_order_id,
},
stream_first_byte_timeout: resolve_stream_first_byte_timeout(plan),
stream_idle_timeout: resolve_stream_idle_timeout(plan),
upstream_target_permit: None,
}))
}
@@ -2593,6 +2634,8 @@ fn build_direct_h2c_client_from_cache_key(
builder.http2_only(true);
builder.http2_adaptive_window(true);
builder.pool_max_idle_per_host(cache_key.pool_max_idle_per_host);
builder.pool_timer(TokioTimer::new());
builder.pool_idle_timeout(Duration::from_millis(upstream_pool_idle_timeout_ms()));
builder.build(connector)
}
@@ -2854,8 +2897,18 @@ async fn send_via_browser_wreq_transport(
let profile = plan.transport_profile.as_ref().ok_or_else(|| {
ExecutionRuntimeTransportError::UnsupportedTransportProfile(String::new())
})?;
let mut client_timeouts = plan.timeouts.clone();
if plan.stream {
if let Some(timeouts) = client_timeouts.as_mut() {
// Streamed responses use the shared idle reader; sync collectors retain
// their existing client read timeout. Zero explicitly disables either.
if apply_request_total_timeout || timeouts.read_ms == Some(0) {
timeouts.read_ms = None;
}
}
}
let client = build_browser_wreq_client(
plan.timeouts.as_ref(),
client_timeouts.as_ref(),
plan.proxy.as_ref(),
profile,
transport_controls,
@@ -4214,6 +4267,7 @@ fn build_direct_reqwest_client_from_cache_key(
&HttpClientConfig {
connect_timeout_ms: cache_key.connect_timeout_ms,
pool_max_idle_per_host: Some(direct_reqwest_pool_max_idle_per_host()),
pool_idle_timeout_ms: Some(upstream_pool_idle_timeout_ms()),
..HttpClientConfig::default()
},
);
@@ -4233,12 +4287,22 @@ fn build_direct_reqwest_client_from_cache_key(
}
fn direct_reqwest_pool_max_idle_per_host() -> usize {
const DEFAULT_MAX_IDLE_PER_HOST: usize = 1024;
const DEFAULT_MAX_IDLE_PER_HOST: usize = 32;
std::env::var("AETHER_GATEWAY_UPSTREAM_POOL_MAX_IDLE_PER_HOST")
.ok()
.and_then(|value| value.trim().parse::<usize>().ok())
.filter(|value| *value > 0)
.unwrap_or(DEFAULT_MAX_IDLE_PER_HOST)
.min(1024)
}
fn upstream_pool_idle_timeout_ms() -> u64 {
std::env::var("AETHER_GATEWAY_UPSTREAM_POOL_IDLE_TIMEOUT_MS")
.ok()
.and_then(|value| value.trim().parse::<u64>().ok())
.filter(|value| *value > 0)
.unwrap_or(15_000)
.min(300_000)
}
pub(crate) fn direct_reqwest_client_cache_metric_samples() -> Vec<MetricSample> {
@@ -4585,7 +4649,11 @@ pub(crate) fn build_browser_wreq_client(
) -> Result<wreq::Client, ExecutionRuntimeTransportError> {
let emulation = browser_wreq_emulation_from_profile(transport_profile)?;
let proxy_url = resolve_proxy_url(proxy)?;
let mut builder = wreq::Client::builder().no_proxy().emulation(emulation);
let mut builder = wreq::Client::builder()
.no_proxy()
.emulation(emulation)
.pool_max_idle_per_host(direct_reqwest_pool_max_idle_per_host())
.pool_idle_timeout(Duration::from_millis(upstream_pool_idle_timeout_ms()));
if proxy_url.is_none() {
builder = builder.dns_resolver(ExecutionSafeDnsResolver);
}
@@ -40,20 +40,6 @@ pub(super) fn pool_stream_timeout_key(provider_id: &str, key_id: &str) -> String
format!("ap:{provider_id}:stream_timeout:{key_id}")
}
pub(super) fn parse_pool_cost_member(member: &str) -> u64 {
member
.rsplit_once(':')
.and_then(|(_, suffix)| suffix.parse::<u64>().ok())
.unwrap_or(0)
}
pub(super) fn parse_pool_latency_member(member: &str) -> u64 {
member
.rsplit_once(':')
.and_then(|(_, suffix)| suffix.parse::<u64>().ok())
.unwrap_or(0)
}
pub(super) fn pool_cooldown_keys(provider_id: &str, key_ids: &[String]) -> Vec<String> {
key_ids
.iter()
@@ -12,7 +12,8 @@ pub(crate) use self::mutations::{
pub(crate) use self::reads::{
read_admin_provider_pool_cooldown_count, read_admin_provider_pool_cooldown_counts,
read_admin_provider_pool_cooldown_key_ids, read_admin_provider_pool_key_cooldown_reason,
read_admin_provider_pool_runtime_state,
read_admin_provider_pool_runtime_state, read_provider_pool_scheduling_runtime_state,
read_provider_pool_sticky_bound_key_id,
};
pub(crate) use self::status::build_admin_provider_pool_status_payload;
pub(crate) use self::writes::{
@@ -1,7 +1,6 @@
use super::keys::{
parse_pool_cost_member, parse_pool_latency_member, pool_cooldown_index_key, pool_cooldown_key,
pool_cooldown_keys, pool_cost_keys, pool_latency_keys, pool_lru_key, pool_sticky_key,
pool_sticky_pattern,
pool_cooldown_index_key, pool_cooldown_key, pool_cooldown_keys, pool_cost_keys,
pool_latency_keys, pool_lru_key, pool_sticky_key, pool_sticky_pattern,
};
use crate::handlers::admin::provider::pool::config::admin_provider_pool_cache_affinity_enabled;
use crate::handlers::admin::provider::shared::support::{
@@ -12,7 +11,8 @@ use crate::maintenance::PoolQuotaProbeWorkerConfig;
use crate::provider_pool_demand::{
provider_pool_burst_pending, read_provider_pool_demand_snapshot,
};
use aether_runtime_state::{DataLayerError, RuntimeState};
use aether_pool_core::{normalize_enabled_pool_presets, PoolSchedulingPreset};
use aether_runtime_state::{DataLayerError, RuntimeState, ScoreWindowU64Stats};
use futures_util::future::join_all;
use std::collections::{BTreeMap, BTreeSet};
use std::time::{SystemTime, UNIX_EPOCH};
@@ -48,6 +48,103 @@ fn bounded_runtime_window_metric_key_ids(key_ids: &[String], limit: usize) -> &[
&key_ids[..end]
}
#[derive(Clone, Copy, PartialEq, Eq)]
enum PoolRuntimeReadPurpose {
Admin,
Scheduling,
}
fn scheduling_window_metrics(pool_config: &AdminProviderPoolConfig) -> (bool, bool) {
let presets = pool_config
.scheduling_presets
.iter()
.map(|preset| PoolSchedulingPreset {
preset: preset.preset.clone(),
enabled: preset.enabled,
mode: preset.mode.clone(),
})
.collect::<Vec<_>>();
let active = normalize_enabled_pool_presets(&presets);
let cost = pool_config.cost_limit_per_key_tokens.is_some()
|| active
.iter()
.any(|preset| matches!(preset.as_str(), "cost_first" | "quota_balanced"));
let latency = active.iter().any(|preset| preset == "latency_first");
(cost, latency)
}
async fn read_window_stats(
runtime: &RuntimeState,
keys: &[String],
min_score: f64,
) -> Vec<ScoreWindowU64Stats> {
let aggregates = match runtime.score_window_u64_stats_by_min(keys, min_score).await {
Ok(values) => values,
Err(err) => {
warn!(
"gateway provider pool: bounded window aggregation failed, using exact range reads: {err:?}"
);
vec![None; keys.len()]
}
};
join_all(keys.iter().zip(aggregates).map(|(key, stats)| async move {
match stats {
Some(stats) => stats,
// Large windows and failed aggregation retain the original exact
// read. A missing aggregate must never be treated as zero cost.
None => {
let members = runtime
.score_range_by_min(key, min_score)
.await
.unwrap_or_default();
ScoreWindowU64Stats::from_members(members.iter().map(String::as_str))
}
}
}))
.await
}
pub(crate) async fn read_provider_pool_sticky_bound_key_id(
runtime: &RuntimeState,
provider_id: &str,
pool_config: &AdminProviderPoolConfig,
sticky_session_token: Option<&str>,
) -> Option<String> {
if pool_config.sticky_session_ttl_seconds == 0
|| !admin_provider_pool_cache_affinity_enabled(pool_config)
{
return None;
}
let sticky_session_token = sticky_session_token
.map(str::trim)
.filter(|value| !value.is_empty())?;
let sticky_key = pool_sticky_key(provider_id, sticky_session_token);
let bound_key_id = runtime.kv_get(&sticky_key).await.ok().flatten()?;
let cooldown_key = pool_cooldown_key(provider_id, &bound_key_id);
match runtime.kv_exists(&cooldown_key).await {
Ok(false) => {
let _ = runtime
.key_expire(
&sticky_key,
std::time::Duration::from_secs(pool_config.sticky_session_ttl_seconds),
)
.await;
Some(bound_key_id)
}
Ok(true) => {
let _ = runtime.kv_delete(&sticky_key).await;
None
}
Err(err) => {
warn!(
"gateway admin provider pool: failed to validate sticky cooldown for provider {provider_id}: {:?}",
err
);
Some(bound_key_id)
}
}
}
pub(crate) async fn read_admin_provider_pool_cooldown_counts(
runtime: &RuntimeState,
provider_ids: &[String],
@@ -71,9 +168,52 @@ pub(crate) async fn read_admin_provider_pool_runtime_state(
pool_config: &AdminProviderPoolConfig,
sticky_session_token: Option<&str>,
) -> AdminProviderPoolRuntimeState {
read_provider_pool_runtime_state(
runtime,
provider_id,
key_ids,
pool_config,
sticky_session_token,
PoolRuntimeReadPurpose::Admin,
)
.await
}
pub(crate) async fn read_provider_pool_scheduling_runtime_state(
runtime: &RuntimeState,
provider_id: &str,
key_ids: &[String],
pool_config: &AdminProviderPoolConfig,
sticky_session_token: Option<&str>,
) -> AdminProviderPoolRuntimeState {
read_provider_pool_runtime_state(
runtime,
provider_id,
key_ids,
pool_config,
sticky_session_token,
PoolRuntimeReadPurpose::Scheduling,
)
.await
}
async fn read_provider_pool_runtime_state(
runtime: &RuntimeState,
provider_id: &str,
key_ids: &[String],
pool_config: &AdminProviderPoolConfig,
sticky_session_token: Option<&str>,
purpose: PoolRuntimeReadPurpose,
) -> AdminProviderPoolRuntimeState {
let include_admin_metrics = purpose == PoolRuntimeReadPurpose::Admin;
let mut state = AdminProviderPoolRuntimeState::default();
let cooldown_keys = pool_cooldown_keys(provider_id, key_ids);
let metric_key_limit = pool_runtime_window_metric_key_limit();
let metric_key_limit = if include_admin_metrics {
pool_runtime_window_metric_key_limit()
} else {
key_ids.len()
};
// The admin display cap must not hide a candidate's strict cost limit.
let metric_key_ids = bounded_runtime_window_metric_key_ids(key_ids, metric_key_limit);
if metric_key_ids.len() < key_ids.len() {
info!(
@@ -86,44 +226,33 @@ pub(crate) async fn read_admin_provider_pool_runtime_state(
"gateway limited admin pool runtime cost/latency window reads"
);
}
let cost_keys = pool_cost_keys(provider_id, metric_key_ids);
let latency_keys = pool_latency_keys(provider_id, metric_key_ids);
let (load_cost, load_latency) = if include_admin_metrics {
(true, true)
} else {
scheduling_window_metrics(pool_config)
};
let cost_keys = if load_cost {
pool_cost_keys(provider_id, metric_key_ids)
} else {
Vec::new()
};
let latency_keys = if load_latency {
pool_latency_keys(provider_id, metric_key_ids)
} else {
Vec::new()
};
let sticky_sessions_enabled = pool_config.sticky_session_ttl_seconds > 0
&& admin_provider_pool_cache_affinity_enabled(pool_config);
if let Some(sticky_session_token) = sticky_session_token
.map(str::trim)
.filter(|value| !value.is_empty())
.filter(|_| sticky_sessions_enabled)
{
let sticky_key = pool_sticky_key(provider_id, sticky_session_token);
if let Ok(Some(bound_key_id)) = runtime.kv_get(&sticky_key).await {
let cooldown_key = pool_cooldown_key(provider_id, &bound_key_id);
match runtime.kv_exists(&cooldown_key).await {
Ok(false) => {
let _ = runtime
.key_expire(
&sticky_key,
std::time::Duration::from_secs(pool_config.sticky_session_ttl_seconds),
)
.await;
state.sticky_bound_key_id = Some(bound_key_id);
}
Ok(true) => {
let _ = runtime.kv_delete(&sticky_key).await;
}
Err(err) => {
warn!(
"gateway admin provider pool: failed to validate sticky cooldown for provider {provider_id}: {:?}",
err
);
state.sticky_bound_key_id = Some(bound_key_id);
}
}
}
}
state.sticky_bound_key_id = read_provider_pool_sticky_bound_key_id(
runtime,
provider_id,
pool_config,
sticky_session_token,
)
.await;
if sticky_sessions_enabled {
if include_admin_metrics && sticky_sessions_enabled {
let sticky_keys = runtime
.scan_keys(&pool_sticky_pattern(provider_id), 200)
.await
@@ -161,23 +290,25 @@ pub(crate) async fn read_admin_provider_pool_runtime_state(
.unwrap_or_default();
}
let probe_config = PoolQuotaProbeWorkerConfig::from_env();
let demand_snapshot = read_provider_pool_demand_snapshot(
runtime,
provider_id,
key_ids.len(),
probe_config.max_keys_per_provider,
)
.await;
state.provider_in_flight = demand_snapshot.in_flight;
state.provider_ema_in_flight = demand_snapshot.ema_in_flight;
state.provider_desired_hot = if pool_config.probing_enabled {
demand_snapshot.desired_hot
} else {
0
};
state.provider_burst_pending =
pool_config.probing_enabled && provider_pool_burst_pending(runtime, provider_id).await;
if include_admin_metrics || pool_config.probing_enabled {
let probe_config = PoolQuotaProbeWorkerConfig::from_env();
let demand_snapshot = read_provider_pool_demand_snapshot(
runtime,
provider_id,
key_ids.len(),
probe_config.max_keys_per_provider,
)
.await;
state.provider_in_flight = demand_snapshot.in_flight;
state.provider_ema_in_flight = demand_snapshot.ema_in_flight;
state.provider_desired_hot = if pool_config.probing_enabled {
demand_snapshot.desired_hot
} else {
0
};
state.provider_burst_pending =
pool_config.probing_enabled && provider_pool_burst_pending(runtime, provider_id).await;
}
if !cooldown_keys.is_empty() {
let cooldown_reasons = runtime
@@ -190,12 +321,14 @@ pub(crate) async fn read_admin_provider_pool_runtime_state(
{
if let Some(reason) = reason {
state.cooldown_reason_by_key.insert(key_id.clone(), reason);
if let Ok(Some(ttl)) = runtime.kv_ttl_seconds(cooldown_key).await {
if let Ok(ttl_seconds) = u64::try_from(ttl) {
if ttl_seconds > 0 {
state
.cooldown_ttl_by_key
.insert(key_id.clone(), ttl_seconds);
if include_admin_metrics {
if let Ok(Some(ttl)) = runtime.kv_ttl_seconds(cooldown_key).await {
if let Ok(ttl_seconds) = u64::try_from(ttl) {
if ttl_seconds > 0 {
state
.cooldown_ttl_by_key
.insert(key_id.clone(), ttl_seconds);
}
}
}
}
@@ -205,42 +338,24 @@ pub(crate) async fn read_admin_provider_pool_runtime_state(
let now = current_unix_secs();
let cost_window_start = now.saturating_sub(pool_config.cost_window_seconds) as f64;
let cost_results = join_all(
cost_keys
.iter()
.map(|cost_key| runtime.score_range_by_min(cost_key, cost_window_start)),
)
.await;
for (key_id, members) in metric_key_ids.iter().zip(cost_results) {
let total = members
.unwrap_or_default()
.iter()
.map(|member| parse_pool_cost_member(member))
.sum::<u64>();
if total > 0 {
state.cost_window_usage_by_key.insert(key_id.clone(), total);
let latency_window_start = now.saturating_sub(pool_config.latency_window_seconds) as f64;
let (cost_results, latency_results) = tokio::join!(
read_window_stats(runtime, &cost_keys, cost_window_start),
read_window_stats(runtime, &latency_keys, latency_window_start),
);
for (key_id, stats) in metric_key_ids.iter().zip(cost_results) {
if stats.sum > 0 {
state
.cost_window_usage_by_key
.insert(key_id.clone(), stats.sum);
}
}
let latency_window_start = now.saturating_sub(pool_config.latency_window_seconds) as f64;
let latency_results = join_all(
latency_keys
.iter()
.map(|latency_key| runtime.score_range_by_min(latency_key, latency_window_start)),
)
.await;
for (key_id, members) in metric_key_ids.iter().zip(latency_results) {
let samples = members
.unwrap_or_default()
.iter()
.map(|member| parse_pool_latency_member(member))
.filter(|value| *value > 0)
.collect::<Vec<_>>();
if samples.is_empty() {
for (key_id, stats) in metric_key_ids.iter().zip(latency_results) {
if stats.positive_count == 0 {
continue;
}
let total = samples.iter().sum::<u64>() as f64;
let average = total / samples.len() as f64;
let average = stats.sum as f64 / stats.positive_count as f64;
if average.is_finite() && average >= 0.0 {
state.latency_avg_ms_by_key.insert(key_id.clone(), average);
}
@@ -300,7 +415,358 @@ pub(crate) async fn read_admin_provider_pool_key_cooldown_reason(
#[cfg(test)]
mod tests {
use super::bounded_runtime_window_metric_key_ids;
use super::super::keys::{pool_cooldown_key, pool_cost_key, pool_latency_key, pool_sticky_key};
use super::{
bounded_runtime_window_metric_key_ids, current_unix_secs,
read_admin_provider_pool_runtime_state, read_provider_pool_scheduling_runtime_state,
read_provider_pool_sticky_bound_key_id,
};
use crate::handlers::admin::provider::pool::config::admin_provider_pool_config_from_config_value;
use crate::handlers::admin::provider::shared::support::AdminProviderPoolConfig;
use aether_runtime_state::{MemoryRuntimeStateConfig, RedisClientConfig, RuntimeState};
use aether_test_support::ManagedRedisServer;
use serde_json::json;
use std::time::Duration;
fn config(value: serde_json::Value) -> AdminProviderPoolConfig {
admin_provider_pool_config_from_config_value(Some(&json!({ "pool_advanced": value })))
.expect("pool config")
}
async fn seed_window_metrics(runtime: &RuntimeState, provider_id: &str, key_id: &str) {
let now = current_unix_secs() as f64;
for (key, member, timestamp) in [
(pool_cost_key(provider_id, key_id), "current:70", now),
(pool_cost_key(provider_id, key_id), "earlier:30", now - 1.0),
(
pool_cost_key(provider_id, key_id),
"expired:999",
now - 20_000.0,
),
(pool_latency_key(provider_id, key_id), "first:10", now),
(
pool_latency_key(provider_id, key_id),
"second:30",
now - 1.0,
),
] {
runtime
.score_set(&key, member, timestamp)
.await
.expect("seed window");
}
}
async fn admin_command_count(runtime: &RuntimeState) -> u64 {
runtime
.redis_diagnostics()
.await
.expect("diagnostics")
.expect("Redis runtime")
.lanes
.into_iter()
.find(|lane| lane.lane == "admin")
.expect("admin lane")
.command_count
}
#[tokio::test]
async fn scheduling_runtime_aggregates_bounded_windows_and_falls_back_for_large_windows() {
let redis = match ManagedRedisServer::start().await {
Ok(server) => server,
Err(err) if err.to_string().contains("No such file or directory") => {
eprintln!("skipping redis-backed scheduling runtime test: {err}");
return;
}
Err(err) => panic!("start Redis: {err}"),
};
let runtime = RuntimeState::redis(
RedisClientConfig {
url: redis.redis_url().to_string(),
key_prefix: Some("pool-window-aggregation-test".to_string()),
},
Some(1_000),
)
.await
.expect("runtime Redis");
let keys = vec!["bounded".to_string(), "large".to_string()];
let now = current_unix_secs() as f64;
for (key_id, count) in [(&keys[0], 512), (&keys[1], 2048)] {
let cost_key = pool_cost_key("pool", key_id);
for index in 0..count {
runtime
.score_set(&cost_key, &format!("{index}:100"), now)
.await
.expect("seed cost window");
}
runtime
.score_set(&cost_key, "expired:9999999", now - 20_000.0)
.await
.expect("expired cost");
for (member, score) in [("first:10", now), ("second:30", now), ("zero:0", now)] {
runtime
.score_set(&pool_latency_key("pool", key_id), member, score)
.await
.expect("seed latency");
}
}
let pool_config = config(json!({
"cost_limit_per_key_tokens": 50_000,
"cost_window_seconds": 600,
"latency_window_seconds": 600,
"scheduling_presets": [{"preset": "latency_first", "enabled": true}]
}));
let scheduled = read_provider_pool_scheduling_runtime_state(
&runtime,
"pool",
&keys,
&pool_config,
None,
)
.await;
assert_eq!(
scheduled.cost_window_usage_by_key.get("bounded"),
Some(&51_200)
);
assert_eq!(
scheduled.cost_window_usage_by_key.get("large"),
Some(&204_800)
);
assert_eq!(scheduled.latency_avg_ms_by_key.get("bounded"), Some(&20.0));
assert_eq!(scheduled.latency_avg_ms_by_key.get("large"), Some(&20.0));
runtime
.score_remove_by_score(&pool_cost_key("pool", "large"), f64::INFINITY)
.await
.expect("reset window");
runtime
.score_set(&pool_cost_key("pool", "large"), "after-reset:75", now)
.await
.expect("post-reset cost");
let reset = read_provider_pool_scheduling_runtime_state(
&runtime,
"pool",
&keys,
&pool_config,
None,
)
.await;
assert_eq!(reset.cost_window_usage_by_key.get("large"), Some(&75));
runtime
.kv_set(&pool_cost_key("pool", "large"), "wrong-type", None)
.await
.expect("simulate invalid metric key");
let partial = read_provider_pool_scheduling_runtime_state(
&runtime,
"pool",
&keys,
&pool_config,
None,
)
.await;
assert_eq!(
partial.cost_window_usage_by_key.get("bounded"),
Some(&51_200),
"one failed aggregate must not discard another key's strict cost check"
);
}
#[tokio::test]
async fn scheduling_runtime_skips_admin_scan_and_unused_window_queries() {
let redis = match ManagedRedisServer::start().await {
Ok(server) => server,
Err(err) if err.to_string().contains("No such file or directory") => {
eprintln!("skipping redis-backed scheduling runtime test: {err}");
return;
}
Err(err) => panic!("Redis server should start: {err}"),
};
let runtime = RuntimeState::redis(
RedisClientConfig {
url: redis.redis_url().to_string(),
key_prefix: Some("scheduling-runtime-reads".to_string()),
},
Some(2_000),
)
.await
.expect("Redis runtime");
let pool_config = config(json!({}));
let keys = vec!["ready".to_string(), "cooling".to_string()];
seed_window_metrics(&runtime, "pool", "ready").await;
for session in ["current", "other"] {
runtime
.kv_set(
&pool_sticky_key("pool", session),
"ready".to_string(),
Some(Duration::from_secs(60)),
)
.await
.expect("seed sticky session");
}
runtime
.kv_set(
&pool_cooldown_key("pool", "cooling"),
"rate_limit".to_string(),
Some(Duration::from_secs(60)),
)
.await
.expect("seed cooldown");
let before = admin_command_count(&runtime).await;
let scheduled = read_provider_pool_scheduling_runtime_state(
&runtime,
"pool",
&keys,
&pool_config,
Some("current"),
)
.await;
let after = admin_command_count(&runtime).await;
assert_eq!(
after - before,
1,
"only the diagnostics INFO may use the admin lane"
);
assert_eq!(scheduled.sticky_bound_key_id.as_deref(), Some("ready"));
assert_eq!(
scheduled
.cooldown_reason_by_key
.get("cooling")
.map(String::as_str),
Some("rate_limit")
);
assert!(scheduled.cooldown_ttl_by_key.is_empty());
assert!(scheduled.cost_window_usage_by_key.is_empty());
assert!(scheduled.latency_avg_ms_by_key.is_empty());
assert_eq!(scheduled.total_sticky_sessions, 0);
let admin = read_admin_provider_pool_runtime_state(
&runtime,
"pool",
&keys,
&pool_config,
Some("current"),
)
.await;
assert_eq!(admin.total_sticky_sessions, 2);
assert_eq!(admin.sticky_sessions_by_key.get("ready"), Some(&2));
assert_eq!(admin.cost_window_usage_by_key.get("ready"), Some(&100));
assert_eq!(admin.latency_avg_ms_by_key.get("ready"), Some(&20.0));
assert!(admin
.cooldown_ttl_by_key
.get("cooling")
.is_some_and(|ttl| *ttl > 0));
}
#[tokio::test]
async fn scheduling_runtime_loads_only_metrics_used_by_enabled_strategies() {
let runtime = RuntimeState::memory(MemoryRuntimeStateConfig::default());
let keys = vec!["key".to_string()];
seed_window_metrics(&runtime, "pool", "key").await;
for (value, expected_cost, expected_latency) in [
(json!({}), false, false),
(json!({"cost_limit_per_key_tokens": 100}), true, false),
(json!({"cost_limit_per_key_tokens": 0}), true, false),
(
json!({"scheduling_presets": [{"preset": "cost_first", "enabled": true}]}),
true,
false,
),
(
json!({"scheduling_presets": [{"preset": "quota_balanced", "enabled": true}]}),
true,
false,
),
(
json!({"scheduling_presets": [{"preset": "latency_first", "enabled": true}]}),
false,
true,
),
(
json!({"scheduling_presets": [
{"preset": "cost_first", "enabled": false},
{"preset": "latency_first", "enabled": false}
]}),
false,
false,
),
] {
let pool_config = config(value.clone());
let snapshot = read_provider_pool_scheduling_runtime_state(
&runtime,
"pool",
&keys,
&pool_config,
None,
)
.await;
assert_eq!(
snapshot.cost_window_usage_by_key.get("key").copied(),
expected_cost.then_some(100),
"config: {value}"
);
assert_eq!(
snapshot.latency_avg_ms_by_key.get("key").copied(),
expected_latency.then_some(20.0),
"config: {value}"
);
}
}
#[tokio::test]
async fn scheduling_runtime_checks_cost_for_candidates_beyond_admin_display_limit() {
let runtime = RuntimeState::memory(MemoryRuntimeStateConfig::default());
let keys = (0..513)
.map(|index| format!("key-{index}"))
.collect::<Vec<_>>();
seed_window_metrics(&runtime, "pool", &keys[512]).await;
let pool_config = config(json!({ "cost_limit_per_key_tokens": 100 }));
let snapshot = read_provider_pool_scheduling_runtime_state(
&runtime,
"pool",
&keys,
&pool_config,
None,
)
.await;
assert_eq!(
snapshot.cost_window_usage_by_key.get(&keys[512]),
Some(&100)
);
}
#[tokio::test]
async fn scheduling_sticky_lookup_invalidates_a_cooled_down_binding() {
let runtime = RuntimeState::memory(MemoryRuntimeStateConfig::default());
let pool_config = config(json!({}));
let sticky_key = pool_sticky_key("pool", "session");
runtime
.kv_set(&sticky_key, "key".to_string(), None)
.await
.expect("sticky session");
runtime
.kv_set(
&pool_cooldown_key("pool", "key"),
"rate_limit".to_string(),
None,
)
.await
.expect("cooldown");
assert!(read_provider_pool_sticky_bound_key_id(
&runtime,
"pool",
&pool_config,
Some("session")
)
.await
.is_none());
assert!(!runtime
.kv_exists(&sticky_key)
.await
.expect("sticky existence"));
}
#[test]
fn runtime_window_metric_key_ids_are_bounded() {
@@ -1937,6 +1937,7 @@ pub(crate) async fn start_admin_system_rollback_task(
}
fn request_process_restart() -> ! {
let _ = aether_runtime::shutdown_logging(std::time::Duration::from_secs(2));
std::process::exit(RESTART_EXIT_CODE);
}
@@ -225,6 +225,9 @@ impl RequestBodyBufferError {
RequestBodyNormalizationError::RequestBodyTooLarge { .. } => {
"request_body_too_large"
}
RequestBodyNormalizationError::BodyBufferOverloaded { .. } => {
"request_body_buffer_overloaded"
}
},
Self::TooLarge { .. } => "request_body_too_large",
Self::Overloaded { .. } => "request_body_buffer_overloaded",
@@ -292,15 +295,37 @@ pub(super) async fn buffer_and_normalize_request_body(
.await
.map_err(RequestBodyBufferError::from)?;
let elapsed_ms = buffered.elapsed().as_millis() as u64;
let retained_input_capacity = buffered
.requested_bytes()
.saturating_sub(buffered.bytes().len());
let normalized = buffered
.try_map(|body| {
crate::headers::normalize_request_body_headers_and_bytes_with_limit(
.try_map_with_budget(|body, memory| {
crate::headers::normalize_request_body_headers_and_bytes_with_budget(
headers,
body,
policy.effective_max_bytes(),
&mut |requested_bytes| {
let requested_bytes = requested_bytes.saturating_add(retained_input_capacity);
memory.try_reserve_bytes(requested_bytes).map_err(|_| {
RequestBodyNormalizationError::BodyBufferOverloaded {
requested_bytes,
budget_bytes: policy.budget_bytes(),
}
})
},
)
})
.map_err(RequestBodyBufferError::Normalization)?;
.map_err(|error| match error {
RequestBodyNormalizationError::BodyBufferOverloaded {
requested_bytes,
budget_bytes,
} => RequestBodyBufferError::Overloaded {
requested_bytes,
budget_bytes,
timeout_ms: 0,
},
error => RequestBodyBufferError::Normalization(error),
})?;
info!(
event_name = "frontdoor_request_body_buffer_completed",
log_type = "event",
+111 -5
View File
@@ -1035,11 +1035,10 @@ pub(crate) async fn proxy_request(
ConnectInfo(remote_addr): ConnectInfo<std::net::SocketAddr>,
request: Request,
) -> Result<Response<Body>, GatewayError> {
crate::request_lifecycle::run_request(Box::pin(proxy_request_inner(
state,
remote_addr,
request,
)))
crate::request_lifecycle::run_request_with_usage(
state.usage_runtime.clone(),
Box::pin(proxy_request_inner(state, remote_addr, request)),
)
.await
}
@@ -3269,6 +3268,113 @@ mod tests {
));
}
#[tokio::test]
async fn request_body_buffer_allows_parallel_compressed_uploads() {
let mut encoder = GzEncoder::new(Vec::new(), Compression::default());
encoder.write_all(br#"{"model":"test"}"#).unwrap();
let encoded = Bytes::from(encoder.finish().unwrap());
let budget_bytes = 2 * crate::state::REQUEST_BODY_BUFFER_PERMIT_BYTES;
let budget = Arc::new(Semaphore::new(2));
let policy = RequestBodyBufferPolicy::for_tests_with_budget(
budget_bytes as u64,
Duration::from_secs(1),
Duration::from_millis(50),
budget_bytes,
Arc::clone(&budget),
);
let mut headers = HeaderMap::new();
headers.insert(header::CONTENT_ENCODING, HeaderValue::from_static("gzip"));
headers.insert(header::CONTENT_LENGTH, HeaderValue::from(encoded.len()));
let (started_tx, started_rx) = tokio::sync::oneshot::channel();
let (finish_tx, finish_rx) = tokio::sync::oneshot::channel();
let first_policy = policy.clone();
let mut first_headers = headers.clone();
let first_encoded = encoded.clone();
let first = async move {
let stream = async_stream::stream! {
let middle = first_encoded.len() / 2;
yield Ok::<_, std::io::Error>(first_encoded.slice(..middle));
let _ = started_tx.send(());
let _ = finish_rx.await;
yield Ok(first_encoded.slice(middle..));
};
buffer_and_normalize_request_body(
&mut Some(Body::from_stream(stream)),
&mut first_headers,
"test owns body",
"trace-compressed-first",
&Method::POST,
"/v1/responses",
"test",
first_policy,
)
.await
};
let second = async move {
started_rx.await.unwrap();
let result = buffer_and_normalize_request_body(
&mut Some(Body::from(encoded)),
&mut headers,
"test owns body",
"trace-compressed-second",
&Method::POST,
"/v1/responses",
"test",
policy,
)
.await;
let _ = finish_tx.send(());
result
};
let (first, second) = tokio::time::timeout(Duration::from_secs(2), async {
tokio::join!(first, second)
})
.await
.expect("concurrent compressed requests should finish");
assert_eq!(first.unwrap().as_ref(), br#"{"model":"test"}"#);
assert_eq!(second.unwrap().as_ref(), br#"{"model":"test"}"#);
assert_eq!(budget.available_permits(), 2);
}
#[tokio::test]
async fn request_body_buffer_rejects_decompression_growth_when_budget_is_busy() {
let mut encoder = GzEncoder::new(Vec::new(), Compression::default());
encoder.write_all(&vec![b'a'; 100_000]).unwrap();
let encoded = encoder.finish().unwrap();
let budget_bytes = 2 * crate::state::REQUEST_BODY_BUFFER_PERMIT_BYTES;
let budget = Arc::new(Semaphore::new(2));
let held = Arc::clone(&budget).acquire_owned().await.unwrap();
let mut headers = HeaderMap::new();
headers.insert(header::CONTENT_ENCODING, HeaderValue::from_static("gzip"));
headers.insert(header::CONTENT_LENGTH, HeaderValue::from(encoded.len()));
let result = buffer_and_normalize_request_body(
&mut Some(Body::from(encoded)),
&mut headers,
"test owns body",
"trace-decompression-overload",
&Method::POST,
"/v1/responses",
"test",
RequestBodyBufferPolicy::for_tests_with_budget(
budget_bytes as u64,
Duration::from_secs(1),
Duration::from_secs(1),
budget_bytes,
Arc::clone(&budget),
),
)
.await
.unwrap_err();
assert!(matches!(
result,
RequestBodyBufferError::Overloaded { timeout_ms: 0, .. }
));
assert_eq!(result.http_status(), http::StatusCode::SERVICE_UNAVAILABLE);
assert_eq!(budget.available_permits(), 1);
drop(held);
assert_eq!(budget.available_permits(), 2);
}
#[tokio::test]
async fn request_body_buffer_times_out_instead_of_waiting_forever() {
let stream = async_stream::stream! {
@@ -508,7 +508,9 @@ async fn persist_live_audit_event(
let usage_runtime = std::sync::Arc::clone(&state.usage_runtime);
let usage_data = std::sync::Arc::clone(state.usage_lifecycle_data_state());
let write_request_id = request_id.clone();
let usage_producer = usage_runtime.track_producer();
let task = tokio::spawn(async move {
let _usage_producer = usage_producer;
if tokio::time::timeout(
LIVE_AUDIT_WRITE_HARD_TIMEOUT,
usage_runtime.record_terminal_event_direct(usage_data.as_ref(), event),
@@ -572,7 +574,9 @@ fn spawn_live_audit_event_detached(state: &AppState, event: UsageEvent, audit_sc
};
let usage_runtime = std::sync::Arc::clone(&state.usage_runtime);
let usage_data = std::sync::Arc::clone(state.usage_lifecycle_data_state());
let usage_producer = usage_runtime.track_producer();
runtime.spawn(async move {
let _usage_producer = usage_producer;
if tokio::time::timeout(
LIVE_AUDIT_WRITE_HARD_TIMEOUT,
usage_runtime.record_terminal_event_direct(usage_data.as_ref(), event),
@@ -3,7 +3,8 @@ pub(crate) use super::super::admin::provider::pool::config::{
};
pub(crate) use super::super::admin::provider::pool::runtime::{
admin_provider_pool_key_terminal_error_reason, read_admin_provider_pool_key_cooldown_reason,
read_admin_provider_pool_runtime_state, record_admin_provider_pool_error,
read_admin_provider_pool_runtime_state, read_provider_pool_scheduling_runtime_state,
read_provider_pool_sticky_bound_key_id, record_admin_provider_pool_error,
record_admin_provider_pool_stream_timeout, record_admin_provider_pool_success,
release_admin_provider_pool_key_lease,
};
+401 -30
View File
@@ -402,9 +402,21 @@ pub(crate) enum RequestBodyNormalizationError {
InvalidBodyFraming,
AmbiguousBodyFraming,
UnsupportedContentEncoding(String),
DecodeFailed { encoding: String, reason: String },
DecompressedBodyTooLarge { encoding: String, limit_bytes: u64 },
RequestBodyTooLarge { limit_bytes: u64 },
DecodeFailed {
encoding: String,
reason: String,
},
DecompressedBodyTooLarge {
encoding: String,
limit_bytes: u64,
},
RequestBodyTooLarge {
limit_bytes: u64,
},
BodyBufferOverloaded {
requested_bytes: usize,
budget_bytes: usize,
},
}
impl RequestBodyNormalizationError {
@@ -428,6 +440,9 @@ impl RequestBodyNormalizationError {
Self::RequestBodyTooLarge { limit_bytes } => {
format!("Request body exceeds {limit_bytes} bytes")
}
Self::BodyBufferOverloaded { .. } => {
"Request body buffering capacity is temporarily exhausted".to_string()
}
}
}
@@ -440,6 +455,7 @@ impl RequestBodyNormalizationError {
Self::UnsupportedContentEncoding(_) | Self::DecodeFailed { .. } => {
http::StatusCode::BAD_REQUEST
}
Self::BodyBufferOverloaded { .. } => http::StatusCode::SERVICE_UNAVAILABLE,
}
}
}
@@ -468,6 +484,10 @@ impl fmt::Display for RequestBodyNormalizationError {
Self::RequestBodyTooLarge { limit_bytes } => {
write!(f, "request body exceeds {limit_bytes} bytes")
}
Self::BodyBufferOverloaded { requested_bytes, budget_bytes } => write!(
f,
"request body buffering needs {requested_bytes} bytes of a {budget_bytes} byte budget"
),
}
}
}
@@ -489,9 +509,24 @@ pub(crate) fn normalize_request_body_headers_and_bytes_with_limit(
headers: &mut http::HeaderMap,
body_bytes: Bytes,
limit_bytes: u64,
) -> Result<Bytes, RequestBodyNormalizationError> {
normalize_request_body_headers_and_bytes_with_budget(
headers,
body_bytes,
limit_bytes,
&mut |_| Ok(()),
)
}
pub(crate) fn normalize_request_body_headers_and_bytes_with_budget(
headers: &mut http::HeaderMap,
body_bytes: Bytes,
limit_bytes: u64,
budget: &mut impl FnMut(usize) -> Result<(), RequestBodyNormalizationError>,
) -> Result<Bytes, RequestBodyNormalizationError> {
let body_was_encoded = !request_content_encodings(headers).is_empty();
let decoded = decoded_request_body_bytes_with_limit(headers, body_bytes.as_ref(), limit_bytes)?;
let decoded =
decoded_request_body_bytes_with_budget(headers, body_bytes.as_ref(), limit_bytes, budget)?;
if !body_was_encoded {
return Ok(body_bytes);
}
@@ -533,6 +568,15 @@ pub(crate) fn decoded_request_body_bytes_with_limit<'a>(
headers: &http::HeaderMap,
body_bytes: &'a [u8],
limit: u64,
) -> Result<Cow<'a, [u8]>, RequestBodyNormalizationError> {
decoded_request_body_bytes_with_budget(headers, body_bytes, limit, &mut |_| Ok(()))
}
fn decoded_request_body_bytes_with_budget<'a>(
headers: &http::HeaderMap,
body_bytes: &'a [u8],
limit: u64,
budget: &mut impl FnMut(usize) -> Result<(), RequestBodyNormalizationError>,
) -> Result<Cow<'a, [u8]>, RequestBodyNormalizationError> {
validate_request_body_framing(headers)?;
let encodings = request_content_encodings(headers);
@@ -543,11 +587,21 @@ pub(crate) fn decoded_request_body_bytes_with_limit<'a>(
return Ok(Cow::Borrowed(body_bytes));
}
let mut decoded = body_bytes.to_vec();
let mut decoded = Cow::Borrowed(body_bytes);
for encoding in encodings.iter().rev() {
decoded = decode_single_request_body_with_limit(encoding, decoded.as_slice(), limit)?;
let retained_input_bytes = body_bytes.len().saturating_add(match &decoded {
Cow::Borrowed(_) => 0,
Cow::Owned(bytes) => bytes.capacity(),
});
decoded = Cow::Owned(decode_single_request_body_with_budget(
encoding,
decoded.as_ref(),
limit,
retained_input_bytes,
budget,
)?);
}
Ok(Cow::Owned(decoded))
Ok(decoded)
}
fn request_content_encodings(headers: &http::HeaderMap) -> Vec<String> {
@@ -631,11 +685,56 @@ fn decode_single_request_body_with_limit(
encoding: &str,
body_bytes: &[u8],
limit: u64,
) -> Result<Vec<u8>, RequestBodyNormalizationError> {
decode_single_request_body_with_budget(
encoding,
body_bytes,
limit,
body_bytes.len(),
&mut |_| Ok(()),
)
}
fn decode_single_request_body_with_budget(
encoding: &str,
body_bytes: &[u8],
limit: u64,
retained_input_bytes: usize,
budget: &mut impl FnMut(usize) -> Result<(), RequestBodyNormalizationError>,
) -> Result<Vec<u8>, RequestBodyNormalizationError> {
match encoding {
"gzip" | "x-gzip" => decode_gzip_body_with_limit(encoding, body_bytes, limit),
"deflate" => decode_deflate_body_with_limit(encoding, body_bytes, limit),
"zstd" => decode_zstd_body_with_limit(encoding, body_bytes, limit),
"gzip" | "x-gzip" => {
let mut decoder = GzDecoder::new(body_bytes);
read_request_decoder_to_end_with_budget(
encoding,
&mut decoder,
limit,
retained_input_bytes,
budget,
)
}
"deflate" => decode_deflate_body_with_budget(
encoding,
body_bytes,
limit,
retained_input_bytes,
budget,
),
"zstd" => {
let mut decoder = zstd::stream::read::Decoder::new(body_bytes).map_err(|err| {
RequestBodyNormalizationError::DecodeFailed {
encoding: encoding.to_string(),
reason: err.to_string(),
}
})?;
read_request_decoder_to_end_with_budget(
encoding,
&mut decoder,
limit,
retained_input_bytes,
budget,
)
}
_ => Err(RequestBodyNormalizationError::UnsupportedContentEncoding(
encoding.to_string(),
)),
@@ -669,20 +768,48 @@ fn decode_deflate_body_with_limit(
encoding: &str,
body_bytes: &[u8],
limit: u64,
) -> Result<Vec<u8>, RequestBodyNormalizationError> {
decode_deflate_body_with_budget(encoding, body_bytes, limit, body_bytes.len(), &mut |_| {
Ok(())
})
}
fn decode_deflate_body_with_budget(
encoding: &str,
body_bytes: &[u8],
limit: u64,
retained_input_bytes: usize,
budget: &mut impl FnMut(usize) -> Result<(), RequestBodyNormalizationError>,
) -> Result<Vec<u8>, RequestBodyNormalizationError> {
let mut zlib_decoder = ZlibDecoder::new(body_bytes);
match read_request_decoder_to_end_with_limit(encoding, &mut zlib_decoder, limit) {
match read_request_decoder_to_end_with_budget(
encoding,
&mut zlib_decoder,
limit,
retained_input_bytes,
budget,
) {
Ok(decoded) => Ok(decoded),
Err(err @ RequestBodyNormalizationError::DecompressedBodyTooLarge { .. }) => Err(err),
Err(zlib_error) => {
Err(zlib_error @ RequestBodyNormalizationError::DecodeFailed { .. }) => {
let mut raw_decoder = DeflateDecoder::new(body_bytes);
read_request_decoder_to_end_with_limit(encoding, &mut raw_decoder, limit).map_err(
|raw_error| RequestBodyNormalizationError::DecodeFailed {
encoding: encoding.to_string(),
reason: format!("{zlib_error}; raw deflate fallback failed: {raw_error}"),
},
read_request_decoder_to_end_with_budget(
encoding,
&mut raw_decoder,
limit,
retained_input_bytes,
budget,
)
.map_err(|raw_error| match raw_error {
RequestBodyNormalizationError::DecodeFailed { .. } => {
RequestBodyNormalizationError::DecodeFailed {
encoding: encoding.to_string(),
reason: format!("{zlib_error}; raw deflate fallback failed: {raw_error}"),
}
}
error => error,
})
}
Err(error) => Err(error),
}
}
@@ -719,21 +846,70 @@ fn read_request_decoder_to_end_with_limit(
decoder: &mut impl Read,
limit: u64,
) -> Result<Vec<u8>, RequestBodyNormalizationError> {
let mut limited = decoder.take(limit.saturating_add(1));
read_request_decoder_to_end_with_budget(encoding, decoder, limit, 0, &mut |_| Ok(()))
}
fn read_request_decoder_to_end_with_budget(
encoding: &str,
decoder: &mut impl Read,
limit: u64,
retained_input_bytes: usize,
budget: &mut impl FnMut(usize) -> Result<(), RequestBodyNormalizationError>,
) -> Result<Vec<u8>, RequestBodyNormalizationError> {
let capacity_limit = usize::try_from(limit).unwrap_or(usize::MAX);
let mut scratch = [0_u8; 8 * 1024];
let mut out = Vec::new();
limited
.read_to_end(&mut out)
.map_err(|err| RequestBodyNormalizationError::DecodeFailed {
encoding: encoding.to_string(),
reason: err.to_string(),
loop {
let remaining = limit.saturating_sub(out.len() as u64).saturating_add(1);
let read_limit = scratch
.len()
.min(usize::try_from(remaining).unwrap_or(usize::MAX));
let read = decoder.read(&mut scratch[..read_limit]).map_err(|err| {
RequestBodyNormalizationError::DecodeFailed {
encoding: encoding.to_string(),
reason: err.to_string(),
}
})?;
if out.len() as u64 > limit {
return Err(RequestBodyNormalizationError::DecompressedBodyTooLarge {
encoding: encoding.to_string(),
limit_bytes: limit,
});
if read == 0 {
return Ok(out);
}
let next_len = out.len().saturating_add(read);
if next_len as u64 > limit {
return Err(RequestBodyNormalizationError::DecompressedBodyTooLarge {
encoding: encoding.to_string(),
limit_bytes: limit,
});
}
if next_len > out.capacity() {
let mut capacity = out
.capacity()
.saturating_mul(2)
.max(next_len)
.min(capacity_limit);
// The encoded body and previous decoding layer remain alive during growth.
match budget(retained_input_bytes.saturating_add(capacity)) {
Ok(()) => {}
Err(RequestBodyNormalizationError::BodyBufferOverloaded { .. })
if capacity > next_len =>
{
// Rejected reservations leave the budget unchanged. Spare capacity
// must not reject a body whose actual bytes still fit.
capacity = next_len;
budget(retained_input_bytes.saturating_add(capacity))?;
}
Err(error) => return Err(error),
}
out.try_reserve_exact(capacity.saturating_sub(out.len()))
.map_err(|err| RequestBodyNormalizationError::DecodeFailed {
encoding: encoding.to_string(),
reason: err.to_string(),
})?;
if out.capacity() > capacity {
budget(retained_input_bytes.saturating_add(out.capacity()))?;
}
}
out.extend_from_slice(&scratch[..read]);
}
Ok(out)
}
pub(crate) fn header_equals(
@@ -1223,6 +1399,14 @@ mod tests {
#[test]
fn request_body_normalization_error_maps_http_status() {
assert_eq!(
RequestBodyNormalizationError::BodyBufferOverloaded {
requested_bytes: 2,
budget_bytes: 1,
}
.http_status(),
http::StatusCode::SERVICE_UNAVAILABLE
);
assert_eq!(
RequestBodyNormalizationError::RequestBodyTooLarge { limit_bytes: 1 }.http_status(),
http::StatusCode::PAYLOAD_TOO_LARGE
@@ -1250,6 +1434,193 @@ mod tests {
);
}
#[test]
fn budgeted_normalization_accounts_for_encoded_input_and_output_capacity() {
let payload = vec![b'a'; 150_000];
let mut encoder = GzEncoder::new(Vec::new(), Compression::default());
encoder.write_all(&payload).expect("gzip payload");
let encoded = encoder.finish().expect("gzip finish");
let encoded_len = encoded.len();
let mut headers = HeaderMap::new();
headers.insert(
http::header::CONTENT_ENCODING,
HeaderValue::from_static("gzip"),
);
let mut reservations = Vec::new();
let decoded = super::normalize_request_body_headers_and_bytes_with_budget(
&mut headers,
encoded.into(),
256 * 1024,
&mut |bytes| {
reservations.push(bytes);
Ok(())
},
)
.expect("budgeted gzip should decode");
assert_eq!(decoded.as_ref(), payload.as_slice());
assert!(!headers.contains_key(http::header::CONTENT_ENCODING));
assert_eq!(reservations[0], encoded_len + 8 * 1024);
assert!(reservations.windows(2).all(|pair| pair[0] <= pair[1]));
assert!(reservations.last().copied().unwrap() >= encoded_len + payload.len());
assert!(reservations.last().copied().unwrap() <= encoded_len + payload.len() * 2);
assert!(
reservations.len() <= 8,
"output growth should remain geometric"
);
}
#[test]
fn budgeted_normalization_accounts_for_retained_encoding_layers() {
let payload = b"small chained request body";
let mut inner = GzEncoder::new(Vec::new(), Compression::default());
inner.write_all(payload).expect("inner gzip payload");
let inner = inner.finish().expect("inner gzip finish");
let mut outer = GzEncoder::new(Vec::new(), Compression::default());
outer.write_all(&inner).expect("outer gzip payload");
let encoded = outer.finish().expect("outer gzip finish");
let encoded_len = encoded.len();
let mut headers = HeaderMap::new();
headers.insert(
http::header::CONTENT_ENCODING,
HeaderValue::from_static("gzip, gzip"),
);
let mut reservations = Vec::new();
let decoded = super::normalize_request_body_headers_and_bytes_with_budget(
&mut headers,
encoded.into(),
256,
&mut |bytes| {
reservations.push(bytes);
Ok(())
},
)
.expect("chained gzip should decode");
assert_eq!(decoded.as_ref(), payload);
assert_eq!(
reservations,
vec![
encoded_len + inner.len(),
encoded_len + inner.len() + payload.len(),
]
);
}
#[test]
fn budgeted_decoder_stops_before_collecting_when_capacity_is_exhausted() {
let source = vec![b'a'; 150_000];
let mut decoder = std::io::Cursor::new(source);
let error = super::read_request_decoder_to_end_with_budget(
"test",
&mut decoder,
256 * 1024,
100,
&mut |requested_bytes| {
Err(RequestBodyNormalizationError::BodyBufferOverloaded {
requested_bytes,
budget_bytes: 100,
})
},
)
.expect_err("budget rejection must stop output growth");
assert_eq!(decoder.position(), 8 * 1024);
assert_eq!(
error,
RequestBodyNormalizationError::BodyBufferOverloaded {
requested_bytes: 100 + 8 * 1024,
budget_bytes: 100,
}
);
}
#[test]
fn budgeted_decoder_accepts_actual_output_when_geometric_growth_does_not_fit() {
let source = vec![b'a'; 100_000];
let mut decoder = std::io::Cursor::new(&source);
let retained_input_bytes = 100;
let budget_bytes = retained_input_bytes + source.len();
let mut reserved_bytes = retained_input_bytes;
let mut rejected_growth = false;
let decoded = super::read_request_decoder_to_end_with_budget(
"test",
&mut decoder,
256 * 1024,
retained_input_bytes,
&mut |requested_bytes| {
if requested_bytes > budget_bytes {
rejected_growth = true;
return Err(RequestBodyNormalizationError::BodyBufferOverloaded {
requested_bytes,
budget_bytes,
});
}
reserved_bytes = reserved_bytes.max(requested_bytes);
Ok(())
},
)
.expect("actual output within the budget should finish decoding");
assert!(
rejected_growth,
"test must exercise oversized spare capacity"
);
assert_eq!(decoded, source);
assert_eq!(reserved_bytes, budget_bytes);
assert_eq!(decoded.capacity() + retained_input_bytes, budget_bytes);
}
#[test]
fn budgeted_deflate_preserves_capacity_and_size_rejections() {
let payload = [b'a'; 128];
let mut wrapped = ZlibEncoder::new(Vec::new(), Compression::default());
wrapped.write_all(&payload).expect("zlib payload");
let wrapped = wrapped.finish().expect("zlib finish");
let mut raw = DeflateEncoder::new(Vec::new(), Compression::default());
raw.write_all(&payload).expect("raw deflate payload");
let raw = raw.finish().expect("raw deflate finish");
for encoded in [wrapped, raw] {
let mut headers = HeaderMap::new();
headers.insert(
http::header::CONTENT_ENCODING,
HeaderValue::from_static("deflate"),
);
let rejection = RequestBodyNormalizationError::BodyBufferOverloaded {
requested_bytes: 300,
budget_bytes: 200,
};
let error = super::normalize_request_body_headers_and_bytes_with_budget(
&mut headers,
encoded.clone().into(),
256,
&mut |_| Err(rejection.clone()),
)
.expect_err("capacity rejection must remain an overload error");
assert_eq!(error, rejection);
assert!(headers.contains_key(http::header::CONTENT_ENCODING));
let error = super::normalize_request_body_headers_and_bytes_with_budget(
&mut headers,
encoded.into(),
64,
&mut |_| Ok(()),
)
.expect_err("size rejection must remain a size error");
assert!(matches!(
error,
RequestBodyNormalizationError::DecompressedBodyTooLarge {
limit_bytes: 64,
..
}
));
}
}
#[test]
fn check_request_content_length_allows_missing_or_within_limit() {
let empty = HeaderMap::new();
+203 -26
View File
@@ -18,6 +18,7 @@ use hyper_util::{
server::conn::auto::Builder as HyperServerBuilder,
service::TowerToHyperService,
};
use tokio_util::sync::CancellationToken;
use tower::{Service as _, ServiceExt as _};
use tracing::{debug, info, warn};
@@ -127,6 +128,7 @@ use aether_gateway::{
FrontdoorCorsConfig, FrontdoorUserRpmConfig, GatewayDataConfig, UsageRuntimeConfig,
VideoTaskTruthSourceMode,
};
use aether_gateway_frontdoor::{http_connection_limit, HttpConnectionBudget};
use aether_runtime::{
init_service_runtime, FileLoggingConfig, LogDestination, LogFormat, LogRotation,
ServiceRuntimeConfig,
@@ -1007,6 +1009,13 @@ struct GatewayUsageArgs {
)]
queue_stream_maxlen: usize,
#[arg(
long,
env = "AETHER_GATEWAY_USAGE_QUEUE_PAYLOAD_MAX_BYTES",
default_value_t = 1024 * 1024
)]
queue_payload_max_bytes: usize,
#[arg(
long,
env = "AETHER_GATEWAY_USAGE_QUEUE_BATCH_SIZE",
@@ -1220,6 +1229,7 @@ impl GatewayUsageArgs {
consumer_group: self.queue_group.trim().to_string(),
dlq_stream_key: self.queue_dlq_stream_key.trim().to_string(),
stream_maxlen: self.queue_stream_maxlen.max(1),
queue_payload_max_bytes: self.queue_payload_max_bytes,
consumer_batch_size: self.queue_batch_size.max(1),
consumer_block_ms: self.queue_block_ms.max(1),
reclaim_idle_ms: self.queue_reclaim_idle_ms.max(1),
@@ -1486,6 +1496,22 @@ struct Args {
/// Maximum number of HTTP/1 request header fields.
http_max_headers: usize,
#[arg(
long,
env = "AETHER_GATEWAY_HTTP_SHUTDOWN_TIMEOUT_MS",
default_value_t = 30_000
)]
/// Grace period for HTTP requests and upgraded connections before forced close.
http_shutdown_timeout_ms: u64,
#[arg(
long,
env = "AETHER_GATEWAY_USAGE_SHUTDOWN_TIMEOUT_MS",
default_value_t = 30_000
)]
/// Additional time for request finalizers and local usage buffers to persist.
usage_shutdown_timeout_ms: u64,
/// 容器内健康检查入口:根据当前 bind 端口探测本地 /health。
#[arg(long, hide = true, default_value_t = false)]
healthcheck: bool,
@@ -1567,6 +1593,11 @@ struct Args {
#[arg(long, env = "AETHER_GATEWAY_MAX_IN_FLIGHT_REQUESTS")]
max_in_flight_requests: Option<usize>,
/// Maximum accepted HTTP TCP connections across all listener shards, including upgrades.
/// Unset or 0 follows request plus WebSocket capacity, bounded by the FD allowance.
#[arg(long, env = "AETHER_GATEWAY_MAX_HTTP_CONNECTIONS")]
max_http_connections: Option<usize>,
/// Maximum number of long-lived public WebSocket connections. When unset,
/// this follows `max_in_flight_requests` while remaining an independent
/// gate. Set `AETHER_GATEWAY_MAX_WEBSOCKET_CONNECTIONS` to override it.
@@ -1851,10 +1882,12 @@ fn gateway_listeners(
async fn serve_gateway_router(
listeners: Vec<tokio::net::TcpListener>,
router: axum::Router,
connection_budget: Arc<HttpConnectionBudget>,
http2_max_concurrent_streams: u32,
http_header_read_timeout_ms: u64,
http_header_max_bytes: usize,
http_max_headers: usize,
shutdown: CancellationToken,
) -> Result<(), Box<dyn std::error::Error>> {
let http2_max_concurrent_streams =
gateway_http2_max_concurrent_streams(http2_max_concurrent_streams);
@@ -1865,23 +1898,41 @@ async fn serve_gateway_router(
let mut servers = tokio::task::JoinSet::new();
for listener in listeners {
let router = router.clone();
let connection_budget = Arc::clone(&connection_budget);
let shutdown = shutdown.clone();
servers.spawn(async move {
serve_gateway_listener(
listener,
router,
connection_budget,
http2_max_concurrent_streams,
http_header_read_timeout_ms,
http_header_max_bytes,
http_max_headers,
shutdown,
)
.await
});
}
if let Some(result) = servers.join_next().await {
servers.abort_all();
let serve_result = result
.map_err(|err| std::io::Error::other(format!("gateway listener task failed: {err}")))?;
serve_result?;
let mut failure = None;
while let Some(result) = servers.join_next().await {
let result = result.unwrap_or_else(|err| {
Err(std::io::Error::other(format!(
"gateway listener task failed: {err}"
)))
});
if let Err(error) = result {
failure.get_or_insert(error);
shutdown.cancel();
connection_budget.force_close();
}
}
if let Some(error) = failure {
return Err(error.into());
}
// Hyper hands upgrades to application tasks; their IO still owns this budget.
while connection_budget.snapshot().in_flight != 0 {
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
}
Ok(())
}
@@ -1889,14 +1940,26 @@ async fn serve_gateway_router(
async fn serve_gateway_listener(
listener: tokio::net::TcpListener,
router: axum::Router,
connection_budget: Arc<HttpConnectionBudget>,
http2_max_concurrent_streams: u32,
http_header_read_timeout_ms: u64,
http_header_max_bytes: usize,
http_max_headers: usize,
shutdown: CancellationToken,
) -> Result<(), std::io::Error> {
let mut make_service = router.into_make_service_with_connect_info::<std::net::SocketAddr>();
let mut connections = tokio::task::JoinSet::new();
loop {
let (io, remote_addr) = listener.accept().await?;
let (io, remote_addr) = tokio::select! {
biased;
_ = shutdown.cancelled() => break,
_ = connections.join_next(), if !connections.is_empty() => continue,
accepted = connection_budget.accept(&listener) => accepted,
};
let Ok(io) = connection_budget.try_admit(io) else {
tokio::task::yield_now().await;
continue;
};
let tower_service = make_service
.call(remote_addr)
.await
@@ -1909,7 +1972,9 @@ async fn serve_gateway_listener(
});
let io = TokioIo::new(io);
tokio::spawn(async move {
let shutdown = shutdown.clone();
let connection_budget = Arc::clone(&connection_budget);
connections.spawn(async move {
let mut builder = HyperServerBuilder::new(TokioExecutor::new());
// Hyper's HTTP/1 header timer is opt-in when using the custom
// connection builder. Configure both protocol parsers explicitly:
@@ -1938,17 +2003,34 @@ async fn serve_gateway_listener(
// the service so a peer cannot hold a socket open while dribbling
// protocol bytes or an initial header block. Once the gate opens,
// request and response bodies remain fully streaming.
let connection_result = drive_gateway_connection(
builder.serve_connection_with_upgrades(io, hyper_service),
first_request_gate,
std::time::Duration::from_millis(http_header_read_timeout_ms),
)
.await;
let connection = builder.serve_connection_with_upgrades(io, hyper_service);
tokio::pin!(connection);
let draining_connection = async {
tokio::select! {
result = &mut connection => result,
_ = shutdown.cancelled() => {
connection.as_mut().graceful_shutdown();
connection.await
}
}
};
let connection_result = tokio::select! {
biased;
_ = connection_budget.wait_for_forced_close() => Ok(()),
result = drive_gateway_connection(
draining_connection,
first_request_gate,
std::time::Duration::from_millis(http_header_read_timeout_ms),
) => result,
};
if let Err(err) = connection_result {
tracing::trace!(error = ?err, "gateway connection closed with error");
}
});
}
drop(listener);
while connections.join_next().await.is_some() {}
Ok(())
}
fn resolve_local_http_base_url(app_port: u16) -> Result<String, std::io::Error> {
@@ -2065,11 +2147,14 @@ fn validate_deployment_topology(
}
fn main() -> Result<(), Box<dyn std::error::Error>> {
tokio::runtime::Builder::new_multi_thread()
let _log_shutdown = aether_runtime::LogShutdownGuard::new();
let runtime = tokio::runtime::Builder::new_multi_thread()
.enable_all()
.thread_stack_size(GATEWAY_TOKIO_WORKER_STACK_SIZE_BYTES)
.build()?
.block_on(run())
.build()?;
let result = runtime.block_on(run());
aether_usage_runtime::shutdown_usage_background_runtime(std::time::Duration::from_secs(5));
result
}
async fn run() -> Result<(), Box<dyn std::error::Error>> {
@@ -2133,6 +2218,13 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
.max_websocket_connections
.filter(|limit| *limit > 0)
.unwrap_or(request_concurrency_limit);
let http_connection_limit = http_connection_limit(
args.max_http_connections,
request_concurrency_limit,
websocket_connection_limit,
soft_fd_limit(),
);
let http_connection_budget = Arc::new(HttpConnectionBudget::new(http_connection_limit));
let distributed_websocket_connection_limit = match args.distributed_websocket_connection_limit {
Some(limit) if limit > 0 => Some(limit),
Some(_) => None,
@@ -2323,7 +2415,8 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
}
state = state
.with_request_concurrency_limit(request_concurrency_limit)
.with_websocket_connection_limit(websocket_connection_limit);
.with_websocket_connection_limit(websocket_connection_limit)
.with_http_connection_budget(Arc::clone(&http_connection_budget));
if let Some(limit) = args.distributed_request_limit.filter(|limit| *limit > 0) {
let distributed_gate = state
.runtime_state()
@@ -2467,6 +2560,7 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
let listeners = gateway_listeners(bind_addr, listen_backlog, listener_shards)?;
let public_base_url = resolve_local_http_base_url(app_port)?;
let frontdoor_health_url = format!("{public_base_url}/_gateway/health");
let shutdown_state = state.clone();
let api_router = build_router_with_state(state);
// Compose the final router: API routes + optional static file serving.
@@ -2486,6 +2580,7 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
app_port,
listen_backlog,
listener_shards,
max_http_connections = http_connection_limit,
http2_max_concurrent_streams = gateway_http2_max_concurrent_streams(args.http2_max_concurrent_streams),
public_url = %public_base_url,
healthcheck_url = %frontdoor_health_url,
@@ -2493,18 +2588,61 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
"aether-gateway ready"
);
serve_gateway_router(
listeners,
router,
args.http2_max_concurrent_streams,
args.http_header_read_timeout_ms,
args.http_header_max_bytes,
args.http_max_headers,
)
.await?;
let shutdown = CancellationToken::new();
let serve_result = {
let server = serve_gateway_router(
listeners,
router,
Arc::clone(&http_connection_budget),
args.http2_max_concurrent_streams,
args.http_header_read_timeout_ms,
args.http_header_max_bytes,
args.http_max_headers,
shutdown.clone(),
);
tokio::pin!(server);
tokio::select! {
result = &mut server => result,
signal = aether_runtime::wait_for_shutdown_signal() => {
signal?;
info!("shutdown signal received, draining gateway requests");
shutdown.cancel();
match tokio::time::timeout(
std::time::Duration::from_millis(args.http_shutdown_timeout_ms),
&mut server,
).await {
Ok(result) => result,
Err(_) => {
warn!(
event_name = "gateway_http_shutdown_deadline",
connections = http_connection_budget.snapshot().in_flight,
"HTTP drain deadline reached; closing remaining sockets"
);
http_connection_budget.force_close();
match tokio::time::timeout(std::time::Duration::from_secs(5), &mut server).await {
Ok(result) => result,
Err(_) => Err(std::io::Error::new(std::io::ErrorKind::TimedOut,
"gateway connection tasks did not stop after forced close").into()),
}
}
}
}
}
};
let usage_result = shutdown_state
.shutdown_usage_runtime(std::time::Duration::from_millis(
args.usage_shutdown_timeout_ms,
))
.await;
if let Some(background_tasks) = background_tasks {
background_tasks.shutdown().await;
}
serve_result?;
usage_result?;
info!(
event_name = "gateway_shutdown_complete",
"gateway local persistence drained"
);
Ok(())
}
@@ -3437,6 +3575,10 @@ fn pending_backfills_error(
#[cfg(test)]
mod tests {
mod shutdown {
include!("shutdown_tests.rs");
}
use super::{
automatic_gateway_request_concurrency_for_capacity,
automatic_gateway_request_concurrency_for_parallelism, automatic_sql_pool_config,
@@ -3478,6 +3620,8 @@ mod tests {
http_header_read_timeout_ms: DEFAULT_GATEWAY_HTTP_HEADER_READ_TIMEOUT_MS,
http_header_max_bytes: DEFAULT_GATEWAY_HTTP_HEADER_MAX_BYTES,
http_max_headers: DEFAULT_GATEWAY_HTTP_MAX_HEADERS,
http_shutdown_timeout_ms: 30_000,
usage_shutdown_timeout_ms: 30_000,
healthcheck: false,
healthcheck_timeout_ms: 3_000,
deployment_topology: DeploymentTopologyArg::SingleNode,
@@ -3492,6 +3636,7 @@ mod tests {
video_task_poller_batch_size: 32,
video_task_store_path: None,
max_in_flight_requests: None,
max_http_connections: None,
max_websocket_connections: None,
distributed_request_limit: None,
distributed_websocket_connection_limit: None,
@@ -3532,6 +3677,7 @@ mod tests {
queue_group: "usage_consumers".to_string(),
queue_dlq_stream_key: "usage:events:dlq".to_string(),
queue_stream_maxlen: 200_000,
queue_payload_max_bytes: 1024 * 1024,
queue_batch_size: 128,
queue_block_ms: 500,
queue_reclaim_idle_ms: 60_000,
@@ -3925,6 +4071,37 @@ mod tests {
}
}
#[test]
fn gateway_usage_queue_payload_limit_preserves_cli_override_and_rejects_zero() {
let command = <Args as clap::CommandFactory>::command();
let argument = command
.get_arguments()
.find(|argument| argument.get_id() == "queue_payload_max_bytes")
.expect("usage payload argument must be registered");
assert_eq!(
argument.get_env(),
Some(std::ffi::OsStr::new(
"AETHER_GATEWAY_USAGE_QUEUE_PAYLOAD_MAX_BYTES"
))
);
assert_eq!(argument.get_default_values()[0].to_str(), Some("1048576"));
let args = Args::try_parse_from(["aether-gateway", "--queue-payload-max-bytes", "32768"])
.expect("explicit usage payload limit should parse");
let config = args.usage.to_config(4, 8, Some(4));
assert_eq!(config.queue_payload_max_bytes, 32_768);
assert!(config.validate().is_ok());
let mut args = test_args();
assert_eq!(
args.usage.to_config(4, 8, Some(4)).queue_payload_max_bytes,
1024 * 1024
);
args.usage.queue_payload_max_bytes = 0;
let config = args.usage.to_config(4, 8, Some(4));
assert_eq!(config.queue_payload_max_bytes, 0);
assert!(config.validate().is_err());
}
#[test]
fn gateway_usage_queue_workers_manual_override_wins_and_is_capped() {
let mut args = test_args();
+7 -7
View File
@@ -26,11 +26,11 @@ pub(crate) use runtime::{
start_manual_usage_cleanup_task, start_proxy_upgrade_rollout, AccountSelfCheckRunSummary,
AdminCleanupRunRecord, AdminCleanupTaskKind, AdminStatsRebuildSummary,
AdminSystemCleanupSummary, ManualUsageCleanupError, ManualUsageCleanupMode,
ManualUsageCleanupOptions, OAuthTokenRefreshRunSummary, PoolQuotaProbeRunSummary,
PoolQuotaProbeWorkerConfig, ProviderCheckinRunSummary, ProviderQuotaAlertRunSummary,
ProxyUpgradeRolloutCancelSummary, ProxyUpgradeRolloutConflictClearSummary,
ProxyUpgradeRolloutNodeActionSummary, ProxyUpgradeRolloutProbeConfig,
ProxyUpgradeRolloutSkippedRestoreSummary, ProxyUpgradeRolloutStatus,
ProxyUpgradeRolloutTrackedNodeState, UsageCounterFlushRuntimeMetrics,
UsageCounterFlushWorkerConfig,
ManualUsageCleanupOptions, OAuthTokenRefreshRunSummary, PoolQuotaProbeReplenishCoordinator,
PoolQuotaProbeRunSummary, PoolQuotaProbeWorkerConfig, ProviderCheckinRunSummary,
ProviderQuotaAlertRunSummary, ProxyUpgradeRolloutCancelSummary,
ProxyUpgradeRolloutConflictClearSummary, ProxyUpgradeRolloutNodeActionSummary,
ProxyUpgradeRolloutProbeConfig, ProxyUpgradeRolloutSkippedRestoreSummary,
ProxyUpgradeRolloutStatus, ProxyUpgradeRolloutTrackedNodeState,
UsageCounterFlushRuntimeMetrics, UsageCounterFlushWorkerConfig,
};
@@ -84,7 +84,8 @@ pub(crate) use pool_quota_probe::{
perform_pool_quota_probe_once, perform_pool_quota_probe_once_for_provider_with_config,
perform_pool_quota_probe_once_with_config, pool_quota_probe_target_count,
select_pool_quota_probe_key_ids, spawn_pool_quota_probe_replenish_for_request,
spawn_pool_quota_probe_worker, PoolQuotaProbeRunSummary, PoolQuotaProbeWorkerConfig,
spawn_pool_quota_probe_worker, PoolQuotaProbeReplenishCoordinator, PoolQuotaProbeRunSummary,
PoolQuotaProbeWorkerConfig,
};
pub(crate) use pool_score_rebuild::{
ensure_provider_key_pool_scores_for_keys, perform_pool_score_rebuild_once,
@@ -1,4 +1,6 @@
use std::collections::{BTreeMap, BTreeSet};
use std::collections::{BTreeMap, BTreeSet, HashMap};
use std::future::Future;
use std::sync::{Arc, Mutex};
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use aether_data_contracts::repository::pool_scores::{
@@ -47,6 +49,166 @@ const POOL_QUOTA_PROBE_BURST_RETRY_GUARD_SECONDS: u64 = 15;
const POOL_QUOTA_PROBE_AUTO_MIN_INTERVAL_SECONDS: u64 = 30;
const POOL_QUOTA_PROBE_AUTO_MAX_INTERVAL_SECONDS: u64 = 10 * 60;
const POOL_QUOTA_PROBE_AUTO_MAX_PRESSURE: u64 = 64;
const POOL_QUOTA_PROBE_LOCAL_MAX_PROVIDERS: usize = 1024;
#[derive(Debug)]
pub(crate) struct PoolQuotaProbeReplenishCoordinator {
capacity: usize,
state: Mutex<PoolQuotaProbeReplenishState>,
}
#[derive(Debug, Default)]
struct PoolQuotaProbeReplenishState {
providers: HashMap<String, PoolQuotaProbeReplenishEntry>,
started_total: u64,
coalesced_total: u64,
capacity_rejected_total: u64,
}
#[derive(Debug)]
struct PoolQuotaProbeReplenishEntry {
identity: Arc<()>,
pending: bool,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) struct PoolQuotaProbeReplenishSnapshot {
pub(crate) capacity: usize,
pub(crate) active: usize,
pub(crate) started_total: u64,
pub(crate) coalesced_total: u64,
pub(crate) capacity_rejected_total: u64,
}
impl Default for PoolQuotaProbeReplenishCoordinator {
fn default() -> Self {
Self::new(POOL_QUOTA_PROBE_LOCAL_MAX_PROVIDERS)
}
}
impl PoolQuotaProbeReplenishCoordinator {
fn new(capacity: usize) -> Self {
Self {
capacity,
state: Mutex::new(PoolQuotaProbeReplenishState::default()),
}
}
pub(crate) fn snapshot(&self) -> PoolQuotaProbeReplenishSnapshot {
let state = self
.state
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
PoolQuotaProbeReplenishSnapshot {
capacity: self.capacity,
active: state.providers.len(),
started_total: state.started_total,
coalesced_total: state.coalesced_total,
capacity_rejected_total: state.capacity_rejected_total,
}
}
fn request(self: &Arc<Self>, provider_id: String) -> Option<PoolQuotaProbeReplenishGuard> {
let mut state = self
.state
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
if let Some(entry) = state.providers.get_mut(&provider_id) {
entry.pending = true;
state.coalesced_total = state.coalesced_total.saturating_add(1);
return None;
}
if state.providers.len() >= self.capacity {
// Replenishment is best effort; the periodic base scan remains available.
state.capacity_rejected_total = state.capacity_rejected_total.saturating_add(1);
return None;
}
let identity = Arc::new(());
state.providers.insert(
provider_id.clone(),
PoolQuotaProbeReplenishEntry {
identity: Arc::clone(&identity),
pending: true,
},
);
state.started_total = state.started_total.saturating_add(1);
Some(PoolQuotaProbeReplenishGuard {
coordinator: Arc::clone(self),
provider_id,
identity,
finished: false,
})
}
fn spawn<F, Fut>(
self: &Arc<Self>,
provider_id: String,
mut replenish: F,
) -> Option<tokio::task::JoinHandle<()>>
where
F: FnMut() -> Fut + Send + 'static,
Fut: Future<Output = ()> + Send + 'static,
{
// Own the guard before spawn so cancellation before the first poll also cleans up.
let mut guard = self.request(provider_id)?;
Some(tokio::spawn(async move {
while guard.next_pass() {
replenish().await;
}
}))
}
}
struct PoolQuotaProbeReplenishGuard {
coordinator: Arc<PoolQuotaProbeReplenishCoordinator>,
provider_id: String,
identity: Arc<()>,
finished: bool,
}
impl PoolQuotaProbeReplenishGuard {
fn next_pass(&mut self) -> bool {
let mut state = self
.coordinator
.state
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
if let Some(entry) = state
.providers
.get_mut(&self.provider_id)
.filter(|entry| Arc::ptr_eq(&entry.identity, &self.identity))
{
if entry.pending {
entry.pending = false;
return true;
}
// Check for a follow-up and release ownership in one critical section.
state.providers.remove(&self.provider_id);
}
self.finished = true;
false
}
}
impl Drop for PoolQuotaProbeReplenishGuard {
fn drop(&mut self) {
if self.finished {
return;
}
let mut state = self
.coordinator
.state
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
if state
.providers
.get(&self.provider_id)
.is_some_and(|entry| Arc::ptr_eq(&entry.identity, &self.identity))
{
state.providers.remove(&self.provider_id);
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum PoolQuotaProbeMode {
@@ -1672,38 +1834,62 @@ pub(crate) fn spawn_pool_quota_probe_replenish_for_request(
return None;
}
Some(tokio::spawn(async move {
let runtime = state.runtime_state.clone();
mark_probe_burst_pending(runtime.as_ref(), &provider_id).await;
let lease =
acquire_pool_quota_probe_burst_trigger_lock(runtime.as_ref(), &provider_id).await;
if lease.is_none() {
return;
}
let coordinator = Arc::clone(&state.pool_quota_probe_replenish);
coordinator.spawn(provider_id.clone(), move || {
run_pool_quota_probe_replenish(state.clone(), provider_id.clone())
})
}
let config = PoolQuotaProbeWorkerConfig::from_env();
loop {
let pending = runtime
.kv_take(&probe_burst_pending_key(&provider_id))
.await
.ok()
.flatten()
.is_some();
if !pending {
break;
}
match perform_pool_quota_probe_once_for_provider_with_mode(
async fn run_pool_quota_probe_replenish(state: AppState, provider_id: String) {
let runtime = state.runtime_state.clone();
let runtime = runtime.as_ref();
let config = PoolQuotaProbeWorkerConfig::from_env();
run_pool_quota_probe_replenish_with(
runtime,
&provider_id,
|| {
perform_pool_quota_probe_once_for_provider_with_mode(
&state,
&provider_id,
config,
PoolQuotaProbeMode::Burst,
)
.await
{
},
|lease| release_pool_quota_probe_burst_trigger_lock(runtime, Some(lease)),
)
.await;
}
async fn run_pool_quota_probe_replenish_with<Probe, ProbeFuture, Release, ReleaseFuture>(
runtime: &RuntimeState,
provider_id: &str,
mut probe: Probe,
mut release: Release,
) where
Probe: FnMut() -> ProbeFuture,
ProbeFuture: Future<Output = Result<PoolQuotaProbeRunSummary, GatewayError>>,
Release: FnMut(RuntimeLockLease) -> ReleaseFuture,
ReleaseFuture: Future<Output = ()>,
{
mark_probe_burst_pending(runtime, provider_id).await;
loop {
let Some(lease) = acquire_pool_quota_probe_burst_trigger_lock(runtime, provider_id).await
else {
return;
};
let recheck_after_release = loop {
let pending = match runtime.kv_take(&probe_burst_pending_key(provider_id)).await {
Ok(pending) => pending.is_some(),
Err(_) => break false,
};
if !pending {
break true;
}
match probe().await {
Ok(summary) => {
if summary.providers_busy > 0 {
mark_probe_burst_pending(runtime.as_ref(), &provider_id).await;
mark_probe_burst_pending(runtime, provider_id).await;
tokio::time::sleep(Duration::from_millis(250)).await;
continue;
}
@@ -1717,17 +1903,30 @@ pub(crate) fn spawn_pool_quota_probe_replenish_for_request(
}
}
let still_pending = runtime
.kv_exists(&probe_burst_pending_key(&provider_id))
match runtime
.kv_exists(&probe_burst_pending_key(provider_id))
.await
.unwrap_or(false);
if !still_pending {
break;
{
Ok(true) => {}
Ok(false) => break true,
Err(_) => break false,
}
}
};
release_pool_quota_probe_burst_trigger_lock(runtime.as_ref(), lease).await;
}))
release(lease).await;
// A different instance may publish after the final pending check and fail
// to acquire our old lease. Recheck after release, then acquire a fresh token
// before consuming that signal. Read failures terminate instead of spinning.
if !recheck_after_release
|| !runtime
.kv_exists(&probe_burst_pending_key(provider_id))
.await
.unwrap_or(false)
{
return;
}
tokio::task::yield_now().await;
}
}
pub(crate) fn spawn_pool_quota_probe_worker(
@@ -1764,7 +1963,430 @@ pub(crate) fn spawn_pool_quota_probe_worker(
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicUsize, Ordering};
use aether_runtime_state::MemoryRuntimeStateConfig;
use serde_json::json;
use tokio::sync::Notify;
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn pool_quota_probe_local_coalesces_before_spawn_and_keeps_one_follow_up() {
let coordinator = Arc::new(PoolQuotaProbeReplenishCoordinator::new(8));
let runtime = Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default()));
let probe_calls = Arc::new(AtomicUsize::new(0));
let release_calls = Arc::new(AtomicUsize::new(0));
let started = Arc::new(Notify::new());
let finish_first = Arc::new(Notify::new());
let leader = coordinator
.spawn("provider".to_string(), {
let runtime = Arc::clone(&runtime);
let probe_calls = Arc::clone(&probe_calls);
let release_calls = Arc::clone(&release_calls);
let started = Arc::clone(&started);
let finish_first = Arc::clone(&finish_first);
move || {
let runtime = Arc::clone(&runtime);
let probe_calls = Arc::clone(&probe_calls);
let release_calls = Arc::clone(&release_calls);
let started = Arc::clone(&started);
let finish_first = Arc::clone(&finish_first);
async move {
run_pool_quota_probe_replenish_with(
runtime.as_ref(),
"provider",
|| async {
if probe_calls.fetch_add(1, Ordering::AcqRel) == 0 {
started.notify_one();
finish_first.notified().await;
}
Ok(PoolQuotaProbeRunSummary::empty())
},
|lease| {
release_calls.fetch_add(1, Ordering::AcqRel);
release_pool_quota_probe_burst_trigger_lock(
runtime.as_ref(),
Some(lease),
)
},
)
.await;
}
}
})
.expect("one leader");
tokio::time::timeout(Duration::from_secs(2), started.notified())
.await
.expect("first probe starts");
let barrier = Arc::new(tokio::sync::Barrier::new(65));
let mut triggers = tokio::task::JoinSet::new();
for _ in 0..64 {
let coordinator = Arc::clone(&coordinator);
let barrier = Arc::clone(&barrier);
triggers.spawn(async move {
barrier.wait().await;
assert!(coordinator
.spawn("provider".to_string(), || async {
panic!("a duplicate trigger must not spawn work")
})
.is_none());
});
}
barrier.wait().await;
while let Some(result) = triggers.join_next().await {
result.expect("concurrent trigger");
}
assert_eq!(coordinator.snapshot().started_total, 1);
assert_eq!(coordinator.snapshot().coalesced_total, 64);
assert_eq!(probe_calls.load(Ordering::Acquire), 1);
assert!(
!runtime
.kv_exists(&probe_burst_pending_key("provider"))
.await
.expect("pending read"),
"local duplicates must not each write Redis pending while the leader is running"
);
finish_first.notify_one();
tokio::time::timeout(Duration::from_secs(2), leader)
.await
.expect("leader finishes")
.expect("leader task");
assert_eq!(probe_calls.load(Ordering::Acquire), 2);
assert_eq!(
release_calls.load(Ordering::Acquire),
2,
"64 retriggers produce one additional Redis lock/drain cycle"
);
assert_eq!(coordinator.snapshot().active, 0);
}
#[tokio::test]
async fn pool_quota_probe_local_different_providers_run_independently() {
let coordinator = Arc::new(PoolQuotaProbeReplenishCoordinator::new(2));
let started = Arc::new(tokio::sync::Semaphore::new(0));
let finish = Arc::new(tokio::sync::Semaphore::new(0));
let mut tasks = Vec::new();
for provider_id in ["provider-a", "provider-b"] {
tasks.push(
coordinator
.spawn(provider_id.to_string(), {
let started = Arc::clone(&started);
let finish = Arc::clone(&finish);
move || {
let started = Arc::clone(&started);
let finish = Arc::clone(&finish);
async move {
started.add_permits(1);
finish.acquire().await.expect("finish signal").forget();
}
}
})
.expect("independent provider leader"),
);
}
tokio::time::timeout(Duration::from_secs(2), started.acquire_many(2))
.await
.expect("both providers start")
.expect("started permits")
.forget();
assert_eq!(coordinator.snapshot().active, 2);
finish.add_permits(2);
for task in tasks {
task.await.expect("provider finishes");
}
assert_eq!(coordinator.snapshot().active, 0);
}
#[test]
fn pool_quota_probe_local_exit_handoff_keeps_exactly_one_owner() {
let coordinator = Arc::new(PoolQuotaProbeReplenishCoordinator::new(1));
for _ in 0..100 {
let mut first = coordinator
.request("provider".to_string())
.expect("first owner");
assert!(first.next_pass());
let barrier = Arc::new(std::sync::Barrier::new(2));
let (mut first, continues, replacement) = std::thread::scope(|scope| {
let first_barrier = Arc::clone(&barrier);
let exit = scope.spawn(move || {
first_barrier.wait();
let continues = first.next_pass();
(first, continues)
});
let trigger = scope.spawn(|| {
barrier.wait();
coordinator.request("provider".to_string())
});
let (first, continues) = exit.join().expect("exit thread");
(first, continues, trigger.join().expect("trigger thread"))
});
assert_ne!(
continues,
replacement.is_some(),
"the signal is consumed by exactly one owner"
);
assert_eq!(coordinator.snapshot().active, 1);
if continues {
assert!(!first.next_pass());
}
drop(first);
if replacement.is_some() {
assert_eq!(
coordinator.snapshot().active,
1,
"old cleanup must not delete the replacement"
);
}
drop(replacement);
assert_eq!(coordinator.snapshot().active, 0);
}
}
#[tokio::test]
async fn pool_quota_probe_local_abort_and_panic_release_admission() {
let coordinator = Arc::new(PoolQuotaProbeReplenishCoordinator::new(1));
let unpolled = coordinator
.spawn("provider".to_string(), || async {
std::future::pending::<()>().await;
})
.expect("unpolled owner");
unpolled.abort();
assert!(unpolled.await.expect_err("aborted").is_cancelled());
assert_eq!(coordinator.snapshot().active, 0);
let started = Arc::new(Notify::new());
let running = coordinator
.spawn("provider".to_string(), {
let started = Arc::clone(&started);
move || {
let started = Arc::clone(&started);
async move {
started.notify_one();
std::future::pending::<()>().await;
}
}
})
.expect("running owner");
started.notified().await;
assert!(coordinator.request("provider".to_string()).is_none());
running.abort();
assert!(running.await.expect_err("aborted").is_cancelled());
assert_eq!(coordinator.snapshot().active, 0);
let panicked = coordinator
.spawn("provider".to_string(), || async {
panic!("probe panicked")
})
.expect("panic owner");
assert!(panicked.await.expect_err("probe panic").is_panic());
assert_eq!(coordinator.snapshot().active, 0);
coordinator
.spawn("provider".to_string(), || std::future::ready(()))
.expect("later trigger can run")
.await
.expect("recovered probe");
assert_eq!(coordinator.snapshot().active, 0);
}
#[test]
fn pool_quota_probe_local_capacity_is_bounded_and_completed_keys_are_removed() {
let coordinator = Arc::new(PoolQuotaProbeReplenishCoordinator::new(2));
let first = coordinator.request("a".to_string()).expect("first");
let second = coordinator.request("b".to_string()).expect("second");
assert!(coordinator.request("c".to_string()).is_none());
assert!(coordinator.request("a".to_string()).is_none());
assert_eq!(coordinator.snapshot().capacity_rejected_total, 1);
assert_eq!(coordinator.snapshot().coalesced_total, 1);
drop(first);
let replacement = coordinator
.request("c".to_string())
.expect("freed capacity");
drop((second, replacement));
for index in 0..1000 {
drop(
coordinator
.request(format!("provider-{index}"))
.expect("new provider"),
);
assert_eq!(coordinator.snapshot().active, 0);
}
assert_eq!(coordinator.snapshot().started_total, 1003);
}
#[tokio::test]
async fn pool_quota_probe_local_state_clones_share_only_the_same_runtime_binding() {
let state = AppState::new().expect("state");
let cloned = state.clone();
assert!(Arc::ptr_eq(
&state.pool_quota_probe_replenish,
&cloned.pool_quota_probe_replenish
));
let rebound_same = cloned.with_runtime_state(Arc::clone(&state.runtime_state));
assert!(Arc::ptr_eq(
&state.pool_quota_probe_replenish,
&rebound_same.pool_quota_probe_replenish
));
let rebound = state
.clone()
.with_runtime_state(Arc::new(RuntimeState::memory(
MemoryRuntimeStateConfig::default(),
)));
assert!(!Arc::ptr_eq(
&state.pool_quota_probe_replenish,
&rebound.pool_quota_probe_replenish
));
let first = state
.pool_quota_probe_replenish
.request("provider".to_string())
.expect("first runtime");
assert!(rebound_same
.pool_quota_probe_replenish
.request("provider".to_string())
.is_none());
let second = rebound
.pool_quota_probe_replenish
.request("provider".to_string())
.expect("other runtime is independent");
drop((first, second));
}
#[tokio::test]
async fn pool_quota_probe_replenish_rechecks_remote_pending_after_unlock() {
let runtime = Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default()));
let first = Arc::new(PoolQuotaProbeReplenishCoordinator::new(1));
let second = Arc::new(PoolQuotaProbeReplenishCoordinator::new(1));
let before_unlock = Arc::new(Notify::new());
let finish_unlock = Arc::new(Notify::new());
let first_probes = Arc::new(AtomicUsize::new(0));
let first_releases = Arc::new(AtomicUsize::new(0));
let first_task = first
.spawn("provider".to_string(), {
let runtime = Arc::clone(&runtime);
let before_unlock = Arc::clone(&before_unlock);
let finish_unlock = Arc::clone(&finish_unlock);
let first_probes = Arc::clone(&first_probes);
let first_releases = Arc::clone(&first_releases);
move || {
let runtime = Arc::clone(&runtime);
let before_unlock = Arc::clone(&before_unlock);
let finish_unlock = Arc::clone(&finish_unlock);
let first_probes = Arc::clone(&first_probes);
let first_releases = Arc::clone(&first_releases);
async move {
run_pool_quota_probe_replenish_with(
runtime.as_ref(),
"provider",
|| {
first_probes.fetch_add(1, Ordering::AcqRel);
std::future::ready(Ok(PoolQuotaProbeRunSummary::empty()))
},
|lease| {
let runtime = Arc::clone(&runtime);
let before_unlock = Arc::clone(&before_unlock);
let finish_unlock = Arc::clone(&finish_unlock);
let first_release =
first_releases.fetch_add(1, Ordering::AcqRel) == 0;
async move {
if first_release {
before_unlock.notify_one();
finish_unlock.notified().await;
}
release_pool_quota_probe_burst_trigger_lock(
runtime.as_ref(),
Some(lease),
)
.await;
}
},
)
.await;
}
}
})
.expect("first instance leader");
tokio::time::timeout(Duration::from_secs(2), before_unlock.notified())
.await
.expect("first instance drained but still owns Redis lease");
assert_eq!(first_probes.load(Ordering::Acquire), 1);
assert!(!runtime
.kv_exists(&probe_burst_pending_key("provider"))
.await
.expect("drained pending"));
let second_task = second.spawn("provider".to_string(), {
let runtime = Arc::clone(&runtime);
move || {
let runtime = Arc::clone(&runtime);
async move {
run_pool_quota_probe_replenish_with(
runtime.as_ref(), "provider",
|| async { panic!("second instance must not consume pending without the Redis lease") },
|lease| release_pool_quota_probe_burst_trigger_lock(runtime.as_ref(), Some(lease)),
).await;
}
}
}).expect("independent second instance leader");
second_task
.await
.expect("second instance leaves pending for the current owner");
assert_eq!(second.snapshot().active, 0);
assert!(runtime
.kv_exists(&probe_burst_pending_key("provider"))
.await
.expect("new pending signal"));
finish_unlock.notify_one();
tokio::time::timeout(Duration::from_secs(2), first_task)
.await
.expect("handoff drains")
.expect("first instance finishes");
assert_eq!(
first_probes.load(Ordering::Acquire),
2,
"the cross-instance exit signal must trigger a second probe"
);
assert_eq!(first_releases.load(Ordering::Acquire), 2);
assert_eq!(first.snapshot().active, 0);
assert!(!runtime
.kv_exists(&probe_burst_pending_key("provider"))
.await
.expect("all pending consumed"));
let lease = acquire_pool_quota_probe_burst_trigger_lock(runtime.as_ref(), "provider").await;
assert!(lease.is_some(), "the replacement Redis lease is released");
release_pool_quota_probe_burst_trigger_lock(runtime.as_ref(), lease).await;
}
#[tokio::test]
async fn pool_quota_probe_replenish_public_spawn_returns_none_for_merged_triggers() {
let repository = Arc::new(
aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository::seed(
Vec::new(),
Vec::new(),
Vec::new(),
),
);
let state = AppState::new().expect("state").with_data_state_for_tests(
crate::data::GatewayDataState::with_provider_catalog_repository_for_tests(repository),
);
let leader =
spawn_pool_quota_probe_replenish_for_request(state.clone(), "provider".to_string())
.expect("leader handle");
for _ in 0..64 {
assert!(spawn_pool_quota_probe_replenish_for_request(
state.clone(),
"provider".to_string()
)
.is_none());
}
leader.await.expect("leader completes");
assert_eq!(state.pool_quota_probe_replenish.snapshot().started_total, 1);
assert_eq!(
state.pool_quota_probe_replenish.snapshot().coalesced_total,
64
);
assert_eq!(state.pool_quota_probe_replenish.snapshot().active, 0);
spawn_pool_quota_probe_replenish_for_request(state.clone(), "provider".to_string())
.expect("later leader handle")
.await
.expect("later leader finishes");
}
#[test]
fn worker_error_score_reason_drops_runtime_error_details() {
@@ -973,7 +973,7 @@ async fn record_adaptive_rate_limit_effect(
let _effect_guard = effect_lock.lock().await;
let observed_at_unix_secs = current_unix_secs();
let current_rpm = state
.read_recent_request_candidates(ADAPTIVE_RPM_RECENT_CANDIDATE_LIMIT)
.read_recent_runtime_request_candidates(ADAPTIVE_RPM_RECENT_CANDIDATE_LIMIT)
.await
.ok()
.map(|recent_candidates| {
@@ -1123,7 +1123,7 @@ async fn record_adaptive_success_effect(
return;
}
let Some(recent_candidates) = state
.read_recent_request_candidates(ADAPTIVE_RPM_RECENT_CANDIDATE_LIMIT)
.read_recent_runtime_request_candidates(ADAPTIVE_RPM_RECENT_CANDIDATE_LIMIT)
.await
.ok()
else {
@@ -643,6 +643,24 @@ impl RequestCandidateQueueRuntime {
}
}
pub(crate) fn pending_writes(&self) -> usize {
[
self.metrics.pending_current.load(Ordering::Acquire),
self.metrics
.priority_pending_current
.load(Ordering::Acquire),
self.metrics.active_pending_current.load(Ordering::Acquire),
self.metrics
.terminal_pending_current
.load(Ordering::Acquire),
self.metrics
.terminal_barrier_pending
.load(Ordering::Acquire),
]
.into_iter()
.fold(0_usize, usize::saturating_add)
}
pub(crate) fn metric_samples(&self) -> Vec<MetricSample> {
vec![
MetricSample::new(
+115 -2
View File
@@ -5,6 +5,7 @@ use std::sync::Arc;
use std::task::{Context, Poll};
use aether_routing_core::RoutingExecutionPolicy;
use aether_usage_runtime::{UsageProducerGuard, UsageRuntime};
use axum::body::{Body, Bytes, HttpBody};
use http::Response;
use http_body::{Frame, SizeHint};
@@ -29,24 +30,49 @@ pub(crate) fn cancel_on_client_disconnect() -> bool {
.unwrap_or(false)
}
#[cfg(test)]
pub(crate) async fn run_request<F>(future: F) -> Result<Response<Body>, GatewayError>
where
F: Future<Output = Result<Response<Body>, GatewayError>> + Send + 'static,
{
run_tracked_request(future, None).await
}
pub(crate) async fn run_request_with_usage<F>(
usage: Arc<UsageRuntime>,
future: F,
) -> Result<Response<Body>, GatewayError>
where
F: Future<Output = Result<Response<Body>, GatewayError>> + Send + 'static,
{
run_tracked_request(future, Some(Arc::new(usage.track_producer()))).await
}
async fn run_tracked_request<F>(
future: F,
producer: Option<Arc<UsageProducerGuard>>,
) -> Result<Response<Body>, GatewayError>
where
F: Future<Output = Result<Response<Body>, GatewayError>> + Send + 'static,
{
let cancel = Arc::new(AtomicBool::new(true));
let diagnostics = Arc::new(RequestDiagnostics::default());
let cancel_for_response = Arc::clone(&cancel);
let producer_for_request = producer.clone();
let future = CANCEL_ON_CLIENT_DISCONNECT.scope(
Arc::clone(&cancel),
scope_request_diagnostics_with(Some(Arc::clone(&diagnostics)), async move {
let response = future.await?;
if cancel_for_response.load(Ordering::Acquire) {
let complete_on_disconnect = !cancel_for_response.load(Ordering::Acquire);
if !complete_on_disconnect && producer.is_none() {
return Ok(response);
}
Ok(response.map(|body| {
Body::new(CompleteOnDisconnectBody {
body: Some(body),
diagnostics,
complete_on_disconnect,
producer,
})
}))
}),
@@ -54,6 +80,7 @@ where
CompleteOnDisconnectRequest {
future: Some(Box::pin(future)),
cancel,
producer: producer_for_request,
}
.await
}
@@ -64,6 +91,7 @@ where
{
future: Option<Pin<Box<F>>>,
cancel: Arc<AtomicBool>,
producer: Option<Arc<UsageProducerGuard>>,
}
impl<F> Future for CompleteOnDisconnectRequest<F>
@@ -97,7 +125,9 @@ where
if let (Some(future), Ok(runtime)) =
(self.future.take(), tokio::runtime::Handle::try_current())
{
let producer = self.producer.take();
runtime.spawn(async move {
let _producer = producer;
if let Ok(response) = future.await {
drain_body(response.into_body()).await;
}
@@ -109,6 +139,9 @@ where
struct CompleteOnDisconnectBody {
body: Option<Body>,
diagnostics: Arc<RequestDiagnostics>,
complete_on_disconnect: bool,
// Drop the body first so its terminal handoff registers before this guard ends.
producer: Option<Arc<UsageProducerGuard>>,
}
impl HttpBody for CompleteOnDisconnectBody {
@@ -125,6 +158,7 @@ impl HttpBody for CompleteOnDisconnectBody {
let result = Pin::new(body).poll_frame(context);
if matches!(result, Poll::Ready(None | Some(Err(_)))) {
self.body.take();
self.producer.take();
}
result
}
@@ -143,13 +177,20 @@ impl HttpBody for CompleteOnDisconnectBody {
impl Drop for CompleteOnDisconnectBody {
fn drop(&mut self) {
if !self.complete_on_disconnect {
return;
}
let Some(body) = self.body.take().filter(|body| !body.is_end_stream()) else {
return;
};
if let Ok(runtime) = tokio::runtime::Handle::try_current() {
let producer = self.producer.take();
runtime.spawn(scope_request_diagnostics_with(
Some(Arc::clone(&self.diagnostics)),
drain_body(body),
async move {
let _producer = producer;
drain_body(body).await;
},
));
}
}
@@ -297,6 +338,78 @@ mod tests {
assert!(sender.is_closed());
}
#[tokio::test]
async fn usage_shutdown_waits_for_a_disconnected_request_before_headers() {
let usage = Arc::new(UsageRuntime::disabled());
let (started_tx, started_rx) = oneshot::channel();
let (release_tx, release_rx) = oneshot::channel::<()>();
let request = tokio::spawn(run_request_with_usage(usage.clone(), async move {
configure_client_disconnect(RoutingExecutionPolicy::default());
started_tx.send(()).unwrap();
release_rx.await.unwrap();
Ok(Response::new(Body::empty()))
}));
started_rx.await.unwrap();
request.abort();
assert!(request.await.unwrap_err().is_cancelled());
assert!(usage.shutdown(Duration::from_millis(30)).await.is_err());
assert_eq!(usage.metrics_snapshot().producers_in_flight, 1);
release_tx.send(()).unwrap();
usage.shutdown(Duration::from_secs(1)).await.unwrap();
assert_eq!(usage.metrics_snapshot().producers_in_flight, 0);
}
#[tokio::test]
async fn usage_shutdown_waits_for_disconnected_body_drain() {
let usage = Arc::new(UsageRuntime::disabled());
let (sender, receiver) = mpsc::channel::<Result<Bytes, io::Error>>(1);
let response = run_request_with_usage(usage.clone(), async move {
configure_client_disconnect(RoutingExecutionPolicy::default());
Ok(Response::new(Body::from_stream(stream::unfold(
receiver,
|mut receiver| async { receiver.recv().await.map(|item| (item, receiver)) },
))))
})
.await
.unwrap();
drop(response);
assert!(usage.shutdown(Duration::from_millis(30)).await.is_err());
sender.send(Ok(Bytes::from_static(b"last"))).await.unwrap();
drop(sender);
usage.shutdown(Duration::from_secs(1)).await.unwrap();
assert_eq!(usage.metrics_snapshot().producers_in_flight, 0);
}
#[tokio::test]
async fn tracked_bodies_release_shutdown_on_cancellation_or_eof() {
for cancel_on_client_disconnect in [false, true] {
let usage = Arc::new(UsageRuntime::disabled());
let (sender, receiver) = mpsc::channel::<Result<Bytes, io::Error>>(1);
let response = run_request_with_usage(usage.clone(), async move {
configure_client_disconnect(RoutingExecutionPolicy {
cancel_on_client_disconnect,
..Default::default()
});
Ok(Response::new(Body::from_stream(stream::unfold(
receiver,
|mut receiver| async { receiver.recv().await.map(|item| (item, receiver)) },
))))
})
.await
.unwrap();
let mut body = response.into_body();
if cancel_on_client_disconnect {
drop(body);
assert!(sender.is_closed());
} else {
drop(sender);
assert!(body.frame().await.is_none());
assert_eq!(usage.metrics_snapshot().producers_in_flight, 0);
}
usage.shutdown(Duration::from_secs(1)).await.unwrap();
}
}
#[tokio::test]
async fn connected_response_preserves_headers_size_hint_and_trailers() {
let response = run_request(async {
+16 -1
View File
@@ -18,6 +18,7 @@ use tower::{Service as _, ServiceExt};
use tower_http::services::{ServeDir, ServeFile};
use tracing::warn;
use aether_gateway_frontdoor::{http_connection_limit, HttpConnectionBudget};
use aether_runtime::{prometheus_response, ConcurrencyError};
use aether_runtime_state::RuntimeSemaphoreError;
@@ -244,9 +245,23 @@ pub(crate) enum RequestAdmissionError {
pub async fn serve_tcp(bind: &str) -> Result<(), Box<dyn std::error::Error>> {
let listener = tokio::net::TcpListener::bind(bind).await?;
let router = build_router()?;
let configured_connection_limit = std::env::var("AETHER_GATEWAY_MAX_HTTP_CONNECTIONS")
.ok()
.and_then(|value| value.trim().parse::<usize>().ok());
// This compatibility entry point has no configured request capacities or FD probe.
let connection_budget = Arc::new(HttpConnectionBudget::new(http_connection_limit(
configured_connection_limit,
2048,
2048,
None,
)));
let mut make_service = router.into_make_service_with_connect_info::<std::net::SocketAddr>();
loop {
let (io, remote_addr) = listener.accept().await?;
let (io, remote_addr) = connection_budget.accept(&listener).await;
let Ok(io) = connection_budget.try_admit(io) else {
tokio::task::yield_now().await;
continue;
};
let tower_service = make_service
.call(remote_addr)
.await
@@ -32,6 +32,9 @@ use regex::Regex;
use sha2::{Digest, Sha256};
use std::collections::BTreeMap;
pub(crate) use self::runtime::{
select_with_auth_concurrency_wait, wait_for_auth_api_key_concurrency_retry,
};
pub(crate) use self::selection::{
is_auth_api_key_concurrency_limit_skip_reason, SchedulerSkippedCandidate,
API_KEY_CONCURRENCY_LIMIT_SKIP_REASON, AUTH_API_KEY_CONCURRENCY_LIMIT_SKIP_REASON,
@@ -1,4 +1,5 @@
use std::collections::{BTreeMap, BTreeSet};
use std::future::Future;
use aether_admin::provider::{
pool as admin_provider_pool_pure, status as admin_provider_status_pure,
@@ -10,7 +11,8 @@ use aether_scheduler_core::{
candidate_is_selectable_with_runtime_state, candidate_runtime_skip_reason_with_state,
effective_provider_key_rpm_limit, CandidateRuntimeSelectabilityInput,
};
use std::time::{SystemTime, UNIX_EPOCH};
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use tokio::time::Instant;
use crate::data::auth::GatewayAuthApiKeySnapshot;
use crate::GatewayError;
@@ -109,25 +111,92 @@ pub(super) fn auth_snapshot_concurrency_limit_reached(
snapshot: &CandidateRuntimeSelectionSnapshot,
now_unix_secs: u64,
) -> bool {
auth_snapshot
.and_then(|snapshot| {
usize::try_from(snapshot.api_key_concurrent_limit?)
.ok()
.and_then(|limit| {
if limit == 0 {
return None;
}
Some((snapshot.api_key_id.as_str(), limit))
})
})
.is_some_and(|(api_key_id, limit)| {
auth_api_key_concurrency_limit_reached(
&snapshot.recent_candidates,
now_unix_secs,
api_key_id,
limit,
auth_snapshot_concurrency_limit(auth_snapshot).is_some_and(|(api_key_id, limit)| {
auth_api_key_concurrency_limit_reached(
&snapshot.recent_candidates,
now_unix_secs,
api_key_id,
limit,
)
})
}
fn auth_snapshot_concurrency_limit(
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
) -> Option<(&str, usize)> {
let snapshot = auth_snapshot?;
let limit = usize::try_from(snapshot.api_key_concurrent_limit?).ok()?;
(limit > 0).then_some((snapshot.api_key_id.as_str(), limit))
}
async fn read_auth_api_key_concurrency_limit_reached(
state: &(impl SchedulerRuntimeState + ?Sized),
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
) -> Result<bool, GatewayError> {
let Some((api_key_id, limit)) = auth_snapshot_concurrency_limit(auth_snapshot) else {
return Ok(false);
};
let recent_candidates = state.read_recent_request_candidates(128).await?;
Ok(auth_api_key_concurrency_limit_reached(
&recent_candidates,
crate::clock::current_unix_secs(),
api_key_id,
limit,
))
}
/// A retry always rebuilds candidates, including at the deadline. Only the
/// intervening polls omit catalog, quota and ranking work while auth is blocked.
pub(crate) async fn wait_for_auth_api_key_concurrency_retry(
state: &(impl SchedulerRuntimeState + ?Sized),
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
deadline: Instant,
poll_interval: Duration,
) -> Result<bool, GatewayError> {
if Instant::now() >= deadline {
return Ok(false);
}
let poll_interval = poll_interval.max(Duration::from_millis(1));
loop {
let remaining = deadline.saturating_duration_since(Instant::now());
tokio::time::sleep(poll_interval.min(remaining)).await;
if Instant::now() >= deadline
|| !read_auth_api_key_concurrency_limit_reached(state, auth_snapshot).await?
{
return Ok(true);
}
}
}
pub(crate) async fn select_with_auth_concurrency_wait<T, Select, Selection>(
state: &(impl SchedulerRuntimeState + ?Sized),
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
now_unix_secs: u64,
wait_timeout: Duration,
poll_interval: Duration,
mut select: Select,
) -> Result<T, GatewayError>
where
Select: FnMut(u64) -> Selection,
Selection: Future<Output = Result<(T, bool), GatewayError>>,
{
let deadline = Instant::now() + wait_timeout;
let mut attempt_now_unix_secs = now_unix_secs;
loop {
let (result, auth_limit_blocked) = select(attempt_now_unix_secs).await?;
if !auth_limit_blocked
|| !wait_for_auth_api_key_concurrency_retry(
state,
auth_snapshot,
deadline,
poll_interval,
)
})
.await?
{
return Ok(result);
}
attempt_now_unix_secs = crate::clock::current_unix_secs();
}
}
pub(super) fn is_candidate_selectable(
@@ -0,0 +1,548 @@
use std::collections::VecDeque;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Mutex;
use std::time::Duration;
use aether_data::DataLayerError;
use aether_data_contracts::repository::candidate_selection::{
StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateRowsQuery,
StoredRequestedModelCandidateRowsQuery,
};
use aether_data_contracts::repository::candidates::{
RequestCandidateStatus, StoredRequestCandidate,
};
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use aether_data_contracts::repository::quota::StoredProviderQuotaSnapshot;
use aether_scheduler_core::{SchedulerAffinityTarget, SchedulerMinimalCandidateSelectionCandidate};
use async_trait::async_trait;
use tokio::sync::Notify;
use crate::data::auth::GatewayAuthApiKeySnapshot;
use crate::data::candidate_selection::MinimalCandidateSelectionRowSource;
use crate::scheduler::config::SchedulerOrderingConfig;
use crate::scheduler::state::SchedulerRuntimeState;
use crate::GatewayError;
use super::super::{
is_exact_all_skipped_by_auth_limit,
list_selectable_candidates_for_required_capability_without_requested_model_with_auth_limit_signal,
list_selectable_candidates_with_skip_reasons, select_with_auth_concurrency_wait,
SchedulerSkippedCandidate,
};
use super::support::{sample_auth_snapshot, sample_provider, sample_row};
const POLL_INTERVAL: Duration = Duration::from_millis(2);
enum RecentReadAction {
Keep,
Release,
ReleaseAndReplaceCandidate,
ReleaseAndRemoveCandidates,
Fail,
}
struct CountingState {
rows: Mutex<Vec<StoredMinimalCandidateSelectionRow>>,
recent: Mutex<Vec<StoredRequestCandidate>>,
recent_actions: Mutex<VecDeque<RecentReadAction>>,
row_reads: AtomicUsize,
format_reads: AtomicUsize,
provider_reads: AtomicUsize,
key_reads: AtomicUsize,
quota_reads: AtomicUsize,
recent_reads: AtomicUsize,
row_error_at: Option<usize>,
first_row_delay: Duration,
poll_observed: Notify,
}
impl CountingState {
fn blocked() -> Self {
Self {
rows: Mutex::new(vec![sample_row()]),
recent: Mutex::new(vec![active_candidate()]),
recent_actions: Mutex::new(VecDeque::new()),
row_reads: AtomicUsize::new(0),
format_reads: AtomicUsize::new(0),
provider_reads: AtomicUsize::new(0),
key_reads: AtomicUsize::new(0),
quota_reads: AtomicUsize::new(0),
recent_reads: AtomicUsize::new(0),
row_error_at: None,
first_row_delay: Duration::ZERO,
poll_observed: Notify::new(),
}
}
fn on_recent_reads(self, actions: impl IntoIterator<Item = RecentReadAction>) -> Self {
*self.recent_actions.lock().unwrap() = actions.into_iter().collect();
self
}
}
fn active_candidate() -> StoredRequestCandidate {
let now_ms = i64::try_from(crate::clock::current_unix_secs() * 1000).unwrap();
StoredRequestCandidate::new(
"active-candidate".to_string(),
"active-request".to_string(),
Some("user-1".to_string()),
Some("api-key-1".to_string()),
None,
None,
0,
0,
Some("provider-1".to_string()),
Some("endpoint-1".to_string()),
Some("key-1".to_string()),
RequestCandidateStatus::Streaming,
None,
false,
None,
None,
None,
None,
None,
None,
None,
now_ms,
Some(now_ms),
None,
)
.unwrap()
}
fn limited_auth() -> GatewayAuthApiKeySnapshot {
let mut auth = sample_auth_snapshot("api-key-1");
auth.api_key_concurrent_limit = Some(1);
auth
}
#[async_trait]
impl MinimalCandidateSelectionRowSource for CountingState {
async fn read_minimal_candidate_selection_rows_for_api_format_and_global_model(
&self,
_api_format: &str,
_global_model_name: &str,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
panic!("requested-model selection should use its paged query")
}
async fn read_minimal_candidate_selection_rows_for_api_format_and_requested_model(
&self,
_api_format: &str,
_requested_model_name: &str,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
panic!("requested-model selection should use its paged query")
}
async fn read_minimal_candidate_selection_rows_for_api_format_and_requested_model_page(
&self,
query: &StoredRequestedModelCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
let read = self.row_reads.fetch_add(1, Ordering::SeqCst) + 1;
if read == 1 && !self.first_row_delay.is_zero() {
tokio::time::sleep(self.first_row_delay).await;
}
if self.row_error_at == Some(read) {
return Err(DataLayerError::Postgres(
"candidate query failed".to_string(),
));
}
Ok(self
.rows
.lock()
.unwrap()
.iter()
.filter(|row| {
row.endpoint_api_format == query.api_format
&& row.global_model_name == query.requested_model_name
})
.skip(query.offset as usize)
.take(query.limit as usize)
.cloned()
.collect())
}
async fn read_minimal_candidate_selection_rows_for_api_format(
&self,
_api_format: &str,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
self.format_reads.fetch_add(1, Ordering::SeqCst);
Ok(self.rows.lock().unwrap().clone())
}
async fn read_pool_key_candidate_rows_for_group(
&self,
_query: &StoredPoolKeyCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
panic!("the test has no provider pool")
}
}
#[async_trait]
impl SchedulerRuntimeState for CountingState {
async fn read_provider_quota_snapshot(
&self,
_provider_id: &str,
) -> Result<Option<StoredProviderQuotaSnapshot>, GatewayError> {
self.quota_reads.fetch_add(1, Ordering::SeqCst);
Ok(None)
}
async fn read_provider_catalog_providers_by_ids(
&self,
provider_ids: &[String],
) -> Result<Vec<StoredProviderCatalogProvider>, GatewayError> {
self.provider_reads.fetch_add(1, Ordering::SeqCst);
Ok(provider_ids
.iter()
.map(|id| sample_provider(id, None))
.collect())
}
async fn read_provider_catalog_keys_by_ids(
&self,
_key_ids: &[String],
) -> Result<Vec<StoredProviderCatalogKey>, GatewayError> {
self.key_reads.fetch_add(1, Ordering::SeqCst);
Ok(Vec::new())
}
async fn read_recent_request_candidates(
&self,
limit: usize,
) -> Result<Vec<StoredRequestCandidate>, GatewayError> {
assert_eq!(
limit, 128,
"polls must use the same sample as full selection"
);
let read = self.recent_reads.fetch_add(1, Ordering::SeqCst) + 1;
let action = self
.recent_actions
.lock()
.unwrap()
.pop_front()
.unwrap_or(RecentReadAction::Keep);
if matches!(action, RecentReadAction::Fail) {
return Err(GatewayError::Internal(
"recent candidates failed".to_string(),
));
}
if matches!(
action,
RecentReadAction::Release
| RecentReadAction::ReleaseAndReplaceCandidate
| RecentReadAction::ReleaseAndRemoveCandidates
) {
for candidate in self.recent.lock().unwrap().iter_mut() {
candidate.status = RequestCandidateStatus::Success;
candidate.finished_at_unix_ms = Some(crate::clock::current_unix_secs() * 1000);
}
}
match action {
RecentReadAction::ReleaseAndReplaceCandidate => {
self.rows.lock().unwrap()[0].key_id = "replacement-key".to_string();
}
RecentReadAction::ReleaseAndRemoveCandidates => self.rows.lock().unwrap().clear(),
_ => {}
}
if read > 1 {
self.poll_observed.notify_one();
}
Ok(self.recent.lock().unwrap().clone())
}
fn provider_key_rpm_reset_at(&self, _key_id: &str, _now_unix_secs: u64) -> Option<u64> {
None
}
fn read_cached_scheduler_affinity_target(
&self,
_cache_key: &str,
_ttl: Duration,
) -> Option<SchedulerAffinityTarget> {
None
}
fn scheduler_affinity_epoch(&self) -> u64 {
0
}
fn remember_scheduler_affinity_target(
&self,
_cache_key: &str,
_target: SchedulerAffinityTarget,
_ttl: Duration,
_max_entries: usize,
) {
}
fn remember_scheduler_affinity_target_for_epoch(
&self,
_cache_key: &str,
_target: SchedulerAffinityTarget,
_ttl: Duration,
_max_entries: usize,
_expected_epoch: Option<u64>,
) -> bool {
true
}
}
type Selection = (
Vec<SchedulerMinimalCandidateSelectionCandidate>,
Vec<SchedulerSkippedCandidate>,
);
async fn select_requested_model(
state: &CountingState,
auth: Option<&GatewayAuthApiKeySnapshot>,
timeout: Duration,
) -> Result<Selection, GatewayError> {
select_with_auth_concurrency_wait(
state,
auth,
crate::clock::current_unix_secs(),
timeout,
POLL_INTERVAL,
|now| async move {
let result = list_selectable_candidates_with_skip_reasons(
state,
state,
"openai:chat",
"gpt-4.1",
false,
None,
auth,
None,
now,
false,
SchedulerOrderingConfig::default(),
)
.await?;
let blocked = is_exact_all_skipped_by_auth_limit(&result.0, &result.1);
Ok((result, blocked))
},
)
.await
}
#[tokio::test]
async fn concurrent_blocked_selectors_only_prepare_at_start_and_deadline() {
let state = CountingState::blocked();
let auth = limited_auth();
let outcomes = futures_util::future::join_all(
(0..8).map(|_| select_requested_model(&state, Some(&auth), Duration::from_millis(80))),
)
.await;
for outcome in outcomes {
let (selected, skipped) = outcome.unwrap();
assert!(is_exact_all_skipped_by_auth_limit(&selected, &skipped));
}
assert_eq!(state.row_reads.load(Ordering::SeqCst), 16);
assert_eq!(state.provider_reads.load(Ordering::SeqCst), 32);
assert_eq!(state.key_reads.load(Ordering::SeqCst), 16);
assert_eq!(state.quota_reads.load(Ordering::SeqCst), 16);
assert!(state.recent_reads.load(Ordering::SeqCst) > 16);
assert_eq!(state.format_reads.load(Ordering::SeqCst), 0);
}
#[tokio::test]
async fn released_auth_slot_rebuilds_changed_candidates() {
let state = CountingState::blocked().on_recent_reads([
RecentReadAction::Keep,
RecentReadAction::Keep,
RecentReadAction::ReleaseAndReplaceCandidate,
]);
let auth = limited_auth();
let (selected, skipped) = select_requested_model(&state, Some(&auth), Duration::from_secs(1))
.await
.unwrap();
assert!(skipped.is_empty());
assert_eq!(selected.len(), 1);
assert_eq!(selected[0].key_id, "replacement-key");
assert_eq!(state.row_reads.load(Ordering::SeqCst), 2);
assert_eq!(state.recent_reads.load(Ordering::SeqCst), 4);
}
#[tokio::test]
async fn released_auth_slot_does_not_reuse_removed_candidates() {
let state = CountingState::blocked().on_recent_reads([
RecentReadAction::Keep,
RecentReadAction::ReleaseAndRemoveCandidates,
]);
let (selected, skipped) =
select_requested_model(&state, Some(&limited_auth()), Duration::from_secs(1))
.await
.unwrap();
assert!(selected.is_empty());
assert!(skipped.is_empty());
assert_eq!(state.row_reads.load(Ordering::SeqCst), 2);
}
#[tokio::test]
async fn lightweight_poll_errors_are_propagated_without_another_full_query() {
let state =
CountingState::blocked().on_recent_reads([RecentReadAction::Keep, RecentReadAction::Fail]);
let error = select_requested_model(&state, Some(&limited_auth()), Duration::from_secs(1))
.await
.unwrap_err();
assert!(
matches!(error, GatewayError::Internal(message) if message == "recent candidates failed")
);
assert_eq!(state.row_reads.load(Ordering::SeqCst), 1);
assert_eq!(state.recent_reads.load(Ordering::SeqCst), 2);
}
#[tokio::test]
async fn recovered_selection_errors_are_propagated() {
let mut state = CountingState::blocked()
.on_recent_reads([RecentReadAction::Keep, RecentReadAction::Release]);
state.row_error_at = Some(2);
let error = select_requested_model(&state, Some(&limited_auth()), Duration::from_secs(1))
.await
.unwrap_err();
assert!(
matches!(error, GatewayError::Internal(message) if message.contains("candidate query failed"))
);
assert_eq!(state.row_reads.load(Ordering::SeqCst), 2);
}
#[tokio::test]
async fn first_selection_time_consumes_the_wait_budget() {
let mut state = CountingState::blocked();
state.first_row_delay = Duration::from_millis(30);
let (selected, skipped) =
select_requested_model(&state, Some(&limited_auth()), Duration::from_millis(5))
.await
.unwrap();
assert!(is_exact_all_skipped_by_auth_limit(&selected, &skipped));
assert_eq!(state.row_reads.load(Ordering::SeqCst), 1);
assert_eq!(state.recent_reads.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn cancellation_stops_polling_and_candidate_queries() {
let state = CountingState::blocked();
let auth = limited_auth();
let mut selection = Box::pin(select_requested_model(
&state,
Some(&auth),
Duration::from_secs(1),
));
tokio::select! {
result = &mut selection => panic!("selection completed before cancellation: {result:?}"),
_ = state.poll_observed.notified() => {}
}
drop(selection);
let recent_reads = state.recent_reads.load(Ordering::SeqCst);
tokio::time::sleep(Duration::from_millis(10)).await;
assert!(recent_reads >= 2);
assert_eq!(state.recent_reads.load(Ordering::SeqCst), recent_reads);
assert_eq!(state.row_reads.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn no_model_capability_selection_does_not_reenumerate_during_polls() {
let state = CountingState::blocked();
let auth = limited_auth();
let selected = select_with_auth_concurrency_wait(
&state,
Some(&auth),
crate::clock::current_unix_secs(),
Duration::from_millis(80),
POLL_INTERVAL,
|now| {
list_selectable_candidates_for_required_capability_without_requested_model_with_auth_limit_signal(
&state,
&state,
"openai:chat",
"cache_1h",
false,
Some(&auth),
None,
now,
SchedulerOrderingConfig::default(),
)
},
)
.await
.unwrap();
assert!(selected.is_empty());
assert_eq!(state.format_reads.load(Ordering::SeqCst), 2);
assert_eq!(state.row_reads.load(Ordering::SeqCst), 2);
assert!(state.recent_reads.load(Ordering::SeqCst) > 2);
}
#[tokio::test]
async fn absent_or_disabled_auth_limits_do_not_poll() {
for limit in [None, Some(0), Some(-1)] {
let state = CountingState::blocked();
let mut auth = limited_auth();
auth.api_key_concurrent_limit = limit;
let (selected, _) = select_requested_model(&state, Some(&auth), Duration::from_secs(1))
.await
.unwrap();
assert_eq!(selected.len(), 1);
assert_eq!(state.row_reads.load(Ordering::SeqCst), 1);
assert_eq!(state.recent_reads.load(Ordering::SeqCst), 0);
}
let state = CountingState::blocked();
assert_eq!(
select_requested_model(&state, None, Duration::from_secs(1))
.await
.unwrap()
.0
.len(),
1
);
assert_eq!(state.recent_reads.load(Ordering::SeqCst), 0);
}
#[tokio::test]
async fn auth_wait_keeps_candidate_row_and_lifecycle_counting_semantics() {
let state = CountingState::blocked();
let mut duplicate = active_candidate();
duplicate.id = "second-attempt-same-request".to_string();
state.recent.lock().unwrap().push(duplicate);
let mut auth = limited_auth();
auth.api_key_concurrent_limit = Some(2);
let (selected, skipped) = select_requested_model(&state, Some(&auth), Duration::ZERO)
.await
.unwrap();
assert!(is_exact_all_skipped_by_auth_limit(&selected, &skipped));
let state = CountingState::blocked();
state.recent.lock().unwrap()[0].finished_at_unix_ms =
Some(crate::clock::current_unix_secs() * 1000);
assert_eq!(
select_requested_model(&state, Some(&limited_auth()), Duration::ZERO)
.await
.unwrap()
.0
.len(),
1
);
let state = CountingState::blocked();
state.recent.lock().unwrap()[0].started_at_unix_ms =
Some(crate::clock::current_unix_secs().saturating_sub(301) * 1000);
assert_eq!(
select_requested_model(&state, Some(&limited_auth()), Duration::ZERO)
.await
.unwrap()
.0
.len(),
1
);
}
@@ -1,4 +1,5 @@
mod affinity;
mod concurrency_wait;
mod model;
mod required_capability;
mod selection;
+209
View File
@@ -0,0 +1,209 @@
use std::sync::Arc;
use std::time::Duration;
use axum::{body::Body, routing::get, Router};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpListener, TcpStream};
use tokio::sync::Notify;
use tokio_util::sync::CancellationToken;
use super::super::{serve_gateway_router, HttpConnectionBudget};
async fn start(
router: Router,
) -> (
std::net::SocketAddr,
CancellationToken,
Arc<HttpConnectionBudget>,
tokio::task::JoinHandle<Result<(), String>>,
) {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();
let shutdown = CancellationToken::new();
let budget = Arc::new(HttpConnectionBudget::new(16));
let stop = shutdown.clone();
let shared = Arc::clone(&budget);
let server = tokio::spawn(async move {
serve_gateway_router(
vec![listener],
router,
shared,
16,
10_000,
32_768,
100,
stop,
)
.await
.map_err(|error| error.to_string())
});
(address, shutdown, budget, server)
}
async fn within<T>(future: impl std::future::Future<Output = T>) -> T {
tokio::time::timeout(Duration::from_secs(3), future)
.await
.expect("shutdown deadline")
}
#[tokio::test]
async fn gateway_shutdown_drains_in_flight_http1_and_http2_responses() {
for http2 in [false, true] {
let started = Arc::new(Notify::new());
let release = Arc::new(Notify::new());
let handler_started = Arc::clone(&started);
let handler_release = Arc::clone(&release);
let router = Router::new().route(
"/",
get(move || {
let started = Arc::clone(&handler_started);
let release = Arc::clone(&handler_release);
async move {
started.notify_one();
release.notified().await;
"complete response"
}
}),
);
let (address, shutdown, budget, server) = start(router).await;
let client = if http2 {
reqwest::Client::builder().http2_prior_knowledge()
} else {
reqwest::Client::builder().http1_only()
}
.build()
.unwrap();
let request = tokio::spawn(async move {
client
.get(format!("http://{address}/"))
.send()
.await
.unwrap()
.text()
.await
.unwrap()
});
within(started.notified()).await;
shutdown.cancel();
tokio::time::sleep(Duration::from_millis(20)).await;
assert!(!server.is_finished());
assert!(!request.is_finished());
assert!(TcpStream::connect(address).await.is_err());
release.notify_one();
assert_eq!(within(request).await.unwrap(), "complete response");
within(server).await.unwrap().unwrap();
assert_eq!(budget.snapshot().in_flight, 0);
}
}
#[tokio::test]
async fn gateway_shutdown_closes_idle_protocol_detection_connections() {
let (address, shutdown, budget, server) = start(Router::new()).await;
let mut peer = TcpStream::connect(address).await.unwrap();
within(async {
while budget.snapshot().in_flight == 0 {
tokio::task::yield_now().await;
}
})
.await;
shutdown.cancel();
within(server).await.unwrap().unwrap();
assert_eq!(within(peer.read(&mut [0_u8; 1])).await.unwrap(), 0);
assert_eq!(budget.snapshot().in_flight, 0);
}
#[tokio::test]
async fn gateway_shutdown_force_cancels_a_handler_without_socket_io() {
struct HandlerDrop(Arc<Notify>);
impl Drop for HandlerDrop {
fn drop(&mut self) {
self.0.notify_one();
}
}
for http2 in [false, true] {
let started = Arc::new(Notify::new());
let dropped = Arc::new(Notify::new());
let request_started = Arc::clone(&started);
let request_dropped = Arc::clone(&dropped);
let router = Router::new().route(
"/",
get(move || {
let started = Arc::clone(&request_started);
let dropped = Arc::clone(&request_dropped);
async move {
let _guard = HandlerDrop(dropped);
started.notify_one();
std::future::pending::<&'static str>().await
}
}),
);
let (address, shutdown, budget, server) = start(router).await;
let client = if http2 {
reqwest::Client::builder().http2_prior_knowledge()
} else {
reqwest::Client::builder().http1_only()
}
.build()
.unwrap();
let request =
tokio::spawn(async move { client.get(format!("http://{address}/")).send().await });
within(started.notified()).await;
shutdown.cancel();
budget.force_close();
within(server).await.unwrap().unwrap();
within(dropped.notified()).await;
assert!(within(request).await.unwrap().is_err());
assert_eq!(budget.snapshot().in_flight, 0);
}
}
#[tokio::test]
async fn gateway_shutdown_force_closes_upgraded_io_and_waits_for_release() {
let upgraded_done = Arc::new(Notify::new());
let done = Arc::clone(&upgraded_done);
let router = Router::new().route(
"/",
get(move |mut request: axum::extract::Request| {
let done = Arc::clone(&done);
async move {
let upgrade = hyper::upgrade::on(&mut request);
tokio::spawn(async move {
let upgraded = upgrade.await.unwrap();
let mut io = hyper_util::rt::TokioIo::new(upgraded);
let error = io.read_u8().await.unwrap_err();
assert_eq!(error.kind(), std::io::ErrorKind::ConnectionAborted);
drop(io);
done.notify_one();
});
axum::http::Response::builder()
.status(101)
.header("connection", "upgrade")
.header("upgrade", "echo")
.body(Body::empty())
.unwrap()
}
}),
);
let (address, shutdown, budget, server) = start(router).await;
let mut peer = TcpStream::connect(address).await.unwrap();
peer.write_all(
b"GET / HTTP/1.1\r\nHost: localhost\r\nConnection: upgrade\r\nUpgrade: echo\r\n\r\n",
)
.await
.unwrap();
let mut header = Vec::new();
within(async {
while !header.ends_with(b"\r\n\r\n") {
header.push(peer.read_u8().await.unwrap());
}
})
.await;
assert!(header.starts_with(b"HTTP/1.1 101"));
shutdown.cancel();
tokio::time::sleep(Duration::from_millis(20)).await;
assert!(!server.is_finished());
budget.force_close();
within(upgraded_done.notified()).await;
within(server).await.unwrap().unwrap();
assert_eq!(budget.snapshot().in_flight, 0);
}
+21 -12
View File
@@ -8,6 +8,7 @@ use aether_data::repository::users::StoredUserGroup;
use aether_data_contracts::repository::billing::UserDailyQuotaAvailabilityRecord;
use aether_data_contracts::repository::quota::StoredProviderQuotaSnapshot;
use aether_data_contracts::repository::usage::UsageCounterHealthSnapshot;
use aether_gateway_frontdoor::HttpConnectionBudget;
use aether_runtime::ConcurrencyGate;
use aether_runtime_state::{RuntimeSemaphore, RuntimeState};
use dashmap::DashMap;
@@ -35,6 +36,7 @@ use super::{
};
const MIN_REQUEST_BODY_READ_TIMEOUT_MS: u64 = 1_000;
const DEFAULT_REQUEST_BODY_READ_TIMEOUT_MS: u64 = 120_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;
@@ -193,7 +195,9 @@ fn optional_env_duration_ms(key: &str, min_ms: u64, max_ms: u64) -> Option<Durat
}
fn parse_optional_duration_ms(raw: Option<&str>, min_ms: u64, max_ms: u64) -> Option<Duration> {
let parsed = raw?.trim().parse::<u64>().ok()?;
let parsed = raw
.and_then(|value| value.trim().parse::<u64>().ok())
.unwrap_or(DEFAULT_REQUEST_BODY_READ_TIMEOUT_MS);
if parsed == 0 {
return None;
}
@@ -388,6 +392,7 @@ pub struct AppState {
pub(crate) video_task_poller: Option<VideoTaskPollerConfig>,
pub(crate) frontdoor_runtime_guards: Arc<FrontdoorRuntimeGuardConfig>,
pub(crate) request_body_buffer_budget: Arc<Semaphore>,
pub(crate) http_connection_budget: Option<Arc<HttpConnectionBudget>>,
pub(crate) request_gate: Option<Arc<ConcurrencyGate>>,
pub(crate) websocket_connection_gate: Option<Arc<ConcurrencyGate>>,
pub(crate) auth_snapshot_load_gate: Option<Arc<ConcurrencyGate>>,
@@ -456,6 +461,8 @@ pub struct AppState {
Arc<DashMap<String, LocalExecutionRuntimeMissDiagnostic>>,
pub(crate) admin_monitoring_error_stats_reset_at: Arc<StdMutex<Option<u64>>>,
pub(crate) provider_delete_tasks: Arc<StdMutex<HashMap<String, LocalProviderDeleteTaskState>>>,
pub(crate) pool_quota_probe_replenish:
Arc<crate::maintenance::PoolQuotaProbeReplenishCoordinator>,
#[cfg(test)]
pub(crate) turnstile_siteverify_url_override: Option<String>,
#[cfg(test)]
@@ -533,20 +540,22 @@ mod tests {
};
#[test]
fn request_body_read_timeout_parser_defaults_to_disabled() {
assert_eq!(
parse_optional_duration_ms(
None,
MIN_REQUEST_BODY_READ_TIMEOUT_MS,
MAX_REQUEST_BODY_READ_TIMEOUT_MS,
),
None
);
fn request_body_read_timeout_parser_uses_a_finite_default() {
for value in [None, Some(""), Some("invalid"), Some("-1")] {
assert_eq!(
parse_optional_duration_ms(
value,
MIN_REQUEST_BODY_READ_TIMEOUT_MS,
MAX_REQUEST_BODY_READ_TIMEOUT_MS,
),
Some(Duration::from_millis(DEFAULT_REQUEST_BODY_READ_TIMEOUT_MS)),
);
}
}
#[test]
fn request_body_read_timeout_parser_disables_zero_and_invalid_values() {
for value in ["", "invalid", "-1", "0", " 0 "] {
fn request_body_read_timeout_parser_only_disables_explicit_zero() {
for value in ["0", " 0 "] {
assert_eq!(
parse_optional_duration_ms(
Some(value),
+380
View File
@@ -13,6 +13,7 @@ use aether_data::repository::proxy_nodes::{
use aether_data_contracts::repository::usage::{
UsageCounterHealthSnapshot, UsageCounterPendingHealthSnapshot,
};
use aether_gateway_frontdoor::{HttpConnectionBudget, HttpConnectionBudgetSnapshot};
use aether_http::{apply_http_client_config, HttpClientConfig};
use aether_runtime::{
service_up_sample, AdmissionPermit, ConcurrencyGate, ConcurrencySnapshot, MetricKind,
@@ -358,6 +359,7 @@ impl AppState {
request_body_buffer_budget: Arc::new(tokio::sync::Semaphore::new(
frontdoor_runtime_guards.request_body_buffer_budget_permits,
)),
http_connection_budget: None,
request_gate: None,
websocket_connection_gate: None,
auth_snapshot_load_gate: frontdoor_runtime_guards
@@ -438,6 +440,9 @@ impl AppState {
local_execution_runtime_miss_diagnostics: Arc::new(DashMap::new()),
admin_monitoring_error_stats_reset_at: Arc::new(StdMutex::new(None)),
provider_delete_tasks: Arc::new(StdMutex::new(HashMap::new())),
pool_quota_probe_replenish: Arc::new(
crate::maintenance::PoolQuotaProbeReplenishCoordinator::default(),
),
#[cfg(test)]
turnstile_siteverify_url_override: None,
#[cfg(test)]
@@ -614,6 +619,11 @@ impl AppState {
self
}
pub fn with_http_connection_budget(mut self, budget: Arc<HttpConnectionBudget>) -> Self {
self.http_connection_budget = Some(budget);
self
}
pub fn with_request_concurrency_limit(mut self, limit: usize) -> Self {
let limit = limit.max(1);
self.request_gate = Some(Arc::new(ConcurrencyGate::new("gateway_requests", limit)));
@@ -635,6 +645,10 @@ impl AppState {
}
pub fn with_runtime_state(mut self, runtime_state: Arc<RuntimeState>) -> Self {
if !Arc::ptr_eq(&self.runtime_state, &runtime_state) {
self.pool_quota_probe_replenish =
Arc::new(crate::maintenance::PoolQuotaProbeReplenishCoordinator::default());
}
self.runtime_state = runtime_state;
self.admin_security_blacklist_cache.clear();
self.admin_security_whitelist_cache.clear();
@@ -1666,6 +1680,9 @@ impl AppState {
.unwrap_or(u64::MAX),
),
]);
if let Some(budget) = &self.http_connection_budget {
samples.extend(http_connection_metric_samples(&budget.snapshot()));
}
if let Some(snapshot) = self.request_concurrency_snapshot() {
samples.extend(snapshot.to_metric_samples("gateway_requests"));
}
@@ -1820,10 +1837,44 @@ impl AppState {
&self.task_supervisor_metrics.snapshot(),
));
samples.extend(crate::tokio_metrics::gateway_tokio_runtime_metric_samples());
samples.extend(aether_runtime::logging_metric_samples());
samples.extend(
crate::execution_runtime::transport::direct_reqwest_client_cache_metric_samples(),
);
samples.extend(self.upstream_target_admission.metric_samples());
let probe_replenish = self.pool_quota_probe_replenish.snapshot();
samples.extend([
MetricSample::new(
"pool_quota_probe_replenish_provider_capacity",
"Maximum providers with an active local request-triggered probe task.",
MetricKind::Gauge,
probe_replenish.capacity as u64,
),
MetricSample::new(
"pool_quota_probe_replenish_active_providers",
"Providers with an active local request-triggered probe task.",
MetricKind::Gauge,
probe_replenish.active as u64,
),
MetricSample::new(
"pool_quota_probe_replenish_started_total",
"Local request-triggered provider probe tasks admitted.",
MetricKind::Counter,
probe_replenish.started_total,
),
MetricSample::new(
"pool_quota_probe_replenish_coalesced_total",
"Request triggers merged into an existing local provider probe task.",
MetricKind::Counter,
probe_replenish.coalesced_total,
),
MetricSample::new(
"pool_quota_probe_replenish_capacity_rejected_total",
"Request-triggered probe tasks skipped because the local provider limit was reached.",
MetricKind::Counter,
probe_replenish.capacity_rejected_total,
),
]);
samples.extend(crate::cache::candidate_page_cache_metric_samples());
samples.extend(crate::stage_metrics::gateway_stage_metric_samples());
samples.extend(self.tunnel.metric_samples());
@@ -2118,6 +2169,36 @@ impl AppState {
state
}
pub async fn shutdown_usage_runtime(
&self,
timeout: Duration,
) -> Result<(), aether_data_contracts::DataLayerError> {
tokio::time::timeout(timeout, async {
let local_queue = self.runtime_state.is_memory().then(|| {
let queue: Arc<dyn RuntimeQueueStore> = self.runtime_state.clone();
queue
});
self.usage_runtime
.shutdown_with_local_queue(timeout, local_queue)
.await?;
while self
.request_candidate_queue
.as_ref()
.is_some_and(|queue| queue.pending_writes() != 0)
{
tokio::time::sleep(Duration::from_millis(10)).await;
}
Ok(())
})
.await
.map_err(|_| {
aether_data_contracts::DataLayerError::TimedOut(
"gateway usage or candidate persistence did not drain before shutdown deadline"
.to_string(),
)
})?
}
pub fn spawn_background_tasks(&self) -> crate::task_runtime::TaskSupervisor {
let background_state = self.background_worker_state();
let mut supervisor =
@@ -2277,6 +2358,41 @@ fn database_bounded_auth_load_limit(
})
}
fn http_connection_metric_samples(snapshot: &HttpConnectionBudgetSnapshot) -> Vec<MetricSample> {
vec![
MetricSample::new(
"gateway_http_connections_limit",
"Maximum admitted inbound TCP connections across gateway listeners.",
MetricKind::Gauge,
u64::try_from(snapshot.limit).unwrap_or(u64::MAX),
),
MetricSample::new(
"gateway_http_connections_in_flight",
"Currently admitted inbound TCP connections, including upgraded connections.",
MetricKind::Gauge,
u64::try_from(snapshot.in_flight).unwrap_or(u64::MAX),
),
MetricSample::new(
"gateway_http_connections_high_watermark",
"Highest simultaneous admitted inbound TCP connection count.",
MetricKind::Gauge,
u64::try_from(snapshot.high_watermark).unwrap_or(u64::MAX),
),
MetricSample::new(
"gateway_http_connections_rejected_total",
"Inbound TCP connections closed because the connection budget was full.",
MetricKind::Counter,
snapshot.rejected_total,
),
MetricSample::new(
"gateway_http_connections_accept_errors_total",
"Listener accept errors retried with bounded backoff.",
MetricKind::Counter,
snapshot.accept_errors_total,
),
]
}
fn task_supervisor_metric_samples(
snapshot: &aether_task_runtime::TaskSupervisorMetricsSnapshot,
) -> Vec<MetricSample> {
@@ -3190,6 +3306,18 @@ fn usage_runtime_metric_samples(
snapshot: &usage::UsageRuntimeMetricsSnapshot,
) -> Vec<MetricSample> {
vec![
MetricSample::new(
"usage_runtime_shutdown_started", "Whether local usage admission is closed for shutdown.",
MetricKind::Gauge, u64::from(snapshot.shutdown_started),
),
MetricSample::new(
"usage_runtime_producers_in_flight", "Tracked requests and finalizers that can still submit usage events.",
MetricKind::Gauge, snapshot.producers_in_flight as u64,
),
MetricSample::new(
"usage_runtime_delayed_lifecycle_pending", "Delayed lifecycle events still held in this process.",
MetricKind::Gauge, snapshot.delayed_lifecycle_pending as u64,
),
MetricSample::new(
"usage_runtime_enabled",
"Whether the gateway usage runtime is enabled.",
@@ -3322,6 +3450,132 @@ fn usage_runtime_metric_samples(
MetricKind::Gauge,
u64::from(snapshot.retry_deferred_lifecycle_events),
),
MetricSample::new(
"usage_runtime_queue_payload_max_bytes",
"Maximum serialized JSON payload bytes for a new usage queue message.",
MetricKind::Gauge,
snapshot.queue_payload_max_bytes as u64,
),
MetricSample::new(
"usage_runtime_queue_payload_downgraded_total",
"Process-wide enqueue and retry validation encoding attempts that omitted diagnostic data after exceeding the payload limit; not unique events.",
MetricKind::Counter,
snapshot.queue_payload_downgraded_total,
),
MetricSample::new(
"usage_runtime_queue_payload_rejected_total",
"Process-wide enqueue and retry validation encoding attempts rejected because the payload exceeded its limit or billing facts could not be preserved; not unique events.",
MetricKind::Counter,
snapshot.queue_payload_rejected_total,
),
MetricSample::new(
"usage_runtime_queue_read_payload_budget_bytes",
"Process-wide payload reservation limit for usage worker reads and reclaims; not a wire or heap limit.",
MetricKind::Gauge,
snapshot.queue_read_payload_budget_bytes as u64,
),
MetricSample::new(
"usage_runtime_queue_read_batch_payload_bytes",
"Target payload reservation per usage worker batch, allowing at least one configured maximum payload.",
MetricKind::Gauge,
snapshot.queue_read_batch_payload_bytes as u64,
),
MetricSample::new(
"usage_runtime_queue_read_payload_reserved_bytes",
"Process-wide logical payload bytes reserved by usage worker reads, reclaims and unprocessed batches.",
MetricKind::Gauge,
snapshot.queue_read_payload_reserved_bytes as u64,
),
MetricSample::new(
"usage_runtime_queue_read_payload_waiters",
"Usage workers currently waiting for shared payload reservation capacity.",
MetricKind::Gauge,
snapshot.queue_read_payload_waiters as u64,
),
MetricSample::new(
"usage_runtime_queue_read_payload_wait_total",
"Process-wide usage worker batch reservation attempts that had to wait for capacity.",
MetricKind::Counter,
snapshot.queue_read_payload_wait_total,
),
MetricSample::new(
"usage_runtime_queue_read_actual_field_bytes_total",
"Cumulative key and value bytes observed in reserved usage worker batches, including repeated reclaims; excludes Redis framing and allocations.",
MetricKind::Counter,
snapshot.queue_read_actual_field_bytes_total,
),
MetricSample::new(
"usage_runtime_queue_read_oversized_entries_total",
"Observed usage entries whose combined field value bytes exceed the consumer payload estimate, including repeated reclaims.",
MetricKind::Counter,
snapshot.queue_read_oversized_entries_total,
),
MetricSample::new(
"usage_runtime_queue_read_oversized_batches_total",
"Observed usage batches whose combined field value bytes exceed their initial reservation; entries remain eligible for processing.",
MetricKind::Counter,
snapshot.queue_read_oversized_batches_total,
),
MetricSample::new(
"usage_runtime_event_capture_memory_budget_bytes",
"Process-wide diagnostic JSON heap estimate budget for retained usage events.",
MetricKind::Gauge,
snapshot.event_capture_memory_budget_bytes as u64,
),
MetricSample::new(
"usage_runtime_dlq_encoding_budget_bytes",
"Process-wide logical raw string and JSON reservation limit for dead letter encoding; excludes Redis buffers.",
MetricKind::Gauge,
snapshot.dlq_encoding_budget_bytes as u64,
),
MetricSample::new(
"usage_runtime_dlq_encoding_max_jobs",
"Maximum admitted dead letter encoding and write jobs per process.",
MetricKind::Gauge,
snapshot.dlq_encoding_max_jobs as u64,
),
MetricSample::new(
"usage_runtime_dlq_encoding_reserved_bytes",
"Logical raw string and JSON bytes reserved by admitted dead letter jobs.",
MetricKind::Gauge,
snapshot.dlq_encoding_reserved_bytes as u64,
),
MetricSample::new(
"usage_runtime_dlq_encoding_active_jobs",
"Admitted dead letter jobs awaiting or performing encoding or queue writes.",
MetricKind::Gauge,
snapshot.dlq_encoding_active_jobs as u64,
),
MetricSample::new(
"usage_runtime_dlq_encoding_capacity_rejected_total",
"Dead letter attempts rejected because encoding byte or job capacity was occupied; source remains pending.",
MetricKind::Counter,
snapshot.dlq_encoding_capacity_rejected_total,
),
MetricSample::new(
"usage_runtime_dlq_encoding_oversized_rejected_total",
"Dead letter attempts rejected because the conservative encoding reservation exceeded the total budget or overflowed.",
MetricKind::Counter,
snapshot.dlq_encoding_oversized_rejected_total,
),
MetricSample::new(
"usage_runtime_dlq_encoding_encoded_total",
"Completed dead letter JSON encodings, including repeated attempts; not successful archives.",
MetricKind::Counter,
snapshot.dlq_encoding_encoded_total,
),
MetricSample::new(
"usage_runtime_event_capture_memory_retained_bytes",
"Estimated diagnostic JSON heap retained by budgeted usage events and their clones.",
MetricKind::Gauge,
snapshot.event_capture_memory_retained_bytes as u64,
),
MetricSample::new(
"usage_runtime_event_capture_memory_downgraded_total",
"Usage event diagnostic captures omitted after memory budget exhaustion.",
MetricKind::Counter,
snapshot.event_capture_memory_downgraded_total,
),
MetricSample::new(
"usage_runtime_terminal_submission_limit",
"Maximum concurrent end-to-end terminal usage submissions.",
@@ -3700,6 +3954,12 @@ fn usage_runtime_metric_samples(
MetricKind::Counter,
snapshot.enqueue_retry_failed_total,
),
MetricSample::new(
"usage_runtime_enqueue_retry_permanent_failure_total",
"Usage enqueue retry submissions rejected or queued events terminated because of permanent input errors.",
MetricKind::Counter,
snapshot.enqueue_retry_permanent_failure_total,
),
MetricSample::new(
"usage_runtime_enqueue_retry_closed_or_unavailable_total",
"Total usage events rejected because the local enqueue dispatcher was full, closed, or unavailable.",
@@ -3972,6 +4232,126 @@ mod tests {
assert_eq!(database_bounded_auth_load_limit(Some(64), None), Some(64));
}
#[tokio::test]
async fn http_connection_budget_is_optional_and_shared_by_app_state_clones() {
let state = AppState::new().expect("app state should build");
assert!(state.http_connection_budget.is_none());
let budget = Arc::new(super::HttpConnectionBudget::new(1));
let state = state.with_http_connection_budget(Arc::clone(&budget));
let cloned = state.clone();
assert!(Arc::ptr_eq(
cloned.http_connection_budget.as_ref().unwrap(),
&budget,
));
let connection = budget.try_admit(()).expect("first connection admitted");
assert!(cloned
.http_connection_budget
.as_ref()
.unwrap()
.try_admit(())
.is_err());
let samples = state.collect_metric_samples().await;
for (name, expected) in [
("gateway_http_connections_limit", 1),
("gateway_http_connections_in_flight", 1),
("gateway_http_connections_high_watermark", 1),
("gateway_http_connections_rejected_total", 1),
("gateway_http_connections_accept_errors_total", 0),
] {
assert_eq!(
samples
.iter()
.find(|sample| sample.name == name)
.unwrap()
.value,
expected,
);
}
drop(connection);
assert_eq!(budget.snapshot().in_flight, 0);
}
#[test]
fn http_connection_metrics_export_all_budget_counters() {
let samples = super::http_connection_metric_samples(&super::HttpConnectionBudgetSnapshot {
limit: 4096,
in_flight: 11,
high_watermark: 30,
rejected_total: 7,
accept_errors_total: 3,
});
for (name, kind, value) in [
("gateway_http_connections_limit", MetricKind::Gauge, 4096),
("gateway_http_connections_in_flight", MetricKind::Gauge, 11),
(
"gateway_http_connections_high_watermark",
MetricKind::Gauge,
30,
),
(
"gateway_http_connections_rejected_total",
MetricKind::Counter,
7,
),
(
"gateway_http_connections_accept_errors_total",
MetricKind::Counter,
3,
),
] {
let matching = samples
.iter()
.filter(|sample| sample.name == name)
.collect::<Vec<_>>();
assert_eq!(matching.len(), 1);
assert_eq!((matching[0].kind, matching[0].value), (kind, value));
}
}
#[test]
fn usage_runtime_metrics_export_queue_payload_limit_and_attempt_counters() {
let mut snapshot = crate::usage::UsageRuntimeMetricsSnapshot::default();
snapshot.queue_payload_max_bytes = 1024 * 1024;
snapshot.queue_payload_downgraded_total = 11;
snapshot.queue_payload_rejected_total = 3;
snapshot.enqueue_retry_permanent_failure_total = 2;
let samples = usage_runtime_metric_samples(&snapshot);
for (name, kind, value) in [
(
"usage_runtime_queue_payload_max_bytes",
MetricKind::Gauge,
1024 * 1024,
),
(
"usage_runtime_queue_payload_downgraded_total",
MetricKind::Counter,
11,
),
(
"usage_runtime_queue_payload_rejected_total",
MetricKind::Counter,
3,
),
(
"usage_runtime_enqueue_retry_permanent_failure_total",
MetricKind::Counter,
2,
),
] {
let matching = samples
.iter()
.filter(|sample| sample.name == name)
.collect::<Vec<_>>();
assert_eq!(
matching.len(),
1,
"each payload metric must be emitted once"
);
assert_eq!((matching[0].kind, matching[0].value), (kind, value));
}
}
#[test]
fn usage_runtime_metrics_export_first_byte_batch_counters() {
let mut snapshot = crate::usage::UsageRuntimeMetricsSnapshot::default();
@@ -671,7 +671,7 @@ impl SchedulerRuntimeState for AppState {
&self,
limit: usize,
) -> Result<Vec<StoredRequestCandidate>, GatewayError> {
AppState::read_recent_request_candidates(self, limit).await
AppState::read_recent_runtime_request_candidates(self, limit).await
}
fn provider_key_rpm_reset_at(&self, key_id: &str, now_unix_secs: u64) -> Option<u64> {
@@ -109,6 +109,16 @@ impl AppState {
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn read_recent_runtime_request_candidates(
&self,
limit: usize,
) -> Result<Vec<candidates::StoredRequestCandidate>, GatewayError> {
self.data
.list_recent_runtime_request_candidates(limit)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn upsert_request_candidate(
&self,
mut candidate: candidates::UpsertRequestCandidateRecord,
+11
View File
@@ -19,6 +19,17 @@ use sha2::{Digest, Sha256};
use crate::data::GatewayDataState;
use crate::AppState;
pub use crate::tunnel::build_tunnel_pressure_router;
pub async fn gateway_metric_samples(
state: &AppState,
) -> Result<Vec<aether_runtime::MetricSample>, String> {
if !state.prewarm_metric_snapshot().await {
return Err("gateway harness metric refresh timed out".to_string());
}
Ok(state.metric_samples().await)
}
#[derive(Debug, Clone)]
pub struct OpenAiChatPressureTarget {
pub base_url: String,
@@ -183,7 +183,8 @@ fn gateway_tunnel_protocol_path_is_a_thin_compatibility_facade() {
fn frontdoor_owns_bounded_request_body_buffering() {
let frontdoor = read_workspace_file("crates/aether-gateway/frontdoor/src/body.rs");
assert!(frontdoor.contains("acquire_many_owned"));
assert!(frontdoor.contains("to_bytes(body, body_limit)"));
assert!(frontdoor.contains("body.into_data_stream()"));
assert!(frontdoor.contains("try_reserve_bytes"));
assert!(frontdoor.contains("BodyBufferReservation"));
let gateway = read_workspace_file("apps/aether-gateway/src/handlers/proxy/body_buffer.rs");
+31
View File
@@ -59,6 +59,37 @@ pub use embedded::{
ControlPlaneClient as TunnelControlPlaneClient,
};
#[cfg(feature = "testkit")]
pub fn build_tunnel_pressure_router(
state: TunnelRuntimeState,
instance_id: &str,
secret: &[u8],
) -> Result<axum::Router, String> {
validate_tunnel_relay_auth_secret(secret)?;
let directory = TunnelAttachmentDirectory::from_parts(instance_id, None::<String>, 90);
let state = state.with_relay_auth(
instance_id,
Some(secret.to_vec()),
Arc::clone(&directory.runtime_state),
);
let runtime_router = build_tunnel_runtime_router_with_state(state.clone());
let mut gateway = AppState::new().map_err(|error| error.to_string())?;
gateway.tunnel = EmbeddedTunnelState {
inner: state,
attachment_directory: directory,
relay_auth_secret: Ok(Arc::from(secret)),
};
// HTTP relay requests must pass the same authentication and verified spool
// preparation as the gateway before reaching the embedded local dispatcher.
Ok(axum::Router::new()
.route(
TUNNEL_RELAY_PATH_PATTERN,
axum::routing::post(relay_request),
)
.with_state(gateway)
.fallback_service(runtime_router))
}
const DEFAULT_ATTACHMENT_TTL_SECS: u64 = 90;
const TUNNEL_ATTACHMENT_KEY_PREFIX: &str = "tunnel.attachments.";
const TUNNEL_ATTACHMENT_REDIS_KEY_PREFIX: &str = "tunnel:attachments:";