mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-04 00:17:45 +08:00
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:
@@ -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],
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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,
|
||||
};
|
||||
|
||||
@@ -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
@@ -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();
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
@@ -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),
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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");
|
||||
|
||||
@@ -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:";
|
||||
|
||||
Reference in New Issue
Block a user