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
@@ -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 {