mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-08 10:27: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,
|
||||
|
||||
Reference in New Issue
Block a user