feat(providers): add provider transfer limits

This commit is contained in:
elky
2026-07-26 15:06:56 +08:00
parent 2ef7ac79bc
commit 10d369f59c
36 changed files with 1764 additions and 119 deletions
@@ -16,7 +16,7 @@ use aether_scheduler_core::{
};
use async_trait::async_trait;
use serde_json::Value;
use std::collections::VecDeque;
use std::collections::{BTreeSet, VecDeque};
use std::convert::Infallible;
use std::sync::Arc;
use std::time::Duration;
@@ -67,6 +67,7 @@ pub(crate) struct LocalExecutionCandidateAttempt {
pub(crate) struct LocalExecutionCandidateAttemptSource<'a> {
items: VecDeque<LocalExecutionCandidateAttemptSourceItem<'a>>,
skipped_provider_ids: BTreeSet<String>,
}
type DecorateSkippedCandidateFn<'a> = Arc<
@@ -78,6 +79,8 @@ pub(crate) trait LocalExecutionAttemptSource<T>: Send {
async fn next_execution_attempt(&mut self) -> Result<Option<T>, GatewayError>;
async fn drain_execution_attempts(&mut self) -> Result<Vec<T>, GatewayError>;
async fn skip_provider(&mut self, provider_id: &str) -> Result<(), GatewayError>;
}
enum LocalExecutionCandidateAttemptSourceItem<'a> {
@@ -105,7 +108,10 @@ impl<'a> LocalExecutionCandidateAttemptSource<'a> {
attempts: dispatch_sequence_from_attempts(attempts),
});
}
Self { items }
Self {
items,
skipped_provider_ids: BTreeSet::new(),
}
}
pub(crate) async fn next_attempt(
@@ -117,6 +123,10 @@ impl<'a> LocalExecutionCandidateAttemptSource<'a> {
};
match front {
LocalExecutionCandidateAttemptSourceItem::Static { attempts } => {
if dispatch_sequence_provider_is_skipped(attempts, &self.skipped_provider_ids) {
self.items.pop_front();
continue;
}
if let Some(attempt) = next_attempt_from_dispatch_sequence(attempts) {
if dispatch_sequence_exhausted(attempts) {
self.items.pop_front();
@@ -131,6 +141,10 @@ impl<'a> LocalExecutionCandidateAttemptSource<'a> {
pending_attempts,
pool_exhaustion_persistence,
} => {
if self.skipped_provider_ids.contains(cursor.provider_id()) {
self.items.pop_front();
continue;
}
if let Some(attempt) = next_attempt_from_dispatch_sequence(pending_attempts) {
return Ok(Some(attempt));
}
@@ -157,6 +171,9 @@ impl<'a> LocalExecutionCandidateAttemptSource<'a> {
);
}
LocalExecutionCandidateAttemptSourceItem::RequestedModelPage { cursor } => {
for provider_id in &self.skipped_provider_ids {
cursor.skip_provider(provider_id);
}
let Some(attempt) = cursor.next_attempt().await? else {
self.items.pop_front();
continue;
@@ -171,6 +188,19 @@ impl<'a> LocalExecutionCandidateAttemptSource<'a> {
self.items.clear();
Vec::new()
}
pub(crate) fn skip_provider(&mut self, provider_id: &str) {
let provider_id = provider_id.trim();
if provider_id.is_empty() {
return;
}
self.skipped_provider_ids.insert(provider_id.to_string());
for item in &mut self.items {
if let LocalExecutionCandidateAttemptSourceItem::RequestedModelPage { cursor } = item {
cursor.skip_provider(provider_id);
}
}
}
}
impl LocalExecutionCandidateAttempt {
@@ -635,7 +665,10 @@ where
);
(
LocalExecutionCandidateAttemptSource { items },
LocalExecutionCandidateAttemptSource {
items,
skipped_provider_ids: BTreeSet::new(),
},
candidate_count,
)
}
@@ -768,6 +801,7 @@ where
decorate_skipped_candidate,
page_cursor,
pending_items: VecDeque::new(),
skipped_provider_ids: BTreeSet::new(),
candidate_count: 0,
next_candidate_index: 0,
remembered_affinity: false,
@@ -788,7 +822,10 @@ where
);
}
(
LocalExecutionCandidateAttemptSource { items },
LocalExecutionCandidateAttemptSource {
items,
skipped_provider_ids: BTreeSet::new(),
},
candidate_count,
)
}
@@ -813,6 +850,7 @@ struct RequestedModelAttemptPageCursor<'a> {
decorate_skipped_candidate: DecorateSkippedCandidateFn<'a>,
page_cursor: LocalCandidatePreselectionPageCursor<'a>,
pending_items: VecDeque<LocalExecutionCandidateAttemptSourceItem<'a>>,
skipped_provider_ids: BTreeSet<String>,
candidate_count: usize,
next_candidate_index: u32,
remembered_affinity: bool,
@@ -822,6 +860,10 @@ struct RequestedModelAttemptPageCursor<'a> {
}
impl<'a> RequestedModelAttemptPageCursor<'a> {
fn skip_provider(&mut self, provider_id: &str) {
self.skipped_provider_ids.insert(provider_id.to_string());
}
async fn next_attempt(
&mut self,
) -> Result<Option<LocalExecutionCandidateAttempt>, GatewayError> {
@@ -829,7 +871,9 @@ impl<'a> RequestedModelAttemptPageCursor<'a> {
return Err(error);
}
loop {
if let Some(attempt) = pop_attempt_from_items(&mut self.pending_items).await {
if let Some(attempt) =
pop_attempt_from_items(&mut self.pending_items, &self.skipped_provider_ids).await
{
return Ok(Some(attempt));
}
if !self.load_next_page().await? {
@@ -1030,11 +1074,16 @@ fn page_is_exact_auth_api_key_concurrency_limited(
async fn pop_attempt_from_items(
items: &mut VecDeque<LocalExecutionCandidateAttemptSourceItem<'_>>,
skipped_provider_ids: &BTreeSet<String>,
) -> Option<LocalExecutionCandidateAttempt> {
loop {
let front = items.front_mut()?;
match front {
LocalExecutionCandidateAttemptSourceItem::Static { attempts } => {
if dispatch_sequence_provider_is_skipped(attempts, skipped_provider_ids) {
items.pop_front();
continue;
}
if let Some(attempt) = next_attempt_from_dispatch_sequence(attempts) {
if dispatch_sequence_exhausted(attempts) {
items.pop_front();
@@ -1049,6 +1098,10 @@ async fn pop_attempt_from_items(
pending_attempts,
pool_exhaustion_persistence,
} => {
if skipped_provider_ids.contains(cursor.provider_id()) {
items.pop_front();
continue;
}
if let Some(attempt) = next_attempt_from_dispatch_sequence(pending_attempts) {
return Some(attempt);
}
@@ -1727,6 +1780,15 @@ fn next_attempt_from_dispatch_sequence(
Some(attempt)
}
fn dispatch_sequence_provider_is_skipped(
sequence: &DispatchSequence<LocalExecutionCandidateAttempt>,
skipped_provider_ids: &BTreeSet<String>,
) -> bool {
sequence.peek_current().is_some_and(|item| {
skipped_provider_ids.contains(&item.candidate.eligible.candidate.provider_id)
})
}
fn dispatch_sequence_exhausted(
sequence: &mut DispatchSequence<LocalExecutionCandidateAttempt>,
) -> bool {
@@ -2278,6 +2340,7 @@ mod tests {
decorate_skipped_candidate: Arc::new(identity_skipped_candidate),
page_cursor,
pending_items: VecDeque::new(),
skipped_provider_ids: BTreeSet::new(),
candidate_count: 0,
next_candidate_index: 0,
remembered_affinity: false,
@@ -2372,6 +2435,7 @@ mod tests {
decorate_skipped_candidate: Arc::new(identity_skipped_candidate),
page_cursor,
pending_items: VecDeque::new(),
skipped_provider_ids: BTreeSet::new(),
candidate_count: 0,
next_candidate_index: 0,
remembered_affinity: false,
@@ -2569,6 +2633,7 @@ mod tests {
.into(),
),
}]),
skipped_provider_ids: BTreeSet::new(),
};
let first = source
@@ -2587,6 +2652,53 @@ mod tests {
.is_none());
}
#[tokio::test]
async fn skipped_provider_discards_pool_cursor_and_continues_with_next_provider() {
let app = AppState::new().expect("state should build");
let mut pool_group = sample_eligible("pool-group", None);
pool_group.kind = LocalExecutionCandidateKind::PoolGroup;
pool_group.transport = sample_transport("pool-group", Some(json!({ "pool_advanced": {} })));
let pool_cursor = PoolKeyCursor::new(
PlannerAppState::new(&app),
pool_group,
None,
Some("gpt-5"),
None,
);
let mut fallback = sample_eligible("fallback-key", None);
fallback.candidate.provider_id = "provider-b".to_string();
Arc::make_mut(&mut fallback.transport).provider.id = "provider-b".to_string();
Arc::make_mut(&mut fallback.transport).key.provider_id = "provider-b".to_string();
let fallback_attempts = dispatch_sequence_from_attempts(
build_unpersisted_local_execution_candidate_attempts(fallback, 1).into(),
);
let mut source = LocalExecutionCandidateAttemptSource {
items: VecDeque::from([
LocalExecutionCandidateAttemptSourceItem::Pool {
cursor: pool_cursor,
candidate_index: 0,
pending_attempts: DispatchSequence::new(Vec::new()),
pool_exhaustion_persistence: None,
},
LocalExecutionCandidateAttemptSourceItem::Static {
attempts: fallback_attempts,
},
]),
skipped_provider_ids: BTreeSet::new(),
};
source.skip_provider("provider-1");
let attempt = source
.next_attempt()
.await
.expect("candidate source should succeed")
.expect("fallback provider should remain");
assert_eq!(attempt.eligible.candidate.provider_id, "provider-b");
assert_eq!(attempt.eligible.candidate.key_id, "fallback-key");
}
#[tokio::test]
async fn dynamic_pool_exhaustion_persists_group_skip_summary() {
let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default());
@@ -2639,6 +2751,7 @@ mod tests {
pending_attempts: DispatchSequence::new(Vec::new()),
pool_exhaustion_persistence: Some(pool_exhaustion_persistence),
}]),
skipped_provider_ids: BTreeSet::new(),
};
assert!(source
@@ -213,6 +213,11 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalSameFormatProviderSyncA
}
Ok(drained)
}
async fn skip_provider(&mut self, provider_id: &str) -> Result<(), GatewayError> {
self.candidates.skip_provider(provider_id);
Ok(())
}
}
#[async_trait]
@@ -243,6 +248,11 @@ impl LocalExecutionAttemptSource<AiStreamAttempt>
}
Ok(drained)
}
async fn skip_provider(&mut self, provider_id: &str) -> Result<(), GatewayError> {
self.candidates.skip_provider(provider_id);
Ok(())
}
}
impl LocalSameFormatProviderSyncAttemptSource<'_> {
@@ -193,6 +193,11 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalGeminiFilesSyncAttemptS
}
Ok(drained)
}
async fn skip_provider(&mut self, provider_id: &str) -> Result<(), GatewayError> {
self.candidates.skip_provider(provider_id);
Ok(())
}
}
#[async_trait]
@@ -216,6 +221,11 @@ impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalGeminiFilesStreamAtte
}
Ok(drained)
}
async fn skip_provider(&mut self, provider_id: &str) -> Result<(), GatewayError> {
self.candidates.skip_provider(provider_id);
Ok(())
}
}
impl LocalGeminiFilesSyncAttemptSource<'_> {
@@ -14,7 +14,7 @@ use crate::ai_serving::transport::{
resolve_transport_execution_timeouts, resolve_transport_profile,
};
use crate::ai_serving::{ai_local_execution_contract_for_formats, PlannerAppState};
use crate::{AiExecutionDecision, AppState, GatewayError};
use crate::{append_local_failover_policy_to_value, AiExecutionDecision, AppState, GatewayError};
use super::request::resolve_local_gemini_files_candidate_payload_parts;
use super::support::{
@@ -114,6 +114,7 @@ pub(super) async fn maybe_build_local_gemini_files_decision_payload_for_candidat
needs_conversion: false,
extra_fields,
});
let report_context = append_local_failover_policy_to_value(report_context, &transport);
let super::request::LocalGeminiFilesCandidatePayloadParts {
transport: _,
auth_header,
@@ -271,6 +271,11 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalOpenAiImageSyncAttemptS
}
Ok(drained)
}
async fn skip_provider(&mut self, provider_id: &str) -> Result<(), GatewayError> {
self.candidates.skip_provider(provider_id);
Ok(())
}
}
#[async_trait]
@@ -294,6 +299,11 @@ impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalOpenAiImageStreamAtte
}
Ok(drained)
}
async fn skip_provider(&mut self, provider_id: &str) -> Result<(), GatewayError> {
self.candidates.skip_provider(provider_id);
Ok(())
}
}
impl LocalOpenAiImageSyncAttemptSource<'_> {
@@ -13,7 +13,8 @@ use crate::ai_serving::transport::{
};
use crate::ai_serving::{ai_local_execution_contract_for_formats, PlannerAppState};
use crate::{
append_execution_contract_fields_to_value, AiExecutionDecision, AppState, GatewayError,
append_execution_contract_fields_to_value, append_local_failover_policy_to_value,
AiExecutionDecision, AppState, GatewayError,
};
use super::request::resolve_local_openai_image_candidate_payload_parts;
@@ -131,6 +132,7 @@ pub(super) async fn maybe_build_local_openai_image_decision_payload_for_candidat
spec_metadata.api_format,
provider_api_format.as_str(),
);
let report_context = append_local_failover_policy_to_value(report_context, &transport);
let request_encoding = resolve_transport_request_encoding_policy(&transport);
let mut decision = build_ai_execution_decision_response(AiExecutionDecisionResponseParts {
@@ -123,6 +123,11 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalVideoCreateSyncAttemptS
}
Ok(drained)
}
async fn skip_provider(&mut self, provider_id: &str) -> Result<(), GatewayError> {
self.candidates.skip_provider(provider_id);
Ok(())
}
}
impl LocalVideoCreateSyncAttemptSource<'_> {
@@ -12,7 +12,7 @@ use crate::ai_serving::transport::{
resolve_transport_execution_timeouts, resolve_transport_profile,
};
use crate::ai_serving::{ai_local_execution_contract_for_formats, PlannerAppState};
use crate::{AiExecutionDecision, AppState, GatewayError};
use crate::{append_local_failover_policy_to_value, AiExecutionDecision, AppState, GatewayError};
use super::request::resolve_local_video_create_candidate_payload_parts;
use super::support::{LocalVideoCreateCandidateAttempt, LocalVideoCreateDecisionInput};
@@ -95,6 +95,7 @@ pub(super) async fn maybe_build_local_video_create_decision_payload_for_candidat
needs_conversion: false,
extra_fields,
});
let report_context = append_local_failover_policy_to_value(report_context, &transport);
let super::request::LocalVideoCreateCandidatePayloadParts {
transport: _,
auth_header,
@@ -201,6 +201,11 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalStandardSyncAttemptSour
}
Ok(drained)
}
async fn skip_provider(&mut self, provider_id: &str) -> Result<(), GatewayError> {
self.candidates.skip_provider(provider_id);
Ok(())
}
}
#[async_trait]
@@ -229,6 +234,11 @@ impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalStandardStreamAttempt
}
Ok(drained)
}
async fn skip_provider(&mut self, provider_id: &str) -> Result<(), GatewayError> {
self.candidates.skip_provider(provider_id);
Ok(())
}
}
impl LocalStandardSyncAttemptSource<'_> {
@@ -138,6 +138,13 @@ impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalOpenAiChatStreamAttem
}
Ok(drained)
}
async fn skip_provider(&mut self, provider_id: &str) -> Result<(), GatewayError> {
self.prefetched_attempts
.retain(|attempt| attempt.eligible.candidate.provider_id != provider_id);
self.candidates.skip_provider(provider_id);
Ok(())
}
}
impl LocalOpenAiChatStreamAttemptSource<'_> {
@@ -116,6 +116,11 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalOpenAiChatSyncAttemptSo
}
Ok(drained)
}
async fn skip_provider(&mut self, provider_id: &str) -> Result<(), GatewayError> {
self.candidates.skip_provider(provider_id);
Ok(())
}
}
impl LocalOpenAiChatSyncAttemptSource<'_> {
@@ -185,6 +185,11 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalOpenAiResponsesSyncAtte
}
Ok(drained)
}
async fn skip_provider(&mut self, provider_id: &str) -> Result<(), GatewayError> {
self.candidates.skip_provider(provider_id);
Ok(())
}
}
#[async_trait]
@@ -213,6 +218,11 @@ impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalOpenAiResponsesStream
}
Ok(drained)
}
async fn skip_provider(&mut self, provider_id: &str) -> Result<(), GatewayError> {
self.candidates.skip_provider(provider_id);
Ok(())
}
}
impl LocalOpenAiResponsesSyncAttemptSource<'_> {
@@ -371,6 +371,10 @@ pub(crate) struct PoolKeyCursor<'a> {
}
impl<'a> PoolKeyCursor<'a> {
pub(crate) fn provider_id(&self) -> &str {
self.group.candidate.provider_id.as_str()
}
pub(crate) fn new(
state: PlannerAppState<'a>,
group: EligibleLocalExecutionCandidate,
@@ -936,6 +936,8 @@ mod tests {
policy,
LocalFailoverPolicy {
max_retries: Some(1),
max_transfer_count: 0,
max_transfer_timeout_seconds: 0,
stop_status_codes: [503].into_iter().collect(),
continue_status_codes: [409, 429].into_iter().collect(),
success_failover_patterns: Vec::new(),
@@ -1,3 +1,5 @@
use std::collections::{BTreeMap, BTreeSet};
use aether_ai_serving::{
run_ai_attempt_loop, AiAttemptLoopOutcome, AiAttemptLoopPort, AiExecutionAttempt,
};
@@ -10,7 +12,7 @@ use async_trait::async_trait;
use axum::body::Body;
use axum::http::Response;
use futures_util::StreamExt;
use tokio::time::{timeout, Duration};
use tokio::time::{timeout, Duration, Instant};
use tracing::{debug, warn, Instrument};
use crate::ai_serving::LocalExecutionAttemptSource;
@@ -20,7 +22,10 @@ use crate::execution_runtime::{execute_execution_runtime_stream, execute_executi
use crate::executor::{build_local_execution_exhaustion, LocalExecutionRequestOutcome};
use crate::handlers::shared::provider_pool::release_admin_provider_pool_key_lease;
use crate::log_ids::short_request_id;
use crate::orchestration::local_execution_candidate_metadata_from_report_context;
use crate::orchestration::{
local_execution_candidate_metadata_from_report_context,
local_failover_policy_from_report_context, resolve_local_failover_policy, LocalFailoverPolicy,
};
use crate::privacy::RedactionExecutionCandidateId;
use crate::request_candidate_runtime::{
record_local_request_candidate_status, RequestCandidateRuntimeWriter,
@@ -55,6 +60,31 @@ pub(crate) async fn execute_sync_plan_and_reports<T>(
plan_kind: &str,
plan_and_reports: Vec<T>,
) -> Result<LocalExecutionRequestOutcome, GatewayError>
where
T: AiExecutionAttempt + Send + Sync + 'static,
{
let transfer_tracker = ProviderTransferTracker::default();
execute_sync_plan_and_reports_with_transfer_tracker(
state,
parts,
trace_id,
decision,
plan_kind,
plan_and_reports,
&transfer_tracker,
)
.await
}
pub(crate) async fn execute_sync_plan_and_reports_with_transfer_tracker<T>(
state: &AppState,
parts: &http::request::Parts,
trace_id: &str,
decision: &GatewayControlDecision,
plan_kind: &str,
plan_and_reports: Vec<T>,
transfer_tracker: &ProviderTransferTracker,
) -> Result<LocalExecutionRequestOutcome, GatewayError>
where
T: AiExecutionAttempt + Send + Sync + 'static,
{
@@ -88,6 +118,7 @@ where
trace_id,
decision,
plan_kind,
transfer_tracker,
};
match run_ai_attempt_loop(&port, plan_and_reports).await? {
AiAttemptLoopOutcome::Responded(response) => {
@@ -104,12 +135,38 @@ where
}
pub(crate) async fn execute_sync_attempt_source<T, S>(
state: &AppState,
parts: &http::request::Parts,
trace_id: &str,
decision: &GatewayControlDecision,
plan_kind: &str,
source: S,
) -> Result<LocalExecutionRequestOutcome, GatewayError>
where
T: AiExecutionAttempt + Send + Sync + 'static,
S: LocalExecutionAttemptSource<T>,
{
let transfer_tracker = ProviderTransferTracker::default();
execute_sync_attempt_source_with_transfer_tracker(
state,
parts,
trace_id,
decision,
plan_kind,
source,
&transfer_tracker,
)
.await
}
pub(crate) async fn execute_sync_attempt_source_with_transfer_tracker<T, S>(
state: &AppState,
parts: &http::request::Parts,
trace_id: &str,
decision: &GatewayControlDecision,
plan_kind: &str,
mut source: S,
transfer_tracker: &ProviderTransferTracker,
) -> Result<LocalExecutionRequestOutcome, GatewayError>
where
T: AiExecutionAttempt + Send + Sync + 'static,
@@ -132,6 +189,7 @@ where
trace_id,
decision,
plan_kind,
transfer_tracker,
};
run_dynamic_attempt_loop(
&port,
@@ -154,6 +212,7 @@ struct SyncAttemptLoopPort<'a> {
trace_id: &'a str,
decision: &'a GatewayControlDecision,
plan_kind: &'a str,
transfer_tracker: &'a ProviderTransferTracker,
}
#[async_trait]
@@ -165,6 +224,33 @@ where
type Exhaustion = crate::executor::LocalExecutionExhaustion;
type Error = GatewayError;
async fn should_skip_attempt(&self, attempt: &T) -> Result<bool, Self::Error> {
Ok(should_skip_provider_transfer_attempt(
self.transfer_tracker,
self.trace_id,
self.plan_kind,
attempt,
)
.await)
}
async fn record_attempt_started(&self, attempt: &T) -> Result<(), Self::Error> {
record_provider_transfer_attempt_started(self.transfer_tracker, attempt).await;
Ok(())
}
async fn record_attempt_failed(&self, attempt: &T) -> Result<(), Self::Error> {
record_provider_transfer_attempt_failed(
self.state,
self.transfer_tracker,
self.trace_id,
self.plan_kind,
attempt,
)
.await;
Ok(())
}
async fn execute_attempt(&self, attempt: &T) -> Result<Option<Self::Response>, Self::Error> {
let plan = attempt.execution_plan();
let report_context = attempt.report_context();
@@ -242,6 +328,29 @@ pub(crate) async fn execute_stream_plan_and_reports<T>(
plan_kind: &str,
plan_and_reports: Vec<T>,
) -> Result<LocalExecutionRequestOutcome, GatewayError>
where
T: AiExecutionAttempt + Send + Sync + 'static,
{
let transfer_tracker = ProviderTransferTracker::default();
execute_stream_plan_and_reports_with_transfer_tracker(
state,
trace_id,
decision,
plan_kind,
plan_and_reports,
&transfer_tracker,
)
.await
}
pub(crate) async fn execute_stream_plan_and_reports_with_transfer_tracker<T>(
state: &AppState,
trace_id: &str,
decision: &GatewayControlDecision,
plan_kind: &str,
plan_and_reports: Vec<T>,
transfer_tracker: &ProviderTransferTracker,
) -> Result<LocalExecutionRequestOutcome, GatewayError>
where
T: AiExecutionAttempt + Send + Sync + 'static,
{
@@ -274,6 +383,7 @@ where
trace_id,
decision,
plan_kind,
transfer_tracker,
};
match run_ai_attempt_loop(&port, plan_and_reports).await? {
AiAttemptLoopOutcome::Responded(response) => {
@@ -290,11 +400,35 @@ where
}
pub(crate) async fn execute_stream_attempt_source<T, S>(
state: &AppState,
trace_id: &str,
decision: &GatewayControlDecision,
plan_kind: &str,
source: S,
) -> Result<LocalExecutionRequestOutcome, GatewayError>
where
T: AiExecutionAttempt + Send + Sync + 'static,
S: LocalExecutionAttemptSource<T>,
{
let transfer_tracker = ProviderTransferTracker::default();
execute_stream_attempt_source_with_transfer_tracker(
state,
trace_id,
decision,
plan_kind,
source,
&transfer_tracker,
)
.await
}
pub(crate) async fn execute_stream_attempt_source_with_transfer_tracker<T, S>(
state: &AppState,
trace_id: &str,
decision: &GatewayControlDecision,
plan_kind: &str,
mut source: S,
transfer_tracker: &ProviderTransferTracker,
) -> Result<LocalExecutionRequestOutcome, GatewayError>
where
T: AiExecutionAttempt + Send + Sync + 'static,
@@ -316,6 +450,7 @@ where
trace_id,
decision,
plan_kind,
transfer_tracker,
};
run_dynamic_attempt_loop(
&port,
@@ -332,6 +467,281 @@ where
.await
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
struct ProviderTransferLimits {
max_transfer_count: u64,
max_transfer_timeout_seconds: u64,
}
impl From<&LocalFailoverPolicy> for ProviderTransferLimits {
fn from(policy: &LocalFailoverPolicy) -> Self {
Self {
max_transfer_count: policy.max_transfer_count,
max_transfer_timeout_seconds: policy.max_transfer_timeout_seconds,
}
}
}
#[derive(Debug)]
struct ProviderTransferState {
first_attempt_started_at: Instant,
last_key_id: String,
transfer_count: u64,
limits: Option<ProviderTransferLimits>,
}
#[derive(Debug, Default)]
struct ProviderTransferStateTracker {
by_provider: BTreeMap<String, ProviderTransferState>,
exhausted_provider_ids: BTreeSet<String>,
}
#[derive(Clone, Debug, Default)]
pub(crate) struct ProviderTransferTracker {
state: std::sync::Arc<tokio::sync::Mutex<ProviderTransferStateTracker>>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
struct ProviderTransferLimitReached {
provider_id: String,
transfer_count: u64,
elapsed_ms: u64,
limits: ProviderTransferLimits,
count_reached: bool,
timeout_reached: bool,
}
impl ProviderTransferStateTracker {
fn record_attempt_started(&mut self, plan: &aether_contracts::ExecutionPlan, now: Instant) {
match self.by_provider.entry(plan.provider_id.clone()) {
std::collections::btree_map::Entry::Vacant(entry) => {
entry.insert(ProviderTransferState {
first_attempt_started_at: now,
last_key_id: plan.key_id.clone(),
transfer_count: 0,
limits: None,
});
}
std::collections::btree_map::Entry::Occupied(mut entry) => {
let state = entry.get_mut();
if state.last_key_id != plan.key_id {
state.transfer_count = state.transfer_count.saturating_add(1);
state.last_key_id.clone_from(&plan.key_id);
}
}
}
}
fn needs_limits(&self, provider_id: &str) -> bool {
self.by_provider
.get(provider_id)
.is_some_and(|state| state.limits.is_none())
}
fn set_limits(&mut self, provider_id: &str, limits: ProviderTransferLimits) {
if let Some(state) = self.by_provider.get_mut(provider_id) {
state.limits = Some(limits);
}
}
fn check_before_attempt(
&mut self,
plan: &aether_contracts::ExecutionPlan,
now: Instant,
) -> Option<ProviderTransferLimitReached> {
if self.exhausted_provider_ids.contains(&plan.provider_id) {
return Some(self.reached_snapshot(plan.provider_id.as_str(), now, false, false)?);
}
let state = self.by_provider.get(&plan.provider_id)?;
let limits = state.limits?;
let elapsed = now.saturating_duration_since(state.first_attempt_started_at);
let timeout_reached = limits.max_transfer_timeout_seconds > 0
&& elapsed >= Duration::from_secs(limits.max_transfer_timeout_seconds);
let count_reached = state.last_key_id != plan.key_id
&& limits.max_transfer_count > 0
&& state.transfer_count >= limits.max_transfer_count;
if !count_reached && !timeout_reached {
return None;
}
let reached = self.reached_snapshot(
plan.provider_id.as_str(),
now,
count_reached,
timeout_reached,
)?;
self.exhausted_provider_ids.insert(plan.provider_id.clone());
Some(reached)
}
fn check_timeout_after_failure(
&mut self,
provider_id: &str,
now: Instant,
) -> Option<ProviderTransferLimitReached> {
if self.exhausted_provider_ids.contains(provider_id) {
return None;
}
let state = self.by_provider.get(provider_id)?;
let limits = state.limits?;
let elapsed = now.saturating_duration_since(state.first_attempt_started_at);
let timeout_reached = limits.max_transfer_timeout_seconds > 0
&& elapsed >= Duration::from_secs(limits.max_transfer_timeout_seconds);
if !timeout_reached {
return None;
}
let reached = self.reached_snapshot(provider_id, now, false, true)?;
self.exhausted_provider_ids.insert(provider_id.to_string());
Some(reached)
}
fn reached_snapshot(
&self,
provider_id: &str,
now: Instant,
count_reached: bool,
timeout_reached: bool,
) -> Option<ProviderTransferLimitReached> {
let state = self.by_provider.get(provider_id)?;
let limits = state.limits?;
let elapsed = now.saturating_duration_since(state.first_attempt_started_at);
Some(ProviderTransferLimitReached {
provider_id: provider_id.to_string(),
transfer_count: state.transfer_count,
elapsed_ms: elapsed.as_millis().min(u128::from(u64::MAX)) as u64,
limits,
count_reached,
timeout_reached,
})
}
}
async fn load_provider_transfer_limits<Attempt>(
state: &AppState,
tracker: &mut ProviderTransferStateTracker,
attempt: &Attempt,
) where
Attempt: AiExecutionAttempt + Send + Sync + 'static,
{
let plan = attempt.execution_plan();
if !tracker.needs_limits(plan.provider_id.as_str()) {
return;
}
let owned_report_context = if attempt.report_context_ref().is_none() {
attempt.report_context()
} else {
None
};
let report_context = attempt
.report_context_ref()
.or(owned_report_context.as_ref());
let embedded_policy_has_transfer_limits = report_context
.and_then(serde_json::Value::as_object)
.and_then(|object| object.get("local_failover_policy"))
.and_then(serde_json::Value::as_object)
.is_some_and(|policy| {
policy.contains_key("max_transfer_count")
|| policy.contains_key("max_transfer_timeout_seconds")
});
let policy = if embedded_policy_has_transfer_limits {
local_failover_policy_from_report_context(report_context).unwrap_or_default()
} else {
resolve_local_failover_policy(state, plan, report_context).await
};
tracker.set_limits(
plan.provider_id.as_str(),
ProviderTransferLimits::from(&policy),
);
}
async fn provider_transfer_timeout_after_failure<Attempt>(
state: &AppState,
tracker: &mut ProviderTransferStateTracker,
attempt: &Attempt,
) -> Option<ProviderTransferLimitReached>
where
Attempt: AiExecutionAttempt + Send + Sync + 'static,
{
let plan = attempt.execution_plan();
load_provider_transfer_limits(state, tracker, attempt).await;
tracker.check_timeout_after_failure(plan.provider_id.as_str(), Instant::now())
}
fn log_provider_transfer_limit_reached(
trace_id: &str,
plan_kind: &str,
reached: &ProviderTransferLimitReached,
) {
warn!(
event_name = "provider_transfer_limit_reached",
log_type = "event",
trace_id,
plan_kind,
provider_id = %reached.provider_id,
transfer_count = reached.transfer_count,
elapsed_ms = reached.elapsed_ms,
max_transfer_count = reached.limits.max_transfer_count,
max_transfer_timeout_seconds = reached.limits.max_transfer_timeout_seconds,
count_reached = reached.count_reached,
timeout_reached = reached.timeout_reached,
"gateway exhausted the provider transfer budget and will skip its remaining candidates"
);
}
async fn should_skip_provider_transfer_attempt<Attempt>(
tracker: &ProviderTransferTracker,
trace_id: &str,
plan_kind: &str,
attempt: &Attempt,
) -> bool
where
Attempt: AiExecutionAttempt + Send + Sync + 'static,
{
let reached = tracker
.state
.lock()
.await
.check_before_attempt(attempt.execution_plan(), Instant::now());
let Some(reached) = reached else {
return false;
};
if reached.count_reached || reached.timeout_reached {
log_provider_transfer_limit_reached(trace_id, plan_kind, &reached);
}
true
}
async fn record_provider_transfer_attempt_started<Attempt>(
tracker: &ProviderTransferTracker,
attempt: &Attempt,
) where
Attempt: AiExecutionAttempt + Send + Sync + 'static,
{
tracker
.state
.lock()
.await
.record_attempt_started(attempt.execution_plan(), Instant::now());
}
async fn record_provider_transfer_attempt_failed<Attempt>(
state: &AppState,
tracker: &ProviderTransferTracker,
trace_id: &str,
plan_kind: &str,
attempt: &Attempt,
) where
Attempt: AiExecutionAttempt + Send + Sync + 'static,
{
let mut tracker = tracker.state.lock().await;
let reached = provider_transfer_timeout_after_failure(state, &mut tracker, attempt).await;
if let Some(reached) = reached {
log_provider_transfer_limit_reached(trace_id, plan_kind, &reached);
}
}
async fn run_dynamic_attempt_loop<Port, Source, Attempt>(
port: &Port,
source: &mut Source,
@@ -363,6 +773,13 @@ where
let Some(attempt) = next_attempt else {
break;
};
if port.should_skip_attempt(&attempt).await? {
let provider_id = attempt.execution_plan().provider_id.clone();
port.mark_unused_attempts(vec![attempt]).await?;
source.skip_provider(provider_id.as_str()).await?;
continue;
}
port.record_attempt_started(&attempt).await?;
let execute_started_at = std::time::Instant::now();
let response = match port.execute_attempt(&attempt).await {
Ok(response) => response,
@@ -387,6 +804,13 @@ where
return Ok(LocalExecutionRequestOutcome::responded(response));
}
port.record_attempt_failed(&attempt).await?;
if port.should_skip_attempt(&attempt).await? {
source
.skip_provider(attempt.execution_plan().provider_id.as_str())
.await?;
}
// Only retain a deep plan/context snapshot when this candidate really
// failed and exhaustion reporting will need it.
last_attempted = Some((attempt.execution_plan().clone(), attempt.report_context()));
@@ -438,6 +862,7 @@ struct StreamAttemptLoopPort<'a> {
trace_id: &'a str,
decision: &'a GatewayControlDecision,
plan_kind: &'a str,
transfer_tracker: &'a ProviderTransferTracker,
}
#[async_trait]
@@ -449,6 +874,33 @@ where
type Exhaustion = crate::executor::LocalExecutionExhaustion;
type Error = GatewayError;
async fn should_skip_attempt(&self, attempt: &T) -> Result<bool, Self::Error> {
Ok(should_skip_provider_transfer_attempt(
self.transfer_tracker,
self.trace_id,
self.plan_kind,
attempt,
)
.await)
}
async fn record_attempt_started(&self, attempt: &T) -> Result<(), Self::Error> {
record_provider_transfer_attempt_started(self.transfer_tracker, attempt).await;
Ok(())
}
async fn record_attempt_failed(&self, attempt: &T) -> Result<(), Self::Error> {
record_provider_transfer_attempt_failed(
self.state,
self.transfer_tracker,
self.trace_id,
self.plan_kind,
attempt,
)
.await;
Ok(())
}
async fn execute_attempt(&self, attempt: &T) -> Result<Option<Self::Response>, Self::Error> {
let plan = attempt.execution_plan();
let report_context = attempt.report_context();
@@ -1071,7 +1523,7 @@ pub(crate) async fn mark_unused_local_candidate_items<T, FPlan, FContext>(
#[cfg(test)]
mod tests {
use std::sync::Arc;
use std::sync::{Arc, Mutex as StdMutex};
use aether_contracts::{ExecutionPlan, ExecutionTimeouts, RequestBody};
use aether_data_contracts::repository::candidates::{
@@ -1152,6 +1604,365 @@ mod tests {
async fn drain_execution_attempts(&mut self) -> Result<Vec<()>, GatewayError> {
Ok(Vec::new())
}
async fn skip_provider(&mut self, _provider_id: &str) -> Result<(), GatewayError> {
Ok(())
}
}
#[derive(Clone)]
struct TransferTestAttempt {
label: &'static str,
plan: ExecutionPlan,
report_context: serde_json::Value,
}
impl AiExecutionAttempt for TransferTestAttempt {
fn execution_plan(&self) -> &ExecutionPlan {
&self.plan
}
fn report_kind(&self) -> Option<String> {
None
}
fn report_context(&self) -> Option<serde_json::Value> {
Some(self.report_context.clone())
}
fn report_context_ref(&self) -> Option<&serde_json::Value> {
Some(&self.report_context)
}
}
struct TransferTestPort<'a> {
state: &'a AppState,
tracker: ProviderTransferTracker,
executed: StdMutex<Vec<&'static str>>,
unused: StdMutex<Vec<&'static str>>,
}
impl<'a> TransferTestPort<'a> {
fn new(state: &'a AppState) -> Self {
Self::with_tracker(state, ProviderTransferTracker::default())
}
fn with_tracker(state: &'a AppState, tracker: ProviderTransferTracker) -> Self {
Self {
state,
tracker,
executed: StdMutex::new(Vec::new()),
unused: StdMutex::new(Vec::new()),
}
}
}
#[async_trait]
impl AiAttemptLoopPort<TransferTestAttempt> for TransferTestPort<'_> {
type Response = Response<Body>;
type Exhaustion = crate::executor::LocalExecutionExhaustion;
type Error = GatewayError;
async fn should_skip_attempt(
&self,
attempt: &TransferTestAttempt,
) -> Result<bool, Self::Error> {
Ok(should_skip_provider_transfer_attempt(
&self.tracker,
"trace-transfer-test",
"transfer_test",
attempt,
)
.await)
}
async fn record_attempt_started(
&self,
attempt: &TransferTestAttempt,
) -> Result<(), Self::Error> {
record_provider_transfer_attempt_started(&self.tracker, attempt).await;
Ok(())
}
async fn record_attempt_failed(
&self,
attempt: &TransferTestAttempt,
) -> Result<(), Self::Error> {
record_provider_transfer_attempt_failed(
self.state,
&self.tracker,
"trace-transfer-test",
"transfer_test",
attempt,
)
.await;
Ok(())
}
async fn execute_attempt(
&self,
attempt: &TransferTestAttempt,
) -> Result<Option<Self::Response>, Self::Error> {
self.executed.lock().unwrap().push(attempt.label);
Ok((attempt.plan.provider_id == "provider-b").then(|| Response::new(Body::from("ok"))))
}
async fn mark_unused_attempts(
&self,
attempts: Vec<TransferTestAttempt>,
) -> Result<(), Self::Error> {
self.unused
.lock()
.unwrap()
.extend(attempts.into_iter().map(|attempt| attempt.label));
Ok(())
}
async fn build_exhaustion(
&self,
last_plan: ExecutionPlan,
last_report_context: Option<serde_json::Value>,
) -> Result<Self::Exhaustion, Self::Error> {
Ok(build_local_execution_exhaustion(
self.state,
&last_plan,
last_report_context.as_ref(),
)
.await)
}
}
struct TransferTestAttemptSource {
attempts: std::collections::VecDeque<TransferTestAttempt>,
skipped_providers: Vec<String>,
}
#[async_trait]
impl LocalExecutionAttemptSource<TransferTestAttempt> for TransferTestAttemptSource {
async fn next_execution_attempt(
&mut self,
) -> Result<Option<TransferTestAttempt>, GatewayError> {
Ok(self.attempts.pop_front())
}
async fn drain_execution_attempts(
&mut self,
) -> Result<Vec<TransferTestAttempt>, GatewayError> {
Ok(self.attempts.drain(..).collect())
}
async fn skip_provider(&mut self, provider_id: &str) -> Result<(), GatewayError> {
self.skipped_providers.push(provider_id.to_string());
self.attempts
.retain(|attempt| attempt.plan.provider_id != provider_id);
Ok(())
}
}
fn transfer_test_attempts() -> Vec<TransferTestAttempt> {
fn attempt(label: &'static str, provider_id: &str, key_id: &str) -> TransferTestAttempt {
let mut plan = test_plan(None);
plan.provider_id = provider_id.to_string();
plan.key_id = key_id.to_string();
TransferTestAttempt {
label,
plan,
report_context: json!({
"local_failover_policy": {
"max_transfer_count": 1,
"max_transfer_timeout_seconds": 0
}
}),
}
}
vec![
attempt("a-key1-retry0", "provider-a", "key-1"),
attempt("a-key2-retry0", "provider-a", "key-2"),
attempt("a-key2-retry1", "provider-a", "key-2"),
attempt("a-key3-retry0", "provider-a", "key-3"),
attempt("b-key1-retry0", "provider-b", "key-b"),
]
}
#[tokio::test]
async fn static_loop_allows_same_key_retries_then_skips_next_transfer() {
let state = AppState::new().expect("state should build");
let port = TransferTestPort::new(&state);
let outcome = run_ai_attempt_loop(&port, transfer_test_attempts())
.await
.expect("attempt loop should succeed");
assert!(matches!(outcome, AiAttemptLoopOutcome::Responded(_)));
assert_eq!(
port.executed.lock().unwrap().as_slice(),
[
"a-key1-retry0",
"a-key2-retry0",
"a-key2-retry1",
"b-key1-retry0"
]
);
assert_eq!(port.unused.lock().unwrap().as_slice(), ["a-key3-retry0"]);
}
#[tokio::test]
async fn cloned_tracker_preserves_transfer_budget_across_candidate_loops() {
let state = AppState::new().expect("state should build");
let tracker = ProviderTransferTracker::default();
let mut attempts = transfer_test_attempts();
let provider_b = attempts.pop().expect("provider-b attempt should exist");
let key_3 = attempts.pop().expect("third provider-a key should exist");
let first_port = TransferTestPort::with_tracker(&state, tracker.clone());
let first_outcome = run_ai_attempt_loop(&first_port, attempts)
.await
.expect("first candidate loop should exhaust");
assert!(matches!(first_outcome, AiAttemptLoopOutcome::Exhausted(_)));
let second_port = TransferTestPort::with_tracker(&state, tracker);
let second_outcome = run_ai_attempt_loop(&second_port, vec![key_3, provider_b])
.await
.expect("second candidate loop should succeed");
assert!(matches!(second_outcome, AiAttemptLoopOutcome::Responded(_)));
assert_eq!(
second_port.executed.lock().unwrap().as_slice(),
["b-key1-retry0"]
);
assert_eq!(
second_port.unused.lock().unwrap().as_slice(),
["a-key3-retry0"]
);
}
#[tokio::test]
async fn dynamic_loop_skips_exhausted_provider_at_candidate_source() {
let state = AppState::new().expect("state should build");
let port = TransferTestPort::new(&state);
let mut source = TransferTestAttemptSource {
attempts: transfer_test_attempts().into(),
skipped_providers: Vec::new(),
};
let outcome = run_dynamic_attempt_loop(
&port,
&mut source,
"trace-transfer-test",
"transfer_test",
Duration::from_secs(1),
)
.await
.expect("dynamic attempt loop should succeed");
assert!(matches!(
outcome,
LocalExecutionRequestOutcome::Responded(_)
));
assert_eq!(
port.executed.lock().unwrap().as_slice(),
[
"a-key1-retry0",
"a-key2-retry0",
"a-key2-retry1",
"b-key1-retry0"
]
);
assert_eq!(source.skipped_providers, ["provider-a"]);
}
#[test]
fn transfer_timeout_is_checked_at_candidate_boundary_and_zero_disables_limits() {
let started_at = Instant::now();
let mut first = test_plan(None);
first.provider_id = "provider-a".to_string();
first.key_id = "key-1".to_string();
let mut timeout_tracker = ProviderTransferStateTracker::default();
timeout_tracker.record_attempt_started(&first, started_at);
timeout_tracker.set_limits(
"provider-a",
ProviderTransferLimits {
max_transfer_count: 0,
max_transfer_timeout_seconds: 60,
},
);
assert!(timeout_tracker
.check_before_attempt(&first, started_at + Duration::from_secs(59))
.is_none());
let reached = timeout_tracker
.check_before_attempt(&first, started_at + Duration::from_secs(60))
.expect("timeout should stop the provider at the next candidate boundary");
assert!(reached.timeout_reached);
assert!(!reached.count_reached);
let mut count_tracker = ProviderTransferStateTracker::default();
count_tracker.record_attempt_started(&first, started_at);
count_tracker.set_limits(
"provider-a",
ProviderTransferLimits {
max_transfer_count: 1,
max_transfer_timeout_seconds: 60,
},
);
let mut second_key = first.clone();
second_key.key_id = "key-2".to_string();
assert!(count_tracker
.check_before_attempt(&second_key, started_at + Duration::from_secs(1))
.is_none());
count_tracker.record_attempt_started(&second_key, started_at + Duration::from_secs(1));
let mut third_key = first.clone();
third_key.key_id = "key-3".to_string();
let reached = count_tracker
.check_before_attempt(&third_key, started_at + Duration::from_secs(2))
.expect("count should stop the provider before another key transfer");
assert!(reached.count_reached);
assert!(!reached.timeout_reached);
let mut unlimited_tracker = ProviderTransferStateTracker::default();
unlimited_tracker.record_attempt_started(&first, started_at);
unlimited_tracker.set_limits("provider-a", ProviderTransferLimits::default());
let mut another_key = first.clone();
another_key.key_id = "key-2".to_string();
assert!(unlimited_tracker
.check_before_attempt(&another_key, started_at + Duration::from_secs(3_600))
.is_none());
}
#[test]
fn transfer_count_and_timeout_limits_use_or_semantics() {
let started_at = Instant::now();
let mut first = test_plan(None);
first.provider_id = "provider-a".to_string();
first.key_id = "key-1".to_string();
let limits = ProviderTransferLimits {
max_transfer_count: 1,
max_transfer_timeout_seconds: 60,
};
let mut count_first = ProviderTransferStateTracker::default();
count_first.record_attempt_started(&first, started_at);
count_first.set_limits("provider-a", limits);
let mut second = first.clone();
second.key_id = "key-2".to_string();
count_first.record_attempt_started(&second, started_at + Duration::from_secs(1));
let mut third = first.clone();
third.key_id = "key-3".to_string();
let count_reached = count_first
.check_before_attempt(&third, started_at + Duration::from_secs(2))
.expect("count should independently exhaust a provider before timeout");
assert!(count_reached.count_reached);
assert!(!count_reached.timeout_reached);
let mut timeout_first = ProviderTransferStateTracker::default();
timeout_first.record_attempt_started(&first, started_at);
timeout_first.set_limits("provider-a", limits);
let timeout_reached = timeout_first
.check_before_attempt(&first, started_at + Duration::from_secs(60))
.expect("timeout should independently exhaust a provider before count");
assert!(!timeout_reached.count_reached);
assert!(timeout_reached.timeout_reached);
}
fn test_plan(timeouts: Option<ExecutionTimeouts>) -> ExecutionPlan {
+3 -2
View File
@@ -11,8 +11,9 @@ pub(crate) use crate::request_candidate_runtime::{
persist_available_local_candidate, persist_skipped_local_candidate,
};
pub(crate) use candidate_loop::{
execute_stream_plan_and_reports, execute_sync_plan_and_reports,
mark_unused_local_candidate_items,
execute_stream_plan_and_reports, execute_stream_plan_and_reports_with_transfer_tracker,
execute_sync_plan_and_reports, execute_sync_plan_and_reports_with_transfer_tracker,
mark_unused_local_candidate_items, ProviderTransferTracker,
};
pub(crate) use orchestration::*;
pub(crate) use outcome::{
+203 -95
View File
@@ -43,15 +43,16 @@ use crate::constants::{CONTROL_CANDIDATE_ID_HEADER, EXECUTION_PATH_LOCAL_EXECUTI
use crate::control::GatewayControlDecision;
use crate::execution_runtime::sync::{
build_openai_image_sync_json_whitespace_heartbeat_stream,
build_sync_json_whitespace_heartbeat_stream, execute_execution_runtime_sync,
build_sync_json_whitespace_heartbeat_stream,
};
use crate::executor::candidate_loop::{
execute_stream_attempt_source, execute_sync_attempt_source, execute_sync_plan_and_reports,
mark_unused_local_candidates,
execute_stream_attempt_source_with_transfer_tracker, execute_sync_attempt_source,
execute_sync_attempt_source_with_transfer_tracker,
execute_sync_plan_and_reports_with_transfer_tracker, ProviderTransferTracker,
};
use crate::executor::{
build_local_execution_exhaustion, record_failed_usage_for_exhausted_request,
LocalExecutionExhaustion, LocalExecutionRequestOutcome,
record_failed_usage_for_exhausted_request, LocalExecutionExhaustion,
LocalExecutionRequestOutcome,
};
use crate::handlers::shared::system_config_bool;
use crate::stage_metrics::observe_gateway_stage_ms;
@@ -94,6 +95,7 @@ pub(crate) async fn maybe_execute_sync_via_local_decision(
decision: &GatewayControlDecision,
body_json: &serde_json::Value,
plan_kind: &str,
transfer_tracker: &ProviderTransferTracker,
) -> Result<LocalExecutionRequestOutcome, GatewayError> {
let Some((attempt_source, candidate_count)) =
build_local_openai_chat_sync_attempt_source_for_kind(
@@ -107,6 +109,7 @@ pub(crate) async fn maybe_execute_sync_via_local_decision(
if standard_text_sync_heartbeat_should_wrap(state, plan_kind).await {
let parts_for_task = parts.clone();
let body_json_for_task = body_json.clone();
let transfer_tracker_for_task = transfer_tracker.clone();
return Ok(LocalExecutionRequestOutcome::responded(
build_standard_text_sync_heartbeat_shell_response(
state.clone(),
@@ -129,15 +132,17 @@ pub(crate) async fn maybe_execute_sync_via_local_decision(
return Ok(LocalExecutionRequestOutcome::NoPath);
};
let outcome = execute_sync_attempt_source::<AiSyncAttempt, _>(
&state,
&parts,
trace_id.as_str(),
&decision,
plan_kind.as_str(),
attempt_source,
)
.await?;
let outcome =
execute_sync_attempt_source_with_transfer_tracker::<AiSyncAttempt, _>(
&state,
&parts,
trace_id.as_str(),
&decision,
plan_kind.as_str(),
attempt_source,
&transfer_tracker_for_task,
)
.await?;
match outcome {
LocalExecutionRequestOutcome::Exhausted(exhaustion) => {
set_local_openai_chat_execution_exhausted_diagnostic(
@@ -163,13 +168,14 @@ pub(crate) async fn maybe_execute_sync_via_local_decision(
));
}
let outcome = execute_sync_attempt_source::<AiSyncAttempt, _>(
let outcome = execute_sync_attempt_source_with_transfer_tracker::<AiSyncAttempt, _>(
state,
parts,
trace_id,
decision,
plan_kind,
attempt_source,
transfer_tracker,
)
.await?;
@@ -194,6 +200,7 @@ pub(crate) async fn maybe_execute_stream_via_local_decision(
decision: &GatewayControlDecision,
body_json: &serde_json::Value,
plan_kind: &str,
transfer_tracker: &ProviderTransferTracker,
) -> Result<LocalExecutionRequestOutcome, GatewayError> {
let attempt_source_started_at = std::time::Instant::now();
let attempt_source = build_local_openai_chat_stream_attempt_source_for_kind(
@@ -209,12 +216,13 @@ pub(crate) async fn maybe_execute_stream_via_local_decision(
};
let attempt_source_execute_started_at = std::time::Instant::now();
let outcome = execute_stream_attempt_source::<AiStreamAttempt, _>(
let outcome = execute_stream_attempt_source_with_transfer_tracker::<AiStreamAttempt, _>(
state,
trace_id,
decision,
plan_kind,
attempt_source,
transfer_tracker,
)
.await;
observe_gateway_stage_ms(
@@ -244,6 +252,7 @@ pub(crate) async fn maybe_execute_sync_via_local_openai_responses_decision(
decision: &GatewayControlDecision,
body_json: &serde_json::Value,
plan_kind: &str,
transfer_tracker: &ProviderTransferTracker,
) -> Result<LocalExecutionRequestOutcome, GatewayError> {
let Some((attempt_source, _candidate_count)) =
build_local_openai_responses_sync_attempt_source_for_kind(
@@ -257,6 +266,7 @@ pub(crate) async fn maybe_execute_sync_via_local_openai_responses_decision(
if standard_text_sync_heartbeat_should_wrap(state, plan_kind).await {
let parts_for_task = parts.clone();
let body_json_for_task = body_json.clone();
let transfer_tracker_for_task = transfer_tracker.clone();
return Ok(LocalExecutionRequestOutcome::responded(
build_standard_text_sync_heartbeat_shell_response(
state.clone(),
@@ -279,15 +289,17 @@ pub(crate) async fn maybe_execute_sync_via_local_openai_responses_decision(
return Ok(LocalExecutionRequestOutcome::NoPath);
};
let outcome = execute_sync_attempt_source::<AiSyncAttempt, _>(
&state,
&parts,
trace_id.as_str(),
&decision,
plan_kind.as_str(),
attempt_source,
)
.await?;
let outcome =
execute_sync_attempt_source_with_transfer_tracker::<AiSyncAttempt, _>(
&state,
&parts,
trace_id.as_str(),
&decision,
plan_kind.as_str(),
attempt_source,
&transfer_tracker_for_task,
)
.await?;
match outcome {
LocalExecutionRequestOutcome::Exhausted(exhaustion) => {
record_standard_text_sync_heartbeat_exhaustion(
@@ -305,13 +317,14 @@ pub(crate) async fn maybe_execute_sync_via_local_openai_responses_decision(
));
}
execute_sync_attempt_source::<AiSyncAttempt, _>(
execute_sync_attempt_source_with_transfer_tracker::<AiSyncAttempt, _>(
state,
parts,
trace_id,
decision,
plan_kind,
attempt_source,
transfer_tracker,
)
.await
}
@@ -323,6 +336,7 @@ pub(crate) async fn maybe_execute_stream_via_local_openai_responses_decision(
decision: &GatewayControlDecision,
body_json: &serde_json::Value,
plan_kind: &str,
transfer_tracker: &ProviderTransferTracker,
) -> Result<LocalExecutionRequestOutcome, GatewayError> {
let Some((attempt_source, _candidate_count)) =
build_local_openai_responses_stream_attempt_source_for_kind(
@@ -333,12 +347,13 @@ pub(crate) async fn maybe_execute_stream_via_local_openai_responses_decision(
return Ok(LocalExecutionRequestOutcome::NoPath);
};
execute_stream_attempt_source::<AiStreamAttempt, _>(
execute_stream_attempt_source_with_transfer_tracker::<AiStreamAttempt, _>(
state,
trace_id,
decision,
plan_kind,
attempt_source,
transfer_tracker,
)
.await
}
@@ -351,6 +366,7 @@ pub(crate) async fn maybe_execute_sync_via_standard_family_decision(
body_json: &serde_json::Value,
plan_kind: &str,
resolve_sync_spec: fn(&str) -> Option<LocalStandardSpec>,
transfer_tracker: &ProviderTransferTracker,
) -> Result<LocalExecutionRequestOutcome, GatewayError> {
let Some(spec) = resolve_sync_spec(plan_kind) else {
return Ok(LocalExecutionRequestOutcome::NoPath);
@@ -367,6 +383,7 @@ pub(crate) async fn maybe_execute_sync_via_standard_family_decision(
if standard_text_sync_heartbeat_should_wrap(state, plan_kind).await {
let parts_for_task = parts.clone();
let body_json_for_task = body_json.clone();
let transfer_tracker_for_task = transfer_tracker.clone();
return Ok(LocalExecutionRequestOutcome::responded(
build_standard_text_sync_heartbeat_shell_response(
state.clone(),
@@ -389,15 +406,17 @@ pub(crate) async fn maybe_execute_sync_via_standard_family_decision(
return Ok(LocalExecutionRequestOutcome::NoPath);
};
let outcome = execute_sync_attempt_source::<AiSyncAttempt, _>(
&state,
&parts,
trace_id.as_str(),
&decision,
plan_kind.as_str(),
attempt_source,
)
.await?;
let outcome =
execute_sync_attempt_source_with_transfer_tracker::<AiSyncAttempt, _>(
&state,
&parts,
trace_id.as_str(),
&decision,
plan_kind.as_str(),
attempt_source,
&transfer_tracker_for_task,
)
.await?;
match outcome {
LocalExecutionRequestOutcome::Exhausted(exhaustion) => {
record_standard_text_sync_heartbeat_exhaustion(
@@ -415,13 +434,14 @@ pub(crate) async fn maybe_execute_sync_via_standard_family_decision(
));
}
execute_sync_attempt_source::<AiSyncAttempt, _>(
execute_sync_attempt_source_with_transfer_tracker::<AiSyncAttempt, _>(
state,
parts,
trace_id,
decision,
plan_kind,
attempt_source,
transfer_tracker,
)
.await
}
@@ -434,6 +454,7 @@ pub(crate) async fn maybe_execute_stream_via_standard_family_decision(
body_json: &serde_json::Value,
plan_kind: &str,
resolve_stream_spec: fn(&str) -> Option<LocalStandardSpec>,
transfer_tracker: &ProviderTransferTracker,
) -> Result<LocalExecutionRequestOutcome, GatewayError> {
let Some(spec) = resolve_stream_spec(plan_kind) else {
return Ok(LocalExecutionRequestOutcome::NoPath);
@@ -447,12 +468,13 @@ pub(crate) async fn maybe_execute_stream_via_standard_family_decision(
return Ok(LocalExecutionRequestOutcome::NoPath);
};
execute_stream_attempt_source::<AiStreamAttempt, _>(
execute_stream_attempt_source_with_transfer_tracker::<AiStreamAttempt, _>(
state,
trace_id,
decision,
plan_kind,
attempt_source,
transfer_tracker,
)
.await
}
@@ -464,6 +486,7 @@ pub(crate) async fn maybe_execute_sync_via_local_standard_decision(
decision: &GatewayControlDecision,
body_json: &serde_json::Value,
plan_kind: &str,
transfer_tracker: &ProviderTransferTracker,
) -> Result<LocalExecutionRequestOutcome, GatewayError> {
let mut exhausted = None;
@@ -475,6 +498,7 @@ pub(crate) async fn maybe_execute_sync_via_local_standard_decision(
body_json,
plan_kind,
resolve_claude_sync_spec,
transfer_tracker,
)
.await?
{
@@ -493,6 +517,7 @@ pub(crate) async fn maybe_execute_sync_via_local_standard_decision(
body_json,
plan_kind,
resolve_gemini_sync_spec,
transfer_tracker,
)
.await?
{
@@ -515,6 +540,7 @@ pub(crate) async fn maybe_execute_stream_via_local_standard_decision(
decision: &GatewayControlDecision,
body_json: &serde_json::Value,
plan_kind: &str,
transfer_tracker: &ProviderTransferTracker,
) -> Result<LocalExecutionRequestOutcome, GatewayError> {
let mut exhausted = None;
@@ -526,6 +552,7 @@ pub(crate) async fn maybe_execute_stream_via_local_standard_decision(
body_json,
plan_kind,
resolve_claude_stream_spec,
transfer_tracker,
)
.await?
{
@@ -544,6 +571,7 @@ pub(crate) async fn maybe_execute_stream_via_local_standard_decision(
body_json,
plan_kind,
resolve_gemini_stream_spec,
transfer_tracker,
)
.await?
{
@@ -566,6 +594,7 @@ pub(crate) async fn maybe_execute_sync_via_local_same_format_provider_decision(
decision: &GatewayControlDecision,
body_json: &serde_json::Value,
plan_kind: &str,
transfer_tracker: &ProviderTransferTracker,
) -> Result<LocalExecutionRequestOutcome, GatewayError> {
let Some(spec) = resolve_local_same_format_sync_spec(plan_kind) else {
return Ok(LocalExecutionRequestOutcome::NoPath);
@@ -582,6 +611,7 @@ pub(crate) async fn maybe_execute_sync_via_local_same_format_provider_decision(
if standard_text_sync_heartbeat_should_wrap(state, plan_kind).await {
let parts_for_task = parts.clone();
let body_json_for_task = body_json.clone();
let transfer_tracker_for_task = transfer_tracker.clone();
return Ok(LocalExecutionRequestOutcome::responded(
build_standard_text_sync_heartbeat_shell_response(
state.clone(),
@@ -604,15 +634,17 @@ pub(crate) async fn maybe_execute_sync_via_local_same_format_provider_decision(
return Ok(LocalExecutionRequestOutcome::NoPath);
};
let outcome = execute_sync_attempt_source::<AiSyncAttempt, _>(
&state,
&parts,
trace_id.as_str(),
&decision,
plan_kind.as_str(),
attempt_source,
)
.await?;
let outcome =
execute_sync_attempt_source_with_transfer_tracker::<AiSyncAttempt, _>(
&state,
&parts,
trace_id.as_str(),
&decision,
plan_kind.as_str(),
attempt_source,
&transfer_tracker_for_task,
)
.await?;
match outcome {
LocalExecutionRequestOutcome::Exhausted(exhaustion) => {
record_standard_text_sync_heartbeat_exhaustion(
@@ -630,13 +662,14 @@ pub(crate) async fn maybe_execute_sync_via_local_same_format_provider_decision(
));
}
execute_sync_attempt_source::<AiSyncAttempt, _>(
execute_sync_attempt_source_with_transfer_tracker::<AiSyncAttempt, _>(
state,
parts,
trace_id,
decision,
plan_kind,
attempt_source,
transfer_tracker,
)
.await
}
@@ -648,6 +681,7 @@ pub(crate) async fn maybe_execute_stream_via_local_same_format_provider_decision
decision: &GatewayControlDecision,
body_json: &serde_json::Value,
plan_kind: &str,
transfer_tracker: &ProviderTransferTracker,
) -> Result<LocalExecutionRequestOutcome, GatewayError> {
let Some(spec) = resolve_local_same_format_stream_spec(plan_kind) else {
return Ok(LocalExecutionRequestOutcome::NoPath);
@@ -661,12 +695,13 @@ pub(crate) async fn maybe_execute_stream_via_local_same_format_provider_decision
return Ok(LocalExecutionRequestOutcome::NoPath);
};
execute_stream_attempt_source::<AiStreamAttempt, _>(
execute_stream_attempt_source_with_transfer_tracker::<AiStreamAttempt, _>(
state,
trace_id,
decision,
plan_kind,
attempt_source,
transfer_tracker,
)
.await
}
@@ -680,6 +715,7 @@ pub(crate) async fn maybe_execute_sync_via_local_gemini_files_decision(
trace_id: &str,
decision: &GatewayControlDecision,
plan_kind: &str,
transfer_tracker: &ProviderTransferTracker,
) -> Result<LocalExecutionRequestOutcome, GatewayError> {
let Some((attempt_source, _candidate_count)) =
build_local_gemini_files_sync_attempt_source_for_kind(
@@ -697,13 +733,14 @@ pub(crate) async fn maybe_execute_sync_via_local_gemini_files_decision(
return Ok(LocalExecutionRequestOutcome::NoPath);
};
execute_sync_attempt_source::<AiSyncAttempt, _>(
execute_sync_attempt_source_with_transfer_tracker::<AiSyncAttempt, _>(
state,
parts,
trace_id,
decision,
plan_kind,
attempt_source,
transfer_tracker,
)
.await
}
@@ -1068,6 +1105,7 @@ fn build_openai_image_sync_heartbeat_shell_response(
decision: GatewayControlDecision,
plan_kind: String,
attempts: Vec<AiSyncAttempt>,
transfer_tracker: ProviderTransferTracker,
) -> Result<Response<Body>, GatewayError> {
let request_id = attempts
.first()
@@ -1087,6 +1125,7 @@ fn build_openai_image_sync_heartbeat_shell_response(
decision,
plan_kind,
attempts,
transfer_tracker,
started_at,
)
.await,
@@ -1129,51 +1168,39 @@ async fn execute_openai_image_sync_heartbeat_attempts(
decision: GatewayControlDecision,
plan_kind: String,
attempts: Vec<AiSyncAttempt>,
transfer_tracker: ProviderTransferTracker,
started_at: Instant,
) -> Result<LocalExecutionRequestOutcome, GatewayError> {
let mut attempts = VecDeque::from(attempts);
let mut last_attempted = None;
while let Some(attempt) = attempts.pop_front() {
let plan = attempt.plan;
let report_kind = attempt.report_kind;
let report_context = attempt.report_context;
last_attempted = Some((plan.clone(), report_context.clone()));
match execute_execution_runtime_sync(
&state,
request_path.as_str(),
plan,
trace_id.as_str(),
&decision,
plan_kind.as_str(),
report_kind,
report_context,
)
.await?
{
Some(response) => {
mark_unused_local_candidates(&state, attempts.into_iter().collect()).await;
return Ok(LocalExecutionRequestOutcome::responded(response));
}
None => continue,
}
}
let Some((last_plan, last_report_context)) = last_attempted else {
return Ok(LocalExecutionRequestOutcome::NoPath);
};
let exhaustion =
build_local_execution_exhaustion(&state, &last_plan, last_report_context.as_ref()).await;
record_failed_usage_for_exhausted_request(
let (parts, _) = http::Request::builder()
.uri(request_path.as_str())
.body(())
.map_err(|err| GatewayError::Internal(err.to_string()))?
.into_parts();
match execute_sync_plan_and_reports_with_transfer_tracker(
&state,
exhaustion,
&started_at,
"OpenAI image sync heartbeat exhausted all local candidates",
EXECUTION_PATH_LOCAL_EXECUTION_RUNTIME_MISS,
None,
&parts,
trace_id.as_str(),
&decision,
plan_kind.as_str(),
attempts,
&transfer_tracker,
)
.await;
Ok(LocalExecutionRequestOutcome::NoPath)
.await?
{
LocalExecutionRequestOutcome::Exhausted(exhaustion) => {
record_failed_usage_for_exhausted_request(
&state,
exhaustion,
&started_at,
"OpenAI image sync heartbeat exhausted all local candidates",
EXECUTION_PATH_LOCAL_EXECUTION_RUNTIME_MISS,
None,
)
.await;
Ok(LocalExecutionRequestOutcome::NoPath)
}
outcome => Ok(outcome),
}
}
async fn openai_image_sync_heartbeat_final_bytes(
@@ -1276,6 +1303,7 @@ pub(crate) async fn maybe_execute_sync_via_local_image_decision(
trace_id: &str,
decision: &GatewayControlDecision,
plan_kind: &str,
transfer_tracker: &ProviderTransferTracker,
) -> Result<LocalExecutionRequestOutcome, GatewayError> {
let Some((mut attempt_source, candidate_count)) =
build_local_image_sync_attempt_source_for_kind(
@@ -1305,17 +1333,19 @@ pub(crate) async fn maybe_execute_sync_via_local_image_decision(
decision.clone(),
plan_kind.to_string(),
attempts,
transfer_tracker.clone(),
)?,
));
}
let outcome = execute_sync_attempt_source::<AiSyncAttempt, _>(
let outcome = execute_sync_attempt_source_with_transfer_tracker::<AiSyncAttempt, _>(
state,
parts,
trace_id,
decision,
plan_kind,
attempt_source,
transfer_tracker,
)
.await?;
@@ -1339,6 +1369,7 @@ pub(crate) async fn maybe_execute_stream_via_local_gemini_files_decision(
trace_id: &str,
decision: &GatewayControlDecision,
plan_kind: &str,
transfer_tracker: &ProviderTransferTracker,
) -> Result<LocalExecutionRequestOutcome, GatewayError> {
let Some((attempt_source, _candidate_count)) =
build_local_gemini_files_stream_attempt_source_for_kind(
@@ -1349,12 +1380,13 @@ pub(crate) async fn maybe_execute_stream_via_local_gemini_files_decision(
return Ok(LocalExecutionRequestOutcome::NoPath);
};
execute_stream_attempt_source::<AiStreamAttempt, _>(
execute_stream_attempt_source_with_transfer_tracker::<AiStreamAttempt, _>(
state,
trace_id,
decision,
plan_kind,
attempt_source,
transfer_tracker,
)
.await
}
@@ -1367,6 +1399,7 @@ pub(crate) async fn maybe_execute_stream_via_local_image_decision(
trace_id: &str,
decision: &GatewayControlDecision,
plan_kind: &str,
transfer_tracker: &ProviderTransferTracker,
) -> Result<LocalExecutionRequestOutcome, GatewayError> {
let Some((attempt_source, candidate_count)) = build_local_image_stream_attempt_source_for_kind(
state,
@@ -1382,12 +1415,13 @@ pub(crate) async fn maybe_execute_stream_via_local_image_decision(
return Ok(LocalExecutionRequestOutcome::NoPath);
};
let outcome = execute_stream_attempt_source::<AiStreamAttempt, _>(
let outcome = execute_stream_attempt_source_with_transfer_tracker::<AiStreamAttempt, _>(
state,
trace_id,
decision,
plan_kind,
attempt_source,
transfer_tracker,
)
.await?;
@@ -1412,6 +1446,7 @@ pub(crate) async fn maybe_execute_sync_via_local_video_decision(
trace_id: &str,
decision: &GatewayControlDecision,
plan_kind: &str,
transfer_tracker: &ProviderTransferTracker,
) -> Result<LocalExecutionRequestOutcome, GatewayError> {
let Some((attempt_source, _candidate_count)) = build_local_video_sync_attempt_source_for_kind(
state, parts, body_json, trace_id, decision, plan_kind,
@@ -1421,13 +1456,14 @@ pub(crate) async fn maybe_execute_sync_via_local_video_decision(
return Ok(LocalExecutionRequestOutcome::NoPath);
};
execute_sync_attempt_source::<AiSyncAttempt, _>(
execute_sync_attempt_source_with_transfer_tracker::<AiSyncAttempt, _>(
state,
parts,
trace_id,
decision,
plan_kind,
attempt_source,
transfer_tracker,
)
.await
}
@@ -1549,6 +1585,12 @@ mod tests {
async fn drain_execution_attempts(&mut self) -> Result<Vec<AiSyncAttempt>, GatewayError> {
Ok(self.attempts.drain(..).collect())
}
async fn skip_provider(&mut self, provider_id: &str) -> Result<(), GatewayError> {
self.attempts
.retain(|attempt| attempt.plan.provider_id != provider_id);
Ok(())
}
}
fn test_openai_image_heartbeat_decision() -> GatewayControlDecision {
@@ -1779,6 +1821,7 @@ mod tests {
test_openai_image_heartbeat_decision(),
TEST_OPENAI_IMAGE_SYNC_PLAN_KIND.to_string(),
attempts,
ProviderTransferTracker::default(),
Instant::now(),
)
.await
@@ -1793,6 +1836,70 @@ mod tests {
assert_eq!(body, json!({"data": [{"b64_json": "second-candidate"}]}));
}
#[tokio::test]
async fn openai_image_sync_heartbeat_honors_provider_transfer_limit() {
let call_count = Arc::new(AtomicUsize::new(0));
let call_count_for_override = Arc::clone(&call_count);
let state = AppState::new()
.expect("state should build")
.with_execution_runtime_sync_override_for_tests(move |plan| {
call_count_for_override.fetch_add(1, Ordering::SeqCst);
if plan.provider_id == "provider-fallback" {
Ok(test_openai_image_execution_result(
plan,
StatusCode::OK.as_u16(),
json!({"data": [{"b64_json": "fallback-provider"}]}),
))
} else {
Ok(test_openai_image_execution_result(
plan,
StatusCode::TOO_MANY_REQUESTS.as_u16(),
json!({"error": {"message": "retry another key"}}),
))
}
});
let mut attempts = vec![
test_openai_image_heartbeat_attempt(0, "endpoint-key-1", "candidate-key-1"),
test_openai_image_heartbeat_attempt(1, "endpoint-key-2", "candidate-key-2"),
test_openai_image_heartbeat_attempt(2, "endpoint-key-3", "candidate-key-3"),
test_openai_image_heartbeat_attempt(3, "endpoint-fallback", "candidate-fallback"),
];
for (index, attempt) in attempts.iter_mut().take(3).enumerate() {
attempt.plan.key_id = format!("key-{}", index + 1);
attempt.report_context = Some(json!({
"candidate_index": index,
"retry_index": 0,
"local_failover_policy": {
"max_transfer_count": 1,
"max_transfer_timeout_seconds": 0
}
}));
}
attempts[3].plan.provider_id = "provider-fallback".to_string();
attempts[3].plan.key_id = "key-fallback".to_string();
let outcome = execute_openai_image_sync_heartbeat_attempts(
state,
"/v1/images/generations".to_string(),
"trace-image-heartbeat-transfer-limit".to_string(),
test_openai_image_heartbeat_decision(),
TEST_OPENAI_IMAGE_SYNC_PLAN_KIND.to_string(),
attempts,
ProviderTransferTracker::default(),
Instant::now(),
)
.await
.expect("heartbeat attempts should execute");
let LocalExecutionRequestOutcome::Responded(response) = outcome else {
panic!("fallback provider should return a response");
};
let bytes = openai_image_sync_heartbeat_response_body_bytes(response).await;
let body: Value = serde_json::from_slice(&bytes).expect("body should decode");
assert_eq!(call_count.load(Ordering::SeqCst), 3);
assert_eq!(body, json!({"data": [{"b64_json": "fallback-provider"}]}));
}
#[tokio::test]
async fn standard_text_sync_heartbeat_missing_config_defaults_disabled() {
let state = AppState::new().expect("state should build");
@@ -1824,6 +1931,7 @@ mod tests {
&test_standard_text_heartbeat_decision(),
&json!({"model": "missing-local-candidate"}),
TEST_STANDARD_TEXT_SYNC_PLAN_KIND,
&ProviderTransferTracker::default(),
)
.await
.expect("heartbeat no-path check should execute");
@@ -3,7 +3,9 @@ use crate::ai_serving::api::{
};
use crate::control::GatewayControlDecision;
use crate::executor::{
execute_stream_plan_and_reports, execute_sync_plan_and_reports, LocalExecutionRequestOutcome,
execute_stream_plan_and_reports_with_transfer_tracker,
execute_sync_plan_and_reports_with_transfer_tracker, LocalExecutionRequestOutcome,
ProviderTransferTracker,
};
use crate::{AiExecutionPlanPayload, AppState, GatewayError, GatewayFallbackReason};
@@ -17,6 +19,7 @@ pub(crate) async fn maybe_execute_sync_via_plan_fallback(
_plan_kind: &str,
_bypass_cache_key: String,
_fallback_reason: GatewayFallbackReason,
transfer_tracker: &ProviderTransferTracker,
) -> Result<LocalExecutionRequestOutcome, GatewayError> {
let body_is_empty =
body_base64.is_none() && body_json.as_object().is_some_and(|value| value.is_empty());
@@ -47,7 +50,7 @@ pub(crate) async fn maybe_execute_sync_via_plan_fallback(
return Ok(LocalExecutionRequestOutcome::NoPath);
};
execute_sync_plan_and_reports(
execute_sync_plan_and_reports_with_transfer_tracker(
state,
parts,
trace_id,
@@ -58,6 +61,7 @@ pub(crate) async fn maybe_execute_sync_via_plan_fallback(
report_kind,
report_context,
}],
transfer_tracker,
)
.await
}
@@ -72,6 +76,7 @@ pub(crate) async fn maybe_execute_stream_via_plan_fallback(
_plan_kind: &str,
_bypass_cache_key: String,
_fallback_reason: GatewayFallbackReason,
transfer_tracker: &ProviderTransferTracker,
) -> Result<LocalExecutionRequestOutcome, GatewayError> {
let Some(payload) = maybe_build_stream_plan_payload(
state,
@@ -99,7 +104,7 @@ pub(crate) async fn maybe_execute_stream_via_plan_fallback(
return Ok(LocalExecutionRequestOutcome::NoPath);
};
execute_stream_plan_and_reports(
execute_stream_plan_and_reports_with_transfer_tracker(
state,
trace_id,
decision,
@@ -109,6 +114,7 @@ pub(crate) async fn maybe_execute_stream_via_plan_fallback(
report_kind,
report_context,
}],
transfer_tracker,
)
.await
}
@@ -21,14 +21,14 @@ use crate::stage_metrics::observe_gateway_stage_ms;
use crate::{AppState, GatewayError, GatewayFallbackReason};
use super::{
build_direct_plan_bypass_cache_key, execute_stream_plan_and_reports,
build_direct_plan_bypass_cache_key, execute_stream_plan_and_reports_with_transfer_tracker,
maybe_execute_stream_via_local_decision, maybe_execute_stream_via_local_gemini_files_decision,
maybe_execute_stream_via_local_image_decision,
maybe_execute_stream_via_local_openai_responses_decision,
maybe_execute_stream_via_local_same_format_provider_decision,
maybe_execute_stream_via_local_standard_decision, maybe_execute_stream_via_plan_fallback,
maybe_execute_stream_via_remote_decision, parse_local_request_body, should_skip_direct_plan,
LocalExecutionRequestOutcome,
LocalExecutionRequestOutcome, ProviderTransferTracker,
};
pub(crate) async fn maybe_execute_via_stream_decision_path(
@@ -86,6 +86,7 @@ pub(crate) async fn maybe_execute_via_stream_decision_path(
if skip_direct_plan {
return Ok(LocalExecutionRequestOutcome::NoPath);
}
let transfer_tracker = ProviderTransferTracker::default();
if plan_kind == OPENAI_CHAT_STREAM_PLAN_KIND
&& supports_stream_execution_decision_kind(plan_kind)
@@ -101,6 +102,7 @@ pub(crate) async fn maybe_execute_via_stream_decision_path(
body_base64,
plan_kind,
bypass_cache_key,
&transfer_tracker,
)
.await;
observe_gateway_stage_ms(
@@ -120,6 +122,7 @@ pub(crate) async fn maybe_execute_via_stream_decision_path(
plan_kind,
bypass_cache_key,
scheduler_supported: supports_stream_execution_decision_kind(plan_kind),
transfer_tracker,
};
Ok(from_ai_serving_outcome(
@@ -137,10 +140,17 @@ async fn execute_openai_chat_stream_fast_path(
body_base64: Option<String>,
plan_kind: &str,
bypass_cache_key: String,
transfer_tracker: &ProviderTransferTracker,
) -> Result<LocalExecutionRequestOutcome, GatewayError> {
let started_at = std::time::Instant::now();
let local_outcome = maybe_execute_stream_via_local_decision(
state, parts, trace_id, decision, body_json, plan_kind,
state,
parts,
trace_id,
decision,
body_json,
plan_kind,
transfer_tracker,
)
.await?;
observe_gateway_stage_ms(
@@ -176,6 +186,7 @@ async fn execute_openai_chat_stream_fast_path(
plan_kind,
bypass_cache_key,
GatewayFallbackReason::RemoteDecisionMiss,
transfer_tracker,
)
.await?;
observe_gateway_stage_ms(
@@ -195,6 +206,7 @@ struct GatewayStreamExecutionPathPort<'a> {
plan_kind: &'a str,
bypass_cache_key: String,
scheduler_supported: bool,
transfer_tracker: ProviderTransferTracker,
}
#[async_trait]
@@ -224,6 +236,7 @@ impl AiStreamExecutionPathPort for GatewayStreamExecutionPathPort<'_> {
self.trace_id,
self.decision,
self.plan_kind,
&self.transfer_tracker,
)
.await?
}
@@ -236,6 +249,7 @@ impl AiStreamExecutionPathPort for GatewayStreamExecutionPathPort<'_> {
self.trace_id,
self.decision,
self.plan_kind,
&self.transfer_tracker,
)
.await?
}
@@ -247,6 +261,7 @@ impl AiStreamExecutionPathPort for GatewayStreamExecutionPathPort<'_> {
self.decision,
self.body_json,
self.plan_kind,
&self.transfer_tracker,
)
.await?
}
@@ -258,6 +273,7 @@ impl AiStreamExecutionPathPort for GatewayStreamExecutionPathPort<'_> {
self.decision,
self.body_json,
self.plan_kind,
&self.transfer_tracker,
)
.await?
}
@@ -269,6 +285,7 @@ impl AiStreamExecutionPathPort for GatewayStreamExecutionPathPort<'_> {
self.decision,
self.body_json,
self.plan_kind,
&self.transfer_tracker,
)
.await?
}
@@ -280,6 +297,7 @@ impl AiStreamExecutionPathPort for GatewayStreamExecutionPathPort<'_> {
self.decision,
self.body_json,
self.plan_kind,
&self.transfer_tracker,
)
.await?
}
@@ -290,6 +308,7 @@ impl AiStreamExecutionPathPort for GatewayStreamExecutionPathPort<'_> {
self.trace_id,
self.decision,
self.plan_kind,
&self.transfer_tracker,
)
.await?
}
@@ -335,6 +354,7 @@ impl AiStreamExecutionPathPort for GatewayStreamExecutionPathPort<'_> {
self.plan_kind,
self.bypass_cache_key.clone(),
gateway_fallback_reason(reason),
&self.transfer_tracker,
)
.await?;
Ok(to_ai_serving_outcome(outcome))
@@ -440,6 +460,7 @@ async fn maybe_execute_local_video_task_content_stream(
trace_id: &str,
decision: &GatewayControlDecision,
plan_kind: &str,
transfer_tracker: &ProviderTransferTracker,
) -> Result<LocalExecutionRequestOutcome, GatewayError> {
if plan_kind != OPENAI_VIDEO_CONTENT_PLAN_KIND
|| decision.route_family.as_deref() != Some("openai")
@@ -481,7 +502,7 @@ async fn maybe_execute_local_video_task_content_stream(
)),
crate::video_tasks::LocalVideoTaskContentAction::StreamPlan(plan) => {
let plan = *plan;
execute_stream_plan_and_reports(
execute_stream_plan_and_reports_with_transfer_tracker(
state,
trace_id,
decision,
@@ -491,6 +512,7 @@ async fn maybe_execute_local_video_task_content_stream(
report_kind: None,
report_context: None,
}],
transfer_tracker,
)
.await
}
+16 -2
View File
@@ -19,7 +19,7 @@ use crate::control::GatewayControlDecision;
use crate::{AppState, GatewayError, GatewayFallbackReason};
use super::{
build_direct_plan_bypass_cache_key, execute_sync_plan_and_reports,
build_direct_plan_bypass_cache_key, execute_sync_plan_and_reports_with_transfer_tracker,
maybe_execute_sync_via_local_decision, maybe_execute_sync_via_local_gemini_files_decision,
maybe_execute_sync_via_local_image_decision,
maybe_execute_sync_via_local_openai_responses_decision,
@@ -27,6 +27,7 @@ use super::{
maybe_execute_sync_via_local_standard_decision, maybe_execute_sync_via_local_video_decision,
maybe_execute_sync_via_plan_fallback, maybe_execute_sync_via_remote_decision,
parse_local_request_body, should_skip_direct_plan, LocalExecutionRequestOutcome,
ProviderTransferTracker,
};
pub(crate) async fn maybe_execute_via_sync_decision_path(
@@ -73,6 +74,7 @@ pub(crate) async fn maybe_execute_via_sync_decision_path(
plan_kind,
bypass_cache_key,
scheduler_supported: supports_sync_execution_decision_kind(plan_kind),
transfer_tracker: ProviderTransferTracker::default(),
};
Ok(from_ai_serving_outcome(
@@ -91,6 +93,7 @@ struct GatewaySyncExecutionPathPort<'a> {
plan_kind: &'a str,
bypass_cache_key: String,
scheduler_supported: bool,
transfer_tracker: ProviderTransferTracker,
}
#[async_trait]
@@ -116,6 +119,7 @@ impl AiSyncExecutionPathPort for GatewaySyncExecutionPathPort<'_> {
self.trace_id,
self.decision,
self.plan_kind,
&self.transfer_tracker,
)
.await?
}
@@ -127,6 +131,7 @@ impl AiSyncExecutionPathPort for GatewaySyncExecutionPathPort<'_> {
self.trace_id,
self.decision,
self.plan_kind,
&self.transfer_tracker,
)
.await?
}
@@ -139,6 +144,7 @@ impl AiSyncExecutionPathPort for GatewaySyncExecutionPathPort<'_> {
self.trace_id,
self.decision,
self.plan_kind,
&self.transfer_tracker,
)
.await?
}
@@ -150,6 +156,7 @@ impl AiSyncExecutionPathPort for GatewaySyncExecutionPathPort<'_> {
self.decision,
self.body_json,
self.plan_kind,
&self.transfer_tracker,
)
.await?
}
@@ -161,6 +168,7 @@ impl AiSyncExecutionPathPort for GatewaySyncExecutionPathPort<'_> {
self.decision,
self.body_json,
self.plan_kind,
&self.transfer_tracker,
)
.await?
}
@@ -172,6 +180,7 @@ impl AiSyncExecutionPathPort for GatewaySyncExecutionPathPort<'_> {
self.decision,
self.body_json,
self.plan_kind,
&self.transfer_tracker,
)
.await?
}
@@ -183,6 +192,7 @@ impl AiSyncExecutionPathPort for GatewaySyncExecutionPathPort<'_> {
self.decision,
self.body_json,
self.plan_kind,
&self.transfer_tracker,
)
.await?
}
@@ -196,6 +206,7 @@ impl AiSyncExecutionPathPort for GatewaySyncExecutionPathPort<'_> {
self.trace_id,
self.decision,
self.plan_kind,
&self.transfer_tracker,
)
.await?
}
@@ -233,6 +244,7 @@ impl AiSyncExecutionPathPort for GatewaySyncExecutionPathPort<'_> {
self.plan_kind,
self.bypass_cache_key.clone(),
gateway_fallback_reason(reason),
&self.transfer_tracker,
)
.await?;
Ok(to_ai_serving_outcome(outcome))
@@ -342,6 +354,7 @@ async fn maybe_execute_local_video_task_follow_up_sync(
trace_id: &str,
decision: &GatewayControlDecision,
plan_kind: &str,
transfer_tracker: &ProviderTransferTracker,
) -> Result<LocalExecutionRequestOutcome, GatewayError> {
if !matches!(
plan_kind,
@@ -375,7 +388,7 @@ async fn maybe_execute_local_video_task_follow_up_sync(
return Ok(LocalExecutionRequestOutcome::NoPath);
};
execute_sync_plan_and_reports(
execute_sync_plan_and_reports_with_transfer_tracker(
state,
parts,
trace_id,
@@ -386,6 +399,7 @@ async fn maybe_execute_local_video_task_follow_up_sync(
report_kind: follow_up.report_kind,
report_context: follow_up.report_context,
}],
transfer_tracker,
)
.await
}
@@ -157,6 +157,10 @@ pub(crate) struct AdminProviderCreateRequest {
#[serde(default)]
pub(crate) max_retries: Option<i32>,
#[serde(default)]
pub(crate) max_transfer_count: Option<i64>,
#[serde(default)]
pub(crate) max_transfer_timeout_seconds: Option<i64>,
#[serde(default)]
pub(crate) proxy: Option<serde_json::Value>,
#[serde(
default,
@@ -212,6 +216,10 @@ pub(crate) struct AdminProviderUpdateRequest {
#[serde(default)]
pub(crate) max_retries: Option<i32>,
#[serde(default)]
pub(crate) max_transfer_count: Option<i64>,
#[serde(default)]
pub(crate) max_transfer_timeout_seconds: Option<i64>,
#[serde(default)]
pub(crate) proxy: Option<serde_json::Value>,
#[serde(
default,
@@ -12,6 +12,35 @@ pub(crate) const ADMIN_PROVIDER_POOL_QUOTA_PROBE_ACTIVE_SET_PREFIX: &str =
"ap:quota_probe:active_members";
pub(crate) const ADMIN_PROVIDER_OAUTH_DATA_UNAVAILABLE_DETAIL: &str =
"Admin provider OAuth data unavailable";
pub(crate) const PROVIDER_MAX_TRANSFER_COUNT_CONFIG_KEY: &str = "max_transfer_count";
pub(crate) const PROVIDER_MAX_TRANSFER_TIMEOUT_SECONDS_CONFIG_KEY: &str =
"max_transfer_timeout_seconds";
pub(crate) fn normalize_provider_transfer_limit(
value: i64,
field_name: &str,
) -> Result<u64, String> {
u64::try_from(value).map_err(|_| format!("{field_name} 必须是非负整数"))
}
pub(crate) fn normalize_provider_transfer_limit_json(
value: &serde_json::Value,
field_name: &str,
) -> Result<u64, String> {
value
.as_u64()
.ok_or_else(|| format!("{field_name} 必须是非负整数"))
}
pub(crate) fn provider_transfer_limit_from_config(
config: Option<&serde_json::Map<String, serde_json::Value>>,
field_name: &str,
) -> u64 {
config
.and_then(|config| config.get(field_name))
.and_then(serde_json::Value::as_u64)
.unwrap_or(0)
}
pub(crate) fn admin_provider_pool_quota_probe_active_members_key(provider_id: &str) -> String {
format!("{ADMIN_PROVIDER_POOL_QUOTA_PROBE_ACTIVE_SET_PREFIX}:{provider_id}")
@@ -1,3 +1,7 @@
use crate::handlers::admin::provider::shared::support::{
provider_transfer_limit_from_config, PROVIDER_MAX_TRANSFER_COUNT_CONFIG_KEY,
PROVIDER_MAX_TRANSFER_TIMEOUT_SECONDS_CONFIG_KEY,
};
use crate::handlers::admin::shared::unix_secs_to_rfc3339;
use crate::handlers::public::{request_candidate_event_unix_ms, request_candidate_status_label};
use crate::orchestration::codex_cyber_flag_passthrough_enabled;
@@ -126,6 +130,12 @@ pub(crate) fn build_admin_provider_summary_value(
let config = provider_config
.as_ref()
.and_then(serde_json::Value::as_object);
let max_transfer_count =
provider_transfer_limit_from_config(config, PROVIDER_MAX_TRANSFER_COUNT_CONFIG_KEY);
let max_transfer_timeout_seconds = provider_transfer_limit_from_config(
config,
PROVIDER_MAX_TRANSFER_TIMEOUT_SECONDS_CONFIG_KEY,
);
let provider_ops_config = config.and_then(|cfg| cfg.get("provider_ops"));
let ops_configured = provider_ops_config.is_some_and(json_truthy);
let ops_architecture_id = provider_ops_config
@@ -184,6 +194,8 @@ pub(crate) fn build_admin_provider_summary_value(
"quota_last_reset_at": quota_last_reset_at,
"quota_expires_at": quota_expires_at,
"max_retries": provider.max_retries,
"max_transfer_count": max_transfer_count,
"max_transfer_timeout_seconds": max_transfer_timeout_seconds,
"proxy": provider.proxy.clone(),
"stream_first_byte_timeout": provider.stream_first_byte_timeout_secs,
"request_timeout": provider.request_timeout_secs,
@@ -1,6 +1,8 @@
use crate::handlers::admin::provider::shared::payloads::AdminProviderCreateRequest;
use crate::handlers::admin::provider::shared::support::{
normalize_provider_billing_type, parse_optional_rfc3339_unix_secs,
normalize_provider_billing_type, normalize_provider_transfer_limit,
normalize_provider_transfer_limit_json, parse_optional_rfc3339_unix_secs,
PROVIDER_MAX_TRANSFER_COUNT_CONFIG_KEY, PROVIDER_MAX_TRANSFER_TIMEOUT_SECONDS_CONFIG_KEY,
};
use crate::handlers::admin::provider::write::normalize::normalize_chat_pii_redaction_config;
use crate::handlers::admin::provider::write::normalize::normalize_pool_advanced_config;
@@ -114,6 +116,27 @@ pub(crate) async fn build_admin_create_provider_record(
let mut config_map = normalize_json_object(payload.config, "config")?
.and_then(|value| value.as_object().cloned())
.unwrap_or_default();
for (field_name, payload_value) in [
(
PROVIDER_MAX_TRANSFER_COUNT_CONFIG_KEY,
payload.max_transfer_count,
),
(
PROVIDER_MAX_TRANSFER_TIMEOUT_SECONDS_CONFIG_KEY,
payload.max_transfer_timeout_seconds,
),
] {
let value = match payload_value {
Some(value) => Some(normalize_provider_transfer_limit(value, field_name)?),
None => config_map
.get(field_name)
.map(|value| normalize_provider_transfer_limit_json(value, field_name))
.transpose()?,
};
if let Some(value) = value {
config_map.insert(field_name.to_string(), json!(value));
}
}
if let Some(value) = normalize_pool_advanced_config(payload.pool_advanced)? {
config_map.insert("pool_advanced".to_string(), value);
}
@@ -1,6 +1,8 @@
use crate::handlers::admin::provider::shared::payloads::AdminProviderUpdatePatch;
use crate::handlers::admin::provider::shared::support::{
normalize_provider_billing_type, parse_optional_rfc3339_unix_secs,
normalize_provider_billing_type, normalize_provider_transfer_limit,
normalize_provider_transfer_limit_json, parse_optional_rfc3339_unix_secs,
PROVIDER_MAX_TRANSFER_COUNT_CONFIG_KEY, PROVIDER_MAX_TRANSFER_TIMEOUT_SECONDS_CONFIG_KEY,
};
use crate::handlers::admin::provider::write::normalize::normalize_chat_pii_redaction_config;
use crate::handlers::admin::provider::write::normalize::normalize_pool_advanced_config;
@@ -242,6 +244,30 @@ pub(crate) async fn build_admin_update_provider_record(
}
}
for (field_name, payload_value) in [
(
PROVIDER_MAX_TRANSFER_COUNT_CONFIG_KEY,
payload.max_transfer_count,
),
(
PROVIDER_MAX_TRANSFER_TIMEOUT_SECONDS_CONFIG_KEY,
payload.max_transfer_timeout_seconds,
),
] {
if fields.contains(field_name) {
let value = payload_value
.map(|value| normalize_provider_transfer_limit(value, field_name))
.transpose()?
.unwrap_or(0);
config_map.insert(field_name.to_string(), json!(value));
} else if fields.contains("config") {
if let Some(value) = config_map.get(field_name) {
let value = normalize_provider_transfer_limit_json(value, field_name)?;
config_map.insert(field_name.to_string(), json!(value));
}
}
}
if fields.contains("claude_code_advanced") {
if fields.is_null("claude_code_advanced") {
config_map.remove("claude_code_advanced");
@@ -12,6 +12,8 @@ pub(crate) const CYBER_CONTINUE_FAILOVER_CONFIG_KEY: &str = "cyber_continue_fail
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct LocalFailoverPolicy {
pub(crate) max_retries: Option<u64>,
pub(crate) max_transfer_count: u64,
pub(crate) max_transfer_timeout_seconds: u64,
pub(crate) stop_status_codes: BTreeSet<u16>,
pub(crate) continue_status_codes: BTreeSet<u16>,
pub(crate) success_failover_patterns: Vec<LocalFailoverRegexRule>,
@@ -24,6 +26,8 @@ impl Default for LocalFailoverPolicy {
fn default() -> Self {
Self {
max_retries: None,
max_transfer_count: 0,
max_transfer_timeout_seconds: 0,
stop_status_codes: BTreeSet::new(),
continue_status_codes: BTreeSet::new(),
success_failover_patterns: Vec::new(),
@@ -63,6 +67,8 @@ pub(crate) async fn resolve_local_failover_policy(
key_id = %plan.key_id,
source = "transport_snapshot",
max_retries = ?policy.max_retries,
max_transfer_count = policy.max_transfer_count,
max_transfer_timeout_seconds = policy.max_transfer_timeout_seconds,
stop_status_code_count = policy.stop_status_codes.len(),
continue_status_code_count = policy.continue_status_codes.len(),
success_failover_pattern_count = policy.success_failover_patterns.len(),
@@ -87,6 +93,7 @@ pub(crate) async fn cyber_continue_failover_enabled(state: &AppState) -> bool {
pub(crate) fn local_failover_policy_from_transport(
transport: &GatewayProviderTransportSnapshot,
) -> LocalFailoverPolicy {
let provider_config = transport.provider.config.as_ref();
let rules = transport
.provider
.config
@@ -111,6 +118,14 @@ pub(crate) fn local_failover_policy_from_transport(
LocalFailoverPolicy {
max_retries,
max_transfer_count: provider_config
.and_then(|value| value.get("max_transfer_count"))
.and_then(parse_u64_value)
.unwrap_or(0),
max_transfer_timeout_seconds: provider_config
.and_then(|value| value.get("max_transfer_timeout_seconds"))
.and_then(parse_u64_value)
.unwrap_or(0),
retry_client_errors_by_default:
crate::ai_serving::api_format_defaults_to_client_error_failover(
&transport.endpoint.api_format,
@@ -161,6 +176,14 @@ pub(crate) fn local_failover_policy_from_report_context(
Some(LocalFailoverPolicy {
max_retries: object.get("max_retries").and_then(parse_u64_value),
max_transfer_count: object
.get("max_transfer_count")
.and_then(parse_u64_value)
.unwrap_or(0),
max_transfer_timeout_seconds: object
.get("max_transfer_timeout_seconds")
.and_then(parse_u64_value)
.unwrap_or(0),
stop_status_codes: object
.get("stop_status_codes")
.map(parse_status_code_list)
@@ -208,6 +231,8 @@ fn parse_status_code_list(value: &Value) -> BTreeSet<u16> {
fn local_failover_policy_to_value(policy: &LocalFailoverPolicy) -> Value {
json!({
"max_retries": policy.max_retries,
"max_transfer_count": policy.max_transfer_count,
"max_transfer_timeout_seconds": policy.max_transfer_timeout_seconds,
"stop_status_codes": policy.stop_status_codes.iter().copied().collect::<Vec<_>>(),
"continue_status_codes": policy.continue_status_codes.iter().copied().collect::<Vec<_>>(),
"success_failover_patterns": policy.success_failover_patterns.iter().map(local_failover_regex_rule_to_value).collect::<Vec<_>>(),
@@ -385,6 +410,8 @@ mod tests {
Some(5),
Some(4),
Some(json!({
"max_transfer_count": 10,
"max_transfer_timeout_seconds": 60,
"failover_rules": {
"max_retries": 2,
"continue_status_codes": [429],
@@ -400,6 +427,8 @@ mod tests {
local_failover_policy_from_report_context(Some(&report_context)),
Some(LocalFailoverPolicy {
max_retries: Some(2),
max_transfer_count: 10,
max_transfer_timeout_seconds: 60,
stop_status_codes: [400].into_iter().collect(),
continue_status_codes: [429].into_iter().collect(),
success_failover_patterns: vec![LocalFailoverRegexRule {
@@ -416,6 +445,33 @@ mod tests {
);
}
#[test]
fn transfer_limits_are_read_only_from_top_level_provider_config() {
let top_level = local_failover_policy_from_transport(&sample_transport(
None,
None,
Some(json!({
"max_transfer_count": 3,
"max_transfer_timeout_seconds": 45,
})),
));
assert_eq!(top_level.max_transfer_count, 3);
assert_eq!(top_level.max_transfer_timeout_seconds, 45);
let nested = local_failover_policy_from_transport(&sample_transport(
None,
None,
Some(json!({
"failover_rules": {
"max_transfer_count": 8,
"max_transfer_timeout_seconds": 90,
}
})),
));
assert_eq!(nested.max_transfer_count, 0);
assert_eq!(nested.max_transfer_timeout_seconds, 0);
}
#[test]
fn search_transport_disables_default_client_error_failover() {
let mut transport = sample_transport(None, None, None);
@@ -1,5 +1,20 @@
use super::*;
#[test]
fn specialized_decisions_embed_provider_failover_policy_in_report_context() {
for path in [
"apps/aether-gateway/src/ai_serving/planner/specialized/files/decision.rs",
"apps/aether-gateway/src/ai_serving/planner/specialized/image/decision.rs",
"apps/aether-gateway/src/ai_serving/planner/specialized/video/decision.rs",
] {
let source = read_workspace_file(path);
assert!(
source.contains("append_local_failover_policy_to_value(report_context, &transport)"),
"{path} should embed provider failover policy in every generated report context"
);
}
}
#[test]
fn ai_serving_target_structure_removes_legacy_pipeline_boundary() {
assert!(
@@ -340,6 +340,8 @@ async fn gateway_handles_admin_provider_summary_locally_with_trusted_admin_princ
assert_eq!(payload["billing_type"], "monthly_quota");
assert_eq!(payload["monthly_quota_usd"], 100.0);
assert_eq!(payload["monthly_used_usd"], 12.5);
assert_eq!(payload["max_transfer_count"], 0);
assert_eq!(payload["max_transfer_timeout_seconds"], 0);
assert_eq!(payload["total_endpoints"], 2);
assert_eq!(payload["active_endpoints"], 2);
assert_eq!(payload["total_keys"], 2);
@@ -832,6 +834,8 @@ async fn gateway_updates_admin_provider_locally_with_trusted_admin_principal() {
"is_active": false,
"concurrent_limit": 8,
"max_retries": 6,
"max_transfer_count": 10,
"max_transfer_timeout_seconds": 60,
"request_timeout": aether_contracts::MAX_EXECUTION_REQUEST_TIMEOUT_SECS,
"stream_first_byte_timeout": 11.0,
"enable_format_conversion": false,
@@ -860,6 +864,8 @@ async fn gateway_updates_admin_provider_locally_with_trusted_admin_principal() {
assert_eq!(payload["enable_format_conversion"], false);
assert_eq!(payload["is_active"], false);
assert_eq!(payload["max_retries"], 6);
assert_eq!(payload["max_transfer_count"], 10);
assert_eq!(payload["max_transfer_timeout_seconds"], 60);
assert_eq!(
payload["request_timeout"].as_f64(),
Some(aether_contracts::MAX_EXECUTION_REQUEST_TIMEOUT_SECS as f64)
@@ -888,6 +894,29 @@ async fn gateway_updates_admin_provider_locally_with_trusted_admin_principal() {
.expect("request should succeed");
assert_eq!(invalid_timeout_response.status(), StatusCode::BAD_REQUEST);
for (field_name, value) in [
("max_transfer_count", -1),
("max_transfer_timeout_seconds", -1),
] {
let invalid_transfer_response = reqwest::Client::new()
.patch(format!("{gateway_url}/api/admin/providers/provider-openai"))
.header(crate::constants::GATEWAY_HEADER, "rust-phase3b")
.header(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user-123")
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin")
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123")
.json(&json!({ (field_name): value }))
.send()
.await
.expect("request should succeed");
let status = invalid_transfer_response.status();
let payload: serde_json::Value = invalid_transfer_response
.json()
.await
.expect("json body should parse");
assert_eq!(status, StatusCode::BAD_REQUEST);
assert_eq!(payload["detail"], format!("{field_name} 必须是非负整数"));
}
let disable_response = reqwest::Client::new()
.patch(format!("{gateway_url}/api/admin/providers/provider-openai"))
.header(crate::constants::GATEWAY_HEADER, "rust-phase3b")
@@ -895,6 +924,7 @@ async fn gateway_updates_admin_provider_locally_with_trusted_admin_principal() {
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin")
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123")
.json(&json!({
"max_transfer_count": null,
"config": {
"chat_pii_redaction": {"enabled": false}
}
@@ -912,6 +942,8 @@ async fn gateway_updates_admin_provider_locally_with_trusted_admin_principal() {
json!({"enabled": false})
);
assert_eq!(disable_payload["pool_advanced"], json!({}));
assert_eq!(disable_payload["max_transfer_count"], 0);
assert_eq!(disable_payload["max_transfer_timeout_seconds"], 60);
assert_eq!(
disable_payload["failover_rules"],
json!({"strategy": "ordered"})
@@ -946,6 +978,22 @@ async fn gateway_updates_admin_provider_locally_with_trusted_admin_principal() {
updated_provider.request_timeout_secs,
Some(aether_contracts::MAX_EXECUTION_REQUEST_TIMEOUT_SECS as f64)
);
assert_eq!(
updated_provider
.config
.as_ref()
.and_then(|value| value.get("max_transfer_count"))
.and_then(serde_json::Value::as_u64),
Some(0)
);
assert_eq!(
updated_provider
.config
.as_ref()
.and_then(|value| value.get("max_transfer_timeout_seconds"))
.and_then(serde_json::Value::as_u64),
Some(60)
);
assert_eq!(
updated_provider
.config
@@ -1022,6 +1070,8 @@ async fn gateway_creates_admin_provider_locally_with_trusted_admin_principal() {
"website": "codex.example",
"keep_priority_on_conversion": true,
"max_retries": 7,
"max_transfer_count": 12,
"max_transfer_timeout_seconds": 90,
"request_timeout": aether_contracts::MAX_EXECUTION_REQUEST_TIMEOUT_SECS,
"config": {"chat_pii_redaction": {"enabled": true}},
"pool_advanced": {},
@@ -1057,6 +1107,22 @@ async fn gateway_creates_admin_provider_locally_with_trusted_admin_principal() {
assert_eq!(created.website.as_deref(), Some("https://codex.example"));
assert!(created.enable_format_conversion);
assert_eq!(created.max_retries, Some(7));
assert_eq!(
created
.config
.as_ref()
.and_then(|value| value.get("max_transfer_count"))
.and_then(serde_json::Value::as_u64),
Some(12)
);
assert_eq!(
created
.config
.as_ref()
.and_then(|value| value.get("max_transfer_timeout_seconds"))
.and_then(serde_json::Value::as_u64),
Some(90)
);
assert_eq!(
created.request_timeout_secs,
Some(aether_contracts::MAX_EXECUTION_REQUEST_TIMEOUT_SECS as f64)