fix(routing): harden routed pool scheduling

This commit is contained in:
elky
2026-07-27 22:06:28 +08:00
parent 550cc36760
commit 4148ab1931
62 changed files with 5510 additions and 850 deletions
@@ -4,10 +4,10 @@ use async_trait::async_trait;
use sqlx::{sqlite::SqliteRow, QueryBuilder, Row, Sqlite};
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;
@@ -68,6 +68,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 SqliteMinimalCandidateSelectionReadRepository {
pool: SqlitePool,
@@ -77,7 +80,6 @@ pub struct SqliteMinimalCandidateSelectionReadRepository {
struct CandidateSelectionRow {
row: StoredMinimalCandidateSelectionRow,
key_auth_config: Option<String>,
key_last_used_at_unix_secs: Option<u64>,
}
#[derive(Debug, Clone, Copy)]
@@ -99,6 +101,58 @@ struct SqlPage {
offset: i64,
}
#[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()
}
}
#[derive(Debug)]
struct RequestedModelRawPage {
rows: Vec<CandidateSelectionRow>,
raw_len: u32,
}
impl SqliteMinimalCandidateSelectionReadRepository {
pub fn new(pool: SqlitePool) -> Self {
Self { pool }
@@ -148,44 +202,7 @@ impl SqliteMinimalCandidateSelectionReadRepository {
);
}
}
builder.push(
r#"
),
pool_rows AS (
SELECT candidate.*
FROM candidate_rows candidate
WHERE candidate.provider_pool_enabled = 1
AND NOT EXISTS (
SELECT 1
FROM candidate_rows other
WHERE other.provider_pool_enabled = 1
AND other.provider_id = candidate.provider_id
AND other.endpoint_id = candidate.endpoint_id
AND other.model_id = candidate.model_id
AND (
other.key_internal_priority < candidate.key_internal_priority
OR (
other.key_internal_priority = candidate.key_internal_priority
AND other.key_id < candidate.key_id
)
)
)
),
selected_rows AS (
SELECT * FROM candidate_rows WHERE provider_pool_enabled = 0
UNION ALL
SELECT * FROM pool_rows
)
SELECT * FROM selected_rows
"#,
);
push_selected_rows_order(&mut builder, order);
if let Some(page) = page {
builder.push(" LIMIT ");
builder.push_bind(page.limit);
builder.push(" OFFSET ");
builder.push_bind(page.offset);
}
push_selected_rows_query_tail(&mut builder, order, page);
let query_rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
let mut items = query_rows
@@ -211,6 +228,42 @@ SELECT * FROM selected_rows
};
Ok(dedupe_candidate_selection_rows(rows))
}
async fn load_requested_model_raw_page(
&self,
api_format: &str,
requested_model_name: &str,
page: SqlPage,
) -> Result<RequestedModelRawPage, DataLayerError> {
let canonical_api_format = normalize_api_format(api_format);
let storage_aliases = api_format_aliases(&canonical_api_format);
let match_aliases =
sql_match_aliases(&api_format_permission_aliases(&canonical_api_format));
let mut builder = QueryBuilder::<Sqlite>::new("WITH candidate_rows AS (");
builder.push(CANDIDATE_SELECTION_COLUMNS);
push_candidate_sql_filters_for_aliases(
&mut builder,
&storage_aliases,
&match_aliases,
&canonical_api_format,
);
push_requested_model_sql_filter(&mut builder, requested_model_name, &match_aliases);
push_selected_rows_query_tail(&mut builder, SelectedRowsOrder::WithGlobalModel, Some(page));
let query_rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
let raw_len = u32::try_from(query_rows.len()).unwrap_or(u32::MAX);
let mut rows = query_rows
.iter()
.map(map_candidate_selection_row)
.collect::<Result<Vec<_>, _>>()?;
rows.retain(|item| {
api_format_matches(&item.row.endpoint_api_format, &canonical_api_format)
&& item.row.key_supports_api_format(&canonical_api_format)
&& key_auth_channel_matches(item, &canonical_api_format)
});
Ok(RequestedModelRawPage { rows, raw_len })
}
}
#[async_trait]
@@ -222,6 +275,33 @@ impl MinimalCandidateSelectionReadRepository for SqliteMinimalCandidateSelection
self.selected_rows_for_api_format(api_format).await
}
async fn list_for_exact_api_format_page(
&self,
query: &StoredApiFormatCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
if query.limit == 0 {
return Ok(Vec::new());
}
let fetch_limit = query.offset.saturating_add(query.limit);
let mut rows = self
.load_selected_rows_for_api_format(
&query.api_format,
SelectedRowsFilter::None,
SelectedRowsOrder::WithGlobalModel,
Some(SqlPage {
limit: i64::from(fetch_limit),
offset: 0,
}),
)
.await?;
sort_candidate_selection_rows(&mut rows, true);
Ok(rows
.into_iter()
.skip(query.offset as usize)
.take(query.limit as usize)
.collect())
}
async fn list_for_exact_api_format_and_global_model(
&self,
api_format: &str,
@@ -254,28 +334,58 @@ impl MinimalCandidateSelectionReadRepository for SqliteMinimalCandidateSelection
&self,
query: &StoredRequestedModelCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
self.load_selected_rows_for_api_format(
&query.api_format,
SelectedRowsFilter::RequestedModel(&query.requested_model_name),
SelectedRowsOrder::WithGlobalModel,
Some(SqlPage {
limit: i64::from(query.limit.max(1)),
offset: i64::from(query.offset),
}),
)
.await
if query.limit == 0 {
return Ok(Vec::new());
}
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 raw_page = self
.load_requested_model_raw_page(
&query.api_format,
&query.requested_model_name,
SqlPage {
limit: i64::from(raw_limit),
offset: i64::from(raw_offset),
},
)
.await?;
exact_page.push_matching(raw_page.rows, |item| {
row_matches_requested_model(
&item.row,
&query.requested_model_name,
&query.api_format,
)
});
raw_offset = raw_offset.saturating_add(raw_page.raw_len);
if raw_page.raw_len < raw_limit || raw_page.raw_len == 0 {
break;
}
}
let rows = exact_page
.into_page()
.into_iter()
.map(|item| item.row)
.collect();
Ok(dedupe_candidate_selection_rows(rows))
}
async fn list_pool_key_rows_for_group(
&self,
query: &StoredPoolKeyCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
if query.limit == 0 {
return Ok(Vec::new());
}
let canonical_api_format = normalize_api_format(&query.api_format);
let storage_aliases = api_format_aliases(&canonical_api_format);
let match_aliases =
sql_match_aliases(&api_format_permission_aliases(&canonical_api_format));
let mut rows = Vec::<CandidateSelectionRow>::new();
let page_in_sql = !matches!(query.order, StoredPoolKeyCandidateOrder::LoadBalance { .. });
for storage_api_format in storage_aliases {
let mut builder = QueryBuilder::<Sqlite>::new(CANDIDATE_SELECTION_COLUMNS);
@@ -286,15 +396,11 @@ impl MinimalCandidateSelectionReadRepository for SqliteMinimalCandidateSelection
builder.push_bind(&query.endpoint_id);
builder.push(" AND m.id = ");
builder.push_bind(&query.model_id);
if page_in_sql {
push_pool_key_order(&mut builder, &query.order);
builder.push(" LIMIT ");
builder.push_bind(i64::from(query.limit.max(1)));
builder.push(" OFFSET ");
builder.push_bind(i64::from(query.offset));
} else {
builder.push(" ORDER BY pak.id ASC");
}
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));
let query_rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
let mut items = query_rows
@@ -309,20 +415,9 @@ impl MinimalCandidateSelectionReadRepository for SqliteMinimalCandidateSelection
rows.extend(items);
}
if page_in_sql {
Ok(dedupe_candidate_selection_rows(
rows.into_iter().map(|item| item.row).collect(),
))
} else {
Ok(dedupe_candidate_selection_rows(
sort_pool_key_rows(rows, &query.order)
.into_iter()
.skip(query.offset as usize)
.take(query.limit as usize)
.map(|item| item.row)
.collect(),
))
}
Ok(dedupe_candidate_selection_rows(
rows.into_iter().map(|item| item.row).collect(),
))
}
async fn list_pool_key_rows_for_group_key_ids(
@@ -411,6 +506,19 @@ fn push_candidate_sql_filters(
push_key_auth_channel_sql_filter(builder, storage_api_format);
}
fn push_candidate_sql_filters_for_aliases(
builder: &mut QueryBuilder<'_, Sqlite>,
storage_api_formats: &[String],
match_aliases: &[String],
requested_api_format: &str,
) {
builder.push(" AND LOWER(COALESCE(pe.api_format, '')) IN (");
push_bind_list(builder, storage_api_formats);
builder.push(")");
push_key_api_format_sql_filter(builder, match_aliases);
push_key_auth_channel_sql_filter(builder, requested_api_format);
}
fn push_key_api_format_sql_filter(
builder: &mut QueryBuilder<'_, Sqlite>,
match_aliases: &[String],
@@ -626,6 +734,51 @@ fn push_requested_model_sql_filter(
);
}
fn push_selected_rows_query_tail(
builder: &mut QueryBuilder<'_, Sqlite>,
order: SelectedRowsOrder,
page: Option<SqlPage>,
) {
builder.push(
r#"
),
pool_rows AS (
SELECT candidate.*
FROM candidate_rows candidate
WHERE candidate.provider_pool_enabled = 1
AND NOT EXISTS (
SELECT 1
FROM candidate_rows other
WHERE other.provider_pool_enabled = 1
AND other.provider_id = candidate.provider_id
AND other.endpoint_id = candidate.endpoint_id
AND other.model_id = candidate.model_id
AND (
other.key_internal_priority < candidate.key_internal_priority
OR (
other.key_internal_priority = candidate.key_internal_priority
AND other.key_id < candidate.key_id
)
)
)
),
selected_rows AS (
SELECT * FROM candidate_rows WHERE provider_pool_enabled = 0
UNION ALL
SELECT * FROM pool_rows
)
SELECT * FROM selected_rows
"#,
);
push_selected_rows_order(builder, order);
if let Some(page) = page {
builder.push(" LIMIT ");
builder.push_bind(page.limit);
builder.push(" OFFSET ");
builder.push_bind(page.offset);
}
}
fn push_selected_rows_order(builder: &mut QueryBuilder<'_, Sqlite>, order: SelectedRowsOrder) {
builder.push(" ORDER BY ");
if matches!(order, SelectedRowsOrder::WithGlobalModel) {
@@ -660,12 +813,52 @@ fn push_pool_key_order(
);
}
StoredPoolKeyCandidateOrder::LoadBalance { seed } => {
let _ = seed;
builder.push(" ORDER BY pak.id ASC");
push_seeded_pool_key_order(builder, seed);
}
}
}
fn push_seeded_pool_key_order(builder: &mut QueryBuilder<'_, Sqlite>, seed: &str) {
const HEX_DIGITS: &[u8; 16] = b"0123456789abcdef";
const PLACEHOLDERS: &[u8; 16] = b"ghijklmnopqrstuv";
let mut digits_by_rank = HEX_DIGITS.map(char::from);
digits_by_rank.sort_by(|left, right| {
stable_pool_key_hash(seed, &left.to_string())
.cmp(&stable_pool_key_hash(seed, &right.to_string()))
.then(left.cmp(right))
});
let mut rank_by_digit = ['0'; 16];
for (rank, digit) in digits_by_rank.into_iter().enumerate() {
let digit_index = digit
.to_digit(16)
.expect("seeded pool-key rank input must be hexadecimal")
as usize;
rank_by_digit[digit_index] = char::from(HEX_DIGITS[rank]);
}
// SQLite has no built-in hash; placeholders avoid cascading replacements
// while remapping every key-id nibble to a seed-derived rank.
builder.push(" ORDER BY ");
for _ in 0..(HEX_DIGITS.len() + PLACEHOLDERS.len()) {
builder.push("replace(");
}
builder.push("lower(hex(pak.id))");
for (digit, placeholder) in HEX_DIGITS.iter().zip(PLACEHOLDERS) {
builder.push(format!(
", '{}', '{}')",
char::from(*digit),
char::from(*placeholder)
));
}
for (placeholder, rank) in PLACEHOLDERS.iter().zip(rank_by_digit) {
builder.push(format!(", '{}', ", char::from(*placeholder)));
builder.push_bind(rank.to_string());
builder.push(")");
}
builder.push(" ASC, pak.id ASC");
}
fn push_bind_list(builder: &mut QueryBuilder<'_, Sqlite>, values: &[String]) {
let mut separated = builder.separated(", ");
for value in values {
@@ -673,52 +866,6 @@ fn push_bind_list(builder: &mut QueryBuilder<'_, Sqlite>, values: &[String]) {
}
}
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
@@ -875,6 +1022,24 @@ fn dedupe_candidate_selection_rows(
.collect()
}
fn sort_candidate_selection_rows(
rows: &mut [StoredMinimalCandidateSelectionRow],
include_global_model: bool,
) {
rows.sort_by(|left, right| {
let global_model_order = include_global_model
.then(|| left.global_model_name.cmp(&right.global_model_name))
.unwrap_or(std::cmp::Ordering::Equal);
global_model_order
.then(left.provider_priority.cmp(&right.provider_priority))
.then(left.key_internal_priority.cmp(&right.key_internal_priority))
.then(left.provider_id.cmp(&right.provider_id))
.then(left.endpoint_id.cmp(&right.endpoint_id))
.then(left.key_id.cmp(&right.key_id))
.then(left.model_id.cmp(&right.model_id))
});
}
fn map_candidate_selection_row(row: &SqliteRow) -> Result<CandidateSelectionRow, DataLayerError> {
let _provider_config = parse_json(row.try_get("provider_config").ok().flatten())?;
let global_model_config = parse_json(row.try_get("global_model_config").ok().flatten())?;
@@ -931,10 +1096,6 @@ fn map_candidate_selection_row(row: &SqliteRow) -> Result<CandidateSelectionRow,
model_is_available: row.try_get("model_is_available").map_sql_err()?,
},
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()),
})
}
@@ -1177,8 +1338,9 @@ fn sql_match_aliases(api_formats: &[String]) -> Vec<String> {
#[cfg(test)]
mod tests {
use super::{
push_key_auth_channel_sql_filter, vertex_key_auth_channel_matches,
SqliteMinimalCandidateSelectionReadRepository,
push_key_auth_channel_sql_filter, push_pool_key_order, vertex_key_auth_channel_matches,
ExactPageAccumulator, SqliteMinimalCandidateSelectionReadRepository,
REQUESTED_MODEL_RAW_SCAN_LIMIT,
};
use crate::run_migrations;
use aether_data_contracts::repository::candidate_selection::{
@@ -1216,6 +1378,57 @@ mod tests {
assert!(vertex_clause.contains("gemini:embedding"));
}
#[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 load_balance_pool_key_order_is_seeded_and_pageable_in_sql() {
let sql_for_seed = |seed: &str| {
let mut builder =
sqlx::QueryBuilder::<sqlx::Sqlite>::new("SELECT pak.id FROM provider_api_keys pak");
push_pool_key_order(
&mut builder,
&StoredPoolKeyCandidateOrder::LoadBalance {
seed: seed.to_string(),
},
);
builder.push(" LIMIT ");
builder.push_bind(64_i64);
builder.push(" OFFSET ");
builder.push_bind(128_i64);
builder.sql().to_string()
};
let first_seed_sql = sql_for_seed("seed-a");
let second_seed_sql = sql_for_seed("seed-b");
assert!(first_seed_sql.contains("lower(hex(pak.id))"));
assert!(first_seed_sql.contains("ASC, pak.id ASC LIMIT ? OFFSET ?"));
assert_eq!(first_seed_sql, second_seed_sql);
assert_eq!(first_seed_sql.matches('?').count(), 18);
}
#[tokio::test]
async fn sqlite_repository_reads_candidate_selection_rows() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
@@ -1336,6 +1549,97 @@ mod tests {
assert_eq!(search_rows[0].endpoint_api_format, "openai:search");
}
#[tokio::test]
async fn sqlite_requested_model_page_crosses_coarse_false_positive_windows() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_migrations(&pool)
.await
.expect("sqlite migrations should run");
seed_requested_model_pagination(&pool).await;
let repository = SqliteMinimalCandidateSelectionReadRepository::new(pool);
let rows = repository
.list_for_exact_api_format_and_requested_model_page(
&StoredRequestedModelCandidateRowsQuery {
api_format: "openai:chat".to_string(),
requested_model_name: "sqlite-page-target".to_string(),
offset: 1,
limit: 1,
},
)
.await
.expect("requested model page should cross the coarse-only window");
assert_eq!(rows.len(), 1);
assert_eq!(rows[0].model_id, "model-pagination-exact-1");
}
#[tokio::test]
async fn sqlite_load_balance_pool_key_pages_use_stable_seeded_order() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_migrations(&pool)
.await
.expect("sqlite migrations should run");
seed_candidate_selection(&pool).await;
let repository = SqliteMinimalCandidateSelectionReadRepository::new(pool);
let load_page = |seed: &str, offset, limit| 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: "provider-model".to_string(),
order: StoredPoolKeyCandidateOrder::LoadBalance {
seed: seed.to_string(),
},
offset,
limit,
};
let seed_a_first = repository
.list_pool_key_rows_for_group(&load_page("seed-a", 0, 1))
.await
.expect("first load-balance page should load");
let seed_a_second = repository
.list_pool_key_rows_for_group(&load_page("seed-a", 1, 1))
.await
.expect("second load-balance page should load");
let seed_a_replay = repository
.list_pool_key_rows_for_group(&load_page("seed-a", 0, 2))
.await
.expect("replayed load-balance window should load");
let seed_b = repository
.list_pool_key_rows_for_group(&load_page("seed-b", 0, 2))
.await
.expect("alternate load-balance seed should load");
let seed_a_pages = seed_a_first
.iter()
.chain(&seed_a_second)
.map(|row| row.key_id.as_str())
.collect::<Vec<_>>();
let seed_a_replay = seed_a_replay
.iter()
.map(|row| row.key_id.as_str())
.collect::<Vec<_>>();
let seed_b = seed_b
.iter()
.map(|row| row.key_id.as_str())
.collect::<Vec<_>>();
assert_eq!(seed_a_pages, vec!["key-1", "key-2"]);
assert_eq!(seed_a_pages, seed_a_replay);
assert_eq!(seed_b, vec!["key-2", "key-1"]);
}
async fn seed_candidate_selection(pool: &sqlx::SqlitePool) {
sqlx::query(
r#"
@@ -1455,4 +1759,94 @@ VALUES (
.await
.expect("candidate selection rows should seed");
}
async fn seed_requested_model_pagination(pool: &sqlx::SqlitePool) {
sqlx::query(
r#"
INSERT INTO providers (
id, name, provider_type, provider_priority, is_active, created_at, updated_at
)
VALUES (
'provider-pagination', 'Pagination Provider', 'custom', 10, 1, 1, 1
);
INSERT INTO provider_endpoints (
id, provider_id, name, base_url, api_format, is_active, created_at, updated_at
)
VALUES (
'endpoint-pagination', 'provider-pagination', 'Pagination Endpoint',
'https://example.test', 'openai:chat', 1, 1, 1
);
INSERT INTO provider_api_keys (
id, provider_id, name, auth_type, api_formats, internal_priority,
is_active, created_at, updated_at
)
VALUES (
'key-pagination', 'provider-pagination', 'Pagination Key', 'api_key',
'["openai:chat"]', 10, 1, 1, 1
);
WITH RECURSIVE sequence(value) AS (
SELECT 0
UNION ALL
SELECT value + 1 FROM sequence WHERE value < 255
)
INSERT INTO global_models (
id, name, is_active, created_at, updated_at
)
SELECT
printf('global-pagination-false-%03d', value),
printf('a-pagination-false-%03d', value),
1, 1, 1
FROM sequence;
INSERT INTO global_models (
id, name, is_active, created_at, updated_at
)
VALUES
('global-pagination-exact-0', 'z-pagination-exact-0', 1, 1, 1),
('global-pagination-exact-1', 'z-pagination-exact-1', 1, 1, 1);
WITH RECURSIVE sequence(value) AS (
SELECT 0
UNION ALL
SELECT value + 1 FROM sequence WHERE value < 255
)
INSERT INTO models (
id, provider_id, global_model_id, provider_model_name, provider_model_mappings,
is_active, is_available, created_at, updated_at
)
SELECT
printf('model-pagination-false-%03d', value),
'provider-pagination',
printf('global-pagination-false-%03d', value),
'upstream-false',
'[{"name":"sqlite-page-target-noise","api_formats":["openai:chat"],"priority":1}]',
1, 1, 1, 1
FROM sequence;
INSERT INTO models (
id, provider_id, global_model_id, provider_model_name, provider_model_mappings,
is_active, is_available, created_at, updated_at
)
VALUES
(
'model-pagination-exact-0', 'provider-pagination', 'global-pagination-exact-0',
'upstream-exact-0',
'[{"name":"sqlite-page-target","api_formats":["openai:chat"],"priority":1}]',
1, 1, 1, 1
),
(
'model-pagination-exact-1', 'provider-pagination', 'global-pagination-exact-1',
'upstream-exact-1',
'[{"name":"sqlite-page-target","api_formats":["openai:chat"],"priority":1}]',
1, 1, 1, 1
);
"#,
)
.execute(pool)
.await
.expect("requested model pagination rows should seed");
}
}
@@ -56,24 +56,6 @@ impl SqliteRoutingGroupRepository {
pub fn new(pool: SqlitePool) -> 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,19 @@ impl RoutingGroupWriteRepository for SqliteRoutingGroupRepository {
record: CreateRoutingGroupRecord,
) -> Result<StoredRoutingGroup, DataLayerError> {
let group = StoredRoutingGroup::new(record)?;
let mut tx = self.pool.begin().await.map_sql_err()?;
sqlx::query("UPDATE routing_groups SET is_system_default = is_system_default WHERE 0")
.execute(&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 +187,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 +199,29 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
id: &str,
patch: UpdateRoutingGroupRecord,
) -> Result<Option<StoredRoutingGroup>, DataLayerError> {
let Some(mut group) = self.reload_group(id).await? else {
let mut tx = self.pool.begin().await.map_sql_err()?;
sqlx::query("UPDATE routing_groups SET is_system_default = is_system_default WHERE 0")
.execute(&mut *tx)
.await
.map_sql_err()?;
let row = sqlx::query(&format!("{ROUTING_GROUP_SELECT} WHERE id = ? LIMIT 1"))
.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 +248,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 +282,25 @@ WHERE id = ?
record: CreateRoutingGroupBindingRecord,
) -> Result<StoredRoutingGroupBinding, DataLayerError> {
let binding = StoredRoutingGroupBinding::new(record)?;
let mut tx = self.pool.begin().await.map_sql_err()?;
sqlx::query("UPDATE routing_group_bindings SET is_default = is_default WHERE 0")
.execute(&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 +318,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 +342,40 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?)
id: &str,
patch: UpdateRoutingGroupBindingRecord,
) -> Result<Option<StoredRoutingGroupBinding>, DataLayerError> {
let Some(mut binding) = self.find_binding_by_id(id).await? else {
let mut tx = self.pool.begin().await.map_sql_err()?;
sqlx::query("UPDATE routing_group_bindings SET is_default = is_default WHERE 0")
.execute(&mut *tx)
.await
.map_sql_err()?;
let row = sqlx::query(&format!(
"{ROUTING_GROUP_BINDING_SELECT} WHERE id = ? LIMIT 1"
))
.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 +395,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))
}
@@ -526,4 +593,255 @@ mod tests {
1
);
}
#[tokio::test]
async fn sqlite_keeps_system_and_subject_defaults_unique() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_sqlite_migrations(&pool)
.await
.expect("sqlite migrations should run");
let repository = SqliteRoutingGroupRepository::new(pool);
for (id, is_system_default) in [("group-1", true), ("group-2", true), ("group-3", false)] {
repository
.create_routing_group(group_record(id, is_system_default))
.await
.expect("group should create");
}
assert_eq!(system_default_ids(&repository).await, vec!["group-2"]);
repository
.update_routing_group(
"group-1",
UpdateRoutingGroupRecord {
is_system_default: Some(true),
updated_at: 2,
..UpdateRoutingGroupRecord::default()
},
)
.await
.expect("group should update");
assert_eq!(system_default_ids(&repository).await, vec!["group-1"]);
repository
.create_routing_group_binding(binding_record("binding-1", "group-1", "subject-1", true))
.await
.expect("binding should create");
repository
.create_routing_group_binding(binding_record("binding-2", "group-2", "subject-1", true))
.await
.expect("binding should create");
repository
.create_routing_group_binding(binding_record("binding-3", "group-3", "subject-2", true))
.await
.expect("binding should create");
assert_eq!(
default_binding_ids(&repository, "subject-1").await,
vec!["binding-2"]
);
assert_eq!(
default_binding_ids(&repository, "subject-2").await,
vec!["binding-3"]
);
repository
.update_routing_group_binding(
"binding-1",
UpdateRoutingGroupBindingRecord {
is_default: Some(true),
updated_at: 2,
..UpdateRoutingGroupBindingRecord::default()
},
)
.await
.expect("binding should update");
assert_eq!(
default_binding_ids(&repository, "subject-1").await,
vec!["binding-1"]
);
assert_eq!(
default_binding_ids(&repository, "subject-2").await,
vec!["binding-3"]
);
repository
.update_routing_group_binding(
"binding-3",
UpdateRoutingGroupBindingRecord {
subject_id: Some("subject-1".to_string()),
updated_at: 3,
..UpdateRoutingGroupBindingRecord::default()
},
)
.await
.expect("binding should move");
assert_eq!(
default_binding_ids(&repository, "subject-1").await,
vec!["binding-3"]
);
assert!(default_binding_ids(&repository, "subject-2")
.await
.is_empty());
}
#[tokio::test]
async fn sqlite_repair_migration_resolves_existing_duplicate_defaults() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_sqlite_migrations(&pool)
.await
.expect("sqlite migrations should run");
let repository = SqliteRoutingGroupRepository::new(pool.clone());
sqlx::raw_sql(
r#"
DROP INDEX routing_groups_one_system_default_key;
DROP INDEX routing_group_bindings_subject_default_key;
"#,
)
.execute(&pool)
.await
.expect("unique indexes should be removable to simulate a pre-repair database");
for id in ["group-1", "group-2", "group-3"] {
repository
.create_routing_group(group_record(id, false))
.await
.expect("group should create");
}
sqlx::query(
r#"
UPDATE routing_groups
SET is_system_default = 1,
enabled = CASE id WHEN 'group-3' THEN 0 ELSE 1 END,
updated_at = CASE id
WHEN 'group-1' THEN 1
WHEN 'group-2' THEN 2
ELSE 3
END
"#,
)
.execute(&pool)
.await
.expect("duplicate system defaults should seed");
for (id, subject_id) in [
("binding-3", "subject-1"),
("binding-2", "subject-1"),
("binding-1", "subject-1"),
("binding-4", "subject-2"),
] {
repository
.create_routing_group_binding(binding_record(id, "group-1", subject_id, false))
.await
.expect("binding should create");
}
sqlx::query(
r#"
UPDATE routing_group_bindings
SET is_default = 1,
created_at = CASE id WHEN 'binding-3' THEN 2 ELSE 1 END
"#,
)
.execute(&pool)
.await
.expect("duplicate binding defaults should seed");
let repair_migration =
include_str!("../migrations/20260727000000_repair_routing_default_uniqueness.sql");
for _ in 0..2 {
sqlx::raw_sql(repair_migration)
.execute(&pool)
.await
.expect("repair migration should be idempotent");
}
assert_eq!(system_default_ids(&repository).await, vec!["group-2"]);
assert_eq!(
default_binding_ids(&repository, "subject-1").await,
vec!["binding-1"]
);
assert_eq!(
default_binding_ids(&repository, "subject-2").await,
vec!["binding-4"]
);
sqlx::query("UPDATE routing_groups SET is_system_default = 1 WHERE id = 'group-3'")
.execute(&pool)
.await
.expect_err("database should reject a second system default");
sqlx::query("UPDATE routing_group_bindings SET is_default = 1 WHERE id = 'binding-2'")
.execute(&pool)
.await
.expect_err("database should reject a second default for the same subject");
}
fn group_record(id: &str, is_system_default: bool) -> CreateRoutingGroupRecord {
CreateRoutingGroupRecord {
id: id.to_string(),
name: id.to_string(),
description: None,
enabled: true,
is_system_default,
config_json: json!({}),
version: 1,
created_at: 1,
updated_at: 1,
published_at: None,
}
}
fn binding_record(
id: &str,
group_id: &str,
subject_id: &str,
is_default: bool,
) -> CreateRoutingGroupBindingRecord {
CreateRoutingGroupBindingRecord {
id: id.to_string(),
group_id: group_id.to_string(),
subject_type: RoutingGroupBindingSubject::ApiKey,
subject_id: subject_id.to_string(),
is_default,
allow_explicit_select: true,
created_at: 1,
updated_at: 1,
}
}
async fn system_default_ids(repository: &SqliteRoutingGroupRepository) -> Vec<String> {
repository
.list_routing_groups()
.await
.expect("groups should list")
.into_iter()
.filter(|group| group.is_system_default)
.map(|group| group.id)
.collect()
}
async fn default_binding_ids(
repository: &SqliteRoutingGroupRepository,
subject_id: &str,
) -> Vec<String> {
repository
.list_routing_group_bindings(&RoutingGroupBindingQuery {
group_id: None,
subject_type: Some(RoutingGroupBindingSubject::ApiKey),
subject_id: Some(subject_id.to_string()),
})
.await
.expect("bindings should list")
.into_iter()
.filter(|binding| binding.is_default)
.map(|binding| binding.id)
.collect()
}
}