mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
Refactor pool candidate scheduling
This commit is contained in:
@@ -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(
|
||||
|
||||
@@ -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!(
|
||||
|
||||
@@ -2,5 +2,6 @@ mod types;
|
||||
|
||||
pub use types::{
|
||||
MinimalCandidateSelectionReadRepository, MinimalCandidateSelectionRepository,
|
||||
StoredMinimalCandidateSelectionRow, StoredProviderModelMapping,
|
||||
StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateRowsQuery,
|
||||
StoredProviderModelMapping, StoredRequestedModelCandidateRowsQuery,
|
||||
};
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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"]
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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> {
|
||||
|
||||
Reference in New Issue
Block a user