mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-06 17:37:47 +08:00
feat(providers): add provider transfer limits
This commit is contained in:
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user