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

@@ -141,7 +141,7 @@ where
}
pub fn ai_should_persist_available_candidate_for_pool_key(pool_key_index: Option<u32>) -> bool {
pool_key_index.is_none_or(|index| index == 0)
pool_key_index.is_none()
}
pub fn ai_should_persist_skipped_candidate_for_pool_membership(is_pool_candidate: bool) -> bool {
@@ -387,9 +387,9 @@ mod tests {
}
#[test]
fn pool_candidate_persistence_policy_persists_representatives_only() {
fn pool_candidate_persistence_policy_skips_pool_keys_until_execution() {
assert!(ai_should_persist_available_candidate_for_pool_key(None));
assert!(ai_should_persist_available_candidate_for_pool_key(Some(0)));
assert!(!ai_should_persist_available_candidate_for_pool_key(Some(0)));
assert!(!ai_should_persist_available_candidate_for_pool_key(Some(1)));
assert!(ai_should_persist_skipped_candidate_for_pool_membership(

View File

@@ -11,6 +11,7 @@ pub struct AiCandidateResolutionRequest<'a> {
pub client_api_format: &'a str,
pub requested_model: Option<&'a str>,
pub mode: AiCandidateResolutionMode,
pub expand_pool_groups: bool,
}
impl<'a> AiCandidateResolutionRequest<'a> {
@@ -19,6 +20,7 @@ impl<'a> AiCandidateResolutionRequest<'a> {
client_api_format,
requested_model,
mode: AiCandidateResolutionMode::Standard,
expand_pool_groups: true,
}
}
@@ -30,8 +32,14 @@ impl<'a> AiCandidateResolutionRequest<'a> {
client_api_format,
requested_model,
mode: AiCandidateResolutionMode::WithoutTransportPairGate,
expand_pool_groups: true,
}
}
pub fn logical_pool_groups(mut self) -> Self {
self.expand_pool_groups = false;
self
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
@@ -140,8 +148,13 @@ where
let ranked = port
.rank_eligible_candidates(eligible, normalized_client_api_format.as_str())
.await?;
let (ranked, pool_skipped) = port.apply_pool_scheduler(ranked).await?;
skipped.extend(pool_skipped);
let ranked = if request.expand_pool_groups {
let (ranked, pool_skipped) = port.apply_pool_scheduler(ranked).await?;
skipped.extend(pool_skipped);
ranked
} else {
ranked
};
Ok(AiCandidateResolutionOutcome {
eligible_candidates: ranked,
@@ -385,6 +398,38 @@ mod tests {
);
}
#[tokio::test]
async fn resolution_can_keep_pool_groups_logical() {
let port = TestPort::default();
let outcome = run_ai_candidate_resolution(
&port,
vec!["first", "second"],
AiCandidateResolutionRequest::standard("openai:chat", Some("gpt-4.1"))
.logical_pool_groups(),
)
.await
.unwrap();
assert_eq!(
outcome.eligible_candidates,
["eligible:second", "eligible:first"]
);
assert!(outcome.skipped_candidates.is_empty());
assert_eq!(
port.calls.lock().unwrap().as_slice(),
[
"transport:first",
"common:first:gpt-4.1",
"pair:first:openai:chat:gpt-4.1",
"transport:second",
"common:second:gpt-4.1",
"pair:second:openai:chat:gpt-4.1",
"rank:openai:chat",
]
);
}
#[test]
fn sticky_session_token_is_extracted_from_known_request_fields() {
assert_eq!(

View File

@@ -2,5 +2,6 @@ mod types;
pub use types::{
MinimalCandidateSelectionReadRepository, MinimalCandidateSelectionRepository,
StoredMinimalCandidateSelectionRow, StoredProviderModelMapping,
StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateRowsQuery,
StoredProviderModelMapping, StoredRequestedModelCandidateRowsQuery,
};

View File

@@ -40,6 +40,25 @@ pub struct StoredMinimalCandidateSelectionRow {
pub model_is_available: bool,
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct StoredPoolKeyCandidateRowsQuery {
pub api_format: String,
pub provider_id: String,
pub endpoint_id: String,
pub model_id: String,
pub selected_provider_model_name: String,
pub offset: u32,
pub limit: u32,
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct StoredRequestedModelCandidateRowsQuery {
pub api_format: String,
pub requested_model_name: String,
pub offset: u32,
pub limit: u32,
}
impl StoredMinimalCandidateSelectionRow {
pub fn supports_streaming(&self) -> bool {
self.model_supports_streaming
@@ -73,6 +92,22 @@ pub trait MinimalCandidateSelectionReadRepository: Send + Sync {
api_format: &str,
global_model_name: &str,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, crate::DataLayerError>;
async fn list_for_exact_api_format_and_requested_model(
&self,
api_format: &str,
requested_model_name: &str,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, crate::DataLayerError>;
async fn list_for_exact_api_format_and_requested_model_page(
&self,
query: &StoredRequestedModelCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, crate::DataLayerError>;
async fn list_pool_key_rows_for_group(
&self,
query: &StoredPoolKeyCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, crate::DataLayerError>;
}
pub trait MinimalCandidateSelectionRepository:

View File

@@ -2,7 +2,10 @@ use std::sync::RwLock;
use async_trait::async_trait;
use super::{MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow};
use super::{
MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow,
StoredPoolKeyCandidateRowsQuery, StoredRequestedModelCandidateRowsQuery,
};
use crate::DataLayerError;
#[derive(Debug, Default)]
@@ -68,6 +71,76 @@ impl MinimalCandidateSelectionReadRepository for InMemoryMinimalCandidateSelecti
.filter(|row| row.global_model_name == global_model_name)
.collect())
}
async fn list_for_exact_api_format_and_requested_model(
&self,
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
}
async fn list_for_exact_api_format_and_requested_model_page(
&self,
query: &StoredRequestedModelCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
let rows = self.list_for_exact_api_format(&query.api_format).await?;
let mut rows = rows
.into_iter()
.filter(|row| {
row_matches_requested_model(row, &query.requested_model_name, &query.api_format)
})
.collect::<Vec<_>>();
rows.sort_by(|left, right| {
left.global_model_name
.cmp(&right.global_model_name)
.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))
});
Ok(rows
.into_iter()
.skip(query.offset as usize)
.take(query.limit as usize)
.collect())
}
async fn list_pool_key_rows_for_group(
&self,
query: &StoredPoolKeyCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
let mut rows = self
.list_for_exact_api_format(&query.api_format)
.await?
.into_iter()
.filter(|row| {
row.provider_id == query.provider_id
&& row.endpoint_id == query.endpoint_id
&& 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))
});
Ok(rows
.into_iter()
.skip(query.offset as usize)
.take(query.limit as usize)
.collect())
}
}
fn normalize_api_format(value: &str) -> String {
@@ -78,6 +151,27 @@ fn api_format_matches(left: &str, right: &str) -> bool {
aether_ai_formats::api_format_alias_matches(left, right)
}
fn row_matches_requested_model(
row: &StoredMinimalCandidateSelectionRow,
requested_model_name: &str,
api_format: &str,
) -> bool {
row.global_model_name == requested_model_name
|| row.model_provider_model_name == requested_model_name
|| row
.model_provider_model_mappings
.as_ref()
.is_some_and(|mappings| {
mappings.iter().any(|mapping| {
mapping.api_formats.as_ref().is_none_or(|formats| {
formats
.iter()
.any(|value| api_format_matches(value, api_format))
}) && mapping.name == requested_model_name
})
})
}
fn key_auth_channel_matches(row: &StoredMinimalCandidateSelectionRow, api_format: &str) -> bool {
let provider_type = row.provider_type.trim().to_ascii_lowercase();
let auth_type = row.key_auth_type.trim().to_ascii_lowercase();
@@ -114,6 +208,7 @@ mod tests {
use super::InMemoryMinimalCandidateSelectionReadRepository;
use crate::repository::candidate_selection::{
MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow,
StoredPoolKeyCandidateRowsQuery, StoredRequestedModelCandidateRowsQuery,
};
fn sample_row(
@@ -174,6 +269,61 @@ mod tests {
assert_eq!(rows[1].provider_id, "provider-2");
}
#[tokio::test]
async fn filters_by_exact_api_format_and_requested_model_aliases() {
let mut mapped = sample_row("provider-1", "openai:chat", "gpt-4.1", 10);
mapped.model_provider_model_name = "provider-gpt-4.1".to_string();
mapped.model_provider_model_mappings = Some(vec![
crate::repository::candidate_selection::StoredProviderModelMapping {
name: "alias-gpt-4.1".to_string(),
priority: 0,
api_formats: Some(vec!["openai:chat".to_string()]),
},
]);
let repository = InMemoryMinimalCandidateSelectionReadRepository::seed(vec![
mapped,
sample_row("provider-2", "openai:chat", "gpt-4.1-mini", 20),
]);
let rows = repository
.list_for_exact_api_format_and_requested_model("openai:chat", "alias-gpt-4.1")
.await
.expect("list should succeed");
assert_eq!(rows.len(), 1);
assert_eq!(rows[0].provider_id, "provider-1");
}
#[tokio::test]
async fn requested_model_page_returns_requested_slice_only() {
let mut rows = Vec::new();
for index in 0..5 {
let mut row = sample_row(&format!("provider-{index}"), "openai:chat", "gpt-5", index);
row.key_internal_priority = index;
rows.push(row);
}
let repository = InMemoryMinimalCandidateSelectionReadRepository::seed(rows);
let page = repository
.list_for_exact_api_format_and_requested_model_page(
&StoredRequestedModelCandidateRowsQuery {
api_format: "openai:chat".to_string(),
requested_model_name: "gpt-5".to_string(),
offset: 2,
limit: 2,
},
)
.await
.expect("page should load");
assert_eq!(
page.iter()
.map(|row| row.provider_id.as_str())
.collect::<Vec<_>>(),
vec!["provider-2", "provider-3"]
);
}
#[tokio::test]
async fn filters_by_exact_api_format_only() {
let repository = InMemoryMinimalCandidateSelectionReadRepository::seed(vec![
@@ -191,4 +341,38 @@ mod tests {
assert_eq!(rows[0].provider_id, "provider-1");
assert_eq!(rows[1].provider_id, "provider-2");
}
#[tokio::test]
async fn list_pool_key_rows_for_group_returns_requested_page_only() {
let mut rows = Vec::new();
for index in 0..5 {
let mut row = sample_row("provider-pool", "openai:chat", "gpt-5", 10);
row.endpoint_id = "endpoint-pool".to_string();
row.model_id = "model-pool".to_string();
row.key_id = format!("key-{index}");
row.key_internal_priority = index;
rows.push(row);
}
let repository = InMemoryMinimalCandidateSelectionReadRepository::seed(rows);
let page = repository
.list_pool_key_rows_for_group(&StoredPoolKeyCandidateRowsQuery {
api_format: "openai:chat".to_string(),
provider_id: "provider-pool".to_string(),
endpoint_id: "endpoint-pool".to_string(),
model_id: "model-pool".to_string(),
selected_provider_model_name: "gpt-5".to_string(),
offset: 2,
limit: 2,
})
.await
.expect("pool key page should load");
assert_eq!(
page.iter()
.map(|row| row.key_id.as_str())
.collect::<Vec<_>>(),
vec!["key-2", "key-3"]
);
}
}

View File

@@ -4,7 +4,8 @@ mod sql;
#[allow(unused_imports)]
pub(crate) use aether_data_contracts::repository::candidate_selection::{
MinimalCandidateSelectionReadRepository, MinimalCandidateSelectionRepository,
StoredMinimalCandidateSelectionRow, StoredProviderModelMapping,
StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateRowsQuery,
StoredProviderModelMapping, StoredRequestedModelCandidateRowsQuery,
};
pub use memory::InMemoryMinimalCandidateSelectionReadRepository;
pub use sql::SqlxMinimalCandidateSelectionReadRepository;

View File

@@ -5,11 +5,13 @@ use std::collections::BTreeSet;
use super::{
MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow,
StoredProviderModelMapping,
StoredPoolKeyCandidateRowsQuery, StoredProviderModelMapping,
StoredRequestedModelCandidateRowsQuery,
};
use crate::{error::SqlxResultExt, DataLayerError};
const LIST_FOR_EXACT_API_FORMAT_SQL: &str = r#"
WITH candidate_rows AS (
SELECT
p.id AS provider_id,
p.name AS provider_name,
@@ -46,7 +48,8 @@ SELECT
m.provider_model_mappings AS model_provider_model_mappings,
m.supports_streaming AS model_supports_streaming,
m.is_active AS model_is_active,
m.is_available AS model_is_available
m.is_available AS model_is_available,
(p.config -> 'pool_advanced') IS NOT NULL AS provider_pool_enabled
FROM providers p
INNER JOIN provider_endpoints pe
ON pe.provider_id = p.id
@@ -124,17 +127,67 @@ WHERE p.is_active = TRUE
AND LOWER(BTRIM(pak.auth_type)) <> 'oauth'
)
)
),
pool_rows AS (
SELECT DISTINCT ON (provider_id, endpoint_id, model_id)
*
FROM candidate_rows
WHERE provider_pool_enabled
ORDER BY
provider_id ASC,
endpoint_id ASC,
model_id ASC,
key_internal_priority ASC,
key_id ASC
),
selected_rows AS (
SELECT * FROM candidate_rows WHERE NOT provider_pool_enabled
UNION ALL
SELECT * FROM pool_rows
)
SELECT
provider_id,
provider_name,
provider_type,
provider_priority,
provider_is_active,
endpoint_id,
endpoint_api_format,
endpoint_api_family,
endpoint_kind,
endpoint_is_active,
key_id,
key_name,
key_auth_type,
key_is_active,
key_api_formats,
key_allowed_models,
key_capabilities,
key_internal_priority,
key_global_priority_by_format,
model_id,
global_model_id,
global_model_name,
global_model_mappings,
global_model_supports_streaming,
model_provider_model_name,
model_provider_model_mappings,
model_supports_streaming,
model_is_active,
model_is_available
FROM selected_rows
ORDER BY
gm.name ASC,
p.provider_priority ASC,
pak.internal_priority ASC,
p.id ASC,
pe.id ASC,
pak.id ASC,
m.id ASC
global_model_name ASC,
provider_priority ASC,
key_internal_priority ASC,
provider_id ASC,
endpoint_id ASC,
key_id ASC,
model_id ASC
"#;
const LIST_FOR_EXACT_API_FORMAT_AND_GLOBAL_MODEL_SQL: &str = r#"
WITH candidate_rows AS (
SELECT
p.id AS provider_id,
p.name AS provider_name,
@@ -171,7 +224,8 @@ SELECT
m.provider_model_mappings AS model_provider_model_mappings,
m.supports_streaming AS model_supports_streaming,
m.is_active AS model_is_active,
m.is_available AS model_is_available
m.is_available AS model_is_available,
(p.config -> 'pool_advanced') IS NOT NULL AS provider_pool_enabled
FROM providers p
INNER JOIN provider_endpoints pe
ON pe.provider_id = p.id
@@ -250,13 +304,187 @@ WHERE p.is_active = TRUE
AND LOWER(BTRIM(pak.auth_type)) <> 'oauth'
)
)
),
pool_rows AS (
SELECT DISTINCT ON (provider_id, endpoint_id, model_id)
*
FROM candidate_rows
WHERE provider_pool_enabled
ORDER BY
provider_id ASC,
endpoint_id ASC,
model_id ASC,
key_internal_priority ASC,
key_id ASC
),
selected_rows AS (
SELECT * FROM candidate_rows WHERE NOT provider_pool_enabled
UNION ALL
SELECT * FROM pool_rows
)
SELECT
provider_id,
provider_name,
provider_type,
provider_priority,
provider_is_active,
endpoint_id,
endpoint_api_format,
endpoint_api_family,
endpoint_kind,
endpoint_is_active,
key_id,
key_name,
key_auth_type,
key_is_active,
key_api_formats,
key_allowed_models,
key_capabilities,
key_internal_priority,
key_global_priority_by_format,
model_id,
global_model_id,
global_model_name,
global_model_mappings,
global_model_supports_streaming,
model_provider_model_name,
model_provider_model_mappings,
model_supports_streaming,
model_is_active,
model_is_available
FROM selected_rows
ORDER BY
provider_priority ASC,
key_internal_priority ASC,
provider_id ASC,
endpoint_id ASC,
key_id ASC,
model_id ASC
"#;
const LIST_POOL_KEYS_FOR_GROUP_SQL: &str = r#"
SELECT
p.id AS provider_id,
p.name AS provider_name,
p.provider_type AS provider_type,
p.provider_priority AS provider_priority,
p.is_active AS provider_is_active,
pe.id AS endpoint_id,
pe.api_format AS endpoint_api_format,
pe.api_family AS endpoint_api_family,
pe.endpoint_kind AS endpoint_kind,
pe.is_active AS endpoint_is_active,
pak.id AS key_id,
pak.name AS key_name,
pak.auth_type AS key_auth_type,
pak.is_active AS key_is_active,
pak.api_formats AS key_api_formats,
pak.allowed_models AS key_allowed_models,
pak.capabilities AS key_capabilities,
pak.internal_priority AS key_internal_priority,
pak.global_priority_by_format AS key_global_priority_by_format,
m.id AS model_id,
m.global_model_id AS global_model_id,
gm.name AS global_model_name,
CASE
WHEN gm.config IS NOT NULL THEN gm.config -> 'model_mappings'
ELSE NULL
END AS global_model_mappings,
CASE
WHEN gm.config IS NOT NULL AND gm.config ? 'streaming'
THEN (gm.config ->> 'streaming')::BOOLEAN
ELSE NULL
END AS global_model_supports_streaming,
m.provider_model_name AS model_provider_model_name,
m.provider_model_mappings AS model_provider_model_mappings,
m.supports_streaming AS model_supports_streaming,
m.is_active AS model_is_active,
m.is_available AS model_is_available
FROM providers p
INNER JOIN provider_endpoints pe
ON pe.provider_id = p.id
INNER JOIN provider_api_keys pak
ON pak.provider_id = p.id
INNER JOIN models m
ON m.provider_id = p.id
INNER JOIN global_models gm
ON gm.id = m.global_model_id
WHERE p.is_active = TRUE
AND pe.is_active = TRUE
AND pak.is_active = TRUE
AND m.is_active = TRUE
AND m.is_available = TRUE
AND gm.is_active = TRUE
AND LOWER(pe.api_format) = LOWER($1)
AND p.id = $2
AND pe.id = $3
AND m.id = $4
AND (
pak.api_formats IS NULL
OR EXISTS (
SELECT 1
FROM json_array_elements_text(pak.api_formats) AS fmt(value)
WHERE LOWER(BTRIM(fmt.value)) = ANY($5::text[])
)
)
AND (
(
LOWER(BTRIM(p.provider_type)) = 'codex'
AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
AND LOWER($6) IN ('openai:responses', 'openai:responses:compact', 'openai:image')
)
OR (
LOWER(BTRIM(p.provider_type)) = 'claude_code'
AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
AND LOWER($6) = 'claude:messages'
)
OR (
LOWER(BTRIM(p.provider_type)) = 'kiro'
AND LOWER($6) = 'claude:messages'
AND (
LOWER(BTRIM(pak.auth_type)) = 'oauth'
OR (
LOWER(BTRIM(pak.auth_type)) = 'bearer'
AND pak.auth_config IS NOT NULL
AND BTRIM(pak.auth_config) <> ''
)
)
)
OR (
LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity')
AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
AND LOWER($6) = 'gemini:generate_content'
)
OR (
LOWER(BTRIM(p.provider_type)) = 'vertex_ai'
AND (
(
LOWER(BTRIM(pak.auth_type)) = 'api_key'
AND LOWER($6) = 'gemini:generate_content'
)
OR (
LOWER(BTRIM(pak.auth_type)) IN ('service_account', 'vertex_ai')
AND LOWER($6) IN ('claude:messages', 'gemini:generate_content')
)
)
)
OR (
LOWER(BTRIM(p.provider_type)) NOT IN (
'claude_code',
'codex',
'gemini_cli',
'vertex_ai',
'antigravity',
'kiro'
)
AND LOWER(BTRIM(pak.auth_type)) <> 'oauth'
)
)
ORDER BY
p.provider_priority ASC,
pak.internal_priority ASC,
p.id ASC,
pe.id ASC,
pak.id ASC,
m.id ASC
pak.id ASC
LIMIT $7
OFFSET $8
"#;
#[derive(Debug, Clone)]
@@ -336,6 +564,130 @@ impl SqlxMinimalCandidateSelectionReadRepository {
}
Ok(dedupe_candidate_selection_rows(rows))
}
pub async fn list_for_exact_api_format_and_requested_model(
&self,
api_format: &str,
requested_model_name: &str,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
let mut rows = Vec::new();
let canonical_api_format = normalize_api_format(api_format);
let storage_aliases = api_format_aliases(&canonical_api_format);
let sql_match_aliases = sql_match_aliases(&storage_aliases);
let sql = requested_model_selection_sql();
for api_format in storage_aliases {
rows.extend(
Self::collect_query_rows(
sqlx::query(sql.as_str())
.bind(api_format)
.bind(requested_model_name)
.bind(sql_match_aliases.clone())
.bind(canonical_api_format.clone())
.fetch(&self.pool),
map_candidate_selection_row,
)
.await?,
);
}
Ok(dedupe_candidate_selection_rows(rows))
}
pub async fn list_for_exact_api_format_and_requested_model_page(
&self,
query: &StoredRequestedModelCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
let mut rows = Vec::new();
let canonical_api_format = normalize_api_format(&query.api_format);
let storage_aliases = api_format_aliases(&canonical_api_format);
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 = requested_model_selection_page_sql();
for api_format in storage_aliases {
rows.extend(
Self::collect_query_rows(
sqlx::query(sql.as_str())
.bind(api_format)
.bind(query.requested_model_name.as_str())
.bind(sql_match_aliases.clone())
.bind(canonical_api_format.clone())
.bind(limit)
.bind(offset)
.fetch(&self.pool),
map_candidate_selection_row,
)
.await?,
);
}
Ok(dedupe_candidate_selection_rows(rows))
}
pub async fn list_pool_key_rows_for_group(
&self,
query: &StoredPoolKeyCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
let mut rows = Vec::new();
let canonical_api_format = normalize_api_format(&query.api_format);
let storage_aliases = api_format_aliases(&canonical_api_format);
let sql_match_aliases = sql_match_aliases(&storage_aliases);
let limit = i64::from(query.limit.max(1));
let offset = i64::from(query.offset);
for api_format in storage_aliases {
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),
map_candidate_selection_row,
)
.await?,
);
}
Ok(dedupe_candidate_selection_rows(rows))
}
}
fn requested_model_selection_sql() -> String {
LIST_FOR_EXACT_API_FORMAT_AND_GLOBAL_MODEL_SQL
.replace(
"AND gm.name = $2",
r#"AND (
gm.name = $2
OR m.provider_model_name = $2
OR (
jsonb_typeof(m.provider_model_mappings) = 'array'
AND EXISTS (
SELECT 1
FROM jsonb_array_elements(m.provider_model_mappings) AS mapping(value)
WHERE mapping.value ->> 'name' = $2
AND (
mapping.value -> 'api_formats' IS NULL
OR jsonb_typeof(mapping.value -> 'api_formats') <> 'array'
OR EXISTS (
SELECT 1
FROM jsonb_array_elements_text(mapping.value -> 'api_formats') AS fmt(value)
WHERE LOWER(BTRIM(fmt.value)) = ANY($3::text[])
)
)
)
)
)"#,
)
.replace(
"ORDER BY\n provider_priority ASC,",
"ORDER BY\n global_model_name ASC,\n provider_priority ASC,",
)
}
fn requested_model_selection_page_sql() -> String {
format!("{}\nLIMIT $5\nOFFSET $6", requested_model_selection_sql())
}
fn api_format_aliases(api_format: &str) -> Vec<String> {
@@ -384,6 +736,29 @@ impl MinimalCandidateSelectionReadRepository for SqlxMinimalCandidateSelectionRe
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
Self::list_for_exact_api_format_and_global_model(self, api_format, global_model_name).await
}
async fn list_for_exact_api_format_and_requested_model(
&self,
api_format: &str,
requested_model_name: &str,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
Self::list_for_exact_api_format_and_requested_model(self, api_format, requested_model_name)
.await
}
async fn list_for_exact_api_format_and_requested_model_page(
&self,
query: &StoredRequestedModelCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
Self::list_for_exact_api_format_and_requested_model_page(self, query).await
}
async fn list_pool_key_rows_for_group(
&self,
query: &StoredPoolKeyCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
Self::list_pool_key_rows_for_group(self, query).await
}
}
fn map_candidate_selection_row(
@@ -630,8 +1005,8 @@ mod tests {
use serde_json::json;
use super::{
parse_provider_model_mappings, parse_string_list,
SqlxMinimalCandidateSelectionReadRepository,
parse_provider_model_mappings, parse_string_list, requested_model_selection_page_sql,
requested_model_selection_sql, SqlxMinimalCandidateSelectionReadRepository,
};
use crate::postgres::{PostgresPoolConfig, PostgresPoolFactory};
use crate::repository::candidate_selection::StoredProviderModelMapping;
@@ -655,6 +1030,25 @@ mod tests {
let _ = repository.pool();
}
#[test]
fn requested_model_selection_sql_filters_before_row_materialization() {
let sql = requested_model_selection_sql();
assert!(sql.contains("m.provider_model_name = $2"));
assert!(sql.contains("jsonb_array_elements(m.provider_model_mappings)"));
assert!(!sql.contains("json_typeof(m.provider_model_mappings)"));
assert!(!sql.contains("json_array_elements_text(gm.config -> 'model_mappings')"));
assert!(sql.contains("ORDER BY\n global_model_name ASC,"));
assert!(!sql.contains("AND gm.name = $2\n AND"));
}
#[test]
fn requested_model_selection_page_sql_adds_limit_and_offset() {
let sql = requested_model_selection_page_sql();
assert!(sql.ends_with("LIMIT $5\nOFFSET $6"));
}
#[test]
fn parse_string_list_accepts_stringified_array() {
let parsed = parse_string_list(

View File

@@ -309,6 +309,29 @@ pub fn build_execution_request_candidate_seed(
Value::String(plan.endpoint_id.clone()),
);
context.insert("key_id".to_string(), Value::String(plan.key_id.clone()));
let mut extra_data = parse_request_candidate_report_context(Some(&Value::Object(
context.clone(),
)))
.and_then(|metadata| {
build_report_candidate_extra_data(ReportCandidateExtraDataInput {
client_api_format: metadata.client_api_format,
provider_api_format: metadata.provider_api_format,
upstream_url: metadata.upstream_url,
mapped_model: metadata.mapped_model,
key_name: metadata.key_name,
header_rules: metadata.header_rules,
body_rules: metadata.body_rules,
proxy: metadata.proxy,
error_flow: metadata.error_flow,
ranking_mode: metadata.ranking_mode,
priority_mode: metadata.priority_mode,
ranking_index: metadata.ranking_index,
priority_slot: metadata.priority_slot,
promoted_by: metadata.promoted_by,
demoted_by: metadata.demoted_by,
})
});
append_seed_extra_data_from_report_context(&mut extra_data, &context);
SchedulerExecutionRequestCandidateSeed {
upsert_record: UpsertRequestCandidateRecord {
@@ -331,7 +354,7 @@ pub fn build_execution_request_candidate_seed(
error_message: None,
latency_ms: None,
concurrent_requests: None,
extra_data: None,
extra_data,
required_capabilities: None,
created_at_unix_ms: Some(started_at_unix_ms),
started_at_unix_ms: Some(started_at_unix_ms),
@@ -341,6 +364,33 @@ pub fn build_execution_request_candidate_seed(
}
}
fn append_seed_extra_data_from_report_context(
extra_data: &mut Option<Value>,
context: &Map<String, Value>,
) {
const PASSTHROUGH_FIELDS: &[&str] = &[
"execution_strategy",
"conversion_mode",
"client_contract",
"provider_contract",
"transport_diagnostics",
];
let mut object = extra_data
.take()
.and_then(|value| match value {
Value::Object(object) => Some(object),
_ => None,
})
.unwrap_or_default();
for field in PASSTHROUGH_FIELDS {
if let Some(value) = context.get(*field).filter(|value| !value.is_null()) {
object.insert((*field).to_string(), value.clone());
}
}
*extra_data = (!object.is_empty()).then_some(Value::Object(object));
}
pub fn build_local_request_candidate_status_record(
input: LocalRequestCandidateStatusRecordInput<'_>,
) -> Option<UpsertRequestCandidateRecord> {