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

View File

@@ -1,5 +1,8 @@
use aether_data::DataLayerError;
use aether_data_contracts::repository::candidate_selection::StoredMinimalCandidateSelectionRow;
use aether_data_contracts::repository::candidate_selection::{
StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateRowsQuery,
StoredRequestedModelCandidateRowsQuery,
};
use aether_scheduler_core::{
auth_constraints_allow_api_format, collect_global_model_names_for_required_capability,
enumerate_minimal_candidate_selection_with_model_directives, normalize_api_format,
@@ -20,32 +23,61 @@ pub(crate) trait MinimalCandidateSelectionRowSource {
global_model_name: &str,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError>;
async fn read_minimal_candidate_selection_rows_for_api_format_and_requested_model(
&self,
api_format: &str,
requested_model_name: &str,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError>;
async fn read_minimal_candidate_selection_rows_for_api_format_and_requested_model_page(
&self,
query: &StoredRequestedModelCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError>;
async fn read_minimal_candidate_selection_rows_for_api_format(
&self,
api_format: &str,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError>;
async fn read_pool_key_candidate_rows_for_group(
&self,
query: &StoredPoolKeyCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError>;
}
const REQUESTED_MODEL_CANDIDATE_PAGE_SIZE: u32 = 256;
const REQUESTED_MODEL_MAX_SCANNED_ROWS: u32 = 2048;
pub(crate) async fn read_requested_model_rows(
state: &(impl MinimalCandidateSelectionRowSource + Sync),
api_format: &str,
requested_model_name: &str,
enable_model_directives: bool,
) -> Result<Option<(String, Vec<StoredMinimalCandidateSelectionRow>)>, DataLayerError> {
let rows = state
.read_minimal_candidate_selection_rows_for_api_format(api_format)
.await?;
let rows = rows
.into_iter()
.filter(|row| {
row_supports_requested_model_with_model_directives(
row,
requested_model_name,
api_format,
enable_model_directives,
)
})
.collect::<Vec<_>>();
let fast_rows = read_requested_model_rows_fast_path(
state,
api_format,
requested_model_name,
enable_model_directives,
)
.await?;
let mut rows = filter_rows_for_requested_model(
fast_rows,
requested_model_name,
api_format,
enable_model_directives,
);
if rows.is_empty() {
let fallback_rows = state
.read_minimal_candidate_selection_rows_for_api_format(api_format)
.await?;
rows = filter_rows_for_requested_model(
fallback_rows,
requested_model_name,
api_format,
enable_model_directives,
);
}
if rows.is_empty() {
return Ok(None);
}
@@ -69,6 +101,85 @@ pub(crate) async fn read_requested_model_rows(
Ok(Some((resolved_global_model_name, resolved_rows)))
}
fn filter_rows_for_requested_model(
rows: Vec<StoredMinimalCandidateSelectionRow>,
requested_model_name: &str,
api_format: &str,
enable_model_directives: bool,
) -> Vec<StoredMinimalCandidateSelectionRow> {
rows.into_iter()
.filter(|row| {
row_supports_requested_model_with_model_directives(
row,
requested_model_name,
api_format,
enable_model_directives,
)
})
.collect()
}
async fn read_requested_model_rows_fast_path(
state: &(impl MinimalCandidateSelectionRowSource + Sync),
api_format: &str,
requested_model_name: &str,
enable_model_directives: bool,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
let mut requested_names = vec![requested_model_name.trim().to_string()];
if enable_model_directives {
if let Some(base_model) =
crate::ai_serving::model_directive_base_model(requested_model_name)
{
if !requested_names.iter().any(|value| value == &base_model) {
requested_names.push(base_model);
}
}
}
let mut rows = Vec::new();
let mut seen = BTreeSet::new();
for requested_name in requested_names {
if requested_name.is_empty() {
continue;
}
let mut offset = 0;
let mut scanned = 0;
while scanned < REQUESTED_MODEL_MAX_SCANNED_ROWS {
let limit =
REQUESTED_MODEL_CANDIDATE_PAGE_SIZE.min(REQUESTED_MODEL_MAX_SCANNED_ROWS - scanned);
let page = state
.read_minimal_candidate_selection_rows_for_api_format_and_requested_model_page(
&StoredRequestedModelCandidateRowsQuery {
api_format: api_format.to_string(),
requested_model_name: requested_name.clone(),
offset,
limit,
},
)
.await?;
if page.is_empty() {
break;
}
let page_len = page.len() as u32;
for row in page {
if seen.insert((
row.endpoint_id.clone(),
row.key_id.clone(),
row.model_id.clone(),
)) {
rows.push(row);
}
}
scanned = scanned.saturating_add(page_len);
if page_len < limit {
break;
}
offset = offset.saturating_add(limit);
}
}
Ok(rows)
}
pub(crate) async fn enumerate_minimal_candidate_selection_with_required_capabilities(
state: &(impl MinimalCandidateSelectionRowSource + Sync),
api_format: &str,
@@ -210,3 +321,180 @@ fn auth_snapshot_constraints(snapshot: &GatewayAuthApiKeySnapshot) -> SchedulerA
.map(|items| items.to_vec()),
}
}
#[cfg(test)]
mod tests {
use super::{
read_requested_model_rows, MinimalCandidateSelectionRowSource,
StoredMinimalCandidateSelectionRow,
};
use aether_data::DataLayerError;
use aether_data_contracts::repository::candidate_selection::{
StoredPoolKeyCandidateRowsQuery, StoredRequestedModelCandidateRowsQuery,
};
use async_trait::async_trait;
use std::sync::atomic::{AtomicUsize, Ordering};
struct CountingSelectionSource {
fast_rows: Vec<StoredMinimalCandidateSelectionRow>,
fallback_rows: Vec<StoredMinimalCandidateSelectionRow>,
fast_calls: AtomicUsize,
fallback_calls: AtomicUsize,
}
impl CountingSelectionSource {
fn new(
fast_rows: Vec<StoredMinimalCandidateSelectionRow>,
fallback_rows: Vec<StoredMinimalCandidateSelectionRow>,
) -> Self {
Self {
fast_rows,
fallback_rows,
fast_calls: AtomicUsize::new(0),
fallback_calls: AtomicUsize::new(0),
}
}
}
#[async_trait]
impl MinimalCandidateSelectionRowSource for CountingSelectionSource {
async fn read_minimal_candidate_selection_rows_for_api_format_and_global_model(
&self,
_api_format: &str,
_global_model_name: &str,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
Ok(Vec::new())
}
async fn read_minimal_candidate_selection_rows_for_api_format_and_requested_model(
&self,
_api_format: &str,
_requested_model_name: &str,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
Ok(Vec::new())
}
async fn read_minimal_candidate_selection_rows_for_api_format_and_requested_model_page(
&self,
query: &StoredRequestedModelCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
self.fast_calls.fetch_add(1, Ordering::SeqCst);
Ok(self
.fast_rows
.iter()
.skip(query.offset as usize)
.take(query.limit as usize)
.cloned()
.collect())
}
async fn read_minimal_candidate_selection_rows_for_api_format(
&self,
_api_format: &str,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
self.fallback_calls.fetch_add(1, Ordering::SeqCst);
Ok(self.fallback_rows.clone())
}
async fn read_pool_key_candidate_rows_for_group(
&self,
_query: &StoredPoolKeyCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
Ok(Vec::new())
}
}
fn sample_row(global_model_name: &str) -> StoredMinimalCandidateSelectionRow {
StoredMinimalCandidateSelectionRow {
provider_id: "provider-1".to_string(),
provider_name: "provider".to_string(),
provider_type: "custom".to_string(),
provider_priority: 10,
provider_is_active: true,
endpoint_id: "endpoint-1".to_string(),
endpoint_api_format: "openai:chat".to_string(),
endpoint_api_family: Some("openai".to_string()),
endpoint_kind: Some("chat".to_string()),
endpoint_is_active: true,
key_id: "key-1".to_string(),
key_name: "key".to_string(),
key_auth_type: "api_key".to_string(),
key_is_active: true,
key_api_formats: Some(vec!["openai:chat".to_string()]),
key_allowed_models: None,
key_capabilities: None,
key_internal_priority: 10,
key_global_priority_by_format: None,
model_id: "model-1".to_string(),
global_model_id: "global-model-1".to_string(),
global_model_name: global_model_name.to_string(),
global_model_mappings: None,
global_model_supports_streaming: Some(true),
model_provider_model_name: global_model_name.to_string(),
model_provider_model_mappings: None,
model_supports_streaming: Some(true),
model_is_active: true,
model_is_available: true,
}
}
#[tokio::test]
async fn requested_model_rows_use_fast_path_without_full_format_scan() {
let source = CountingSelectionSource::new(vec![sample_row("gpt-5")], Vec::new());
let result = read_requested_model_rows(&source, "openai:chat", "gpt-5", false)
.await
.expect("read should succeed")
.expect("rows should resolve");
assert_eq!(result.0, "gpt-5");
assert_eq!(result.1.len(), 1);
assert_eq!(source.fast_calls.load(Ordering::SeqCst), 1);
assert_eq!(source.fallback_calls.load(Ordering::SeqCst), 0);
}
#[tokio::test]
async fn requested_model_rows_fall_back_to_full_format_scan_when_fast_path_misses() {
let source = CountingSelectionSource::new(Vec::new(), vec![sample_row("gpt-5")]);
let result = read_requested_model_rows(&source, "openai:chat", "gpt-5", false)
.await
.expect("read should succeed")
.expect("rows should resolve");
assert_eq!(result.0, "gpt-5");
assert_eq!(result.1.len(), 1);
assert_eq!(source.fast_calls.load(Ordering::SeqCst), 1);
assert_eq!(source.fallback_calls.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn requested_model_rows_fast_path_stops_at_scan_limit() {
let mut rows = Vec::new();
for index in 0..(super::REQUESTED_MODEL_MAX_SCANNED_ROWS + 5) {
let mut row = sample_row("gpt-5");
row.provider_id = format!("provider-{index}");
row.endpoint_id = format!("endpoint-{index}");
row.key_id = format!("key-{index}");
row.model_id = format!("model-{index}");
rows.push(row);
}
let source = CountingSelectionSource::new(rows, Vec::new());
let result = read_requested_model_rows(&source, "openai:chat", "gpt-5", false)
.await
.expect("read should succeed")
.expect("rows should resolve");
assert_eq!(
result.1.len(),
super::REQUESTED_MODEL_MAX_SCANNED_ROWS as usize
);
assert_eq!(
source.fast_calls.load(Ordering::SeqCst),
(super::REQUESTED_MODEL_MAX_SCANNED_ROWS / super::REQUESTED_MODEL_CANDIDATE_PAGE_SIZE)
as usize
);
assert_eq!(source.fallback_calls.load(Ordering::SeqCst), 0);
}
}

View File

@@ -146,6 +146,26 @@ impl MinimalCandidateSelectionRowSource for GatewayDataState {
.await
}
async fn read_minimal_candidate_selection_rows_for_api_format_and_requested_model(
&self,
api_format: &str,
requested_model_name: &str,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
self.list_minimal_candidate_selection_rows_for_requested_model(
api_format,
requested_model_name,
)
.await
}
async fn read_minimal_candidate_selection_rows_for_api_format_and_requested_model_page(
&self,
query: &aether_data_contracts::repository::candidate_selection::StoredRequestedModelCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
self.list_minimal_candidate_selection_rows_for_requested_model_page(query)
.await
}
async fn read_minimal_candidate_selection_rows_for_api_format(
&self,
api_format: &str,
@@ -153,6 +173,13 @@ impl MinimalCandidateSelectionRowSource for GatewayDataState {
self.list_minimal_candidate_selection_rows_for_api_format(api_format)
.await
}
async fn read_pool_key_candidate_rows_for_group(
&self,
query: &aether_data_contracts::repository::candidate_selection::StoredPoolKeyCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
self.list_pool_key_candidate_rows_for_group(query).await
}
}
#[async_trait]

View File

@@ -77,6 +77,7 @@ use aether_data_contracts::repository::billing::{
};
use aether_data_contracts::repository::candidate_selection::{
MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow,
StoredPoolKeyCandidateRowsQuery, StoredRequestedModelCandidateRowsQuery,
};
use aether_data_contracts::repository::candidates::{
PublicHealthStatusCount, PublicHealthTimelineBucket, RequestCandidateReadRepository,

View File

@@ -2,9 +2,10 @@ use super::{
AdminGlobalModelListQuery, AdminProviderModelListQuery, CreateAdminGlobalModelRecord,
DataLayerError, GatewayDataState, PublicCatalogModelListQuery, PublicCatalogModelSearchQuery,
PublicGlobalModelQuery, StoredAdminGlobalModel, StoredAdminGlobalModelPage,
StoredAdminProviderModel, StoredMinimalCandidateSelectionRow, StoredProviderActiveGlobalModel,
StoredProviderModelStats, StoredPublicCatalogModel, StoredPublicGlobalModel,
StoredPublicGlobalModelPage, UpdateAdminGlobalModelRecord, UpsertAdminProviderModelRecord,
StoredAdminProviderModel, StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateRowsQuery,
StoredProviderActiveGlobalModel, StoredProviderModelStats, StoredPublicCatalogModel,
StoredPublicGlobalModel, StoredPublicGlobalModelPage, StoredRequestedModelCandidateRowsQuery,
UpdateAdminGlobalModelRecord, UpsertAdminProviderModelRecord,
};
impl GatewayDataState {
@@ -23,6 +24,35 @@ impl GatewayDataState {
}
}
pub(crate) async fn list_minimal_candidate_selection_rows_for_requested_model(
&self,
api_format: &str,
requested_model_name: &str,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
match &self.minimal_candidate_selection_reader {
Some(repository) => {
repository
.list_for_exact_api_format_and_requested_model(api_format, requested_model_name)
.await
}
None => Ok(Vec::new()),
}
}
pub(crate) async fn list_minimal_candidate_selection_rows_for_requested_model_page(
&self,
query: &StoredRequestedModelCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
match &self.minimal_candidate_selection_reader {
Some(repository) => {
repository
.list_for_exact_api_format_and_requested_model_page(query)
.await
}
None => Ok(Vec::new()),
}
}
pub(crate) async fn list_minimal_candidate_selection_rows_for_api_format(
&self,
api_format: &str,
@@ -33,6 +63,16 @@ impl GatewayDataState {
}
}
pub(crate) async fn list_pool_key_candidate_rows_for_group(
&self,
query: &StoredPoolKeyCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
match &self.minimal_candidate_selection_reader {
Some(repository) => repository.list_pool_key_rows_for_group(query).await,
None => Ok(Vec::new()),
}
}
pub(crate) async fn list_public_global_models(
&self,
query: &PublicGlobalModelQuery,