Improve gateway scheduling and runtime admission

This commit is contained in:
elky
2026-06-24 01:53:45 +08:00
parent cf0af8fa1e
commit d336d1a7fa
87 changed files with 9671 additions and 804 deletions
@@ -8,6 +8,12 @@ use crate::scheduler::affinity::SCHEDULER_AFFINITY_TTL;
const PLANNER_SCHEDULER_AFFINITY_MAX_ENTRIES: usize = 10_000;
pub(crate) fn has_explicit_session_affinity(
client_session_affinity: Option<&ClientSessionAffinity>,
) -> bool {
client_session_affinity.is_some_and(ClientSessionAffinity::has_session_key)
}
pub(crate) fn read_cached_scheduler_affinity_target(
state: PlannerAppState<'_>,
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
@@ -15,6 +21,9 @@ pub(crate) fn read_cached_scheduler_affinity_target(
client_api_format: &str,
requested_model: Option<&str>,
) -> Option<SchedulerAffinityTarget> {
if !has_explicit_session_affinity(client_session_affinity) {
return None;
}
let requested_model = requested_model
.map(str::trim)
.filter(|value| !value.is_empty())?;
@@ -41,6 +50,9 @@ pub(crate) fn remember_scheduler_affinity_for_candidate(
requested_model: &str,
candidate: &SchedulerMinimalCandidateSelectionCandidate,
) {
if !has_explicit_session_affinity(client_session_affinity) {
return;
}
remember_scheduler_affinity_for_candidate_at_epoch(
state,
auth_snapshot,
@@ -61,6 +73,9 @@ pub(crate) fn remember_scheduler_affinity_for_candidate_at_epoch(
candidate: &SchedulerMinimalCandidateSelectionCandidate,
expected_epoch: Option<u64>,
) {
if !has_explicit_session_affinity(client_session_affinity) {
return;
}
let Some(api_key_id) = auth_snapshot
.map(|snapshot| snapshot.api_key_id.trim())
.filter(|value| !value.is_empty())
@@ -38,12 +38,19 @@ use crate::ai_serving::planner::pool_scheduler::PoolKeyCursor;
use crate::ai_serving::planner::runtime_miss::record_local_runtime_candidate_skip_reason;
use crate::ai_serving::planner::CandidateFailureDiagnostic;
use crate::ai_serving::{GatewayAuthApiKeySnapshot, PlannerAppState};
use crate::cache::{
candidate_page_cache_stale_ttl, candidate_page_cache_ttl_from_env,
record_candidate_page_resolve_cache_follower_wait, record_candidate_page_resolve_cache_hit,
record_candidate_page_resolve_cache_load, record_candidate_page_resolve_cache_miss,
CacheLoadObserver, CandidateResolvedPageCacheKey, CandidateResolvedPageSnapshot,
};
use crate::clock::current_unix_ms;
use crate::dispatch::refs::dispatch_ref_for_local_candidate;
use crate::handlers::shared::provider_pool::admin_provider_pool_config_from_config_value;
use crate::orchestration::{local_attempt_slot_count, ExecutionAttemptIdentity};
use crate::scheduler::candidate::is_auth_api_key_concurrency_limit_skip_reason;
use crate::scheduler::config::SchedulerSchedulingMode;
use crate::stage_metrics::observe_gateway_stage_ms;
use crate::{AppState, GatewayError};
const POOL_KEY_RETRY_INDEX_STRIDE: u32 = 100;
@@ -101,16 +108,20 @@ impl<'a> LocalExecutionCandidateAttemptSource<'a> {
Self { items }
}
pub(crate) async fn next_attempt(&mut self) -> Option<LocalExecutionCandidateAttempt> {
pub(crate) async fn next_attempt(
&mut self,
) -> Result<Option<LocalExecutionCandidateAttempt>, GatewayError> {
loop {
let front = self.items.front_mut()?;
let Some(front) = self.items.front_mut() else {
return Ok(None);
};
match front {
LocalExecutionCandidateAttemptSourceItem::Static { attempts } => {
if let Some(attempt) = next_attempt_from_dispatch_sequence(attempts) {
if dispatch_sequence_exhausted(attempts) {
self.items.pop_front();
}
return Some(attempt);
return Ok(Some(attempt));
}
self.items.pop_front();
}
@@ -121,7 +132,7 @@ impl<'a> LocalExecutionCandidateAttemptSource<'a> {
pool_exhaustion_persistence,
} => {
if let Some(attempt) = next_attempt_from_dispatch_sequence(pending_attempts) {
return Some(attempt);
return Ok(Some(attempt));
}
let Some(candidate) = cursor.next_key().await else {
if let Some(skipped) = cursor.exhausted_group_skipped_candidate() {
@@ -146,11 +157,11 @@ impl<'a> LocalExecutionCandidateAttemptSource<'a> {
);
}
LocalExecutionCandidateAttemptSourceItem::RequestedModelPage { cursor } => {
let Some(attempt) = cursor.next_attempt().await else {
let Some(attempt) = cursor.next_attempt().await? else {
self.items.pop_front();
continue;
};
return Some(attempt);
return Ok(Some(attempt));
}
}
}
@@ -726,8 +737,10 @@ where
auth_snapshot,
routing_policy,
client_session_affinity,
request_auth_channel,
use_api_format_alias_match,
key_mode,
Some(trace_id),
)
.await;
let mut cursor = RequestedModelAttemptPageCursor {
@@ -755,11 +768,14 @@ where
remembered_affinity: false,
scheduler_cache_affinity_enabled,
auth_api_key_concurrency_wait_deadline: None,
deferred_error: None,
};
cursor.load_next_page().await;
if let Err(error) = cursor.load_next_page().await {
cursor.deferred_error = Some(error);
}
let candidate_count = cursor.candidate_count;
let mut items = VecDeque::new();
if !cursor.pending_items.is_empty() {
if !cursor.pending_items.is_empty() || cursor.deferred_error.is_some() {
items.push_back(
LocalExecutionCandidateAttemptSourceItem::RequestedModelPage {
cursor: Box::new(cursor),
@@ -797,34 +813,58 @@ struct RequestedModelAttemptPageCursor<'a> {
remembered_affinity: bool,
scheduler_cache_affinity_enabled: bool,
auth_api_key_concurrency_wait_deadline: Option<Instant>,
deferred_error: Option<GatewayError>,
}
impl<'a> RequestedModelAttemptPageCursor<'a> {
async fn next_attempt(&mut self) -> Option<LocalExecutionCandidateAttempt> {
async fn next_attempt(
&mut self,
) -> Result<Option<LocalExecutionCandidateAttempt>, GatewayError> {
if let Some(error) = self.deferred_error.take() {
return Err(error);
}
loop {
if let Some(attempt) = pop_attempt_from_items(&mut self.pending_items).await {
return Some(attempt);
return Ok(Some(attempt));
}
if !self.load_next_page().await {
return None;
if !self.load_next_page().await? {
return Ok(None);
}
}
}
async fn load_next_page(&mut self) -> bool {
async fn load_next_page(&mut self) -> Result<bool, GatewayError> {
loop {
let page_started_at = std::time::Instant::now();
let page = match self.page_cursor.next_page().await {
Ok(Some(page)) => page,
Ok(None) => return false,
Ok(None) => {
observe_gateway_stage_ms(
"candidate_page_load",
page_started_at.elapsed().as_millis() as u64,
);
return Ok(false);
}
Err(error) => {
observe_gateway_stage_ms(
"candidate_page_load",
page_started_at.elapsed().as_millis() as u64,
);
if matches!(error, GatewayError::AdmissionTimeout { .. }) {
return Err(error);
}
warn!(
trace_id = %self.trace_id,
error = ?error,
"gateway lazy requested-model candidate page read failed"
);
return false;
return Ok(false);
}
};
observe_gateway_stage_ms(
"candidate_page_load",
page_started_at.elapsed().as_millis() as u64,
);
if page_is_exact_auth_api_key_concurrency_limited(&page) {
if self.wait_for_auth_api_key_concurrency_retry().await {
@@ -832,24 +872,16 @@ impl<'a> RequestedModelAttemptPageCursor<'a> {
}
self.persist_final_auth_api_key_concurrency_skips(page.skipped_candidates)
.await;
return false;
return Ok(false);
}
let resolve_started_at = std::time::Instant::now();
let (candidates, resolved_skipped) =
resolve_and_rank_logical_local_execution_candidates(
self.state,
page.candidates,
&self.client_api_format,
Some(&self.requested_model),
Some(&self.auth_snapshot),
self.client_session_affinity.as_ref(),
self.required_capabilities.as_ref(),
self.routing_policy.as_ref(),
self.sticky_session_token.as_deref(),
self.request_auth_channel.as_deref(),
self.resolution_mode,
)
.await;
resolve_priority_candidate_page_with_cache(self, page.candidates).await;
observe_gateway_stage_ms(
"candidate_page_resolve",
resolve_started_at.elapsed().as_millis() as u64,
);
let skipped_candidates = page
.skipped_candidates
.into_iter()
@@ -899,7 +931,7 @@ impl<'a> RequestedModelAttemptPageCursor<'a> {
.saturating_add(u32::try_from(skipped_candidate_count).unwrap_or(u32::MAX));
if !items.is_empty() {
self.pending_items = items;
return true;
return Ok(true);
}
let skipped_starting_candidate_index = next_candidate_index;
let skipped_persistence = LocalSkippedCandidatePersistenceContext {
@@ -1080,6 +1112,120 @@ pub(crate) fn remember_first_local_candidate_affinity(
);
}
async fn resolve_priority_candidate_page_with_cache(
cursor: &RequestedModelAttemptPageCursor<'_>,
page_candidates: Vec<SchedulerMinimalCandidateSelectionCandidate>,
) -> (
Vec<EligibleLocalExecutionCandidate>,
Vec<SkippedLocalExecutionCandidate>,
) {
if !should_cache_resolved_candidate_page(cursor) {
return resolve_and_rank_logical_local_execution_candidates(
cursor.state,
page_candidates,
&cursor.client_api_format,
Some(&cursor.requested_model),
Some(&cursor.auth_snapshot),
cursor.client_session_affinity.as_ref(),
cursor.required_capabilities.as_ref(),
cursor.routing_policy.as_ref(),
cursor.sticky_session_token.as_deref(),
cursor.request_auth_channel.as_deref(),
cursor.resolution_mode,
)
.await;
}
let key = CandidateResolvedPageCacheKey::new(
&cursor.requested_model,
&cursor.client_api_format,
true,
&cursor.auth_snapshot,
cursor.required_capabilities.as_ref(),
cursor.routing_policy.as_ref(),
cursor.request_auth_channel.as_deref(),
cursor.state.app().scheduler_affinity_epoch(),
cursor.page_cursor.resolved_page_cache_preselection_mode(),
cursor
.page_cursor
.resolved_page_cache_use_api_format_alias_match(),
cursor.client_session_affinity.as_ref(),
cursor.resolution_mode,
);
let page_candidates_for_fallback = page_candidates.clone();
let page_candidates_for_load = page_candidates;
let cache = cursor.state.app().candidate_resolved_page_cache.clone();
let ttl = candidate_page_cache_ttl_from_env();
let stale_ttl = candidate_page_cache_stale_ttl(ttl);
let cached = cache
.get_or_load_once_stale_while_refreshing(
key,
ttl,
stale_ttl,
|| async move {
let (candidates, resolved_skipped) =
resolve_and_rank_logical_local_execution_candidates(
cursor.state,
page_candidates_for_load,
&cursor.client_api_format,
Some(&cursor.requested_model),
Some(&cursor.auth_snapshot),
cursor.client_session_affinity.as_ref(),
cursor.required_capabilities.as_ref(),
cursor.routing_policy.as_ref(),
cursor.sticky_session_token.as_deref(),
cursor.request_auth_channel.as_deref(),
cursor.resolution_mode,
)
.await;
Ok::<_, GatewayError>(Some(Arc::new(CandidateResolvedPageSnapshot {
candidates,
resolved_skipped,
})))
},
CacheLoadObserver::new()
.on_hit(record_candidate_page_resolve_cache_hit)
.on_miss(record_candidate_page_resolve_cache_miss)
.on_load(record_candidate_page_resolve_cache_load)
.on_follower_wait(record_candidate_page_resolve_cache_follower_wait),
)
.await
.unwrap_or(None);
match cached {
Some(snapshot) => (
snapshot.candidates.clone(),
snapshot.resolved_skipped.clone(),
),
None => {
if page_candidates_for_fallback.is_empty() {
return (Vec::new(), Vec::new());
}
resolve_and_rank_logical_local_execution_candidates(
cursor.state,
page_candidates_for_fallback,
&cursor.client_api_format,
Some(&cursor.requested_model),
Some(&cursor.auth_snapshot),
cursor.client_session_affinity.as_ref(),
cursor.required_capabilities.as_ref(),
cursor.routing_policy.as_ref(),
cursor.sticky_session_token.as_deref(),
cursor.request_auth_channel.as_deref(),
cursor.resolution_mode,
)
.await
}
}
}
fn should_cache_resolved_candidate_page(cursor: &RequestedModelAttemptPageCursor<'_>) -> bool {
cursor.sticky_session_token.is_none()
&& cursor
.page_cursor
.should_cache_current_priority_resolved_page()
}
fn should_persist_available_local_candidate(eligible: &EligibleLocalExecutionCandidate) -> bool {
ai_should_persist_available_candidate_for_pool_key(eligible.orchestration.pool_key_index)
}
@@ -1937,7 +2083,8 @@ mod tests {
GatewayDataState::with_request_candidate_repository_for_tests(Arc::clone(
&repository,
)),
);
)
.without_request_candidate_queue_for_tests();
let attempts = persist_available_local_execution_candidates(
PlannerAppState::new(&app),
@@ -2025,6 +2172,145 @@ mod tests {
.is_none());
}
#[tokio::test]
async fn resolved_candidate_page_cache_requires_fixed_order_or_explicit_affinity() {
let app = AppState::new().expect("state should build");
let auth_snapshot = sample_auth_snapshot();
let mut page_cursor = LocalCandidatePreselectionPageCursor::new(
PlannerAppState::new(&app),
"openai:chat",
"gpt-5",
true,
None,
&auth_snapshot,
None,
None,
None,
false,
LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModel,
Some("trace-no-session-affinity"),
)
.await;
page_cursor.mark_priority_page_emitted_for_tests();
let cursor = RequestedModelAttemptPageCursor {
state: PlannerAppState::new(&app),
trace_id: "trace-no-session-affinity".to_string(),
client_api_format: "openai:chat".to_string(),
requested_model: "gpt-5".to_string(),
auth_snapshot: auth_snapshot.clone(),
client_session_affinity: None,
required_capabilities: None,
routing_policy: None,
sticky_session_token: None,
request_auth_channel: None,
skipped_user_id: "user-1".to_string(),
skipped_api_key_id: "api-key-1".to_string(),
skipped_required_capabilities: None,
skipped_error_context: "test skipped",
record_runtime_miss_diagnostic: false,
resolution_mode: LocalCandidateResolutionMode::Standard,
decorate_skipped_candidate: Arc::new(identity_skipped_candidate),
page_cursor,
pending_items: VecDeque::new(),
candidate_count: 0,
next_candidate_index: 0,
remembered_affinity: false,
scheduler_cache_affinity_enabled: false,
auth_api_key_concurrency_wait_deadline: None,
deferred_error: None,
};
assert!(!should_cache_resolved_candidate_page(&cursor));
let sticky_cursor = RequestedModelAttemptPageCursor {
sticky_session_token: Some("sticky-token".to_string()),
..cursor
};
assert!(!should_cache_resolved_candidate_page(&sticky_cursor));
let mut page_cursor = LocalCandidatePreselectionPageCursor::new(
PlannerAppState::new(&app),
"openai:chat",
"gpt-5",
true,
None,
&auth_snapshot,
None,
Some(&ClientSessionAffinity::from_session_key("chat-session-1")),
None,
false,
LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModel,
Some("trace-session-affinity"),
)
.await;
page_cursor.mark_priority_page_emitted_for_tests();
let cursor = RequestedModelAttemptPageCursor {
client_session_affinity: Some(ClientSessionAffinity::from_session_key(
"chat-session-1",
)),
page_cursor,
sticky_session_token: None,
..sticky_cursor
};
assert!(should_cache_resolved_candidate_page(&cursor));
let fixed_order_app = AppState::new()
.expect("state should build")
.with_data_state_for_tests(
GatewayDataState::disabled().with_system_config_values_for_tests([(
"scheduling_mode".to_string(),
json!("fixed_order"),
)]),
);
let mut page_cursor = LocalCandidatePreselectionPageCursor::new(
PlannerAppState::new(&fixed_order_app),
"openai:chat",
"gpt-5",
true,
None,
&auth_snapshot,
None,
None,
None,
false,
LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModel,
Some("trace-fixed-order"),
)
.await;
page_cursor.mark_priority_page_emitted_for_tests();
let cursor = RequestedModelAttemptPageCursor {
state: PlannerAppState::new(&fixed_order_app),
trace_id: "trace-fixed-order".to_string(),
client_api_format: "openai:chat".to_string(),
requested_model: "gpt-5".to_string(),
auth_snapshot,
client_session_affinity: None,
required_capabilities: None,
routing_policy: None,
sticky_session_token: None,
request_auth_channel: None,
skipped_user_id: "user-1".to_string(),
skipped_api_key_id: "api-key-1".to_string(),
skipped_required_capabilities: None,
skipped_error_context: "test skipped",
record_runtime_miss_diagnostic: false,
resolution_mode: LocalCandidateResolutionMode::Standard,
decorate_skipped_candidate: Arc::new(identity_skipped_candidate),
page_cursor,
pending_items: VecDeque::new(),
candidate_count: 0,
next_candidate_index: 0,
remembered_affinity: false,
scheduler_cache_affinity_enabled: false,
auth_api_key_concurrency_wait_deadline: None,
deferred_error: None,
};
assert!(should_cache_resolved_candidate_page(&cursor));
}
#[tokio::test]
async fn logical_materialization_does_not_persist_pool_group_representative() {
let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default());
@@ -2042,7 +2328,8 @@ mod tests {
Arc::clone(&request_candidate_repository),
"test-encryption-key",
),
);
)
.without_request_candidate_queue_for_tests();
let mut pool_group = sample_eligible("pool-group", None);
pool_group.kind = LocalExecutionCandidateKind::PoolGroup;
pool_group.transport = sample_transport(
@@ -2140,7 +2427,8 @@ mod tests {
GatewayDataState::with_request_candidate_repository_for_tests(Arc::clone(
&repository,
)),
);
)
.without_request_candidate_queue_for_tests();
let mut eligible = sample_eligible("ranked-key", None);
eligible.ranking = Some(SchedulerRankingOutcome {
original_index: 1,
@@ -2214,12 +2502,17 @@ mod tests {
let first = source
.next_attempt()
.await
.expect("first attempt read should succeed")
.expect("first attempt should be available");
assert_eq!(first.eligible.candidate.key_id, "normal-key");
let remaining = source.drain_static_attempts();
assert!(remaining.is_empty());
assert!(source.next_attempt().await.is_none());
assert!(source
.next_attempt()
.await
.expect("remaining attempt read should succeed")
.is_none());
}
#[tokio::test]
@@ -2239,7 +2532,8 @@ mod tests {
Arc::clone(&request_candidate_repository),
"test-encryption-key",
),
);
)
.without_request_candidate_queue_for_tests();
let mut pool_group = sample_eligible("pool-group", None);
pool_group.kind = LocalExecutionCandidateKind::PoolGroup;
pool_group.transport = sample_transport(
@@ -2275,7 +2569,11 @@ mod tests {
}]),
};
assert!(source.next_attempt().await.is_none());
assert!(source
.next_attempt()
.await
.expect("pool attempt read should succeed")
.is_none());
let stored = app
.read_request_candidates_by_request_id("trace-dynamic-pool")
@@ -2308,7 +2606,8 @@ mod tests {
GatewayDataState::with_request_candidate_repository_for_tests(Arc::clone(
&repository,
)),
);
)
.without_request_candidate_queue_for_tests();
persist_skipped_local_execution_candidates(
&app,
@@ -1,5 +1,3 @@
use std::collections::BTreeMap;
use aether_ai_serving::{
ai_ranking_context, build_ai_rankable_candidate, run_ai_candidate_ranking,
AiCandidateRankingPort, AiRankableCandidateParts, AiRankingContextConfig,
@@ -7,6 +5,7 @@ use aether_ai_serving::{
};
use aether_routing_core::{ResolvedRoutingPolicy, RoutingSchedulingMode, RoutingSetPriorityMode};
use async_trait::async_trait;
use tokio::sync::Mutex;
use tracing::warn;
use crate::ai_serving::{GatewayAuthApiKeySnapshot, PlannerAppState};
@@ -24,7 +23,7 @@ use aether_scheduler_core::{
use super::candidate_affinity_cache::read_cached_scheduler_affinity_target;
use super::candidate_resolution::{EligibleLocalExecutionCandidate, LocalExecutionCandidateKind};
use super::candidate_transport_ranking_facts::{
resolve_cached_transport_ranking_facts, CandidateTransportRankingFacts,
resolve_cached_transport_ranking_facts, CandidateTransportRankingFactsCache,
};
struct GatewayLocalCandidateRankingPort<'a> {
@@ -35,6 +34,7 @@ struct GatewayLocalCandidateRankingPort<'a> {
required_capabilities: Option<&'a serde_json::Value>,
ordering_config: SchedulerOrderingConfig,
routing_policy: Option<&'a ResolvedRoutingPolicy>,
transport_ranking_facts_cache: Mutex<CandidateTransportRankingFactsCache>,
}
#[async_trait]
@@ -84,13 +84,17 @@ impl AiCandidateRankingPort for GatewayLocalCandidateRankingPort<'_> {
normalized_client_api_format: &str,
cached_affinity_match: bool,
) -> Result<SchedulerRankableCandidate, Self::Error> {
let ranking_facts = resolve_transport_ranking_facts_for_candidate(
self.state,
&candidate.candidate,
candidate.transport.as_ref(),
self.ordering_config,
)
.await;
let ranking_facts = {
let mut cache = self.transport_ranking_facts_cache.lock().await;
resolve_cached_transport_ranking_facts(
self.state,
&mut cache,
&candidate.candidate,
candidate.transport.as_ref(),
self.ordering_config,
)
.await
};
let routing_overlaid_candidate =
routing_overlaid_candidate(self.routing_policy, candidate.kind, &candidate.candidate);
Ok(build_ai_rankable_candidate(AiRankableCandidateParts {
@@ -137,6 +141,7 @@ pub(crate) async fn rank_eligible_local_execution_candidates(
required_capabilities,
ordering_config,
routing_policy,
transport_ranking_facts_cache: Mutex::new(CandidateTransportRankingFactsCache::default()),
};
match run_ai_candidate_ranking(&port, candidates, normalized_client_api_format).await {
@@ -145,23 +150,6 @@ pub(crate) async fn rank_eligible_local_execution_candidates(
}
}
async fn resolve_transport_ranking_facts_for_candidate(
state: PlannerAppState<'_>,
candidate: &SchedulerMinimalCandidateSelectionCandidate,
transport: &crate::ai_serving::GatewayProviderTransportSnapshot,
ordering_config: SchedulerOrderingConfig,
) -> CandidateTransportRankingFacts {
let mut ordering_cache = BTreeMap::new();
resolve_cached_transport_ranking_facts(
state,
&mut ordering_cache,
candidate,
transport,
ordering_config,
)
.await
}
fn cached_affinity_matches_local_execution_scope(
eligible: &EligibleLocalExecutionCandidate,
target: &SchedulerAffinityTarget,
@@ -284,7 +272,9 @@ mod tests {
use serde_json::json;
use super::super::candidate_affinity_cache::remember_scheduler_affinity_for_candidate;
use super::super::candidate_transport_ranking_facts::resolve_cached_candidate_transport_ranking_facts;
use super::super::candidate_transport_ranking_facts::{
resolve_cached_candidate_transport_ranking_facts, CandidateTransportRankingFactsCache,
};
use super::{PlannerAppState, SchedulerMinimalCandidateSelectionCandidate};
use crate::ai_serving::planner::candidate_resolution::{
resolve_and_rank_local_execution_candidates,
@@ -306,7 +296,7 @@ mod tests {
let ordering_config = super::read_scheduler_ordering_config_or_default(state).await;
let mut candidates = candidates;
let mut rankables = Vec::with_capacity(candidates.len());
let mut ordering_cache = BTreeMap::new();
let mut ordering_cache = CandidateTransportRankingFactsCache::default();
for (original_index, candidate) in candidates.iter().enumerate() {
let ranking_facts = resolve_cached_candidate_transport_ranking_facts(
@@ -1540,6 +1530,7 @@ mod tests {
.expect("state should build")
.with_data_state_for_tests(data_state);
let auth_snapshot = sample_auth_snapshot();
let client_session_affinity = ClientSessionAffinity::from_session_key("session-1");
let cached_candidate = sample_priority_candidate(
"provider-cached",
"endpoint-cached",
@@ -1551,7 +1542,7 @@ mod tests {
remember_scheduler_affinity_for_candidate(
PlannerAppState::new(&state),
Some(&auth_snapshot),
None,
Some(&client_session_affinity),
"openai:chat",
"gpt-4.1",
&cached_candidate,
@@ -1573,7 +1564,7 @@ mod tests {
"openai:chat",
"gpt-4.1",
Some(&auth_snapshot),
None,
Some(&client_session_affinity),
None,
None,
None,
@@ -1714,6 +1705,7 @@ mod tests {
.expect("state should build")
.with_data_state_for_tests(data_state);
let auth_snapshot = sample_auth_snapshot();
let client_session_affinity = ClientSessionAffinity::from_session_key("session-1");
let cached_cross_format = sample_priority_candidate(
"provider-shared",
"endpoint-openai",
@@ -1725,7 +1717,7 @@ mod tests {
remember_scheduler_affinity_for_candidate(
PlannerAppState::new(&state),
Some(&auth_snapshot),
None,
Some(&client_session_affinity),
"claude:messages",
"gpt-4.1",
&cached_cross_format,
@@ -1747,7 +1739,7 @@ mod tests {
"claude:messages",
"gpt-4.1",
Some(&auth_snapshot),
None,
Some(&client_session_affinity),
None,
None,
None,
@@ -1915,6 +1907,7 @@ mod tests {
.expect("state should build")
.with_data_state_for_tests(data_state);
let auth_snapshot = sample_auth_snapshot();
let client_session_affinity = ClientSessionAffinity::from_session_key("session-1");
let cached_candidate = sample_priority_candidate(
"provider-pool",
"endpoint-pool",
@@ -1926,7 +1919,7 @@ mod tests {
remember_scheduler_affinity_for_candidate(
PlannerAppState::new(&state),
Some(&auth_snapshot),
None,
Some(&client_session_affinity),
"openai:chat",
"gpt-4.1",
&cached_candidate,
@@ -1948,7 +1941,7 @@ mod tests {
"openai:chat",
Some("gpt-4.1"),
Some(&auth_snapshot),
None,
Some(&client_session_affinity),
None,
None,
None,
@@ -2008,6 +2001,7 @@ mod tests {
.expect("state should build")
.with_data_state_for_tests(data_state);
let auth_snapshot = sample_auth_snapshot();
let client_session_affinity = ClientSessionAffinity::from_session_key("session-1");
let cached_candidate = sample_priority_candidate(
"provider-pool",
"endpoint-pool",
@@ -2019,7 +2013,7 @@ mod tests {
remember_scheduler_affinity_for_candidate(
PlannerAppState::new(&state),
Some(&auth_snapshot),
None,
Some(&client_session_affinity),
"openai:chat",
"gpt-4.1",
&cached_candidate,
@@ -2041,7 +2035,7 @@ mod tests {
"openai:chat",
Some("gpt-4.1"),
Some(&auth_snapshot),
None,
Some(&client_session_affinity),
None,
None,
None,
@@ -2064,7 +2058,7 @@ mod tests {
}
#[tokio::test]
async fn remembers_scheduler_affinity_for_candidate_using_requested_model_key() {
async fn ignores_scheduler_affinity_without_client_session_scope() {
let state = AppState::new().expect("state should build");
let auth_snapshot = sample_auth_snapshot();
let candidate = sample_candidate("endpoint-1", "key-1");
@@ -2078,15 +2072,12 @@ mod tests {
&candidate,
);
let remembered = state
assert!(state
.read_scheduler_affinity_target(
"scheduler_affinity:api-key-1:openai:chat:gpt-5",
SCHEDULER_AFFINITY_TTL,
)
.expect("affinity target should be cached");
assert_eq!(remembered.provider_id, "provider-1");
assert_eq!(remembered.endpoint_id, "endpoint-1");
assert_eq!(remembered.key_id, "key-1");
.is_none());
}
#[tokio::test]
@@ -7,6 +7,7 @@ use aether_ai_serving::{
use aether_routing_core::ResolvedRoutingPolicy;
use async_trait::async_trait;
use std::convert::Infallible;
use std::time::Instant;
use tracing::warn;
use aether_scheduler_core::{
@@ -20,6 +21,7 @@ use crate::ai_serving::{
PlannerAppState,
};
use crate::orchestration::LocalExecutionCandidateMetadata;
use crate::stage_metrics::observe_gateway_stage_ms;
use super::candidate_ranking::rank_eligible_local_execution_candidates;
@@ -68,7 +70,7 @@ struct GatewayLocalCandidateResolutionPort<'a> {
#[async_trait]
impl AiCandidateResolutionPort for GatewayLocalCandidateResolutionPort<'_> {
type Candidate = SchedulerMinimalCandidateSelectionCandidate;
type Transport = GatewayProviderTransportSnapshot;
type Transport = Arc<GatewayProviderTransportSnapshot>;
type Eligible = EligibleLocalExecutionCandidate;
type Skipped = SkippedLocalExecutionCandidate;
type Error = Infallible;
@@ -77,7 +79,12 @@ impl AiCandidateResolutionPort for GatewayLocalCandidateResolutionPort<'_> {
&self,
candidate: &Self::Candidate,
) -> Result<Option<Self::Transport>, Self::Error> {
Ok(read_candidate_transport_snapshot(self.state, candidate).await)
let started_at = Instant::now();
let transport = read_candidate_transport_snapshot_arc(self.state, candidate).await;
let elapsed_ms = started_at.elapsed().as_millis() as u64;
observe_gateway_stage_ms("candidate_transport_snapshot", elapsed_ms);
observe_gateway_stage_ms("candidate_resolution_transport_read", elapsed_ms);
Ok(transport)
}
fn build_missing_transport_skipped_candidate(
@@ -139,7 +146,7 @@ impl AiCandidateResolutionPort for GatewayLocalCandidateResolutionPort<'_> {
SkippedLocalExecutionCandidate {
candidate,
skip_reason,
transport: Some(Arc::new(transport)),
transport: Some(transport),
ranking: None,
extra_data: None,
}
@@ -159,7 +166,7 @@ impl AiCandidateResolutionPort for GatewayLocalCandidateResolutionPort<'_> {
EligibleLocalExecutionCandidate {
kind,
candidate,
transport: Arc::new(transport),
transport,
provider_api_format,
orchestration: LocalExecutionCandidateMetadata::default(),
ranking: None,
@@ -171,7 +178,8 @@ impl AiCandidateResolutionPort for GatewayLocalCandidateResolutionPort<'_> {
candidates: Vec<Self::Eligible>,
normalized_client_api_format: &str,
) -> Result<Vec<Self::Eligible>, Self::Error> {
Ok(rank_eligible_local_execution_candidates(
let started_at = Instant::now();
let ranked = rank_eligible_local_execution_candidates(
self.state,
candidates,
normalized_client_api_format,
@@ -181,7 +189,12 @@ impl AiCandidateResolutionPort for GatewayLocalCandidateResolutionPort<'_> {
self.required_capabilities,
self.routing_policy,
)
.await)
.await;
observe_gateway_stage_ms(
"candidate_resolution_rank",
started_at.elapsed().as_millis() as u64,
);
Ok(ranked)
}
async fn apply_pool_scheduler(
@@ -358,8 +371,13 @@ async fn resolve_and_rank_local_execution_candidates_with_pool_expansion(
expand_pool_groups,
};
let started_at = Instant::now();
match run_ai_candidate_resolution(&port, candidates, request).await {
Ok(mut outcome) => {
observe_gateway_stage_ms(
"candidate_resolution_core",
started_at.elapsed().as_millis() as u64,
);
for candidate in &mut outcome.eligible_candidates {
candidate.orchestration.scheduler_affinity_epoch = Some(scheduler_affinity_epoch);
}
@@ -506,8 +524,17 @@ pub(crate) async fn read_candidate_transport_snapshot(
state: PlannerAppState<'_>,
candidate: &SchedulerMinimalCandidateSelectionCandidate,
) -> Option<GatewayProviderTransportSnapshot> {
read_candidate_transport_snapshot_arc(state, candidate)
.await
.map(|transport| (*transport).clone())
}
pub(crate) async fn read_candidate_transport_snapshot_arc(
state: PlannerAppState<'_>,
candidate: &SchedulerMinimalCandidateSelectionCandidate,
) -> Option<Arc<GatewayProviderTransportSnapshot>> {
match state
.read_provider_transport_snapshot(
.read_provider_transport_snapshot_arc(
&candidate.provider_id,
&candidate.endpoint_id,
&candidate.key_id,
@@ -3,6 +3,7 @@ use aether_ai_serving::{
};
use aether_data_contracts::repository::candidate_selection::StoredMinimalCandidateSelectionRow;
use aether_routing_core::ResolvedRoutingPolicy;
use aether_runtime::ConcurrencyPermit;
use aether_scheduler_core::{
enumerate_minimal_candidate_selection_with_model_directives, normalize_api_format,
resolve_requested_global_model_name_with_model_directives,
@@ -12,16 +13,17 @@ use aether_scheduler_core::{
use async_trait::async_trait;
use std::collections::{BTreeMap, BTreeSet, VecDeque};
use crate::ai_serving::planner::candidate_affinity_cache::has_explicit_session_affinity;
use crate::ai_serving::planner::candidate_resolution::SkippedLocalExecutionCandidate;
use crate::ai_serving::{GatewayAuthApiKeySnapshot, PlannerAppState};
use crate::clock::current_unix_secs;
use crate::clock::request_distribution_seed;
use crate::data::candidate_selection::{
read_requested_model_rows_fast_path_page, requested_model_candidate_names,
MinimalCandidateSelectionRowSource, REQUESTED_MODEL_CANDIDATE_PAGE_SIZE,
REQUESTED_MODEL_MAX_SCANNED_ROWS,
};
use crate::scheduler::candidate::SchedulerSkippedCandidate;
use crate::scheduler::config::SchedulerOrderingConfig;
use crate::scheduler::config::{SchedulerOrderingConfig, SchedulerSchedulingMode};
use crate::GatewayError;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
@@ -30,6 +32,15 @@ pub(crate) enum LocalCandidatePreselectionKeyMode {
ProviderEndpointKeyModelAndApiFormat,
}
impl LocalCandidatePreselectionKeyMode {
pub(crate) fn cache_key_name(self) -> &'static str {
match self {
Self::ProviderEndpointKeyModel => "provider_endpoint_key_model",
Self::ProviderEndpointKeyModelAndApiFormat => "provider_endpoint_key_model_api_format",
}
}
}
struct GatewayLocalCandidatePreselectionPort<'a> {
state: PlannerAppState<'a>,
client_api_format: &'a str,
@@ -43,6 +54,7 @@ struct GatewayLocalCandidatePreselectionPort<'a> {
key_mode: LocalCandidatePreselectionKeyMode,
candidate_api_formats: Vec<String>,
model_directive_enabled_api_formats: BTreeSet<String>,
ranking_seed: u64,
}
#[async_trait]
@@ -81,7 +93,7 @@ impl AiCandidatePreselectionPort for GatewayLocalCandidatePreselectionPort<'_> {
self.required_capabilities,
auth_snapshot,
self.client_session_affinity,
current_unix_secs(),
self.ranking_seed,
)
.await?;
@@ -227,6 +239,7 @@ pub(crate) async fn preselect_local_execution_candidates_for_api_formats_with_se
key_mode,
candidate_api_formats,
model_directive_enabled_api_formats,
ranking_seed: request_distribution_seed(),
};
run_ai_candidate_preselection(&port).await
@@ -234,6 +247,7 @@ pub(crate) async fn preselect_local_execution_candidates_for_api_formats_with_se
pub(crate) struct LocalCandidatePreselectionPageCursor<'a> {
state: PlannerAppState<'a>,
trace_id: String,
client_api_format: String,
requested_model: String,
require_streaming: bool,
@@ -241,11 +255,13 @@ pub(crate) struct LocalCandidatePreselectionPageCursor<'a> {
auth_snapshot: GatewayAuthApiKeySnapshot,
routing_policy: Option<ResolvedRoutingPolicy>,
client_session_affinity: Option<ClientSessionAffinity>,
request_auth_channel: Option<String>,
use_api_format_alias_match: bool,
key_mode: LocalCandidatePreselectionKeyMode,
candidate_api_formats: Vec<String>,
model_directive_enabled_api_formats: BTreeSet<String>,
ordering_config: SchedulerOrderingConfig,
ranking_seed: u64,
priority_page_emitted: bool,
deferred_pages_by_format: BTreeMap<
String,
@@ -276,8 +292,10 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
auth_snapshot: &GatewayAuthApiKeySnapshot,
routing_policy: Option<&ResolvedRoutingPolicy>,
client_session_affinity: Option<&ClientSessionAffinity>,
request_auth_channel: Option<&str>,
use_api_format_alias_match: bool,
key_mode: LocalCandidatePreselectionKeyMode,
trace_id: Option<&str>,
) -> Self {
let candidate_api_formats =
crate::ai_serving::request_candidate_api_formats(client_api_format, require_streaming)
@@ -307,6 +325,7 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
Self {
state,
trace_id: trace_id.unwrap_or_default().to_string(),
client_api_format: client_api_format.to_string(),
requested_model: requested_model.to_string(),
require_streaming,
@@ -314,11 +333,13 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
auth_snapshot: auth_snapshot.clone(),
routing_policy: routing_policy.cloned(),
client_session_affinity: client_session_affinity.cloned(),
request_auth_channel: request_auth_channel.map(str::to_string),
use_api_format_alias_match,
key_mode,
candidate_api_formats,
model_directive_enabled_api_formats,
ordering_config,
ranking_seed: request_distribution_seed(),
priority_page_emitted: false,
deferred_pages_by_format: BTreeMap::new(),
format_index: 0,
@@ -344,7 +365,7 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
> {
if !self.priority_page_emitted {
self.priority_page_emitted = true;
let priority_page = self.next_priority_page().await?;
let priority_page = self.cached_next_priority_page().await?;
if !priority_page.candidates.is_empty() || !priority_page.skipped_candidates.is_empty()
{
return Ok(Some(priority_page));
@@ -356,7 +377,10 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
if let Some(outcome) = self.pop_deferred_page(&candidate_api_format) {
return Ok(Some(outcome));
}
let Some(outcome) = self.next_page_for_api_format(&candidate_api_format).await? else {
let Some(outcome) = self
.next_page_for_api_format_with_planning_gate(&candidate_api_format)
.await?
else {
self.format_index += 1;
continue;
};
@@ -380,6 +404,70 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
self.deferred_pages_by_format.clear();
}
pub(crate) fn resolved_page_cache_preselection_mode(&self) -> &'static str {
self.key_mode.cache_key_name()
}
pub(crate) fn resolved_page_cache_use_api_format_alias_match(&self) -> bool {
self.use_api_format_alias_match
}
pub(crate) fn should_cache_current_priority_resolved_page(&self) -> bool {
if !(self.priority_page_emitted
&& self.format_index == 0
&& self.deferred_pages_by_format.is_empty())
{
return false;
}
match self.ordering_config.scheduling_mode {
SchedulerSchedulingMode::FixedOrder => true,
SchedulerSchedulingMode::CacheAffinity => {
has_explicit_session_affinity(self.client_session_affinity.as_ref())
}
SchedulerSchedulingMode::LoadBalance => false,
}
}
#[cfg(test)]
pub(crate) fn mark_priority_page_emitted_for_tests(&mut self) {
self.priority_page_emitted = true;
}
async fn cached_next_priority_page(
&mut self,
) -> Result<
AiCandidatePreselectionOutcome<
SchedulerMinimalCandidateSelectionCandidate,
SkippedLocalExecutionCandidate,
>,
GatewayError,
> {
let page = self.next_priority_page_with_planning_gate().await?;
self.remember_seen_candidates_from_page(&page);
Ok(page)
}
fn remember_seen_candidates_from_page(
&mut self,
page: &AiCandidatePreselectionOutcome<
SchedulerMinimalCandidateSelectionCandidate,
SkippedLocalExecutionCandidate,
>,
) {
for candidate in &page.candidates {
self.seen_candidate_keys
.insert(local_candidate_preselection_key(candidate, self.key_mode));
}
for skipped_candidate in &page.skipped_candidates {
self.seen_candidate_keys
.insert(local_candidate_preselection_key(
&skipped_candidate.candidate,
self.key_mode,
));
}
}
async fn next_priority_page(
&mut self,
) -> Result<
@@ -423,6 +511,35 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
Ok(priority_page)
}
async fn next_priority_page_with_planning_gate(
&mut self,
) -> Result<
AiCandidatePreselectionOutcome<
SchedulerMinimalCandidateSelectionCandidate,
SkippedLocalExecutionCandidate,
>,
GatewayError,
> {
let _permit = acquire_candidate_planning_gate(self.state, &self.trace_id).await?;
self.next_priority_page().await
}
async fn next_page_for_api_format_with_planning_gate(
&mut self,
candidate_api_format: &str,
) -> Result<
Option<
AiCandidatePreselectionOutcome<
SchedulerMinimalCandidateSelectionCandidate,
SkippedLocalExecutionCandidate,
>,
>,
GatewayError,
> {
let _permit = acquire_candidate_planning_gate(self.state, &self.trace_id).await?;
self.next_page_for_api_format(candidate_api_format).await
}
async fn split_priority_conversion_page(
&self,
candidate_api_format: &str,
@@ -797,7 +914,7 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
self.required_capabilities.as_ref(),
auth_snapshot,
self.client_session_affinity.as_ref(),
current_unix_secs(),
self.ranking_seed,
)
.await?;
let skipped_candidates = skipped_candidates
@@ -894,6 +1011,28 @@ fn local_candidate_preselection_key(
}
}
async fn acquire_candidate_planning_gate(
state: PlannerAppState<'_>,
trace_id: &str,
) -> Result<Option<ConcurrencyPermit>, GatewayError> {
let Some(gate) = state.app().candidate_planning_gate.as_ref() else {
return Ok(None);
};
let budget = state
.app()
.frontdoor_runtime_guards
.internal_gate_queue_budget;
match tokio::time::timeout(budget, gate.acquire()).await {
Ok(Ok(permit)) => Ok(Some(permit)),
Ok(Err(err)) => Err(GatewayError::Internal(err.to_string())),
Err(_) => Err(GatewayError::AdmissionTimeout {
trace_id: trace_id.to_string(),
gate: "gateway_candidate_planning",
queue_budget_ms: budget.as_millis() as u64,
}),
}
}
fn matches_client_api_format(
use_api_format_alias_match: bool,
candidate_api_format: &str,
@@ -1220,8 +1359,10 @@ mod tests {
&auth_snapshot,
None,
None,
None,
true,
LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModelAndApiFormat,
None,
)
.await;
@@ -1277,8 +1418,10 @@ mod tests {
&auth_snapshot,
None,
None,
None,
true,
LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModelAndApiFormat,
None,
)
.await;
@@ -1349,8 +1492,10 @@ mod tests {
&auth_snapshot,
None,
None,
None,
true,
LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModelAndApiFormat,
None,
)
.await;
@@ -1,8 +1,10 @@
use std::collections::BTreeMap;
use aether_contracts::ProxySnapshot;
use aether_scheduler_core::{
SchedulerMinimalCandidateSelectionCandidate, SchedulerTunnelAffinityBucket,
};
use serde_json::Value;
use tracing::warn;
use crate::ai_serving::{GatewayProviderTransportSnapshot, PlannerAppState};
@@ -10,7 +12,9 @@ use crate::scheduler::config::SchedulerOrderingConfig;
use super::candidate_resolution::read_candidate_transport_snapshot;
pub(super) type CandidateTransportIdentity<'a> = (&'a str, &'a str, &'a str);
const TUNNEL_OWNER_INSTANCE_ID_EXTRA_KEY: &str = "tunnel_owner_instance_id";
pub(super) type CandidateTransportIdentity = (String, String, String);
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(super) struct CandidateTransportRankingFacts {
@@ -18,38 +22,51 @@ pub(super) struct CandidateTransportRankingFacts {
pub(super) keep_priority_on_conversion: bool,
}
pub(super) async fn resolve_cached_candidate_transport_ranking_facts<'a>(
state: PlannerAppState<'_>,
cache: &mut BTreeMap<CandidateTransportIdentity<'a>, CandidateTransportRankingFacts>,
candidate: &'a SchedulerMinimalCandidateSelectionCandidate,
ordering_config: SchedulerOrderingConfig,
) -> CandidateTransportRankingFacts {
let identity = candidate_transport_identity(candidate);
if let Some(facts) = cache.get(&identity).copied() {
return facts;
}
let facts = resolve_candidate_transport_ranking_facts(state, candidate, ordering_config).await;
cache.insert(identity, facts);
facts
#[derive(Debug, Default)]
pub(super) struct CandidateTransportRankingFactsCache {
candidate_facts: BTreeMap<CandidateTransportIdentity, CandidateTransportRankingFacts>,
configured_proxy_snapshots: BTreeMap<String, Option<ProxySnapshot>>,
system_proxy_snapshot: Option<Option<ProxySnapshot>>,
tunnel_buckets_by_node_id: BTreeMap<String, SchedulerTunnelAffinityBucket>,
}
pub(super) async fn resolve_cached_transport_ranking_facts<'a>(
pub(super) async fn resolve_cached_candidate_transport_ranking_facts(
state: PlannerAppState<'_>,
cache: &mut BTreeMap<CandidateTransportIdentity<'a>, CandidateTransportRankingFacts>,
candidate: &'a SchedulerMinimalCandidateSelectionCandidate,
transport: &GatewayProviderTransportSnapshot,
cache: &mut CandidateTransportRankingFactsCache,
candidate: &SchedulerMinimalCandidateSelectionCandidate,
ordering_config: SchedulerOrderingConfig,
) -> CandidateTransportRankingFacts {
let identity = candidate_transport_identity(candidate);
if let Some(facts) = cache.get(&identity).copied() {
if let Some(facts) = cache.candidate_facts.get(&identity).copied() {
return facts;
}
let facts =
resolve_candidate_transport_ranking_facts_from_transport(state, transport, ordering_config)
.await;
cache.insert(identity, facts);
resolve_candidate_transport_ranking_facts(state, cache, candidate, ordering_config).await;
cache.candidate_facts.insert(identity, facts);
facts
}
pub(super) async fn resolve_cached_transport_ranking_facts(
state: PlannerAppState<'_>,
cache: &mut CandidateTransportRankingFactsCache,
candidate: &SchedulerMinimalCandidateSelectionCandidate,
transport: &GatewayProviderTransportSnapshot,
ordering_config: SchedulerOrderingConfig,
) -> CandidateTransportRankingFacts {
let identity = candidate_transport_identity(candidate);
if let Some(facts) = cache.candidate_facts.get(&identity).copied() {
return facts;
}
let facts = resolve_candidate_transport_ranking_facts_from_transport(
state,
cache,
transport,
ordering_config,
)
.await;
cache.candidate_facts.insert(identity, facts);
facts
}
@@ -58,13 +75,15 @@ pub(super) async fn candidate_keeps_priority_on_conversion(
candidate: &SchedulerMinimalCandidateSelectionCandidate,
ordering_config: SchedulerOrderingConfig,
) -> bool {
resolve_candidate_transport_ranking_facts(state, candidate, ordering_config)
let mut cache = CandidateTransportRankingFactsCache::default();
resolve_candidate_transport_ranking_facts(state, &mut cache, candidate, ordering_config)
.await
.keep_priority_on_conversion
}
async fn resolve_candidate_transport_ranking_facts(
state: PlannerAppState<'_>,
cache: &mut CandidateTransportRankingFactsCache,
candidate: &SchedulerMinimalCandidateSelectionCandidate,
ordering_config: SchedulerOrderingConfig,
) -> CandidateTransportRankingFacts {
@@ -75,17 +94,23 @@ async fn resolve_candidate_transport_ranking_facts(
};
};
resolve_candidate_transport_ranking_facts_from_transport(state, &transport, ordering_config)
.await
resolve_candidate_transport_ranking_facts_from_transport(
state,
cache,
&transport,
ordering_config,
)
.await
}
async fn resolve_candidate_transport_ranking_facts_from_transport(
state: PlannerAppState<'_>,
cache: &mut CandidateTransportRankingFactsCache,
transport: &GatewayProviderTransportSnapshot,
ordering_config: SchedulerOrderingConfig,
) -> CandidateTransportRankingFacts {
CandidateTransportRankingFacts {
tunnel_bucket: resolve_tunnel_owner_affinity_from_transport(state, transport).await,
tunnel_bucket: resolve_tunnel_owner_affinity_from_transport(state, cache, transport).await,
keep_priority_on_conversion: ordering_config.keep_priority_on_conversion
|| transport.provider.keep_priority_on_conversion,
}
@@ -93,12 +118,11 @@ async fn resolve_candidate_transport_ranking_facts_from_transport(
async fn resolve_tunnel_owner_affinity_from_transport(
state: PlannerAppState<'_>,
cache: &mut CandidateTransportRankingFactsCache,
transport: &GatewayProviderTransportSnapshot,
) -> SchedulerTunnelAffinityBucket {
let Some(proxy) = state
.app()
.resolve_transport_proxy_snapshot_with_tunnel_affinity(transport)
.await
let Some(proxy) =
resolve_transport_proxy_snapshot_with_tunnel_affinity_cached(state, cache, transport).await
else {
return SchedulerTunnelAffinityBucket::Neutral;
};
@@ -114,10 +138,74 @@ async fn resolve_tunnel_owner_affinity_from_transport(
return SchedulerTunnelAffinityBucket::Neutral;
};
if let Some(bucket) = cache.tunnel_buckets_by_node_id.get(node_id).copied() {
return bucket;
}
let bucket = resolve_tunnel_owner_affinity_from_proxy(state, &proxy, node_id).await;
cache
.tunnel_buckets_by_node_id
.insert(node_id.to_string(), bucket);
bucket
}
async fn resolve_transport_proxy_snapshot_with_tunnel_affinity_cached(
state: PlannerAppState<'_>,
cache: &mut CandidateTransportRankingFactsCache,
transport: &GatewayProviderTransportSnapshot,
) -> Option<ProxySnapshot> {
for raw in [
transport.key.proxy.as_ref(),
transport.endpoint.proxy.as_ref(),
transport.provider.proxy.as_ref(),
]
.into_iter()
.flatten()
{
let cache_key = proxy_config_cache_key(raw);
if let Some(snapshot) = cache.configured_proxy_snapshots.get(&cache_key) {
if snapshot.is_some() {
return snapshot.clone();
}
continue;
}
let snapshot = state
.app()
.resolve_configured_proxy_snapshot_with_tunnel_affinity(Some(raw))
.await;
cache
.configured_proxy_snapshots
.insert(cache_key, snapshot.clone());
if snapshot.is_some() {
return snapshot;
}
}
if let Some(snapshot) = cache.system_proxy_snapshot.as_ref() {
return snapshot.clone();
}
let snapshot = state.app().resolve_system_proxy_snapshot().await;
cache.system_proxy_snapshot = Some(snapshot.clone());
snapshot
}
async fn resolve_tunnel_owner_affinity_from_proxy(
state: PlannerAppState<'_>,
proxy: &ProxySnapshot,
node_id: &str,
) -> SchedulerTunnelAffinityBucket {
if state.app().tunnel.has_local_proxy(node_id) {
return SchedulerTunnelAffinityBucket::LocalTunnel;
}
if let Some(owner_instance_id) = proxy_tunnel_owner_instance_id(proxy) {
return if owner_instance_id == state.app().tunnel.local_instance_id() {
SchedulerTunnelAffinityBucket::LocalTunnel
} else {
SchedulerTunnelAffinityBucket::RemoteTunnel
};
}
match state
.app()
.tunnel
@@ -144,10 +232,25 @@ async fn resolve_tunnel_owner_affinity_from_transport(
fn candidate_transport_identity(
candidate: &SchedulerMinimalCandidateSelectionCandidate,
) -> CandidateTransportIdentity<'_> {
) -> CandidateTransportIdentity {
(
candidate.provider_id.as_str(),
candidate.endpoint_id.as_str(),
candidate.key_id.as_str(),
candidate.provider_id.clone(),
candidate.endpoint_id.clone(),
candidate.key_id.clone(),
)
}
fn proxy_config_cache_key(raw: &Value) -> String {
serde_json::to_string(raw).unwrap_or_else(|_| raw.to_string())
}
fn proxy_tunnel_owner_instance_id(proxy: &ProxySnapshot) -> Option<&str> {
proxy
.extra
.as_ref()
.and_then(Value::as_object)
.and_then(|extra| extra.get(TUNNEL_OWNER_INSTANCE_ID_EXTRA_KEY))
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
}
@@ -66,7 +66,7 @@ pub(crate) async fn maybe_build_sync_local_same_format_provider_decision_payload
candidate_count,
);
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
if let Some(payload) =
maybe_build_local_same_format_provider_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec,
@@ -134,7 +134,7 @@ pub(crate) async fn maybe_build_stream_local_same_format_provider_decision_paylo
candidate_count,
);
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
if let Some(payload) =
maybe_build_local_same_format_provider_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec,
@@ -190,7 +190,7 @@ pub(crate) async fn build_local_stream_attempt_source<'a>(
#[async_trait]
impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalSameFormatProviderSyncAttemptSource<'_> {
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
while let Some(attempt) = self.candidates.next_attempt().await {
while let Some(attempt) = self.candidates.next_attempt().await? {
match self.build_sync_attempt(attempt).await? {
Some(attempt) => return Ok(Some(attempt)),
None => continue,
@@ -220,7 +220,7 @@ impl LocalExecutionAttemptSource<AiStreamAttempt>
for LocalSameFormatProviderStreamAttemptSource<'_>
{
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
while let Some(attempt) = self.candidates.next_attempt().await {
while let Some(attempt) = self.candidates.next_attempt().await? {
match self.build_stream_attempt(attempt).await? {
Some(attempt) => return Ok(Some(attempt)),
None => continue,
@@ -372,7 +372,7 @@ pub(crate) async fn build_local_sync_plan_and_reports(
}
let mut plans = Vec::new();
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
let Some(payload) = maybe_build_local_same_format_provider_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec,
)
@@ -458,7 +458,7 @@ pub(crate) async fn build_local_stream_plan_and_reports(
}
let mut plans = Vec::new();
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
let Some(payload) = maybe_build_local_same_format_provider_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec,
)
@@ -1,5 +1,5 @@
use std::borrow::Cow;
use std::time::{SystemTime, UNIX_EPOCH};
use std::time::{Instant, SystemTime, UNIX_EPOCH};
use serde_json::Value;
use tracing::warn;
@@ -7,9 +7,11 @@ use tracing::warn;
use crate::ai_serving::ExecutionRuntimeAuthContext;
use crate::privacy::{
build_redaction_session_config, read_chat_pii_redaction_runtime_config,
try_mask_chat_pii_request_json_with_cache_options, ChatPiiRedactionRequestFormat,
MaskChatRequestOptions, RedactionMaskError, RedactionSessionSlot, RedisRedactionMappingCache,
try_mask_chat_pii_request_value_with_cache_options, CachedRequestRedaction,
ChatPiiRedactionRequestFormat, MaskChatRequestOptions, RedactionMaskError,
RedactionSessionSlot, RedisRedactionMappingCache,
};
use crate::stage_metrics::observe_gateway_stage_ms;
use crate::{AppState, GatewayError};
pub(crate) struct ProviderRequestRedaction<'a> {
@@ -73,6 +75,18 @@ pub(crate) async fn resolve_provider_chat_pii_redaction<'a>(
let Some(slot) = parts.extensions.get::<RedactionSessionSlot>() else {
return Ok(ProviderRequestRedaction::disabled(body_json));
};
let request_cache_key = request_redaction_cache_key(format, body_json);
if let Some(cached) = slot.cached_request_redaction(&request_cache_key) {
observe_gateway_stage_ms("chat_pii_redaction_request_cache_hit", 0);
return Ok(provider_redaction_from_cached(
slot,
candidate_id,
body_json,
cached,
));
}
let runtime_config_started_at = Instant::now();
let runtime_config = read_chat_pii_redaction_runtime_config(state)
.await
.map_err(|err| {
@@ -82,11 +96,22 @@ pub(crate) async fn resolve_provider_chat_pii_redaction<'a>(
);
GatewayError::Internal("chat pii redaction setup failed".to_string())
})?;
observe_gateway_stage_ms(
"chat_pii_redaction_runtime_config",
runtime_config_started_at.elapsed().as_millis() as u64,
);
if !runtime_config.enabled {
slot.put_cached_request_redaction(request_cache_key, CachedRequestRedaction::unredacted());
return Ok(ProviderRequestRedaction::disabled(body_json));
}
let feature_settings_started_at = Instant::now();
let feature_settings = resolve_chat_pii_redaction_feature_settings(state, auth_context).await?;
observe_gateway_stage_ms(
"chat_pii_redaction_feature_settings",
feature_settings_started_at.elapsed().as_millis() as u64,
);
if !feature_settings.effective_enabled() {
slot.put_cached_request_redaction(request_cache_key, CachedRequestRedaction::unredacted());
return Ok(ProviderRequestRedaction::disabled(body_json));
}
let Some(hmac_key) = state.encryption_key().map(str::as_bytes).map(Vec::from) else {
@@ -95,20 +120,14 @@ pub(crate) async fn resolve_provider_chat_pii_redaction<'a>(
"chat pii redaction setup failed".to_string(),
));
};
let body_bytes = serde_json::to_vec(body_json).map_err(|err| {
warn!(
error = ?err,
"gateway failed to serialize provider chat pii redaction body"
);
GatewayError::Internal("chat pii redaction setup failed".to_string())
})?;
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
let cache = RedisRedactionMappingCache::new(state.runtime_state.as_ref());
let masked = try_mask_chat_pii_request_json_with_cache_options(
&body_bytes,
let mask_started_at = Instant::now();
let masked = try_mask_chat_pii_request_value_with_cache_options(
body_json,
format,
build_redaction_session_config(hmac_key, &runtime_config, now_unix_secs),
MaskChatRequestOptions::runtime(),
@@ -116,19 +135,27 @@ pub(crate) async fn resolve_provider_chat_pii_redaction<'a>(
)
.await
.map_err(redaction_mask_error_to_gateway_error)?;
observe_gateway_stage_ms(
"chat_pii_redaction_mask_body",
mask_started_at.elapsed().as_millis() as u64,
);
if !masked.redacted {
slot.put_cached_request_redaction(request_cache_key, CachedRequestRedaction::unredacted());
return Ok(ProviderRequestRedaction {
body_json: Cow::Borrowed(body_json),
redacted: false,
});
}
let masked_body_json = serde_json::from_slice::<Value>(&masked.body).map_err(|err| {
warn!(
error = ?err,
"gateway failed to decode redacted provider chat pii body"
);
GatewayError::Internal("chat pii redaction setup failed".to_string())
})?;
let Some(masked_body_json) = masked.body_json else {
warn!("gateway pii redaction reported redacted without masked body");
return Err(GatewayError::Internal(
"chat pii redaction setup failed".to_string(),
));
};
slot.put_cached_request_redaction(
request_cache_key,
CachedRequestRedaction::redacted(masked_body_json.clone(), masked.session.clone()),
);
slot.put_for_candidate(candidate_id, masked.session);
Ok(ProviderRequestRedaction {
body_json: Cow::Owned(masked_body_json),
@@ -136,31 +163,46 @@ pub(crate) async fn resolve_provider_chat_pii_redaction<'a>(
})
}
fn request_redaction_cache_key(format: ChatPiiRedactionRequestFormat, body_json: &Value) -> String {
format!("{format:?}:{:p}", body_json)
}
fn provider_redaction_from_cached<'a>(
slot: &RedactionSessionSlot,
candidate_id: &str,
body_json: &'a Value,
cached: CachedRequestRedaction,
) -> ProviderRequestRedaction<'a> {
if !cached.redacted {
return ProviderRequestRedaction::disabled(body_json);
}
let Some(masked_body_json) = cached.body_json else {
return ProviderRequestRedaction::disabled(body_json);
};
if let Some(session) = cached.session {
slot.put_for_candidate(candidate_id, session);
}
ProviderRequestRedaction {
body_json: Cow::Owned(masked_body_json),
redacted: true,
}
}
async fn resolve_chat_pii_redaction_feature_settings(
state: &AppState,
auth_context: &ExecutionRuntimeAuthContext,
) -> Result<ChatPiiRedactionFeatureSettings, GatewayError> {
let user_settings = state
.read_user_feature_settings(&auth_context.user_id)
.await
let user_settings_fut = state.read_user_feature_settings(&auth_context.user_id);
let key_settings_fut = state.read_auth_api_key_feature_settings(
&auth_context.user_id,
&auth_context.api_key_id,
auth_context.api_key_is_standalone,
);
let (user_settings, key_settings) = tokio::try_join!(user_settings_fut, key_settings_fut)
.map_err(|err| {
warn!(
error = ?err,
"gateway failed to read user chat pii redaction feature settings"
);
GatewayError::Internal("chat pii redaction setup failed".to_string())
})?;
let key_settings = state
.read_auth_api_key_feature_settings(
&auth_context.user_id,
&auth_context.api_key_id,
auth_context.api_key_is_standalone,
)
.await
.map_err(|err| {
warn!(
error = ?err,
"gateway failed to read api key chat pii redaction feature settings"
"gateway failed to read chat pii redaction feature settings"
);
GatewayError::Internal("chat pii redaction setup failed".to_string())
})?;
@@ -175,7 +175,7 @@ pub(crate) async fn build_local_gemini_files_stream_attempt_source_for_kind<'a>(
#[async_trait]
impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalGeminiFilesSyncAttemptSource<'_> {
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
while let Some(attempt) = self.candidates.next_attempt().await {
while let Some(attempt) = self.candidates.next_attempt().await? {
match self.build_sync_attempt(attempt).await? {
Some(attempt) => return Ok(Some(attempt)),
None => continue,
@@ -198,7 +198,7 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalGeminiFilesSyncAttemptS
#[async_trait]
impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalGeminiFilesStreamAttemptSource<'_> {
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
while let Some(attempt) = self.candidates.next_attempt().await {
while let Some(attempt) = self.candidates.next_attempt().await? {
match self.build_stream_attempt(attempt).await? {
Some(attempt) => return Ok(Some(attempt)),
None => continue,
@@ -323,7 +323,7 @@ pub(crate) async fn maybe_build_sync_local_gemini_files_decision_payload(
let (mut source, _) =
build_local_gemini_files_candidate_attempt_source(state, trace_id, &input).await?;
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
if let Some(payload) = maybe_build_local_gemini_files_decision_payload_for_candidate(
state,
parts,
@@ -365,7 +365,7 @@ pub(crate) async fn maybe_build_stream_local_gemini_files_decision_payload(
build_local_gemini_files_candidate_attempt_source(state, trace_id, &input).await?;
let empty_body_json = serde_json::Value::Null;
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
if let Some(payload) = maybe_build_local_gemini_files_decision_payload_for_candidate(
state,
parts,
@@ -414,7 +414,7 @@ async fn build_local_sync_plan_and_reports(
build_local_gemini_files_candidate_attempt_source(state, trace_id, &input).await?;
let mut plans = Vec::new();
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
let Some(payload) = maybe_build_local_gemini_files_decision_payload_for_candidate(
state,
parts,
@@ -467,7 +467,7 @@ async fn build_local_stream_plan_and_reports(
let mut plans = Vec::new();
let empty_body_json = serde_json::Value::Null;
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
let Some(payload) = maybe_build_local_gemini_files_decision_payload_for_candidate(
state,
parts,
@@ -253,7 +253,7 @@ pub(crate) async fn build_local_image_stream_attempt_source_for_kind<'a>(
#[async_trait]
impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalOpenAiImageSyncAttemptSource<'_> {
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
while let Some(attempt) = self.candidates.next_attempt().await {
while let Some(attempt) = self.candidates.next_attempt().await? {
match self.build_sync_attempt(attempt).await? {
Some(attempt) => return Ok(Some(attempt)),
None => continue,
@@ -276,7 +276,7 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalOpenAiImageSyncAttemptS
#[async_trait]
impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalOpenAiImageStreamAttemptSource<'_> {
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
while let Some(attempt) = self.candidates.next_attempt().await {
while let Some(attempt) = self.candidates.next_attempt().await? {
match self.build_stream_attempt(attempt).await? {
Some(attempt) => return Ok(Some(attempt)),
None => continue,
@@ -421,7 +421,7 @@ pub(crate) async fn maybe_build_sync_local_image_decision_payload(
return Ok(None);
};
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
if let Some(payload) = maybe_build_local_openai_image_decision_payload_for_candidate(
state,
parts,
@@ -482,7 +482,7 @@ pub(crate) async fn maybe_build_stream_local_image_decision_payload(
return Ok(None);
};
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
if let Some(payload) = maybe_build_local_openai_image_decision_payload_for_candidate(
state,
parts,
@@ -540,7 +540,7 @@ async fn build_local_sync_plan_and_reports(
};
let mut plans = Vec::new();
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
let Some(payload) = maybe_build_local_openai_image_decision_payload_for_candidate(
state,
parts,
@@ -617,7 +617,7 @@ async fn build_local_stream_plan_and_reports(
};
let mut plans = Vec::new();
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
let Some(payload) = maybe_build_local_openai_image_decision_payload_for_candidate(
state,
parts,
@@ -105,7 +105,7 @@ pub(crate) async fn build_local_video_sync_attempt_source_for_kind<'a>(
#[async_trait]
impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalVideoCreateSyncAttemptSource<'_> {
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
while let Some(attempt) = self.candidates.next_attempt().await {
while let Some(attempt) = self.candidates.next_attempt().await? {
match self.build_sync_attempt(attempt).await? {
Some(attempt) => return Ok(Some(attempt)),
None => continue,
@@ -195,7 +195,7 @@ pub(crate) async fn maybe_build_sync_local_video_decision_payload(
return Ok(None);
};
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
if let Some(payload) = maybe_build_local_video_create_decision_payload_for_candidate(
state, parts, body_json, trace_id, &input, attempt, spec,
)
@@ -240,7 +240,7 @@ async fn build_local_sync_plan_and_reports(
};
let mut plans = Vec::new();
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
let Some(payload) = maybe_build_local_video_create_decision_payload_for_candidate(
state, parts, body_json, trace_id, &input, attempt, spec,
)
@@ -178,7 +178,7 @@ pub(crate) async fn build_local_stream_attempt_source<'a>(
#[async_trait]
impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalStandardSyncAttemptSource<'_> {
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
while let Some(attempt) = self.candidates.next_attempt().await {
while let Some(attempt) = self.candidates.next_attempt().await? {
match self.build_sync_attempt(attempt).await? {
Some(attempt) => return Ok(Some(attempt)),
None => continue,
@@ -206,7 +206,7 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalStandardSyncAttemptSour
#[async_trait]
impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalStandardStreamAttemptSource<'_> {
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
while let Some(attempt) = self.candidates.next_attempt().await {
while let Some(attempt) = self.candidates.next_attempt().await? {
match self.build_stream_attempt(attempt).await? {
Some(attempt) => return Ok(Some(attempt)),
None => continue,
@@ -340,7 +340,7 @@ pub(crate) async fn maybe_build_sync_via_standard_family_payload(
.await?;
apply_local_runtime_candidate_evaluation_progress(state, trace_id, candidate_count);
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
if let Some(payload) = maybe_build_local_standard_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec,
)
@@ -390,7 +390,7 @@ pub(crate) async fn maybe_build_stream_via_standard_family_payload(
.await?;
apply_local_runtime_candidate_evaluation_progress(state, trace_id, candidate_count);
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
if let Some(payload) = maybe_build_local_standard_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec,
)
@@ -449,7 +449,7 @@ pub(crate) async fn build_local_sync_plan_and_reports(
return Ok(Vec::new());
}
let mut plans = Vec::new();
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
let Some(payload) = maybe_build_local_standard_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec,
)
@@ -524,7 +524,7 @@ pub(crate) async fn build_local_stream_plan_and_reports(
return Ok(Vec::new());
}
let mut plans = Vec::new();
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
let Some(payload) = maybe_build_local_standard_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec,
)
@@ -12,6 +12,7 @@ use crate::ai_serving::planner::{
use crate::ai_serving::transport::{
resolve_transport_execution_timeouts, resolve_transport_profile,
};
use crate::stage_metrics::observe_gateway_stage_ms;
use crate::{
append_execution_contract_fields_to_value, append_local_failover_policy_to_value,
AiExecutionDecision, AppState, GatewayError,
@@ -40,6 +41,7 @@ pub(crate) async fn maybe_build_local_openai_chat_decision_payload_for_candidate
candidate_id,
..
} = attempt;
let payload_started_at = std::time::Instant::now();
let Some(resolved) = resolve_local_openai_chat_candidate_payload_parts(
state,
parts,
@@ -55,8 +57,16 @@ pub(crate) async fn maybe_build_local_openai_chat_decision_payload_for_candidate
)
.await?
else {
observe_gateway_stage_ms(
"stream_candidate_payload_parts",
payload_started_at.elapsed().as_millis() as u64,
);
return Ok(None);
};
observe_gateway_stage_ms(
"stream_candidate_payload_parts",
payload_started_at.elapsed().as_millis() as u64,
);
let candidate = &eligible.candidate;
let prompt_cache_key = resolved
@@ -66,9 +76,14 @@ pub(crate) async fn maybe_build_local_openai_chat_decision_payload_for_candidate
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned);
let proxy_started_at = std::time::Instant::now();
let proxy = state
.resolve_transport_proxy_snapshot_with_tunnel_affinity(&resolved.transport)
.await;
observe_gateway_stage_ms(
"stream_candidate_proxy",
proxy_started_at.elapsed().as_millis() as u64,
);
let transport_profile = resolved
.transport_profile
.clone()
@@ -130,6 +145,7 @@ pub(crate) async fn maybe_build_local_openai_chat_decision_payload_for_candidate
Some(body_json)
};
let effective_headers = input.effective_headers(&parts.headers);
let report_context_started_at = std::time::Instant::now();
let report_context = append_local_failover_policy_to_value(
append_execution_contract_fields_to_value(
build_local_execution_report_context(LocalExecutionReportContextParts {
@@ -184,8 +200,13 @@ pub(crate) async fn maybe_build_local_openai_chat_decision_payload_for_candidate
),
&transport,
);
observe_gateway_stage_ms(
"stream_candidate_report_context",
report_context_started_at.elapsed().as_millis() as u64,
);
let request_gzip = resolve_transport_request_gzip_policy(&transport);
let decision_started_at = std::time::Instant::now();
let mut decision = build_ai_execution_decision_response(AiExecutionDecisionResponseParts {
decision_is_stream,
decision_kind: decision_kind.to_string(),
@@ -222,5 +243,9 @@ pub(crate) async fn maybe_build_local_openai_chat_decision_payload_for_candidate
auth_context: input.auth_context.clone(),
});
apply_provider_request_routing_policy_to_decision(input, &mut decision)?;
observe_gateway_stage_ms(
"stream_candidate_decision_build",
decision_started_at.elapsed().as_millis() as u64,
);
Ok(Some(decision))
}
@@ -54,6 +54,7 @@ use crate::ai_serving::{
LocalResolvedOAuthRequestAuth,
};
use crate::ai_serving::{ConversionMode, ExecutionStrategy};
use crate::stage_metrics::observe_gateway_stage_ms;
use crate::{AppState, GatewayError};
use super::support::{
@@ -102,6 +103,7 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
report_kind: &str,
upstream_is_stream: bool,
) -> Result<Option<LocalOpenAiChatCandidatePayloadParts>, GatewayError> {
let prepare_started_at = std::time::Instant::now();
let planner_state = crate::ai_serving::PlannerAppState::new(state);
let candidate = &eligible.candidate;
let provider_api_format = eligible.provider_api_format.as_str();
@@ -109,6 +111,7 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
let transport_profile = crate::ai_serving::transport::resolve_transport_profile(transport);
let force_body_stream_field =
endpoint_config_forces_body_stream_field(transport.endpoint.config.as_ref());
let model_directives_started_at = std::time::Instant::now();
let enable_model_directives =
crate::system_features::reasoning_model_directive_enabled_for_api_format_and_model(
state,
@@ -116,6 +119,11 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
Some(&input.requested_model),
)
.await;
observe_gateway_stage_ms(
"openai_chat_payload_model_directives",
model_directives_started_at.elapsed().as_millis() as u64,
);
let redaction_started_at = std::time::Instant::now();
let redaction = resolve_provider_chat_pii_redaction(
state,
parts,
@@ -125,6 +133,10 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
candidate_id,
)
.await?;
observe_gateway_stage_ms(
"openai_chat_payload_redaction",
redaction_started_at.elapsed().as_millis() as u64,
);
let body_json = redaction.body_json.as_ref();
let effective_headers = input.effective_headers(&parts.headers);
let is_grok = transport
@@ -234,7 +246,7 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
redaction.redacted,
);
return Ok(Some(LocalOpenAiChatCandidatePayloadParts {
let result = Ok(Some(LocalOpenAiChatCandidatePayloadParts {
client_api_format: "openai:chat".to_string(),
auth_header: prepared_candidate.auth_header,
auth_value: prepared_candidate.auth_value,
@@ -252,6 +264,11 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
transport_profile,
image_request_summary: None,
}));
observe_gateway_stage_ms(
"openai_chat_payload_parts_prepare",
prepare_started_at.elapsed().as_millis() as u64,
);
return result;
}
if provider_api_format == "openai:chat" && is_windsurf_provider_transport(transport) {
@@ -288,6 +305,7 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
return Ok(None);
};
let auth_prepare_started_at = std::time::Instant::now();
let prepared_candidate = match prepare_header_authenticated_candidate(
planner_state,
transport,
@@ -316,7 +334,12 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
return Ok(None);
}
};
observe_gateway_stage_ms(
"openai_chat_payload_auth_prepare",
auth_prepare_started_at.elapsed().as_millis() as u64,
);
let body_build_started_at = std::time::Instant::now();
let Some(mut provider_request_body) = build_local_openai_chat_request_body(
body_json,
&prepared_candidate.mapped_model,
@@ -343,6 +366,10 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
.await;
return Ok(None);
};
observe_gateway_stage_ms(
"openai_chat_payload_body_build",
body_build_started_at.elapsed().as_millis() as u64,
);
apply_deepseek_tool_call_thinking_compat(
&mut provider_request_body,
transport.provider.provider_type.as_str(),
@@ -147,7 +147,7 @@ pub(crate) async fn maybe_build_sync_local_decision_payload(
)
.await;
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
let upstream_is_stream = self::plans::openai_chat_upstream_is_stream_for_candidate(
&attempt.eligible.transport,
attempt.eligible.provider_api_format.as_str(),
@@ -199,7 +199,7 @@ pub(crate) async fn maybe_build_stream_local_decision_payload(
)
.await;
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
let upstream_is_stream = self::plans::openai_chat_upstream_is_stream_for_candidate(
&attempt.eligible.transport,
attempt.eligible.provider_api_format.as_str(),
@@ -18,6 +18,7 @@ use crate::ai_serving::planner::plan_builders::{
build_openai_chat_stream_plan_from_decision, AiStreamAttempt,
};
use crate::ai_serving::planner::runtime_miss::apply_local_runtime_candidate_terminal_reason;
use crate::stage_metrics::observe_gateway_stage_ms;
pub(crate) struct LocalOpenAiChatStreamAttemptSource<'a> {
state: &'a AppState,
@@ -93,10 +94,36 @@ pub(crate) async fn build_local_openai_chat_stream_attempt_source<'a>(
#[async_trait]
impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalOpenAiChatStreamAttemptSource<'_> {
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
while let Some(attempt) = self.candidates.next_attempt().await {
loop {
let source_started_at = std::time::Instant::now();
let Some(attempt) = self.candidates.next_attempt().await? else {
observe_gateway_stage_ms(
"stream_candidate_source_next",
source_started_at.elapsed().as_millis() as u64,
);
break;
};
observe_gateway_stage_ms(
"stream_candidate_source_next",
source_started_at.elapsed().as_millis() as u64,
);
let plan_started_at = std::time::Instant::now();
match self.build_stream_attempt(attempt).await? {
Some(attempt) => return Ok(Some(attempt)),
None => continue,
Some(attempt) => {
observe_gateway_stage_ms(
"stream_candidate_plan_build",
plan_started_at.elapsed().as_millis() as u64,
);
return Ok(Some(attempt));
}
None => {
observe_gateway_stage_ms(
"stream_candidate_plan_build",
plan_started_at.elapsed().as_millis() as u64,
);
continue;
}
}
}
apply_local_runtime_candidate_terminal_reason(
@@ -93,7 +93,7 @@ pub(crate) async fn build_local_openai_chat_sync_attempt_source<'a>(
#[async_trait]
impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalOpenAiChatSyncAttemptSource<'_> {
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
while let Some(attempt) = self.candidates.next_attempt().await {
while let Some(attempt) = self.candidates.next_attempt().await? {
match self.build_sync_attempt(attempt).await? {
Some(attempt) => return Ok(Some(attempt)),
None => continue,
@@ -114,7 +114,7 @@ pub(crate) async fn maybe_build_sync_local_openai_responses_decision_payload(
)
.await?;
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
if let Some(payload) = maybe_build_local_openai_responses_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec,
)
@@ -153,7 +153,7 @@ pub(crate) async fn maybe_build_stream_local_openai_responses_decision_payload(
)
.await?;
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
if let Some(payload) = maybe_build_local_openai_responses_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec,
)
@@ -162,7 +162,7 @@ pub(super) async fn build_local_stream_attempt_source<'a>(
#[async_trait]
impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalOpenAiResponsesSyncAttemptSource<'_> {
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
while let Some(attempt) = self.candidates.next_attempt().await {
while let Some(attempt) = self.candidates.next_attempt().await? {
match self.build_sync_attempt(attempt).await? {
Some(attempt) => return Ok(Some(attempt)),
None => continue,
@@ -190,7 +190,7 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalOpenAiResponsesSyncAtte
#[async_trait]
impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalOpenAiResponsesStreamAttemptSource<'_> {
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
while let Some(attempt) = self.candidates.next_attempt().await {
while let Some(attempt) = self.candidates.next_attempt().await? {
match self.build_stream_attempt(attempt).await? {
Some(attempt) => return Ok(Some(attempt)),
None => continue,
@@ -331,7 +331,7 @@ pub(super) async fn build_local_sync_plan_and_reports(
}
let mut plans = Vec::new();
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
let Some(payload) = maybe_build_local_openai_responses_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec,
)
@@ -403,7 +403,7 @@ pub(super) async fn build_local_stream_plan_and_reports(
}
let mut plans = Vec::new();
while let Some(attempt) = source.next_attempt().await {
while let Some(attempt) = source.next_attempt().await? {
let Some(payload) = maybe_build_local_openai_responses_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec,
)
@@ -3,8 +3,20 @@ pub(crate) use crate::ai_serving::transport::{
GatewayProviderTransportSnapshot, LocalResolvedOAuthRequestAuth,
};
use crate::GatewayError;
use std::sync::Arc;
impl<'a> PlannerAppState<'a> {
pub(crate) async fn read_provider_transport_snapshot_arc(
self,
provider_id: &str,
endpoint_id: &str,
key_id: &str,
) -> Result<Option<Arc<GatewayProviderTransportSnapshot>>, GatewayError> {
self.app()
.read_provider_transport_snapshot_arc(provider_id, endpoint_id, key_id)
.await
}
pub(crate) async fn read_provider_transport_snapshot(
self,
provider_id: &str,