refactor: lazy pool key scheduling

This commit is contained in:
fawney19
2026-05-11 00:12:05 +08:00
parent 1a0f1a7b72
commit bacb14e5f0
26 changed files with 1856 additions and 223 deletions

View File

@@ -4,7 +4,8 @@ use async_trait::async_trait;
use super::{
MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow,
StoredPoolKeyCandidateRowsQuery, StoredRequestedModelCandidateRowsQuery,
StoredPoolKeyCandidateOrder, StoredPoolKeyCandidateRowsQuery,
StoredRequestedModelCandidateRowsQuery,
};
use crate::DataLayerError;
@@ -130,11 +131,7 @@ impl MinimalCandidateSelectionReadRepository for InMemoryMinimalCandidateSelecti
&& row.model_id == query.model_id
})
.collect::<Vec<_>>();
rows.sort_by(|left, right| {
left.key_internal_priority
.cmp(&right.key_internal_priority)
.then(left.key_id.cmp(&right.key_id))
});
sort_pool_key_rows(&mut rows, &query.order);
Ok(rows
.into_iter()
.skip(query.offset as usize)
@@ -143,6 +140,38 @@ impl MinimalCandidateSelectionReadRepository for InMemoryMinimalCandidateSelecti
}
}
fn sort_pool_key_rows(
rows: &mut [StoredMinimalCandidateSelectionRow],
order: &StoredPoolKeyCandidateOrder,
) {
rows.sort_by(|left, right| match order {
StoredPoolKeyCandidateOrder::LoadBalance { seed } => {
stable_pool_key_hash(seed.as_str(), left.key_id.as_str())
.cmp(&stable_pool_key_hash(seed.as_str(), right.key_id.as_str()))
.then(left.key_id.cmp(&right.key_id))
}
_ => left
.key_internal_priority
.cmp(&right.key_internal_priority)
.then(left.key_id.cmp(&right.key_id)),
});
}
fn stable_pool_key_hash(seed: &str, key_id: &str) -> u64 {
let mut hash = 0xcbf29ce484222325u64;
for byte in seed
.as_bytes()
.iter()
.copied()
.chain(std::iter::once(b':'))
.chain(key_id.as_bytes().iter().copied())
{
hash ^= u64::from(byte);
hash = hash.wrapping_mul(0x100000001b3);
}
hash
}
fn normalize_api_format(value: &str) -> String {
aether_ai_formats::normalize_api_format_alias(value)
}
@@ -215,7 +244,8 @@ mod tests {
use super::InMemoryMinimalCandidateSelectionReadRepository;
use crate::repository::candidate_selection::{
MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow,
StoredPoolKeyCandidateRowsQuery, StoredRequestedModelCandidateRowsQuery,
StoredPoolKeyCandidateOrder, StoredPoolKeyCandidateRowsQuery,
StoredRequestedModelCandidateRowsQuery,
};
fn sample_row(
@@ -402,6 +432,7 @@ mod tests {
endpoint_id: "endpoint-pool".to_string(),
model_id: "model-pool".to_string(),
selected_provider_model_name: "gpt-5".to_string(),
order: StoredPoolKeyCandidateOrder::InternalPriority,
offset: 2,
limit: 2,
})

View File

@@ -6,8 +6,9 @@ mod sqlite;
#[allow(unused_imports)]
pub(crate) use aether_data_contracts::repository::candidate_selection::{
MinimalCandidateSelectionReadRepository, MinimalCandidateSelectionRepository,
StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateRowsQuery,
StoredProviderModelMapping, StoredRequestedModelCandidateRowsQuery,
StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateOrder,
StoredPoolKeyCandidateRowsQuery, StoredProviderModelMapping,
StoredRequestedModelCandidateRowsQuery,
};
pub use memory::InMemoryMinimalCandidateSelectionReadRepository;
pub use mysql::MysqlMinimalCandidateSelectionReadRepository;

View File

@@ -5,7 +5,7 @@ use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row};
use super::{
MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow,
StoredPoolKeyCandidateRowsQuery, StoredProviderModelMapping,
StoredPoolKeyCandidateOrder, StoredPoolKeyCandidateRowsQuery, StoredProviderModelMapping,
StoredRequestedModelCandidateRowsQuery,
};
use crate::driver::mysql::MysqlPool;
@@ -35,6 +35,7 @@ SELECT
pak.capabilities AS key_capabilities,
pak.internal_priority AS key_internal_priority,
pak.global_priority_by_format AS key_global_priority_by_format,
pak.last_used_at AS key_last_used_at_unix_secs,
m.id AS model_id,
m.global_model_id AS global_model_id,
gm.name AS global_model_name,
@@ -67,6 +68,7 @@ struct CandidateSelectionRow {
row: StoredMinimalCandidateSelectionRow,
provider_pool_enabled: bool,
key_auth_config: Option<String>,
key_last_used_at_unix_secs: Option<u64>,
}
impl MysqlMinimalCandidateSelectionReadRepository {
@@ -180,18 +182,18 @@ impl MinimalCandidateSelectionReadRepository for MysqlMinimalCandidateSelectionR
.load_rows_for_api_format(&query.api_format)
.await?
.into_iter()
.map(|item| item.row)
.filter(|row| {
row.provider_id == query.provider_id
&& row.endpoint_id == query.endpoint_id
&& row.model_id == query.model_id
row.row.provider_id == query.provider_id
&& row.row.endpoint_id == query.endpoint_id
&& row.row.model_id == query.model_id
})
.collect::<Vec<_>>();
let mut rows = sort_pool_key_rows(rows);
let mut rows = sort_pool_key_rows(rows, &query.order);
Ok(rows
.drain(..)
.skip(query.offset as usize)
.take(query.limit as usize)
.map(|item| item.row)
.collect())
}
}
@@ -246,16 +248,66 @@ fn sort_rows(
}
fn sort_pool_key_rows(
mut rows: Vec<StoredMinimalCandidateSelectionRow>,
) -> Vec<StoredMinimalCandidateSelectionRow> {
rows.sort_by(|left, right| {
left.key_internal_priority
.cmp(&right.key_internal_priority)
.then(left.key_id.cmp(&right.key_id))
mut rows: Vec<CandidateSelectionRow>,
order: &StoredPoolKeyCandidateOrder,
) -> Vec<CandidateSelectionRow> {
rows.sort_by(|left, right| match order {
StoredPoolKeyCandidateOrder::InternalPriority => compare_pool_key_internal(left, right),
StoredPoolKeyCandidateOrder::Lru => left
.key_last_used_at_unix_secs
.cmp(&right.key_last_used_at_unix_secs)
.then_with(|| compare_pool_key_internal(left, right)),
StoredPoolKeyCandidateOrder::CacheAffinity => right
.key_last_used_at_unix_secs
.cmp(&left.key_last_used_at_unix_secs)
.then_with(|| compare_pool_key_internal(left, right)),
StoredPoolKeyCandidateOrder::SingleAccount => left
.row
.key_internal_priority
.cmp(&right.row.key_internal_priority)
.then_with(|| {
right
.key_last_used_at_unix_secs
.cmp(&left.key_last_used_at_unix_secs)
})
.then(left.row.key_id.cmp(&right.row.key_id)),
StoredPoolKeyCandidateOrder::LoadBalance { seed } => {
stable_pool_key_hash(seed.as_str(), left.row.key_id.as_str())
.cmp(&stable_pool_key_hash(
seed.as_str(),
right.row.key_id.as_str(),
))
.then(left.row.key_id.cmp(&right.row.key_id))
}
});
rows
}
fn compare_pool_key_internal(
left: &CandidateSelectionRow,
right: &CandidateSelectionRow,
) -> std::cmp::Ordering {
left.row
.key_internal_priority
.cmp(&right.row.key_internal_priority)
.then(left.row.key_id.cmp(&right.row.key_id))
}
fn stable_pool_key_hash(seed: &str, key_id: &str) -> u64 {
let mut hash = 0xcbf29ce484222325u64;
for byte in seed
.as_bytes()
.iter()
.copied()
.chain(std::iter::once(b':'))
.chain(key_id.as_bytes().iter().copied())
{
hash ^= u64::from(byte);
hash = hash.wrapping_mul(0x100000001b3);
}
hash
}
fn row_matches_requested_model(
row: &StoredMinimalCandidateSelectionRow,
requested_model_name: &str,
@@ -394,6 +446,10 @@ fn map_candidate_selection_row(row: &MySqlRow) -> Result<CandidateSelectionRow,
},
provider_pool_enabled,
key_auth_config: row.try_get("key_auth_config").map_sql_err()?,
key_last_used_at_unix_secs: row
.try_get::<Option<i64>, _>("key_last_used_at_unix_secs")
.map_sql_err()?
.and_then(|value| u64::try_from(value).ok()),
})
}

