Fix provider pool exhaustion scheduling

This commit is contained in:
fawney19
2026-05-20 20:13:07 +08:00
parent d0981c2fd5
commit c972bbd397
9 changed files with 424 additions and 48 deletions

View File

@@ -709,31 +709,39 @@ impl<'a> PoolKeyCursor<'a> {
}
async fn refill_queued_candidates(&mut self) -> bool {
let mut candidates = Vec::new();
let refill_target = self.window_size.max(1) as usize;
// Keep pool expansion bounded; the cursor freezes one small window at a time.
while candidates.len() < refill_target {
let Some(mut page_candidates) = self.next_page_candidates().await else {
break;
};
candidates.append(&mut page_candidates);
}
if candidates.is_empty() {
return false;
}
loop {
let mut candidates = Vec::new();
// Keep pool expansion bounded; the cursor freezes one small window at a time.
while candidates.len() < refill_target {
let Some(mut page_candidates) = self.next_page_candidates().await else {
break;
};
candidates.append(&mut page_candidates);
}
let (mut scheduled, mut skipped) = schedule_pool_page_candidates(
self.state,
candidates,
self.sticky_session_token.as_deref(),
)
.await;
scheduled.truncate(refill_target);
self.record_skipped_candidates(&skipped);
self.queued_candidates.extend(scheduled.drain(..));
self.skipped_candidates.append(&mut skipped);
!self.queued_candidates.is_empty()
if candidates.is_empty() {
return false;
}
let (mut scheduled, mut skipped) = schedule_pool_page_candidates(
self.state,
candidates,
self.sticky_session_token.as_deref(),
)
.await;
self.record_skipped_candidates(&skipped);
self.skipped_candidates.append(&mut skipped);
if scheduled.is_empty() {
continue;
}
scheduled.truncate(refill_target);
self.queued_candidates.extend(scheduled.drain(..));
return true;
}
}
async fn next_queued_candidate(&mut self) -> Option<EligibleLocalExecutionCandidate> {
@@ -2798,6 +2806,78 @@ mod tests {
assert_eq!(cursor.skip_reason_counts.get("pool_cooldown"), Some(&1));
}
#[tokio::test]
async fn pool_key_cursor_continues_after_exhausted_window() {
let provider_config = Some(json!({
"pool_advanced": {
"skip_exhausted_accounts": true
}
}));
let (provider, endpoint, mut keys, rows) = large_pool_fixture(3, provider_config.clone());
for key in keys.iter_mut().take(2) {
key.status_snapshot = Some(json!({
"quota": {
"provider_type": "openai",
"exhausted": true,
"usage_ratio": 1.0,
"windows": [
{
"code": "daily",
"used_ratio": 1.0,
"remaining_ratio": 0.0
}
]
}
}));
}
let data_state =
GatewayDataState::with_provider_catalog_and_minimal_candidate_selection_for_tests(
Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![provider],
vec![endpoint],
keys,
)),
Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(rows)),
)
.with_encryption_key_for_tests(aether_crypto::DEVELOPMENT_ENCRYPTION_KEY);
let app = AppState::new()
.expect("state should build")
.with_data_state_for_tests(data_state);
let group = sample_eligible_candidate(
"provider-pool",
"endpoint-1",
"pool-group",
10,
provider_config,
);
let mut cursor = PoolKeyCursor::new(PlannerAppState::new(&app), group, None, None, None);
cursor.window_size = 2;
cursor.page_size = 2;
cursor.max_scanned_keys = 4;
let candidate = cursor
.next_key()
.await
.expect("cursor should scan past an exhausted window");
assert_eq!(candidate.candidate.key_id, "key-00002");
assert_eq!(candidate.orchestration.pool_key_index, Some(0));
assert!(candidate.orchestration.pool_key_lease.is_none());
assert_eq!(
cursor
.skip_reason_counts
.get(aether_pool_core::POOL_ACCOUNT_EXHAUSTED_SKIP_REASON),
Some(&2)
);
let skipped = cursor.take_skipped_candidates();
assert_eq!(skipped.len(), 2);
assert!(skipped.iter().all(|candidate| {
candidate.skip_reason == aether_pool_core::POOL_ACCOUNT_EXHAUSTED_SKIP_REASON
}));
}
#[tokio::test]
async fn pool_key_cursor_simulates_large_lru_pool_with_lazy_pages_and_dynamic_skips() {
const KEY_COUNT: usize = 2048;

View File

@@ -2,7 +2,9 @@ use super::state::{
decode_jwt_claims, enrich_admin_provider_oauth_auth_config, json_non_empty_string,
json_u64_value,
};
use crate::handlers::admin::admin_provider_pool_config;
use crate::handlers::admin::request::AdminAppState;
use crate::maintenance::ensure_provider_key_pool_scores_for_keys;
use crate::provider_key_auth::provider_active_api_formats;
use crate::GatewayError;
use aether_data_contracts::repository::provider_catalog::{
@@ -168,6 +170,7 @@ pub(crate) async fn create_provider_oauth_catalog_key(
.app()
.invalidate_local_oauth_refresh_entry(&key.id)
.await;
seed_provider_oauth_pool_score(state, provider_id, key, now_unix_secs).await;
}
Ok(created)
}
@@ -221,10 +224,75 @@ pub(crate) async fn update_existing_provider_oauth_catalog_key(
.app()
.invalidate_local_oauth_refresh_entry(&key.id)
.await;
seed_provider_oauth_pool_score(state, &existing_key.provider_id, key, now_unix_secs).await;
}
Ok(persisted)
}
async fn seed_provider_oauth_pool_score(
state: &AdminAppState<'_>,
provider_id: &str,
key: &StoredProviderCatalogKey,
now_unix_secs: u64,
) {
let provider_id = provider_id.to_string();
let provider = match state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await
{
Ok(mut providers) => providers.pop(),
Err(err) => {
tracing::debug!(
provider_id = %provider_id,
key_id = %key.id,
error = ?err,
"gateway provider oauth provisioning: failed to read provider for pool score seed"
);
return;
}
};
let Some(provider) = provider else {
return;
};
let Some(pool_config) = admin_provider_pool_config(&provider) else {
return;
};
let endpoints = match state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
.await
{
Ok(endpoints) => endpoints,
Err(err) => {
tracing::debug!(
provider_id = %provider_id,
key_id = %key.id,
error = ?err,
"gateway provider oauth provisioning: failed to read endpoints for pool score seed"
);
return;
}
};
let score_ensure_budget = (pool_config.score_fallback_scan_limit as usize).clamp(1, 50_000);
if let Err(err) = ensure_provider_key_pool_scores_for_keys(
state.as_ref(),
&provider,
&pool_config,
&endpoints,
std::slice::from_ref(key),
now_unix_secs,
score_ensure_budget,
)
.await
{
tracing::debug!(
provider_id = %provider_id,
key_id = %key.id,
error = ?err,
"gateway provider oauth provisioning: failed to seed pool score row"
);
}
}
fn provider_oauth_catalog_key_api_formats(
provider_type: &str,
api_formats: &[String],

View File

@@ -12,8 +12,12 @@ use aether_data::repository::management_tokens::{
use aether_data::repository::oauth_providers::{
InMemoryOAuthProviderRepository, OAuthProviderReadRepository,
};
use aether_data::repository::pool_scores::InMemoryPoolMemberScoreRepository;
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
use aether_data::repository::proxy_nodes::InMemoryProxyNodeRepository;
use aether_data_contracts::repository::pool_scores::{
GetPoolMemberScoresByIdsQuery, PoolMemberHardState, PoolMemberIdentity, PoolScoreReadRepository,
};
use aether_data_contracts::repository::provider_catalog::{
ProviderCatalogReadRepository, ProviderCatalogWriteRepository, StoredProviderCatalogEndpoint,
};
@@ -32,6 +36,7 @@ use super::super::{
use crate::admin_api::{
maybe_build_local_admin_provider_oauth_response, AdminAppState, AdminRequestContext,
};
use crate::ai_serving::{provider_key_pool_score_id, provider_key_pool_score_scope};
use crate::audit::AdminAuditEvent;
use crate::constants::{
GATEWAY_HEADER, TRUSTED_ADMIN_MANAGEMENT_TOKEN_ID_HEADER, TRUSTED_ADMIN_SESSION_ID_HEADER,
@@ -1930,6 +1935,7 @@ async fn gateway_batch_imports_chatgpt_web_access_tokens_with_pool_hints() {
let mut provider = sample_provider("provider-chatgpt-web", "chatgpt_web", 10);
provider.provider_type = "chatgpt_web".to_string();
provider.config = Some(json!({"pool_advanced": {}}));
let endpoint = sample_endpoint(
"endpoint-chatgpt-web-image",
"provider-chatgpt-web",
@@ -1941,6 +1947,7 @@ async fn gateway_batch_imports_chatgpt_web_access_tokens_with_pool_hints() {
vec![endpoint],
vec![],
));
let pool_score_repository = Arc::new(InMemoryPoolMemberScoreRepository::default());
let (token_url, token_handle) = start_server(token_server).await;
let gateway = build_router_with_state(
@@ -1950,6 +1957,7 @@ async fn gateway_batch_imports_chatgpt_web_access_tokens_with_pool_hints() {
GatewayDataState::with_provider_catalog_repository_for_tests(
provider_catalog_repository.clone(),
)
.with_pool_score_repository_for_tests(Arc::clone(&pool_score_repository))
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
)
.with_provider_oauth_token_url_for_tests(
@@ -2030,6 +2038,20 @@ async fn gateway_batch_imports_chatgpt_web_access_tokens_with_pool_hints() {
assert_eq!(auth_config["plan_type"], "plus");
assert_eq!(auth_config["user_id"], "user-pool-image");
let score_scope = provider_key_pool_score_scope();
let score_identity =
PoolMemberIdentity::provider_api_key("provider-chatgpt-web", persisted.id.clone());
let scores = pool_score_repository
.get_pool_member_scores_by_ids(&GetPoolMemberScoresByIdsQuery {
ids: vec![provider_key_pool_score_id(&score_identity, &score_scope)],
})
.await
.expect("pool score should load");
assert_eq!(scores.len(), 1);
assert_eq!(scores[0].member_id, persisted.id);
assert_eq!(scores[0].hard_state, PoolMemberHardState::Unknown);
assert!(scores[0].score > 0.0);
gateway_handle.abort();
token_handle.abort();
}