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