View File

@@ -5,7 +5,7 @@ use std::collections::BTreeSet;
use super::{
MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow,
StoredPoolKeyCandidateRowsQuery, StoredProviderModelMapping,
StoredPoolKeyCandidateOrder, StoredPoolKeyCandidateRowsQuery, StoredProviderModelMapping,
StoredRequestedModelCandidateRowsQuery,
};
use crate::{error::SqlxResultExt, DataLayerError};
@@ -505,6 +505,36 @@ LIMIT $7
OFFSET $8
"#;
fn pool_key_candidate_order_by_sql(order: &StoredPoolKeyCandidateOrder) -> &'static str {
match order {
StoredPoolKeyCandidateOrder::InternalPriority => {
"ORDER BY\n pak.internal_priority ASC,\n pak.id ASC"
}
StoredPoolKeyCandidateOrder::Lru => {
"ORDER BY\n pak.last_used_at ASC NULLS FIRST,\n pak.internal_priority ASC,\n pak.id ASC"
}
StoredPoolKeyCandidateOrder::CacheAffinity => {
"ORDER BY\n pak.last_used_at DESC NULLS LAST,\n pak.internal_priority ASC,\n pak.id ASC"
}
StoredPoolKeyCandidateOrder::SingleAccount => {
"ORDER BY\n pak.internal_priority ASC,\n pak.last_used_at DESC NULLS LAST,\n pak.id ASC"
}
StoredPoolKeyCandidateOrder::LoadBalance { .. } => {
"ORDER BY\n md5($9 || ':' || pak.id) ASC,\n pak.id ASC"
}
}
}
fn pool_key_candidate_selection_sql(order: &StoredPoolKeyCandidateOrder) -> String {
let default_order =
"ORDER BY\n pak.internal_priority ASC,\n pak.id ASC\nLIMIT $7\nOFFSET $8\n";
let replacement = format!(
"{}\nLIMIT $7\nOFFSET $8\n",
pool_key_candidate_order_by_sql(order)
);
LIST_POOL_KEYS_FOR_GROUP_SQL.replace(default_order, &replacement)
}
#[derive(Debug, Clone)]
pub struct SqlxMinimalCandidateSelectionReadRepository {
pool: PgPool,
@@ -650,19 +680,23 @@ impl SqlxMinimalCandidateSelectionReadRepository {
let sql_match_aliases = sql_match_aliases(&storage_aliases);
let limit = i64::from(query.limit.max(1));
let offset = i64::from(query.offset);
let sql = pool_key_candidate_selection_sql(&query.order);
for api_format in storage_aliases {
let mut query_builder = sqlx::query(sql.as_str())
.bind(api_format)
.bind(query.provider_id.as_str())
.bind(query.endpoint_id.as_str())
.bind(query.model_id.as_str())
.bind(sql_match_aliases.clone())
.bind(canonical_api_format.clone())
.bind(limit)
.bind(offset);
if let StoredPoolKeyCandidateOrder::LoadBalance { seed } = &query.order {
query_builder = query_builder.bind(seed.as_str());
}
rows.extend(
Self::collect_query_rows(
sqlx::query(LIST_POOL_KEYS_FOR_GROUP_SQL)
.bind(api_format)
.bind(query.provider_id.as_str())
.bind(query.endpoint_id.as_str())
.bind(query.model_id.as_str())
.bind(sql_match_aliases.clone())
.bind(canonical_api_format.clone())
.bind(limit)
.bind(offset)
.fetch(&self.pool),
query_builder.fetch(&self.pool),
map_candidate_selection_row,
)
.await?,
@@ -1039,13 +1073,16 @@ mod tests {
use serde_json::json;
use super::{
parse_provider_model_mappings, parse_string_list, requested_model_selection_page_sql,
requested_model_selection_sql, SqlxMinimalCandidateSelectionReadRepository,
parse_provider_model_mappings, parse_string_list, pool_key_candidate_selection_sql,
requested_model_selection_page_sql, requested_model_selection_sql,
SqlxMinimalCandidateSelectionReadRepository,
LIST_FOR_EXACT_API_FORMAT_AND_GLOBAL_MODEL_SQL, LIST_FOR_EXACT_API_FORMAT_SQL,
LIST_POOL_KEYS_FOR_GROUP_SQL,
};
use crate::driver::postgres::{PostgresPoolConfig, PostgresPoolFactory};
use crate::repository::candidate_selection::StoredProviderModelMapping;
use crate::repository::candidate_selection::{
StoredPoolKeyCandidateOrder, StoredProviderModelMapping,
};
#[tokio::test]
async fn repository_constructs_from_lazy_pool() {
@@ -1100,6 +1137,23 @@ mod tests {
assert!(sql.ends_with("LIMIT $5\nOFFSET $6"));
}
#[test]
fn pool_key_selection_sql_applies_query_order() {
let load_balance_sql =
pool_key_candidate_selection_sql(&StoredPoolKeyCandidateOrder::LoadBalance {
seed: "seed".to_string(),
});
assert!(load_balance_sql.contains("md5($9 || ':' || pak.id) ASC"));
assert!(load_balance_sql.ends_with("LIMIT $7\nOFFSET $8\n"));
let lru_sql = pool_key_candidate_selection_sql(&StoredPoolKeyCandidateOrder::Lru);
assert!(lru_sql.contains("pak.last_used_at ASC NULLS FIRST"));
let cache_affinity_sql =
pool_key_candidate_selection_sql(&StoredPoolKeyCandidateOrder::CacheAffinity);
assert!(cache_affinity_sql.contains("pak.last_used_at DESC NULLS LAST"));
}
#[test]
fn parse_string_list_accepts_stringified_array() {
let parsed = parse_string_list(

View File

@@ -5,7 +5,7 @@ use sqlx::{sqlite::SqliteRow, QueryBuilder, Row, Sqlite};
use super::{
MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow,
StoredPoolKeyCandidateRowsQuery, StoredProviderModelMapping,
StoredPoolKeyCandidateOrder, StoredPoolKeyCandidateRowsQuery, StoredProviderModelMapping,
StoredRequestedModelCandidateRowsQuery,
};
use crate::driver::sqlite::SqlitePool;
@@ -35,6 +35,7 @@ SELECT
pak.capabilities AS key_capabilities,
pak.internal_priority AS key_internal_priority,
pak.global_priority_by_format AS key_global_priority_by_format,
pak.last_used_at AS key_last_used_at_unix_secs,
m.id AS model_id,
m.global_model_id AS global_model_id,
gm.name AS global_model_name,
@@ -67,6 +68,7 @@ struct CandidateSelectionRow {
row: StoredMinimalCandidateSelectionRow,
provider_pool_enabled: bool,
key_auth_config: Option<String>,
key_last_used_at_unix_secs: Option<u64>,
}
impl SqliteMinimalCandidateSelectionReadRepository {
@@ -180,18 +182,18 @@ impl MinimalCandidateSelectionReadRepository for SqliteMinimalCandidateSelection
.load_rows_for_api_format(&query.api_format)
.await?
.into_iter()
.map(|item| item.row)
.filter(|row| {
row.provider_id == query.provider_id
&& row.endpoint_id == query.endpoint_id
&& row.model_id == query.model_id
row.row.provider_id == query.provider_id
&& row.row.endpoint_id == query.endpoint_id
&& row.row.model_id == query.model_id
})
.collect::<Vec<_>>();
let mut rows = sort_pool_key_rows(rows);
let mut rows = sort_pool_key_rows(rows, &query.order);
Ok(rows
.drain(..)
.skip(query.offset as usize)
.take(query.limit as usize)
.map(|item| item.row)
.collect())
}
}
@@ -246,16 +248,66 @@ fn sort_rows(
}
fn sort_pool_key_rows(
mut rows: Vec<StoredMinimalCandidateSelectionRow>,
) -> Vec<StoredMinimalCandidateSelectionRow> {
rows.sort_by(|left, right| {
left.key_internal_priority
.cmp(&right.key_internal_priority)
.then(left.key_id.cmp(&right.key_id))
mut rows: Vec<CandidateSelectionRow>,
order: &StoredPoolKeyCandidateOrder,
) -> Vec<CandidateSelectionRow> {
rows.sort_by(|left, right| match order {
StoredPoolKeyCandidateOrder::InternalPriority => compare_pool_key_internal(left, right),
StoredPoolKeyCandidateOrder::Lru => left
.key_last_used_at_unix_secs
.cmp(&right.key_last_used_at_unix_secs)
.then_with(|| compare_pool_key_internal(left, right)),
StoredPoolKeyCandidateOrder::CacheAffinity => right
.key_last_used_at_unix_secs
.cmp(&left.key_last_used_at_unix_secs)
.then_with(|| compare_pool_key_internal(left, right)),
StoredPoolKeyCandidateOrder::SingleAccount => left
.row
.key_internal_priority
.cmp(&right.row.key_internal_priority)
.then_with(|| {
right
.key_last_used_at_unix_secs
.cmp(&left.key_last_used_at_unix_secs)
})
.then(left.row.key_id.cmp(&right.row.key_id)),
StoredPoolKeyCandidateOrder::LoadBalance { seed } => {
stable_pool_key_hash(seed.as_str(), left.row.key_id.as_str())
.cmp(&stable_pool_key_hash(
seed.as_str(),
right.row.key_id.as_str(),
))
.then(left.row.key_id.cmp(&right.row.key_id))
}
});
rows
}
fn compare_pool_key_internal(
left: &CandidateSelectionRow,
right: &CandidateSelectionRow,
) -> std::cmp::Ordering {
left.row
.key_internal_priority
.cmp(&right.row.key_internal_priority)
.then(left.row.key_id.cmp(&right.row.key_id))
}
fn stable_pool_key_hash(seed: &str, key_id: &str) -> u64 {
let mut hash = 0xcbf29ce484222325u64;
for byte in seed
.as_bytes()
.iter()
.copied()
.chain(std::iter::once(b':'))
.chain(key_id.as_bytes().iter().copied())
{
hash ^= u64::from(byte);
hash = hash.wrapping_mul(0x100000001b3);
}
hash
}
fn row_matches_requested_model(
row: &StoredMinimalCandidateSelectionRow,
requested_model_name: &str,
@@ -394,6 +446,10 @@ fn map_candidate_selection_row(row: &SqliteRow) -> Result<CandidateSelectionRow,
},
provider_pool_enabled,
key_auth_config: row.try_get("key_auth_config").map_sql_err()?,
key_last_used_at_unix_secs: row
.try_get::<Option<i64>, _>("key_last_used_at_unix_secs")
.map_sql_err()?
.and_then(|value| u64::try_from(value).ok()),
})
}
@@ -620,8 +676,8 @@ mod tests {
use super::SqliteMinimalCandidateSelectionReadRepository;
use crate::lifecycle::migrate::run_sqlite_migrations;
use crate::repository::candidate_selection::{
MinimalCandidateSelectionReadRepository, StoredPoolKeyCandidateRowsQuery,
StoredRequestedModelCandidateRowsQuery,
MinimalCandidateSelectionReadRepository, StoredPoolKeyCandidateOrder,
StoredPoolKeyCandidateRowsQuery, StoredRequestedModelCandidateRowsQuery,
};
#[tokio::test]
@@ -673,6 +729,7 @@ mod tests {
endpoint_id: "endpoint-1".to_string(),
model_id: "model-1".to_string(),
selected_provider_model_name: "provider-model".to_string(),
order: StoredPoolKeyCandidateOrder::InternalPriority,
offset: 1,
limit: 1,
})