Merge branch 'fawney19:main' into main

This commit is contained in:
ZheFox
2026-07-28 13:56:34 +08:00
committed by GitHub
117 changed files with 10467 additions and 1253 deletions
@@ -0,0 +1,115 @@
UPDATE routing_groups
SET is_system_default = 0
WHERE id IN (
SELECT id
FROM (
SELECT
id,
ROW_NUMBER() OVER (ORDER BY enabled DESC, updated_at DESC, id ASC) AS default_rank
FROM routing_groups
WHERE is_system_default = 1
) AS ranked_defaults
WHERE default_rank > 1
);
UPDATE routing_group_bindings
SET is_default = 0
WHERE id IN (
SELECT id
FROM (
SELECT
id,
ROW_NUMBER() OVER (
PARTITION BY subject_type, subject_id
ORDER BY created_at ASC, id ASC
) AS default_rank
FROM routing_group_bindings
WHERE is_default = 1
) AS ranked_defaults
WHERE default_rank > 1
);
SET @aether_routing_groups_default_guard_column_sql := IF(
(
SELECT COUNT(*)
FROM information_schema.columns
WHERE table_schema = DATABASE()
AND table_name = 'routing_groups'
AND column_name = 'system_default_unique_guard'
) = 0,
'ALTER TABLE routing_groups ADD COLUMN system_default_unique_guard TINYINT GENERATED ALWAYS AS (CASE WHEN is_system_default = 1 THEN 1 ELSE NULL END) VIRTUAL INVISIBLE',
'DO 0'
);
PREPARE aether_routing_groups_default_guard_column_stmt
FROM @aether_routing_groups_default_guard_column_sql;
EXECUTE aether_routing_groups_default_guard_column_stmt;
DEALLOCATE PREPARE aether_routing_groups_default_guard_column_stmt;
SET @aether_routing_bindings_default_type_guard_column_sql := IF(
(
SELECT COUNT(*)
FROM information_schema.columns
WHERE table_schema = DATABASE()
AND table_name = 'routing_group_bindings'
AND column_name = 'default_subject_type_guard'
) = 0,
'ALTER TABLE routing_group_bindings ADD COLUMN default_subject_type_guard VARCHAR(32) GENERATED ALWAYS AS (CASE WHEN is_default = 1 THEN subject_type ELSE NULL END) VIRTUAL INVISIBLE',
'DO 0'
);
PREPARE aether_routing_bindings_default_type_guard_column_stmt
FROM @aether_routing_bindings_default_type_guard_column_sql;
EXECUTE aether_routing_bindings_default_type_guard_column_stmt;
DEALLOCATE PREPARE aether_routing_bindings_default_type_guard_column_stmt;
SET @aether_routing_bindings_default_id_guard_column_sql := IF(
(
SELECT COUNT(*)
FROM information_schema.columns
WHERE table_schema = DATABASE()
AND table_name = 'routing_group_bindings'
AND column_name = 'default_subject_id_guard'
) = 0,
'ALTER TABLE routing_group_bindings ADD COLUMN default_subject_id_guard VARCHAR(64) GENERATED ALWAYS AS (CASE WHEN is_default = 1 THEN subject_id ELSE NULL END) VIRTUAL INVISIBLE',
'DO 0'
);
PREPARE aether_routing_bindings_default_id_guard_column_stmt
FROM @aether_routing_bindings_default_id_guard_column_sql;
EXECUTE aether_routing_bindings_default_id_guard_column_stmt;
DEALLOCATE PREPARE aether_routing_bindings_default_id_guard_column_stmt;
SET @aether_routing_groups_default_unique_index_sql := IF(
(
SELECT COUNT(*)
FROM information_schema.statistics
WHERE table_schema = DATABASE()
AND table_name = 'routing_groups'
AND index_name = 'routing_groups_one_system_default_key'
) = 0,
'CREATE UNIQUE INDEX routing_groups_one_system_default_key ON routing_groups (system_default_unique_guard)',
'DO 0'
);
PREPARE aether_routing_groups_default_unique_index_stmt
FROM @aether_routing_groups_default_unique_index_sql;
EXECUTE aether_routing_groups_default_unique_index_stmt;
DEALLOCATE PREPARE aether_routing_groups_default_unique_index_stmt;
SET @aether_routing_bindings_default_unique_index_sql := IF(
(
SELECT COUNT(*)
FROM information_schema.statistics
WHERE table_schema = DATABASE()
AND table_name = 'routing_group_bindings'
AND index_name = 'routing_group_bindings_subject_default_key'
) = 0,
'CREATE UNIQUE INDEX routing_group_bindings_subject_default_key ON routing_group_bindings (default_subject_type_guard, default_subject_id_guard)',
'DO 0'
);
PREPARE aether_routing_bindings_default_unique_index_stmt
FROM @aether_routing_bindings_default_unique_index_sql;
EXECUTE aether_routing_bindings_default_unique_index_stmt;
DEALLOCATE PREPARE aether_routing_bindings_default_unique_index_stmt;
@@ -4,10 +4,10 @@ use async_trait::async_trait;
use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row};
use aether_data_contracts::repository::candidate_selection::{
MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow,
StoredPoolKeyCandidateOrder, StoredPoolKeyCandidateRowsByKeyIdsQuery,
StoredPoolKeyCandidateRowsQuery, StoredProviderModelMapping,
StoredRequestedModelCandidateRowsQuery,
MinimalCandidateSelectionReadRepository, StoredApiFormatCandidateRowsQuery,
StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateOrder,
StoredPoolKeyCandidateRowsByKeyIdsQuery, StoredPoolKeyCandidateRowsQuery,
StoredProviderModelMapping, StoredRequestedModelCandidateRowsQuery,
};
use aether_data_contracts::DataLayerError;
@@ -37,7 +37,6 @@ 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,
@@ -60,6 +59,9 @@ WHERE p.is_active = 1
AND gm.is_active = 1
"#;
const REQUESTED_MODEL_RAW_PAGE_SIZE: u32 = 256;
const REQUESTED_MODEL_RAW_SCAN_LIMIT: u32 = 2048;
#[derive(Debug, Clone)]
pub struct MysqlMinimalCandidateSelectionReadRepository {
pool: MysqlPool,
@@ -70,7 +72,52 @@ struct CandidateSelectionRow {
row: StoredMinimalCandidateSelectionRow,
provider_pool_enabled: bool,
key_auth_config: Option<String>,
key_last_used_at_unix_secs: Option<u64>,
}
#[derive(Debug)]
struct ExactPageAccumulator<T> {
rows: Vec<T>,
offset: usize,
limit: usize,
target_len: usize,
}
impl<T> ExactPageAccumulator<T> {
fn new(offset: u32, limit: u32) -> Self {
let offset = usize::try_from(offset).unwrap_or(usize::MAX);
let limit = usize::try_from(limit).unwrap_or(usize::MAX);
Self {
rows: Vec::new(),
offset,
limit,
target_len: offset.saturating_add(limit),
}
}
fn is_full(&self) -> bool {
self.rows.len() >= self.target_len
}
fn push_matching<I, F>(&mut self, rows: I, mut predicate: F)
where
I: IntoIterator<Item = T>,
F: FnMut(&T) -> bool,
{
let remaining = self.target_len.saturating_sub(self.rows.len());
self.rows.extend(
rows.into_iter()
.filter(|row| predicate(row))
.take(remaining),
);
}
fn into_page(self) -> Vec<T> {
self.rows
.into_iter()
.skip(self.offset)
.take(self.limit)
.collect()
}
}
impl MysqlMinimalCandidateSelectionReadRepository {
@@ -115,6 +162,263 @@ impl MysqlMinimalCandidateSelectionReadRepository {
let rows = self.load_rows_for_api_format(api_format).await?;
Ok(sort_rows(select_pool_rows(rows), true))
}
async fn load_selected_rows_for_api_format_page(
&self,
api_format: &str,
limit: u32,
offset: u32,
) -> Result<Vec<CandidateSelectionRow>, DataLayerError> {
let mut builder = api_format_page_query(api_format, limit, offset);
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_candidate_selection_row).collect()
}
async fn selected_rows_for_api_format_page(
&self,
query: &StoredApiFormatCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
if query.limit == 0 {
return Ok(Vec::new());
}
let target_len = query.offset.saturating_add(query.limit);
let mut raw_offset = 0_u32;
let mut selected = Vec::new();
while selected.len() < target_len as usize && raw_offset < REQUESTED_MODEL_RAW_SCAN_LIMIT {
let raw_limit =
target_len.min(REQUESTED_MODEL_RAW_SCAN_LIMIT.saturating_sub(raw_offset));
let rows = self
.load_selected_rows_for_api_format_page(&query.api_format, raw_limit, raw_offset)
.await?;
let raw_len = rows.len() as u32;
selected.extend(rows.into_iter().filter(|item| {
api_format_matches(&item.row.endpoint_api_format, &query.api_format)
&& item.row.key_supports_api_format(&query.api_format)
&& key_auth_channel_matches(item, &query.api_format)
}));
if raw_len < raw_limit || raw_len == 0 {
break;
}
let next_offset = raw_offset.saturating_add(raw_len);
if next_offset == raw_offset {
break;
}
raw_offset = next_offset;
}
let rows = selected.into_iter().map(|item| item.row).collect();
Ok(sort_rows(dedupe_candidate_selection_rows(rows), true)
.into_iter()
.skip(query.offset as usize)
.take(query.limit as usize)
.collect())
}
}
fn api_format_page_query(
api_format: &str,
limit: u32,
offset: u32,
) -> QueryBuilder<'static, MySql> {
let mut builder = QueryBuilder::<MySql>::new("WITH candidate_rows AS (");
builder.push(CANDIDATE_SELECTION_COLUMNS);
push_candidate_api_format_filters(&mut builder, api_format);
push_selected_pool_rows(&mut builder);
builder.push(
r#"
ORDER BY
global_model_name ASC,
provider_priority ASC,
key_internal_priority ASC,
provider_id ASC,
endpoint_id ASC,
key_id ASC,
model_id ASC
LIMIT "#,
);
builder.push_bind(i64::from(limit));
builder.push(" OFFSET ");
builder.push_bind(i64::from(offset));
builder
}
fn requested_model_page_query(
api_format: &str,
requested_model_name: &str,
limit: u32,
offset: u32,
) -> QueryBuilder<'static, MySql> {
let mut builder = QueryBuilder::<MySql>::new("WITH candidate_rows AS (");
builder.push(CANDIDATE_SELECTION_COLUMNS);
push_candidate_api_format_filters(&mut builder, api_format);
push_requested_model_sql_filter(&mut builder, requested_model_name);
push_selected_pool_rows(&mut builder);
builder.push(
r#"
ORDER BY
global_model_name ASC,
provider_priority ASC,
key_internal_priority ASC,
provider_id ASC,
endpoint_id ASC,
key_id ASC,
model_id ASC
LIMIT "#,
);
builder.push_bind(i64::from(limit));
builder.push(" OFFSET ");
builder.push_bind(i64::from(offset));
builder
}
fn pool_key_group_query(query: &StoredPoolKeyCandidateRowsQuery) -> QueryBuilder<'static, MySql> {
let mut builder = QueryBuilder::<MySql>::new(CANDIDATE_SELECTION_COLUMNS);
push_candidate_api_format_filters(&mut builder, &query.api_format);
push_pool_key_group_filters(
&mut builder,
&query.provider_id,
&query.endpoint_id,
&query.model_id,
);
push_pool_key_order(&mut builder, &query.order);
builder.push(" LIMIT ");
builder.push_bind(i64::from(query.limit));
builder.push(" OFFSET ");
builder.push_bind(i64::from(query.offset));
builder
}
fn pool_key_group_by_key_ids_query(
query: &StoredPoolKeyCandidateRowsByKeyIdsQuery,
) -> QueryBuilder<'static, MySql> {
let mut builder = QueryBuilder::<MySql>::new(CANDIDATE_SELECTION_COLUMNS);
push_candidate_api_format_filters(&mut builder, &query.api_format);
push_pool_key_group_filters(
&mut builder,
&query.provider_id,
&query.endpoint_id,
&query.model_id,
);
builder.push(" AND pak.id IN (");
{
let mut separated = builder.separated(", ");
for key_id in &query.key_ids {
separated.push_bind(key_id.clone());
}
}
builder.push(") ORDER BY FIELD(pak.id, ");
{
let mut separated = builder.separated(", ");
for key_id in &query.key_ids {
separated.push_bind(key_id.clone());
}
}
builder.push(") ASC, pak.id ASC");
builder
}
fn push_candidate_api_format_filters(builder: &mut QueryBuilder<'_, MySql>, api_format: &str) {
let canonical_api_format = normalize_api_format(api_format);
let storage_aliases = sql_match_aliases(&api_format_aliases(&canonical_api_format));
let permission_aliases =
sql_match_aliases(&api_format_permission_aliases(&canonical_api_format));
builder.push(" AND LOWER(pe.api_format) IN (");
{
let mut separated = builder.separated(", ");
for alias in storage_aliases {
separated.push_bind(alias);
}
}
builder.push(") AND (pak.api_formats IS NULL OR TRIM(pak.api_formats) = ''");
for alias in permission_aliases {
builder.push(" OR JSON_SEARCH(LOWER(pak.api_formats), 'one', ");
builder.push_bind(alias);
builder.push(") IS NOT NULL");
}
builder.push(")");
push_key_auth_channel_filter(builder, &canonical_api_format);
}
fn push_requested_model_sql_filter(
builder: &mut QueryBuilder<'_, MySql>,
requested_model_name: &str,
) {
builder.push(" AND (gm.name = ");
builder.push_bind(requested_model_name.to_string());
builder.push(" OR m.provider_model_name = ");
builder.push_bind(requested_model_name.to_string());
builder.push(" OR (m.provider_model_mappings IS NOT NULL AND LOCATE(");
builder.push_bind(requested_model_name.to_string());
builder.push(", m.provider_model_mappings) > 0))");
}
fn push_selected_pool_rows(builder: &mut QueryBuilder<'_, MySql>) {
builder.push(
r#"
),
ranked_rows AS (
SELECT
candidate_rows.*,
CASE
WHEN JSON_EXTRACT(provider_config, '$.pool_advanced') IS NOT NULL
AND JSON_TYPE(JSON_EXTRACT(provider_config, '$.pool_advanced')) <> 'NULL'
THEN 1 ELSE 0
END AS provider_pool_enabled,
ROW_NUMBER() OVER (
PARTITION BY provider_id, endpoint_id, model_id
ORDER BY key_internal_priority ASC, key_id ASC
) AS pool_rank
FROM candidate_rows
),
selected_rows AS (
SELECT *
FROM ranked_rows
WHERE provider_pool_enabled = 0 OR pool_rank = 1
)
SELECT *
FROM selected_rows"#,
);
}
fn push_pool_key_group_filters(
builder: &mut QueryBuilder<'_, MySql>,
provider_id: &str,
endpoint_id: &str,
model_id: &str,
) {
builder.push(" AND p.id = ");
builder.push_bind(provider_id.to_string());
builder.push(" AND pe.id = ");
builder.push_bind(endpoint_id.to_string());
builder.push(" AND m.id = ");
builder.push_bind(model_id.to_string());
}
fn push_pool_key_order(builder: &mut QueryBuilder<'_, MySql>, order: &StoredPoolKeyCandidateOrder) {
match order {
StoredPoolKeyCandidateOrder::InternalPriority => {
builder.push(" ORDER BY pak.internal_priority ASC, pak.id ASC");
}
StoredPoolKeyCandidateOrder::Lru => {
builder.push(
" ORDER BY pak.last_used_at IS NOT NULL ASC, pak.last_used_at ASC, pak.internal_priority ASC, pak.id ASC",
);
}
StoredPoolKeyCandidateOrder::CacheAffinity => {
builder.push(
" ORDER BY pak.last_used_at IS NULL ASC, pak.last_used_at DESC, pak.internal_priority ASC, pak.id ASC",
);
}
StoredPoolKeyCandidateOrder::SingleAccount => {
builder.push(
" ORDER BY pak.internal_priority ASC, pak.last_used_at IS NULL ASC, pak.last_used_at DESC, pak.id ASC",
);
}
StoredPoolKeyCandidateOrder::LoadBalance { seed } => {
builder.push(" ORDER BY MD5(CONCAT(");
builder.push_bind(seed.clone());
builder.push(", ':', pak.id)) ASC, pak.id ASC");
}
}
}
#[async_trait]
@@ -126,6 +430,13 @@ impl MinimalCandidateSelectionReadRepository for MysqlMinimalCandidateSelectionR
self.selected_rows_for_api_format(api_format).await
}
async fn list_for_exact_api_format_page(
&self,
query: &StoredApiFormatCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
self.selected_rows_for_api_format_page(query).await
}
async fn list_for_exact_api_format_and_global_model(
&self,
api_format: &str,
@@ -146,57 +457,76 @@ impl MinimalCandidateSelectionReadRepository for MysqlMinimalCandidateSelectionR
api_format: &str,
requested_model_name: &str,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
self.list_for_exact_api_format_and_requested_model_page(
&StoredRequestedModelCandidateRowsQuery {
api_format: api_format.to_string(),
requested_model_name: requested_model_name.to_string(),
offset: 0,
limit: u32::MAX,
},
)
.await
let rows = self
.selected_rows_for_api_format(api_format)
.await?
.into_iter()
.filter(|row| row_matches_requested_model(row, requested_model_name, api_format))
.collect::<Vec<_>>();
Ok(sort_rows(rows, true))
}
async fn list_for_exact_api_format_and_requested_model_page(
&self,
query: &StoredRequestedModelCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
let rows = self
.selected_rows_for_api_format(&query.api_format)
.await?
if query.limit == 0 {
return Ok(Vec::new());
}
// The SQL model predicate is a coarse superset, so fill the exact page across raw windows.
let mut exact_page = ExactPageAccumulator::new(query.offset, query.limit);
let mut raw_offset = 0_u32;
while !exact_page.is_full() && raw_offset < REQUESTED_MODEL_RAW_SCAN_LIMIT {
let raw_limit = REQUESTED_MODEL_RAW_PAGE_SIZE
.min(REQUESTED_MODEL_RAW_SCAN_LIMIT.saturating_sub(raw_offset));
let mut builder = requested_model_page_query(
&query.api_format,
&query.requested_model_name,
raw_limit,
raw_offset,
);
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
let raw_len = u32::try_from(rows.len()).unwrap_or(u32::MAX);
let items = rows
.iter()
.map(map_candidate_selection_row)
.collect::<Result<Vec<_>, _>>()?;
exact_page.push_matching(items, |item| {
row_matches_requested_model(
&item.row,
&query.requested_model_name,
&query.api_format,
)
});
raw_offset = raw_offset.saturating_add(raw_len);
if raw_len < raw_limit || raw_len == 0 {
break;
}
}
let rows = exact_page
.into_page()
.into_iter()
.filter(|row| {
row_matches_requested_model(row, &query.requested_model_name, &query.api_format)
})
.collect::<Vec<_>>();
Ok(sort_rows(rows, true)
.into_iter()
.skip(query.offset as usize)
.take(query.limit as usize)
.collect())
.map(|item| item.row)
.collect();
Ok(sort_rows(dedupe_candidate_selection_rows(rows), true))
}
async fn list_pool_key_rows_for_group(
&self,
query: &StoredPoolKeyCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
let rows = self
.load_rows_for_api_format(&query.api_format)
.await?
.into_iter()
.filter(|row| {
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, &query.order);
Ok(rows
.drain(..)
.skip(query.offset as usize)
.take(query.limit as usize)
.map(|item| item.row)
.collect())
if query.limit == 0 {
return Ok(Vec::new());
}
let mut builder = pool_key_group_query(query);
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
let rows = rows
.iter()
.map(map_candidate_selection_row)
.collect::<Result<Vec<_>, _>>()?;
Ok(dedupe_candidate_selection_rows(
rows.into_iter().map(|item| item.row).collect(),
))
}
async fn list_pool_key_rows_for_group_key_ids(
@@ -212,25 +542,23 @@ impl MinimalCandidateSelectionReadRepository for MysqlMinimalCandidateSelectionR
.enumerate()
.map(|(index, key_id)| (key_id.as_str(), index))
.collect::<BTreeMap<_, _>>();
let mut rows = self
.load_rows_for_api_format(&query.api_format)
.await?
let mut builder = pool_key_group_by_key_ids_query(query);
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
let mut rows = rows
.iter()
.map(map_candidate_selection_row)
.collect::<Result<Vec<_>, _>>()?
.into_iter()
.filter(|row| {
row.row.provider_id == query.provider_id
&& row.row.endpoint_id == query.endpoint_id
&& row.row.model_id == query.model_id
&& key_order.contains_key(row.row.key_id.as_str())
})
.map(|item| item.row)
.collect::<Vec<_>>();
rows = dedupe_candidate_selection_rows(rows);
rows.sort_by(|left, right| {
key_order
.get(left.key_id.as_str())
.cmp(&key_order.get(right.key_id.as_str()))
.then(left.key_id.cmp(&right.key_id))
});
Ok(dedupe_candidate_selection_rows(rows))
Ok(rows)
}
}
@@ -283,67 +611,6 @@ fn sort_rows(
rows
}
fn sort_pool_key_rows(
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,
@@ -418,6 +685,57 @@ fn mapping_scope_matches(
})
}
fn push_key_auth_channel_filter(builder: &mut QueryBuilder<'_, MySql>, api_format: &str) {
builder.push(" AND ((LOWER(TRIM(p.provider_type)) = 'codex'");
builder.push(" AND LOWER(TRIM(pak.auth_type)) = 'oauth' AND ");
builder.push_bind(api_format.to_string());
builder.push(
" IN ('openai:responses', 'openai:responses:compact', 'openai:search', 'openai:image'))",
);
builder.push(" OR (LOWER(TRIM(p.provider_type)) = 'chatgpt_web'");
builder.push(" AND LOWER(TRIM(pak.auth_type)) IN ('oauth', 'bearer') AND ");
builder.push_bind(api_format.to_string());
builder.push(" = 'openai:image')");
builder.push(" OR (LOWER(TRIM(p.provider_type)) = 'claude_code'");
builder.push(" AND LOWER(TRIM(pak.auth_type)) = 'oauth' AND ");
builder.push_bind(api_format.to_string());
builder.push(" = 'claude:messages')");
builder.push(" OR (LOWER(TRIM(p.provider_type)) = 'kiro' AND ");
builder.push_bind(api_format.to_string());
builder.push(" = 'claude:messages' AND (LOWER(TRIM(pak.auth_type)) = 'oauth'");
builder.push(" OR (LOWER(TRIM(pak.auth_type)) = 'bearer'");
builder.push(" AND pak.auth_config IS NOT NULL AND TRIM(pak.auth_config) <> '')))");
builder.push(" OR (LOWER(TRIM(p.provider_type)) IN ('gemini_cli', 'antigravity')");
builder.push(" AND LOWER(TRIM(pak.auth_type)) = 'oauth' AND ");
builder.push_bind(api_format.to_string());
builder.push(" = 'gemini:generate_content')");
builder.push(" OR (LOWER(TRIM(p.provider_type)) = 'grok'");
builder.push(" AND LOWER(TRIM(pak.auth_type)) = 'oauth' AND ");
builder.push_bind(api_format.to_string());
builder.push(" IN ('openai:chat', 'openai:responses', 'claude:messages', 'openai:image'))");
builder.push(" OR (LOWER(TRIM(p.provider_type)) = 'windsurf'");
builder.push(" AND LOWER(TRIM(pak.auth_type)) IN ('oauth', 'api_key', 'bearer') AND ");
builder.push_bind(api_format.to_string());
builder.push(" = 'openai:chat')");
builder.push(" OR (LOWER(TRIM(p.provider_type)) = 'vertex_ai'");
builder.push(
" AND LOWER(TRIM(pak.auth_type)) IN ('api_key', 'service_account', 'vertex_ai') AND ",
);
builder.push_bind(api_format.to_string());
builder.push(" IN ('gemini:generate_content', 'gemini:embedding'))");
builder.push(
" OR (LOWER(TRIM(p.provider_type)) NOT IN ('chatgpt_web', 'claude_code', 'codex', 'gemini_cli', 'grok', 'vertex_ai', 'antigravity', 'kiro', 'windsurf') AND LOWER(TRIM(pak.auth_type)) <> 'oauth'))",
);
}
fn key_auth_channel_matches(row: &CandidateSelectionRow, api_format: &str) -> bool {
let provider_type = row.row.provider_type.trim().to_ascii_lowercase();
let auth_type = row.row.key_auth_type.trim().to_ascii_lowercase();
@@ -543,10 +861,6 @@ 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()),
})
}
@@ -770,6 +1084,10 @@ fn api_format_aliases(api_format: &str) -> Vec<String> {
aether_ai_formats::api_format_storage_aliases(api_format)
}
fn api_format_permission_aliases(api_format: &str) -> Vec<String> {
aether_ai_formats::api_format_permission_storage_aliases(api_format)
}
fn normalize_api_format(api_format: &str) -> String {
aether_ai_formats::normalize_api_format_alias(api_format)
}
@@ -791,7 +1109,134 @@ fn sql_match_aliases(api_formats: &[String]) -> Vec<String> {
#[cfg(test)]
mod tests {
use super::{vertex_key_auth_channel_matches, MysqlMinimalCandidateSelectionReadRepository};
use super::{
api_format_page_query, pool_key_group_by_key_ids_query, pool_key_group_query,
requested_model_page_query, vertex_key_auth_channel_matches, ExactPageAccumulator,
MysqlMinimalCandidateSelectionReadRepository, REQUESTED_MODEL_RAW_SCAN_LIMIT,
};
use aether_data_contracts::repository::candidate_selection::{
StoredPoolKeyCandidateOrder, StoredPoolKeyCandidateRowsByKeyIdsQuery,
StoredPoolKeyCandidateRowsQuery,
};
#[test]
fn api_format_page_query_uses_portable_sql_pagination_and_stable_order() {
let query = api_format_page_query("openai:chat", 256, 512);
let sql = query.sql();
assert!(sql.contains("ROW_NUMBER() OVER ("));
assert!(sql.contains("JSON_SEARCH(LOWER(pak.api_formats), 'one', ?"));
assert!(sql.contains("WHERE provider_pool_enabled = 0 OR pool_rank = 1"));
assert!(sql.contains(
"ORDER BY\n global_model_name ASC,\n provider_priority ASC,\n key_internal_priority ASC,"
));
assert!(sql.contains("LIMIT ? OFFSET ?"));
}
#[test]
fn requested_model_page_query_filters_and_pages_before_fetch() {
let query = requested_model_page_query("openai:chat", "gpt-5", 256, 256);
let sql = query.sql();
assert!(sql.contains("LOWER(pe.api_format) IN ("));
assert!(sql.contains("JSON_SEARCH(LOWER(pak.api_formats), 'one', ?"));
assert!(sql.contains("LOWER(TRIM(p.provider_type)) = 'codex'"));
assert!(sql.contains("AND (gm.name = ? OR m.provider_model_name = ?"));
assert!(sql.contains("LOCATE(?, m.provider_model_mappings) > 0"));
assert!(sql.contains("ROW_NUMBER() OVER ("));
assert!(sql.contains("LIMIT ? OFFSET ?"));
}
#[test]
fn exact_page_accumulator_continues_after_coarse_false_positives() {
let mut accumulator = ExactPageAccumulator::new(1, 2);
accumulator.push_matching(vec![("coarse-1", false), ("coarse-2", false)], |row| row.1);
assert!(!accumulator.is_full());
accumulator.push_matching(
vec![
("exact-1", true),
("coarse-3", false),
("exact-2", true),
("exact-3", true),
],
|row| row.1,
);
assert!(accumulator.is_full());
assert_eq!(
accumulator.into_page(),
vec![("exact-2", true), ("exact-3", true)]
);
assert_eq!(REQUESTED_MODEL_RAW_SCAN_LIMIT, 2048);
}
#[test]
fn pool_group_query_pushes_group_filters_order_and_page_into_sql() {
let query = pool_key_group_query(&StoredPoolKeyCandidateRowsQuery {
api_format: "openai:chat".to_string(),
provider_id: "provider-1".to_string(),
endpoint_id: "endpoint-1".to_string(),
model_id: "model-1".to_string(),
selected_provider_model_name: "gpt-5".to_string(),
order: StoredPoolKeyCandidateOrder::Lru,
offset: 64,
limit: 64,
});
let sql = query.sql();
assert!(sql.contains("LOWER(pe.api_format) IN ("));
assert!(sql.contains("JSON_SEARCH(LOWER(pak.api_formats), 'one', ?"));
assert!(sql.contains("LOWER(TRIM(p.provider_type)) = 'codex'"));
assert!(sql.contains("AND p.id = ? AND pe.id = ? AND m.id = ?"));
assert!(sql.contains(
"ORDER BY pak.last_used_at IS NOT NULL ASC, pak.last_used_at ASC, pak.internal_priority ASC, pak.id ASC"
));
assert!(sql.contains("LIMIT ? OFFSET ?"));
}
#[test]
fn load_balance_pool_group_query_uses_seeded_sql_order_before_paging() {
let query = pool_key_group_query(&StoredPoolKeyCandidateRowsQuery {
api_format: "openai:chat".to_string(),
provider_id: "provider-1".to_string(),
endpoint_id: "endpoint-1".to_string(),
model_id: "model-1".to_string(),
selected_provider_model_name: "gpt-5".to_string(),
order: StoredPoolKeyCandidateOrder::LoadBalance {
seed: "request-1".to_string(),
},
offset: 128,
limit: 64,
});
let sql = query.sql();
assert!(sql.contains("AND p.id = ? AND pe.id = ? AND m.id = ?"));
assert!(
sql.contains("ORDER BY MD5(CONCAT(?, ':', pak.id)) ASC, pak.id ASC LIMIT ? OFFSET ?")
);
assert!(!sql.contains("ROW_NUMBER() OVER ("));
}
#[test]
fn pool_group_by_key_ids_query_filters_ids_and_preserves_requested_order() {
let query = pool_key_group_by_key_ids_query(&StoredPoolKeyCandidateRowsByKeyIdsQuery {
api_format: "openai:chat".to_string(),
provider_id: "provider-1".to_string(),
endpoint_id: "endpoint-1".to_string(),
model_id: "model-1".to_string(),
selected_provider_model_name: "gpt-5".to_string(),
key_ids: vec!["key-2".to_string(), "key-1".to_string()],
});
let sql = query.sql();
assert!(sql.contains("LOWER(pe.api_format) IN ("));
assert!(sql.contains("JSON_SEARCH(LOWER(pak.api_formats), 'one', ?"));
assert!(sql.contains("LOWER(TRIM(p.provider_type)) = 'codex'"));
assert!(sql.contains("AND p.id = ? AND pe.id = ? AND m.id = ?"));
assert!(sql.contains("AND pak.id IN (?, ?)"));
assert!(sql.contains("ORDER BY FIELD(pak.id, ?, ?) ASC, pak.id ASC"));
}
#[tokio::test]
async fn repository_builds_from_lazy_pool() {
@@ -1,6 +1,6 @@
use async_trait::async_trait;
use serde_json::Value;
use sqlx::{mysql::MySqlRow, Row};
use sqlx::{mysql::MySqlRow, Acquire, Row};
use aether_data_contracts::repository::routing_profiles::*;
use aether_data_contracts::DataLayerError;
@@ -56,24 +56,6 @@ impl MysqlRoutingGroupRepository {
pub fn new(pool: MysqlPool) -> Self {
Self { pool }
}
async fn reload_group(&self, id: &str) -> Result<Option<StoredRoutingGroup>, DataLayerError> {
self.find_routing_group(RoutingGroupLookupKey::Id(id)).await
}
async fn find_binding_by_id(
&self,
id: &str,
) -> Result<Option<StoredRoutingGroupBinding>, DataLayerError> {
let row = sqlx::query(&format!(
"{ROUTING_GROUP_BINDING_SELECT} WHERE id = ? LIMIT 1"
))
.bind(id)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_binding_row).transpose()
}
}
#[async_trait]
@@ -170,6 +152,24 @@ impl RoutingGroupWriteRepository for MysqlRoutingGroupRepository {
record: CreateRoutingGroupRecord,
) -> Result<StoredRoutingGroup, DataLayerError> {
let group = StoredRoutingGroup::new(record)?;
let mut connection = self.pool.acquire().await.map_sql_err()?;
sqlx::query("SET TRANSACTION ISOLATION LEVEL SERIALIZABLE")
.execute(&mut *connection)
.await
.map_sql_err()?;
let mut tx = connection.begin().await.map_sql_err()?;
sqlx::query("SELECT id FROM routing_groups ORDER BY id FOR UPDATE")
.fetch_all(&mut *tx)
.await
.map_sql_err()?;
if group.is_system_default {
sqlx::query(
"UPDATE routing_groups SET is_system_default = 0 WHERE is_system_default = 1",
)
.execute(&mut *tx)
.await
.map_sql_err()?;
}
sqlx::query(
r#"
INSERT INTO routing_groups (
@@ -192,9 +192,10 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
.bind(group.created_at)
.bind(group.updated_at)
.bind(group.published_at)
.execute(&self.pool)
.execute(&mut *tx)
.await
.map_sql_err()?;
tx.commit().await.map_sql_err()?;
Ok(group)
}
@@ -203,10 +204,36 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
id: &str,
patch: UpdateRoutingGroupRecord,
) -> Result<Option<StoredRoutingGroup>, DataLayerError> {
let Some(mut group) = self.reload_group(id).await? else {
let mut connection = self.pool.acquire().await.map_sql_err()?;
sqlx::query("SET TRANSACTION ISOLATION LEVEL SERIALIZABLE")
.execute(&mut *connection)
.await
.map_sql_err()?;
let mut tx = connection.begin().await.map_sql_err()?;
sqlx::query("SELECT id FROM routing_groups ORDER BY id FOR UPDATE")
.fetch_all(&mut *tx)
.await
.map_sql_err()?;
let row = sqlx::query(&format!(
"{ROUTING_GROUP_SELECT} WHERE id = ? LIMIT 1 FOR UPDATE"
))
.bind(id)
.fetch_optional(&mut *tx)
.await
.map_sql_err()?;
let Some(mut group) = row.as_ref().map(map_group_row).transpose()? else {
return Ok(None);
};
apply_group_patch(&mut group, patch)?;
if group.is_system_default {
sqlx::query(
"UPDATE routing_groups SET is_system_default = 0 WHERE is_system_default = 1 AND id <> ?",
)
.bind(id)
.execute(&mut *tx)
.await
.map_sql_err()?;
}
sqlx::query(
r#"
UPDATE routing_groups
@@ -233,9 +260,10 @@ WHERE id = ?
.bind(group.updated_at)
.bind(group.published_at)
.bind(id)
.execute(&self.pool)
.execute(&mut *tx)
.await
.map_sql_err()?;
tx.commit().await.map_sql_err()?;
Ok(Some(group))
}
@@ -266,6 +294,30 @@ WHERE id = ?
record: CreateRoutingGroupBindingRecord,
) -> Result<StoredRoutingGroupBinding, DataLayerError> {
let binding = StoredRoutingGroupBinding::new(record)?;
let mut connection = self.pool.acquire().await.map_sql_err()?;
sqlx::query("SET TRANSACTION ISOLATION LEVEL SERIALIZABLE")
.execute(&mut *connection)
.await
.map_sql_err()?;
let mut tx = connection.begin().await.map_sql_err()?;
sqlx::query("SELECT id FROM routing_group_bindings ORDER BY id FOR UPDATE")
.fetch_all(&mut *tx)
.await
.map_sql_err()?;
if binding.is_default {
sqlx::query(
r#"
UPDATE routing_group_bindings
SET is_default = 0
WHERE is_default = 1 AND subject_type = ? AND subject_id = ?
"#,
)
.bind(binding_subject_to_database(binding.subject_type))
.bind(&binding.subject_id)
.execute(&mut *tx)
.await
.map_sql_err()?;
}
sqlx::query(
r#"
INSERT INTO routing_group_bindings (
@@ -283,9 +335,10 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?)
.bind(binding.allow_explicit_select)
.bind(binding.created_at)
.bind(binding.updated_at)
.execute(&self.pool)
.execute(&mut *tx)
.await
.map_sql_err()?;
tx.commit().await.map_sql_err()?;
Ok(binding)
}
@@ -306,10 +359,45 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?)
id: &str,
patch: UpdateRoutingGroupBindingRecord,
) -> Result<Option<StoredRoutingGroupBinding>, DataLayerError> {
let Some(mut binding) = self.find_binding_by_id(id).await? else {
let mut connection = self.pool.acquire().await.map_sql_err()?;
sqlx::query("SET TRANSACTION ISOLATION LEVEL SERIALIZABLE")
.execute(&mut *connection)
.await
.map_sql_err()?;
let mut tx = connection.begin().await.map_sql_err()?;
sqlx::query("SELECT id FROM routing_group_bindings ORDER BY id FOR UPDATE")
.fetch_all(&mut *tx)
.await
.map_sql_err()?;
let row = sqlx::query(&format!(
"{ROUTING_GROUP_BINDING_SELECT} WHERE id = ? LIMIT 1 FOR UPDATE"
))
.bind(id)
.fetch_optional(&mut *tx)
.await
.map_sql_err()?;
let Some(mut binding) = row.as_ref().map(map_binding_row).transpose()? else {
return Ok(None);
};
apply_binding_patch(&mut binding, patch)?;
if binding.is_default {
sqlx::query(
r#"
UPDATE routing_group_bindings
SET is_default = 0
WHERE is_default = 1
AND subject_type = ?
AND subject_id = ?
AND id <> ?
"#,
)
.bind(binding_subject_to_database(binding.subject_type))
.bind(&binding.subject_id)
.bind(id)
.execute(&mut *tx)
.await
.map_sql_err()?;
}
sqlx::query(
r#"
UPDATE routing_group_bindings
@@ -329,9 +417,10 @@ WHERE id = ?
.bind(binding.allow_explicit_select)
.bind(binding.updated_at)
.bind(id)
.execute(&self.pool)
.execute(&mut *tx)
.await
.map_sql_err()?;
tx.commit().await.map_sql_err()?;
Ok(Some(binding))
}