mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-08 02:17:46 +08:00
Improve gateway scheduling and runtime admission
This commit is contained in:
@@ -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))
|
||||
}
|
||||
|
||||
+28
-1
@@ -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
@@ -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
@@ -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
@@ -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);
|
||||
}
|
||||
}
|
||||
Vendored
+14
-3
@@ -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
@@ -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();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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";
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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())
|
||||
{
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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>,
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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>>>,
|
||||
|
||||
@@ -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>,
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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::{
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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",
|
||||
..
|
||||
}
|
||||
));
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user