Improve gateway scheduling and runtime admission

This commit is contained in:
elky
2026-06-24 01:53:45 +08:00
parent cf0af8fa1e
commit d336d1a7fa
87 changed files with 9671 additions and 804 deletions
@@ -8,6 +8,12 @@ use crate::scheduler::affinity::SCHEDULER_AFFINITY_TTL;
const PLANNER_SCHEDULER_AFFINITY_MAX_ENTRIES: usize = 10_000;
pub(crate) fn has_explicit_session_affinity(
client_session_affinity: Option<&ClientSessionAffinity>,
) -> bool {
client_session_affinity.is_some_and(ClientSessionAffinity::has_session_key)
}
pub(crate) fn read_cached_scheduler_affinity_target(
state: PlannerAppState<'_>,
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
@@ -15,6 +21,9 @@ pub(crate) fn read_cached_scheduler_affinity_target(
client_api_format: &str,
requested_model: Option<&str>,
) -> Option<SchedulerAffinityTarget> {
if !has_explicit_session_affinity(client_session_affinity) {
return None;
}
let requested_model = requested_model
.map(str::trim)
.filter(|value| !value.is_empty())?;
@@ -41,6 +50,9 @@ pub(crate) fn remember_scheduler_affinity_for_candidate(
requested_model: &str,
candidate: &SchedulerMinimalCandidateSelectionCandidate,
) {
if !has_explicit_session_affinity(client_session_affinity) {
return;
}
remember_scheduler_affinity_for_candidate_at_epoch(
state,
auth_snapshot,
@@ -61,6 +73,9 @@ pub(crate) fn remember_scheduler_affinity_for_candidate_at_epoch(
candidate: &SchedulerMinimalCandidateSelectionCandidate,
expected_epoch: Option<u64>,
) {
if !has_explicit_session_affinity(client_session_affinity) {
return;
}
let Some(api_key_id) = auth_snapshot
.map(|snapshot| snapshot.api_key_id.trim())
.filter(|value| !value.is_empty())
@@ -38,12 +38,19 @@ use crate::ai_serving::planner::pool_scheduler::PoolKeyCursor;
use crate::ai_serving::planner::runtime_miss::record_local_runtime_candidate_skip_reason;
use crate::ai_serving::planner::CandidateFailureDiagnostic;
use crate::ai_serving::{GatewayAuthApiKeySnapshot, PlannerAppState};
use crate::cache::{
candidate_page_cache_stale_ttl, candidate_page_cache_ttl_from_env,
record_candidate_page_resolve_cache_follower_wait, record_candidate_page_resolve_cache_hit,
record_candidate_page_resolve_cache_load, record_candidate_page_resolve_cache_miss,
CacheLoadObserver, CandidateResolvedPageCacheKey, CandidateResolvedPageSnapshot,
};
use crate::clock::current_unix_ms;
use crate::dispatch::refs::dispatch_ref_for_local_candidate;
use crate::handlers::shared::provider_pool::admin_provider_pool_config_from_config_value;
use crate::orchestration::{local_attempt_slot_count, ExecutionAttemptIdentity};
use crate::scheduler::candidate::is_auth_api_key_concurrency_limit_skip_reason;
use crate::scheduler::config::SchedulerSchedulingMode;
use crate::stage_metrics::observe_gateway_stage_ms;
use crate::{AppState, GatewayError};
const POOL_KEY_RETRY_INDEX_STRIDE: u32 = 100;
@@ -101,16 +108,20 @@ impl<'a> LocalExecutionCandidateAttemptSource<'a> {
Self { items }
}
pub(crate) async fn next_attempt(&mut self) -> Option<LocalExecutionCandidateAttempt> {
pub(crate) async fn next_attempt(
&mut self,
) -> Result<Option<LocalExecutionCandidateAttempt>, GatewayError> {
loop {
let front = self.items.front_mut()?;
let Some(front) = self.items.front_mut() else {
return Ok(None);
};
match front {
LocalExecutionCandidateAttemptSourceItem::Static { attempts } => {
if let Some(attempt) = next_attempt_from_dispatch_sequence(attempts) {
if dispatch_sequence_exhausted(attempts) {
self.items.pop_front();
}
return Some(attempt);
return Ok(Some(attempt));
}
self.items.pop_front();
}
@@ -121,7 +132,7 @@ impl<'a> LocalExecutionCandidateAttemptSource<'a> {
pool_exhaustion_persistence,
} => {
if let Some(attempt) = next_attempt_from_dispatch_sequence(pending_attempts) {
return Some(attempt);
return Ok(Some(attempt));
}
let Some(candidate) = cursor.next_key().await else {
if let Some(skipped) = cursor.exhausted_group_skipped_candidate() {
@@ -146,11 +157,11 @@ impl<'a> LocalExecutionCandidateAttemptSource<'a> {
);
}
LocalExecutionCandidateAttemptSourceItem::RequestedModelPage { cursor } => {
let Some(attempt) = cursor.next_attempt().await else {
let Some(attempt) = cursor.next_attempt().await? else {
self.items.pop_front();
continue;
};
return Some(attempt);
return Ok(Some(attempt));
}
}
}
@@ -726,8 +737,10 @@ where
auth_snapshot,
routing_policy,
client_session_affinity,
request_auth_channel,
use_api_format_alias_match,
key_mode,
Some(trace_id),
)
.await;
let mut cursor = RequestedModelAttemptPageCursor {
@@ -755,11 +768,14 @@ where
remembered_affinity: false,
scheduler_cache_affinity_enabled,
auth_api_key_concurrency_wait_deadline: None,
deferred_error: None,
};
cursor.load_next_page().await;
if let Err(error) = cursor.load_next_page().await {
cursor.deferred_error = Some(error);
}
let candidate_count = cursor.candidate_count;
let mut items = VecDeque::new();
if !cursor.pending_items.is_empty() {
if !cursor.pending_items.is_empty() || cursor.deferred_error.is_some() {
items.push_back(
LocalExecutionCandidateAttemptSourceItem::RequestedModelPage {
cursor: Box::new(cursor),
@@ -797,34 +813,58 @@ struct RequestedModelAttemptPageCursor<'a> {
remembered_affinity: bool,
scheduler_cache_affinity_enabled: bool,
auth_api_key_concurrency_wait_deadline: Option<Instant>,
deferred_error: Option<GatewayError>,
}
impl<'a> RequestedModelAttemptPageCursor<'a> {
async fn next_attempt(&mut self) -> Option<LocalExecutionCandidateAttempt> {
async fn next_attempt(
&mut self,
) -> Result<Option<LocalExecutionCandidateAttempt>, GatewayError> {
if let Some(error) = self.deferred_error.take() {
return Err(error);
}
loop {
if let Some(attempt) = pop_attempt_from_items(&mut self.pending_items).await {
return Some(attempt);
return Ok(Some(attempt));
}
if !self.load_next_page().await {
return None;
if !self.load_next_page().await? {
return Ok(None);
}
}
}
async fn load_next_page(&mut self) -> bool {
async fn load_next_page(&mut self) -> Result<bool, GatewayError> {
loop {
let page_started_at = std::time::Instant::now();
let page = match self.page_cursor.next_page().await {
Ok(Some(page)) => page,
Ok(None) => return false,
Ok(None) => {
observe_gateway_stage_ms(
"candidate_page_load",
page_started_at.elapsed().as_millis() as u64,
);
return Ok(false);
}
Err(error) => {
observe_gateway_stage_ms(
"candidate_page_load",
page_started_at.elapsed().as_millis() as u64,
);
if matches!(error, GatewayError::AdmissionTimeout { .. }) {
return Err(error);
}
warn!(
trace_id = %self.trace_id,
error = ?error,
"gateway lazy requested-model candidate page read failed"
);
return false;
return Ok(false);
}
};
observe_gateway_stage_ms(
"candidate_page_load",
page_started_at.elapsed().as_millis() as u64,
);
if page_is_exact_auth_api_key_concurrency_limited(&page) {
if self.wait_for_auth_api_key_concurrency_retry().await {
@@ -832,24 +872,16 @@ impl<'a> RequestedModelAttemptPageCursor<'a> {
}
self.persist_final_auth_api_key_concurrency_skips(page.skipped_candidates)
.await;
return false;
return Ok(false);
}
let resolve_started_at = std::time::Instant::now();
let (candidates, resolved_skipped) =
resolve_and_rank_logical_local_execution_candidates(
self.state,
page.candidates,
&self.client_api_format,
Some(&self.requested_model),
Some(&self.auth_snapshot),
self.client_session_affinity.as_ref(),
self.required_capabilities.as_ref(),
self.routing_policy.as_ref(),
self.sticky_session_token.as_deref(),
self.request_auth_channel.as_deref(),
self.resolution_mode,
)
.await;
resolve_priority_candidate_page_with_cache(self, page.candidates).await;
observe_gateway_stage_ms(
"candidate_page_resolve",
resolve_started_at.elapsed().as_millis() as u64,
);
let skipped_candidates = page
.skipped_candidates
.into_iter()
@@ -899,7 +931,7 @@ impl<'a> RequestedModelAttemptPageCursor<'a> {
.saturating_add(u32::try_from(skipped_candidate_count).unwrap_or(u32::MAX));
if !items.is_empty() {
self.pending_items = items;
return true;
return Ok(true);
}
let skipped_starting_candidate_index = next_candidate_index;
let skipped_persistence = LocalSkippedCandidatePersistenceContext {
@@ -1080,6 +1112,120 @@ pub(crate) fn remember_first_local_candidate_affinity(
);
}
async fn resolve_priority_candidate_page_with_cache(
cursor: &RequestedModelAttemptPageCursor<'_>,
page_candidates: Vec<SchedulerMinimalCandidateSelectionCandidate>,
) -> (
Vec<EligibleLocalExecutionCandidate>,
Vec<SkippedLocalExecutionCandidate>,
) {
if !should_cache_resolved_candidate_page(cursor) {
return resolve_and_rank_logical_local_execution_candidates(
cursor.state,
page_candidates,
&cursor.client_api_format,
Some(&cursor.requested_model),
Some(&cursor.auth_snapshot),
cursor.client_session_affinity.as_ref(),
cursor.required_capabilities.as_ref(),
cursor.routing_policy.as_ref(),
cursor.sticky_session_token.as_deref(),
cursor.request_auth_channel.as_deref(),
cursor.resolution_mode,
)
.await;
}
let key = CandidateResolvedPageCacheKey::new(
&cursor.requested_model,
&cursor.client_api_format,
true,
&cursor.auth_snapshot,
cursor.required_capabilities.as_ref(),
cursor.routing_policy.as_ref(),
cursor.request_auth_channel.as_deref(),
cursor.state.app().scheduler_affinity_epoch(),
cursor.page_cursor.resolved_page_cache_preselection_mode(),
cursor
.page_cursor
.resolved_page_cache_use_api_format_alias_match(),
cursor.client_session_affinity.as_ref(),
cursor.resolution_mode,
);
let page_candidates_for_fallback = page_candidates.clone();
let page_candidates_for_load = page_candidates;
let cache = cursor.state.app().candidate_resolved_page_cache.clone();
let ttl = candidate_page_cache_ttl_from_env();
let stale_ttl = candidate_page_cache_stale_ttl(ttl);
let cached = cache
.get_or_load_once_stale_while_refreshing(
key,
ttl,
stale_ttl,
|| async move {
let (candidates, resolved_skipped) =
resolve_and_rank_logical_local_execution_candidates(
cursor.state,
page_candidates_for_load,
&cursor.client_api_format,
Some(&cursor.requested_model),
Some(&cursor.auth_snapshot),
cursor.client_session_affinity.as_ref(),
cursor.required_capabilities.as_ref(),
cursor.routing_policy.as_ref(),
cursor.sticky_session_token.as_deref(),
cursor.request_auth_channel.as_deref(),
cursor.resolution_mode,
)
.await;
Ok::<_, GatewayError>(Some(Arc::new(CandidateResolvedPageSnapshot {
candidates,
resolved_skipped,
})))
},
CacheLoadObserver::new()
.on_hit(record_candidate_page_resolve_cache_hit)
.on_miss(record_candidate_page_resolve_cache_miss)
.on_load(record_candidate_page_resolve_cache_load)
.on_follower_wait(record_candidate_page_resolve_cache_follower_wait),
)
.await
.unwrap_or(None);
match cached {
Some(snapshot) => (
snapshot.candidates.clone(),
snapshot.resolved_skipped.clone(),
),
None => {
if page_candidates_for_fallback.is_empty() {
return (Vec::new(), Vec::new());
}
resolve_and_rank_logical_local_execution_candidates(
cursor.state,
page_candidates_for_fallback,
&cursor.client_api_format,
Some(&cursor.requested_model),
Some(&cursor.auth_snapshot),
cursor.client_session_affinity.as_ref(),
cursor.required_capabilities.as_ref(),
cursor.routing_policy.as_ref(),
cursor.sticky_session_token.as_deref(),
cursor.request_auth_channel.as_deref(),
cursor.resolution_mode,
)
.await
}
}
}
fn should_cache_resolved_candidate_page(cursor: &RequestedModelAttemptPageCursor<'_>) -> bool {
cursor.sticky_session_token.is_none()
&& cursor
.page_cursor
.should_cache_current_priority_resolved_page()
}
fn should_persist_available_local_candidate(eligible: &EligibleLocalExecutionCandidate) -> bool {
ai_should_persist_available_candidate_for_pool_key(eligible.orchestration.pool_key_index)
}
@@ -1937,7 +2083,8 @@ mod tests {
GatewayDataState::with_request_candidate_repository_for_tests(Arc::clone(
&repository,
)),
);
)
.without_request_candidate_queue_for_tests();
let attempts = persist_available_local_execution_candidates(
PlannerAppState::new(&app),
@@ -2025,6 +2172,145 @@ mod tests {
.is_none());
}
#[tokio::test]
async fn resolved_candidate_page_cache_requires_fixed_order_or_explicit_affinity() {
let app = AppState::new().expect("state should build");
let auth_snapshot = sample_auth_snapshot();
let mut page_cursor = LocalCandidatePreselectionPageCursor::new(
PlannerAppState::new(&app),
"openai:chat",
"gpt-5",
true,
None,
&auth_snapshot,
None,
None,
None,
false,
LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModel,
Some("trace-no-session-affinity"),
)
.await;
page_cursor.mark_priority_page_emitted_for_tests();
let cursor = RequestedModelAttemptPageCursor {
state: PlannerAppState::new(&app),
trace_id: "trace-no-session-affinity".to_string(),
client_api_format: "openai:chat".to_string(),
requested_model: "gpt-5".to_string(),
auth_snapshot: auth_snapshot.clone(),
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 skipped",
record_runtime_miss_diagnostic: false,
resolution_mode: LocalCandidateResolutionMode::Standard,
decorate_skipped_candidate: Arc::new(identity_skipped_candidate),
page_cursor,
pending_items: VecDeque::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,
};
assert!(!should_cache_resolved_candidate_page(&cursor));
let sticky_cursor = RequestedModelAttemptPageCursor {
sticky_session_token: Some("sticky-token".to_string()),
..cursor
};
assert!(!should_cache_resolved_candidate_page(&sticky_cursor));
let mut page_cursor = LocalCandidatePreselectionPageCursor::new(
PlannerAppState::new(&app),
"openai:chat",
"gpt-5",
true,
None,
&auth_snapshot,
None,
Some(&ClientSessionAffinity::from_session_key("chat-session-1")),
None,
false,
LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModel,
Some("trace-session-affinity"),
)
.await;
page_cursor.mark_priority_page_emitted_for_tests();
let cursor = RequestedModelAttemptPageCursor {
client_session_affinity: Some(ClientSessionAffinity::from_session_key(
"chat-session-1",
)),
page_cursor,
sticky_session_token: None,
..sticky_cursor
};
assert!(should_cache_resolved_candidate_page(&cursor));
let fixed_order_app = AppState::new()
.expect("state should build")
.with_data_state_for_tests(
GatewayDataState::disabled().with_system_config_values_for_tests([(
"scheduling_mode".to_string(),
json!("fixed_order"),
)]),
);
let mut page_cursor = LocalCandidatePreselectionPageCursor::new(
PlannerAppState::new(&fixed_order_app),
"openai:chat",
"gpt-5",
true,
None,
&auth_snapshot,
None,
None,
None,
false,
LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModel,
Some("trace-fixed-order"),
)
.await;
page_cursor.mark_priority_page_emitted_for_tests();
let cursor = RequestedModelAttemptPageCursor {
state: PlannerAppState::new(&fixed_order_app),
trace_id: "trace-fixed-order".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 skipped",
record_runtime_miss_diagnostic: false,
resolution_mode: LocalCandidateResolutionMode::Standard,
decorate_skipped_candidate: Arc::new(identity_skipped_candidate),
page_cursor,
pending_items: VecDeque::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,
};
assert!(should_cache_resolved_candidate_page(&cursor));
}
#[tokio::test]
async fn logical_materialization_does_not_persist_pool_group_representative() {
let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default());
@@ -2042,7 +2328,8 @@ mod tests {
Arc::clone(&request_candidate_repository),
"test-encryption-key",
),
);
)
.without_request_candidate_queue_for_tests();
let mut pool_group = sample_eligible("pool-group", None);
pool_group.kind = LocalExecutionCandidateKind::PoolGroup;
pool_group.transport = sample_transport(
@@ -2140,7 +2427,8 @@ mod tests {
GatewayDataState::with_request_candidate_repository_for_tests(Arc::clone(
&repository,
)),
);
)
.without_request_candidate_queue_for_tests();
let mut eligible = sample_eligible("ranked-key", None);
eligible.ranking = Some(SchedulerRankingOutcome {
original_index: 1,
@@ -2214,12 +2502,17 @@ mod tests {
let first = source
.next_attempt()
.await
.expect("first attempt read should succeed")
.expect("first attempt should be available");
assert_eq!(first.eligible.candidate.key_id, "normal-key");
let remaining = source.drain_static_attempts();
assert!(remaining.is_empty());
assert!(source.next_attempt().await.is_none());
assert!(source
.next_attempt()
.await
.expect("remaining attempt read should succeed")
.is_none());
}
#[tokio::test]
@@ -2239,7 +2532,8 @@ mod tests {
Arc::clone(&request_candidate_repository),
"test-encryption-key",
),
);
)
.without_request_candidate_queue_for_tests();
let mut pool_group = sample_eligible("pool-group", None);
pool_group.kind = LocalExecutionCandidateKind::PoolGroup;
pool_group.transport = sample_transport(
@@ -2275,7 +2569,11 @@ mod tests {
}]),
};
assert!(source.next_attempt().await.is_none());
assert!(source
.next_attempt()
.await
.expect("pool attempt read should succeed")
.is_none());
let stored = app
.read_request_candidates_by_request_id("trace-dynamic-pool")
@@ -2308,7 +2606,8 @@ mod tests {
GatewayDataState::with_request_candidate_repository_for_tests(Arc::clone(
&repository,
)),
);
)
.without_request_candidate_queue_for_tests();
persist_skipped_local_execution_candidates(
&app,
@@ -1,5 +1,3 @@
use std::collections::BTreeMap;
use aether_ai_serving::{
ai_ranking_context, build_ai_rankable_candidate, run_ai_candidate_ranking,
AiCandidateRankingPort, AiRankableCandidateParts, AiRankingContextConfig,
@@ -7,6 +5,7 @@ use aether_ai_serving::{
};
use aether_routing_core::{ResolvedRoutingPolicy, RoutingSchedulingMode, RoutingSetPriorityMode};
use async_trait::async_trait;
use tokio::sync::Mutex;
use tracing::warn;
use crate::ai_serving::{GatewayAuthApiKeySnapshot, PlannerAppState};
@@ -24,7 +23,7 @@ use aether_scheduler_core::{
use super::candidate_affinity_cache::read_cached_scheduler_affinity_target;
use super::candidate_resolution::{EligibleLocalExecutionCandidate, LocalExecutionCandidateKind};
use super::candidate_transport_ranking_facts::{
resolve_cached_transport_ranking_facts, CandidateTransportRankingFacts,
resolve_cached_transport_ranking_facts, CandidateTransportRankingFactsCache,
};
struct GatewayLocalCandidateRankingPort<'a> {
@@ -35,6 +34,7 @@ struct GatewayLocalCandidateRankingPort<'a> {
required_capabilities: Option<&'a serde_json::Value>,
ordering_config: SchedulerOrderingConfig,
routing_policy: Option<&'a ResolvedRoutingPolicy>,
transport_ranking_facts_cache: Mutex<CandidateTransportRankingFactsCache>,
}
#[async_trait]
@@ -84,13 +84,17 @@ impl AiCandidateRankingPort for GatewayLocalCandidateRankingPort<'_> {
normalized_client_api_format: &str,
cached_affinity_match: bool,
) -> Result<SchedulerRankableCandidate, Self::Error> {
let ranking_facts = resolve_transport_ranking_facts_for_candidate(
self.state,
&candidate.candidate,
candidate.transport.as_ref(),
self.ordering_config,
)
.await;
let ranking_facts = {
let mut cache = self.transport_ranking_facts_cache.lock().await;
resolve_cached_transport_ranking_facts(
self.state,
&mut cache,
&candidate.candidate,
candidate.transport.as_ref(),
self.ordering_config,
)
.await
};
let routing_overlaid_candidate =
routing_overlaid_candidate(self.routing_policy, candidate.kind, &candidate.candidate);
Ok(build_ai_rankable_candidate(AiRankableCandidateParts {
@@ -137,6 +141,7 @@ pub(crate) async fn rank_eligible_local_execution_candidates(
required_capabilities,
ordering_config,
routing_policy,
transport_ranking_facts_cache: Mutex::new(CandidateTransportRankingFactsCache::default()),
};
match run_ai_candidate_ranking(&port, candidates, normalized_client_api_format).await {
@@ -145,23 +150,6 @@ pub(crate) async fn rank_eligible_local_execution_candidates(
}
}
async fn resolve_transport_ranking_facts_for_candidate(
state: PlannerAppState<'_>,
candidate: &SchedulerMinimalCandidateSelectionCandidate,
transport: &crate::ai_serving::GatewayProviderTransportSnapshot,
ordering_config: SchedulerOrderingConfig,
) -> CandidateTransportRankingFacts {
let mut ordering_cache = BTreeMap::new();
resolve_cached_transport_ranking_facts(
state,
&mut ordering_cache,
candidate,
transport,
ordering_config,
)
.await
}
fn cached_affinity_matches_local_execution_scope(
eligible: &EligibleLocalExecutionCandidate,
target: &SchedulerAffinityTarget,
@@ -284,7 +272,9 @@ mod tests {
use serde_json::json;
use super::super::candidate_affinity_cache::remember_scheduler_affinity_for_candidate;
use super::super::candidate_transport_ranking_facts::resolve_cached_candidate_transport_ranking_facts;
use super::super::candidate_transport_ranking_facts::{
resolve_cached_candidate_transport_ranking_facts, CandidateTransportRankingFactsCache,
};
use super::{PlannerAppState, SchedulerMinimalCandidateSelectionCandidate};
use crate::ai_serving::planner::candidate_resolution::{
resolve_and_rank_local_execution_candidates,
@@ -306,7 +296,7 @@ mod tests {
let ordering_config = super::read_scheduler_ordering_config_or_default(state).await;
let mut candidates = candidates;
let mut rankables = Vec::with_capacity(candidates.len());
let mut ordering_cache = BTreeMap::new();
let mut ordering_cache = CandidateTransportRankingFactsCache::default();
for (original_index, candidate) in candidates.iter().enumerate() {
let ranking_facts = resolve_cached_candidate_transport_ranking_facts(
@@ -1540,6 +1530,7 @@ mod tests {
.expect("state should build")
.with_data_state_for_tests(data_state);
let auth_snapshot = sample_auth_snapshot();
let client_session_affinity = ClientSessionAffinity::from_session_key("session-1");
let cached_candidate = sample_priority_candidate(
"provider-cached",
"endpoint-cached",
@@ -1551,7 +1542,7 @@ mod tests {
remember_scheduler_affinity_for_candidate(
PlannerAppState::new(&state),
Some(&auth_snapshot),
None,
Some(&client_session_affinity),
"openai:chat",
"gpt-4.1",
&cached_candidate,
@@ -1573,7 +1564,7 @@ mod tests {
"openai:chat",
"gpt-4.1",
Some(&auth_snapshot),
None,
Some(&client_session_affinity),
None,
None,
None,
@@ -1714,6 +1705,7 @@ mod tests {
.expect("state should build")
.with_data_state_for_tests(data_state);
let auth_snapshot = sample_auth_snapshot();
let client_session_affinity = ClientSessionAffinity::from_session_key("session-1");
let cached_cross_format = sample_priority_candidate(
"provider-shared",
"endpoint-openai",
@@ -1725,7 +1717,7 @@ mod tests {
remember_scheduler_affinity_for_candidate(
PlannerAppState::new(&state),
Some(&auth_snapshot),
None,
Some(&client_session_affinity),
"claude:messages",
"gpt-4.1",
&cached_cross_format,
@@ -1747,7 +1739,7 @@ mod tests {
"claude:messages",
"gpt-4.1",
Some(&auth_snapshot),
None,
Some(&client_session_affinity),
None,
None,
None,
@@ -1915,6 +1907,7 @@ mod tests {
.expect("state should build")
.with_data_state_for_tests(data_state);
let auth_snapshot = sample_auth_snapshot();
let client_session_affinity = ClientSessionAffinity::from_session_key("session-1");
let cached_candidate = sample_priority_candidate(
"provider-pool",
"endpoint-pool",
@@ -1926,7 +1919,7 @@ mod tests {
remember_scheduler_affinity_for_candidate(
PlannerAppState::new(&state),
Some(&auth_snapshot),
None,
Some(&client_session_affinity),
"openai:chat",
"gpt-4.1",
&cached_candidate,
@@ -1948,7 +1941,7 @@ mod tests {
"openai:chat",
Some("gpt-4.1"),
Some(&auth_snapshot),
None,
Some(&client_session_affinity),
None,
None,
None,
@@ -2008,6 +2001,7 @@ mod tests {
.expect("state should build")
.with_data_state_for_tests(data_state);
let auth_snapshot = sample_auth_snapshot();
let client_session_affinity = ClientSessionAffinity::from_session_key("session-1");
let cached_candidate = sample_priority_candidate(
"provider-pool",
"endpoint-pool",
@@ -2019,7 +2013,7 @@ mod tests {
remember_scheduler_affinity_for_candidate(
PlannerAppState::new(&state),
Some(&auth_snapshot),
None,
Some(&client_session_affinity),
"openai:chat",
"gpt-4.1",
&cached_candidate,
@@ -2041,7 +2035,7 @@ mod tests {
"openai:chat",
Some("gpt-4.1"),
Some(&auth_snapshot),
None,
Some(&client_session_affinity),
None,
None,
None,
@@ -2064,7 +2058,7 @@ mod tests {
}
#[tokio::test]
async fn remembers_scheduler_affinity_for_candidate_using_requested_model_key() {
async fn ignores_scheduler_affinity_without_client_session_scope() {
let state = AppState::new().expect("state should build");
let auth_snapshot = sample_auth_snapshot();
let candidate = sample_candidate("endpoint-1", "key-1");
@@ -2078,15 +2072,12 @@ mod tests {
&candidate,
);
let remembered = state
assert!(state
.read_scheduler_affinity_target(
"scheduler_affinity:api-key-1:openai:chat:gpt-5",
SCHEDULER_AFFINITY_TTL,
)
.expect("affinity target should be cached");
assert_eq!(remembered.provider_id, "provider-1");
assert_eq!(remembered.endpoint_id, "endpoint-1");
assert_eq!(remembered.key_id, "key-1");
.is_none());
}
#[tokio::test]
@@ -7,6 +7,7 @@ use aether_ai_serving::{
use aether_routing_core::ResolvedRoutingPolicy;
use async_trait::async_trait;
use std::convert::Infallible;
use std::time::Instant;
use tracing::warn;
use aether_scheduler_core::{
@@ -20,6 +21,7 @@ use crate::ai_serving::{
PlannerAppState,
};
use crate::orchestration::LocalExecutionCandidateMetadata;
use crate::stage_metrics::observe_gateway_stage_ms;
use super::candidate_ranking::rank_eligible_local_execution_candidates;
@@ -68,7 +70,7 @@ struct GatewayLocalCandidateResolutionPort<'a> {
#[async_trait]
impl AiCandidateResolutionPort for GatewayLocalCandidateResolutionPort<'_> {
type Candidate = SchedulerMinimalCandidateSelectionCandidate;
type Transport = GatewayProviderTransportSnapshot;
type Transport = Arc<GatewayProviderTransportSnapshot>;
type Eligible = EligibleLocalExecutionCandidate;
type Skipped = SkippedLocalExecutionCandidate;
type Error = Infallible;
@@ -77,7 +79,12 @@ impl AiCandidateResolutionPort for GatewayLocalCandidateResolutionPort<'_> {
&self,
candidate: &Self::Candidate,
) -> Result<Option<Self::Transport>, Self::Error> {
Ok(read_candidate_transport_snapshot(self.state, candidate).await)
let started_at = Instant::now();
let transport = read_candidate_transport_snapshot_arc(self.state, candidate).await;
let elapsed_ms = started_at.elapsed().as_millis() as u64;
observe_gateway_stage_ms("candidate_transport_snapshot", elapsed_ms);
observe_gateway_stage_ms("candidate_resolution_transport_read", elapsed_ms);
Ok(transport)
}
fn build_missing_transport_skipped_candidate(
@@ -139,7 +146,7 @@ impl AiCandidateResolutionPort for GatewayLocalCandidateResolutionPort<'_> {
SkippedLocalExecutionCandidate {
candidate,
skip_reason,
transport: Some(Arc::new(transport)),
transport: Some(transport),
ranking: None,
extra_data: None,
}
@@ -159,7 +166,7 @@ impl AiCandidateResolutionPort for GatewayLocalCandidateResolutionPort<'_> {
EligibleLocalExecutionCandidate {
kind,
candidate,
transport: Arc::new(transport),
transport,
provider_api_format,
orchestration: LocalExecutionCandidateMetadata::default(),
ranking: None,
@@ -171,7 +178,8 @@ impl AiCandidateResolutionPort for GatewayLocalCandidateResolutionPort<'_> {
candidates: Vec<Self::Eligible>,
normalized_client_api_format: &str,
) -> Result<Vec<Self::Eligible>, Self::Error> {
Ok(rank_eligible_local_execution_candidates(
let started_at = Instant::now();
let ranked = rank_eligible_local_execution_candidates(
self.state,
candidates,
normalized_client_api_format,
@@ -181,7 +189,12 @@ impl AiCandidateResolutionPort for GatewayLocalCandidateResolutionPort<'_> {
self.required_capabilities,
self.routing_policy,
)
.await)
.await;
observe_gateway_stage_ms(
"candidate_resolution_rank",
started_at.elapsed().as_millis() as u64,
);
Ok(ranked)
}
async fn apply_pool_scheduler(
@@ -358,8 +371,13 @@ async fn resolve_and_rank_local_execution_candidates_with_pool_expansion(
expand_pool_groups,
};
let started_at = Instant::now();
match run_ai_candidate_resolution(&port, candidates, request).await {
Ok(mut outcome) => {
observe_gateway_stage_ms(
"candidate_resolution_core",
started_at.elapsed().as_millis() as u64,
);
for candidate in &mut outcome.eligible_candidates {
candidate.orchestration.scheduler_affinity_epoch = Some(scheduler_affinity_epoch);
}
@@ -506,8 +524,17 @@ pub(crate) async fn read_candidate_transport_snapshot(
state: PlannerAppState<'_>,
candidate: &SchedulerMinimalCandidateSelectionCandidate,
) -> Option<GatewayProviderTransportSnapshot> {
read_candidate_transport_snapshot_arc(state, candidate)
.await
.map(|transport| (*transport).clone())
}
pub(crate) async fn read_candidate_transport_snapshot_arc(
state: PlannerAppState<'_>,
candidate: &SchedulerMinimalCandidateSelectionCandidate,
) -> Option<Arc<GatewayProviderTransportSnapshot>> {
match state
.read_provider_transport_snapshot(
.read_provider_transport_snapshot_arc(
&candidate.provider_id,
&candidate.endpoint_id,
&candidate.key_id,
@@ -3,6 +3,7 @@ use aether_ai_serving::{
};
use aether_data_contracts::repository::candidate_selection::StoredMinimalCandidateSelectionRow;
use aether_routing_core::ResolvedRoutingPolicy;
use aether_runtime::ConcurrencyPermit;
use aether_scheduler_core::{
enumerate_minimal_candidate_selection_with_model_directives, normalize_api_format,
resolve_requested_global_model_name_with_model_directives,
@@ -12,16 +13,17 @@ use aether_scheduler_core::{
use async_trait::async_trait;
use std::collections::{BTreeMap, BTreeSet, VecDeque};
use crate::ai_serving::planner::candidate_affinity_cache::has_explicit_session_affinity;
use crate::ai_serving::planner::candidate_resolution::SkippedLocalExecutionCandidate;
use crate::ai_serving::{GatewayAuthApiKeySnapshot, PlannerAppState};
use crate::clock::current_unix_secs;
use crate::clock::request_distribution_seed;
use crate::data::candidate_selection::{
read_requested_model_rows_fast_path_page, requested_model_candidate_names,
MinimalCandidateSelectionRowSource, REQUESTED_MODEL_CANDIDATE_PAGE_SIZE,
REQUESTED_MODEL_MAX_SCANNED_ROWS,
};
use crate::scheduler::candidate::SchedulerSkippedCandidate;
use crate::scheduler::config::SchedulerOrderingConfig;
use crate::scheduler::config::{SchedulerOrderingConfig, SchedulerSchedulingMode};
use crate::GatewayError;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
@@ -30,6 +32,15 @@ pub(crate) enum LocalCandidatePreselectionKeyMode {
ProviderEndpointKeyModelAndApiFormat,
}
impl LocalCandidatePreselectionKeyMode {
pub(crate) fn cache_key_name(self) -> &'static str {
match self {
Self::ProviderEndpointKeyModel => "provider_endpoint_key_model",
Self::ProviderEndpointKeyModelAndApiFormat => "provider_endpoint_key_model_api_format",
}
}
}
struct GatewayLocalCandidatePreselectionPort<'a> {
state: PlannerAppState<'a>,
client_api_format: &'a str,
@@ -43,6 +54,7 @@ struct GatewayLocalCandidatePreselectionPort<'a> {
key_mode: LocalCandidatePreselectionKeyMode,
candidate_api_formats: Vec<String>,
model_directive_enabled_api_formats: BTreeSet<String>,
ranking_seed: u64,
}
#[async_trait]
@@ -81,7 +93,7 @@ impl AiCandidatePreselectionPort for GatewayLocalCandidatePreselectionPort<'_> {
self.required_capabilities,
auth_snapshot,
self.client_session_affinity,
current_unix_secs(),
self.ranking_seed,
)
.await?;
@@ -227,6 +239,7 @@ pub(crate) async fn preselect_local_execution_candidates_for_api_formats_with_se
key_mode,
candidate_api_formats,
model_directive_enabled_api_formats,
ranking_seed: request_distribution_seed(),
};
run_ai_candidate_preselection(&port).await
@@ -234,6 +247,7 @@ pub(crate) async fn preselect_local_execution_candidates_for_api_formats_with_se
pub(crate) struct LocalCandidatePreselectionPageCursor<'a> {
state: PlannerAppState<'a>,
trace_id: String,
client_api_format: String,
requested_model: String,
require_streaming: bool,
@@ -241,11 +255,13 @@ pub(crate) struct LocalCandidatePreselectionPageCursor<'a> {
auth_snapshot: GatewayAuthApiKeySnapshot,
routing_policy: Option<ResolvedRoutingPolicy>,
client_session_affinity: Option<ClientSessionAffinity>,
request_auth_channel: Option<String>,
use_api_format_alias_match: bool,
key_mode: LocalCandidatePreselectionKeyMode,
candidate_api_formats: Vec<String>,
model_directive_enabled_api_formats: BTreeSet<String>,
ordering_config: SchedulerOrderingConfig,
ranking_seed: u64,
priority_page_emitted: bool,
deferred_pages_by_format: BTreeMap<
String,
@@ -276,8 +292,10 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
auth_snapshot: &GatewayAuthApiKeySnapshot,
routing_policy: Option<&ResolvedRoutingPolicy>,
client_session_affinity: Option<&ClientSessionAffinity>,
request_auth_channel: Option<&str>,
use_api_format_alias_match: bool,
key_mode: LocalCandidatePreselectionKeyMode,
trace_id: Option<&str>,
) -> Self {
let candidate_api_formats =
crate::ai_serving::request_candidate_api_formats(client_api_format, require_streaming)
@@ -307,6 +325,7 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
Self {
state,
trace_id: trace_id.unwrap_or_default().to_string(),
client_api_format: client_api_format.to_string(),
requested_model: requested_model.to_string(),
require_streaming,
@@ -314,11 +333,13 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
auth_snapshot: auth_snapshot.clone(),
routing_policy: routing_policy.cloned(),
client_session_affinity: client_session_affinity.cloned(),
request_auth_channel: request_auth_channel.map(str::to_string),
use_api_format_alias_match,
key_mode,
candidate_api_formats,
model_directive_enabled_api_formats,
ordering_config,
ranking_seed: request_distribution_seed(),
priority_page_emitted: false,
deferred_pages_by_format: BTreeMap::new(),
format_index: 0,
@@ -344,7 +365,7 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
> {
if !self.priority_page_emitted {
self.priority_page_emitted = true;
let priority_page = self.next_priority_page().await?;
let priority_page = self.cached_next_priority_page().await?;
if !priority_page.candidates.is_empty() || !priority_page.skipped_candidates.is_empty()
{
return Ok(Some(priority_page));
@@ -356,7 +377,10 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
if let Some(outcome) = self.pop_deferred_page(&candidate_api_format) {
return Ok(Some(outcome));
}
let Some(outcome) = self.next_page_for_api_format(&candidate_api_format).await? else {
let Some(outcome) = self
.next_page_for_api_format_with_planning_gate(&candidate_api_format)
.await?
else {
self.format_index += 1;
continue;
};
@@ -380,6 +404,70 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
self.deferred_pages_by_format.clear();
}
pub(crate) fn resolved_page_cache_preselection_mode(&self) -> &'static str {
self.key_mode.cache_key_name()
}
pub(crate) fn resolved_page_cache_use_api_format_alias_match(&self) -> bool {
self.use_api_format_alias_match
}
pub(crate) fn should_cache_current_priority_resolved_page(&self) -> bool {
if !(self.priority_page_emitted
&& self.format_index == 0
&& self.deferred_pages_by_format.is_empty())
{
return false;
}
match self.ordering_config.scheduling_mode {
SchedulerSchedulingMode::FixedOrder => true,
SchedulerSchedulingMode::CacheAffinity => {
has_explicit_session_affinity(self.client_session_affinity.as_ref())
}
SchedulerSchedulingMode::LoadBalance => false,
}
}
#[cfg(test)]
pub(crate) fn mark_priority_page_emitted_for_tests(&mut self) {
self.priority_page_emitted = true;
}
async fn cached_next_priority_page(
&mut self,
) -> Result<
AiCandidatePreselectionOutcome<
SchedulerMinimalCandidateSelectionCandidate,
SkippedLocalExecutionCandidate,
>,
GatewayError,
> {
let page = self.next_priority_page_with_planning_gate().await?;
self.remember_seen_candidates_from_page(&page);
Ok(page)
}
fn remember_seen_candidates_from_page(
&mut self,
page: &AiCandidatePreselectionOutcome<
SchedulerMinimalCandidateSelectionCandidate,
SkippedLocalExecutionCandidate,
>,
) {
for candidate in &page.candidates {
self.seen_candidate_keys
.insert(local_candidate_preselection_key(candidate, self.key_mode));
}
for skipped_candidate in &page.skipped_candidates {
self.seen_candidate_keys
.insert(local_candidate_preselection_key(
&skipped_candidate.candidate,
self.key_mode,
));
}
}
async fn next_priority_page(
&mut self,
) -> Result<
@@ -423,6 +511,35 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
Ok(priority_page)
}
async fn next_priority_page_with_planning_gate(
&mut self,
) -> Result<
AiCandidatePreselectionOutcome<
SchedulerMinimalCandidateSelectionCandidate,
SkippedLocalExecutionCandidate,
>,
GatewayError,
> {
let _permit = acquire_candidate_planning_gate(self.state, &self.trace_id).await?;
self.next_priority_page().await
}
async fn next_page_for_api_format_with_planning_gate(
&mut self,
candidate_api_format: &str,
) -> Result<
Option<
AiCandidatePreselectionOutcome<
SchedulerMinimalCandidateSelectionCandidate,
SkippedLocalExecutionCandidate,
>,
>,
GatewayError,
> {
let _permit = acquire_candidate_planning_gate(self.state, &self.trace_id).await?;
self.next_page_for_api_format(candidate_api_format).await
}
async fn split_priority_conversion_page(
&self,
candidate_api_format: &str,
@@ -797,7 +914,7 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
self.required_capabilities.as_ref(),
auth_snapshot,
self.client_session_affinity.as_ref(),
current_unix_secs(),
self.ranking_seed,
)
.await?;
let skipped_candidates = skipped_candidates
@@ -894,6 +1011,28 @@ fn local_candidate_preselection_key(
}
}
async fn acquire_candidate_planning_gate(
state: PlannerAppState<'_>,
trace_id: &str,
) -> Result<Option<ConcurrencyPermit>, GatewayError> {
let Some(gate) = state.app().candidate_planning_gate.as_ref() else {
return Ok(None);
};
let budget = state
.app()
.frontdoor_runtime_guards
.internal_gate_queue_budget;
match tokio::time::timeout(budget, gate.acquire()).await {
Ok(Ok(permit)) => Ok(Some(permit)),
Ok(Err(err)) => Err(GatewayError::Internal(err.to_string())),
Err(_) => Err(GatewayError::AdmissionTimeout {
trace_id: trace_id.to_string(),
gate: "gateway_candidate_planning",
queue_budget_ms: budget.as_millis() as u64,
}),
}
}
fn matches_client_api_format(
use_api_format_alias_match: bool,
candidate_api_format: &str,
@@ -1220,8 +1359,10 @@ mod tests {
&auth_snapshot,
None,
None,
None,
true,
LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModelAndApiFormat,
None,
)
.await;
@@ -1277,8 +1418,10 @@ mod tests {
&auth_snapshot,
None,
None,
None,
true,
LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModelAndApiFormat,
None,
)
.await;
@@ -1349,8 +1492,10 @@ mod tests {
&auth_snapshot,
None,
None,
None,
true,
LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModelAndApiFormat,
None,
)
.await;
@@ -1,8 +1,10 @@
use std::collections::BTreeMap;
use aether_contracts::ProxySnapshot;
use aether_scheduler_core::{
SchedulerMinimalCandidateSelectionCandidate, SchedulerTunnelAffinityBucket,
};
use serde_json::Value;
use tracing::warn;
use crate::ai_serving::{GatewayProviderTransportSnapshot, PlannerAppState};
@@ -10,7 +12,9 @@ use crate::scheduler::config::SchedulerOrderingConfig;
use super::candidate_resolution::read_candidate_transport_snapshot;
pub(super) type CandidateTransportIdentity<'a> = (&'a str, &'a str, &'a str);
const TUNNEL_OWNER_INSTANCE_ID_EXTRA_KEY: &str = "tunnel_owner_instance_id";
pub(super) type CandidateTransportIdentity = (String, String, String);
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(super) struct CandidateTransportRankingFacts {
@@ -18,38 +22,51 @@ pub(super) struct CandidateTransportRankingFacts {
pub(super) keep_priority_on_conversion: bool,
}
pub(super) async fn resolve_cached_candidate_transport_ranking_facts<'a>(
state: PlannerAppState<'_>,
cache: &mut BTreeMap<CandidateTransportIdentity<'a>, CandidateTransportRankingFacts>,
candidate: &'a SchedulerMinimalCandidateSelectionCandidate,
ordering_config: SchedulerOrderingConfig,
) -> CandidateTransportRankingFacts {
let identity = candidate_transport_identity(candidate);
if let Some(facts) = cache.get(&identity).copied() {
return facts;
}
let facts = resolve_candidate_transport_ranking_facts(state, candidate, ordering_config).await;
cache.insert(identity, facts);
facts
#[derive(Debug, Default)]
pub(super) struct CandidateTransportRankingFactsCache {
candidate_facts: BTreeMap<CandidateTransportIdentity, CandidateTransportRankingFacts>,
configured_proxy_snapshots: BTreeMap<String, Option<ProxySnapshot>>,
system_proxy_snapshot: Option<Option<ProxySnapshot>>,
tunnel_buckets_by_node_id: BTreeMap<String, SchedulerTunnelAffinityBucket>,
}
pub(super) async fn resolve_cached_transport_ranking_facts<'a>(
pub(super) async fn resolve_cached_candidate_transport_ranking_facts(
state: PlannerAppState<'_>,
cache: &mut BTreeMap<CandidateTransportIdentity<'a>, CandidateTransportRankingFacts>,
candidate: &'a SchedulerMinimalCandidateSelectionCandidate,
transport: &GatewayProviderTransportSnapshot,
cache: &mut CandidateTransportRankingFactsCache,
candidate: &SchedulerMinimalCandidateSelectionCandidate,
ordering_config: SchedulerOrderingConfig,
) -> CandidateTransportRankingFacts {
let identity = candidate_transport_identity(candidate);
if let Some(facts) = cache.get(&identity).copied() {
if let Some(facts) = cache.candidate_facts.get(&identity).copied() {
return facts;
}
let facts =
resolve_candidate_transport_ranking_facts_from_transport(state, transport, ordering_config)
.await;
cache.insert(identity, facts);
resolve_candidate_transport_ranking_facts(state, cache, candidate, ordering_config).await;
cache.candidate_facts.insert(identity, facts);
facts
}
pub(super) async fn resolve_cached_transport_ranking_facts(
state: PlannerAppState<'_>,
cache: &mut CandidateTransportRankingFactsCache,
candidate: &SchedulerMinimalCandidateSelectionCandidate,
transport: &GatewayProviderTransportSnapshot,
ordering_config: SchedulerOrderingConfig,
) -> CandidateTransportRankingFacts {
let identity = candidate_transport_identity(candidate);
if let Some(facts) = cache.candidate_facts.get(&identity).copied() {
return facts;
}
let facts = resolve_candidate_transport_ranking_facts_from_transport(
state,
cache,
transport,
ordering_config,
)
.await;
cache.candidate_facts.insert(identity, facts);
facts
}
@@ -58,13 +75,15 @@ pub(super) async fn candidate_keeps_priority_on_conversion(
candidate: &SchedulerMinimalCandidateSelectionCandidate,
ordering_config: SchedulerOrderingConfig,
) -> bool {
resolve_candidate_transport_ranking_facts(state, candidate, ordering_config)
let mut cache = CandidateTransportRankingFactsCache::default();
resolve_candidate_transport_ranking_facts(state, &mut cache, candidate, ordering_config)
.await
.keep_priority_on_conversion
}
async fn resolve_candidate_transport_ranking_facts(
state: PlannerAppState<'_>,
cache: &mut CandidateTransportRankingFactsCache,
candidate: &SchedulerMinimalCandidateSelectionCandidate,
ordering_config: SchedulerOrderingConfig,
) -> CandidateTransportRankingFacts {
@@ -75,17 +94,23 @@ async fn resolve_candidate_transport_ranking_facts(
};
};
resolve_candidate_transport_ranking_facts_from_transport(state, &transport, ordering_config)
.await
resolve_candidate_transport_ranking_facts_from_transport(
state,
cache,
&transport,
ordering_config,
)
.await
}
async fn resolve_candidate_transport_ranking_facts_from_transport(
state: PlannerAppState<'_>,
cache: &mut CandidateTransportRankingFactsCache,
transport: &GatewayProviderTransportSnapshot,
ordering_config: SchedulerOrderingConfig,
) -> CandidateTransportRankingFacts {
CandidateTransportRankingFacts {
tunnel_bucket: resolve_tunnel_owner_affinity_from_transport(state, transport).await,
tunnel_bucket: resolve_tunnel_owner_affinity_from_transport(state, cache, transport).await,
keep_priority_on_conversion: ordering_config.keep_priority_on_conversion
|| transport.provider.keep_priority_on_conversion,
}
@@ -93,12 +118,11 @@ async fn resolve_candidate_transport_ranking_facts_from_transport(
async fn resolve_tunnel_owner_affinity_from_transport(
state: PlannerAppState<'_>,
cache: &mut CandidateTransportRankingFactsCache,
transport: &GatewayProviderTransportSnapshot,
) -> SchedulerTunnelAffinityBucket {
let Some(proxy) = state
.app()
.resolve_transport_proxy_snapshot_with_tunnel_affinity(transport)
.await
let Some(proxy) =
resolve_transport_proxy_snapshot_with_tunnel_affinity_cached(state, cache, transport).await
else {
return SchedulerTunnelAffinityBucket::Neutral;
};
@@ -114,10 +138,74 @@ async fn resolve_tunnel_owner_affinity_from_transport(
return SchedulerTunnelAffinityBucket::Neutral;
};
if let Some(bucket) = cache.tunnel_buckets_by_node_id.get(node_id).copied() {
return bucket;
}
let bucket = resolve_tunnel_owner_affinity_from_proxy(state, &proxy, node_id).await;
cache
.tunnel_buckets_by_node_id
.insert(node_id.to_string(), bucket);
bucket
}
async fn resolve_transport_proxy_snapshot_with_tunnel_affinity_cached(
state: PlannerAppState<'_>,
cache: &mut CandidateTransportRankingFactsCache,
transport: &GatewayProviderTransportSnapshot,
) -> Option<ProxySnapshot> {
for raw in [
transport.key.proxy.as_ref(),
transport.endpoint.proxy.as_ref(),
transport.provider.proxy.as_ref(),
]
.into_iter()
.flatten()
{
let cache_key = proxy_config_cache_key(raw);
if let Some(snapshot) = cache.configured_proxy_snapshots.get(&cache_key) {
if snapshot.is_some() {
return snapshot.clone();
}
continue;
}
let snapshot = state
.app()
.resolve_configured_proxy_snapshot_with_tunnel_affinity(Some(raw))
.await;
cache
.configured_proxy_snapshots
.insert(cache_key, snapshot.clone());
if snapshot.is_some() {
return snapshot;
}
}
if let Some(snapshot) = cache.system_proxy_snapshot.as_ref() {
return snapshot.clone();
}
let snapshot = state.app().resolve_system_proxy_snapshot().await;
cache.system_proxy_snapshot = Some(snapshot.clone());
snapshot
}
async fn resolve_tunnel_owner_affinity_from_proxy(
state: PlannerAppState<'_>,
proxy: &ProxySnapshot,
node_id: &str,
) -> SchedulerTunnelAffinityBucket {
if state.app().tunnel.has_local_proxy(node_id) {
return SchedulerTunnelAffinityBucket::LocalTunnel;
}
if let Some(owner_instance_id) = proxy_tunnel_owner_instance_id(proxy) {
return if owner_instance_id == state.app().tunnel.local_instance_id() {
SchedulerTunnelAffinityBucket::LocalTunnel
} else {
SchedulerTunnelAffinityBucket::RemoteTunnel
};
}
match state
.app()
.tunnel
@@ -144,10 +232,25 @@ async fn resolve_tunnel_owner_affinity_from_transport(
fn candidate_transport_identity(
candidate: &SchedulerMinimalCandidateSelectionCandidate,
) -> CandidateTransportIdentity<'_> {
) -> CandidateTransportIdentity {
(
candidate.provider_id.as_str(),
candidate.endpoint_id.as_str(),
candidate.key_id.as_str(),
candidate.provider_id.clone(),
candidate.endpoint_id.clone(),
candidate.key_id.clone(),
)
}
fn proxy_config_cache_key(raw: &Value) -> String {
serde_json::to_string(raw).unwrap_or_else(|_| raw.to_string())
}
fn proxy_tunnel_owner_instance_id(proxy: &ProxySnapshot) -> Option<&str> {
proxy
.extra
.as_ref()
.and_then(Value::as_object)
.and_then(|extra| extra.get(TUNNEL_OWNER_INSTANCE_ID_EXTRA_KEY))
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
}
@@ -66,7 +66,7 @@ pub(crate) async fn maybe_build_sync_local_same_format_provider_decision_payload
candidate_count,
);
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
if let Some(payload) =
maybe_build_local_same_format_provider_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec,
@@ -134,7 +134,7 @@ pub(crate) async fn maybe_build_stream_local_same_format_provider_decision_paylo
candidate_count,
);
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
if let Some(payload) =
maybe_build_local_same_format_provider_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec,
@@ -190,7 +190,7 @@ pub(crate) async fn build_local_stream_attempt_source<'a>(
#[async_trait]
impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalSameFormatProviderSyncAttemptSource<'_> {
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
while let Some(attempt) = self.candidates.next_attempt().await {
while let Some(attempt) = self.candidates.next_attempt().await? {
match self.build_sync_attempt(attempt).await? {
Some(attempt) => return Ok(Some(attempt)),
None => continue,
@@ -220,7 +220,7 @@ impl LocalExecutionAttemptSource<AiStreamAttempt>
for LocalSameFormatProviderStreamAttemptSource<'_>
{
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
while let Some(attempt) = self.candidates.next_attempt().await {
while let Some(attempt) = self.candidates.next_attempt().await? {
match self.build_stream_attempt(attempt).await? {
Some(attempt) => return Ok(Some(attempt)),
None => continue,
@@ -372,7 +372,7 @@ pub(crate) async fn build_local_sync_plan_and_reports(
}
let mut plans = Vec::new();
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
let Some(payload) = maybe_build_local_same_format_provider_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec,
)
@@ -458,7 +458,7 @@ pub(crate) async fn build_local_stream_plan_and_reports(
}
let mut plans = Vec::new();
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
let Some(payload) = maybe_build_local_same_format_provider_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec,
)
@@ -1,5 +1,5 @@
use std::borrow::Cow;
use std::time::{SystemTime, UNIX_EPOCH};
use std::time::{Instant, SystemTime, UNIX_EPOCH};
use serde_json::Value;
use tracing::warn;
@@ -7,9 +7,11 @@ use tracing::warn;
use crate::ai_serving::ExecutionRuntimeAuthContext;
use crate::privacy::{
build_redaction_session_config, read_chat_pii_redaction_runtime_config,
try_mask_chat_pii_request_json_with_cache_options, ChatPiiRedactionRequestFormat,
MaskChatRequestOptions, RedactionMaskError, RedactionSessionSlot, RedisRedactionMappingCache,
try_mask_chat_pii_request_value_with_cache_options, CachedRequestRedaction,
ChatPiiRedactionRequestFormat, MaskChatRequestOptions, RedactionMaskError,
RedactionSessionSlot, RedisRedactionMappingCache,
};
use crate::stage_metrics::observe_gateway_stage_ms;
use crate::{AppState, GatewayError};
pub(crate) struct ProviderRequestRedaction<'a> {
@@ -73,6 +75,18 @@ pub(crate) async fn resolve_provider_chat_pii_redaction<'a>(
let Some(slot) = parts.extensions.get::<RedactionSessionSlot>() else {
return Ok(ProviderRequestRedaction::disabled(body_json));
};
let request_cache_key = request_redaction_cache_key(format, body_json);
if let Some(cached) = slot.cached_request_redaction(&request_cache_key) {
observe_gateway_stage_ms("chat_pii_redaction_request_cache_hit", 0);
return Ok(provider_redaction_from_cached(
slot,
candidate_id,
body_json,
cached,
));
}
let runtime_config_started_at = Instant::now();
let runtime_config = read_chat_pii_redaction_runtime_config(state)
.await
.map_err(|err| {
@@ -82,11 +96,22 @@ pub(crate) async fn resolve_provider_chat_pii_redaction<'a>(
);
GatewayError::Internal("chat pii redaction setup failed".to_string())
})?;
observe_gateway_stage_ms(
"chat_pii_redaction_runtime_config",
runtime_config_started_at.elapsed().as_millis() as u64,
);
if !runtime_config.enabled {
slot.put_cached_request_redaction(request_cache_key, CachedRequestRedaction::unredacted());
return Ok(ProviderRequestRedaction::disabled(body_json));
}
let feature_settings_started_at = Instant::now();
let feature_settings = resolve_chat_pii_redaction_feature_settings(state, auth_context).await?;
observe_gateway_stage_ms(
"chat_pii_redaction_feature_settings",
feature_settings_started_at.elapsed().as_millis() as u64,
);
if !feature_settings.effective_enabled() {
slot.put_cached_request_redaction(request_cache_key, CachedRequestRedaction::unredacted());
return Ok(ProviderRequestRedaction::disabled(body_json));
}
let Some(hmac_key) = state.encryption_key().map(str::as_bytes).map(Vec::from) else {
@@ -95,20 +120,14 @@ pub(crate) async fn resolve_provider_chat_pii_redaction<'a>(
"chat pii redaction setup failed".to_string(),
));
};
let body_bytes = serde_json::to_vec(body_json).map_err(|err| {
warn!(
error = ?err,
"gateway failed to serialize provider chat pii redaction body"
);
GatewayError::Internal("chat pii redaction setup failed".to_string())
})?;
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
let cache = RedisRedactionMappingCache::new(state.runtime_state.as_ref());
let masked = try_mask_chat_pii_request_json_with_cache_options(
&body_bytes,
let mask_started_at = Instant::now();
let masked = try_mask_chat_pii_request_value_with_cache_options(
body_json,
format,
build_redaction_session_config(hmac_key, &runtime_config, now_unix_secs),
MaskChatRequestOptions::runtime(),
@@ -116,19 +135,27 @@ pub(crate) async fn resolve_provider_chat_pii_redaction<'a>(
)
.await
.map_err(redaction_mask_error_to_gateway_error)?;
observe_gateway_stage_ms(
"chat_pii_redaction_mask_body",
mask_started_at.elapsed().as_millis() as u64,
);
if !masked.redacted {
slot.put_cached_request_redaction(request_cache_key, CachedRequestRedaction::unredacted());
return Ok(ProviderRequestRedaction {
body_json: Cow::Borrowed(body_json),
redacted: false,
});
}
let masked_body_json = serde_json::from_slice::<Value>(&masked.body).map_err(|err| {
warn!(
error = ?err,
"gateway failed to decode redacted provider chat pii body"
);
GatewayError::Internal("chat pii redaction setup failed".to_string())
})?;
let Some(masked_body_json) = masked.body_json else {
warn!("gateway pii redaction reported redacted without masked body");
return Err(GatewayError::Internal(
"chat pii redaction setup failed".to_string(),
));
};
slot.put_cached_request_redaction(
request_cache_key,
CachedRequestRedaction::redacted(masked_body_json.clone(), masked.session.clone()),
);
slot.put_for_candidate(candidate_id, masked.session);
Ok(ProviderRequestRedaction {
body_json: Cow::Owned(masked_body_json),
@@ -136,31 +163,46 @@ pub(crate) async fn resolve_provider_chat_pii_redaction<'a>(
})
}
fn request_redaction_cache_key(format: ChatPiiRedactionRequestFormat, body_json: &Value) -> String {
format!("{format:?}:{:p}", body_json)
}
fn provider_redaction_from_cached<'a>(
slot: &RedactionSessionSlot,
candidate_id: &str,
body_json: &'a Value,
cached: CachedRequestRedaction,
) -> ProviderRequestRedaction<'a> {
if !cached.redacted {
return ProviderRequestRedaction::disabled(body_json);
}
let Some(masked_body_json) = cached.body_json else {
return ProviderRequestRedaction::disabled(body_json);
};
if let Some(session) = cached.session {
slot.put_for_candidate(candidate_id, session);
}
ProviderRequestRedaction {
body_json: Cow::Owned(masked_body_json),
redacted: true,
}
}
async fn resolve_chat_pii_redaction_feature_settings(
state: &AppState,
auth_context: &ExecutionRuntimeAuthContext,
) -> Result<ChatPiiRedactionFeatureSettings, GatewayError> {
let user_settings = state
.read_user_feature_settings(&auth_context.user_id)
.await
let user_settings_fut = state.read_user_feature_settings(&auth_context.user_id);
let key_settings_fut = state.read_auth_api_key_feature_settings(
&auth_context.user_id,
&auth_context.api_key_id,
auth_context.api_key_is_standalone,
);
let (user_settings, key_settings) = tokio::try_join!(user_settings_fut, key_settings_fut)
.map_err(|err| {
warn!(
error = ?err,
"gateway failed to read user chat pii redaction feature settings"
);
GatewayError::Internal("chat pii redaction setup failed".to_string())
})?;
let key_settings = state
.read_auth_api_key_feature_settings(
&auth_context.user_id,
&auth_context.api_key_id,
auth_context.api_key_is_standalone,
)
.await
.map_err(|err| {
warn!(
error = ?err,
"gateway failed to read api key chat pii redaction feature settings"
"gateway failed to read chat pii redaction feature settings"
);
GatewayError::Internal("chat pii redaction setup failed".to_string())
})?;
@@ -175,7 +175,7 @@ pub(crate) async fn build_local_gemini_files_stream_attempt_source_for_kind<'a>(
#[async_trait]
impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalGeminiFilesSyncAttemptSource<'_> {
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
while let Some(attempt) = self.candidates.next_attempt().await {
while let Some(attempt) = self.candidates.next_attempt().await? {
match self.build_sync_attempt(attempt).await? {
Some(attempt) => return Ok(Some(attempt)),
None => continue,
@@ -198,7 +198,7 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalGeminiFilesSyncAttemptS
#[async_trait]
impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalGeminiFilesStreamAttemptSource<'_> {
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
while let Some(attempt) = self.candidates.next_attempt().await {
while let Some(attempt) = self.candidates.next_attempt().await? {
match self.build_stream_attempt(attempt).await? {
Some(attempt) => return Ok(Some(attempt)),
None => continue,
@@ -323,7 +323,7 @@ pub(crate) async fn maybe_build_sync_local_gemini_files_decision_payload(
let (mut source, _) =
build_local_gemini_files_candidate_attempt_source(state, trace_id, &input).await?;
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
if let Some(payload) = maybe_build_local_gemini_files_decision_payload_for_candidate(
state,
parts,
@@ -365,7 +365,7 @@ pub(crate) async fn maybe_build_stream_local_gemini_files_decision_payload(
build_local_gemini_files_candidate_attempt_source(state, trace_id, &input).await?;
let empty_body_json = serde_json::Value::Null;
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
if let Some(payload) = maybe_build_local_gemini_files_decision_payload_for_candidate(
state,
parts,
@@ -414,7 +414,7 @@ async fn build_local_sync_plan_and_reports(
build_local_gemini_files_candidate_attempt_source(state, trace_id, &input).await?;
let mut plans = Vec::new();
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
let Some(payload) = maybe_build_local_gemini_files_decision_payload_for_candidate(
state,
parts,
@@ -467,7 +467,7 @@ async fn build_local_stream_plan_and_reports(
let mut plans = Vec::new();
let empty_body_json = serde_json::Value::Null;
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
let Some(payload) = maybe_build_local_gemini_files_decision_payload_for_candidate(
state,
parts,
@@ -253,7 +253,7 @@ pub(crate) async fn build_local_image_stream_attempt_source_for_kind<'a>(
#[async_trait]
impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalOpenAiImageSyncAttemptSource<'_> {
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
while let Some(attempt) = self.candidates.next_attempt().await {
while let Some(attempt) = self.candidates.next_attempt().await? {
match self.build_sync_attempt(attempt).await? {
Some(attempt) => return Ok(Some(attempt)),
None => continue,
@@ -276,7 +276,7 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalOpenAiImageSyncAttemptS
#[async_trait]
impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalOpenAiImageStreamAttemptSource<'_> {
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
while let Some(attempt) = self.candidates.next_attempt().await {
while let Some(attempt) = self.candidates.next_attempt().await? {
match self.build_stream_attempt(attempt).await? {
Some(attempt) => return Ok(Some(attempt)),
None => continue,
@@ -421,7 +421,7 @@ pub(crate) async fn maybe_build_sync_local_image_decision_payload(
return Ok(None);
};
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
if let Some(payload) = maybe_build_local_openai_image_decision_payload_for_candidate(
state,
parts,
@@ -482,7 +482,7 @@ pub(crate) async fn maybe_build_stream_local_image_decision_payload(
return Ok(None);
};
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
if let Some(payload) = maybe_build_local_openai_image_decision_payload_for_candidate(
state,
parts,
@@ -540,7 +540,7 @@ async fn build_local_sync_plan_and_reports(
};
let mut plans = Vec::new();
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
let Some(payload) = maybe_build_local_openai_image_decision_payload_for_candidate(
state,
parts,
@@ -617,7 +617,7 @@ async fn build_local_stream_plan_and_reports(
};
let mut plans = Vec::new();
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
let Some(payload) = maybe_build_local_openai_image_decision_payload_for_candidate(
state,
parts,
@@ -105,7 +105,7 @@ pub(crate) async fn build_local_video_sync_attempt_source_for_kind<'a>(
#[async_trait]
impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalVideoCreateSyncAttemptSource<'_> {
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
while let Some(attempt) = self.candidates.next_attempt().await {
while let Some(attempt) = self.candidates.next_attempt().await? {
match self.build_sync_attempt(attempt).await? {
Some(attempt) => return Ok(Some(attempt)),
None => continue,
@@ -195,7 +195,7 @@ pub(crate) async fn maybe_build_sync_local_video_decision_payload(
return Ok(None);
};
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
if let Some(payload) = maybe_build_local_video_create_decision_payload_for_candidate(
state, parts, body_json, trace_id, &input, attempt, spec,
)
@@ -240,7 +240,7 @@ async fn build_local_sync_plan_and_reports(
};
let mut plans = Vec::new();
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
let Some(payload) = maybe_build_local_video_create_decision_payload_for_candidate(
state, parts, body_json, trace_id, &input, attempt, spec,
)
@@ -178,7 +178,7 @@ pub(crate) async fn build_local_stream_attempt_source<'a>(
#[async_trait]
impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalStandardSyncAttemptSource<'_> {
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
while let Some(attempt) = self.candidates.next_attempt().await {
while let Some(attempt) = self.candidates.next_attempt().await? {
match self.build_sync_attempt(attempt).await? {
Some(attempt) => return Ok(Some(attempt)),
None => continue,
@@ -206,7 +206,7 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalStandardSyncAttemptSour
#[async_trait]
impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalStandardStreamAttemptSource<'_> {
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
while let Some(attempt) = self.candidates.next_attempt().await {
while let Some(attempt) = self.candidates.next_attempt().await? {
match self.build_stream_attempt(attempt).await? {
Some(attempt) => return Ok(Some(attempt)),
None => continue,
@@ -340,7 +340,7 @@ pub(crate) async fn maybe_build_sync_via_standard_family_payload(
.await?;
apply_local_runtime_candidate_evaluation_progress(state, trace_id, candidate_count);
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
if let Some(payload) = maybe_build_local_standard_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec,
)
@@ -390,7 +390,7 @@ pub(crate) async fn maybe_build_stream_via_standard_family_payload(
.await?;
apply_local_runtime_candidate_evaluation_progress(state, trace_id, candidate_count);
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
if let Some(payload) = maybe_build_local_standard_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec,
)
@@ -449,7 +449,7 @@ pub(crate) async fn build_local_sync_plan_and_reports(
return Ok(Vec::new());
}
let mut plans = Vec::new();
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
let Some(payload) = maybe_build_local_standard_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec,
)
@@ -524,7 +524,7 @@ pub(crate) async fn build_local_stream_plan_and_reports(
return Ok(Vec::new());
}
let mut plans = Vec::new();
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
let Some(payload) = maybe_build_local_standard_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec,
)
@@ -12,6 +12,7 @@ use crate::ai_serving::planner::{
use crate::ai_serving::transport::{
resolve_transport_execution_timeouts, resolve_transport_profile,
};
use crate::stage_metrics::observe_gateway_stage_ms;
use crate::{
append_execution_contract_fields_to_value, append_local_failover_policy_to_value,
AiExecutionDecision, AppState, GatewayError,
@@ -40,6 +41,7 @@ pub(crate) async fn maybe_build_local_openai_chat_decision_payload_for_candidate
candidate_id,
..
} = attempt;
let payload_started_at = std::time::Instant::now();
let Some(resolved) = resolve_local_openai_chat_candidate_payload_parts(
state,
parts,
@@ -55,8 +57,16 @@ pub(crate) async fn maybe_build_local_openai_chat_decision_payload_for_candidate
)
.await?
else {
observe_gateway_stage_ms(
"stream_candidate_payload_parts",
payload_started_at.elapsed().as_millis() as u64,
);
return Ok(None);
};
observe_gateway_stage_ms(
"stream_candidate_payload_parts",
payload_started_at.elapsed().as_millis() as u64,
);
let candidate = &eligible.candidate;
let prompt_cache_key = resolved
@@ -66,9 +76,14 @@ pub(crate) async fn maybe_build_local_openai_chat_decision_payload_for_candidate
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned);
let proxy_started_at = std::time::Instant::now();
let proxy = state
.resolve_transport_proxy_snapshot_with_tunnel_affinity(&resolved.transport)
.await;
observe_gateway_stage_ms(
"stream_candidate_proxy",
proxy_started_at.elapsed().as_millis() as u64,
);
let transport_profile = resolved
.transport_profile
.clone()
@@ -130,6 +145,7 @@ pub(crate) async fn maybe_build_local_openai_chat_decision_payload_for_candidate
Some(body_json)
};
let effective_headers = input.effective_headers(&parts.headers);
let report_context_started_at = std::time::Instant::now();
let report_context = append_local_failover_policy_to_value(
append_execution_contract_fields_to_value(
build_local_execution_report_context(LocalExecutionReportContextParts {
@@ -184,8 +200,13 @@ pub(crate) async fn maybe_build_local_openai_chat_decision_payload_for_candidate
),
&transport,
);
observe_gateway_stage_ms(
"stream_candidate_report_context",
report_context_started_at.elapsed().as_millis() as u64,
);
let request_gzip = resolve_transport_request_gzip_policy(&transport);
let decision_started_at = std::time::Instant::now();
let mut decision = build_ai_execution_decision_response(AiExecutionDecisionResponseParts {
decision_is_stream,
decision_kind: decision_kind.to_string(),
@@ -222,5 +243,9 @@ pub(crate) async fn maybe_build_local_openai_chat_decision_payload_for_candidate
auth_context: input.auth_context.clone(),
});
apply_provider_request_routing_policy_to_decision(input, &mut decision)?;
observe_gateway_stage_ms(
"stream_candidate_decision_build",
decision_started_at.elapsed().as_millis() as u64,
);
Ok(Some(decision))
}
@@ -54,6 +54,7 @@ use crate::ai_serving::{
LocalResolvedOAuthRequestAuth,
};
use crate::ai_serving::{ConversionMode, ExecutionStrategy};
use crate::stage_metrics::observe_gateway_stage_ms;
use crate::{AppState, GatewayError};
use super::support::{
@@ -102,6 +103,7 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
report_kind: &str,
upstream_is_stream: bool,
) -> Result<Option<LocalOpenAiChatCandidatePayloadParts>, GatewayError> {
let prepare_started_at = std::time::Instant::now();
let planner_state = crate::ai_serving::PlannerAppState::new(state);
let candidate = &eligible.candidate;
let provider_api_format = eligible.provider_api_format.as_str();
@@ -109,6 +111,7 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
let transport_profile = crate::ai_serving::transport::resolve_transport_profile(transport);
let force_body_stream_field =
endpoint_config_forces_body_stream_field(transport.endpoint.config.as_ref());
let model_directives_started_at = std::time::Instant::now();
let enable_model_directives =
crate::system_features::reasoning_model_directive_enabled_for_api_format_and_model(
state,
@@ -116,6 +119,11 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
Some(&input.requested_model),
)
.await;
observe_gateway_stage_ms(
"openai_chat_payload_model_directives",
model_directives_started_at.elapsed().as_millis() as u64,
);
let redaction_started_at = std::time::Instant::now();
let redaction = resolve_provider_chat_pii_redaction(
state,
parts,
@@ -125,6 +133,10 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
candidate_id,
)
.await?;
observe_gateway_stage_ms(
"openai_chat_payload_redaction",
redaction_started_at.elapsed().as_millis() as u64,
);
let body_json = redaction.body_json.as_ref();
let effective_headers = input.effective_headers(&parts.headers);
let is_grok = transport
@@ -234,7 +246,7 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
redaction.redacted,
);
return Ok(Some(LocalOpenAiChatCandidatePayloadParts {
let result = Ok(Some(LocalOpenAiChatCandidatePayloadParts {
client_api_format: "openai:chat".to_string(),
auth_header: prepared_candidate.auth_header,
auth_value: prepared_candidate.auth_value,
@@ -252,6 +264,11 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
transport_profile,
image_request_summary: None,
}));
observe_gateway_stage_ms(
"openai_chat_payload_parts_prepare",
prepare_started_at.elapsed().as_millis() as u64,
);
return result;
}
if provider_api_format == "openai:chat" && is_windsurf_provider_transport(transport) {
@@ -288,6 +305,7 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
return Ok(None);
};
let auth_prepare_started_at = std::time::Instant::now();
let prepared_candidate = match prepare_header_authenticated_candidate(
planner_state,
transport,
@@ -316,7 +334,12 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
return Ok(None);
}
};
observe_gateway_stage_ms(
"openai_chat_payload_auth_prepare",
auth_prepare_started_at.elapsed().as_millis() as u64,
);
let body_build_started_at = std::time::Instant::now();
let Some(mut provider_request_body) = build_local_openai_chat_request_body(
body_json,
&prepared_candidate.mapped_model,
@@ -343,6 +366,10 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
.await;
return Ok(None);
};
observe_gateway_stage_ms(
"openai_chat_payload_body_build",
body_build_started_at.elapsed().as_millis() as u64,
);
apply_deepseek_tool_call_thinking_compat(
&mut provider_request_body,
transport.provider.provider_type.as_str(),
@@ -147,7 +147,7 @@ pub(crate) async fn maybe_build_sync_local_decision_payload(
)
.await;
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
let upstream_is_stream = self::plans::openai_chat_upstream_is_stream_for_candidate(
&attempt.eligible.transport,
attempt.eligible.provider_api_format.as_str(),
@@ -199,7 +199,7 @@ pub(crate) async fn maybe_build_stream_local_decision_payload(
)
.await;
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
let upstream_is_stream = self::plans::openai_chat_upstream_is_stream_for_candidate(
&attempt.eligible.transport,
attempt.eligible.provider_api_format.as_str(),
@@ -18,6 +18,7 @@ use crate::ai_serving::planner::plan_builders::{
build_openai_chat_stream_plan_from_decision, AiStreamAttempt,
};
use crate::ai_serving::planner::runtime_miss::apply_local_runtime_candidate_terminal_reason;
use crate::stage_metrics::observe_gateway_stage_ms;
pub(crate) struct LocalOpenAiChatStreamAttemptSource<'a> {
state: &'a AppState,
@@ -93,10 +94,36 @@ pub(crate) async fn build_local_openai_chat_stream_attempt_source<'a>(
#[async_trait]
impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalOpenAiChatStreamAttemptSource<'_> {
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
while let Some(attempt) = self.candidates.next_attempt().await {
loop {
let source_started_at = std::time::Instant::now();
let Some(attempt) = self.candidates.next_attempt().await? else {
observe_gateway_stage_ms(
"stream_candidate_source_next",
source_started_at.elapsed().as_millis() as u64,
);
break;
};
observe_gateway_stage_ms(
"stream_candidate_source_next",
source_started_at.elapsed().as_millis() as u64,
);
let plan_started_at = std::time::Instant::now();
match self.build_stream_attempt(attempt).await? {
Some(attempt) => return Ok(Some(attempt)),
None => continue,
Some(attempt) => {
observe_gateway_stage_ms(
"stream_candidate_plan_build",
plan_started_at.elapsed().as_millis() as u64,
);
return Ok(Some(attempt));
}
None => {
observe_gateway_stage_ms(
"stream_candidate_plan_build",
plan_started_at.elapsed().as_millis() as u64,
);
continue;
}
}
}
apply_local_runtime_candidate_terminal_reason(
@@ -93,7 +93,7 @@ pub(crate) async fn build_local_openai_chat_sync_attempt_source<'a>(
#[async_trait]
impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalOpenAiChatSyncAttemptSource<'_> {
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
while let Some(attempt) = self.candidates.next_attempt().await {
while let Some(attempt) = self.candidates.next_attempt().await? {
match self.build_sync_attempt(attempt).await? {
Some(attempt) => return Ok(Some(attempt)),
None => continue,
@@ -114,7 +114,7 @@ pub(crate) async fn maybe_build_sync_local_openai_responses_decision_payload(
)
.await?;
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
if let Some(payload) = maybe_build_local_openai_responses_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec,
)
@@ -153,7 +153,7 @@ pub(crate) async fn maybe_build_stream_local_openai_responses_decision_payload(
)
.await?;
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
if let Some(payload) = maybe_build_local_openai_responses_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec,
)
@@ -162,7 +162,7 @@ pub(super) async fn build_local_stream_attempt_source<'a>(
#[async_trait]
impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalOpenAiResponsesSyncAttemptSource<'_> {
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
while let Some(attempt) = self.candidates.next_attempt().await {
while let Some(attempt) = self.candidates.next_attempt().await? {
match self.build_sync_attempt(attempt).await? {
Some(attempt) => return Ok(Some(attempt)),
None => continue,
@@ -190,7 +190,7 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalOpenAiResponsesSyncAtte
#[async_trait]
impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalOpenAiResponsesStreamAttemptSource<'_> {
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
while let Some(attempt) = self.candidates.next_attempt().await {
while let Some(attempt) = self.candidates.next_attempt().await? {
match self.build_stream_attempt(attempt).await? {
Some(attempt) => return Ok(Some(attempt)),
None => continue,
@@ -331,7 +331,7 @@ pub(super) async fn build_local_sync_plan_and_reports(
}
let mut plans = Vec::new();
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
let Some(payload) = maybe_build_local_openai_responses_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec,
)
@@ -403,7 +403,7 @@ pub(super) async fn build_local_stream_plan_and_reports(
}
let mut plans = Vec::new();
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
let Some(payload) = maybe_build_local_openai_responses_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec,
)
@@ -3,8 +3,20 @@ pub(crate) use crate::ai_serving::transport::{
GatewayProviderTransportSnapshot, LocalResolvedOAuthRequestAuth,
};
use crate::GatewayError;
use std::sync::Arc;
impl<'a> PlannerAppState<'a> {
pub(crate) async fn read_provider_transport_snapshot_arc(
self,
provider_id: &str,
endpoint_id: &str,
key_id: &str,
) -> Result<Option<Arc<GatewayProviderTransportSnapshot>>, GatewayError> {
self.app()
.read_provider_transport_snapshot_arc(provider_id, endpoint_id, key_id)
.await
}
pub(crate) async fn read_provider_transport_snapshot(
self,
provider_id: &str,
+75 -1
View File
@@ -1,12 +1,54 @@
use std::collections::HashSet;
use std::time::Duration;
use aether_cache::ExpiringMap;
use tokio::sync::Notify;
use crate::control::GatewayControlAuthContext;
#[derive(Debug, Default)]
#[derive(Debug)]
pub(crate) struct AuthContextCache {
entries: ExpiringMap<String, GatewayControlAuthContext>,
inflight: std::sync::Mutex<HashSet<String>>,
notify: Notify,
}
impl Default for AuthContextCache {
fn default() -> Self {
Self {
entries: ExpiringMap::default(),
inflight: std::sync::Mutex::new(HashSet::new()),
notify: Notify::new(),
}
}
}
pub(crate) enum AuthContextInflightRegistration<'a> {
Leader(AuthContextInflightGuard<'a>),
Follower,
Bypass,
}
pub(crate) struct AuthContextInflightGuard<'a> {
cache: &'a AuthContextCache,
cache_key: Option<String>,
}
impl Drop for AuthContextInflightGuard<'_> {
fn drop(&mut self) {
let Some(cache_key) = self.cache_key.take() else {
return;
};
let removed = self
.cache
.inflight
.lock()
.map(|mut inflight| inflight.remove(&cache_key))
.unwrap_or(false);
if removed {
self.cache.notify.notify_waiters();
}
}
}
impl AuthContextCache {
@@ -29,7 +71,39 @@ impl AuthContextCache {
.insert(cache_key, auth_context, ttl, max_entries);
}
pub(crate) fn notified(&self) -> tokio::sync::futures::Notified<'_> {
self.notify.notified()
}
pub(crate) fn register_inflight(&self, cache_key: &str) -> AuthContextInflightRegistration<'_> {
let cache_key = cache_key.trim();
if cache_key.is_empty() {
return AuthContextInflightRegistration::Bypass;
}
match self.inflight.lock() {
Ok(mut inflight) => {
if inflight.contains(cache_key) {
AuthContextInflightRegistration::Follower
} else {
inflight.insert(cache_key.to_string());
AuthContextInflightRegistration::Leader(AuthContextInflightGuard {
cache: self,
cache_key: Some(cache_key.to_string()),
})
}
}
Err(_) => AuthContextInflightRegistration::Bypass,
}
}
pub(crate) fn clear(&self) {
self.entries.clear();
if let Ok(mut inflight) = self.inflight.lock() {
let had_inflight = !inflight.is_empty();
inflight.clear();
if had_inflight {
self.notify.notify_waiters();
}
}
}
}
+303 -7
View File
@@ -198,10 +198,10 @@ impl AuthSnapshotCache {
&self,
key: AuthSnapshotCacheKey,
ttl: Duration,
load: F,
mut load: F,
) -> Result<Option<GatewayAuthApiKeySnapshot>, E>
where
F: Fn() -> Fut,
F: FnMut() -> Fut,
Fut: Future<Output = Result<Option<GatewayAuthApiKeySnapshot>, E>>,
{
if let Some(value) = self.get(&key, ttl) {
@@ -269,10 +269,10 @@ where
&self,
key: K,
ttl: Duration,
load: F,
mut load: F,
) -> Result<Option<Value>, E>
where
F: Fn() -> Fut,
F: FnMut() -> Fut,
Fut: Future<Output = Result<Option<Value>, E>>,
{
if let Some(value) = self.get(&key, ttl) {
@@ -341,10 +341,10 @@ where
&self,
key: K,
ttl: Duration,
load: F,
mut load: F,
) -> Result<Option<V>, E>
where
F: Fn() -> Fut,
F: FnMut() -> Fut,
Fut: Future<Output = Result<Option<V>, E>>,
{
if let Some(value) = self.get(&key, ttl) {
@@ -374,15 +374,200 @@ where
}
}
pub(crate) async fn get_or_load_once<E, F, Fut>(
&self,
key: K,
ttl: Duration,
load: F,
) -> Result<Option<V>, E>
where
F: FnOnce() -> Fut,
Fut: Future<Output = Result<Option<V>, E>>,
{
self.get_or_load_once_with_observer(key, ttl, load, CacheLoadObserver::default())
.await
}
pub(crate) async fn get_or_load_once_with_observer<E, F, Fut>(
&self,
key: K,
ttl: Duration,
load: F,
observer: CacheLoadObserver,
) -> Result<Option<V>, E>
where
F: FnOnce() -> Fut,
Fut: Future<Output = Result<Option<V>, E>>,
{
if let Some(value) = self.get(&key, ttl) {
observer.hit();
return Ok(value);
}
observer.miss();
let mut load = Some(load);
loop {
let notified = self.singleflight.notified();
match self.singleflight.register(&key) {
CacheInflightRegistration::Bypass => {
observer.load();
let value =
load.take().expect("cache load closure should be available")().await?;
self.insert(key, value.clone(), ttl);
return Ok(value);
}
CacheInflightRegistration::Follower => {
observer.follower_wait();
notified.await;
if let Some(value) = self.get(&key, ttl) {
observer.hit();
return Ok(value);
}
}
CacheInflightRegistration::Leader(_guard) => {
observer.load();
let value =
load.take().expect("cache load closure should be available")().await?;
self.insert(key, value.clone(), ttl);
return Ok(value);
}
}
}
}
pub(crate) async fn get_or_load_once_stale_while_refreshing<E, F, Fut>(
&self,
key: K,
ttl: Duration,
stale_ttl: Duration,
load: F,
observer: CacheLoadObserver,
) -> Result<Option<V>, E>
where
F: FnOnce() -> Fut,
Fut: Future<Output = Result<Option<V>, E>>,
{
if let Some((value, age)) = self.entries.get_with_age(&key, stale_ttl) {
if age <= ttl {
observer.hit();
return Ok(value);
}
// Keep stale snapshots off the request critical path. The caller's
// invalidation path clears entries when provider/catalog/routing
// state changes, and the bounded stale TTL limits passive drift.
observer.hit();
return Ok(value);
}
observer.miss();
let mut load = Some(load);
loop {
let notified = self.singleflight.notified();
match self.singleflight.register(&key) {
CacheInflightRegistration::Bypass => {
observer.load();
let value =
load.take().expect("cache load closure should be available")().await?;
self.entries.insert(
key,
value.clone(),
stale_ttl,
AUTH_RUNTIME_CACHE_MAX_ENTRIES,
);
return Ok(value);
}
CacheInflightRegistration::Follower => {
observer.follower_wait();
notified.await;
if let Some((value, _age)) = self.entries.get_with_age(&key, stale_ttl) {
observer.hit();
return Ok(value);
}
}
CacheInflightRegistration::Leader(_guard) => {
observer.load();
let value =
load.take().expect("cache load closure should be available")().await?;
self.entries.insert(
key,
value.clone(),
stale_ttl,
AUTH_RUNTIME_CACHE_MAX_ENTRIES,
);
return Ok(value);
}
}
}
}
pub(crate) fn clear(&self) {
self.entries.clear();
self.singleflight.clear();
}
}
#[derive(Clone, Copy, Default)]
pub(crate) struct CacheLoadObserver {
on_hit: Option<fn()>,
on_miss: Option<fn()>,
on_load: Option<fn()>,
on_follower_wait: Option<fn()>,
}
impl CacheLoadObserver {
pub(crate) fn new() -> Self {
Self::default()
}
pub(crate) fn on_hit(mut self, callback: fn()) -> Self {
self.on_hit = Some(callback);
self
}
pub(crate) fn on_miss(mut self, callback: fn()) -> Self {
self.on_miss = Some(callback);
self
}
pub(crate) fn on_load(mut self, callback: fn()) -> Self {
self.on_load = Some(callback);
self
}
pub(crate) fn on_follower_wait(mut self, callback: fn()) -> Self {
self.on_follower_wait = Some(callback);
self
}
fn hit(self) {
if let Some(callback) = self.on_hit {
callback();
}
}
fn miss(self) {
if let Some(callback) = self.on_miss {
callback();
}
}
fn load(self) {
if let Some(callback) = self.on_load {
callback();
}
}
fn follower_wait(self) {
if let Some(callback) = self.on_follower_wait {
callback();
}
}
}
#[cfg(test)]
mod tests {
use super::ValueCache;
use super::{CacheLoadObserver, ValueCache};
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use std::time::{Duration, Instant};
@@ -472,6 +657,117 @@ mod tests {
assert_eq!(max_active.load(Ordering::Acquire), 2);
}
#[tokio::test]
async fn value_cache_returns_stale_without_refreshing_on_request_path() {
let cache = Arc::new(ValueCache::<String, u64>::default());
let key = "hot-key".to_string();
let calls = Arc::new(AtomicUsize::new(0));
cache.insert(key.clone(), Some(1), Duration::from_millis(10));
tokio::time::sleep(Duration::from_millis(20)).await;
let first_cache = Arc::clone(&cache);
let first_key = key.clone();
let first_calls = Arc::clone(&calls);
let first_started = Instant::now();
let first = tokio::spawn(async move {
first_cache
.get_or_load_once_stale_while_refreshing::<(), _, _>(
first_key,
Duration::from_millis(10),
Duration::from_secs(1),
|| async move {
first_calls.fetch_add(1, Ordering::AcqRel);
tokio::time::sleep(Duration::from_millis(100)).await;
Ok(Some(2))
},
CacheLoadObserver::default(),
)
.await
});
let follower_cache = Arc::clone(&cache);
let follower_started = Instant::now();
let follower_calls = Arc::clone(&calls);
let follower = tokio::spawn(async move {
follower_cache
.get_or_load_once_stale_while_refreshing::<(), _, _>(
key,
Duration::from_millis(10),
Duration::from_secs(1),
|| async move {
follower_calls.fetch_add(1, Ordering::AcqRel);
Ok(Some(3))
},
CacheLoadObserver::default(),
)
.await
});
assert_eq!(first.await.unwrap().unwrap(), Some(1));
assert!(
first_started.elapsed() < Duration::from_millis(80),
"stale value should not wait for request-path refresh"
);
assert_eq!(follower.await.unwrap().unwrap(), Some(1));
assert!(
follower_started.elapsed() < Duration::from_millis(80),
"follower should return stale value without waiting for refresh"
);
assert_eq!(calls.load(Ordering::Acquire), 0);
}
#[tokio::test]
async fn value_cache_cold_stale_followers_do_not_reload_after_fresh_ttl() {
let cache = Arc::new(ValueCache::<String, u64>::default());
let key = "cold-hot-key".to_string();
let calls = Arc::new(AtomicUsize::new(0));
let leader_cache = Arc::clone(&cache);
let leader_key = key.clone();
let leader_calls = Arc::clone(&calls);
let leader = tokio::spawn(async move {
leader_cache
.get_or_load_once_stale_while_refreshing::<(), _, _>(
leader_key,
Duration::from_millis(10),
Duration::from_secs(1),
|| async move {
leader_calls.fetch_add(1, Ordering::AcqRel);
Ok(Some(1))
},
CacheLoadObserver::default(),
)
.await
});
assert_eq!(leader.await.unwrap().unwrap(), Some(1));
tokio::time::sleep(Duration::from_millis(25)).await;
let follower_started = Instant::now();
let follower_cache = Arc::clone(&cache);
let follower_calls = Arc::clone(&calls);
let follower = tokio::spawn(async move {
follower_cache
.get_or_load_once_stale_while_refreshing::<(), _, _>(
key,
Duration::from_millis(10),
Duration::from_secs(1),
|| async move {
follower_calls.fetch_add(1, Ordering::AcqRel);
Ok(Some(2))
},
CacheLoadObserver::default(),
)
.await
});
assert_eq!(follower.await.unwrap().unwrap(), Some(1));
assert_eq!(calls.load(Ordering::Acquire), 1);
assert!(
follower_started.elapsed() < Duration::from_millis(50),
"follower should reuse cold-loaded stale value without reloading"
);
}
#[tokio::test]
async fn value_cache_clear_releases_same_key_followers() {
let cache = Arc::new(ValueCache::<String, u64>::default());
+479
View File
@@ -0,0 +1,479 @@
use std::collections::hash_map::DefaultHasher;
use std::hash::{Hash, Hasher};
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, LazyLock};
use std::time::Duration;
use aether_ai_serving::AiCandidatePreselectionOutcome;
use aether_ai_serving::AiCandidateResolutionMode;
use aether_routing_core::ResolvedRoutingPolicy;
use aether_runtime::{MetricKind, MetricSample};
use aether_scheduler_core::{
normalize_api_format, ClientSessionAffinity, SchedulerMinimalCandidateSelectionCandidate,
};
use serde_json::Value;
use sha2::Digest as _;
use crate::ai_serving::{
EligibleLocalExecutionCandidate, GatewayAuthApiKeySnapshot, SkippedLocalExecutionCandidate,
};
const DEFAULT_CANDIDATE_PAGE_CACHE_TTL_MS: u64 = 250;
const MIN_CANDIDATE_PAGE_CACHE_TTL_MS: u64 = 50;
const MAX_CANDIDATE_PAGE_CACHE_TTL_MS: u64 = 1_000;
const CANDIDATE_PAGE_CACHE_TTL_ENV: &str = "AETHER_GATEWAY_CANDIDATE_PAGE_CACHE_TTL_MS";
pub(crate) type CandidatePageSnapshot = AiCandidatePreselectionOutcome<
SchedulerMinimalCandidateSelectionCandidate,
SkippedLocalExecutionCandidate,
>;
pub(crate) type CandidatePageCache =
super::ValueCache<CandidatePageCacheKey, Arc<CandidatePageSnapshot>>;
#[derive(Debug, Clone)]
pub(crate) struct CandidateResolvedPageSnapshot {
pub(crate) candidates: Vec<EligibleLocalExecutionCandidate>,
pub(crate) resolved_skipped: Vec<SkippedLocalExecutionCandidate>,
}
pub(crate) type CandidateResolvedPageCache =
super::ValueCache<CandidateResolvedPageCacheKey, Arc<CandidateResolvedPageSnapshot>>;
static CANDIDATE_PAGE_CACHE_METRICS: LazyLock<CandidatePageCacheMetrics> =
LazyLock::new(CandidatePageCacheMetrics::default);
#[derive(Debug, Default)]
struct CandidatePageCacheMetrics {
hit_total: AtomicU64,
load_total: AtomicU64,
follower_wait_total: AtomicU64,
miss_total: AtomicU64,
none_total: AtomicU64,
resolve_hit_total: AtomicU64,
resolve_load_total: AtomicU64,
resolve_follower_wait_total: AtomicU64,
resolve_miss_total: AtomicU64,
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub(crate) struct CandidatePageCacheKey {
requested_model: String,
client_api_format: String,
auth_identity: CandidatePageAuthIdentity,
require_streaming: bool,
required_capabilities_hash: String,
routing_policy_hash: String,
request_auth_channel: String,
scheduler_affinity_epoch: u64,
preselection_mode: &'static str,
use_api_format_alias_match: bool,
client_session_affinity_hash: String,
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
enum CandidatePageAuthIdentity {
Standalone { api_key_id: String },
UserApiKey { user_id: String, api_key_id: String },
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub(crate) struct CandidateResolvedPageCacheKey {
page_key: CandidatePageCacheKey,
resolution_mode: &'static str,
}
impl CandidatePageCacheKey {
#[allow(clippy::too_many_arguments)]
pub(crate) fn new(
requested_model: &str,
client_api_format: &str,
require_streaming: bool,
auth_snapshot: &GatewayAuthApiKeySnapshot,
required_capabilities: Option<&Value>,
routing_policy: Option<&ResolvedRoutingPolicy>,
request_auth_channel: Option<&str>,
scheduler_affinity_epoch: u64,
preselection_mode: &'static str,
use_api_format_alias_match: bool,
client_session_affinity: Option<&ClientSessionAffinity>,
) -> Self {
Self {
requested_model: normalize_text_key(requested_model),
client_api_format: normalize_api_format(client_api_format),
auth_identity: CandidatePageAuthIdentity::from_auth_snapshot(auth_snapshot),
require_streaming,
required_capabilities_hash: stable_json_hash(required_capabilities),
routing_policy_hash: stable_json_hash(routing_policy),
request_auth_channel: normalize_text_key(request_auth_channel.unwrap_or_default()),
scheduler_affinity_epoch,
preselection_mode,
use_api_format_alias_match,
client_session_affinity_hash: client_session_affinity_key(client_session_affinity),
}
}
}
impl CandidateResolvedPageCacheKey {
#[allow(clippy::too_many_arguments)]
pub(crate) fn new(
requested_model: &str,
client_api_format: &str,
require_streaming: bool,
auth_snapshot: &GatewayAuthApiKeySnapshot,
required_capabilities: Option<&Value>,
routing_policy: Option<&ResolvedRoutingPolicy>,
request_auth_channel: Option<&str>,
scheduler_affinity_epoch: u64,
preselection_mode: &'static str,
use_api_format_alias_match: bool,
client_session_affinity: Option<&ClientSessionAffinity>,
resolution_mode: AiCandidateResolutionMode,
) -> Self {
Self {
page_key: CandidatePageCacheKey::new(
requested_model,
client_api_format,
require_streaming,
auth_snapshot,
required_capabilities,
routing_policy,
request_auth_channel,
scheduler_affinity_epoch,
preselection_mode,
use_api_format_alias_match,
client_session_affinity,
),
resolution_mode: resolution_mode_name(resolution_mode),
}
}
}
impl CandidatePageAuthIdentity {
fn from_auth_snapshot(auth_snapshot: &GatewayAuthApiKeySnapshot) -> Self {
let api_key_id = normalize_text_key(&auth_snapshot.api_key_id);
if auth_snapshot.api_key_is_standalone {
Self::Standalone { api_key_id }
} else {
Self::UserApiKey {
user_id: normalize_text_key(&auth_snapshot.user_id),
api_key_id,
}
}
}
}
pub(crate) fn candidate_page_cache_ttl_from_env() -> Duration {
let ttl_ms = std::env::var(CANDIDATE_PAGE_CACHE_TTL_ENV)
.ok()
.and_then(|value| value.trim().parse::<u64>().ok())
.filter(|value| *value > 0)
.unwrap_or(DEFAULT_CANDIDATE_PAGE_CACHE_TTL_MS)
.clamp(
MIN_CANDIDATE_PAGE_CACHE_TTL_MS,
MAX_CANDIDATE_PAGE_CACHE_TTL_MS,
);
Duration::from_millis(ttl_ms)
}
pub(crate) fn candidate_page_cache_stale_ttl(ttl: Duration) -> Duration {
let stale_ttl = ttl.saturating_mul(8);
stale_ttl.min(Duration::from_secs(2)).max(ttl)
}
pub(crate) fn record_candidate_page_cache_hit() {
CANDIDATE_PAGE_CACHE_METRICS
.hit_total
.fetch_add(1, Ordering::Relaxed);
}
pub(crate) fn record_candidate_page_cache_miss() {
CANDIDATE_PAGE_CACHE_METRICS
.miss_total
.fetch_add(1, Ordering::Relaxed);
}
pub(crate) fn record_candidate_page_cache_load() {
CANDIDATE_PAGE_CACHE_METRICS
.load_total
.fetch_add(1, Ordering::Relaxed);
}
pub(crate) fn record_candidate_page_cache_follower_wait() {
CANDIDATE_PAGE_CACHE_METRICS
.follower_wait_total
.fetch_add(1, Ordering::Relaxed);
}
pub(crate) fn record_candidate_page_cache_none() {
CANDIDATE_PAGE_CACHE_METRICS
.none_total
.fetch_add(1, Ordering::Relaxed);
}
pub(crate) fn record_candidate_page_resolve_cache_hit() {
CANDIDATE_PAGE_CACHE_METRICS
.resolve_hit_total
.fetch_add(1, Ordering::Relaxed);
}
pub(crate) fn record_candidate_page_resolve_cache_miss() {
CANDIDATE_PAGE_CACHE_METRICS
.resolve_miss_total
.fetch_add(1, Ordering::Relaxed);
}
pub(crate) fn record_candidate_page_resolve_cache_load() {
CANDIDATE_PAGE_CACHE_METRICS
.resolve_load_total
.fetch_add(1, Ordering::Relaxed);
}
pub(crate) fn record_candidate_page_resolve_cache_follower_wait() {
CANDIDATE_PAGE_CACHE_METRICS
.resolve_follower_wait_total
.fetch_add(1, Ordering::Relaxed);
}
pub(crate) fn candidate_page_cache_metric_samples() -> Vec<MetricSample> {
vec![
MetricSample::new(
"candidate_page_cache_hit_total",
"Total candidate page cache hits.",
MetricKind::Counter,
CANDIDATE_PAGE_CACHE_METRICS
.hit_total
.load(Ordering::Relaxed),
),
MetricSample::new(
"candidate_page_cache_miss_total",
"Total candidate page cache misses before singleflight registration.",
MetricKind::Counter,
CANDIDATE_PAGE_CACHE_METRICS
.miss_total
.load(Ordering::Relaxed),
),
MetricSample::new(
"candidate_page_cache_load_total",
"Total candidate page cache loader executions.",
MetricKind::Counter,
CANDIDATE_PAGE_CACHE_METRICS
.load_total
.load(Ordering::Relaxed),
),
MetricSample::new(
"candidate_page_cache_follower_wait_total",
"Total candidate page cache requests that waited for another loader.",
MetricKind::Counter,
CANDIDATE_PAGE_CACHE_METRICS
.follower_wait_total
.load(Ordering::Relaxed),
),
MetricSample::new(
"candidate_page_cache_none_total",
"Total candidate page cache lookups that resolved to no page.",
MetricKind::Counter,
CANDIDATE_PAGE_CACHE_METRICS
.none_total
.load(Ordering::Relaxed),
),
MetricSample::new(
"candidate_page_resolve_cache_hit_total",
"Total resolved candidate page cache hits.",
MetricKind::Counter,
CANDIDATE_PAGE_CACHE_METRICS
.resolve_hit_total
.load(Ordering::Relaxed),
),
MetricSample::new(
"candidate_page_resolve_cache_miss_total",
"Total resolved candidate page cache misses before singleflight registration.",
MetricKind::Counter,
CANDIDATE_PAGE_CACHE_METRICS
.resolve_miss_total
.load(Ordering::Relaxed),
),
MetricSample::new(
"candidate_page_resolve_cache_load_total",
"Total resolved candidate page cache loader executions.",
MetricKind::Counter,
CANDIDATE_PAGE_CACHE_METRICS
.resolve_load_total
.load(Ordering::Relaxed),
),
MetricSample::new(
"candidate_page_resolve_cache_follower_wait_total",
"Total resolved candidate page cache requests that waited for another loader.",
MetricKind::Counter,
CANDIDATE_PAGE_CACHE_METRICS
.resolve_follower_wait_total
.load(Ordering::Relaxed),
),
]
}
fn normalize_text_key(value: &str) -> String {
value.trim().to_string()
}
fn client_session_affinity_key(affinity: Option<&ClientSessionAffinity>) -> String {
let Some(affinity) = affinity else {
return String::new();
};
let family = affinity
.client_family
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.map(str::to_ascii_lowercase)
.unwrap_or_default();
let session = affinity
.session_key
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.map(sha256_hex)
.unwrap_or_default();
format!("{family}:{session}")
}
fn stable_json_hash<T>(value: Option<&T>) -> String
where
T: serde::Serialize,
{
let Some(value) = value else {
return String::new();
};
match serde_json::to_vec(value) {
Ok(serialized) => sha256_hex(&serialized),
Err(_) => {
let mut hasher = DefaultHasher::new();
std::any::type_name::<T>().hash(&mut hasher);
format!("fallback:{:016x}", hasher.finish())
}
}
}
fn sha256_hex(value: impl AsRef<[u8]>) -> String {
let digest = sha2::Sha256::digest(value.as_ref());
digest.iter().map(|byte| format!("{byte:02x}")).collect()
}
fn resolution_mode_name(mode: AiCandidateResolutionMode) -> &'static str {
match mode {
AiCandidateResolutionMode::Standard => "standard",
AiCandidateResolutionMode::WithoutTransportPairGate => "without_transport_pair_gate",
}
}
#[cfg(test)]
mod tests {
use super::*;
use aether_data::repository::auth::ResolvedAuthApiKeySnapshot;
use serde_json::json;
fn auth_snapshot(user_id: &str, api_key_id: &str) -> ResolvedAuthApiKeySnapshot {
ResolvedAuthApiKeySnapshot {
user_id: user_id.to_string(),
username: "user".to_string(),
email: None,
user_role: "user".to_string(),
user_auth_source: "local".to_string(),
user_is_active: true,
user_is_deleted: false,
user_rate_limit: None,
user_allowed_providers: None,
user_allowed_api_formats: None,
user_allowed_models: None,
api_key_id: api_key_id.to_string(),
api_key_name: None,
api_key_is_active: true,
api_key_is_locked: false,
api_key_is_standalone: false,
api_key_rate_limit: None,
api_key_concurrent_limit: None,
api_key_expires_at_unix_secs: None,
api_key_allowed_providers: None,
api_key_allowed_api_formats: None,
api_key_allowed_models: None,
api_key_ip_rules: None,
currently_usable: true,
}
}
#[test]
fn candidate_page_cache_key_isolates_auth_model_format_and_capabilities() {
let auth_a = auth_snapshot("user-a", "key-a");
let auth_b = auth_snapshot("user-b", "key-a");
let base = CandidatePageCacheKey::new(
"gpt-4o",
"openai:chat",
true,
&auth_a,
Some(&json!({"vision": true})),
None,
Some("bearer"),
7,
"provider_endpoint_key_model",
true,
None,
);
let different_user = CandidatePageCacheKey::new(
"gpt-4o",
"openai:chat",
true,
&auth_b,
Some(&json!({"vision": true})),
None,
Some("bearer"),
7,
"provider_endpoint_key_model",
true,
None,
);
let different_model = CandidatePageCacheKey::new(
"gpt-4.1",
"openai:chat",
true,
&auth_a,
Some(&json!({"vision": true})),
None,
Some("bearer"),
7,
"provider_endpoint_key_model",
true,
None,
);
let different_format = CandidatePageCacheKey::new(
"gpt-4o",
"openai:responses",
true,
&auth_a,
Some(&json!({"vision": true})),
None,
Some("bearer"),
7,
"provider_endpoint_key_model",
true,
None,
);
let different_capabilities = CandidatePageCacheKey::new(
"gpt-4o",
"openai:chat",
true,
&auth_a,
Some(&json!({"vision": false})),
None,
Some("bearer"),
7,
"provider_endpoint_key_model",
true,
None,
);
assert_ne!(base, different_user);
assert_ne!(base, different_model);
assert_ne!(base, different_format);
assert_ne!(base, different_capabilities);
}
}
+14 -3
View File
@@ -1,20 +1,31 @@
mod auth_api_key_last_used;
mod auth_context;
mod auth_runtime;
mod candidate_page;
mod dashboard_response;
mod direct_plan_bypass;
mod scheduler_affinity;
mod system_config;
pub(crate) use auth_api_key_last_used::AuthApiKeyLastUsedCache;
pub(crate) use auth_context::AuthContextCache;
pub(crate) use auth_context::{AuthContextCache, AuthContextInflightRegistration};
pub(crate) use auth_runtime::{
AuthApiKeyFeatureCacheKey, AuthApiKeyIdentityCacheKey, AuthSnapshotCache, AuthSnapshotCacheKey,
JsonValueCache, ValueCache,
CacheLoadObserver, JsonValueCache, ValueCache,
};
pub(crate) use candidate_page::{
candidate_page_cache_metric_samples, candidate_page_cache_stale_ttl,
candidate_page_cache_ttl_from_env, record_candidate_page_cache_follower_wait,
record_candidate_page_cache_hit, record_candidate_page_cache_load,
record_candidate_page_cache_miss, record_candidate_page_cache_none,
record_candidate_page_resolve_cache_follower_wait, record_candidate_page_resolve_cache_hit,
record_candidate_page_resolve_cache_load, record_candidate_page_resolve_cache_miss,
CandidatePageCache, CandidatePageCacheKey, CandidatePageSnapshot, CandidateResolvedPageCache,
CandidateResolvedPageCacheKey, CandidateResolvedPageSnapshot,
};
pub(crate) use dashboard_response::DashboardResponseCache;
pub(crate) use direct_plan_bypass::DirectPlanBypassCache;
pub(crate) use scheduler_affinity::{
SchedulerAffinityCache, SchedulerAffinitySnapshotEntry, SchedulerAffinityTarget,
};
pub(crate) use system_config::SystemConfigCache;
pub(crate) use system_config::{SystemConfigCache, SystemConfigInflightRegistration};
+67 -5
View File
@@ -1,21 +1,43 @@
use std::collections::HashSet;
use std::time::Duration;
use aether_cache::ExpiringMap;
use tokio::sync::{Mutex, MutexGuard};
use tokio::sync::Notify;
const MAX_ENTRIES: usize = 512;
#[derive(Debug)]
pub(crate) struct SystemConfigCache {
entries: ExpiringMap<String, Option<serde_json::Value>>,
load_guard: Mutex<()>,
inflight: std::sync::Mutex<HashSet<String>>,
notify: Notify,
}
impl Default for SystemConfigCache {
fn default() -> Self {
Self {
entries: ExpiringMap::new(),
load_guard: Mutex::new(()),
inflight: std::sync::Mutex::new(HashSet::new()),
notify: Notify::new(),
}
}
}
pub(crate) enum SystemConfigInflightRegistration<'a> {
Leader(SystemConfigInflightGuard<'a>),
Follower,
Bypass,
}
pub(crate) struct SystemConfigInflightGuard<'a> {
cache: &'a SystemConfigCache,
key: Option<String>,
}
impl Drop for SystemConfigInflightGuard<'_> {
fn drop(&mut self) {
if let Some(key) = self.key.take() {
self.cache.finish_load(&key);
}
}
}
@@ -29,11 +51,51 @@ impl SystemConfigCache {
self.entries.insert(key, value, ttl, MAX_ENTRIES);
}
pub(crate) async fn load_guard(&self) -> MutexGuard<'_, ()> {
self.load_guard.lock().await
pub(crate) fn register_load(&self, key: &str) -> SystemConfigInflightRegistration<'_> {
match self.inflight.lock() {
Ok(mut inflight) => {
if inflight.contains(key) {
SystemConfigInflightRegistration::Follower
} else {
inflight.insert(key.to_string());
SystemConfigInflightRegistration::Leader(SystemConfigInflightGuard {
cache: self,
key: Some(key.to_string()),
})
}
}
Err(_) => SystemConfigInflightRegistration::Bypass,
}
}
pub(crate) fn notified(&self) -> tokio::sync::futures::Notified<'_> {
self.notify.notified()
}
fn finish_load(&self, key: &str) {
let removed = self
.inflight
.lock()
.map(|mut inflight| inflight.remove(key))
.unwrap_or(false);
if removed {
self.notify.notify_waiters();
}
}
pub(crate) fn clear(&self) {
self.entries.clear();
let cleared = self
.inflight
.lock()
.map(|mut inflight| {
let had_entries = !inflight.is_empty();
inflight.clear();
had_entries
})
.unwrap_or(false);
if cleared {
self.notify.notify_waiters();
}
}
}
+9
View File
@@ -1,5 +1,8 @@
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::{SystemTime, UNIX_EPOCH};
static REQUEST_DISTRIBUTION_COUNTER: AtomicU64 = AtomicU64::new(0);
pub(crate) fn current_unix_secs() -> u64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
@@ -13,3 +16,9 @@ pub(crate) fn current_unix_ms() -> u64 {
.unwrap_or_default()
.as_millis() as u64
}
pub(crate) fn request_distribution_seed() -> u64 {
let now_ms = current_unix_ms();
let counter = REQUEST_DISTRIBUTION_COUNTER.fetch_add(1, Ordering::Relaxed);
now_ms.rotate_left(21) ^ counter
}
@@ -23,6 +23,7 @@ use super::principal::derive_principal_candidate;
use super::types::{
GatewayCredentialCarrier, GatewayPrincipalCandidate, GatewayTrustedAuthHeaders,
};
use crate::cache::AuthContextInflightRegistration;
use crate::headers::header_value_str;
const AUTH_CONTEXT_CACHE_TTL: Duration = Duration::from_secs(60);
@@ -137,14 +138,15 @@ pub(in super::super) async fn resolve_control_decision_auth(
if let Some(cache_key) = auth_context_cache_key.as_deref() {
if let Some(auth_context) = get_cached_auth_context(state, cache_key) {
resolved_auth_context = if auth_context_cache_refresh_on_hit() {
let refreshed = refresh_execution_runtime_auth_context(
state,
auth_context,
decision.auth_endpoint_signature.as_deref(),
Some(
refresh_cached_auth_context_or_reuse(
state,
cache_key,
auth_context,
decision.auth_endpoint_signature.as_deref(),
)
.await?,
)
.await?;
put_cached_auth_context(state, cache_key.to_string(), refreshed.clone());
Some(refreshed)
} else {
Some(auth_context)
};
@@ -152,11 +154,13 @@ pub(in super::super) async fn resolve_control_decision_auth(
}
if resolved_auth_context.is_none() {
resolved_auth_context = resolve_data_backed_auth_context(
resolved_auth_context = resolve_data_backed_auth_context_cached(
state,
auth_context_cache_key.as_deref(),
headers,
uri,
decision.auth_endpoint_signature.as_deref(),
true,
)
.await?;
if let (Some(cache_key), Some(auth_context)) = (
@@ -501,14 +505,15 @@ pub(crate) async fn resolve_execution_runtime_auth_context(
if !auth_context_cache_refresh_on_hit() {
return Ok(Some(auth_context));
}
return Ok(Some(
refresh_execution_runtime_auth_context(
state,
auth_context,
decision.auth_endpoint_signature.as_deref(),
)
.await?,
));
return refresh_decision_auth_context_on_hit(
state,
headers,
uri,
decision.auth_endpoint_signature.as_deref(),
auth_context,
)
.await
.map(Some);
}
let Some(auth_endpoint_signature) = decision.auth_endpoint_signature.as_deref() else {
@@ -524,18 +529,25 @@ pub(crate) async fn resolve_execution_runtime_auth_context(
return Ok(Some(auth_context));
}
let refreshed = refresh_execution_runtime_auth_context(
let refreshed = refresh_cached_auth_context_or_reuse(
state,
&cache_key,
auth_context,
Some(auth_endpoint_signature),
)
.await?;
put_cached_auth_context(state, cache_key, refreshed.clone());
return Ok(Some(refreshed));
}
if let Some(auth_context) =
resolve_data_backed_auth_context(state, headers, uri, Some(auth_endpoint_signature)).await?
if let Some(auth_context) = resolve_data_backed_auth_context_cached(
state,
Some(cache_key.as_str()),
headers,
uri,
Some(auth_endpoint_signature),
true,
)
.await?
{
if auth_context.user_id.is_empty() || auth_context.api_key_id.is_empty() {
return Ok(None);
@@ -547,6 +559,109 @@ pub(crate) async fn resolve_execution_runtime_auth_context(
Ok(None)
}
async fn refresh_decision_auth_context_on_hit(
state: &AppState,
headers: &http::HeaderMap,
uri: &Uri,
auth_endpoint_signature: Option<&str>,
auth_context: GatewayControlAuthContext,
) -> Result<GatewayControlAuthContext, GatewayError> {
let Some(auth_endpoint_signature) = auth_endpoint_signature else {
return Ok(auth_context);
};
let Some(cache_key) = build_auth_context_cache_key(headers, uri, auth_endpoint_signature)
else {
return refresh_execution_runtime_auth_context(
state,
auth_context,
Some(auth_endpoint_signature),
)
.await;
};
refresh_cached_auth_context_or_reuse(
state,
&cache_key,
auth_context,
Some(auth_endpoint_signature),
)
.await
}
async fn refresh_cached_auth_context_or_reuse(
state: &AppState,
cache_key: &str,
auth_context: GatewayControlAuthContext,
auth_endpoint_signature: Option<&str>,
) -> Result<GatewayControlAuthContext, GatewayError> {
match state.auth_context_cache.register_inflight(cache_key) {
AuthContextInflightRegistration::Leader(_guard) => {
let refreshed = refresh_execution_runtime_auth_context(
state,
auth_context,
auth_endpoint_signature,
)
.await?;
put_cached_auth_context(state, cache_key.to_string(), refreshed.clone());
Ok(refreshed)
}
AuthContextInflightRegistration::Follower => Ok(auth_context),
AuthContextInflightRegistration::Bypass => {
refresh_execution_runtime_auth_context(state, auth_context, auth_endpoint_signature)
.await
}
}
}
async fn resolve_data_backed_auth_context_cached(
state: &AppState,
cache_key: Option<&str>,
headers: &http::HeaderMap,
uri: &Uri,
auth_endpoint_signature: Option<&str>,
cache_negative: bool,
) -> Result<Option<GatewayControlAuthContext>, GatewayError> {
let Some(cache_key) = cache_key else {
return resolve_data_backed_auth_context(state, headers, uri, auth_endpoint_signature)
.await;
};
loop {
let notified = state.auth_context_cache.notified();
match state.auth_context_cache.register_inflight(cache_key) {
AuthContextInflightRegistration::Leader(_guard) => {
let resolved =
resolve_data_backed_auth_context(state, headers, uri, auth_endpoint_signature)
.await?;
if let Some(auth_context) = resolved.as_ref() {
if cache_negative
|| (!auth_context.user_id.is_empty() && !auth_context.api_key_id.is_empty())
{
put_cached_auth_context(state, cache_key.to_string(), auth_context.clone());
}
}
return Ok(resolved);
}
AuthContextInflightRegistration::Follower => {
notified.await;
if let Some(auth_context) = get_cached_auth_context(state, cache_key) {
return Ok(Some(auth_context));
}
if !cache_negative {
return Ok(None);
}
}
AuthContextInflightRegistration::Bypass => {
return resolve_data_backed_auth_context(
state,
headers,
uri,
auth_endpoint_signature,
)
.await;
}
}
}
}
pub(crate) async fn refresh_execution_runtime_auth_context(
state: &AppState,
auth_context: GatewayControlAuthContext,
@@ -1176,6 +1291,7 @@ fn get_cached_auth_context(state: &AppState, cache_key: &str) -> Option<GatewayC
#[cfg(test)]
mod tests {
use std::sync::Arc;
use std::time::Duration;
use aether_data::repository::auth::{
InMemoryAuthApiKeySnapshotRepository, StoredAuthApiKeySnapshot,
@@ -1188,6 +1304,7 @@ mod tests {
StoredProviderCatalogEndpoint, StoredProviderCatalogProvider,
};
use axum::http::{HeaderMap, Uri};
use futures_util::future::join_all;
use super::{
get_cached_auth_context, resolve_control_decision_auth, resolve_data_backed_auth_context,
@@ -1347,6 +1464,61 @@ mod tests {
assert_eq!(repository.touch_count("key-1"), 1);
}
#[tokio::test]
async fn control_auth_context_singleflights_concurrent_cache_misses() {
let api_key = "sk-test-concurrent-auth-miss";
let repository = Arc::new(
InMemoryAuthApiKeySnapshotRepository::seed(vec![(
Some(hash_api_key(api_key)),
sample_snapshot("key-concurrent-auth-miss", "user-concurrent-auth-miss"),
)])
.with_lookup_delay_for_tests(Duration::from_millis(20)),
);
let data = GatewayDataState::with_auth_api_key_repository_for_tests(repository.clone());
let state = AppState::new()
.expect("state should build")
.with_data_state_for_tests(data);
let mut headers = HeaderMap::new();
headers.insert(
http::header::AUTHORIZATION,
format!("Bearer {api_key}").parse().unwrap(),
);
let request_uri = uri("/v1/chat/completions");
let tasks = (0..32).map(|index| {
let decision = GatewayControlDecision::synthetic(
"/v1/chat/completions",
Some("ai_public".to_string()),
Some("openai".to_string()),
Some("chat".to_string()),
Some("openai:chat".to_string()),
);
let trace_id = format!("trace-concurrent-auth-miss-{index}");
let state = &state;
let headers = &headers;
let request_uri = &request_uri;
async move {
resolve_control_decision_auth(state, headers, request_uri, &trace_id, decision)
.await
}
});
for result in join_all(tasks).await {
let ControlDecisionAuthResolution::Resolved(decision) =
result.expect("auth resolution should succeed");
let auth_context = decision
.auth_context
.expect("auth context should be resolved");
assert_eq!(auth_context.user_id, "user-concurrent-auth-miss");
assert_eq!(auth_context.api_key_id, "key-concurrent-auth-miss");
}
assert_eq!(
repository.key_hash_lookup_count(&hash_api_key(api_key)),
1,
"concurrent cache misses for one auth context should only load one snapshot"
);
}
#[tokio::test]
async fn data_backed_auth_context_marks_wallet_denial_as_not_allowed() {
let api_key = "sk-test-empty-wallet";
@@ -1496,6 +1668,85 @@ mod tests {
assert!(!second.access_allowed);
}
#[tokio::test]
async fn execution_runtime_auth_context_singleflights_concurrent_cache_refreshes() {
let api_key = "sk-test-runtime-auth-refresh";
let auth_repository = Arc::new(
InMemoryAuthApiKeySnapshotRepository::seed(vec![(
Some(hash_api_key(api_key)),
sample_snapshot("key-runtime-auth-refresh", "user-runtime-auth-refresh"),
)])
.with_lookup_delay_for_tests(Duration::from_millis(20)),
);
let data =
GatewayDataState::with_auth_api_key_repository_for_tests(auth_repository.clone());
let state = AppState::new()
.expect("state should build")
.with_data_state_for_tests(data);
let decision = GatewayControlDecision::synthetic(
"/v1/chat/completions",
Some("ai_public".to_string()),
Some("openai".to_string()),
Some("chat".to_string()),
Some("openai:chat".to_string()),
);
let mut headers = HeaderMap::new();
headers.insert("x-api-key", api_key.parse().unwrap());
let request_uri = uri("/v1/chat/completions");
let first = resolve_execution_runtime_auth_context(
&state,
&decision,
&headers,
&request_uri,
"trace-runtime-auth-refresh-prime",
)
.await
.expect("resolution should succeed")
.expect("auth context should exist");
assert_eq!(first.api_key_id, "key-runtime-auth-refresh");
assert_eq!(
auth_repository.key_hash_lookup_count(&hash_api_key(api_key)),
1
);
let tasks = (0..32).map(|index| {
let trace_id = format!("trace-runtime-auth-refresh-{index}");
let state = &state;
let decision = &decision;
let headers = &headers;
let request_uri = &request_uri;
async move {
resolve_execution_runtime_auth_context(
state,
decision,
headers,
request_uri,
&trace_id,
)
.await
}
});
for result in join_all(tasks).await {
let auth_context = result
.expect("resolution should succeed")
.expect("auth context should exist");
assert_eq!(auth_context.user_id, "user-runtime-auth-refresh");
assert_eq!(auth_context.api_key_id, "key-runtime-auth-refresh");
}
assert_eq!(
auth_repository.key_hash_lookup_count(&hash_api_key(api_key)),
1,
"cache refreshes should reuse the existing auth context under concurrent pressure"
);
assert_eq!(
auth_repository.snapshot_lookup_count("key-runtime-auth-refresh"),
1,
"only one cached auth context refresh should read the snapshot by user/key id"
);
}
#[tokio::test]
async fn data_backed_auth_context_allows_provider_id_for_matching_provider_type() {
let api_key = "sk-test-provider-id";
+22 -8
View File
@@ -1748,11 +1748,15 @@ impl GatewayDataState {
api_key_id: &str,
now_unix_secs: u64,
) -> Result<Option<GatewayAuthApiKeySnapshot>, DataLayerError> {
let snapshot = read_resolved_auth_api_key_snapshot_by_user_api_key_ids(
self,
user_id,
api_key_id,
now_unix_secs,
let snapshot = crate::request_diagnostics::observe_db_operation(
"auth_api_key_snapshot",
self.database_pool_summary(),
read_resolved_auth_api_key_snapshot_by_user_api_key_ids(
self,
user_id,
api_key_id,
now_unix_secs,
),
)
.await?;
self.apply_user_group_effective_policies(snapshot).await
@@ -1763,8 +1767,12 @@ impl GatewayDataState {
key_hash: &str,
now_unix_secs: u64,
) -> Result<Option<GatewayAuthApiKeySnapshot>, DataLayerError> {
let snapshot =
read_resolved_auth_api_key_snapshot_by_key_hash(self, key_hash, now_unix_secs).await?;
let snapshot = crate::request_diagnostics::observe_db_operation(
"auth_api_key_snapshot_by_hash",
self.database_pool_summary(),
read_resolved_auth_api_key_snapshot_by_key_hash(self, key_hash, now_unix_secs),
)
.await?;
self.apply_user_group_effective_policies(snapshot).await
}
@@ -1782,7 +1790,13 @@ impl GatewayDataState {
let Some(repository) = self.user_reader.as_ref() else {
return Ok(Some(snapshot));
};
let Some(user) = repository.find_user_auth_by_id(&snapshot.user_id).await? else {
let Some(user) = crate::request_diagnostics::observe_db_operation(
"auth_user_policy",
self.database_pool_summary(),
repository.find_user_auth_by_id(&snapshot.user_id),
)
.await?
else {
return Ok(Some(snapshot));
};
if user.role.eq_ignore_ascii_case("admin") && !snapshot.api_key_is_standalone {
+11 -4
View File
@@ -107,10 +107,17 @@ impl GatewayDataState {
&self,
candidate: UpsertRequestCandidateRecord,
) -> Result<Option<StoredRequestCandidate>, DataLayerError> {
match &self.request_candidate_writer {
Some(repository) => repository.upsert(candidate).await.map(Some),
None => Ok(None),
}
crate::request_diagnostics::observe_db_operation(
"request_candidate_upsert",
self.database_pool_summary(),
async {
match &self.request_candidate_writer {
Some(repository) => repository.upsert(candidate).await.map(Some),
None => Ok(None),
}
},
)
.await
}
pub(crate) async fn delete_request_candidates_created_before(
+126 -24
View File
@@ -3,17 +3,77 @@ use aether_data_contracts::repository::candidate_selection::MinimalCandidateSele
use aether_data_contracts::repository::candidates::RequestCandidateReadRepository;
use aether_data_contracts::repository::provider_catalog::ProviderCatalogReadRepository;
use aether_runtime_state::RuntimeQueueStore;
use std::sync::{Arc, OnceLock};
use std::collections::HashSet;
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::sync::Mutex;
use tokio::sync::Notify;
use super::{GatewayDataConfig, GatewayDataState, StoredSystemConfigEntry};
const SYSTEM_CONFIG_VALUE_CACHE_TTL: Duration = Duration::from_secs(30);
fn system_config_value_load_guard() -> &'static Mutex<()> {
static GUARD: OnceLock<Mutex<()>> = OnceLock::new();
GUARD.get_or_init(|| Mutex::new(()))
fn system_config_value_load_state() -> &'static SystemConfigValueLoadState {
static STATE: std::sync::OnceLock<SystemConfigValueLoadState> = std::sync::OnceLock::new();
STATE.get_or_init(SystemConfigValueLoadState::default)
}
#[derive(Debug, Default)]
struct SystemConfigValueLoadState {
inflight: std::sync::Mutex<HashSet<String>>,
notify: Notify,
}
enum SystemConfigValueLoadRegistration<'a> {
Leader(SystemConfigValueLoadGuard<'a>),
Follower,
Bypass,
}
struct SystemConfigValueLoadGuard<'a> {
state: &'a SystemConfigValueLoadState,
key: Option<String>,
}
impl Drop for SystemConfigValueLoadGuard<'_> {
fn drop(&mut self) {
if let Some(key) = self.key.take() {
self.state.finish(&key);
}
}
}
impl SystemConfigValueLoadState {
fn register(&self, key: &str) -> SystemConfigValueLoadRegistration<'_> {
match self.inflight.lock() {
Ok(mut inflight) => {
if inflight.contains(key) {
SystemConfigValueLoadRegistration::Follower
} else {
inflight.insert(key.to_string());
SystemConfigValueLoadRegistration::Leader(SystemConfigValueLoadGuard {
state: self,
key: Some(key.to_string()),
})
}
}
Err(_) => SystemConfigValueLoadRegistration::Bypass,
}
}
fn notified(&self) -> tokio::sync::futures::Notified<'_> {
self.notify.notified()
}
fn finish(&self, key: &str) {
let removed = self
.inflight
.lock()
.map(|mut inflight| inflight.remove(key))
.unwrap_or(false);
if removed {
self.notify.notify_waiters();
}
}
}
fn current_system_config_updated_at_unix_secs() -> u64 {
@@ -447,27 +507,69 @@ impl GatewayDataState {
return Ok(value);
}
}
let _guard = system_config_value_load_guard().lock().await;
let cached_value = self
.system_config_value_cache
.read()
.expect("system config value cache lock")
.get(key)
.cloned();
if let Some((cached_at, value)) = cached_value {
if cached_at.elapsed() <= SYSTEM_CONFIG_VALUE_CACHE_TTL {
return Ok(value);
let load_state = system_config_value_load_state();
loop {
let notified = load_state.notified();
match load_state.register(key) {
SystemConfigValueLoadRegistration::Bypass => {
let Some(backends) = self.backends.as_ref() else {
return Ok(None);
};
let value = crate::request_diagnostics::observe_db_operation(
"system_config_value",
self.database_pool_summary(),
backends.find_system_config_value(key),
)
.await?;
self.system_config_value_cache
.write()
.expect("system config value cache lock")
.insert(key.to_string(), (Instant::now(), value.clone()));
return Ok(value);
}
SystemConfigValueLoadRegistration::Follower => {
notified.await;
let cached_value = self
.system_config_value_cache
.read()
.expect("system config value cache lock")
.get(key)
.cloned();
if let Some((cached_at, value)) = cached_value {
if cached_at.elapsed() <= SYSTEM_CONFIG_VALUE_CACHE_TTL {
return Ok(value);
}
}
}
SystemConfigValueLoadRegistration::Leader(_guard) => {
let cached_value = self
.system_config_value_cache
.read()
.expect("system config value cache lock")
.get(key)
.cloned();
if let Some((cached_at, value)) = cached_value {
if cached_at.elapsed() <= SYSTEM_CONFIG_VALUE_CACHE_TTL {
return Ok(value);
}
}
let Some(backends) = self.backends.as_ref() else {
return Ok(None);
};
let value = crate::request_diagnostics::observe_db_operation(
"system_config_value",
self.database_pool_summary(),
backends.find_system_config_value(key),
)
.await?;
self.system_config_value_cache
.write()
.expect("system config value cache lock")
.insert(key.to_string(), (Instant::now(), value.clone()));
return Ok(value);
}
}
}
let Some(backends) = self.backends.as_ref() else {
return Ok(None);
};
let value = backends.find_system_config_value(key).await?;
self.system_config_value_cache
.write()
.expect("system config value cache lock")
.insert(key.to_string(), (Instant::now(), value.clone()));
Ok(value)
}
pub(crate) async fn upsert_system_config_value(
+10 -3
View File
@@ -197,9 +197,7 @@ pub(crate) struct GatewayDataState {
settlement_writer: Option<Arc<dyn SettlementWriteRepository>>,
system_config_values: Option<Arc<RwLock<BTreeMap<String, StoredSystemConfigEntry>>>>,
system_config_value_cache: Arc<RwLock<BTreeMap<String, (Instant, Option<serde_json::Value>)>>>,
billing_model_context_cache: Arc<
RwLock<HashMap<BillingModelContextCacheKey, (Instant, Option<StoredBillingModelContext>)>>,
>,
billing_model_context_cache: Arc<BillingModelContextCacheState>,
}
#[derive(Clone, Debug, Eq, Hash, PartialEq)]
@@ -216,6 +214,15 @@ pub(super) enum BillingModelContextCacheKey {
},
}
#[derive(Default)]
pub(super) struct BillingModelContextCacheState {
pub(super) entries:
RwLock<HashMap<BillingModelContextCacheKey, (Instant, Option<StoredBillingModelContext>)>>,
pub(super) inflight: std::sync::Mutex<HashMap<BillingModelContextCacheKey, u64>>,
pub(super) inflight_notify: tokio::sync::Notify,
pub(super) next_inflight_token: std::sync::atomic::AtomicU64,
}
impl fmt::Debug for GatewayDataState {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("GatewayDataState")
+62 -28
View File
@@ -15,14 +15,24 @@ impl GatewayDataState {
api_format: &str,
global_model_name: &str,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
match &self.minimal_candidate_selection_reader {
Some(repository) => {
repository
.list_for_exact_api_format_and_global_model(api_format, global_model_name)
.await
}
None => Ok(Vec::new()),
}
crate::request_diagnostics::observe_db_operation(
"candidate_selection",
self.database_pool_summary(),
async {
match &self.minimal_candidate_selection_reader {
Some(repository) => {
repository
.list_for_exact_api_format_and_global_model(
api_format,
global_model_name,
)
.await
}
None => Ok(Vec::new()),
}
},
)
.await
}
pub(crate) async fn list_minimal_candidate_selection_rows_for_requested_model(
@@ -30,38 +40,62 @@ impl GatewayDataState {
api_format: &str,
requested_model_name: &str,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
match &self.minimal_candidate_selection_reader {
Some(repository) => {
repository
.list_for_exact_api_format_and_requested_model(api_format, requested_model_name)
.await
}
None => Ok(Vec::new()),
}
crate::request_diagnostics::observe_db_operation(
"candidate_selection",
self.database_pool_summary(),
async {
match &self.minimal_candidate_selection_reader {
Some(repository) => {
repository
.list_for_exact_api_format_and_requested_model(
api_format,
requested_model_name,
)
.await
}
None => Ok(Vec::new()),
}
},
)
.await
}
pub(crate) async fn list_minimal_candidate_selection_rows_for_requested_model_page(
&self,
query: &StoredRequestedModelCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
match &self.minimal_candidate_selection_reader {
Some(repository) => {
repository
.list_for_exact_api_format_and_requested_model_page(query)
.await
}
None => Ok(Vec::new()),
}
crate::request_diagnostics::observe_db_operation(
"candidate_selection",
self.database_pool_summary(),
async {
match &self.minimal_candidate_selection_reader {
Some(repository) => {
repository
.list_for_exact_api_format_and_requested_model_page(query)
.await
}
None => Ok(Vec::new()),
}
},
)
.await
}
pub(crate) async fn list_minimal_candidate_selection_rows_for_api_format(
&self,
api_format: &str,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
match &self.minimal_candidate_selection_reader {
Some(repository) => repository.list_for_exact_api_format(api_format).await,
None => Ok(Vec::new()),
}
crate::request_diagnostics::observe_db_operation(
"candidate_selection",
self.database_pool_summary(),
async {
match &self.minimal_candidate_selection_reader {
Some(repository) => repository.list_for_exact_api_format(api_format).await,
None => Ok(Vec::new()),
}
},
)
.await
}
pub(crate) async fn list_pool_key_candidate_rows_for_group(
+343 -25
View File
@@ -43,6 +43,7 @@ use aether_data_contracts::repository::usage::{
use aether_runtime_state::RuntimeQueueStore;
use aether_video_tasks_core::read_data_backed_video_task_response;
use std::time::{Duration, Instant};
use tokio::time::timeout;
fn normalize_billing_context_cache_part(value: &str) -> String {
value.trim().to_string()
@@ -55,12 +56,51 @@ fn normalize_optional_billing_context_cache_part(value: Option<&str>) -> Option<
.map(ToOwned::to_owned)
}
enum BillingModelContextInflightRegistration {
Leader(u64),
Follower,
Bypass,
}
struct BillingModelContextInflightGuard<'a> {
state: &'a GatewayDataState,
key: Option<BillingModelContextCacheKey>,
token: u64,
}
impl<'a> BillingModelContextInflightGuard<'a> {
fn new(state: &'a GatewayDataState, key: BillingModelContextCacheKey, token: u64) -> Self {
Self {
state,
key: Some(key),
token,
}
}
fn finish(&mut self) {
if let Some(key) = self.key.take() {
self.state
.finish_billing_model_context_inflight(&key, self.token);
}
}
}
impl Drop for BillingModelContextInflightGuard<'_> {
fn drop(&mut self) {
self.finish();
}
}
impl GatewayDataState {
const MAINTENANCE_POOL_IDLE_RESERVE_ENV: &'static str =
"AETHER_GATEWAY_MAINTENANCE_POOL_IDLE_RESERVE";
const MAINTENANCE_POOL_PRESSURE_MAX_DEFER: Duration = Duration::from_secs(30);
const BILLING_MODEL_CONTEXT_CACHE_TTL: Duration = Duration::from_secs(30);
const BILLING_MODEL_CONTEXT_CACHE_MAX_ENTRIES: usize = 4096;
#[cfg(not(test))]
const BILLING_MODEL_CONTEXT_CACHE_INFLIGHT_WAIT_TIMEOUT: Duration = Duration::from_secs(10);
#[cfg(test)]
const BILLING_MODEL_CONTEXT_CACHE_INFLIGHT_WAIT_TIMEOUT: Duration = Duration::from_millis(100);
pub(crate) async fn run_database_maintenance(
&self,
@@ -1041,10 +1081,17 @@ impl GatewayDataState {
&self,
usage: UpsertUsageRecord,
) -> Result<Option<StoredRequestUsageAudit>, DataLayerError> {
match &self.usage_writer {
Some(repository) => repository.upsert(usage).await.map(Some),
None => Ok(None),
}
crate::request_diagnostics::observe_db_operation(
"usage_upsert",
self.database_pool_summary(),
async {
match &self.usage_writer {
Some(repository) => repository.upsert(usage).await.map(Some),
None => Ok(None),
}
},
)
.await
}
#[allow(dead_code)]
@@ -1710,17 +1757,47 @@ impl GatewayDataState {
if let Some(value) = self.cached_billing_model_context(&key) {
return Ok(value);
}
match &self.billing_reader {
Some(repository) => {
let value = repository
.find_model_context(provider_id, provider_api_key_id, global_model_name)
.await?;
self.remember_billing_model_context(key, value.clone());
Ok(value)
}
None => {
self.remember_billing_model_context(key, None);
Ok(None)
loop {
let notified = self.billing_model_context_cache.inflight_notify.notified();
match self.register_billing_model_context_inflight(&key) {
BillingModelContextInflightRegistration::Bypass => {
return self
.load_billing_model_context_by_name(
key,
provider_id,
provider_api_key_id,
global_model_name,
)
.await;
}
BillingModelContextInflightRegistration::Follower => {
if timeout(
Self::BILLING_MODEL_CONTEXT_CACHE_INFLIGHT_WAIT_TIMEOUT,
notified,
)
.await
.is_err()
{
self.expire_billing_model_context_inflight(&key);
}
if let Some(value) = self.cached_billing_model_context(&key) {
return Ok(value);
}
continue;
}
BillingModelContextInflightRegistration::Leader(token) => {
let mut guard = BillingModelContextInflightGuard::new(self, key.clone(), token);
let result = self
.load_billing_model_context_by_name(
key,
provider_id,
provider_api_key_id,
global_model_name,
)
.await;
guard.finish();
return result;
}
}
}
}
@@ -1739,18 +1816,164 @@ impl GatewayDataState {
if let Some(value) = self.cached_billing_model_context(&key) {
return Ok(value);
}
match &self.billing_reader {
Some(repository) => {
let value = repository
.find_model_context_by_model_id(provider_id, provider_api_key_id, model_id)
.await?;
self.remember_billing_model_context(key, value.clone());
Ok(value)
loop {
let notified = self.billing_model_context_cache.inflight_notify.notified();
match self.register_billing_model_context_inflight(&key) {
BillingModelContextInflightRegistration::Bypass => {
return self
.load_billing_model_context_by_model_id(
key,
provider_id,
provider_api_key_id,
model_id,
)
.await;
}
BillingModelContextInflightRegistration::Follower => {
if timeout(
Self::BILLING_MODEL_CONTEXT_CACHE_INFLIGHT_WAIT_TIMEOUT,
notified,
)
.await
.is_err()
{
self.expire_billing_model_context_inflight(&key);
}
if let Some(value) = self.cached_billing_model_context(&key) {
return Ok(value);
}
continue;
}
BillingModelContextInflightRegistration::Leader(token) => {
let mut guard = BillingModelContextInflightGuard::new(self, key.clone(), token);
let result = self
.load_billing_model_context_by_model_id(
key,
provider_id,
provider_api_key_id,
model_id,
)
.await;
guard.finish();
return result;
}
}
None => {
self.remember_billing_model_context(key, None);
Ok(None)
}
}
async fn load_billing_model_context_by_name(
&self,
key: BillingModelContextCacheKey,
provider_id: &str,
provider_api_key_id: Option<&str>,
global_model_name: &str,
) -> Result<Option<StoredBillingModelContext>, DataLayerError> {
crate::request_diagnostics::observe_db_operation(
"billing_model_context",
self.database_pool_summary(),
async {
match &self.billing_reader {
Some(repository) => {
let value = repository
.find_model_context(provider_id, provider_api_key_id, global_model_name)
.await?;
self.remember_billing_model_context(key, value.clone());
Ok(value)
}
None => {
self.remember_billing_model_context(key, None);
Ok(None)
}
}
},
)
.await
}
async fn load_billing_model_context_by_model_id(
&self,
key: BillingModelContextCacheKey,
provider_id: &str,
provider_api_key_id: Option<&str>,
model_id: &str,
) -> Result<Option<StoredBillingModelContext>, DataLayerError> {
crate::request_diagnostics::observe_db_operation(
"billing_model_context",
self.database_pool_summary(),
async {
match &self.billing_reader {
Some(repository) => {
let value = repository
.find_model_context_by_model_id(
provider_id,
provider_api_key_id,
model_id,
)
.await?;
self.remember_billing_model_context(key, value.clone());
Ok(value)
}
None => {
self.remember_billing_model_context(key, None);
Ok(None)
}
}
},
)
.await
}
fn register_billing_model_context_inflight(
&self,
key: &BillingModelContextCacheKey,
) -> BillingModelContextInflightRegistration {
match self.billing_model_context_cache.inflight.lock() {
Ok(mut inflight) => {
if inflight.contains_key(key) {
return BillingModelContextInflightRegistration::Follower;
}
let token = self
.billing_model_context_cache
.next_inflight_token
.fetch_add(1, std::sync::atomic::Ordering::AcqRel);
inflight.insert(key.clone(), token);
BillingModelContextInflightRegistration::Leader(token)
}
Err(_) => BillingModelContextInflightRegistration::Bypass,
}
}
fn finish_billing_model_context_inflight(&self, key: &BillingModelContextCacheKey, token: u64) {
let mut removed = false;
if let Ok(mut inflight) = self.billing_model_context_cache.inflight.lock() {
if inflight.get(key).copied() == Some(token) {
inflight.remove(key);
removed = true;
}
}
if removed {
self.billing_model_context_cache
.inflight_notify
.notify_waiters();
}
}
fn expire_billing_model_context_inflight(&self, key: &BillingModelContextCacheKey) {
let mut removed = false;
if let Ok(mut inflight) = self.billing_model_context_cache.inflight.lock() {
removed = inflight.remove(key).is_some();
}
if removed {
tracing::warn!(
event_name = "billing_model_context_cache_inflight_expired",
log_type = "ops",
cache_key = ?key,
wait_timeout_ms = Self::BILLING_MODEL_CONTEXT_CACHE_INFLIGHT_WAIT_TIMEOUT.as_millis() as u64,
"gateway billing model context cache expired stale inflight load"
);
self.billing_model_context_cache
.inflight_notify
.notify_waiters();
}
}
@@ -1759,6 +1982,7 @@ impl GatewayDataState {
key: &BillingModelContextCacheKey,
) -> Option<Option<StoredBillingModelContext>> {
self.billing_model_context_cache
.entries
.read()
.expect("billing model context cache lock")
.get(key)
@@ -1775,6 +1999,7 @@ impl GatewayDataState {
) {
let mut cache = self
.billing_model_context_cache
.entries
.write()
.expect("billing model context cache lock");
cache.retain(|_, (cached_at, _)| {
@@ -1794,9 +2019,25 @@ impl GatewayDataState {
fn clear_billing_model_context_cache(&self) {
self.billing_model_context_cache
.entries
.write()
.expect("billing model context cache lock")
.clear();
let mut cleared_inflight = false;
if let Ok(mut inflight) = self.billing_model_context_cache.inflight.lock() {
cleared_inflight = !inflight.is_empty();
inflight.clear();
}
if cleared_inflight {
tracing::warn!(
event_name = "billing_model_context_cache_inflight_cleared",
log_type = "ops",
"gateway billing model context cache cleared in-flight loads"
);
self.billing_model_context_cache
.inflight_notify
.notify_waiters();
}
}
pub(crate) async fn admin_billing_enabled_default_value_exists(
@@ -2211,12 +2452,89 @@ impl GatewayDataState {
#[cfg(test)]
mod tests {
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use std::time::Duration;
use aether_data::repository::users::{InMemoryUserReadRepository, StoredUserExportRow};
use aether_data_contracts::repository::billing::{
BillingReadRepository, StoredBillingModelContext,
};
use async_trait::async_trait;
use serde_json::json;
use super::GatewayDataState;
struct SlowBillingContextRepository {
calls: AtomicUsize,
context: StoredBillingModelContext,
}
#[async_trait]
impl BillingReadRepository for SlowBillingContextRepository {
async fn find_model_context(
&self,
_provider_id: &str,
_provider_api_key_id: Option<&str>,
_global_model_name: &str,
) -> Result<Option<StoredBillingModelContext>, aether_data_contracts::DataLayerError>
{
self.calls.fetch_add(1, Ordering::AcqRel);
tokio::time::sleep(Duration::from_millis(25)).await;
Ok(Some(self.context.clone()))
}
}
fn billing_context() -> StoredBillingModelContext {
StoredBillingModelContext::new(
"provider-1".to_string(),
Some("pay_as_you_go".to_string()),
Some("key-1".to_string()),
None,
None,
"global-model-1".to_string(),
"gpt-5".to_string(),
None,
Some(0.02),
Some(json!({"tiers":[{"up_to":null,"input_price_per_1m":3.0,"output_price_per_1m":15.0}]})),
Some("model-1".to_string()),
Some("gpt-5-upstream".to_string()),
None,
None,
None,
)
.expect("billing context should build")
}
#[tokio::test]
async fn billing_model_context_cache_coalesces_concurrent_loads() {
let repository = Arc::new(SlowBillingContextRepository {
calls: AtomicUsize::new(0),
context: billing_context(),
});
let state = Arc::new(GatewayDataState::with_billing_reader_for_tests(
repository.clone(),
));
let mut tasks = Vec::new();
for _ in 0..16 {
let state = Arc::clone(&state);
tasks.push(tokio::spawn(async move {
state
.find_billing_model_context("provider-1", Some("key-1"), "gpt-5")
.await
.expect("billing context lookup should succeed")
.expect("billing context should exist");
}));
}
for task in tasks {
task.await.expect("lookup task should complete");
}
assert_eq!(repository.calls.load(Ordering::Acquire), 1);
}
#[tokio::test]
async fn lists_non_admin_export_users_from_user_reader() {
let repository = Arc::new(InMemoryUserReadRepository::seed_export_users(vec![
@@ -47,6 +47,7 @@ use crate::handlers::shared::provider_pool::{
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;
use crate::stage_metrics::observe_gateway_stage_ms;
static LOAD_BALANCE_SEQUENCE: AtomicU64 = AtomicU64::new(0);
static POOL_SCORE_SCHEDULE_INTEREST_SEMAPHORE: LazyLock<Arc<Semaphore>> =
@@ -133,14 +134,20 @@ async fn schedule_pool_page_candidates(
let runtime = if key_ids.is_empty() {
AdminProviderPoolRuntimeState::default()
} else {
read_admin_provider_pool_runtime_state(
let runtime_started_at = std::time::Instant::now();
let runtime = read_admin_provider_pool_runtime_state(
state.app().runtime_state.as_ref(),
provider_id.as_str(),
&key_ids,
&pool_config,
sticky_session_token,
)
.await
.await;
observe_gateway_stage_ms(
"pool_runtime_state",
runtime_started_at.elapsed().as_millis() as u64,
);
runtime
};
pool_config_by_provider.insert(provider_id.clone(), pool_config);
runtime_by_provider.insert(provider_id, runtime);
@@ -449,9 +456,17 @@ impl<'a> PoolKeyCursor<'a> {
}
pub(crate) async fn next_key(&mut self) -> Option<EligibleLocalExecutionCandidate> {
let started_at = std::time::Instant::now();
let mut observed = false;
loop {
if let Some(candidate) = self.next_queued_candidate().await {
self.returned_key_count = self.returned_key_count.saturating_add(1);
if !observed {
observe_gateway_stage_ms(
"pool_cursor_next_key",
started_at.elapsed().as_millis() as u64,
);
}
return Some(candidate);
}
@@ -464,6 +479,13 @@ impl<'a> PoolKeyCursor<'a> {
}
if !self.refill_queued_candidates().await {
if !observed {
observe_gateway_stage_ms(
"pool_cursor_next_key",
started_at.elapsed().as_millis() as u64,
);
observed = true;
}
return None;
}
}
@@ -634,9 +656,14 @@ impl<'a> PoolKeyCursor<'a> {
offset: 0,
limit: limit as usize,
};
let score_started_at = std::time::Instant::now();
let scores = match self.state.app().data.list_ranked_pool_members(&query).await {
Ok(scores) => scores,
Err(err) => {
observe_gateway_stage_ms(
"pool_score_load",
score_started_at.elapsed().as_millis() as u64,
);
warn!(
event_name = "pool_group_score_load_failed",
log_type = "event",
@@ -650,6 +677,10 @@ impl<'a> PoolKeyCursor<'a> {
return None;
}
};
observe_gateway_stage_ms(
"pool_score_load",
score_started_at.elapsed().as_millis() as u64,
);
if scores.is_empty() {
return None;
}
@@ -668,6 +699,7 @@ impl<'a> PoolKeyCursor<'a> {
selected_provider_model_name: self.group.candidate.selected_provider_model_name.clone(),
key_ids,
};
let rows_started_at = std::time::Instant::now();
let rows = match self
.state
.app()
@@ -676,6 +708,10 @@ impl<'a> PoolKeyCursor<'a> {
{
Ok(rows) => rows,
Err(err) => {
observe_gateway_stage_ms(
"pool_score_key_rows",
rows_started_at.elapsed().as_millis() as u64,
);
warn!(
event_name = "pool_group_score_key_load_failed",
log_type = "event",
@@ -690,6 +726,10 @@ impl<'a> PoolKeyCursor<'a> {
return None;
}
};
observe_gateway_stage_ms(
"pool_score_key_rows",
rows_started_at.elapsed().as_millis() as u64,
);
let materialized_row_count = rows.len() as u32;
let missing_score_count = scores.len().saturating_sub(rows.len());
if missing_score_count > 0 {
@@ -1012,11 +1052,20 @@ impl<'a> PoolKeyCursor<'a> {
return None;
}
let transport_started_at = std::time::Instant::now();
let Some(transport) = read_candidate_transport_snapshot(self.state, &candidate).await
else {
observe_gateway_stage_ms(
"candidate_transport_snapshot",
transport_started_at.elapsed().as_millis() as u64,
);
self.record_skip_reason("transport_snapshot_missing");
return None;
};
observe_gateway_stage_ms(
"candidate_transport_snapshot",
transport_started_at.elapsed().as_millis() as u64,
);
if let Some(skip_reason) =
candidate_auth_channel_skip_reason(&transport, self.request_auth_channel.as_deref())
{
+78
View File
@@ -24,6 +24,11 @@ pub(crate) enum GatewayError {
phase: &'static str,
timeout_ms: u64,
},
AdmissionTimeout {
trace_id: String,
gate: &'static str,
queue_budget_ms: u64,
},
Client {
status: StatusCode,
message: String,
@@ -43,6 +48,13 @@ impl GatewayError {
} => {
format!("local execution planning timed out in {phase} after {timeout_ms}ms")
}
Self::AdmissionTimeout {
gate,
queue_budget_ms,
..
} => {
format!("gateway admission gate {gate} timed out after {queue_budget_ms}ms")
}
}
}
}
@@ -113,6 +125,34 @@ impl IntoResponse for GatewayError {
);
response
}
Self::AdmissionTimeout {
trace_id,
gate,
queue_budget_ms,
} => {
tracing::debug!(
trace_id = %trace_id,
gate,
queue_budget_ms,
"gateway admission gate timed out"
);
let body = Json(json!({
"error": {
"message": "gateway admission queue timed out",
"trace_id": trace_id,
}
}));
let mut response = (StatusCode::TOO_MANY_REQUESTS, body).into_response();
let _ =
insert_header_if_missing(response.headers_mut(), TRACE_ID_HEADER, &trace_id);
let _ = insert_header_if_missing(
response.headers_mut(),
GATEWAY_HEADER,
"rust-phase3b",
);
let _ = insert_header_if_missing(response.headers_mut(), "Retry-After", "1");
response
}
Self::Client { status, message } => (
status,
Json(json!({
@@ -140,3 +180,41 @@ impl From<AiSurfaceFinalizeError> for GatewayError {
GatewayError::Internal(error.0)
}
}
#[cfg(test)]
mod tests {
use axum::http::{header::RETRY_AFTER, StatusCode};
use axum::response::IntoResponse;
use crate::constants::TRACE_ID_HEADER;
use super::GatewayError;
#[test]
fn admission_timeout_returns_429_with_retry_after_without_panicking() {
let trace_id = "trace-admission-timeout".to_string();
let response = GatewayError::AdmissionTimeout {
trace_id: trace_id.clone(),
gate: "gateway_upstream_execution",
queue_budget_ms: 250,
}
.into_response();
assert_eq!(response.status(), StatusCode::TOO_MANY_REQUESTS);
assert_eq!(
response
.headers()
.get(RETRY_AFTER)
.and_then(|v| v.to_str().ok()),
Some("1")
);
assert_eq!(
response
.headers()
.get(TRACE_ID_HEADER)
.and_then(|v| v.to_str().ok()),
Some(trace_id.as_str())
);
}
}
File diff suppressed because it is too large Load Diff
@@ -26,6 +26,7 @@ use crate::orchestration::{
LocalPoolErrorEffect,
};
use crate::request_candidate_runtime::record_report_request_candidate_status;
use crate::request_diagnostics::attach_current_request_diagnostics_to_report_context;
use crate::usage::submit_sync_report;
use crate::{usage::GatewaySyncReportRequest, AppState, GatewayError};
@@ -319,7 +320,12 @@ async fn record_stream_sync_failure(
}),
)
.await;
let context_seed = build_terminal_usage_context_seed(plan, report_context);
let report_context_with_diagnostics =
attach_current_request_diagnostics_to_report_context(report_context);
let context_seed = build_terminal_usage_context_seed(
plan,
report_context_with_diagnostics.as_ref().or(report_context),
);
let payload_seed = build_sync_terminal_usage_payload_seed(payload);
state
.usage_runtime
@@ -39,7 +39,9 @@ pub(crate) fn build_direct_execution_frame_stream(
response,
started_at,
stream_first_byte_timeout,
upstream_target_permit,
} = execution;
let _upstream_target_permit = upstream_target_permit;
let mut observer_context = stream_summary_report_context;
if observer_context
@@ -79,6 +79,7 @@ use crate::request_candidate_runtime::{
ensure_execution_request_candidate_slot, record_local_request_candidate_extra_data,
record_local_request_candidate_status,
};
use crate::request_diagnostics::attach_current_request_diagnostics_to_report_context;
use crate::usage::{spawn_sync_report, submit_sync_report};
use crate::video_tasks::VideoTaskSyncReportMode;
use crate::{usage::GatewaySyncReportRequest, AppState, GatewayError};
@@ -296,7 +297,12 @@ fn record_sync_terminal_usage(
report_context: Option<&serde_json::Value>,
payload: &GatewaySyncReportRequest,
) {
let context_seed = build_terminal_usage_context_seed(plan, report_context);
let report_context_with_diagnostics =
attach_current_request_diagnostics_to_report_context(report_context);
let context_seed = build_terminal_usage_context_seed(
plan,
report_context_with_diagnostics.as_ref().or(report_context),
);
let payload_seed = build_sync_terminal_usage_payload_seed(payload);
state
.usage_runtime
@@ -3,6 +3,7 @@ use std::error::Error as _;
use std::future::Future;
use std::io::Read;
use std::io::Write;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{LazyLock, Mutex as StdMutex};
use std::time::{Duration, Instant};
@@ -11,10 +12,11 @@ use aether_contracts::{
ResponseBody, EXECUTION_REQUEST_ACCEPT_INVALID_CERTS_HEADER,
EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER, EXECUTION_REQUEST_HTTP1_ONLY_HEADER,
TRANSPORT_BACKEND_BROWSER_WREQ, TRANSPORT_BACKEND_REQWEST_RUSTLS,
TRANSPORT_HTTP_MODE_HTTP1_ONLY,
TRANSPORT_HTTP_MODE_H2C_PRIOR_KNOWLEDGE, TRANSPORT_HTTP_MODE_HTTP1_ONLY,
};
use aether_data::repository::proxy_nodes::ProxyNodeTrafficMutation;
use aether_http::{apply_http_client_config, HttpClientConfig};
use aether_runtime::{MetricKind, MetricSample};
use axum::body::Bytes;
use base64::Engine as _;
use flate2::read::{DeflateDecoder, GzDecoder};
@@ -35,7 +37,9 @@ 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,
};
use crate::stage_metrics::observe_gateway_stage_ms;
use crate::tunnel::{self, tunnel_protocol};
use crate::upstream_admission::UpstreamTargetAdmissionPermit;
use crate::{AppState, GatewayError};
const HUB_RELAY_CONTENT_TYPE: &str = "application/vnd.aether.tunnel-envelope";
@@ -46,18 +50,88 @@ const DEFAULT_STREAM_FIRST_BYTE_TIMEOUT_MS: u64 = 30_000;
const DEFAULT_NON_STREAM_TOTAL_TIMEOUT_MS: u64 = 300_000;
const MIN_TUNNEL_TIMEOUT_SECS: u64 = 1;
const MAX_TUNNEL_TIMEOUT_SECS: u64 = 300;
const DIRECT_REQWEST_H2_CLIENT_SHARDS_ENV: &str = "AETHER_GATEWAY_DIRECT_REQWEST_H2_CLIENT_SHARDS";
const DIRECT_REQWEST_H2_TARGET_STREAMS_PER_CLIENT_ENV: &str =
"AETHER_GATEWAY_DIRECT_REQWEST_H2_TARGET_STREAMS_PER_CLIENT";
const DIRECT_REQWEST_SYNC_WARM_CLIENTS_ENV: &str =
"AETHER_GATEWAY_DIRECT_REQWEST_SYNC_WARM_CLIENTS";
const UPSTREAM_TARGET_GATE_LIMIT_ENV: &str = "AETHER_GATEWAY_UPSTREAM_TARGET_GATE_LIMIT";
const DEFAULT_UPSTREAM_TARGET_GATE_LIMIT: usize = 2_000;
const DEFAULT_H2_TARGET_STREAMS_PER_CLIENT: usize = 200;
const DEFAULT_DIRECT_REQWEST_SYNC_WARM_CLIENTS: usize = 4;
const MAX_DIRECT_REQWEST_H2_CLIENT_SHARDS: usize = 128;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
struct DirectReqwestClientCacheKey {
connect_timeout_ms: Option<u64>,
proxy_url: Option<String>,
follow_redirects: bool,
http1_only: bool,
accept_invalid_certs: bool,
transport_profile: Option<DirectReqwestTransportProfileCacheKey>,
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
struct DirectReqwestTransportProfileCacheKey {
profile_id: String,
backend: String,
http_mode: String,
pool_scope: String,
header_fingerprint: Option<String>,
extra: Option<String>,
}
struct DirectReqwestClientCacheEntry {
clients: Vec<reqwest::Client>,
next: AtomicU64,
target_len: usize,
warming: bool,
}
impl DirectReqwestClientCacheEntry {
fn new(clients: Vec<reqwest::Client>, target_len: usize, warming: bool) -> Self {
Self {
clients,
next: AtomicU64::new(0),
target_len: target_len.max(1),
warming,
}
}
fn select(&self) -> reqwest::Client {
if self.clients.len() <= 1 {
return self
.clients
.first()
.expect("direct reqwest client cache entry should contain a client")
.clone();
}
let index = self.next.fetch_add(1, Ordering::Relaxed) as usize % self.clients.len();
self.clients[index].clone()
}
fn len(&self) -> usize {
self.clients.len()
}
fn should_warm(&self) -> bool {
self.clients.len() < self.target_len && !self.warming
}
}
static DIRECT_REQWEST_CLIENT_CACHE: LazyLock<
StdMutex<HashMap<DirectReqwestClientCacheKey, reqwest::Client>>,
StdMutex<HashMap<DirectReqwestClientCacheKey, DirectReqwestClientCacheEntry>>,
> = LazyLock::new(|| StdMutex::new(HashMap::new()));
#[derive(Debug, Default)]
struct DirectReqwestClientCacheMetrics {
hits: AtomicU64,
misses: AtomicU64,
builds: AtomicU64,
}
static DIRECT_REQWEST_CLIENT_CACHE_METRICS: LazyLock<DirectReqwestClientCacheMetrics> =
LazyLock::new(DirectReqwestClientCacheMetrics::default);
pub(crate) fn format_upstream_request_error(err: &reqwest::Error) -> String {
let mut kinds = Vec::new();
if err.is_connect() {
@@ -244,6 +318,7 @@ pub(crate) struct DirectUpstreamStreamExecution {
pub(crate) response: DirectUpstreamResponse,
pub(crate) started_at: Instant,
pub(crate) stream_first_byte_timeout: Option<Duration>,
pub(crate) upstream_target_permit: Option<UpstreamTargetAdmissionPermit>,
}
impl DirectSyncExecutionRuntime {
@@ -302,10 +377,19 @@ impl DirectSyncExecutionRuntime {
return Err(ExecutionRuntimeTransportError::StreamUnsupported);
}
let build_body_started_at = Instant::now();
let body_bytes = build_request_body(plan)?;
observe_gateway_stage_ms(
"direct_build_body",
build_body_started_at.elapsed().as_millis() as u64,
);
let started_at = Instant::now();
let response = send_request(plan, body_bytes).await?;
observe_gateway_stage_ms(
"direct_send_headers",
started_at.elapsed().as_millis() as u64,
);
let status_code = response.status_code();
let headers = response.headers();
@@ -321,6 +405,7 @@ impl DirectSyncExecutionRuntime {
response: response.into_direct_upstream_response(),
started_at,
stream_first_byte_timeout: resolve_stream_first_byte_timeout(plan),
upstream_target_permit: None,
})
}
}
@@ -433,6 +518,7 @@ pub(crate) async fn execute_stream_plan_via_local_tunnel(
response: DirectUpstreamResponse::LocalTunnel(response),
started_at,
stream_first_byte_timeout: resolve_stream_first_byte_timeout(plan),
upstream_target_permit: None,
}))
}
@@ -1366,64 +1452,367 @@ fn build_client(
) -> Result<reqwest::Client, ExecutionRuntimeTransportError> {
validate_reqwest_transport_profile(transport_profile)?;
let resolved_proxy_url = resolve_proxy_url(proxy)?;
if resolved_proxy_url.is_none() && transport_profile.is_none() {
let cache_key = DirectReqwestClientCacheKey {
connect_timeout_ms: timeouts.and_then(|timeouts| timeouts.connect_ms),
follow_redirects: transport_controls.follow_redirects == Some(true),
http1_only: transport_controls.http1_only,
accept_invalid_certs: transport_controls.accept_invalid_certs,
};
if let Ok(cache) = DIRECT_REQWEST_CLIENT_CACHE.lock() {
if let Some(client) = cache.get(&cache_key) {
return Ok(client.clone());
let cache_key = direct_reqwest_client_cache_key(
timeouts,
resolved_proxy_url,
transport_profile,
transport_controls,
);
cached_direct_reqwest_client(cache_key)
}
pub(crate) fn prewarm_direct_reqwest_client_cache_for_plan(plan: &ExecutionPlan) {
match try_prewarm_direct_reqwest_client_cache_for_plan(plan) {
Ok(true) => {}
Ok(false) => {}
Err(err) => {
tracing::debug!(
error = ?err,
request_id = %plan.request_id,
candidate_id = ?plan.candidate_id,
provider_id = %plan.provider_id,
endpoint_id = %plan.endpoint_id,
key_id = %plan.key_id,
"gateway direct reqwest client prewarm skipped"
);
}
}
}
fn try_prewarm_direct_reqwest_client_cache_for_plan(
plan: &ExecutionPlan,
) -> Result<bool, ExecutionRuntimeTransportError> {
if transport_profile_uses_browser_wreq(plan.transport_profile.as_ref()) {
return Ok(false);
}
if resolve_tunnel_node_id(plan.proxy.as_ref()).is_some() {
return Ok(false);
}
let transport_controls = resolve_execution_transport_controls(&plan.headers);
validate_reqwest_transport_profile(plan.transport_profile.as_ref())?;
let resolved_proxy_url = resolve_proxy_url(plan.proxy.as_ref())?;
let cache_key = direct_reqwest_client_cache_key(
plan.timeouts.as_ref(),
resolved_proxy_url,
plan.transport_profile.as_ref(),
transport_controls,
);
prewarm_direct_reqwest_client_cache(cache_key)?;
Ok(true)
}
fn prewarm_direct_reqwest_client_cache(
cache_key: DirectReqwestClientCacheKey,
) -> Result<(), ExecutionRuntimeTransportError> {
let mut warm_after_unlock = None;
if let Ok(mut cache) = DIRECT_REQWEST_CLIENT_CACHE.lock() {
if let Some(entry) = cache.get_mut(&cache_key) {
if entry.should_warm() {
entry.warming = true;
warm_after_unlock = Some((cache_key.clone(), entry.len(), entry.target_len));
}
drop(cache);
if let Some((cache_key, existing_len, target_len)) = warm_after_unlock {
let spawned = spawn_direct_reqwest_client_cache_warm(
cache_key.clone(),
existing_len,
target_len,
);
if !spawned {
mark_direct_reqwest_client_cache_not_warming(&cache_key);
}
}
return Ok(());
}
let client = build_plain_direct_reqwest_client(cache_key)?;
if let Ok(mut cache) = DIRECT_REQWEST_CLIENT_CACHE.lock() {
let client = cache.entry(cache_key).or_insert_with(|| client.clone());
return Ok(client.clone());
let target_len = direct_reqwest_client_shard_count(&cache_key);
let initial_len = direct_reqwest_initial_client_shard_count(target_len);
let mut clients = Vec::with_capacity(initial_len);
for _ in 0..initial_len {
clients.push(build_direct_reqwest_client_from_cache_key(&cache_key)?);
DIRECT_REQWEST_CLIENT_CACHE_METRICS
.builds
.fetch_add(1, Ordering::Relaxed);
}
let entry =
DirectReqwestClientCacheEntry::new(clients, target_len, target_len > initial_len);
let warm_key = (target_len > initial_len).then(|| cache_key.clone());
cache.insert(cache_key, entry);
if let Some(warm_key) = warm_key {
warm_after_unlock = Some((warm_key, initial_len, target_len));
}
drop(cache);
if let Some((cache_key, existing_len, target_len)) = warm_after_unlock {
let spawned =
spawn_direct_reqwest_client_cache_warm(cache_key.clone(), existing_len, target_len);
if !spawned {
mark_direct_reqwest_client_cache_not_warming(&cache_key);
}
}
}
Ok(())
}
fn cached_direct_reqwest_client(
cache_key: DirectReqwestClientCacheKey,
) -> Result<reqwest::Client, ExecutionRuntimeTransportError> {
let mut warm_after_unlock = None;
if let Ok(mut cache) = DIRECT_REQWEST_CLIENT_CACHE.lock() {
if let Some(entry) = cache.get_mut(&cache_key) {
DIRECT_REQWEST_CLIENT_CACHE_METRICS
.hits
.fetch_add(1, Ordering::Relaxed);
let client = entry.select();
if entry.should_warm() {
entry.warming = true;
warm_after_unlock = Some((cache_key.clone(), entry.len(), entry.target_len));
}
drop(cache);
if let Some((cache_key, existing_len, target_len)) = warm_after_unlock {
let spawned = spawn_direct_reqwest_client_cache_warm(
cache_key.clone(),
existing_len,
target_len,
);
if !spawned {
mark_direct_reqwest_client_cache_not_warming(&cache_key);
}
}
return Ok(client);
}
DIRECT_REQWEST_CLIENT_CACHE_METRICS
.misses
.fetch_add(1, Ordering::Relaxed);
let target_len = direct_reqwest_client_shard_count(&cache_key);
let initial_len = direct_reqwest_initial_client_shard_count(target_len);
let mut clients = Vec::with_capacity(initial_len);
for _ in 0..initial_len {
clients.push(build_direct_reqwest_client_from_cache_key(&cache_key)?);
DIRECT_REQWEST_CLIENT_CACHE_METRICS
.builds
.fetch_add(1, Ordering::Relaxed);
}
let entry =
DirectReqwestClientCacheEntry::new(clients, target_len, target_len > initial_len);
let client = entry.select();
let warm_key = (target_len > initial_len).then(|| cache_key.clone());
cache.insert(cache_key, entry);
if let Some(warm_key) = warm_key {
warm_after_unlock = Some((warm_key, initial_len, target_len));
}
drop(cache);
if let Some((cache_key, existing_len, target_len)) = warm_after_unlock {
let spawned =
spawn_direct_reqwest_client_cache_warm(cache_key.clone(), existing_len, target_len);
if !spawned {
mark_direct_reqwest_client_cache_not_warming(&cache_key);
}
}
return Ok(client);
}
let mut builder = reqwest::Client::builder();
if transport_controls.follow_redirects != Some(true) {
builder = builder.redirect(Policy::none());
}
if transport_controls.http1_only || transport_profile_http1_only(transport_profile) {
builder = builder.http1_only();
}
let mut builder = apply_http_client_config(
builder,
&HttpClientConfig {
connect_timeout_ms: timeouts.and_then(|timeouts| timeouts.connect_ms),
..HttpClientConfig::default()
},
);
builder = apply_transport_profile(builder, transport_profile);
if transport_controls.accept_invalid_certs {
builder = builder.danger_accept_invalid_certs(true);
}
if let Some(proxy_url) = resolved_proxy_url {
let proxy = reqwest::Proxy::all(&proxy_url)
.map_err(ExecutionRuntimeTransportError::InvalidProxy)?;
builder = builder.proxy(proxy);
}
builder
.build()
.map_err(ExecutionRuntimeTransportError::ClientBuild)
DIRECT_REQWEST_CLIENT_CACHE_METRICS
.misses
.fetch_add(1, Ordering::Relaxed);
let client = build_direct_reqwest_client_from_cache_key(&cache_key)?;
DIRECT_REQWEST_CLIENT_CACHE_METRICS
.builds
.fetch_add(1, Ordering::Relaxed);
Ok(client)
}
fn build_plain_direct_reqwest_client(
fn spawn_direct_reqwest_client_cache_warm(
cache_key: DirectReqwestClientCacheKey,
existing_len: usize,
target_len: usize,
) -> bool {
if target_len <= existing_len {
return false;
}
let Ok(handle) = tokio::runtime::Handle::try_current() else {
return false;
};
handle.spawn_blocking(move || {
for _ in existing_len..target_len {
match build_direct_reqwest_client_from_cache_key(&cache_key) {
Ok(client) => {
DIRECT_REQWEST_CLIENT_CACHE_METRICS
.builds
.fetch_add(1, Ordering::Relaxed);
let Ok(mut cache) = DIRECT_REQWEST_CLIENT_CACHE.lock() else {
return;
};
let Some(entry) = cache.get_mut(&cache_key) else {
return;
};
if entry.clients.len() >= entry.target_len {
entry.warming = false;
return;
}
entry.clients.push(client);
if entry.clients.len() >= entry.target_len {
entry.warming = false;
return;
}
}
Err(err) => {
tracing::debug!(
error = ?err,
"gateway direct reqwest client cache warm failed"
);
mark_direct_reqwest_client_cache_not_warming(&cache_key);
break;
}
}
}
let Ok(mut cache) = DIRECT_REQWEST_CLIENT_CACHE.lock() else {
return;
};
let Some(entry) = cache.get_mut(&cache_key) else {
return;
};
entry.warming = false;
});
true
}
fn mark_direct_reqwest_client_cache_warming(cache_key: &DirectReqwestClientCacheKey) {
if let Ok(mut cache) = DIRECT_REQWEST_CLIENT_CACHE.lock() {
if let Some(entry) = cache.get_mut(cache_key) {
entry.warming = true;
}
}
}
fn mark_direct_reqwest_client_cache_not_warming(cache_key: &DirectReqwestClientCacheKey) {
if let Ok(mut cache) = DIRECT_REQWEST_CLIENT_CACHE.lock() {
if let Some(entry) = cache.get_mut(cache_key) {
entry.warming = false;
}
}
}
fn direct_reqwest_client_cache_key(
timeouts: Option<&aether_contracts::ExecutionTimeouts>,
proxy_url: Option<String>,
transport_profile: Option<&ResolvedTransportProfile>,
transport_controls: ExecutionTransportControls,
) -> DirectReqwestClientCacheKey {
DirectReqwestClientCacheKey {
connect_timeout_ms: timeouts.and_then(|timeouts| timeouts.connect_ms),
proxy_url,
follow_redirects: transport_controls.follow_redirects == Some(true),
http1_only: transport_controls.http1_only,
accept_invalid_certs: transport_controls.accept_invalid_certs,
transport_profile: transport_profile.map(direct_reqwest_transport_profile_cache_key),
}
}
fn direct_reqwest_transport_profile_cache_key(
profile: &ResolvedTransportProfile,
) -> DirectReqwestTransportProfileCacheKey {
DirectReqwestTransportProfileCacheKey {
profile_id: profile.profile_id.trim().to_string(),
backend: profile.backend.trim().to_ascii_lowercase(),
http_mode: profile.http_mode.trim().to_ascii_lowercase(),
pool_scope: profile.pool_scope.trim().to_ascii_lowercase(),
header_fingerprint: stable_json_cache_key(profile.header_fingerprint.as_ref()),
extra: stable_json_cache_key(profile.extra.as_ref()),
}
}
fn stable_json_cache_key(value: Option<&Value>) -> Option<String> {
value.and_then(|value| serde_json::to_string(value).ok())
}
fn build_direct_reqwest_client_cache_entry_from_cache_key(
cache_key: &DirectReqwestClientCacheKey,
) -> Result<DirectReqwestClientCacheEntry, ExecutionRuntimeTransportError> {
let shard_count = direct_reqwest_client_shard_count(cache_key);
let mut clients = Vec::with_capacity(shard_count);
for _ in 0..shard_count {
clients.push(build_direct_reqwest_client_from_cache_key(cache_key)?);
}
Ok(DirectReqwestClientCacheEntry::new(
clients,
shard_count,
false,
))
}
fn direct_reqwest_client_shard_count(cache_key: &DirectReqwestClientCacheKey) -> usize {
if !direct_reqwest_client_cache_key_uses_http2(cache_key) {
return 1;
}
direct_reqwest_h2_client_shards_from_config(
env_positive_usize(DIRECT_REQWEST_H2_CLIENT_SHARDS_ENV),
env_positive_usize(UPSTREAM_TARGET_GATE_LIMIT_ENV)
.unwrap_or(DEFAULT_UPSTREAM_TARGET_GATE_LIMIT),
env_positive_usize(DIRECT_REQWEST_H2_TARGET_STREAMS_PER_CLIENT_ENV)
.unwrap_or(DEFAULT_H2_TARGET_STREAMS_PER_CLIENT),
)
}
fn direct_reqwest_client_cache_key_uses_http2(cache_key: &DirectReqwestClientCacheKey) -> bool {
if cache_key.http1_only {
return false;
}
let Some(profile) = cache_key.transport_profile.as_ref() else {
return false;
};
profile.http_mode != TRANSPORT_HTTP_MODE_HTTP1_ONLY
}
fn direct_reqwest_h2_client_shards_from_config(
explicit_shards: Option<usize>,
target_gate_limit: usize,
target_streams_per_client: usize,
) -> usize {
if let Some(shards) = explicit_shards {
return shards.clamp(1, MAX_DIRECT_REQWEST_H2_CLIENT_SHARDS);
}
let streams_per_client = target_streams_per_client.max(1);
target_gate_limit
.max(1)
.div_ceil(streams_per_client)
.clamp(1, MAX_DIRECT_REQWEST_H2_CLIENT_SHARDS)
}
fn direct_reqwest_initial_client_shard_count(target_len: usize) -> usize {
env_positive_usize(DIRECT_REQWEST_SYNC_WARM_CLIENTS_ENV)
.unwrap_or(DEFAULT_DIRECT_REQWEST_SYNC_WARM_CLIENTS)
.clamp(1, target_len.max(1))
}
fn env_positive_usize(name: &str) -> Option<usize> {
std::env::var(name)
.ok()
.and_then(|value| value.trim().parse::<usize>().ok())
.filter(|value| *value > 0)
}
fn build_direct_reqwest_client_from_cache_key(
cache_key: &DirectReqwestClientCacheKey,
) -> Result<reqwest::Client, ExecutionRuntimeTransportError> {
let mut builder = reqwest::Client::builder();
if !cache_key.follow_redirects {
builder = builder.redirect(Policy::none());
}
if cache_key.http1_only {
if cache_key.http1_only
|| cache_key
.transport_profile
.as_ref()
.is_some_and(|profile| profile.http_mode == TRANSPORT_HTTP_MODE_HTTP1_ONLY)
{
builder = builder.http1_only();
} else if cache_key
.transport_profile
.as_ref()
.is_some_and(|profile| profile.http_mode == TRANSPORT_HTTP_MODE_H2C_PRIOR_KNOWLEDGE)
{
builder = builder.http2_prior_knowledge();
}
let mut builder = apply_http_client_config(
builder,
@@ -1433,9 +1822,15 @@ fn build_plain_direct_reqwest_client(
..HttpClientConfig::default()
},
);
builder = apply_transport_profile_cache_key(builder, cache_key.transport_profile.as_ref());
if cache_key.accept_invalid_certs {
builder = builder.danger_accept_invalid_certs(true);
}
if let Some(proxy_url) = cache_key.proxy_url.as_deref() {
let proxy =
reqwest::Proxy::all(proxy_url).map_err(ExecutionRuntimeTransportError::InvalidProxy)?;
builder = builder.proxy(proxy);
}
builder
.build()
.map_err(ExecutionRuntimeTransportError::ClientBuild)
@@ -1450,6 +1845,55 @@ fn direct_reqwest_pool_max_idle_per_host() -> usize {
.unwrap_or(DEFAULT_MAX_IDLE_PER_HOST)
}
pub(crate) fn direct_reqwest_client_cache_metric_samples() -> Vec<MetricSample> {
let (entries, clients) = DIRECT_REQWEST_CLIENT_CACHE
.lock()
.map(|cache| {
let entries = cache.len() as u64;
let clients = cache.values().map(|entry| entry.len() as u64).sum();
(entries, clients)
})
.unwrap_or((0, 0));
vec![
MetricSample::new(
"direct_reqwest_client_cache_entries",
"Number of cached direct reqwest clients.",
MetricKind::Gauge,
entries,
),
MetricSample::new(
"direct_reqwest_client_cache_clients",
"Number of direct reqwest clients across all cache entries.",
MetricKind::Gauge,
clients,
),
MetricSample::new(
"direct_reqwest_client_cache_hits_total",
"Number of direct reqwest client cache hits.",
MetricKind::Counter,
DIRECT_REQWEST_CLIENT_CACHE_METRICS
.hits
.load(Ordering::Relaxed),
),
MetricSample::new(
"direct_reqwest_client_cache_misses_total",
"Number of direct reqwest client cache misses.",
MetricKind::Counter,
DIRECT_REQWEST_CLIENT_CACHE_METRICS
.misses
.load(Ordering::Relaxed),
),
MetricSample::new(
"direct_reqwest_client_cache_builds_total",
"Number of direct reqwest clients built after cache misses.",
MetricKind::Counter,
DIRECT_REQWEST_CLIENT_CACHE_METRICS
.builds
.load(Ordering::Relaxed),
),
]
}
pub(crate) fn build_browser_wreq_client(
timeouts: Option<&aether_contracts::ExecutionTimeouts>,
proxy: Option<&ProxySnapshot>,
@@ -1622,6 +2066,19 @@ fn transport_profile_http1_only(transport_profile: Option<&ResolvedTransportProf
.unwrap_or(false)
}
fn transport_profile_h2c_prior_knowledge(
transport_profile: Option<&ResolvedTransportProfile>,
) -> bool {
transport_profile
.map(|profile| {
profile
.http_mode
.trim()
.eq_ignore_ascii_case(TRANSPORT_HTTP_MODE_H2C_PRIOR_KNOWLEDGE)
})
.unwrap_or(false)
}
fn apply_transport_profile(
builder: reqwest::ClientBuilder,
transport_profile: Option<&ResolvedTransportProfile>,
@@ -1630,7 +2087,24 @@ fn apply_transport_profile(
return builder;
};
let profile_id = profile.profile_id.trim();
if profile_id.is_empty() {
if profile_id.is_empty() || transport_profile_h2c_prior_knowledge(Some(profile)) {
return builder;
}
let _ = rustls::crypto::ring::default_provider().install_default();
builder.use_preconfigured_tls(build_best_effort_transport_tls_config())
}
fn apply_transport_profile_cache_key(
builder: reqwest::ClientBuilder,
transport_profile: Option<&DirectReqwestTransportProfileCacheKey>,
) -> reqwest::ClientBuilder {
let Some(profile) = transport_profile else {
return builder;
};
if profile.profile_id.is_empty() || profile.http_mode == TRANSPORT_HTTP_MODE_H2C_PRIOR_KNOWLEDGE
{
return builder;
}
@@ -1915,6 +2389,7 @@ mod tests {
ExecutionPlan, ExecutionTimeouts, ProxySnapshot, RequestBody, ResolvedTransportProfile,
EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER, EXECUTION_REQUEST_HTTP1_ONLY_HEADER,
TRANSPORT_BACKEND_BROWSER_WREQ, TRANSPORT_BACKEND_REQWEST_RUSTLS,
TRANSPORT_HTTP_MODE_H2C_PRIOR_KNOWLEDGE, TRANSPORT_HTTP_MODE_HTTP1_ONLY,
};
use aether_data::repository::proxy_nodes::{
InMemoryProxyNodeRepository, ProxyNodeReadRepository, StoredProxyNode,
@@ -2024,6 +2499,186 @@ mod tests {
}
}
#[test]
fn direct_reqwest_client_cache_key_includes_transport_profile() {
let timeouts = ExecutionTimeouts {
connect_ms: Some(5_000),
..ExecutionTimeouts::default()
};
let h2c_profile = ResolvedTransportProfile {
profile_id: "mock-h2c".into(),
backend: TRANSPORT_BACKEND_REQWEST_RUSTLS.into(),
http_mode: TRANSPORT_HTTP_MODE_H2C_PRIOR_KNOWLEDGE.into(),
pool_scope: "key".into(),
header_fingerprint: None,
extra: Some(json!({"pool": "a"})),
};
let same_h2c_profile = ResolvedTransportProfile {
extra: Some(json!({"pool": "a"})),
..h2c_profile.clone()
};
let http1_profile = ResolvedTransportProfile {
http_mode: TRANSPORT_HTTP_MODE_HTTP1_ONLY.into(),
..h2c_profile.clone()
};
let left = super::direct_reqwest_client_cache_key(
Some(&timeouts),
None,
Some(&h2c_profile),
ExecutionTransportControls::default(),
);
let right = super::direct_reqwest_client_cache_key(
Some(&timeouts),
None,
Some(&same_h2c_profile),
ExecutionTransportControls::default(),
);
let different_mode = super::direct_reqwest_client_cache_key(
Some(&timeouts),
None,
Some(&http1_profile),
ExecutionTransportControls::default(),
);
let different_proxy = super::direct_reqwest_client_cache_key(
Some(&timeouts),
Some("http://127.0.0.1:8080".into()),
Some(&h2c_profile),
ExecutionTransportControls::default(),
);
assert_eq!(left, right);
assert_ne!(left, different_mode);
assert_ne!(left, different_proxy);
assert!(super::direct_reqwest_client_cache_key_uses_http2(&left));
assert!(!super::direct_reqwest_client_cache_key_uses_http2(
&different_mode
));
}
#[test]
fn direct_reqwest_h2_client_shards_scale_from_target_gate() {
assert_eq!(
super::direct_reqwest_h2_client_shards_from_config(None, 12_000, 200),
60
);
assert_eq!(
super::direct_reqwest_h2_client_shards_from_config(None, 2_000, 200),
10
);
assert_eq!(
super::direct_reqwest_h2_client_shards_from_config(Some(4), 12_000, 200),
4
);
assert_eq!(
super::direct_reqwest_h2_client_shards_from_config(None, 100_000, 100),
super::MAX_DIRECT_REQWEST_H2_CLIENT_SHARDS
);
}
#[test]
fn direct_reqwest_initial_client_shards_are_bounded_by_target() {
assert_eq!(super::direct_reqwest_initial_client_shard_count(1), 1);
assert_eq!(super::direct_reqwest_initial_client_shard_count(2), 2);
assert_eq!(
super::direct_reqwest_initial_client_shard_count(21),
super::DEFAULT_DIRECT_REQWEST_SYNC_WARM_CLIENTS
);
}
#[test]
fn direct_reqwest_prewarm_populates_cache_for_plan() {
let profile = ResolvedTransportProfile {
profile_id: "mock-h2c-prewarm".into(),
backend: TRANSPORT_BACKEND_REQWEST_RUSTLS.into(),
http_mode: TRANSPORT_HTTP_MODE_H2C_PRIOR_KNOWLEDGE.into(),
pool_scope: "key".into(),
header_fingerprint: None,
extra: None,
};
let plan = ExecutionPlan {
request_id: "req-prewarm".into(),
candidate_id: Some("candidate-prewarm".into()),
provider_name: Some("mock".into()),
provider_id: "provider-1".into(),
endpoint_id: "endpoint-1".into(),
key_id: "key-1".into(),
method: "POST".into(),
url: "http://127.0.0.1:18184/v1/chat/completions".into(),
headers: BTreeMap::new(),
content_type: Some("application/json".into()),
content_encoding: None,
body: RequestBody::from_json(json!({"stream": true})),
stream: true,
client_api_format: "openai:chat".into(),
provider_api_format: "openai:chat".into(),
model_name: Some("mock-model".into()),
proxy: None,
transport_profile: Some(profile.clone()),
timeouts: Some(ExecutionTimeouts {
connect_ms: Some(5_000),
..ExecutionTimeouts::default()
}),
};
assert!(
super::try_prewarm_direct_reqwest_client_cache_for_plan(&plan)
.expect("prewarm should succeed")
);
let cache_key = super::direct_reqwest_client_cache_key(
plan.timeouts.as_ref(),
None,
Some(&profile),
super::ExecutionTransportControls::default(),
);
let target_len = super::direct_reqwest_client_shard_count(&cache_key);
let expected_initial_len = super::direct_reqwest_initial_client_shard_count(target_len);
let cache = super::DIRECT_REQWEST_CLIENT_CACHE
.lock()
.expect("cache lock");
let entry = cache.get(&cache_key).expect("cache entry");
assert_eq!(entry.len(), expected_initial_len);
assert_eq!(entry.target_len, target_len);
}
#[test]
fn direct_reqwest_prewarm_skips_browser_transport() {
let plan = ExecutionPlan {
request_id: "req-browser".into(),
candidate_id: None,
provider_name: Some("browser".into()),
provider_id: "provider-1".into(),
endpoint_id: "endpoint-1".into(),
key_id: "key-1".into(),
method: "POST".into(),
url: "https://example.com/v1/chat/completions".into(),
headers: BTreeMap::new(),
content_type: Some("application/json".into()),
content_encoding: None,
body: RequestBody::from_json(json!({"stream": true})),
stream: true,
client_api_format: "openai:chat".into(),
provider_api_format: "openai:chat".into(),
model_name: Some("mock-model".into()),
proxy: None,
transport_profile: Some(ResolvedTransportProfile {
profile_id: "chrome_136".into(),
backend: TRANSPORT_BACKEND_BROWSER_WREQ.into(),
http_mode: "auto".into(),
pool_scope: "key".into(),
header_fingerprint: None,
extra: None,
}),
timeouts: None,
};
assert!(
!super::try_prewarm_direct_reqwest_client_cache_for_plan(&plan)
.expect("browser transport should skip prewarm")
);
}
#[test]
fn direct_sync_execution_runtime_strips_accept_invalid_certs_control_header() {
let headers = BTreeMap::from([
@@ -3667,6 +4322,73 @@ mod tests {
);
}
#[tokio::test]
async fn direct_sync_execution_runtime_supports_h2c_prior_knowledge_profile() {
let listener = crate::test_support::bind_loopback_listener()
.await
.expect("listener should bind");
let addr = listener.local_addr().expect("local addr should resolve");
let app = Router::new().route(
"/chat",
post(|| async {
(
axum::http::StatusCode::OK,
Json(json!({"transport_profile": "h2c"})),
)
}),
);
let server = tokio::spawn(async move {
axum::serve(listener, app)
.await
.expect("test server should run");
});
let execution_runtime = DirectSyncExecutionRuntime::new();
let result = execution_runtime
.execute_sync(&ExecutionPlan {
request_id: "req-h2c-1".into(),
candidate_id: Some("cand-h2c-1".into()),
provider_name: Some("mock".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(json!({"model": "mock-model"})),
stream: false,
client_api_format: "openai:chat".into(),
provider_api_format: "openai:chat".into(),
model_name: Some("mock-model".into()),
proxy: None,
transport_profile: Some(ResolvedTransportProfile {
profile_id: "mock-h2c".into(),
backend: TRANSPORT_BACKEND_REQWEST_RUSTLS.into(),
http_mode: TRANSPORT_HTTP_MODE_H2C_PRIOR_KNOWLEDGE.into(),
pool_scope: "key".into(),
header_fingerprint: None,
extra: None,
}),
timeouts: Some(ExecutionTimeouts {
connect_ms: Some(5_000),
total_ms: Some(LOCAL_HTTP_SUCCESS_TIMEOUT_MS),
..ExecutionTimeouts::default()
}),
})
.await
.expect("h2c prior-knowledge execution should succeed");
server.abort();
assert_eq!(result.status_code, 200);
assert_eq!(
result.body.and_then(|body| body.json_body),
Some(json!({"transport_profile": "h2c"}))
);
}
#[test]
fn direct_sync_execution_runtime_rejects_unsupported_transport_backend() {
let profile = ResolvedTransportProfile {
@@ -2,12 +2,14 @@ use aether_ai_serving::{
run_ai_attempt_loop, AiAttemptLoopOutcome, AiAttemptLoopPort, AiExecutionAttempt,
};
use aether_data_contracts::repository::candidates::RequestCandidateStatus;
use aether_runtime::ConcurrencyPermit;
use aether_scheduler_core::{
parse_request_candidate_report_context, SchedulerRequestCandidateStatusUpdate,
};
use async_trait::async_trait;
use axum::body::Body;
use axum::http::Response;
use futures_util::StreamExt;
use tokio::time::{timeout, Duration};
use tracing::{debug, warn, Instrument};
@@ -23,6 +25,7 @@ use crate::privacy::RedactionExecutionCandidateId;
use crate::request_candidate_runtime::{
record_local_request_candidate_status, RequestCandidateRuntimeWriter,
};
use crate::stage_metrics::observe_gateway_stage_ms;
use crate::{AppState, GatewayError};
const DEFAULT_STREAM_FIRST_BYTE_WATCHDOG_TIMEOUT_MS: u64 = 30_000;
@@ -157,6 +160,8 @@ where
type Error = GatewayError;
async fn execute_attempt(&self, attempt: &T) -> Result<Option<Self::Response>, Self::Error> {
prewarm_direct_reqwest_candidate_client(attempt.execution_plan());
let _permit = acquire_upstream_execution_gate(self.state, self.trace_id).await?;
let mut response = execute_execution_runtime_sync(
self.state,
self.parts.uri.path(),
@@ -323,13 +328,33 @@ where
{
let mut last_attempted = None;
while let Some(attempt) =
next_execution_attempt_with_timeout(source, trace_id, plan_kind, planning_timeout).await?
{
loop {
let next_started_at = std::time::Instant::now();
let next_attempt =
next_execution_attempt_with_timeout(source, trace_id, plan_kind, planning_timeout)
.await?;
observe_gateway_stage_ms(
"stream_candidate_next",
next_started_at.elapsed().as_millis() as u64,
);
let Some(attempt) = next_attempt else {
break;
};
last_attempted = Some((attempt.execution_plan().clone(), attempt.report_context()));
if let Some(response) = port.execute_attempt(&attempt).await? {
let execute_started_at = std::time::Instant::now();
let response = port.execute_attempt(&attempt).await?;
observe_gateway_stage_ms(
"stream_candidate_execute",
execute_started_at.elapsed().as_millis() as u64,
);
if let Some(response) = response {
let remaining = source.drain_execution_attempts().await?;
let unused_started_at = std::time::Instant::now();
port.mark_unused_attempts(remaining).await?;
observe_gateway_stage_ms(
"stream_candidate_unused",
unused_started_at.elapsed().as_millis() as u64,
);
return Ok(LocalExecutionRequestOutcome::responded(response));
}
}
@@ -412,6 +437,7 @@ where
candidate_index = candidate_index.as_str(),
"candidate loop attempting stream execution candidate"
);
prewarm_direct_reqwest_candidate_client(&plan);
let watchdog_plan = plan.clone();
let watchdog_report_context = report_context.clone();
let execution_state = self.state.clone();
@@ -475,6 +501,15 @@ where
}
}
fn prewarm_direct_reqwest_candidate_client(plan: &aether_contracts::ExecutionPlan) {
let started_at = std::time::Instant::now();
crate::execution_runtime::transport::prewarm_direct_reqwest_client_cache_for_plan(plan);
observe_gateway_stage_ms(
"direct_reqwest_client_prewarm",
started_at.elapsed().as_millis() as u64,
);
}
pub(crate) async fn mark_unused_local_candidates<T>(state: &AppState, remaining: Vec<T>)
where
T: AiExecutionAttempt,
@@ -543,7 +578,7 @@ fn stream_candidate_watchdog_timeout_message() -> &'static str {
}
async fn execute_stream_candidate_with_watchdog<Fut>(
state: &(impl RequestCandidateRuntimeWriter + ?Sized),
state: &(impl RequestCandidateRuntimeWriter + UpstreamExecutionGateProvider + ?Sized),
trace_id: &str,
plan_kind: &str,
plan: &aether_contracts::ExecutionPlan,
@@ -556,9 +591,12 @@ where
{
let timeout_duration = resolve_stream_candidate_watchdog_timeout(plan, report_context);
let candidate_started_unix_ms = current_unix_ms();
let permit = acquire_upstream_execution_gate(state, trace_id).await?;
let mut join_handle = tokio::spawn(execute());
match timeout(timeout_duration, &mut join_handle).await {
Ok(Ok(result)) => result,
Ok(Ok(result)) => {
result.map(|response| maybe_hold_upstream_execution_permit(response, permit))
}
Ok(Err(join_error)) => Err(GatewayError::Internal(format!(
"local stream candidate task join failed: {join_error}"
))),
@@ -608,6 +646,67 @@ where
}
}
fn maybe_hold_upstream_execution_permit(
response: Option<Response<Body>>,
permit: Option<ConcurrencyPermit>,
) -> Option<Response<Body>> {
match (response, permit) {
(Some(response), Some(permit)) => {
Some(hold_response_upstream_execution_permit(response, permit))
}
(response, _) => response,
}
}
fn hold_response_upstream_execution_permit(
response: Response<Body>,
permit: ConcurrencyPermit,
) -> Response<Body> {
let (parts, body) = response.into_parts();
let stream = async_stream::stream! {
let _permit = permit;
let mut body_stream = body.into_data_stream();
while let Some(item) = body_stream.next().await {
yield item;
}
};
Response::from_parts(parts, Body::from_stream(stream))
}
trait UpstreamExecutionGateProvider {
fn upstream_execution_gate(&self) -> Option<&aether_runtime::ConcurrencyGate>;
fn upstream_execution_gate_queue_budget(&self) -> Duration;
}
impl UpstreamExecutionGateProvider for AppState {
fn upstream_execution_gate(&self) -> Option<&aether_runtime::ConcurrencyGate> {
self.upstream_execution_gate.as_deref()
}
fn upstream_execution_gate_queue_budget(&self) -> Duration {
self.frontdoor_runtime_guards.internal_gate_queue_budget
}
}
async fn acquire_upstream_execution_gate(
state: &(impl UpstreamExecutionGateProvider + ?Sized),
trace_id: &str,
) -> Result<Option<ConcurrencyPermit>, GatewayError> {
let Some(gate) = state.upstream_execution_gate() else {
return Ok(None);
};
let budget = state.upstream_execution_gate_queue_budget();
match timeout(budget, gate.acquire()).await {
Ok(Ok(permit)) => Ok(Some(permit)),
Ok(Err(err)) => Err(GatewayError::Internal(err.to_string())),
Err(_) => Err(GatewayError::AdmissionTimeout {
trace_id: trace_id.to_string(),
gate: "gateway_upstream_execution",
queue_budget_ms: budget.as_millis() as u64,
}),
}
}
pub(crate) async fn mark_unused_local_candidate_items<T, FPlan, FContext>(
state: &AppState,
remaining: Vec<T>,
@@ -677,6 +776,16 @@ mod tests {
}
}
impl UpstreamExecutionGateProvider for TestRequestCandidateWriter {
fn upstream_execution_gate(&self) -> Option<&aether_runtime::ConcurrencyGate> {
None
}
fn upstream_execution_gate_queue_budget(&self) -> Duration {
Duration::from_millis(250)
}
}
struct PendingAttemptSource;
#[async_trait]
+2 -1
View File
@@ -16,7 +16,8 @@ pub(crate) use candidate_loop::{
};
pub(crate) use orchestration::*;
pub(crate) use outcome::{
beautify_local_execution_client_error_message, build_local_execution_exhaustion,
beautify_local_execution_client_error_message, build_fast_local_execution_exhaustion,
build_fast_local_execution_runtime_miss_context, build_local_execution_exhaustion,
build_local_execution_runtime_miss_context, record_failed_usage_for_exhausted_request,
record_failed_usage_for_runtime_miss_request, LocalExecutionExhaustion,
LocalExecutionRequestOutcome, LocalExecutionRuntimeMissContext,
+60 -33
View File
@@ -150,6 +150,7 @@ pub(crate) async fn build_local_execution_exhaustion(
plan: &ExecutionPlan,
report_context: Option<&Value>,
) -> LocalExecutionExhaustion {
let mut exhaustion = build_fast_local_execution_exhaustion(plan, report_context);
let mut data = build_usage_event_data_seed(plan, report_context);
let last_failed_candidate = match state
.read_request_candidates_by_request_id(plan.request_id.as_str())
@@ -180,28 +181,46 @@ pub(crate) async fn build_local_execution_exhaustion(
.or_else(|| candidate.key_id.clone());
}
exhaustion.data = data;
exhaustion.candidate_id = last_failed_candidate
.as_ref()
.map(|candidate| candidate.id.clone());
exhaustion.candidate_index = last_failed_candidate
.as_ref()
.map(|candidate| candidate.candidate_index);
exhaustion.upstream_status_code = last_failed_candidate
.as_ref()
.and_then(|candidate| candidate.status_code);
exhaustion.upstream_error_type = last_failed_candidate
.as_ref()
.and_then(|candidate| candidate.error_type.clone())
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty());
exhaustion.upstream_error_message = last_failed_candidate
.as_ref()
.and_then(|candidate| candidate.error_message.clone())
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty());
exhaustion
}
pub(crate) fn build_fast_local_execution_exhaustion(
plan: &ExecutionPlan,
report_context: Option<&Value>,
) -> LocalExecutionExhaustion {
let data = build_usage_event_data_seed(plan, report_context);
LocalExecutionExhaustion {
request_id: plan.request_id.clone(),
candidate_id: plan.candidate_id.clone(),
candidate_index: report_context
.and_then(Value::as_object)
.and_then(|value| value.get("candidate_index"))
.and_then(Value::as_u64)
.and_then(|value| u32::try_from(value).ok()),
data,
candidate_id: last_failed_candidate
.as_ref()
.map(|candidate| candidate.id.clone()),
candidate_index: last_failed_candidate
.as_ref()
.map(|candidate| candidate.candidate_index),
upstream_status_code: last_failed_candidate
.as_ref()
.and_then(|candidate| candidate.status_code),
upstream_error_type: last_failed_candidate
.as_ref()
.and_then(|candidate| candidate.error_type.clone())
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty()),
upstream_error_message: last_failed_candidate
.as_ref()
.and_then(|candidate| candidate.error_message.clone())
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty()),
upstream_status_code: None,
upstream_error_type: None,
upstream_error_message: None,
}
}
@@ -224,6 +243,20 @@ pub(crate) async fn build_local_execution_runtime_miss_context(
}
}
pub(crate) fn build_fast_local_execution_runtime_miss_context(
decision: Option<&GatewayControlDecision>,
) -> LocalExecutionRuntimeMissContext {
let auth_context = decision.and_then(|value| value.auth_context.as_ref());
LocalExecutionRuntimeMissContext {
auth_user_id: auth_context.map(|value| value.user_id.clone()),
auth_api_key_id: auth_context.map(|value| value.api_key_id.clone()),
auth_username: auth_context.and_then(|value| value.username.clone()),
auth_api_key_name: auth_context.and_then(|value| value.api_key_name.clone()),
candidate_contexts: Vec::new(),
}
}
pub(crate) async fn record_failed_usage_for_exhausted_request(
state: &AppState,
exhaustion: LocalExecutionExhaustion,
@@ -309,13 +342,10 @@ pub(crate) async fn record_failed_usage_for_exhausted_request(
);
data.request_metadata = Some(Value::Object(request_metadata));
state
.usage_runtime
.record_terminal_event_direct(
state.data.as_ref(),
UsageEvent::new(UsageEventType::Failed, request_id, data),
)
.await;
state.usage_runtime.submit_terminal_event(
state.data.as_ref(),
UsageEvent::new(UsageEventType::Failed, request_id, data),
);
}
pub(crate) async fn record_failed_usage_for_runtime_miss_request(
@@ -448,13 +478,10 @@ pub(crate) async fn record_failed_usage_for_runtime_miss_request(
data.request_metadata =
(!request_metadata.is_empty()).then_some(Value::Object(request_metadata));
state
.usage_runtime
.record_terminal_event_direct(
state.data.as_ref(),
UsageEvent::new(UsageEventType::Failed, request_id, data),
)
.await;
state.usage_runtime.submit_terminal_event(
state.data.as_ref(),
UsageEvent::new(UsageEventType::Failed, request_id, data),
);
}
pub(crate) fn beautify_local_execution_client_error_message(message: &str) -> String {
@@ -13,6 +13,7 @@ use crate::ai_serving::api::{
};
use crate::api::response::build_client_response_from_parts;
use crate::control::GatewayControlDecision;
use crate::stage_metrics::observe_gateway_stage_ms;
use crate::{AppState, GatewayError, GatewayFallbackReason};
use super::{
@@ -94,6 +95,7 @@ impl AiStreamExecutionPathPort for GatewayStreamExecutionPathPort<'_> {
&self,
step: AiStreamExecutionStep,
) -> Result<AiServingExecutionOutcome<Self::Response, Self::Exhaustion>, Self::Error> {
let step_started_at = std::time::Instant::now();
let outcome = match step {
AiStreamExecutionStep::LocalVideoContent => {
maybe_execute_local_video_task_content_stream(
@@ -188,6 +190,10 @@ impl AiStreamExecutionPathPort for GatewayStreamExecutionPathPort<'_> {
}
}
};
observe_gateway_stage_ms(
"stream_path_step",
step_started_at.elapsed().as_millis() as u64,
);
Ok(to_ai_serving_outcome(outcome))
}
+88 -9
View File
@@ -60,6 +60,7 @@ use crate::scheduler::candidate::{
LEGACY_API_KEY_CONCURRENCY_LIMIT_SKIP_REASON,
};
use crate::scheduler::config::{read_scheduler_ordering_config, SchedulerSchedulingMode};
use crate::stage_metrics::observe_gateway_stage_ms;
use crate::{
AppState, FrontdoorUserRpmOutcome, GatewayError, GatewayFallbackMetricKind,
GatewayFallbackReason, LocalExecutionRuntimeMissDiagnostic,
@@ -1023,8 +1024,30 @@ pub(crate) async fn proxy_request(
State(state): State<AppState>,
ConnectInfo(remote_addr): ConnectInfo<std::net::SocketAddr>,
request: Request,
) -> Result<Response<Body>, GatewayError> {
crate::request_diagnostics::scope_request_diagnostics(proxy_request_inner(
state,
remote_addr,
request,
))
.await
}
async fn proxy_request_inner(
state: AppState,
remote_addr: std::net::SocketAddr,
request: Request,
) -> Result<Response<Body>, GatewayError> {
let started_at = Instant::now();
if let Some(accepted_at) = request
.extensions()
.get::<crate::middleware::GatewayRequestAcceptedAt>()
{
observe_gateway_stage_ms(
"frontdoor_handler_queue",
started_at.duration_since(accepted_at.0).as_millis() as u64,
);
}
let mut request_permit = match state.try_acquire_request_permit().await {
Ok(permit) => permit,
Err(RequestAdmissionError::Local(aether_runtime::ConcurrencyError::Saturated {
@@ -1085,6 +1108,7 @@ pub(crate) async fn proxy_request(
)) => return Err(GatewayError::Internal(message)),
};
let request_admission_ms = started_at.elapsed().as_millis() as u64;
observe_gateway_stage_ms("frontdoor_admission", request_admission_ms);
let (mut parts, body) = request.into_parts();
let redaction_slot = crate::privacy::RedactionSessionSlot::default();
parts.extensions.insert(redaction_slot.clone());
@@ -1176,6 +1200,7 @@ pub(crate) async fn proxy_request(
}
}
let request_context_ms = request_context_started_at.elapsed().as_millis() as u64;
observe_gateway_stage_ms("frontdoor_context", request_context_ms);
if request_context
.control_decision
.as_ref()
@@ -1202,6 +1227,7 @@ pub(crate) async fn proxy_request(
let mut request_body = Some(body);
let local_proxy_body = if local_proxy_route_requires_buffered_body(&request_context) {
let body_buffer_policy = RequestBodyBufferPolicy::from_state(&state);
let stage_started_at = Instant::now();
let body = buffer_and_normalize_request_body(
&mut request_body,
&mut parts.headers,
@@ -1213,6 +1239,10 @@ pub(crate) async fn proxy_request(
body_buffer_policy,
)
.await;
observe_gateway_stage_ms(
"frontdoor_body_buffer",
stage_started_at.elapsed().as_millis() as u64,
);
match body {
Ok(body) => Some(body),
Err(err) => {
@@ -1388,6 +1418,7 @@ pub(crate) async fn proxy_request(
let buffered_body = if should_buffer_body {
let body_buffer_policy = RequestBodyBufferPolicy::from_state(&state);
let stage_started_at = Instant::now();
let body = buffer_and_normalize_request_body(
&mut request_body,
&mut parts.headers,
@@ -1399,6 +1430,10 @@ pub(crate) async fn proxy_request(
body_buffer_policy,
)
.await;
observe_gateway_stage_ms(
"frontdoor_body_buffer",
stage_started_at.elapsed().as_millis() as u64,
);
match body {
Ok(body) => Some(body),
Err(err) => {
@@ -1417,15 +1452,20 @@ pub(crate) async fn proxy_request(
None
};
if let Some(response) = maybe_forward_public_request_to_tunnel_owner(
let owner_forward_started_at = Instant::now();
let owner_forward_response = maybe_forward_public_request_to_tunnel_owner(
&state,
&remote_addr,
&request_context,
&parts,
buffered_body.as_ref(),
)
.await?
{
.await?;
observe_gateway_stage_ms(
"frontdoor_owner_forward",
owner_forward_started_at.elapsed().as_millis() as u64,
);
if let Some(response) = owner_forward_response {
return Ok(finalize_gateway_response_with_context(
&state,
response,
@@ -1452,15 +1492,20 @@ pub(crate) async fn proxy_request(
}
if let Some(buffered_body) = buffered_body.as_ref() {
if let Some(rejection) = request_model_local_rejection(
let auth_model_started_at = Instant::now();
let model_rejection = request_model_local_rejection(
&state,
control_decision,
&parts.uri,
&parts.headers,
buffered_body,
)
.await?
{
.await?;
observe_gateway_stage_ms(
"frontdoor_auth_model",
auth_model_started_at.elapsed().as_millis() as u64,
);
if let Some(rejection) = model_rejection {
let response =
build_local_auth_rejection_response(&trace_id, control_decision, &rejection)?;
return Ok(finalize_gateway_response_with_context(
@@ -1475,10 +1520,12 @@ pub(crate) async fn proxy_request(
}
}
let rpm_started_at = Instant::now();
let rate_limit_outcome = state
.frontdoor_user_rpm()
.check_and_consume(&state, control_decision)
.await?;
observe_gateway_stage_ms("frontdoor_rpm", rpm_started_at.elapsed().as_millis() as u64);
if let FrontdoorUserRpmOutcome::Rejected(rejection) = &rate_limit_outcome {
let auth_context = control_decision.and_then(|decision| decision.auth_context.as_ref());
let user_id = auth_context
@@ -1514,13 +1561,18 @@ pub(crate) async fn proxy_request(
));
}
if let Some(response) = super::public::maybe_build_local_ai_public_response(
let local_ai_public_started_at = Instant::now();
let local_ai_public_response = super::public::maybe_build_local_ai_public_response(
&state,
&request_context,
buffered_body.as_ref(),
)
.await
{
.await;
observe_gateway_stage_ms(
"frontdoor_local_ai_public",
local_ai_public_started_at.elapsed().as_millis() as u64,
);
if let Some(response) = local_ai_public_response {
return Ok(finalize_gateway_response_with_context(
&state,
response,
@@ -1557,6 +1609,7 @@ pub(crate) async fn proxy_request(
let stream_request = request_wants_stream(&request_context, &parts.headers, buffered_body);
let mut local_execution_exhaustion = None;
if stream_request {
let execute_stream_started_at = Instant::now();
let stream_outcome = match maybe_execute_stream_request(
&state,
&parts,
@@ -1568,6 +1621,10 @@ pub(crate) async fn proxy_request(
{
Ok(outcome) => outcome,
Err(err) => {
observe_gateway_stage_ms(
"frontdoor_execute_stream",
execute_stream_started_at.elapsed().as_millis() as u64,
);
if let Some((phase, timeout_ms)) = local_execution_planning_timeout_parts(&err)
{
return finalize_local_execution_planning_timeout(
@@ -1585,6 +1642,10 @@ pub(crate) async fn proxy_request(
return Err(err);
}
};
observe_gateway_stage_ms(
"frontdoor_execute_stream",
execute_stream_started_at.elapsed().as_millis() as u64,
);
debug!(
event_name = "proxy_stream_local_execute_outcome",
log_type = "debug",
@@ -1622,6 +1683,7 @@ pub(crate) async fn proxy_request(
LocalExecutionRequestOutcome::NoPath => {}
}
}
let execute_sync_started_at = Instant::now();
let sync_outcome = match maybe_execute_sync_request(
&state,
&parts,
@@ -1633,6 +1695,10 @@ pub(crate) async fn proxy_request(
{
Ok(outcome) => outcome,
Err(err) => {
observe_gateway_stage_ms(
"frontdoor_execute_sync",
execute_sync_started_at.elapsed().as_millis() as u64,
);
if let Some((phase, timeout_ms)) = local_execution_planning_timeout_parts(&err) {
return finalize_local_execution_planning_timeout(
&state,
@@ -1649,6 +1715,10 @@ pub(crate) async fn proxy_request(
return Err(err);
}
};
observe_gateway_stage_ms(
"frontdoor_execute_sync",
execute_sync_started_at.elapsed().as_millis() as u64,
);
match sync_outcome {
LocalExecutionRequestOutcome::Responded(execution_runtime_response) => {
let execution_runtime_response = restore_redacted_sync_execution_response(
@@ -1673,6 +1743,7 @@ pub(crate) async fn proxy_request(
LocalExecutionRequestOutcome::NoPath => {}
}
if parts.method != http::Method::POST {
let execute_stream_started_at = Instant::now();
let stream_outcome = match maybe_execute_stream_request(
&state,
&parts,
@@ -1684,6 +1755,10 @@ pub(crate) async fn proxy_request(
{
Ok(outcome) => outcome,
Err(err) => {
observe_gateway_stage_ms(
"frontdoor_execute_stream",
execute_stream_started_at.elapsed().as_millis() as u64,
);
if let Some((phase, timeout_ms)) = local_execution_planning_timeout_parts(&err)
{
return finalize_local_execution_planning_timeout(
@@ -1701,6 +1776,10 @@ pub(crate) async fn proxy_request(
return Err(err);
}
};
observe_gateway_stage_ms(
"frontdoor_execute_stream",
execute_stream_started_at.elapsed().as_millis() as u64,
);
match stream_outcome {
LocalExecutionRequestOutcome::Responded(execution_runtime_response) => {
let execution_runtime_response = restore_redacted_stream_execution_response(
+5 -1
View File
@@ -63,15 +63,18 @@ pub(crate) use aether_provider_transport as provider_transport;
mod rate_limit;
mod request_candidate_queue;
mod request_candidate_runtime;
mod request_diagnostics;
mod roles;
mod router;
mod routing;
mod scheduler;
mod server_chan_push;
mod stage_metrics;
mod state;
mod system_features;
mod task_runtime;
mod tunnel;
mod upstream_admission;
mod usage;
mod video_tasks;
mod wallet_runtime;
@@ -126,7 +129,8 @@ fn insert_header_if_missing(
if headers.contains_key(key) {
return Ok(());
}
let name = HeaderName::from_static(key);
let name = HeaderName::from_bytes(key.as_bytes())
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let value =
HeaderValue::from_str(value).map_err(|err| GatewayError::Internal(err.to_string()))?;
headers.insert(name, value);
+202 -7
View File
@@ -236,6 +236,11 @@ const AUTO_SERVER_SQL_POOL_MIN_CONNECTIONS_FLOOR: u32 = 4;
const AUTO_SERVER_SQL_POOL_MIN_CONNECTIONS_CAP: u32 = 16;
const AUTO_SERVER_SQL_POOL_MAX_CONNECTIONS_FLOOR: u32 = 20;
const AUTO_SERVER_SQL_POOL_MAX_CONNECTIONS_CAP: u32 = 100;
const DEFAULT_GATEWAY_LISTEN_BACKLOG: i32 = 8192;
const MIN_GATEWAY_LISTEN_BACKLOG: i32 = 128;
const MAX_GATEWAY_LISTEN_BACKLOG: i32 = 65_535;
const DEFAULT_GATEWAY_LISTENER_SHARDS: usize = 1;
const MAX_GATEWAY_LISTENER_SHARDS: usize = 64;
fn env_var_trimmed(name: &str) -> Option<String> {
std::env::var(name)
.ok()
@@ -570,6 +575,34 @@ struct GatewayUsageArgs {
default_value_t = 5_000
)]
queue_reclaim_interval_ms: u64,
#[arg(
long,
env = "AETHER_GATEWAY_USAGE_ENQUEUE_RETRY_BUFFER_CAPACITY",
default_value_t = 131_072
)]
enqueue_retry_buffer_capacity: usize,
#[arg(
long,
env = "AETHER_GATEWAY_USAGE_ENQUEUE_RETRY_WORKERS",
default_value_t = 4
)]
enqueue_retry_workers: usize,
#[arg(
long,
env = "AETHER_GATEWAY_USAGE_ENQUEUE_RETRY_INITIAL_BACKOFF_MS",
default_value_t = 10
)]
enqueue_retry_initial_backoff_ms: u64,
#[arg(
long,
env = "AETHER_GATEWAY_USAGE_ENQUEUE_RETRY_MAX_BACKOFF_MS",
default_value_t = 1_000
)]
enqueue_retry_max_backoff_ms: u64,
}
impl GatewayUsageArgs {
@@ -587,6 +620,12 @@ impl GatewayUsageArgs {
reclaim_idle_ms: self.queue_reclaim_idle_ms.max(1),
reclaim_count: self.queue_reclaim_count.max(1),
reclaim_interval_ms: self.queue_reclaim_interval_ms.max(1),
enqueue_retry_buffer_capacity: self.enqueue_retry_buffer_capacity.max(1),
enqueue_retry_workers: self.enqueue_retry_workers.clamp(1, 64),
enqueue_retry_initial_backoff_ms: self.enqueue_retry_initial_backoff_ms.max(1),
enqueue_retry_max_backoff_ms: self
.enqueue_retry_max_backoff_ms
.max(self.enqueue_retry_initial_backoff_ms.max(1)),
}
}
}
@@ -755,6 +794,20 @@ struct Args {
#[arg(long, env = "APP_PORT", default_value_t = 8084)]
app_port: u16,
#[arg(
long,
env = "AETHER_GATEWAY_LISTEN_BACKLOG",
default_value_t = DEFAULT_GATEWAY_LISTEN_BACKLOG
)]
listen_backlog: i32,
#[arg(
long,
env = "AETHER_GATEWAY_LISTENER_SHARDS",
default_value_t = DEFAULT_GATEWAY_LISTENER_SHARDS
)]
listener_shards: usize,
/// 容器内健康检查入口:根据当前 bind 端口探测本地 /health。
#[arg(long, hide = true, default_value_t = false)]
healthcheck: bool,
@@ -1004,6 +1057,98 @@ fn gateway_bind_addr(app_port: u16) -> Result<std::net::SocketAddr, std::io::Err
)))
}
fn gateway_listen_backlog(backlog: i32) -> i32 {
backlog.clamp(MIN_GATEWAY_LISTEN_BACKLOG, MAX_GATEWAY_LISTEN_BACKLOG)
}
fn gateway_listener_shards(shards: usize) -> usize {
shards.clamp(1, MAX_GATEWAY_LISTENER_SHARDS)
}
fn gateway_listener(
bind_addr: std::net::SocketAddr,
backlog: i32,
reuse_port: bool,
) -> Result<tokio::net::TcpListener, std::io::Error> {
let domain = match bind_addr {
std::net::SocketAddr::V4(_) => socket2::Domain::IPV4,
std::net::SocketAddr::V6(_) => socket2::Domain::IPV6,
};
let socket = socket2::Socket::new(domain, socket2::Type::STREAM, Some(socket2::Protocol::TCP))?;
socket.set_reuse_address(true)?;
if reuse_port {
set_gateway_listener_reuse_port(&socket)?;
}
socket.set_nonblocking(true)?;
socket.set_tcp_nodelay(true)?;
socket.bind(&bind_addr.into())?;
socket.listen(gateway_listen_backlog(backlog))?;
tokio::net::TcpListener::from_std(socket.into())
}
#[cfg(unix)]
fn set_gateway_listener_reuse_port(socket: &socket2::Socket) -> Result<(), std::io::Error> {
socket.set_reuse_port(true)
}
#[cfg(not(unix))]
fn set_gateway_listener_reuse_port(_socket: &socket2::Socket) -> Result<(), std::io::Error> {
Err(std::io::Error::new(
std::io::ErrorKind::Unsupported,
"AETHER_GATEWAY_LISTENER_SHARDS > 1 requires SO_REUSEPORT support",
))
}
fn gateway_listeners(
bind_addr: std::net::SocketAddr,
backlog: i32,
shards: usize,
) -> Result<Vec<tokio::net::TcpListener>, std::io::Error> {
let shards = gateway_listener_shards(shards);
let mut listeners = Vec::with_capacity(shards);
for _ in 0..shards {
listeners.push(gateway_listener(bind_addr, backlog, shards > 1)?);
}
Ok(listeners)
}
async fn serve_gateway_router(
listeners: Vec<tokio::net::TcpListener>,
router: axum::Router,
) -> Result<(), Box<dyn std::error::Error>> {
if listeners.len() == 1 {
let listener = listeners
.into_iter()
.next()
.ok_or_else(|| std::io::Error::other("gateway listener set is empty"))?;
axum::serve(
listener,
router.into_make_service_with_connect_info::<std::net::SocketAddr>(),
)
.await?;
return Ok(());
}
let mut servers = tokio::task::JoinSet::new();
for listener in listeners {
let router = router.clone();
servers.spawn(async move {
axum::serve(
listener,
router.into_make_service_with_connect_info::<std::net::SocketAddr>(),
)
.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?;
}
Ok(())
}
fn resolve_local_http_base_url(app_port: u16) -> Result<String, std::io::Error> {
Ok(format!("http://127.0.0.1:{}", validate_app_port(app_port)?))
}
@@ -1323,6 +1468,20 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
);
}
state.bootstrap_admin_from_env().await?;
match state.prewarm_chat_pii_redaction_runtime_config().await {
Ok(enabled) => {
info!(
chat_pii_redaction_enabled = enabled,
"prewarmed chat pii redaction runtime config"
);
}
Err(err) => {
warn!(
error = %err,
"failed to prewarm chat pii redaction runtime config"
);
}
}
let background_tasks = if args.node_role.spawns_background_tasks() {
Some(state.spawn_background_tasks())
@@ -1333,7 +1492,9 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
);
None
};
let listener = tokio::net::TcpListener::bind(bind_addr).await?;
let listen_backlog = gateway_listen_backlog(args.listen_backlog);
let listener_shards = gateway_listener_shards(args.listener_shards);
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 api_router = build_router_with_state(state);
@@ -1353,17 +1514,15 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
log_type = "ops",
bind = %bind_addr,
app_port,
listen_backlog,
listener_shards,
public_url = %public_base_url,
healthcheck_url = %frontdoor_health_url,
legacy_route_policy = "fail_closed",
"aether-gateway ready"
);
axum::serve(
listener,
router.into_make_service_with_connect_info::<std::net::SocketAddr>(),
)
.await?;
serve_gateway_router(listeners, router).await?;
if let Some(background_tasks) = background_tasks {
background_tasks.shutdown().await;
}
@@ -1730,7 +1889,8 @@ mod tests {
DatabaseDriverArg, DeploymentTopologyArg, GatewayDataArgs, GatewayFrontdoorArgs,
GatewayLogDestinationArg, GatewayLogFormatArg, GatewayLogRotationArg, GatewayLoggingArgs,
GatewayRateLimitArgs, GatewayUsageArgs, NodeRoleArg, RuntimeBackendArg,
VideoTaskTruthSourceArg,
VideoTaskTruthSourceArg, DEFAULT_GATEWAY_LISTENER_SHARDS, DEFAULT_GATEWAY_LISTEN_BACKLOG,
MAX_GATEWAY_LISTENER_SHARDS, MAX_GATEWAY_LISTEN_BACKLOG, MIN_GATEWAY_LISTEN_BACKLOG,
};
use aether_data::{DatabaseDriver, SqlDatabaseConfig, SqlPoolConfig};
use aether_gateway::AppState;
@@ -1739,6 +1899,8 @@ mod tests {
Args {
command: None,
app_port: 8084,
listen_backlog: DEFAULT_GATEWAY_LISTEN_BACKLOG,
listener_shards: DEFAULT_GATEWAY_LISTENER_SHARDS,
healthcheck: false,
healthcheck_timeout_ms: 3_000,
deployment_topology: DeploymentTopologyArg::SingleNode,
@@ -1789,6 +1951,10 @@ mod tests {
queue_reclaim_idle_ms: 30_000,
queue_reclaim_count: 500,
queue_reclaim_interval_ms: 5_000,
enqueue_retry_buffer_capacity: 131_072,
enqueue_retry_workers: 4,
enqueue_retry_initial_backoff_ms: 10,
enqueue_retry_max_backoff_ms: 1_000,
},
frontdoor: GatewayFrontdoorArgs {
environment: "development".to_string(),
@@ -1825,6 +1991,35 @@ mod tests {
assert_eq!(error.kind(), std::io::ErrorKind::InvalidInput);
}
#[test]
fn clamps_gateway_listen_backlog() {
assert_eq!(
super::gateway_listen_backlog(MIN_GATEWAY_LISTEN_BACKLOG - 1),
MIN_GATEWAY_LISTEN_BACKLOG
);
assert_eq!(
super::gateway_listen_backlog(DEFAULT_GATEWAY_LISTEN_BACKLOG),
DEFAULT_GATEWAY_LISTEN_BACKLOG
);
assert_eq!(
super::gateway_listen_backlog(MAX_GATEWAY_LISTEN_BACKLOG + 1),
MAX_GATEWAY_LISTEN_BACKLOG
);
}
#[test]
fn clamps_gateway_listener_shards() {
assert_eq!(super::gateway_listener_shards(0), 1);
assert_eq!(
super::gateway_listener_shards(DEFAULT_GATEWAY_LISTENER_SHARDS),
DEFAULT_GATEWAY_LISTENER_SHARDS
);
assert_eq!(
super::gateway_listener_shards(MAX_GATEWAY_LISTENER_SHARDS + 1),
MAX_GATEWAY_LISTENER_SHARDS
);
}
#[test]
fn explicit_migrate_runtime_config_enables_data_logs() {
let mut args = test_args();
@@ -18,6 +18,9 @@ use crate::log_ids::short_request_id;
#[derive(Debug, Clone, Copy)]
pub(crate) struct RequestLogEmitted;
#[derive(Debug, Clone, Copy)]
pub(crate) struct GatewayRequestAcceptedAt(pub(crate) Instant);
fn is_usage_detail_path(path: &str) -> bool {
let Some(detail_id) = path.strip_prefix("/api/admin/usage/") else {
return false;
@@ -59,6 +62,9 @@ pub(crate) fn sanitize_access_log_path(path: &str) -> String {
pub(crate) async fn access_log_middleware(mut request: Request<Body>, next: Next) -> Response {
let started_at = Instant::now();
request
.extensions_mut()
.insert(GatewayRequestAcceptedAt(started_at));
let method = request.method().clone();
let raw_path = request
.uri()
+2 -1
View File
@@ -3,7 +3,8 @@ mod frontdoor_cors;
mod strip_cf_headers;
pub(crate) use access_log::{
access_log_middleware, sanitize_access_log_path, should_downgrade_access_log, RequestLogEmitted,
access_log_middleware, sanitize_access_log_path, should_downgrade_access_log,
GatewayRequestAcceptedAt, RequestLogEmitted,
};
pub(crate) use frontdoor_cors::frontdoor_cors_middleware;
pub use strip_cf_headers::strip_cf_headers_middleware;
@@ -43,16 +43,22 @@ use crate::{
};
const POOL_SCORE_FEEDBACK_GATE_MAX_ENTRIES: usize = 50_000;
const HEALTH_SUCCESS_PERSIST_GATE_MAX_ENTRIES: usize = 50_000;
const POOL_SCORE_SUCCESS_FEEDBACK_MIN_INTERVAL_ENV: &str =
"AETHER_GATEWAY_POOL_SCORE_SUCCESS_FEEDBACK_MIN_INTERVAL_SECS";
const POOL_SCORE_FAILURE_FEEDBACK_MIN_INTERVAL_ENV: &str =
"AETHER_GATEWAY_POOL_SCORE_FAILURE_FEEDBACK_MIN_INTERVAL_SECS";
const HEALTH_SUCCESS_PERSIST_MIN_INTERVAL_ENV: &str =
"AETHER_GATEWAY_PROVIDER_KEY_HEALTH_SUCCESS_PERSIST_MIN_INTERVAL_SECS";
const DEFAULT_POOL_SCORE_SUCCESS_FEEDBACK_MIN_INTERVAL_SECS: u64 = 5;
const DEFAULT_POOL_SCORE_FAILURE_FEEDBACK_MIN_INTERVAL_SECS: u64 = 1;
const DEFAULT_HEALTH_SUCCESS_PERSIST_MIN_INTERVAL_SECS: u64 = 5;
const MAX_POOL_SCORE_FEEDBACK_MIN_INTERVAL_SECS: u64 = 300;
static POOL_SCORE_FEEDBACK_GATE: LazyLock<ExpiringMap<String, ()>> =
LazyLock::new(ExpiringMap::new);
static HEALTH_SUCCESS_PERSIST_GATE: LazyLock<ExpiringMap<String, ()>> =
LazyLock::new(ExpiringMap::new);
static POOL_SCORE_SUCCESS_FEEDBACK_MIN_INTERVAL: LazyLock<Duration> = LazyLock::new(|| {
pool_score_feedback_interval_from_env(
POOL_SCORE_SUCCESS_FEEDBACK_MIN_INTERVAL_ENV,
@@ -65,6 +71,12 @@ static POOL_SCORE_FAILURE_FEEDBACK_MIN_INTERVAL: LazyLock<Duration> = LazyLock::
DEFAULT_POOL_SCORE_FAILURE_FEEDBACK_MIN_INTERVAL_SECS,
)
});
static HEALTH_SUCCESS_PERSIST_MIN_INTERVAL: LazyLock<Duration> = LazyLock::new(|| {
pool_score_feedback_interval_from_env(
HEALTH_SUCCESS_PERSIST_MIN_INTERVAL_ENV,
DEFAULT_HEALTH_SUCCESS_PERSIST_MIN_INTERVAL_SECS,
)
});
#[derive(Debug, Clone, Copy)]
pub(crate) struct LocalExecutionEffectContext<'a> {
@@ -645,6 +657,14 @@ async fn record_health_failure_effect(
.or(current_key.circuit_breaker_by_format.as_ref())
};
if !provider_key_health_success_persist_gate_allows(
&context.plan.key_id,
api_format,
circuit_breaker_update_owned.is_some(),
) {
return;
}
if let Err(err) = state
.update_provider_catalog_key_health_state(
&context.plan.key_id,
@@ -711,6 +731,8 @@ async fn record_health_success_effect(
.or(current_key.circuit_breaker_by_format.as_ref())
};
clear_provider_key_health_success_persist_gate(&context.plan.key_id, api_format);
if let Err(err) = state
.update_provider_catalog_key_health_state(
&context.plan.key_id,
@@ -727,6 +749,36 @@ async fn record_health_success_effect(
}
}
fn provider_key_health_success_persist_gate_allows(
key_id: &str,
api_format: &str,
closes_circuit: bool,
) -> bool {
if closes_circuit {
return true;
}
let min_interval = *HEALTH_SUCCESS_PERSIST_MIN_INTERVAL;
if min_interval.is_zero() {
return true;
}
let key = provider_key_health_success_persist_gate_key(key_id, api_format);
HEALTH_SUCCESS_PERSIST_GATE.insert_if_absent_fresh(
key,
(),
min_interval,
HEALTH_SUCCESS_PERSIST_GATE_MAX_ENTRIES,
)
}
fn clear_provider_key_health_success_persist_gate(key_id: &str, api_format: &str) {
let key = provider_key_health_success_persist_gate_key(key_id, api_format);
HEALTH_SUCCESS_PERSIST_GATE.remove(&key);
}
fn provider_key_health_success_persist_gate_key(key_id: &str, api_format: &str) -> String {
format!("success:{key_id}:{api_format}")
}
async fn record_stream_pool_success_effect(
state: &AppState,
context: LocalExecutionEffectContext<'_>,
@@ -2612,6 +2664,88 @@ mod tests {
);
}
#[tokio::test]
async fn health_success_projection_is_rate_limited_until_failure_resets_gate() {
let state = health_state();
let plan = sample_plan();
apply_local_execution_effect(
&state,
LocalExecutionEffectContext {
plan: &plan,
report_context: None,
},
LocalExecutionEffect::HealthSuccess(LocalHealthSuccessEffect),
)
.await;
let first_updated_at = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&plan.key_id))
.await
.expect("provider catalog keys should load")
.into_iter()
.next()
.expect("stored key should exist")
.updated_at_unix_secs;
apply_local_execution_effect(
&state,
LocalExecutionEffectContext {
plan: &plan,
report_context: None,
},
LocalExecutionEffect::HealthSuccess(LocalHealthSuccessEffect),
)
.await;
let second_updated_at = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&plan.key_id))
.await
.expect("provider catalog keys should load")
.into_iter()
.next()
.expect("stored key should exist")
.updated_at_unix_secs;
assert_eq!(second_updated_at, first_updated_at);
apply_local_execution_effect(
&state,
LocalExecutionEffectContext {
plan: &plan,
report_context: None,
},
LocalExecutionEffect::HealthFailure(LocalHealthFailureEffect {
status_code: 503,
classification: LocalFailoverClassification::RetryUpstreamFailure,
}),
)
.await;
apply_local_execution_effect(
&state,
LocalExecutionEffectContext {
plan: &plan,
report_context: None,
},
LocalExecutionEffect::HealthSuccess(LocalHealthSuccessEffect),
)
.await;
let stored_key = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&plan.key_id))
.await
.expect("provider catalog keys should load")
.into_iter()
.next()
.expect("stored key should exist");
assert_eq!(
stored_key
.health_by_format
.as_ref()
.and_then(|value| value.get("openai:chat"))
.and_then(|value| value.get("consecutive_failures"))
.and_then(Value::as_u64),
Some(0)
);
}
#[tokio::test]
async fn health_success_projection_closes_key_circuit_for_format() {
let mut key = sample_health_key();
+310 -4
View File
@@ -1,8 +1,9 @@
use std::collections::{BTreeMap, HashMap, HashSet};
use std::fmt;
use std::net::{Ipv4Addr, Ipv6Addr};
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, LazyLock, Mutex};
use std::time::Duration;
use std::time::{Duration, Instant};
use aether_data_contracts::DataLayerError;
use aether_runtime_state::RuntimeState;
@@ -26,6 +27,7 @@ const MAX_SENTINEL_NAMESPACE_LEN: usize = 32;
const DIRECT_RESTORE_SENTINEL_LIMIT: usize = 32;
const MAX_CACHE_SENTINEL_BYTES: usize = 128;
const MAX_CACHE_RECORD_BYTES: usize = 512;
const CHAT_PII_REDACTION_RUNTIME_CONFIG_CACHE_TTL: Duration = Duration::from_secs(5);
static EMAIL_REGEX: LazyLock<Regex> = LazyLock::new(|| {
Regex::new(r"(?i)[A-Z0-9._%+-]{1,64}@[A-Z0-9.-]{1,253}\.[A-Z]{2,63}")
@@ -368,6 +370,7 @@ struct MappingKey {
original: String,
}
#[derive(Clone)]
pub(crate) struct RedactionSession {
config: RedactionSessionConfig,
mappings: HashMap<MappingKey, RedactionMapping>,
@@ -717,10 +720,53 @@ impl fmt::Debug for MaskedChatRequest {
}
}
pub(crate) struct MaskedChatRequestValue {
pub(crate) body_json: Option<Value>,
pub(crate) session: RedactionSession,
pub(crate) redacted: bool,
}
impl fmt::Debug for MaskedChatRequestValue {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("MaskedChatRequestValue")
.field("body_json_owned", &self.body_json.is_some())
.field("session", &self.session)
.field("redacted", &self.redacted)
.finish()
}
}
#[derive(Clone)]
pub(crate) struct CachedRequestRedaction {
pub(crate) body_json: Option<Value>,
pub(crate) session: Option<RedactionSession>,
pub(crate) redacted: bool,
}
impl CachedRequestRedaction {
pub(crate) fn unredacted() -> Self {
Self {
body_json: None,
session: None,
redacted: false,
}
}
pub(crate) fn redacted(body_json: Value, session: RedactionSession) -> Self {
Self {
body_json: Some(body_json),
session: Some(session),
redacted: true,
}
}
}
#[derive(Default, Clone)]
pub(crate) struct RedactionSessionSlot {
session: Arc<Mutex<Option<RedactionSession>>>,
sessions_by_candidate: Arc<Mutex<HashMap<String, RedactionSession>>>,
request_redactions: Arc<Mutex<HashMap<String, CachedRequestRedaction>>>,
}
impl RedactionSessionSlot {
@@ -747,6 +793,10 @@ impl RedactionSessionSlot {
.lock()
.expect("redaction session candidate slot should lock")
.clear();
self.request_redactions
.lock()
.expect("redaction request cache slot should lock")
.clear();
}
pub(crate) fn take(&self) -> Option<RedactionSession> {
@@ -794,6 +844,25 @@ impl RedactionSessionSlot {
_ => None,
}
}
pub(crate) fn cached_request_redaction(&self, key: &str) -> Option<CachedRequestRedaction> {
self.request_redactions
.lock()
.expect("redaction request cache slot should lock")
.get(key)
.cloned()
}
pub(crate) fn put_cached_request_redaction(
&self,
key: impl Into<String>,
redaction: CachedRequestRedaction,
) {
self.request_redactions
.lock()
.expect("redaction request cache slot should lock")
.insert(key.into(), redaction);
}
}
#[derive(Clone, Debug, Eq, PartialEq)]
@@ -828,6 +897,150 @@ impl Default for ChatPiiRedactionRuntimeConfig {
}
}
impl ChatPiiRedactionRuntimeConfig {
fn disabled() -> Self {
Self {
enabled: false,
rules: Vec::new(),
ttl_seconds: DEFAULT_REDACTION_TTL_SECONDS,
placeholder_prefix: DEFAULT_SENTINEL_NAMESPACE.to_string(),
}
}
}
#[derive(Debug, Default)]
pub(crate) struct ChatPiiRedactionRuntimeConfigCache {
value: Mutex<Option<(Instant, ChatPiiRedactionRuntimeConfig)>>,
loading_generation: Mutex<Option<u64>>,
generation: AtomicU64,
notify: tokio::sync::Notify,
}
enum ChatPiiRedactionRuntimeConfigLoadRegistration {
Leader(ChatPiiRedactionRuntimeConfigLoadGuard),
Follower,
Bypass,
}
struct ChatPiiRedactionRuntimeConfigLoadGuard {
cache: ChatPiiRedactionRuntimeConfigCacheHandle,
generation: u64,
active: bool,
}
impl ChatPiiRedactionRuntimeConfigLoadGuard {
fn generation(&self) -> u64 {
self.generation
}
}
impl Drop for ChatPiiRedactionRuntimeConfigLoadGuard {
fn drop(&mut self) {
if self.active {
self.cache.finish_load(self.generation);
}
}
}
impl ChatPiiRedactionRuntimeConfigCache {
fn get(&self) -> Option<ChatPiiRedactionRuntimeConfig> {
self.value.lock().ok().and_then(|guard| {
guard.as_ref().and_then(|(loaded_at, value)| {
(loaded_at.elapsed() <= CHAT_PII_REDACTION_RUNTIME_CONFIG_CACHE_TTL)
.then(|| value.clone())
})
})
}
fn get_stale(&self) -> Option<ChatPiiRedactionRuntimeConfig> {
self.value
.lock()
.ok()
.and_then(|guard| guard.as_ref().map(|(_, value)| value.clone()))
}
fn insert(&self, value: ChatPiiRedactionRuntimeConfig) {
if let Ok(mut guard) = self.value.lock() {
*guard = Some((Instant::now(), value));
}
}
fn insert_if_generation(&self, generation: u64, value: ChatPiiRedactionRuntimeConfig) {
if self.generation.load(Ordering::Acquire) == generation {
self.insert(value);
}
}
fn register_load(self: &Arc<Self>) -> ChatPiiRedactionRuntimeConfigLoadRegistration {
let generation = self.generation.load(Ordering::Acquire);
match self.loading_generation.lock() {
Ok(mut loading_generation) => {
if loading_generation.is_some() {
ChatPiiRedactionRuntimeConfigLoadRegistration::Follower
} else {
*loading_generation = Some(generation);
ChatPiiRedactionRuntimeConfigLoadRegistration::Leader(
ChatPiiRedactionRuntimeConfigLoadGuard {
cache: Arc::clone(self),
generation,
active: true,
},
)
}
}
Err(_) => ChatPiiRedactionRuntimeConfigLoadRegistration::Bypass,
}
}
fn notified(&self) -> tokio::sync::futures::Notified<'_> {
self.notify.notified()
}
fn finish_load(&self, generation: u64) {
let finished = self
.loading_generation
.lock()
.map(|mut loading_generation| {
if *loading_generation == Some(generation) {
*loading_generation = None;
true
} else {
false
}
})
.unwrap_or(false);
if finished {
self.notify.notify_waiters();
}
}
}
pub(crate) type ChatPiiRedactionRuntimeConfigCacheHandle = Arc<ChatPiiRedactionRuntimeConfigCache>;
pub(crate) fn new_chat_pii_redaction_runtime_config_cache(
) -> ChatPiiRedactionRuntimeConfigCacheHandle {
Arc::new(ChatPiiRedactionRuntimeConfigCache::default())
}
pub(crate) fn clear_chat_pii_redaction_runtime_config_cache(
cache: &ChatPiiRedactionRuntimeConfigCacheHandle,
) {
cache.clear();
}
impl ChatPiiRedactionRuntimeConfigCache {
fn clear(&self) {
self.generation.fetch_add(1, Ordering::AcqRel);
if let Ok(mut value) = self.value.lock() {
*value = None;
}
if let Ok(mut loading_generation) = self.loading_generation.lock() {
*loading_generation = None;
}
self.notify.notify_waiters();
}
}
pub(crate) struct MaskChatRequestOptions {
pub(crate) scan_limits: RedactionScanLimits,
}
@@ -1127,13 +1340,76 @@ fn parse_chat_pii_redaction_rules(
pub(crate) async fn read_chat_pii_redaction_runtime_config(
state: &crate::AppState,
) -> Result<ChatPiiRedactionRuntimeConfig, GatewayError> {
let mut config = ChatPiiRedactionRuntimeConfig::default();
config.enabled = state
let cache = Arc::clone(&state.chat_pii_redaction_runtime_config_cache);
if let Some(value) = cache.get() {
return Ok(value);
}
if let Some(value) = cache.get_stale() {
if let ChatPiiRedactionRuntimeConfigLoadRegistration::Leader(guard) = cache.register_load()
{
spawn_chat_pii_redaction_runtime_config_refresh(state.clone(), cache, guard);
}
return Ok(value);
}
loop {
let notified = cache.notified();
match cache.register_load() {
ChatPiiRedactionRuntimeConfigLoadRegistration::Bypass => {
let value = load_chat_pii_redaction_runtime_config(state).await?;
cache.insert(value.clone());
return Ok(value);
}
ChatPiiRedactionRuntimeConfigLoadRegistration::Follower => {
notified.await;
if let Some(value) = cache.get() {
return Ok(value);
}
}
ChatPiiRedactionRuntimeConfigLoadRegistration::Leader(_guard) => {
let generation = _guard.generation();
let value = load_chat_pii_redaction_runtime_config(state).await?;
cache.insert_if_generation(generation, value.clone());
return Ok(value);
}
}
}
}
fn spawn_chat_pii_redaction_runtime_config_refresh(
state: crate::AppState,
cache: ChatPiiRedactionRuntimeConfigCacheHandle,
guard: ChatPiiRedactionRuntimeConfigLoadGuard,
) {
tokio::spawn(async move {
let generation = guard.generation();
match load_chat_pii_redaction_runtime_config(&state).await {
Ok(value) => cache.insert_if_generation(generation, value),
Err(err) => {
tracing::warn!(
error = ?err,
"gateway failed to refresh chat pii redaction runtime config"
);
}
}
drop(guard);
});
}
async fn load_chat_pii_redaction_runtime_config(
state: &crate::AppState,
) -> Result<ChatPiiRedactionRuntimeConfig, GatewayError> {
let enabled = state
.read_system_config_json_value("module.chat_pii_redaction.enabled")
.await?
.as_ref()
.and_then(Value::as_bool)
.unwrap_or(config.enabled);
.unwrap_or(false);
if !enabled {
return Ok(ChatPiiRedactionRuntimeConfig::disabled());
}
let mut config = ChatPiiRedactionRuntimeConfig::default();
config.enabled = true;
config.rules = parse_chat_pii_redaction_rules(
state
.read_system_config_json_value("module.chat_pii_redaction.rules")
@@ -1291,6 +1567,35 @@ pub(crate) async fn try_mask_chat_pii_request_json_with_cache_options(
})
}
pub(crate) async fn try_mask_chat_pii_request_value_with_cache_options(
body_json: &Value,
format: ChatPiiRedactionRequestFormat,
config: RedactionSessionConfig,
options: MaskChatRequestOptions,
cache: Option<&RedisRedactionMappingCache<'_>>,
) -> Result<MaskedChatRequestValue, RedactionMaskError> {
let mut session = RedactionSession::new(config);
let mut value = body_json.clone();
session.set_collision_corpus(request_collision_corpus(format, &value));
let mut scan_state = RedactionScanState::new(options.scan_limits);
let redacted = mask_request_value_async(
format,
&mut value,
&mut session,
&mut scan_state,
options,
cache,
)
.await?;
Ok(MaskedChatRequestValue {
body_json: redacted.then_some(value),
session,
redacted,
})
}
fn request_collision_corpus(format: ChatPiiRedactionRequestFormat, value: &Value) -> Vec<String> {
match format {
ChatPiiRedactionRequestFormat::OpenAiChat => value
@@ -2752,6 +3057,7 @@ impl fmt::Debug for RedactionMatch {
}
}
#[derive(Clone)]
pub(crate) struct RedactionMapping {
pub(crate) rule_label: String,
pub(crate) kind: Option<RedactionKind>,
+142 -49
View File
@@ -1,10 +1,11 @@
use std::sync::{
atomic::{AtomicBool, Ordering},
atomic::{AtomicBool, AtomicUsize, Ordering},
Arc,
};
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use aether_runtime_state::RuntimeState;
use dashmap::DashMap;
use serde::{Deserialize, Serialize};
use tokio::task::JoinHandle;
use tracing::debug;
@@ -32,14 +33,24 @@ pub(crate) struct ProviderPoolDemandSnapshot {
}
pub(crate) struct ProviderPoolInFlightGuard {
runtime: Arc<RuntimeState>,
tokens_key: String,
token: String,
stop_renewal: Arc<AtomicBool>,
renew_handle: Option<JoinHandle<()>>,
kind: ProviderPoolInFlightGuardKind,
released: bool,
}
enum ProviderPoolInFlightGuardKind {
Local {
provider_id: String,
counter: Arc<AtomicUsize>,
},
Runtime {
runtime: Arc<RuntimeState>,
tokens_key: String,
token: String,
stop_renewal: Arc<AtomicBool>,
renew_handle: Option<JoinHandle<()>>,
},
}
impl ProviderPoolInFlightGuard {
pub(crate) async fn release(mut self) {
self.release_inner().await;
@@ -50,19 +61,29 @@ impl ProviderPoolInFlightGuard {
return;
}
self.released = true;
self.stop_renewal.store(true, Ordering::Release);
if let Some(handle) = self.renew_handle.take() {
handle.abort();
}
if let Err(err) = self
.runtime
.score_remove(&self.tokens_key, &self.token)
.await
{
debug!(
error = ?err,
"gateway provider pool demand: failed to release in-flight token"
);
match &mut self.kind {
ProviderPoolInFlightGuardKind::Local {
provider_id,
counter,
} => decrement_local_provider_in_flight(provider_id, counter),
ProviderPoolInFlightGuardKind::Runtime {
runtime,
tokens_key,
token,
stop_renewal,
renew_handle,
} => {
stop_renewal.store(true, Ordering::Release);
if let Some(handle) = renew_handle.take() {
handle.abort();
}
if let Err(err) = runtime.score_remove(tokens_key, token).await {
debug!(
error = ?err,
"gateway provider pool demand: failed to release in-flight token"
);
}
}
}
}
}
@@ -73,23 +94,37 @@ impl Drop for ProviderPoolInFlightGuard {
return;
}
self.released = true;
self.stop_renewal.store(true, Ordering::Release);
if let Some(handle) = self.renew_handle.take() {
handle.abort();
}
let runtime = self.runtime.clone();
let tokens_key = self.tokens_key.clone();
let token = self.token.clone();
if let Ok(handle) = tokio::runtime::Handle::try_current() {
handle.spawn(async move {
if let Err(err) = runtime.score_remove(&tokens_key, &token).await {
debug!(
error = ?err,
"gateway provider pool demand: failed to release dropped in-flight token"
);
match &mut self.kind {
ProviderPoolInFlightGuardKind::Local {
provider_id,
counter,
} => decrement_local_provider_in_flight(provider_id, counter),
ProviderPoolInFlightGuardKind::Runtime {
runtime,
tokens_key,
token,
stop_renewal,
renew_handle,
} => {
stop_renewal.store(true, Ordering::Release);
if let Some(handle) = renew_handle.take() {
handle.abort();
}
});
let runtime = runtime.clone();
let tokens_key = tokens_key.clone();
let token = token.clone();
if let Ok(handle) = tokio::runtime::Handle::try_current() {
handle.spawn(async move {
if let Err(err) = runtime.score_remove(&tokens_key, &token).await {
debug!(
error = ?err,
"gateway provider pool demand: failed to release dropped in-flight token"
);
}
});
}
}
}
}
}
@@ -106,6 +141,46 @@ fn in_flight_tokens_key(provider_id: &str) -> String {
format!("{PROVIDER_POOL_IN_FLIGHT_TOKENS_PREFIX}:{provider_id}")
}
fn local_provider_in_flight_counts() -> &'static DashMap<String, Arc<AtomicUsize>> {
static COUNTS: std::sync::OnceLock<DashMap<String, Arc<AtomicUsize>>> =
std::sync::OnceLock::new();
COUNTS.get_or_init(DashMap::new)
}
fn increment_local_provider_in_flight(provider_id: &str) -> Arc<AtomicUsize> {
let counter_ref = local_provider_in_flight_counts()
.entry(provider_id.to_string())
.or_insert_with(|| Arc::new(AtomicUsize::new(0)));
counter_ref.fetch_add(1, Ordering::AcqRel);
counter_ref.clone()
}
fn decrement_local_provider_in_flight(provider_id: &str, counter: &AtomicUsize) {
let mut current = counter.load(Ordering::Acquire);
while current > 0 {
match counter.compare_exchange_weak(
current,
current - 1,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => break,
Err(next) => current = next,
}
}
if counter.load(Ordering::Acquire) == 0 {
let _ = local_provider_in_flight_counts()
.remove_if(provider_id, |_, stored| stored.load(Ordering::Acquire) == 0);
}
}
fn local_provider_live_in_flight_count(provider_id: &str) -> usize {
local_provider_in_flight_counts()
.get(provider_id)
.map(|counter| counter.load(Ordering::Acquire))
.unwrap_or(0)
}
fn demand_snapshot_key(provider_id: &str) -> String {
format!("{PROVIDER_POOL_DEMAND_SNAPSHOT_PREFIX}:{provider_id}")
}
@@ -185,6 +260,17 @@ pub(crate) async fn acquire_provider_pool_in_flight_guard(
return None;
}
if runtime.is_memory() {
let counter = increment_local_provider_in_flight(provider_id);
return Some(ProviderPoolInFlightGuard {
kind: ProviderPoolInFlightGuardKind::Local {
provider_id: provider_id.to_string(),
counter,
},
released: false,
});
}
let tokens_key = in_flight_tokens_key(provider_id);
let token = build_in_flight_token(request_id, candidate_id, key_id);
match tokio::time::timeout(
@@ -221,11 +307,13 @@ pub(crate) async fn acquire_provider_pool_in_flight_guard(
);
Some(ProviderPoolInFlightGuard {
runtime,
tokens_key,
token,
stop_renewal,
renew_handle: Some(renew_handle),
kind: ProviderPoolInFlightGuardKind::Runtime {
runtime,
tokens_key,
token,
stop_renewal,
renew_handle: Some(renew_handle),
},
released: false,
})
}
@@ -238,6 +326,9 @@ pub(crate) async fn provider_pool_live_in_flight_count(
if provider_id.is_empty() {
return 0;
}
if runtime.is_memory() {
return local_provider_live_in_flight_count(provider_id);
}
let key = in_flight_tokens_key(provider_id);
let now_ms = current_unix_ms() as f64;
if let Err(err) = runtime.score_remove_by_score(&key, now_ms).await {
@@ -389,9 +480,10 @@ mod tests {
#[tokio::test]
async fn in_flight_guard_tracks_and_releases_provider_tokens() {
let runtime = Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default()));
let provider_id = "provider-guard-release";
let guard = acquire_provider_pool_in_flight_guard(
runtime.clone(),
"provider-1",
provider_id,
"request-1",
Some("candidate-1"),
"key-1",
@@ -400,14 +492,14 @@ mod tests {
.expect("guard should be acquired");
assert_eq!(
provider_pool_live_in_flight_count(runtime.as_ref(), "provider-1").await,
provider_pool_live_in_flight_count(runtime.as_ref(), provider_id).await,
1
);
guard.release().await;
assert_eq!(
provider_pool_live_in_flight_count(runtime.as_ref(), "provider-1").await,
provider_pool_live_in_flight_count(runtime.as_ref(), provider_id).await,
0
);
}
@@ -415,26 +507,27 @@ mod tests {
#[tokio::test]
async fn demand_snapshot_uses_instant_in_flight_for_fast_rise_and_ema_for_fall() {
let runtime = RuntimeState::memory(MemoryRuntimeStateConfig::default());
let provider_id = "provider-demand-snapshot";
let mut guards = Vec::new();
for idx in 0..10 {
let guard = acquire_provider_pool_in_flight_guard(
Arc::new(runtime.clone()),
"provider-1",
provider_id,
"request-1",
Some(&format!("candidate-{idx}")),
"key-1",
)
.await
.expect("guard");
std::mem::forget(guard);
guards.push(guard);
}
let high = sample_provider_pool_demand(&runtime, "provider-1", 100, 50).await;
let high = sample_provider_pool_demand(&runtime, provider_id, 100, 50).await;
assert_eq!(high.in_flight, 10);
assert_eq!(high.desired_hot, 12);
let key = in_flight_tokens_key("provider-1");
let _ = runtime.score_remove_by_score(&key, f64::INFINITY).await;
let low = sample_provider_pool_demand(&runtime, "provider-1", 100, 50).await;
drop(guards);
let low = sample_provider_pool_demand(&runtime, provider_id, 100, 50).await;
assert_eq!(low.in_flight, 0);
assert!(low.ema_in_flight > 0.0);
assert!(low.desired_hot >= PROVIDER_POOL_DEMAND_FLOOR);
@@ -22,6 +22,7 @@ const DEFAULT_QUEUE_CAPACITY: usize = 65_536;
const DEFAULT_BATCH_SIZE: usize = 512;
const DEFAULT_FLUSH_INTERVAL_MS: u64 = 50;
const DEFAULT_WORKERS: usize = 2;
const FAILED_FLUSH_RETRY_DELAY_MS: u64 = 25;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum RequestCandidateWriteMode {
@@ -53,7 +54,7 @@ impl Default for RequestCandidateQueueConfig {
batch_size: DEFAULT_BATCH_SIZE,
flush_interval: Duration::from_millis(DEFAULT_FLUSH_INTERVAL_MS),
workers: DEFAULT_WORKERS,
full_policy: RequestCandidateQueueFullPolicy::Drop,
full_policy: RequestCandidateQueueFullPolicy::Sync,
}
}
}
@@ -62,8 +63,9 @@ impl RequestCandidateQueueConfig {
pub(crate) fn from_env() -> Self {
let mut config = Self::default();
config.mode = match env_string(MODE_ENV).as_deref() {
Some("sync") | Some("inline") => RequestCandidateWriteMode::Sync,
Some("async") | Some("queued") | Some("queue") => RequestCandidateWriteMode::Async,
_ => RequestCandidateWriteMode::Sync,
_ => RequestCandidateWriteMode::Async,
};
config.capacity = env_usize(QUEUE_CAPACITY_ENV, DEFAULT_QUEUE_CAPACITY).max(1);
config.batch_size = env_usize(BATCH_SIZE_ENV, DEFAULT_BATCH_SIZE).max(1);
@@ -71,10 +73,13 @@ impl RequestCandidateQueueConfig {
Duration::from_millis(env_u64(FLUSH_INTERVAL_MS_ENV, DEFAULT_FLUSH_INTERVAL_MS).max(1));
config.workers = env_usize(WORKERS_ENV, DEFAULT_WORKERS).clamp(1, 32);
config.full_policy = match env_string(QUEUE_FULL_ENV).as_deref() {
Some("drop") | Some("best_effort") | Some("best-effort") => {
RequestCandidateQueueFullPolicy::Drop
}
Some("sync") | Some("fallback_sync") | Some("fallback-sync") => {
RequestCandidateQueueFullPolicy::Sync
}
_ => RequestCandidateQueueFullPolicy::Drop,
_ => RequestCandidateQueueFullPolicy::Sync,
};
config
}
@@ -94,6 +99,7 @@ struct RequestCandidateQueueMetrics {
flush_failed_total: AtomicU64,
flush_batches_total: AtomicU64,
flush_sql_ops_total: AtomicU64,
flush_sql_records_total: AtomicU64,
compacted_total: AtomicU64,
sync_fallback_total: AtomicU64,
}
@@ -248,13 +254,19 @@ impl RequestCandidateQueueRuntime {
),
MetricSample::new(
"request_candidate_queue_flush_sql_ops_total",
"Total repository upsert operations issued by async request candidate persistence workers after compaction.",
"Total repository batch upsert operations issued by async request candidate persistence workers after compaction.",
MetricKind::Counter,
self.metrics.flush_sql_ops_total.load(Ordering::Acquire),
),
MetricSample::new(
"request_candidate_queue_flush_sql_records_total",
"Total request candidate records submitted to repository batch upsert operations after compaction.",
MetricKind::Counter,
self.metrics.flush_sql_records_total.load(Ordering::Acquire),
),
MetricSample::new(
"request_candidate_queue_compacted_total",
"Total request candidate records compacted before async persistence because a later queued record covered the same slot and status.",
"Total request candidate records compacted before async persistence because a later queued record covered the same request candidate slot.",
MetricKind::Counter,
self.metrics.compacted_total.load(Ordering::Acquire),
),
@@ -340,7 +352,7 @@ async fn flush_batch(
return;
}
let source_count = records.len();
let records = compact_same_status_records(records);
let records = compact_records_for_flush(records);
let compacted = source_count.saturating_sub(records.len());
if compacted > 0 {
metrics
@@ -348,30 +360,51 @@ async fn flush_batch(
.fetch_add(compacted as u64, Ordering::AcqRel);
}
metrics.flush_batches_total.fetch_add(1, Ordering::AcqRel);
let source_count = records
.iter()
.map(|record| record.source_count)
.sum::<usize>();
let record_count = records.len();
let upsert_records = records
.into_iter()
.map(|record| record.record)
.collect::<Vec<_>>();
metrics.flush_sql_ops_total.fetch_add(1, Ordering::AcqRel);
metrics
.flush_sql_records_total
.fetch_add(record_count as u64, Ordering::AcqRel);
let mut failed = 0_u64;
for record in records {
metrics.flush_sql_ops_total.fetch_add(1, Ordering::AcqRel);
if let Err(err) = repository.upsert(record.record).await {
failed = failed.saturating_add(record.source_count as u64);
decrement_atomic_usize_by(&metrics.pending_current, record.source_count);
warn!(
event_name = "request_candidate_async_flush_failed",
log_type = "event",
worker_index,
error = ?err,
"gateway failed to asynchronously persist request candidate"
);
} else {
metrics
.flushed_total
.fetch_add(record.source_count as u64, Ordering::AcqRel);
decrement_atomic_usize_by(&metrics.pending_current, record.source_count);
}
let mut retry_records = Vec::new();
if let Err(err) = repository.upsert_many(upsert_records.clone()).await {
failed = source_count as u64;
decrement_atomic_usize_by(
&metrics.pending_current,
source_count.saturating_sub(record_count),
);
warn!(
event_name = "request_candidate_async_flush_failed",
log_type = "event",
worker_index,
record_count,
source_count,
error = ?err,
"gateway failed to asynchronously persist request candidate batch"
);
retry_records = upsert_records;
} else {
metrics
.flushed_total
.fetch_add(source_count as u64, Ordering::AcqRel);
decrement_atomic_usize_by(&metrics.pending_current, source_count);
}
if failed > 0 {
metrics
.flush_failed_total
.fetch_add(failed, Ordering::AcqRel);
tokio::time::sleep(Duration::from_millis(FAILED_FLUSH_RETRY_DELAY_MS)).await;
for record in retry_records {
batch.push(record);
}
}
debug!(
event_name = "request_candidate_async_flush_completed",
@@ -388,10 +421,10 @@ struct CompactedRequestCandidateRecord {
source_count: usize,
}
fn compact_same_status_records(
fn compact_records_for_flush(
records: Vec<UpsertRequestCandidateRecord>,
) -> Vec<CompactedRequestCandidateRecord> {
let mut latest_slot_status = HashMap::<(String, u32, u32), (u8, usize)>::new();
let mut latest_slot = HashMap::<(String, u32, u32), usize>::new();
let mut compacted = Vec::<CompactedRequestCandidateRecord>::with_capacity(records.len());
for record in records {
let slot = (
@@ -399,14 +432,13 @@ fn compact_same_status_records(
record.candidate_index,
record.retry_index,
);
let status = request_candidate_status_discriminant(record.status);
match latest_slot_status.get(&slot).copied() {
Some((latest_status, index)) if latest_status == status => {
merge_request_candidate_record(&mut compacted[index].record, record);
match latest_slot.get(&slot).copied() {
Some(index) => {
merge_request_candidate_record_for_flush(&mut compacted[index].record, record);
compacted[index].source_count = compacted[index].source_count.saturating_add(1);
}
_ => {
latest_slot_status.insert(slot, (status, compacted.len()));
latest_slot.insert(slot, compacted.len());
compacted.push(CompactedRequestCandidateRecord {
record,
source_count: 1,
@@ -489,6 +521,41 @@ fn merge_request_candidate_record(
);
}
fn merge_request_candidate_record_for_flush(
target: &mut UpsertRequestCandidateRecord,
incoming: UpsertRequestCandidateRecord,
) {
let target_status = target.status;
let incoming_status = incoming.status;
let next_status = merged_request_candidate_status(target_status, incoming_status);
merge_request_candidate_record(target, incoming);
target.status = next_status;
}
fn merged_request_candidate_status(
current: RequestCandidateStatus,
incoming: RequestCandidateStatus,
) -> RequestCandidateStatus {
match (
request_candidate_status_is_terminal(current),
request_candidate_status_is_terminal(incoming),
) {
(_, true) => incoming,
(true, false) => current,
(false, false) => incoming,
}
}
fn request_candidate_status_is_terminal(status: RequestCandidateStatus) -> bool {
matches!(
status,
RequestCandidateStatus::Success
| RequestCandidateStatus::Failed
| RequestCandidateStatus::Cancelled
)
}
#[cfg(test)]
fn request_candidate_status_discriminant(status: RequestCandidateStatus) -> u8 {
match status {
RequestCandidateStatus::Available => 0,
@@ -554,7 +621,7 @@ fn decrement_atomic_usize_by(value: &AtomicUsize, amount: usize) {
#[cfg(test)]
mod tests {
use super::{
compact_same_status_records, RequestCandidateQueueConfig, RequestCandidateQueueRuntime,
compact_records_for_flush, RequestCandidateQueueConfig, RequestCandidateQueueRuntime,
};
use aether_data::repository::candidates::InMemoryRequestCandidateRepository;
use aether_data::DataLayerError;
@@ -562,7 +629,7 @@ mod tests {
RequestCandidateReadRepository, RequestCandidateStatus, RequestCandidateWriteRepository,
StoredRequestCandidate, UpsertRequestCandidateRecord,
};
use std::sync::atomic::Ordering;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use std::time::Duration;
@@ -594,6 +661,46 @@ mod tests {
}
}
#[derive(Default)]
struct CountingBatchRequestCandidateRepository {
inner: InMemoryRequestCandidateRepository,
upsert_calls: AtomicUsize,
upsert_many_calls: AtomicUsize,
}
#[async_trait::async_trait]
impl RequestCandidateWriteRepository for CountingBatchRequestCandidateRepository {
async fn upsert(
&self,
candidate: UpsertRequestCandidateRecord,
) -> Result<StoredRequestCandidate, DataLayerError> {
self.upsert_calls.fetch_add(1, Ordering::AcqRel);
self.inner.upsert(candidate).await
}
async fn upsert_many(
&self,
candidates: Vec<UpsertRequestCandidateRecord>,
) -> Result<usize, DataLayerError> {
self.upsert_many_calls.fetch_add(1, Ordering::AcqRel);
let count = candidates.len();
for candidate in candidates {
self.inner.upsert(candidate).await?;
}
Ok(count)
}
async fn delete_created_before(
&self,
created_before_unix_secs: u64,
limit: usize,
) -> Result<usize, DataLayerError> {
self.inner
.delete_created_before(created_before_unix_secs, limit)
.await
}
}
fn record(
request_id: &str,
candidate_index: u32,
@@ -628,8 +735,45 @@ mod tests {
}
}
struct EnvGuard {
key: &'static str,
previous: Option<String>,
}
impl EnvGuard {
fn unset(key: &'static str) -> Self {
let previous = std::env::var(key).ok();
std::env::remove_var(key);
Self { key, previous }
}
}
impl Drop for EnvGuard {
fn drop(&mut self) {
if let Some(previous) = self.previous.as_ref() {
std::env::set_var(self.key, previous);
} else {
std::env::remove_var(self.key);
}
}
}
#[test]
fn compact_merges_same_slot_and_status_without_dropping_state_transitions() {
fn from_env_defaults_to_async_with_sync_full_fallback() {
let _mode = EnvGuard::unset(super::MODE_ENV);
let _full = EnvGuard::unset(super::QUEUE_FULL_ENV);
let config = RequestCandidateQueueConfig::from_env();
assert_eq!(config.mode, super::RequestCandidateWriteMode::Async);
assert_eq!(
config.full_policy,
super::RequestCandidateQueueFullPolicy::Sync
);
}
#[test]
fn compact_merges_same_slot_without_losing_terminal_fields() {
let mut first_success = record("req", 0, 0, RequestCandidateStatus::Success);
first_success.provider_id = Some("provider-a".to_string());
first_success.extra_data = Some(serde_json::json!({"first": true}));
@@ -637,43 +781,40 @@ mod tests {
second_success.latency_ms = Some(123);
second_success.extra_data = Some(serde_json::json!({"second": true}));
let compacted = compact_same_status_records(vec![
let compacted = compact_records_for_flush(vec![
record("req", 0, 0, RequestCandidateStatus::Pending),
first_success,
record("req", 0, 1, RequestCandidateStatus::Failed),
second_success,
]);
assert_eq!(compacted.len(), 3);
assert_eq!(compacted[0].record.status, RequestCandidateStatus::Pending);
assert_eq!(compacted[0].source_count, 1);
assert_eq!(compacted[1].record.status, RequestCandidateStatus::Success);
assert_eq!(compacted[1].source_count, 2);
assert_eq!(compacted.len(), 2);
assert_eq!(compacted[0].record.status, RequestCandidateStatus::Success);
assert_eq!(compacted[0].source_count, 3);
assert_eq!(
compacted[1].record.provider_id.as_deref(),
compacted[0].record.provider_id.as_deref(),
Some("provider-a")
);
assert_eq!(compacted[1].record.latency_ms, Some(123));
assert_eq!(compacted[0].record.latency_ms, Some(123));
assert_eq!(
compacted[1].record.extra_data,
compacted[0].record.extra_data,
Some(serde_json::json!({"first": true, "second": true}))
);
assert_eq!(compacted[2].record.status, RequestCandidateStatus::Failed);
assert_eq!(compacted[2].source_count, 1);
assert_eq!(compacted[1].record.status, RequestCandidateStatus::Failed);
assert_eq!(compacted[1].source_count, 1);
}
#[test]
fn compact_preserves_same_slot_status_order_across_transitions() {
let compacted = compact_same_status_records(vec![
record("req", 0, 0, RequestCandidateStatus::Success),
record("req", 0, 0, RequestCandidateStatus::Failed),
fn compact_keeps_terminal_status_when_later_intermediate_status_arrives() {
let compacted = compact_records_for_flush(vec![
record("req", 0, 0, RequestCandidateStatus::Success),
record("req", 0, 0, RequestCandidateStatus::Streaming),
record("req", 0, 0, RequestCandidateStatus::Unused),
]);
assert_eq!(compacted.len(), 3);
assert_eq!(compacted.len(), 1);
assert_eq!(compacted[0].record.status, RequestCandidateStatus::Success);
assert_eq!(compacted[1].record.status, RequestCandidateStatus::Failed);
assert_eq!(compacted[2].record.status, RequestCandidateStatus::Success);
assert_eq!(compacted[0].source_count, 3);
}
#[tokio::test]
@@ -708,6 +849,56 @@ mod tests {
panic!("async request candidate queue did not flush record in time");
}
#[tokio::test]
async fn async_queue_flushes_records_with_batch_repository_call() {
let repository = Arc::new(CountingBatchRequestCandidateRepository::default());
let runtime = RequestCandidateQueueRuntime::spawn(
repository.clone(),
RequestCandidateQueueConfig {
mode: super::RequestCandidateWriteMode::Async,
capacity: 16,
batch_size: 4,
flush_interval: Duration::from_millis(100),
workers: 1,
full_policy: super::RequestCandidateQueueFullPolicy::Drop,
},
);
for index in 0..4 {
runtime
.enqueue_or_fallback(record(
"req-batch",
index,
0,
RequestCandidateStatus::Success,
))
.await
.unwrap();
}
for _ in 0..50 {
if runtime.metrics.pending_current.load(Ordering::Acquire) == 0 {
assert_eq!(repository.upsert_many_calls.load(Ordering::Acquire), 1);
assert_eq!(repository.upsert_calls.load(Ordering::Acquire), 0);
assert_eq!(
runtime.metrics.flush_sql_ops_total.load(Ordering::Acquire),
1
);
assert_eq!(
runtime
.metrics
.flush_sql_records_total
.load(Ordering::Acquire),
4
);
return;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
panic!("async request candidate queue did not finish batch flush in time");
}
#[tokio::test]
async fn async_queue_preserves_same_slot_order_with_multiple_workers() {
let repository = Arc::new(DelayedPendingRequestCandidateRepository::default());
@@ -0,0 +1,200 @@
use std::collections::BTreeMap;
use std::future::Future;
use std::sync::{Arc, Mutex};
use std::time::Instant;
use aether_data::DatabasePoolSummary;
use serde_json::{Map, Value};
tokio::task_local! {
static REQUEST_DIAGNOSTICS: Arc<RequestDiagnostics>;
}
#[derive(Debug, Default)]
pub(crate) struct RequestDiagnostics {
inner: Mutex<RequestDiagnosticsInner>,
}
#[derive(Debug, Default)]
struct RequestDiagnosticsInner {
db_operations: BTreeMap<&'static str, DbOperationTiming>,
db_pool: Option<DbPoolObservation>,
}
#[derive(Debug, Clone, Copy, Default)]
struct DbOperationTiming {
count: u64,
sum_ms: u64,
max_ms: u64,
}
#[derive(Debug, Clone, Copy)]
struct DbPoolObservation {
max_checked_out: u64,
max_pool_size: u64,
min_idle: u64,
max_connections: u64,
max_usage_rate_x100: u64,
}
impl RequestDiagnostics {
fn record_db_timing_ms(&self, operation: &'static str, elapsed_ms: u64) {
let Ok(mut inner) = self.inner.lock() else {
return;
};
let timing = inner.db_operations.entry(operation).or_default();
timing.count = timing.count.saturating_add(1);
timing.sum_ms = timing.sum_ms.saturating_add(elapsed_ms);
timing.max_ms = timing.max_ms.max(elapsed_ms);
}
fn record_db_pool_summary(&self, summary: DatabasePoolSummary) {
let observation = DbPoolObservation {
max_checked_out: summary.checked_out as u64,
max_pool_size: summary.pool_size as u64,
min_idle: summary.idle as u64,
max_connections: u64::from(summary.max_connections),
max_usage_rate_x100: (summary.usage_rate * 100.0).max(0.0).round() as u64,
};
let Ok(mut inner) = self.inner.lock() else {
return;
};
inner.db_pool = Some(match inner.db_pool {
Some(existing) => DbPoolObservation {
max_checked_out: existing.max_checked_out.max(observation.max_checked_out),
max_pool_size: existing.max_pool_size.max(observation.max_pool_size),
min_idle: existing.min_idle.min(observation.min_idle),
max_connections: existing.max_connections.max(observation.max_connections),
max_usage_rate_x100: existing
.max_usage_rate_x100
.max(observation.max_usage_rate_x100),
},
None => observation,
});
}
pub(crate) fn db_timings_metadata(&self) -> Option<Value> {
let Ok(inner) = self.inner.lock() else {
return None;
};
if inner.db_operations.is_empty() && inner.db_pool.is_none() {
return None;
}
let mut total_count = 0_u64;
let mut query_total_ms = 0_u64;
let mut query_max_ms = 0_u64;
let mut operations = Map::new();
for (operation, timing) in &inner.db_operations {
total_count = total_count.saturating_add(timing.count);
query_total_ms = query_total_ms.saturating_add(timing.sum_ms);
query_max_ms = query_max_ms.max(timing.max_ms);
operations.insert(
(*operation).to_string(),
Value::Object(Map::from_iter([
("count".to_string(), Value::from(timing.count)),
("sum".to_string(), Value::from(timing.sum_ms)),
("max".to_string(), Value::from(timing.max_ms)),
])),
);
}
let mut metadata = Map::new();
if !operations.is_empty() {
metadata.insert("query_count".to_string(), Value::from(total_count));
metadata.insert("query_total".to_string(), Value::from(query_total_ms));
metadata.insert("query_max".to_string(), Value::from(query_max_ms));
metadata.insert("operations".to_string(), Value::Object(operations));
}
if let Some(pool) = inner.db_pool {
metadata.insert(
"pool".to_string(),
Value::Object(Map::from_iter([
(
"max_checked_out".to_string(),
Value::from(pool.max_checked_out),
),
("max_pool_size".to_string(), Value::from(pool.max_pool_size)),
("min_idle".to_string(), Value::from(pool.min_idle)),
(
"max_connections".to_string(),
Value::from(pool.max_connections),
),
(
"max_usage_rate".to_string(),
Value::from(pool.max_usage_rate_x100 as f64 / 100.0),
),
])),
);
}
Some(Value::Object(metadata))
}
}
pub(crate) async fn scope_request_diagnostics<F>(future: F) -> F::Output
where
F: Future,
{
REQUEST_DIAGNOSTICS
.scope(Arc::new(RequestDiagnostics::default()), future)
.await
}
pub(crate) fn current_request_diagnostics() -> Option<Arc<RequestDiagnostics>> {
REQUEST_DIAGNOSTICS.try_with(Arc::clone).ok()
}
pub(crate) async fn observe_db_operation<F>(
operation: &'static str,
pool_summary: Option<DatabasePoolSummary>,
future: F,
) -> F::Output
where
F: Future,
{
if let Some(summary) = pool_summary {
record_db_pool_summary(summary);
}
let started_at = Instant::now();
let output = future.await;
record_db_timing_ms(operation, started_at.elapsed().as_millis() as u64);
output
}
pub(crate) fn record_db_timing_ms(operation: &'static str, elapsed_ms: u64) {
if let Some(diagnostics) = current_request_diagnostics() {
diagnostics.record_db_timing_ms(operation, elapsed_ms);
}
}
pub(crate) fn record_db_pool_summary(summary: DatabasePoolSummary) {
if let Some(diagnostics) = current_request_diagnostics() {
diagnostics.record_db_pool_summary(summary);
}
}
pub(crate) fn attach_request_diagnostics_to_report_context(
report_context: Option<Value>,
diagnostics: Option<&Arc<RequestDiagnostics>>,
) -> Option<Value> {
let Some(db_timings_ms) = diagnostics.and_then(|diagnostics| diagnostics.db_timings_metadata())
else {
return report_context;
};
let mut object = match report_context {
Some(Value::Object(object)) => object,
Some(other) => Map::from_iter([("seed".to_string(), other)]),
None => Map::new(),
};
object.insert("db_timings_ms".to_string(), db_timings_ms);
Some(Value::Object(object))
}
pub(crate) fn attach_current_request_diagnostics_to_report_context(
report_context: Option<&Value>,
) -> Option<Value> {
let diagnostics = current_request_diagnostics()?;
attach_request_diagnostics_to_report_context(report_context.cloned(), Some(&diagnostics))
}
@@ -17,6 +17,9 @@ pub(super) fn build_scheduler_affinity_cache_key(
global_model_name: &str,
client_session_affinity: Option<&ClientSessionAffinity>,
) -> Option<String> {
if !has_explicit_session_affinity(client_session_affinity) {
return None;
}
let api_key_id = auth_snapshot
.map(|snapshot| snapshot.api_key_id.trim())
.filter(|value| !value.is_empty())?;
@@ -28,6 +31,12 @@ pub(super) fn build_scheduler_affinity_cache_key(
)
}
pub(super) fn has_explicit_session_affinity(
client_session_affinity: Option<&ClientSessionAffinity>,
) -> bool {
client_session_affinity.is_some_and(ClientSessionAffinity::has_session_key)
}
pub(super) fn scheduler_candidate_affinity_hash(
affinity_key: &str,
candidate: &SchedulerMinimalCandidateSelectionCandidate,
@@ -138,8 +138,11 @@ pub(crate) async fn list_selectable_enumerated_candidates_with_skip_reasons(
GatewayError,
> {
let ordering_config = runtime_state.read_scheduler_ordering_config().await?;
let priority_affinity_key =
selection::scheduling_priority_affinity_key(auth_snapshot, ordering_config.scheduling_mode);
let priority_affinity_key = selection::scheduling_priority_affinity_key(
auth_snapshot,
client_session_affinity,
ordering_config.scheduling_mode,
);
collect_selectable_enumerated_candidates_with_skip_reasons(
runtime_state,
api_format,
@@ -5,7 +5,9 @@ use crate::scheduler::config::SchedulerSchedulingMode;
use crate::GatewayError;
use aether_scheduler_core::ClientSessionAffinity;
use super::affinity::{build_scheduler_affinity_cache_key, remember_scheduler_affinity};
use super::affinity::{
build_scheduler_affinity_cache_key, has_explicit_session_affinity, remember_scheduler_affinity,
};
use super::enumeration::enumerate_scheduler_candidates;
use super::ranking::rank_scheduler_candidates;
use super::resolution::resolve_scheduler_candidate_selectability;
@@ -66,8 +68,11 @@ pub(super) async fn select_minimal_candidate(
global_model_name,
client_session_affinity,
);
let priority_affinity_key =
scheduling_priority_affinity_key(auth_snapshot, ordering_config.scheduling_mode);
let priority_affinity_key = scheduling_priority_affinity_key(
auth_snapshot,
client_session_affinity,
ordering_config.scheduling_mode,
);
let candidates = enumerate_scheduler_candidates(
selection_row_source,
api_format,
@@ -94,7 +99,9 @@ pub(super) async fn select_minimal_candidate(
.0
.into_iter()
.next();
if ordering_config.scheduling_mode == SchedulerSchedulingMode::CacheAffinity {
if ordering_config.scheduling_mode == SchedulerSchedulingMode::CacheAffinity
&& has_explicit_session_affinity(client_session_affinity)
{
if let Some(candidate) = selected.as_ref() {
remember_scheduler_affinity(
affinity_cache_key.as_deref(),
@@ -154,8 +161,11 @@ pub(super) async fn collect_selectable_candidates_with_skip_reasons(
GatewayError,
> {
let ordering_config = runtime_state.read_scheduler_ordering_config().await?;
let priority_affinity_key =
scheduling_priority_affinity_key(auth_snapshot, ordering_config.scheduling_mode);
let priority_affinity_key = scheduling_priority_affinity_key(
auth_snapshot,
client_session_affinity,
ordering_config.scheduling_mode,
);
let candidates = enumerate_scheduler_candidates(
selection_row_source,
api_format,
@@ -262,11 +272,17 @@ pub(super) async fn collect_selectable_enumerated_candidates_with_skip_reasons(
pub(super) fn scheduling_priority_affinity_key<'a>(
auth_snapshot: Option<&'a GatewayAuthApiKeySnapshot>,
client_session_affinity: Option<&ClientSessionAffinity>,
scheduling_mode: SchedulerSchedulingMode,
) -> Option<&'a str> {
if scheduling_mode == SchedulerSchedulingMode::FixedOrder {
return None;
}
if scheduling_mode == SchedulerSchedulingMode::CacheAffinity
&& !has_explicit_session_affinity(client_session_affinity)
{
return None;
}
auth_snapshot
.map(|snapshot| snapshot.api_key_id.trim())
@@ -13,7 +13,7 @@ use aether_data_contracts::repository::candidates::{
};
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey;
use aether_data_contracts::repository::quota::StoredProviderQuotaSnapshot;
use aether_scheduler_core::SchedulerMinimalCandidateSelectionCandidate;
use aether_scheduler_core::{ClientSessionAffinity, SchedulerMinimalCandidateSelectionCandidate};
use serde_json::json;
use crate::cache::SchedulerAffinityTarget;
@@ -559,9 +559,85 @@ async fn cache_affinity_promotes_cached_scheduler_affinity_candidate_when_enable
second.endpoint_id = "endpoint-b".to_string();
second.key_id = "key-b".to_string();
second.key_name = "beta".to_string();
second.provider_priority = 0;
second.key_internal_priority = 0;
second.key_global_priority_by_format = Some(json!({"openai:chat": 0}));
second.provider_priority = 10;
second.key_internal_priority = 10;
second.key_global_priority_by_format = Some(json!({"openai:chat": 10}));
let candidates = Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(vec![
first, second,
]));
let quotas = Arc::new(InMemoryProviderQuotaRepository::seed(vec![]));
let state = AppState::new()
.expect("state should build")
.with_data_state_for_tests(
GatewayDataState::with_candidate_selection_and_quota_for_tests(candidates, quotas)
.with_system_config_values_for_tests(vec![(
"scheduling_mode".to_string(),
json!("cache_affinity"),
)]),
);
let auth_snapshot = sample_auth_snapshot("affinity-key-1");
let client_session_affinity = ClientSessionAffinity::from_session_key("session-1");
let cache_key = build_scheduler_affinity_cache_key(
Some(&auth_snapshot),
"openai:chat",
"gpt-4.1",
Some(&client_session_affinity),
)
.expect("scheduler affinity cache key should build");
state.remember_scheduler_affinity_target(
&cache_key,
SchedulerAffinityTarget {
provider_id: "provider-b".to_string(),
endpoint_id: "endpoint-b".to_string(),
key_id: "key-b".to_string(),
},
Duration::from_secs(300),
100,
);
let selected = select_candidate_impl(
state.data.as_ref(),
&state,
"openai:chat",
"gpt-4.1",
false,
None,
Some(&auth_snapshot),
Some(&client_session_affinity),
100,
false,
)
.await
.expect("selection should succeed")
.expect("candidate should exist");
assert_eq!(selected.provider_id, "provider-b");
assert_eq!(selected.key_id, "key-b");
}
#[tokio::test]
async fn cache_affinity_ignores_cached_scheduler_affinity_without_client_session() {
let mut first = sample_row();
first.provider_id = "provider-a".to_string();
first.provider_name = "provider-a".to_string();
first.endpoint_id = "endpoint-a".to_string();
first.key_id = "key-a".to_string();
first.key_name = "alpha".to_string();
first.provider_priority = 0;
first.key_internal_priority = 0;
first.key_global_priority_by_format = Some(json!({"openai:chat": 0}));
let mut second = sample_row();
second.provider_id = "provider-b".to_string();
second.provider_name = "provider-b".to_string();
second.endpoint_id = "endpoint-b".to_string();
second.key_id = "key-b".to_string();
second.key_name = "beta".to_string();
second.provider_priority = 10;
second.key_internal_priority = 10;
second.key_global_priority_by_format = Some(json!({"openai:chat": 10}));
let candidates = Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(vec![
first, second,
@@ -602,8 +678,8 @@ async fn cache_affinity_promotes_cached_scheduler_affinity_candidate_when_enable
.expect("selection should succeed")
.expect("candidate should exist");
assert_eq!(selected.provider_id, "provider-b");
assert_eq!(selected.key_id, "key-b");
assert_eq!(selected.provider_id, "provider-a");
assert_eq!(selected.key_id, "key-a");
}
#[tokio::test]
@@ -623,18 +699,26 @@ async fn load_balance_selection_does_not_remember_scheduler_affinity() {
)]),
);
let auth_snapshot = sample_auth_snapshot("affinity-key-1");
let cache_key =
build_scheduler_affinity_cache_key(Some(&auth_snapshot), "openai:chat", "gpt-4.1", None)
.expect("scheduler affinity cache key should build");
let client_session_affinity = ClientSessionAffinity::from_session_key("session-1");
let cache_key = build_scheduler_affinity_cache_key(
Some(&auth_snapshot),
"openai:chat",
"gpt-4.1",
Some(&client_session_affinity),
)
.expect("scheduler affinity cache key should build");
let selected = select_candidate(
let selected = select_candidate_impl(
state.data.as_ref(),
&state,
"openai:chat",
"gpt-4.1",
false,
None,
Some(&auth_snapshot),
Some(&client_session_affinity),
100,
false,
)
.await
.expect("selection should succeed")
+335
View File
@@ -0,0 +1,335 @@
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::LazyLock;
use aether_runtime::{MetricKind, MetricLabel, MetricSample};
use serde_json::{Map, Value};
const BUCKETS_MS: [u64; 12] = [1, 5, 10, 25, 50, 100, 250, 500, 1_000, 2_500, 5_000, 10_000];
const STAGES: [&str; 56] = [
"frontdoor_handler_queue",
"frontdoor_admission",
"frontdoor_context",
"frontdoor_body_buffer",
"frontdoor_owner_forward",
"frontdoor_auth_model",
"frontdoor_rpm",
"frontdoor_local_ai_public",
"frontdoor_execute_stream",
"frontdoor_execute_sync",
"stream_candidate_slot",
"stream_path_step",
"stream_candidate_next",
"stream_candidate_source_next",
"stream_candidate_plan_build",
"stream_candidate_payload_parts",
"stream_candidate_proxy",
"stream_candidate_report_context",
"stream_candidate_decision_build",
"openai_chat_payload_parts_prepare",
"openai_chat_payload_model_directives",
"openai_chat_payload_redaction",
"chat_pii_redaction_request_cache_hit",
"chat_pii_redaction_runtime_config",
"chat_pii_redaction_feature_settings",
"chat_pii_redaction_mask_body",
"openai_chat_payload_auth_prepare",
"openai_chat_payload_body_build",
"candidate_page_load",
"candidate_page_resolve",
"pool_cursor_next_key",
"pool_score_load",
"pool_score_key_rows",
"pool_runtime_state",
"candidate_transport_snapshot",
"candidate_resolution_core",
"candidate_resolution_transport_read",
"candidate_resolution_rank",
"direct_reqwest_client_prewarm",
"stream_candidate_execute",
"stream_candidate_unused",
"stream_usage_pending",
"stream_provider_in_flight",
"stream_upstream_target_admission",
"stream_upstream_headers",
"stream_first_frame",
"stream_first_data",
"stream_response_policy",
"stream_response_ready",
"stream_total",
"direct_passthrough_upstream_body_first",
"direct_passthrough_first_client_send",
"direct_passthrough_body_send_wait",
"direct_passthrough_body_recv_first",
"direct_build_body",
"direct_send_headers",
];
const TRACE_STAGE_CAPACITY: usize = 16;
const STAGE_TRACE_MODE_ENV: &str = "AETHER_GATEWAY_STAGE_TRACE_MODE";
const STAGE_TRACE_SLOW_MS_ENV: &str = "AETHER_GATEWAY_STAGE_TRACE_SLOW_MS";
const STAGE_TRACE_SAMPLE_RATE_ENV: &str = "AETHER_GATEWAY_STAGE_TRACE_SAMPLE_RATE";
const DEFAULT_STAGE_TRACE_SLOW_MS: u64 = 1_000;
static METRICS: LazyLock<Vec<StageMetric>> =
LazyLock::new(|| STAGES.iter().map(|stage| StageMetric::new(stage)).collect());
static STAGE_TRACE_CONFIG: LazyLock<RequestStageTraceConfig> =
LazyLock::new(read_stage_trace_config);
static STAGE_TRACE_SAMPLE_COUNTER: AtomicU64 = AtomicU64::new(0);
struct StageMetric {
stage: &'static str,
count: AtomicU64,
sum_ms: AtomicU64,
max_ms: AtomicU64,
buckets: Vec<AtomicU64>,
}
impl StageMetric {
fn new(stage: &'static str) -> Self {
Self {
stage,
count: AtomicU64::new(0),
sum_ms: AtomicU64::new(0),
max_ms: AtomicU64::new(0),
buckets: BUCKETS_MS.iter().map(|_| AtomicU64::new(0)).collect(),
}
}
fn observe(&self, elapsed_ms: u64) {
self.count.fetch_add(1, Ordering::Relaxed);
self.sum_ms.fetch_add(elapsed_ms, Ordering::Relaxed);
update_max(&self.max_ms, elapsed_ms);
for (index, bucket) in BUCKETS_MS.iter().enumerate() {
if elapsed_ms <= *bucket {
self.buckets[index].fetch_add(1, Ordering::Relaxed);
}
}
}
fn samples(&self) -> Vec<MetricSample> {
let stage_label = vec![MetricLabel::new("stage", self.stage)];
let mut samples = vec![
MetricSample::new(
"gateway_stage_latency_count",
"Number of gateway stage latency observations.",
MetricKind::Counter,
self.count.load(Ordering::Relaxed),
)
.with_labels(stage_label.clone()),
MetricSample::new(
"gateway_stage_latency_sum_ms",
"Total gateway stage latency in milliseconds.",
MetricKind::Counter,
self.sum_ms.load(Ordering::Relaxed),
)
.with_labels(stage_label.clone()),
MetricSample::new(
"gateway_stage_latency_max_ms",
"Maximum observed gateway stage latency in milliseconds since process start.",
MetricKind::Gauge,
self.max_ms.load(Ordering::Relaxed),
)
.with_labels(stage_label.clone()),
];
for (index, upper_bound_ms) in BUCKETS_MS.iter().enumerate() {
samples.push(
MetricSample::new(
"gateway_stage_latency_bucket",
"Cumulative gateway stage latency observations less than or equal to the bucket upper bound.",
MetricKind::Counter,
self.buckets[index].load(Ordering::Relaxed),
)
.with_labels(vec![
MetricLabel::new("stage", self.stage),
MetricLabel::new("le_ms", upper_bound_ms.to_string()),
]),
);
}
samples
}
}
pub(crate) fn observe_gateway_stage_ms(stage: &'static str, elapsed_ms: u64) {
if let Some(metric) = METRICS.iter().find(|metric| metric.stage == stage) {
metric.observe(elapsed_ms);
}
}
pub(crate) fn gateway_stage_metric_samples() -> Vec<MetricSample> {
METRICS.iter().flat_map(StageMetric::samples).collect()
}
fn update_max(max: &AtomicU64, value: u64) {
let mut current = max.load(Ordering::Relaxed);
while value > current {
match max.compare_exchange_weak(current, value, Ordering::Relaxed, Ordering::Relaxed) {
Ok(_) => break,
Err(next) => current = next,
}
}
}
#[derive(Debug, Clone)]
pub(crate) struct RequestStageTrace {
mode: RequestStageTraceMode,
slow_ms: u64,
sampled: bool,
stages: Vec<(&'static str, u64)>,
}
#[derive(Debug, Clone, Copy)]
struct RequestStageTraceConfig {
mode: RequestStageTraceMode,
slow_ms: u64,
sample_rate: f64,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum RequestStageTraceMode {
Off,
Slow,
Sample,
All,
}
impl RequestStageTrace {
pub(crate) fn from_env() -> Self {
let config = *STAGE_TRACE_CONFIG;
let sampled = config.sample_rate > 0.0 && random_unit_sample() < config.sample_rate;
Self {
mode: config.mode,
slow_ms: config.slow_ms,
sampled,
stages: Vec::with_capacity(TRACE_STAGE_CAPACITY),
}
}
pub(crate) fn observe(&mut self, stage: &'static str, elapsed_ms: u64) {
if self.mode == RequestStageTraceMode::Off {
return;
}
if let Some((_, existing)) = self
.stages
.iter_mut()
.find(|(existing_stage, _)| *existing_stage == stage)
{
*existing = elapsed_ms;
return;
}
if self.stages.len() < TRACE_STAGE_CAPACITY {
self.stages.push((stage, elapsed_ms));
}
}
pub(crate) fn into_metadata_value(self, fallback_elapsed_ms: Option<u64>) -> Option<Value> {
if self.mode == RequestStageTraceMode::Off || self.stages.is_empty() {
return None;
}
let max_observed_ms = self
.stages
.iter()
.map(|(_, elapsed_ms)| *elapsed_ms)
.max()
.unwrap_or(0);
let fallback_elapsed_ms = fallback_elapsed_ms.unwrap_or(0);
let slow = max_observed_ms.max(fallback_elapsed_ms) >= self.slow_ms;
let should_emit = match self.mode {
RequestStageTraceMode::Off => false,
RequestStageTraceMode::Slow => slow || self.sampled,
RequestStageTraceMode::Sample => self.sampled,
RequestStageTraceMode::All => true,
};
if !should_emit {
return None;
}
let mut object = Map::new();
for (stage, elapsed_ms) in self.stages {
object.insert(stage.to_string(), Value::from(elapsed_ms));
}
Some(Value::Object(object))
}
}
pub(crate) fn observe_gateway_stage_trace_ms(
trace: &mut RequestStageTrace,
stage: &'static str,
elapsed_ms: u64,
) {
observe_gateway_stage_ms(stage, elapsed_ms);
trace.observe(stage, elapsed_ms);
}
pub(crate) fn attach_stage_trace_to_report_context(
report_context: Option<Value>,
stage_timings_ms: Option<Value>,
) -> Option<Value> {
let Some(stage_timings_ms) = stage_timings_ms else {
return report_context;
};
let mut object = match report_context {
Some(Value::Object(object)) => object,
Some(other) => Map::from_iter([("seed".to_string(), other)]),
None => Map::new(),
};
object.insert("stage_timings_ms".to_string(), stage_timings_ms);
Some(Value::Object(object))
}
fn read_stage_trace_config() -> RequestStageTraceConfig {
RequestStageTraceConfig {
mode: read_stage_trace_mode(),
slow_ms: read_stage_trace_slow_ms(),
sample_rate: read_stage_trace_sample_rate(),
}
}
fn read_stage_trace_mode() -> RequestStageTraceMode {
match std::env::var(STAGE_TRACE_MODE_ENV)
.ok()
.as_deref()
.map(str::trim)
.map(str::to_ascii_lowercase)
.as_deref()
{
Some("all") => RequestStageTraceMode::All,
Some("sample") => RequestStageTraceMode::Sample,
Some("off") | Some("none") | Some("disabled") | Some("0") => RequestStageTraceMode::Off,
_ => RequestStageTraceMode::Slow,
}
}
fn read_stage_trace_slow_ms() -> u64 {
std::env::var(STAGE_TRACE_SLOW_MS_ENV)
.ok()
.as_deref()
.map(str::trim)
.and_then(|value| value.parse::<u64>().ok())
.filter(|value| *value > 0)
.unwrap_or(DEFAULT_STAGE_TRACE_SLOW_MS)
}
fn read_stage_trace_sample_rate() -> f64 {
std::env::var(STAGE_TRACE_SAMPLE_RATE_ENV)
.ok()
.as_deref()
.map(str::trim)
.and_then(|value| value.parse::<f64>().ok())
.filter(|value| value.is_finite())
.map(|value| value.clamp(0.0, 1.0))
.unwrap_or(0.0)
}
fn random_unit_sample() -> f64 {
let mut value = STAGE_TRACE_SAMPLE_COUNTER
.fetch_add(1, Ordering::Relaxed)
.wrapping_add(0x9e37_79b9_7f4a_7c15);
value ^= value >> 12;
value ^= value << 25;
value ^= value >> 27;
let mixed = value.wrapping_mul(0x2545_f491_4f6c_dd1d);
(mixed as f64) / (u64::MAX as f64)
}
+57 -1
View File
@@ -2,6 +2,7 @@ use std::collections::HashMap;
use std::sync::atomic::AtomicU64;
use std::sync::Arc;
use std::sync::Mutex as StdMutex;
use std::sync::RwLock as StdRwLock;
use std::time::Duration;
use aether_data::repository::users::StoredUserGroup;
@@ -37,6 +38,15 @@ const MIN_LOCAL_EXECUTION_PLANNING_TIMEOUT_MS: u64 = 500;
const MAX_LOCAL_EXECUTION_PLANNING_TIMEOUT_MS: u64 = 120_000;
const LOCAL_EXECUTION_PLANNING_TIMEOUT_MS_ENV: &str =
"AETHER_GATEWAY_LOCAL_EXECUTION_PLANNING_TIMEOUT_MS";
const DEFAULT_CANDIDATE_PLANNING_GATE_LIMIT: usize = 1024;
const DEFAULT_UPSTREAM_EXECUTION_GATE_LIMIT: usize = 2000;
const DEFAULT_UPSTREAM_TARGET_GATE_LIMIT: usize = 2000;
const DEFAULT_INTERNAL_GATE_QUEUE_BUDGET_MS: u64 = 250;
const MAX_INTERNAL_GATE_QUEUE_BUDGET_MS: u64 = 5_000;
const CANDIDATE_PLANNING_GATE_LIMIT_ENV: &str = "AETHER_GATEWAY_CANDIDATE_PLANNING_GATE_LIMIT";
const UPSTREAM_EXECUTION_GATE_LIMIT_ENV: &str = "AETHER_GATEWAY_UPSTREAM_EXECUTION_GATE_LIMIT";
const UPSTREAM_TARGET_GATE_LIMIT_ENV: &str = "AETHER_GATEWAY_UPSTREAM_TARGET_GATE_LIMIT";
const INTERNAL_GATE_QUEUE_BUDGET_MS_ENV: &str = "AETHER_GATEWAY_INTERNAL_GATE_QUEUE_BUDGET_MS";
#[cfg(test)]
type TestExecutionRuntimeSyncOverrideFn = dyn Fn(
@@ -62,6 +72,10 @@ impl std::fmt::Debug for TestExecutionRuntimeSyncOverride {
pub(crate) struct FrontdoorRuntimeGuardConfig {
pub(crate) request_body_read_timeout: Duration,
pub(crate) local_execution_planning_timeout: Duration,
pub(crate) internal_gate_queue_budget: Duration,
pub(crate) candidate_planning_gate_limit: Option<usize>,
pub(crate) upstream_execution_gate_limit: Option<usize>,
pub(crate) upstream_target_gate_limit: Option<usize>,
}
impl FrontdoorRuntimeGuardConfig {
@@ -79,6 +93,24 @@ impl FrontdoorRuntimeGuardConfig {
MIN_LOCAL_EXECUTION_PLANNING_TIMEOUT_MS,
MAX_LOCAL_EXECUTION_PLANNING_TIMEOUT_MS,
),
internal_gate_queue_budget: env_duration_ms(
INTERNAL_GATE_QUEUE_BUDGET_MS_ENV,
DEFAULT_INTERNAL_GATE_QUEUE_BUDGET_MS,
1,
MAX_INTERNAL_GATE_QUEUE_BUDGET_MS,
),
candidate_planning_gate_limit: env_optional_usize(
CANDIDATE_PLANNING_GATE_LIMIT_ENV,
DEFAULT_CANDIDATE_PLANNING_GATE_LIMIT,
),
upstream_execution_gate_limit: env_optional_usize(
UPSTREAM_EXECUTION_GATE_LIMIT_ENV,
DEFAULT_UPSTREAM_EXECUTION_GATE_LIMIT,
),
upstream_target_gate_limit: env_optional_usize(
UPSTREAM_TARGET_GATE_LIMIT_ENV,
DEFAULT_UPSTREAM_TARGET_GATE_LIMIT,
),
}
}
@@ -90,6 +122,12 @@ impl FrontdoorRuntimeGuardConfig {
Self {
request_body_read_timeout,
local_execution_planning_timeout,
internal_gate_queue_budget: Duration::from_millis(
DEFAULT_INTERNAL_GATE_QUEUE_BUDGET_MS,
),
candidate_planning_gate_limit: Some(DEFAULT_CANDIDATE_PLANNING_GATE_LIMIT),
upstream_execution_gate_limit: Some(DEFAULT_UPSTREAM_EXECUTION_GATE_LIMIT),
upstream_target_gate_limit: Some(DEFAULT_UPSTREAM_TARGET_GATE_LIMIT),
}
}
}
@@ -104,6 +142,17 @@ fn env_duration_ms(key: &str, default_ms: u64, min_ms: u64, max_ms: u64) -> Dura
Duration::from_millis(ms)
}
fn env_optional_usize(key: &str, default_value: usize) -> Option<usize> {
match std::env::var(key)
.ok()
.and_then(|value| value.trim().parse::<usize>().ok())
{
Some(0) => None,
Some(value) => Some(value.max(1)),
None => Some(default_value.max(1)),
}
}
#[derive(Debug, Clone)]
pub struct AppState {
#[cfg(test)]
@@ -117,6 +166,9 @@ pub struct AppState {
pub(crate) video_task_poller: Option<VideoTaskPollerConfig>,
pub(crate) frontdoor_runtime_guards: Arc<FrontdoorRuntimeGuardConfig>,
pub(crate) request_gate: Option<Arc<ConcurrencyGate>>,
pub(crate) candidate_planning_gate: Option<Arc<ConcurrencyGate>>,
pub(crate) upstream_execution_gate: Option<Arc<ConcurrencyGate>>,
pub(crate) upstream_target_admission: Arc<crate::upstream_admission::UpstreamTargetAdmission>,
pub(crate) distributed_request_gate: Option<Arc<RuntimeSemaphore>>,
pub(crate) client: reqwest::Client,
pub(crate) auth_context_cache: Arc<AuthContextCache>,
@@ -135,13 +187,17 @@ pub struct AppState {
pub(crate) scheduler_affinity_epoch: Arc<AtomicU64>,
pub(crate) dashboard_response_cache: Arc<DashboardResponseCache>,
pub(crate) system_config_cache: Arc<SystemConfigCache>,
pub(crate) candidate_page_cache: Arc<super::super::cache::CandidatePageCache>,
pub(crate) candidate_resolved_page_cache: Arc<super::super::cache::CandidateResolvedPageCache>,
pub(crate) chat_pii_redaction_runtime_config_cache:
crate::privacy::ChatPiiRedactionRuntimeConfigCacheHandle,
pub(crate) fallback_metrics: Arc<fallback_metrics::GatewayFallbackMetrics>,
pub(crate) request_candidate_queue: Option<Arc<RequestCandidateQueueRuntime>>,
pub(crate) frontdoor_cors: Option<Arc<FrontdoorCorsConfig>>,
pub(crate) frontdoor_user_rpm: Arc<FrontdoorUserRpmLimiter>,
pub(crate) tunnel: crate::tunnel::EmbeddedTunnelState,
pub(crate) provider_transport_snapshot_cache:
Arc<StdMutex<HashMap<ProviderTransportSnapshotCacheKey, CachedProviderTransportSnapshot>>>,
Arc<StdRwLock<HashMap<ProviderTransportSnapshotCacheKey, CachedProviderTransportSnapshot>>>,
pub(crate) provider_key_rpm_resets: Arc<StdMutex<HashMap<String, u64>>>,
pub(crate) local_execution_runtime_miss_diagnostics:
Arc<StdMutex<HashMap<String, LocalExecutionRuntimeMissDiagnostic>>>,
+3 -1
View File
@@ -1,3 +1,4 @@
use std::sync::Arc;
use std::time::Duration;
use super::super::provider_transport;
@@ -5,10 +6,11 @@ use super::super::provider_transport;
pub(crate) const AUTH_API_KEY_LAST_USED_TTL: Duration = Duration::from_secs(60);
pub(crate) const AUTH_API_KEY_LAST_USED_MAX_ENTRIES: usize = 10_000;
pub(crate) const PROVIDER_TRANSPORT_SNAPSHOT_CACHE_TTL: Duration = Duration::from_secs(1);
pub(crate) const PROVIDER_TRANSPORT_SNAPSHOT_CACHE_STALE_TTL: Duration = Duration::from_secs(30);
pub(crate) const PROVIDER_TRANSPORT_SNAPSHOT_CACHE_MAX_ENTRIES: usize = 1_024;
#[derive(Debug, Clone)]
pub(crate) struct CachedProviderTransportSnapshot {
pub(crate) loaded_at: std::time::Instant,
pub(crate) snapshot: provider_transport::GatewayProviderTransportSnapshot,
pub(crate) snapshot: Arc<provider_transport::GatewayProviderTransportSnapshot>,
}
+117 -15
View File
@@ -32,7 +32,7 @@ use super::super::async_task::{
use super::super::cache::{
AuthApiKeyLastUsedCache, AuthContextCache, AuthSnapshotCache, DashboardResponseCache,
DirectPlanBypassCache, JsonValueCache, SchedulerAffinityCache, SchedulerAffinitySnapshotEntry,
SchedulerAffinityTarget, SystemConfigCache, ValueCache,
SchedulerAffinityTarget, SystemConfigCache, SystemConfigInflightRegistration, ValueCache,
};
use super::super::data::{GatewayDataConfig, GatewayDataState};
use super::super::fallback_metrics;
@@ -78,6 +78,7 @@ const AUTH_AFFECTING_SYSTEM_CONFIG_KEYS: &[&str] = &[
crate::constants::ANTIGRAVITY_BEARER_BRIDGE_CONFIG_KEY,
];
const FRONTDOOR_RPM_AFFECTING_SYSTEM_CONFIG_KEYS: &[&str] = &["rate_limit_per_minute"];
const CHAT_PII_REDACTION_SYSTEM_CONFIG_PREFIX: &str = "module.chat_pii_redaction.";
fn system_config_key_affects_scheduler(key: &str) -> bool {
let key = key.trim();
@@ -94,7 +95,19 @@ fn system_config_key_affects_frontdoor_rpm(key: &str) -> bool {
FRONTDOOR_RPM_AFFECTING_SYSTEM_CONFIG_KEYS.contains(&key)
}
fn system_config_key_affects_chat_pii_redaction(key: &str) -> bool {
key.trim()
.starts_with(CHAT_PII_REDACTION_SYSTEM_CONFIG_PREFIX)
}
impl AppState {
pub async fn prewarm_chat_pii_redaction_runtime_config(&self) -> Result<bool, String> {
crate::privacy::read_chat_pii_redaction_runtime_config(self)
.await
.map(|config| config.enabled)
.map_err(|err| format!("{err:?}"))
}
fn usage_worker_queue_for(
runtime_state: &Arc<RuntimeState>,
) -> Option<Arc<dyn RuntimeQueueStore>> {
@@ -172,6 +185,7 @@ impl AppState {
self.clear_provider_transport_snapshot_cache();
self.invalidate_scheduler_affinity_cache();
self.invalidate_auth_context_cache();
self.candidate_resolved_page_cache.clear();
self.system_config_cache.clear();
self.frontdoor_user_rpm.clear_system_default_cache();
let data = Arc::new(
@@ -179,6 +193,8 @@ impl AppState {
.clone()
.with_usage_worker_queue(Self::usage_worker_queue_for(&self.runtime_state)),
);
self.candidate_page_cache.clear();
self.candidate_resolved_page_cache.clear();
self.tunnel = crate::tunnel::EmbeddedTunnelState::with_data_and_runtime_state(
Arc::clone(&data),
self.runtime_state.clone(),
@@ -222,6 +238,7 @@ impl AppState {
http2_adaptive_window: true,
..HttpClientConfig::default()
})?;
let frontdoor_runtime_guards = Arc::new(FrontdoorRuntimeGuardConfig::from_env());
Ok(Self {
#[cfg(test)]
execution_runtime_override_base_url: execution_runtime_override_base_url
@@ -236,8 +253,20 @@ impl AppState {
VideoTaskTruthSourceMode::PythonSyncReport,
)),
video_task_poller: None,
frontdoor_runtime_guards: Arc::new(FrontdoorRuntimeGuardConfig::from_env()),
frontdoor_runtime_guards: Arc::clone(&frontdoor_runtime_guards),
request_gate: None,
candidate_planning_gate: frontdoor_runtime_guards
.candidate_planning_gate_limit
.map(|limit| Arc::new(ConcurrencyGate::new("gateway_candidate_planning", limit))),
upstream_execution_gate: frontdoor_runtime_guards
.upstream_execution_gate_limit
.map(|limit| Arc::new(ConcurrencyGate::new("gateway_upstream_execution", limit))),
upstream_target_admission: Arc::new(
crate::upstream_admission::UpstreamTargetAdmission::new(
frontdoor_runtime_guards.upstream_target_gate_limit,
frontdoor_runtime_guards.internal_gate_queue_budget,
),
),
distributed_request_gate: None,
client,
auth_context_cache: Arc::new(AuthContextCache::default()),
@@ -255,6 +284,12 @@ impl AppState {
scheduler_affinity_epoch: Arc::new(AtomicU64::new(0)),
dashboard_response_cache: Arc::new(DashboardResponseCache::default()),
system_config_cache: Arc::new(SystemConfigCache::default()),
candidate_page_cache: Arc::new(crate::cache::CandidatePageCache::default()),
candidate_resolved_page_cache: Arc::new(
crate::cache::CandidateResolvedPageCache::default(),
),
chat_pii_redaction_runtime_config_cache:
crate::privacy::new_chat_pii_redaction_runtime_config_cache(),
fallback_metrics: Arc::new(fallback_metrics::GatewayFallbackMetrics::default()),
request_candidate_queue: None,
frontdoor_cors: None,
@@ -265,7 +300,7 @@ impl AppState {
data,
runtime_state.clone(),
),
provider_transport_snapshot_cache: Arc::new(StdMutex::new(HashMap::new())),
provider_transport_snapshot_cache: Arc::new(std::sync::RwLock::new(HashMap::new())),
provider_key_rpm_resets: Arc::new(StdMutex::new(HashMap::new())),
local_execution_runtime_miss_diagnostics: Arc::new(StdMutex::new(HashMap::new())),
admin_monitoring_error_stats_reset_at: Arc::new(StdMutex::new(None)),
@@ -535,19 +570,44 @@ impl AppState {
return Ok(value);
}
let _guard = self.system_config_cache.load_guard().await;
if let Some(value) = self.system_config_cache.get(key, SYSTEM_CONFIG_CACHE_TTL) {
return Ok(value);
loop {
let notified = self.system_config_cache.notified();
match self.system_config_cache.register_load(key) {
SystemConfigInflightRegistration::Bypass => {
let value = self
.data
.find_system_config_value(key)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
self.system_config_cache.insert(
key.to_string(),
value.clone(),
SYSTEM_CONFIG_CACHE_TTL,
);
return Ok(value);
}
SystemConfigInflightRegistration::Follower => {
notified.await;
if let Some(value) = self.system_config_cache.get(key, SYSTEM_CONFIG_CACHE_TTL)
{
return Ok(value);
}
}
SystemConfigInflightRegistration::Leader(_guard) => {
let value = self
.data
.find_system_config_value(key)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
self.system_config_cache.insert(
key.to_string(),
value.clone(),
SYSTEM_CONFIG_CACHE_TTL,
);
return Ok(value);
}
}
}
let value = self
.data
.find_system_config_value(key)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
self.system_config_cache
.insert(key.to_string(), value.clone(), SYSTEM_CONFIG_CACHE_TTL);
Ok(value)
}
pub(crate) async fn upsert_system_config_json_value(
@@ -606,12 +666,19 @@ impl AppState {
if deleted && system_config_key_affects_frontdoor_rpm(key) {
self.frontdoor_user_rpm.clear_system_default_cache();
}
if deleted && system_config_key_affects_chat_pii_redaction(key) {
crate::privacy::clear_chat_pii_redaction_runtime_config_cache(
&self.chat_pii_redaction_runtime_config_cache,
);
}
Ok(deleted)
}
pub(crate) fn invalidate_provider_routing_caches(&self) {
self.data.clear_minimal_candidate_selection_cache();
self.data.clear_provider_catalog_cache();
self.candidate_page_cache.clear();
self.candidate_resolved_page_cache.clear();
self.clear_provider_transport_snapshot_cache();
self.invalidate_scheduler_affinity_cache();
}
@@ -619,6 +686,8 @@ impl AppState {
pub(crate) fn invalidate_provider_health_routing_caches(&self) {
self.data.clear_minimal_candidate_selection_cache();
self.data.clear_provider_catalog_cache();
self.candidate_page_cache.clear();
self.candidate_resolved_page_cache.clear();
self.clear_provider_transport_snapshot_cache();
}
@@ -631,6 +700,8 @@ impl AppState {
self.auth_api_key_feature_settings_cache.clear();
self.provider_quota_snapshot_cache.clear();
self.user_groups_for_user_cache.clear();
self.candidate_page_cache.clear();
self.candidate_resolved_page_cache.clear();
}
fn remember_system_config_write(&self, key: &str, value: Option<serde_json::Value>) {
@@ -645,6 +716,11 @@ impl AppState {
if system_config_key_affects_frontdoor_rpm(key) {
self.frontdoor_user_rpm.clear_system_default_cache();
}
if system_config_key_affects_chat_pii_redaction(key) {
crate::privacy::clear_chat_pii_redaction_runtime_config_cache(
&self.chat_pii_redaction_runtime_config_cache,
);
}
}
pub(crate) async fn read_admin_system_stats(
@@ -908,6 +984,18 @@ impl AppState {
self.request_gate.as_ref().map(|gate| gate.snapshot())
}
pub(crate) fn candidate_planning_concurrency_snapshot(&self) -> Option<ConcurrencySnapshot> {
self.candidate_planning_gate
.as_ref()
.map(|gate| gate.snapshot())
}
pub(crate) fn upstream_execution_concurrency_snapshot(&self) -> Option<ConcurrencySnapshot> {
self.upstream_execution_gate
.as_ref()
.map(|gate| gate.snapshot())
}
pub(crate) async fn distributed_request_concurrency_snapshot(
&self,
) -> Result<Option<RuntimeSemaphoreSnapshot>, RuntimeSemaphoreError> {
@@ -922,6 +1010,12 @@ impl AppState {
if let Some(snapshot) = self.request_concurrency_snapshot() {
samples.extend(snapshot.to_metric_samples("gateway_requests"));
}
if let Some(snapshot) = self.candidate_planning_concurrency_snapshot() {
samples.extend(snapshot.to_metric_samples("gateway_candidate_planning"));
}
if let Some(snapshot) = self.upstream_execution_concurrency_snapshot() {
samples.extend(snapshot.to_metric_samples("gateway_upstream_execution"));
}
if let Some(gate) = self.distributed_request_gate.as_ref() {
match gate.snapshot().await {
Ok(snapshot) => {
@@ -947,6 +1041,12 @@ impl AppState {
if let Some(queue) = self.request_candidate_queue.as_ref() {
samples.extend(queue.metric_samples());
}
samples.extend(
crate::execution_runtime::transport::direct_reqwest_client_cache_metric_samples(),
);
samples.extend(self.upstream_target_admission.metric_samples());
samples.extend(crate::cache::candidate_page_cache_metric_samples());
samples.extend(crate::stage_metrics::gateway_stage_metric_samples());
samples.extend(self.tunnel.metric_samples());
samples.extend(self.fallback_metrics.metric_samples());
samples
@@ -1131,6 +1231,8 @@ impl AppState {
.fetch_add(1, Ordering::AcqRel)
.saturating_add(1);
self.scheduler_affinity_cache.clear();
self.candidate_page_cache.clear();
self.candidate_resolved_page_cache.clear();
next_epoch
}
+1 -1
View File
@@ -31,7 +31,7 @@ pub(crate) use self::app::FrontdoorRuntimeGuardConfig;
pub(crate) use self::cache::{
CachedProviderTransportSnapshot, AUTH_API_KEY_LAST_USED_MAX_ENTRIES,
AUTH_API_KEY_LAST_USED_TTL, PROVIDER_TRANSPORT_SNAPSHOT_CACHE_MAX_ENTRIES,
PROVIDER_TRANSPORT_SNAPSHOT_CACHE_TTL,
PROVIDER_TRANSPORT_SNAPSHOT_CACHE_STALE_TTL, PROVIDER_TRANSPORT_SNAPSHOT_CACHE_TTL,
};
pub use self::cors::FrontdoorCorsConfig;
pub(crate) use self::types::{
+64 -39
View File
@@ -1,7 +1,7 @@
use super::{
provider_transport_snapshot_looks_refreshed, AppState, CachedProviderTransportSnapshot,
GatewayError, ProviderTransportSnapshotCacheKey, PROVIDER_TRANSPORT_SNAPSHOT_CACHE_MAX_ENTRIES,
PROVIDER_TRANSPORT_SNAPSHOT_CACHE_TTL,
PROVIDER_TRANSPORT_SNAPSHOT_CACHE_STALE_TTL,
};
use crate::handlers::shared::default_provider_key_status_snapshot;
use crate::provider_transport::LocalOAuthHttpExecutor;
@@ -20,6 +20,7 @@ use serde_json::{json, Map, Value};
use sha2::{Digest, Sha256};
use std::collections::BTreeMap;
use std::io::Read;
use std::sync::Arc;
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use aether_crypto::encrypt_python_fernet_plaintext;
@@ -479,39 +480,49 @@ impl<'a> provider_transport::LocalOAuthHttpExecutor for GatewayLocalOAuthHttpExe
impl AppState {
pub(crate) fn clear_provider_transport_snapshot_cache(&self) {
self.provider_transport_snapshot_cache
.lock()
.write()
.expect("provider transport snapshot cache should lock")
.clear();
}
fn get_cached_provider_transport_snapshot(
fn get_cached_provider_transport_snapshot_arc(
&self,
cache_key: &ProviderTransportSnapshotCacheKey,
) -> Option<provider_transport::GatewayProviderTransportSnapshot> {
let mut cache = self
.provider_transport_snapshot_cache
.lock()
.expect("provider transport snapshot cache should lock");
let cached = cache.get(cache_key).cloned()?;
if cached.loaded_at.elapsed() <= PROVIDER_TRANSPORT_SNAPSHOT_CACHE_TTL {
) -> Option<Arc<provider_transport::GatewayProviderTransportSnapshot>> {
let cached = {
let cache = self
.provider_transport_snapshot_cache
.read()
.expect("provider transport snapshot cache should lock");
cache.get(cache_key).cloned()
}?;
if cached.loaded_at.elapsed() <= PROVIDER_TRANSPORT_SNAPSHOT_CACHE_STALE_TTL {
return Some(cached.snapshot);
}
cache.remove(cache_key);
let mut cache = self
.provider_transport_snapshot_cache
.write()
.expect("provider transport snapshot cache should lock");
if cache.get(cache_key).is_some_and(|entry| {
entry.loaded_at.elapsed() > PROVIDER_TRANSPORT_SNAPSHOT_CACHE_STALE_TTL
}) {
cache.remove(cache_key);
}
None
}
fn put_cached_provider_transport_snapshot(
&self,
cache_key: ProviderTransportSnapshotCacheKey,
snapshot: provider_transport::GatewayProviderTransportSnapshot,
snapshot: Arc<provider_transport::GatewayProviderTransportSnapshot>,
) {
let mut cache = self
.provider_transport_snapshot_cache
.lock()
.write()
.expect("provider transport snapshot cache should lock");
if cache.len() >= PROVIDER_TRANSPORT_SNAPSHOT_CACHE_MAX_ENTRIES {
cache.retain(|_, entry| {
entry.loaded_at.elapsed() <= PROVIDER_TRANSPORT_SNAPSHOT_CACHE_TTL
entry.loaded_at.elapsed() <= PROVIDER_TRANSPORT_SNAPSHOT_CACHE_STALE_TTL
});
if cache.len() >= PROVIDER_TRANSPORT_SNAPSHOT_CACHE_MAX_ENTRIES {
cache.clear();
@@ -796,6 +807,41 @@ impl AppState {
None
}
pub(crate) async fn read_provider_transport_snapshot_arc(
&self,
provider_id: &str,
endpoint_id: &str,
key_id: &str,
) -> Result<
Option<Arc<crate::provider_transport::GatewayProviderTransportSnapshot>>,
GatewayError,
> {
let Some(cache_key) =
ProviderTransportSnapshotCacheKey::new(provider_id, endpoint_id, key_id)
else {
return Ok(self
.read_provider_transport_snapshot_uncached(provider_id, endpoint_id, key_id)
.await?
.map(Arc::new));
};
if let Some(snapshot) = self.get_cached_provider_transport_snapshot_arc(&cache_key) {
return Ok(Some(snapshot));
}
let snapshot = self
.read_provider_transport_snapshot_uncached(provider_id, endpoint_id, key_id)
.await?;
match snapshot {
Some(snapshot) => {
let snapshot = self.apply_global_format_conversion_override(snapshot).await;
let snapshot = Arc::new(snapshot);
self.put_cached_provider_transport_snapshot(cache_key, Arc::clone(&snapshot));
Ok(Some(snapshot))
}
None => Ok(None),
}
}
pub(crate) async fn read_provider_transport_snapshot(
&self,
provider_id: &str,
@@ -803,31 +849,10 @@ impl AppState {
key_id: &str,
) -> Result<Option<crate::provider_transport::GatewayProviderTransportSnapshot>, GatewayError>
{
let Some(cache_key) =
ProviderTransportSnapshotCacheKey::new(provider_id, endpoint_id, key_id)
else {
return self
.read_provider_transport_snapshot_uncached(provider_id, endpoint_id, key_id)
.await;
};
if let Some(snapshot) = self.get_cached_provider_transport_snapshot(&cache_key) {
return Ok(Some(
self.apply_global_format_conversion_override(snapshot).await,
));
}
let snapshot = self
.read_provider_transport_snapshot_uncached(provider_id, endpoint_id, key_id)
.await?;
if let Some(snapshot) = snapshot.as_ref() {
self.put_cached_provider_transport_snapshot(cache_key, snapshot.clone());
}
match snapshot {
Some(snapshot) => Ok(Some(
self.apply_global_format_conversion_override(snapshot).await,
)),
None => Ok(None),
}
Ok(self
.read_provider_transport_snapshot_arc(provider_id, endpoint_id, key_id)
.await?
.map(|snapshot| (*snapshot).clone()))
}
pub(crate) async fn update_provider_catalog_key_oauth_credentials(
+5
View File
@@ -21,6 +21,11 @@ impl AppState {
self
}
pub(crate) fn without_request_candidate_queue_for_tests(mut self) -> Self {
self.request_candidate_queue = None;
self
}
pub(crate) fn with_turnstile_siteverify_url_for_tests(mut self, url: &str) -> Self {
self.turnstile_siteverify_url_override = Some(url.trim().to_string());
self
@@ -0,0 +1,269 @@
use std::sync::Arc;
use std::time::{Duration, Instant};
use aether_contracts::ExecutionPlan;
use aether_runtime::{ConcurrencyGate, ConcurrencyPermit, MetricKind, MetricLabel, MetricSample};
use dashmap::DashMap;
use tokio::time::timeout;
use url::Url;
use crate::stage_metrics::observe_gateway_stage_ms;
use crate::GatewayError;
const GATE_NAME: &str = "gateway_upstream_target";
const DEFAULT_METRIC_TARGET_LIMIT: usize = 32;
const METRIC_TARGET_LIMIT_ENV: &str = "AETHER_GATEWAY_UPSTREAM_TARGET_GATE_METRIC_LIMIT";
#[derive(Debug)]
pub(crate) struct UpstreamTargetAdmission {
limit: Option<usize>,
queue_budget: Duration,
gates: DashMap<String, Arc<ConcurrencyGate>>,
}
#[derive(Debug)]
pub(crate) struct UpstreamTargetAdmissionPermit {
_permit: ConcurrencyPermit,
}
impl UpstreamTargetAdmission {
pub(crate) fn new(limit: Option<usize>, queue_budget: Duration) -> Self {
Self {
limit,
queue_budget,
gates: DashMap::new(),
}
}
pub(crate) async fn acquire(
&self,
plan: &ExecutionPlan,
trace_id: &str,
) -> Result<Option<UpstreamTargetAdmissionPermit>, GatewayError> {
let Some(limit) = self.limit else {
return Ok(None);
};
let key = upstream_target_key(plan);
let gate = self
.gates
.entry(key.clone())
.or_insert_with(|| Arc::new(ConcurrencyGate::new(GATE_NAME, limit)))
.clone();
let started_at = Instant::now();
let permit = match timeout(self.queue_budget, gate.acquire()).await {
Ok(Ok(permit)) => permit,
Ok(Err(err)) => return Err(GatewayError::Internal(err.to_string())),
Err(_) => {
tracing::debug!(
event_name = "gateway_upstream_target_admission_timeout",
log_type = "ops",
trace_id,
target = key.as_str(),
limit,
queue_budget_ms = self.queue_budget.as_millis() as u64,
"gateway upstream target admission gate timed out"
);
return Err(GatewayError::AdmissionTimeout {
trace_id: trace_id.to_string(),
gate: GATE_NAME,
queue_budget_ms: self.queue_budget.as_millis() as u64,
});
}
};
observe_gateway_stage_ms(
"stream_upstream_target_admission",
started_at.elapsed().as_millis() as u64,
);
Ok(Some(UpstreamTargetAdmissionPermit { _permit: permit }))
}
pub(crate) fn metric_samples(&self) -> Vec<MetricSample> {
let mut samples = vec![MetricSample::new(
"upstream_target_gate_active_targets",
"Number of upstream targets currently tracked by the gateway upstream target admission gates.",
MetricKind::Gauge,
self.gates.len() as u64,
)];
let Some(limit) = self.limit else {
return samples;
};
samples.push(MetricSample::new(
"upstream_target_gate_limit",
"Configured per-upstream-target admission gate limit.",
MetricKind::Gauge,
limit as u64,
));
let mut snapshots = self
.gates
.iter()
.map(|entry| {
let snapshot = entry.value().snapshot();
(
entry.key().clone(),
snapshot.in_flight,
snapshot.available_permits,
snapshot.high_watermark,
snapshot.rejected,
)
})
.collect::<Vec<_>>();
snapshots.sort_by(|left, right| {
right
.1
.cmp(&left.1)
.then_with(|| right.3.cmp(&left.3))
.then_with(|| right.4.cmp(&left.4))
});
let metric_target_limit = upstream_target_metric_limit();
for (target, in_flight, available, high_watermark, rejected) in
snapshots.into_iter().take(metric_target_limit)
{
let labels = vec![MetricLabel::new("target", target)];
samples.push(
MetricSample::new(
"upstream_target_gate_in_flight",
"Current number of in-flight operations for an upstream target admission gate.",
MetricKind::Gauge,
in_flight as u64,
)
.with_labels(labels.clone()),
);
samples.push(
MetricSample::new(
"upstream_target_gate_available_permits",
"Currently available permits for an upstream target admission gate.",
MetricKind::Gauge,
available as u64,
)
.with_labels(labels.clone()),
);
samples.push(
MetricSample::new(
"upstream_target_gate_high_watermark",
"Highest observed in-flight count for an upstream target admission gate.",
MetricKind::Gauge,
high_watermark as u64,
)
.with_labels(labels.clone()),
);
samples.push(
MetricSample::new(
"upstream_target_gate_rejected_total",
"Number of operations rejected by an upstream target admission gate.",
MetricKind::Counter,
rejected,
)
.with_labels(labels),
);
}
samples
}
}
pub(crate) fn upstream_target_key(plan: &ExecutionPlan) -> String {
let parsed = Url::parse(plan.url.as_str()).ok();
let Some(url) = parsed else {
return fallback_target_key(plan);
};
let scheme = url.scheme().to_ascii_lowercase();
let Some(host) = url.host_str().map(|host| host.to_ascii_lowercase()) else {
return fallback_target_key(plan);
};
let port = url
.port_or_known_default()
.map(|port| port.to_string())
.unwrap_or_else(|| "-".to_string());
let proxy = plan
.proxy
.as_ref()
.and_then(|proxy| proxy.url.as_deref())
.map(str::trim)
.filter(|value| !value.is_empty())
.unwrap_or("-");
format!("{scheme}://{host}:{port}|proxy={proxy}")
}
fn fallback_target_key(plan: &ExecutionPlan) -> String {
format!(
"unparsed|provider={}|endpoint={}|url={}",
plan.provider_id, plan.endpoint_id, plan.url
)
}
fn upstream_target_metric_limit() -> usize {
std::env::var(METRIC_TARGET_LIMIT_ENV)
.ok()
.and_then(|value| value.trim().parse::<usize>().ok())
.unwrap_or(DEFAULT_METRIC_TARGET_LIMIT)
}
#[cfg(test)]
mod tests {
use std::time::Duration;
use aether_contracts::{ExecutionPlan, RequestBody};
use serde_json::json;
use super::*;
fn test_plan(url: &str) -> ExecutionPlan {
ExecutionPlan {
request_id: "req-upstream-target".to_string(),
candidate_id: Some("cand-upstream-target".to_string()),
provider_name: Some("provider".to_string()),
provider_id: "provider_id".to_string(),
endpoint_id: "endpoint_id".to_string(),
key_id: "key_id".to_string(),
method: "POST".to_string(),
url: url.to_string(),
headers: Default::default(),
content_type: Some("application/json".to_string()),
content_encoding: None,
body: RequestBody::from_json(json!({"stream": true})),
stream: true,
client_api_format: "openai".to_string(),
provider_api_format: "openai".to_string(),
model_name: Some("model".to_string()),
proxy: None,
transport_profile: None,
timeouts: None,
}
}
#[test]
fn upstream_target_key_ignores_path_and_query() {
let left = test_plan("http://127.0.0.1:18181/v1/chat/completions?x=1");
let right = test_plan("http://127.0.0.1:18181/v1/responses");
assert_eq!(upstream_target_key(&left), upstream_target_key(&right));
}
#[tokio::test]
async fn acquire_times_out_when_target_gate_is_saturated() {
let admission = UpstreamTargetAdmission::new(Some(1), Duration::from_millis(1));
let plan = test_plan("http://127.0.0.1:18181/v1/chat/completions");
let _first = admission
.acquire(&plan, "trace-upstream-target")
.await
.expect("first acquire should succeed")
.expect("gate enabled");
let err = admission
.acquire(&plan, "trace-upstream-target")
.await
.expect_err("second acquire should time out");
assert!(matches!(
err,
GatewayError::AdmissionTimeout {
gate: "gateway_upstream_target",
..
}
));
}
}