Refactor pool candidate scheduling

This commit is contained in:
fawney19
2026-05-03 20:14:29 +08:00
parent 8ebee9922c
commit a24e4a793d
55 changed files with 4825 additions and 311 deletions
@@ -9,6 +9,7 @@ use aether_ai_serving::{
use aether_scheduler_core::SchedulerMinimalCandidateSelectionCandidate;
use async_trait::async_trait;
use serde_json::Value;
use std::collections::VecDeque;
use std::convert::Infallible;
use uuid::Uuid;
@@ -16,16 +17,18 @@ use crate::ai_serving::planner::candidate_affinity_cache::remember_scheduler_aff
use crate::ai_serving::planner::candidate_resolution::{
resolve_and_rank_local_execution_candidates,
resolve_and_rank_local_execution_candidates_without_transport_pair_gate,
EligibleLocalExecutionCandidate, SkippedLocalExecutionCandidate,
resolve_and_rank_logical_local_execution_candidates, EligibleLocalExecutionCandidate,
LocalExecutionCandidateKind, SkippedLocalExecutionCandidate,
};
use crate::ai_serving::planner::materialization_policy::LocalCandidatePersistencePolicy;
use crate::ai_serving::planner::pool_scheduler::PoolKeyCursor;
use crate::ai_serving::planner::runtime_miss::record_local_runtime_candidate_skip_reason;
use crate::ai_serving::planner::CandidateFailureDiagnostic;
use crate::ai_serving::{GatewayAuthApiKeySnapshot, PlannerAppState};
use crate::clock::current_unix_ms;
use crate::handlers::shared::provider_pool::admin_provider_pool_config_from_config_value;
use crate::orchestration::{local_attempt_slot_count, ExecutionAttemptIdentity};
use crate::AppState;
use crate::{AppState, GatewayError};
#[derive(Debug, Clone)]
pub(crate) struct LocalExecutionCandidateAttempt {
@@ -35,6 +38,71 @@ pub(crate) struct LocalExecutionCandidateAttempt {
pub(crate) candidate_id: String,
}
pub(crate) struct LocalExecutionCandidateAttemptSource<'a> {
items: VecDeque<LocalExecutionCandidateAttemptSourceItem<'a>>,
}
#[async_trait]
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>;
}
enum LocalExecutionCandidateAttemptSourceItem<'a> {
Static {
attempts: VecDeque<LocalExecutionCandidateAttempt>,
},
Pool {
cursor: PoolKeyCursor<'a>,
candidate_index: u32,
pending_attempts: VecDeque<LocalExecutionCandidateAttempt>,
},
}
impl<'a> LocalExecutionCandidateAttemptSource<'a> {
pub(crate) async fn next_attempt(&mut self) -> Option<LocalExecutionCandidateAttempt> {
loop {
let front = self.items.front_mut()?;
match front {
LocalExecutionCandidateAttemptSourceItem::Static { attempts } => {
if let Some(attempt) = attempts.pop_front() {
if attempts.is_empty() {
self.items.pop_front();
}
return Some(attempt);
}
self.items.pop_front();
}
LocalExecutionCandidateAttemptSourceItem::Pool {
cursor,
candidate_index,
pending_attempts,
} => {
if let Some(attempt) = pending_attempts.pop_front() {
return Some(attempt);
}
let Some(candidate) = cursor.next_key().await else {
cursor.log_exhausted();
let _ = cursor.take_skipped_candidates();
self.items.pop_front();
continue;
};
*pending_attempts = build_unpersisted_local_execution_candidate_attempts(
candidate,
*candidate_index,
);
}
}
}
}
pub(crate) fn drain_static_attempts(&mut self) -> Vec<LocalExecutionCandidateAttempt> {
self.items.clear();
Vec::new()
}
}
impl LocalExecutionCandidateAttempt {
pub(crate) fn attempt_identity(&self) -> ExecutionAttemptIdentity {
ExecutionAttemptIdentity::new(self.candidate_index, self.retry_index)
@@ -350,6 +418,99 @@ where
}
}
#[allow(clippy::too_many_arguments)]
pub(crate) async fn build_local_execution_candidate_attempt_source_with_serving<'a, F, G>(
state: PlannerAppState<'a>,
trace_id: &str,
client_api_format: &str,
requested_model: Option<&str>,
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
required_capabilities: Option<&Value>,
sticky_session_token: Option<&str>,
request_auth_channel: Option<&str>,
persistence_policy: LocalCandidatePersistencePolicy<'_>,
candidates: Vec<SchedulerMinimalCandidateSelectionCandidate>,
preselection_skipped: Vec<SkippedLocalExecutionCandidate>,
resolution_mode: LocalCandidateResolutionMode,
build_available_extra_data: F,
decorate_skipped_candidate: G,
) -> (LocalExecutionCandidateAttemptSource<'a>, usize)
where
F: Fn(&EligibleLocalExecutionCandidate) -> Option<Value> + Send + Sync,
G: Fn(SkippedLocalExecutionCandidate) -> SkippedLocalExecutionCandidate + Send + Sync,
{
let (candidates, resolved_skipped) = resolve_and_rank_logical_local_execution_candidates(
state,
candidates,
client_api_format,
requested_model,
auth_snapshot,
required_capabilities,
sticky_session_token,
request_auth_channel,
resolution_mode,
)
.await;
let skipped_candidate_count = preselection_skipped.len() + resolved_skipped.len();
let skipped_candidates = preselection_skipped
.into_iter()
.chain(resolved_skipped)
.map(decorate_skipped_candidate)
.collect::<Vec<_>>();
let candidate_count = candidates.len() + skipped_candidate_count;
if persistence_policy.skipped.record_runtime_miss_diagnostic {
for skipped_candidate in &skipped_candidates {
record_local_runtime_candidate_skip_reason(
state.app(),
trace_id,
skipped_candidate.skip_reason,
);
}
}
remember_first_local_candidate_affinity(
state,
auth_snapshot,
client_api_format,
requested_model,
&candidates,
);
let mut items = VecDeque::new();
for (candidate_index, candidate) in candidates.into_iter().enumerate() {
let candidate_index = candidate_index as u32;
match candidate.kind {
LocalExecutionCandidateKind::SingleKey => {
let attempts = build_unpersisted_local_execution_candidate_attempts(
candidate,
candidate_index,
);
if !attempts.is_empty() {
items.push_back(LocalExecutionCandidateAttemptSourceItem::Static { attempts });
}
}
LocalExecutionCandidateKind::PoolGroup => {
items.push_back(LocalExecutionCandidateAttemptSourceItem::Pool {
cursor: PoolKeyCursor::new(
state,
candidate,
sticky_session_token,
requested_model,
request_auth_channel,
),
candidate_index,
pending_attempts: VecDeque::new(),
});
}
}
}
(
LocalExecutionCandidateAttemptSource { items },
candidate_count,
)
}
pub(crate) fn remember_first_local_candidate_affinity(
state: PlannerAppState<'_>,
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
@@ -438,6 +599,99 @@ where
.await
}
async fn persist_available_local_execution_candidate_at_index<F>(
state: PlannerAppState<'_>,
trace_id: &str,
context: LocalAvailableCandidatePersistenceContext<'_>,
candidate: EligibleLocalExecutionCandidate,
candidate_index: u32,
build_extra_data: &F,
) -> Vec<LocalExecutionCandidateAttempt>
where
F: Fn(&EligibleLocalExecutionCandidate) -> Option<Value> + Send + Sync,
{
let attempt_slots = local_attempt_slot_count(&candidate.transport).max(1);
let extra_data = ai_candidate_extra_data_with_ranking(
build_extra_data(&candidate),
candidate.ranking.as_ref(),
);
let should_persist = should_persist_available_local_candidate(&candidate);
let mut attempts = Vec::with_capacity(attempt_slots as usize);
let mut owned_candidate = Some(candidate);
for retry_index in 0..attempt_slots {
let candidate_ref = owned_candidate
.as_ref()
.expect("candidate should remain available until final retry");
let generated_candidate_id = Uuid::new_v4().to_string();
let candidate_id = if should_persist {
state
.persist_available_local_candidate(
trace_id,
context.user_id,
context.api_key_id,
&candidate_ref.candidate,
candidate_index,
retry_index,
generated_candidate_id.as_str(),
context.required_capabilities,
extra_data.clone(),
current_unix_ms(),
context.error_context,
)
.await
} else {
generated_candidate_id
};
let candidate = if retry_index + 1 == attempt_slots {
owned_candidate
.take()
.expect("final retry should consume owned candidate")
} else {
candidate_ref.clone()
};
attempts.push(LocalExecutionCandidateAttempt {
eligible: candidate,
candidate_index,
retry_index,
candidate_id,
});
}
attempts
}
fn build_unpersisted_local_execution_candidate_attempts(
candidate: EligibleLocalExecutionCandidate,
candidate_index: u32,
) -> VecDeque<LocalExecutionCandidateAttempt> {
let attempt_slots = local_attempt_slot_count(&candidate.transport).max(1);
let mut attempts = VecDeque::with_capacity(attempt_slots as usize);
let mut owned_candidate = Some(candidate);
for retry_index in 0..attempt_slots {
let candidate = if retry_index + 1 == attempt_slots {
owned_candidate
.take()
.expect("final retry should consume owned candidate")
} else {
owned_candidate
.as_ref()
.expect("candidate should remain available until final retry")
.clone()
};
attempts.push_back(LocalExecutionCandidateAttempt {
eligible: candidate,
candidate_index,
retry_index,
candidate_id: Uuid::new_v4().to_string(),
});
}
attempts
}
#[allow(clippy::too_many_arguments)]
pub(crate) async fn persist_skipped_local_execution_candidate(
state: &AppState,
@@ -604,6 +858,7 @@ pub(crate) async fn persist_skipped_local_execution_candidates_with_context(
#[cfg(test)]
mod tests {
use std::collections::VecDeque;
use std::sync::Arc;
use aether_data::repository::candidates::InMemoryRequestCandidateRepository;
@@ -707,6 +962,7 @@ mod tests {
pool_key_index: Option<u32>,
) -> EligibleLocalExecutionCandidate {
EligibleLocalExecutionCandidate {
kind: LocalExecutionCandidateKind::SingleKey,
candidate: sample_candidate(key_id),
transport: sample_transport(
key_id,
@@ -722,7 +978,7 @@ mod tests {
}
#[tokio::test]
async fn pool_group_representatives_are_persisted_as_available_before_attempt() {
async fn pool_group_keys_are_not_persisted_as_available_before_attempt() {
let repository = Arc::new(InMemoryRequestCandidateRepository::default());
let app = AppState::new()
.expect("state should build")
@@ -753,11 +1009,9 @@ mod tests {
.read_request_candidates_by_request_id("trace-pool-lazy")
.await
.expect("request candidates should read");
assert_eq!(stored.len(), 2);
assert_eq!(stored[0].key_id.as_deref(), Some("pool-key"));
assert_eq!(stored[0].candidate_index, 0);
assert_eq!(stored[1].key_id.as_deref(), Some("normal-key"));
assert_eq!(stored[1].candidate_index, 2);
assert_eq!(stored.len(), 1);
assert_eq!(stored[0].key_id.as_deref(), Some("normal-key"));
assert_eq!(stored[0].candidate_index, 2);
}
#[tokio::test]
@@ -818,6 +1072,28 @@ mod tests {
assert_eq!(extra_data.get("demoted_by"), Some(&json!("cross_format")));
}
#[tokio::test]
async fn dynamic_attempt_source_does_not_drain_unexecuted_single_keys() {
let mut source = LocalExecutionCandidateAttemptSource {
items: VecDeque::from([LocalExecutionCandidateAttemptSourceItem::Static {
attempts: build_unpersisted_local_execution_candidate_attempts(
sample_eligible("normal-key", None),
0,
),
}]),
};
let first = source
.next_attempt()
.await
.expect("first attempt should be available");
assert_eq!(first.eligible.candidate.key_id, "normal-key");
let remaining = source.drain_static_attempts();
assert!(remaining.is_empty());
assert!(source.next_attempt().await.is_none());
}
#[tokio::test]
async fn pool_internal_skipped_candidates_are_not_persisted() {
let repository = Arc::new(InMemoryRequestCandidateRepository::default());