mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-09 02:47:45 +08:00
Merge pull request #857 from stabey/codex/fix-global-model-name-reservation
fix(routing): reserve global model names across provider aliases
This commit is contained in:
@@ -6,7 +6,7 @@ use aether_routing_core::ResolvedRoutingPolicy;
|
||||
use aether_runtime::ConcurrencyPermit;
|
||||
use aether_scheduler_core::{
|
||||
enumerate_minimal_candidate_selection_with_model_directives, normalize_api_format,
|
||||
resolve_requested_global_model_name_with_model_directives_and_request_operation,
|
||||
resolve_requested_global_model_name_with_reserved_global_model,
|
||||
row_supports_requested_model_with_model_directives_and_request_operation,
|
||||
ClientSessionAffinity, EnumerateMinimalCandidateSelectionInput,
|
||||
SchedulerMinimalCandidateSelectionCandidate,
|
||||
@@ -378,6 +378,7 @@ pub(crate) struct LocalCandidatePreselectionPageCursor<'a> {
|
||||
requested_name_offsets: BTreeMap<String, u32>,
|
||||
scanned_rows_by_format: BTreeMap<String, u32>,
|
||||
resolved_global_model_names: BTreeMap<String, String>,
|
||||
reserved_global_model_names: BTreeMap<String, Option<String>>,
|
||||
fallback_offsets: BTreeMap<String, u32>,
|
||||
fallback_scan_epoch: u32,
|
||||
exhausted_api_formats: BTreeSet<String>,
|
||||
@@ -457,6 +458,7 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
|
||||
requested_name_offsets: BTreeMap::new(),
|
||||
scanned_rows_by_format: BTreeMap::new(),
|
||||
resolved_global_model_names: BTreeMap::new(),
|
||||
reserved_global_model_names: BTreeMap::new(),
|
||||
fallback_offsets: BTreeMap::new(),
|
||||
fallback_scan_epoch: 0,
|
||||
exhausted_api_formats: BTreeSet::new(),
|
||||
@@ -555,6 +557,7 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
|
||||
self.requested_name_offsets.clear();
|
||||
self.scanned_rows_by_format.clear();
|
||||
self.resolved_global_model_names.clear();
|
||||
self.reserved_global_model_names.clear();
|
||||
self.fallback_offsets.clear();
|
||||
self.fallback_scan_epoch = self.fallback_scan_epoch.wrapping_add(1);
|
||||
self.exhausted_api_formats.clear();
|
||||
@@ -1185,6 +1188,34 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
|
||||
|| self.exhausted_api_formats.contains(&normalized_api_format)
|
||||
}
|
||||
|
||||
/// Global model names are a reserved routing namespace, so a request that
|
||||
/// names one must not be answered by a provider whose own model merely
|
||||
/// carries that name as an upstream alias. Cached per routing model: the
|
||||
/// answer does not change between pages or API formats.
|
||||
async fn reserved_global_model_name(
|
||||
&mut self,
|
||||
rows: &[StoredMinimalCandidateSelectionRow],
|
||||
routing_model: &str,
|
||||
) -> Result<Option<String>, GatewayError> {
|
||||
if let Some(cached) = self.reserved_global_model_names.get(routing_model) {
|
||||
return Ok(cached.clone());
|
||||
}
|
||||
let state = self.state;
|
||||
let reserved_global_model_name =
|
||||
crate::data::candidate_selection::resolve_reserved_global_model_name(
|
||||
state.app().data.as_ref(),
|
||||
rows,
|
||||
routing_model,
|
||||
)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
self.reserved_global_model_names.insert(
|
||||
routing_model.to_string(),
|
||||
reserved_global_model_name.clone(),
|
||||
);
|
||||
Ok(reserved_global_model_name)
|
||||
}
|
||||
|
||||
async fn build_page_outcome_from_rows(
|
||||
&mut self,
|
||||
candidate_api_format: &str,
|
||||
@@ -1216,15 +1247,17 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
|
||||
if let Some(value) = self.resolved_global_model_names.get(normalized_api_format) {
|
||||
value.clone()
|
||||
} else {
|
||||
let Some(value) =
|
||||
resolve_requested_global_model_name_with_model_directives_and_request_operation(
|
||||
&rows,
|
||||
&routing_model,
|
||||
normalized_api_format,
|
||||
false,
|
||||
self.request_operation.as_deref(),
|
||||
)
|
||||
else {
|
||||
let reserved_global_model_name = self
|
||||
.reserved_global_model_name(&rows, &routing_model)
|
||||
.await?;
|
||||
let Some(value) = resolve_requested_global_model_name_with_reserved_global_model(
|
||||
&rows,
|
||||
&routing_model,
|
||||
normalized_api_format,
|
||||
false,
|
||||
self.request_operation.as_deref(),
|
||||
reserved_global_model_name.as_deref(),
|
||||
) else {
|
||||
return Ok(None);
|
||||
};
|
||||
self.resolved_global_model_names
|
||||
@@ -1475,6 +1508,7 @@ mod tests {
|
||||
use crate::AppState;
|
||||
use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY;
|
||||
use aether_data::repository::candidate_selection::InMemoryMinimalCandidateSelectionReadRepository;
|
||||
use aether_data::repository::global_models::InMemoryGlobalModelReadRepository;
|
||||
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
|
||||
use aether_data::DataLayerError;
|
||||
use aether_data_contracts::repository::candidate_selection::{
|
||||
@@ -1482,6 +1516,9 @@ mod tests {
|
||||
StoredPoolKeyCandidateRowsByKeyIdsQuery, StoredPoolKeyCandidateRowsQuery,
|
||||
StoredProviderModelMapping, StoredRequestedModelCandidateRowsQuery,
|
||||
};
|
||||
use aether_data_contracts::repository::global_models::{
|
||||
GlobalModelReadRepository, StoredPublicGlobalModel,
|
||||
};
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
||||
};
|
||||
@@ -2090,6 +2127,96 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
fn public_global_model(name: &str) -> StoredPublicGlobalModel {
|
||||
StoredPublicGlobalModel {
|
||||
id: format!("global-model-{name}"),
|
||||
name: name.to_string(),
|
||||
display_name: None,
|
||||
is_active: true,
|
||||
default_price_per_request: None,
|
||||
default_tiered_pricing: None,
|
||||
supported_capabilities: None,
|
||||
config: None,
|
||||
usage_count: 0,
|
||||
}
|
||||
}
|
||||
|
||||
/// The cursor provider reaches its upstream under a name that belongs to another
|
||||
/// global model. A `claude:messages` client asking for `gemini-3.8-flash` has to
|
||||
/// land on the provider bound to that global model — format conversion and all —
|
||||
/// rather than on the one that only borrows the name on the way out, which is the
|
||||
/// one an API-format-ordered scan reaches first.
|
||||
#[tokio::test]
|
||||
async fn paged_preselection_keeps_a_global_model_name_from_a_provider_alias() {
|
||||
let mut aliasing = standard_candidate_row("ursor", "claude:messages", 1);
|
||||
aliasing.global_model_id = "global-model-gemini-3.8-flash-cursor".to_string();
|
||||
aliasing.global_model_name = "gemini-3.8-flash-cursor".to_string();
|
||||
aliasing.model_provider_model_name = "gemini-3.8-flash-cursor".to_string();
|
||||
aliasing.model_provider_model_mappings = Some(vec![StoredProviderModelMapping {
|
||||
name: "gemini-3.8-flash".to_string(),
|
||||
priority: 1,
|
||||
api_formats: None,
|
||||
endpoint_ids: None,
|
||||
operations: None,
|
||||
}]);
|
||||
|
||||
let mut bound = standard_candidate_row("anti", "gemini:generate_content", 2);
|
||||
bound.global_model_id = "global-model-gemini-3.8-flash".to_string();
|
||||
bound.global_model_name = "gemini-3.8-flash".to_string();
|
||||
bound.model_provider_model_name = "gemini-3.8-flash".to_string();
|
||||
|
||||
let repository: Arc<dyn MinimalCandidateSelectionReadRepository> =
|
||||
Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed([
|
||||
aliasing, bound,
|
||||
]));
|
||||
let global_models: Arc<dyn GlobalModelReadRepository> =
|
||||
Arc::new(InMemoryGlobalModelReadRepository::seed([
|
||||
public_global_model("gemini-3.8-flash"),
|
||||
public_global_model("gemini-3.8-flash-cursor"),
|
||||
]));
|
||||
let data_state =
|
||||
GatewayDataState::with_minimal_candidate_selection_reader_for_tests(repository)
|
||||
.with_global_model_reader(global_models);
|
||||
let app = AppState::new()
|
||||
.expect("gateway state should build")
|
||||
.with_data_state_for_tests(data_state);
|
||||
let auth_snapshot = unrestricted_auth_snapshot();
|
||||
let model_directive_policy =
|
||||
crate::system_features::ModelDirectivePolicySnapshot::load(&app).await;
|
||||
let mut cursor = LocalCandidatePreselectionPageCursor::new(
|
||||
PlannerAppState::new(&app),
|
||||
&model_directive_policy,
|
||||
"claude:messages",
|
||||
"gemini-3.8-flash",
|
||||
None,
|
||||
false,
|
||||
None,
|
||||
&auth_snapshot,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
true,
|
||||
LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModelAndApiFormat,
|
||||
true,
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
|
||||
let page = cursor
|
||||
.next_page()
|
||||
.await
|
||||
.expect("preselection should succeed")
|
||||
.expect("the bound provider should still be reachable");
|
||||
|
||||
assert_eq!(page.candidates.len(), 1);
|
||||
assert_eq!(page.candidates[0].provider_name, "anti");
|
||||
assert_eq!(page.candidates[0].global_model_name, "gemini-3.8-flash");
|
||||
assert_eq!(
|
||||
page.candidates[0].endpoint_api_format,
|
||||
"gemini:generate_content"
|
||||
);
|
||||
}
|
||||
|
||||
fn standard_candidate_row(
|
||||
provider_id: &str,
|
||||
api_format: &str,
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
use axum::body::Bytes;
|
||||
use axum::http::Uri;
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use super::super::GatewayControlDecision;
|
||||
use super::credentials::{contains_string, extract_requested_model};
|
||||
@@ -747,6 +748,11 @@ async fn request_model_resolves_to_allowed_model(
|
||||
return Ok(false);
|
||||
};
|
||||
|
||||
// Global model names are a reserved routing namespace, so authorization has to
|
||||
// resolve a request the same way candidate planning will: a provider whose own
|
||||
// model carries the requested name only as an upstream alias must not make the
|
||||
// request resolve to that provider's global model.
|
||||
let mut reserved_global_model_names: BTreeMap<String, Option<String>> = BTreeMap::new();
|
||||
for api_format in candidate_api_formats_for_model_resolution(&client_api_format) {
|
||||
let resolution = decision
|
||||
.model_directive_policy
|
||||
@@ -762,23 +768,45 @@ async fn request_model_resolves_to_allowed_model(
|
||||
.list_minimal_candidate_selection_rows_for_api_format(&api_format)
|
||||
.await?
|
||||
};
|
||||
let reserved_global_model_name = match reserved_global_model_names.get(routing_model) {
|
||||
Some(cached) => cached.clone(),
|
||||
None => {
|
||||
let reserved_global_model_name =
|
||||
crate::data::candidate_selection::resolve_reserved_global_model_name(
|
||||
state.data.as_ref(),
|
||||
&rows,
|
||||
routing_model,
|
||||
)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
reserved_global_model_names.insert(
|
||||
routing_model.to_string(),
|
||||
reserved_global_model_name.clone(),
|
||||
);
|
||||
reserved_global_model_name
|
||||
}
|
||||
};
|
||||
let matching_rows = rows
|
||||
.into_iter()
|
||||
.filter(|row| {
|
||||
aether_scheduler_core::row_supports_requested_model_with_model_directives(
|
||||
aether_scheduler_core::row_supports_requested_model_with_reserved_global_model(
|
||||
row,
|
||||
routing_model,
|
||||
&api_format,
|
||||
false,
|
||||
None,
|
||||
reserved_global_model_name.as_deref(),
|
||||
)
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
let Some(resolved_global_model) =
|
||||
aether_scheduler_core::resolve_requested_global_model_name_with_model_directives(
|
||||
aether_scheduler_core::resolve_requested_global_model_name_with_reserved_global_model(
|
||||
&matching_rows,
|
||||
routing_model,
|
||||
&api_format,
|
||||
false,
|
||||
None,
|
||||
reserved_global_model_name.as_deref(),
|
||||
)
|
||||
else {
|
||||
continue;
|
||||
|
||||
@@ -6,7 +6,7 @@ use aether_data_contracts::repository::candidate_selection::{
|
||||
use aether_scheduler_core::{
|
||||
auth_constraints_allow_api_format, collect_global_model_names_for_required_capability,
|
||||
enumerate_minimal_candidate_selection_with_model_directives, normalize_api_format,
|
||||
resolve_requested_global_model_name_with_model_directives,
|
||||
resolve_requested_global_model_name_with_reserved_global_model,
|
||||
row_supports_requested_model_with_model_directives, EnumerateMinimalCandidateSelectionInput,
|
||||
SchedulerAuthConstraints, SchedulerMinimalCandidateSelectionCandidate,
|
||||
};
|
||||
@@ -56,6 +56,37 @@ pub(crate) trait MinimalCandidateSelectionRowSource {
|
||||
&self,
|
||||
query: &StoredPoolKeyCandidateRowsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError>;
|
||||
|
||||
/// Returns the canonical global model name when `model_name` is one, so the
|
||||
/// caller can keep provider-side aliases out of a request that names a
|
||||
/// global model. Sources without a global model reader answer `None`, which
|
||||
/// leaves matching unrestricted.
|
||||
async fn read_reserved_global_model_name(
|
||||
&self,
|
||||
_model_name: &str,
|
||||
) -> Result<Option<String>, DataLayerError> {
|
||||
Ok(None)
|
||||
}
|
||||
}
|
||||
|
||||
/// Resolves the reserved global model name for `routing_model`.
|
||||
///
|
||||
/// Rows already in hand answer the question for free whenever one of them is
|
||||
/// bound to a global model of that exact name; only a request that no local row
|
||||
/// claims as a global model needs the lookup, which keeps the extra read off the
|
||||
/// path every ordinary request takes.
|
||||
pub(crate) async fn resolve_reserved_global_model_name(
|
||||
source: &(impl MinimalCandidateSelectionRowSource + Sync),
|
||||
rows: &[StoredMinimalCandidateSelectionRow],
|
||||
routing_model: &str,
|
||||
) -> Result<Option<String>, DataLayerError> {
|
||||
if rows
|
||||
.iter()
|
||||
.any(|row| row.global_model_name == routing_model)
|
||||
{
|
||||
return Ok(Some(routing_model.to_string()));
|
||||
}
|
||||
source.read_reserved_global_model_name(routing_model).await
|
||||
}
|
||||
|
||||
pub(crate) const REQUESTED_MODEL_CANDIDATE_PAGE_SIZE: u32 = 256;
|
||||
@@ -102,12 +133,16 @@ pub(crate) async fn read_requested_model_rows(
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
let reserved_global_model_name =
|
||||
resolve_reserved_global_model_name(state, &rows, requested_model_name).await?;
|
||||
let Some(resolved_global_model_name) =
|
||||
resolve_requested_global_model_name_with_model_directives(
|
||||
resolve_requested_global_model_name_with_reserved_global_model(
|
||||
&rows,
|
||||
requested_model_name,
|
||||
api_format,
|
||||
enable_model_directives,
|
||||
None,
|
||||
reserved_global_model_name.as_deref(),
|
||||
)
|
||||
else {
|
||||
return Ok(None);
|
||||
|
||||
@@ -185,6 +185,16 @@ impl MinimalCandidateSelectionRowSource for GatewayDataState {
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
self.list_pool_key_candidate_rows_for_group(query).await
|
||||
}
|
||||
|
||||
async fn read_reserved_global_model_name(
|
||||
&self,
|
||||
model_name: &str,
|
||||
) -> Result<Option<String>, DataLayerError> {
|
||||
Ok(self
|
||||
.get_public_global_model_by_name(model_name)
|
||||
.await?
|
||||
.map(|global_model| global_model.name))
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
|
||||
@@ -1,8 +1,12 @@
|
||||
use std::sync::Arc;
|
||||
|
||||
use aether_data::repository::candidate_selection::InMemoryMinimalCandidateSelectionReadRepository;
|
||||
use aether_data::repository::global_models::InMemoryGlobalModelReadRepository;
|
||||
use aether_data::repository::quota::InMemoryProviderQuotaRepository;
|
||||
use aether_data_contracts::repository::candidate_selection::StoredProviderModelMapping;
|
||||
use aether_data_contracts::repository::candidate_selection::{
|
||||
StoredMinimalCandidateSelectionRow, StoredProviderModelMapping,
|
||||
};
|
||||
use aether_data_contracts::repository::global_models::StoredPublicGlobalModel;
|
||||
use aether_scheduler_core::{
|
||||
resolve_requested_global_model_name, SchedulerMinimalCandidateSelectionCandidate,
|
||||
};
|
||||
@@ -148,6 +152,157 @@ fn scheduler_candidate_is_serializable() {
|
||||
assert_eq!(json["provider_name"], "OpenAI");
|
||||
}
|
||||
|
||||
/// A provider whose own model reaches the upstream under a borrowed name must not
|
||||
/// answer for that name: `gemini-3.8-flash` belongs to the providers bound to that
|
||||
/// global model, even when the only provider serving the client's own API format is
|
||||
/// the one that merely renames its model on the way out.
|
||||
#[tokio::test]
|
||||
async fn provider_alias_does_not_capture_a_request_naming_another_global_model() {
|
||||
let row = cursor_alias_row();
|
||||
let candidates = Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(vec![
|
||||
row.clone(),
|
||||
]));
|
||||
let quotas = Arc::new(InMemoryProviderQuotaRepository::seed(vec![]));
|
||||
let global_models = Arc::new(InMemoryGlobalModelReadRepository::seed(vec![
|
||||
public_global_model("gemini-3.8-flash"),
|
||||
public_global_model("gemini-3.8-flash-cursor"),
|
||||
]));
|
||||
let state = GatewayDataState::with_candidate_selection_and_quota_for_tests(candidates, quotas)
|
||||
.with_global_model_reader(global_models);
|
||||
|
||||
let hijacked = enumerate_minimal_candidate_selection_with_required_capabilities(
|
||||
&state,
|
||||
"claude:messages",
|
||||
"gemini-3.8-flash",
|
||||
false,
|
||||
None,
|
||||
None,
|
||||
false,
|
||||
)
|
||||
.await
|
||||
.expect("selection should succeed");
|
||||
assert!(hijacked.is_empty());
|
||||
|
||||
let selection = enumerate_minimal_candidate_selection_with_required_capabilities(
|
||||
&state,
|
||||
"claude:messages",
|
||||
"gemini-3.8-flash-cursor",
|
||||
false,
|
||||
None,
|
||||
None,
|
||||
false,
|
||||
)
|
||||
.await
|
||||
.expect("selection should succeed");
|
||||
assert_eq!(selection.len(), 1);
|
||||
assert_eq!(selection[0].global_model_name, "gemini-3.8-flash-cursor");
|
||||
assert_eq!(
|
||||
selection[0].selected_provider_model_name,
|
||||
"gemini-3.8-flash"
|
||||
);
|
||||
}
|
||||
|
||||
/// Pins what the rule is worth: the very same rows hand the request to the aliasing
|
||||
/// provider as soon as nothing can tell that `gemini-3.8-flash` is a global model.
|
||||
#[tokio::test]
|
||||
async fn provider_alias_captures_the_request_without_a_global_model_reader() {
|
||||
let candidates = Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(vec![
|
||||
cursor_alias_row(),
|
||||
]));
|
||||
let quotas = Arc::new(InMemoryProviderQuotaRepository::seed(vec![]));
|
||||
let state = GatewayDataState::with_candidate_selection_and_quota_for_tests(candidates, quotas);
|
||||
|
||||
let selection = enumerate_minimal_candidate_selection_with_required_capabilities(
|
||||
&state,
|
||||
"claude:messages",
|
||||
"gemini-3.8-flash",
|
||||
false,
|
||||
None,
|
||||
None,
|
||||
false,
|
||||
)
|
||||
.await
|
||||
.expect("selection should succeed");
|
||||
assert_eq!(selection.len(), 1);
|
||||
assert_eq!(selection[0].global_model_name, "gemini-3.8-flash-cursor");
|
||||
}
|
||||
|
||||
fn cursor_alias_row() -> StoredMinimalCandidateSelectionRow {
|
||||
let mut row = sample_row();
|
||||
row.endpoint_api_format = "claude:messages".to_string();
|
||||
row.endpoint_api_family = Some("claude".to_string());
|
||||
row.endpoint_kind = Some("messages".to_string());
|
||||
row.key_api_formats = Some(vec!["claude:messages".to_string()]);
|
||||
row.key_global_priority_by_format = None;
|
||||
row.global_model_id = "global-gemini-3.8-flash-cursor".to_string();
|
||||
row.global_model_name = "gemini-3.8-flash-cursor".to_string();
|
||||
row.global_model_mappings = None;
|
||||
row.model_provider_model_name = "gemini-3.8-flash-cursor".to_string();
|
||||
row.model_provider_model_mappings = Some(vec![StoredProviderModelMapping {
|
||||
name: "gemini-3.8-flash".to_string(),
|
||||
priority: 1,
|
||||
api_formats: None,
|
||||
endpoint_ids: None,
|
||||
operations: None,
|
||||
}]);
|
||||
row
|
||||
}
|
||||
|
||||
/// Without a global model of that name the alias stays addressable, which is what a
|
||||
/// provider-scoped variant name relies on.
|
||||
#[tokio::test]
|
||||
async fn provider_alias_stays_addressable_without_a_global_model_of_that_name() {
|
||||
let mut row = sample_row();
|
||||
row.global_model_name = "gpt-5".to_string();
|
||||
row.model_provider_model_name = "gpt-5-upstream".to_string();
|
||||
row.model_provider_model_mappings = Some(vec![StoredProviderModelMapping {
|
||||
name: "gpt-5-alias".to_string(),
|
||||
priority: 1,
|
||||
api_formats: Some(vec!["openai:chat".to_string()]),
|
||||
endpoint_ids: None,
|
||||
operations: None,
|
||||
}]);
|
||||
|
||||
let candidates = Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(vec![
|
||||
row,
|
||||
]));
|
||||
let quotas = Arc::new(InMemoryProviderQuotaRepository::seed(vec![]));
|
||||
let global_models = Arc::new(InMemoryGlobalModelReadRepository::seed(vec![
|
||||
public_global_model("gpt-5"),
|
||||
]));
|
||||
let state = GatewayDataState::with_candidate_selection_and_quota_for_tests(candidates, quotas)
|
||||
.with_global_model_reader(global_models);
|
||||
|
||||
let selection = enumerate_minimal_candidate_selection_with_required_capabilities(
|
||||
&state,
|
||||
"openai:chat",
|
||||
"gpt-5-alias",
|
||||
false,
|
||||
None,
|
||||
None,
|
||||
false,
|
||||
)
|
||||
.await
|
||||
.expect("selection should succeed");
|
||||
assert_eq!(selection.len(), 1);
|
||||
assert_eq!(selection[0].global_model_name, "gpt-5");
|
||||
assert_eq!(selection[0].selected_provider_model_name, "gpt-5-alias");
|
||||
}
|
||||
|
||||
fn public_global_model(name: &str) -> StoredPublicGlobalModel {
|
||||
StoredPublicGlobalModel {
|
||||
id: format!("global-{name}"),
|
||||
name: name.to_string(),
|
||||
display_name: None,
|
||||
is_active: true,
|
||||
default_price_per_request: None,
|
||||
default_tiered_pricing: None,
|
||||
supported_capabilities: None,
|
||||
config: None,
|
||||
usage_count: 0,
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn enumerate_minimal_candidate_selection_resolves_provider_model_alias() {
|
||||
let mut row = sample_row();
|
||||
|
||||
@@ -46,9 +46,11 @@ pub use model::{
|
||||
resolve_provider_model_name_with_model_directives_and_request_operation,
|
||||
resolve_requested_global_model_name, resolve_requested_global_model_name_with_model_directives,
|
||||
resolve_requested_global_model_name_with_model_directives_and_request_operation,
|
||||
row_supports_requested_model, row_supports_requested_model_with_model_directives,
|
||||
resolve_requested_global_model_name_with_reserved_global_model, row_supports_requested_model,
|
||||
row_supports_requested_model_with_model_directives,
|
||||
row_supports_requested_model_with_model_directives_and_request_operation,
|
||||
row_supports_required_capability, select_provider_model_name,
|
||||
row_supports_requested_model_with_reserved_global_model, row_supports_required_capability,
|
||||
select_provider_model_name,
|
||||
};
|
||||
pub use provider::{build_provider_concurrent_limit_map, should_skip_provider_quota};
|
||||
pub use ranking::{
|
||||
|
||||
@@ -42,22 +42,54 @@ pub fn resolve_requested_global_model_name_with_model_directives_and_request_ope
|
||||
api_format: &str,
|
||||
enable_model_directives: bool,
|
||||
request_operation: Option<&str>,
|
||||
) -> Option<String> {
|
||||
resolve_requested_global_model_name_with_reserved_global_model(
|
||||
rows,
|
||||
requested_model_name,
|
||||
api_format,
|
||||
enable_model_directives,
|
||||
request_operation,
|
||||
None,
|
||||
)
|
||||
}
|
||||
|
||||
/// Global model names are a reserved routing namespace.
|
||||
///
|
||||
/// A provider model can also be reached by its upstream name
|
||||
/// (`provider_model_name`) or by one of its `provider_model_mappings` entries.
|
||||
/// Neither of those names is published in the model catalog, which lists global
|
||||
/// model names only, so they must not capture a request that names a global
|
||||
/// model: `gemini-3.8-flash` belongs to the providers bound to that global
|
||||
/// model, not to a provider that merely renames its own model to
|
||||
/// `gemini-3.8-flash` on the way upstream.
|
||||
///
|
||||
/// `reserved_global_model_name` carries the canonical global model name when the
|
||||
/// requested name is one of them, which puts every row bound to a different
|
||||
/// global model out of the running. `None` keeps all matching rules available,
|
||||
/// which is what a request for a provider-side alias needs.
|
||||
pub fn resolve_requested_global_model_name_with_reserved_global_model(
|
||||
rows: &[StoredMinimalCandidateSelectionRow],
|
||||
requested_model_name: &str,
|
||||
api_format: &str,
|
||||
enable_model_directives: bool,
|
||||
request_operation: Option<&str>,
|
||||
reserved_global_model_name: Option<&str>,
|
||||
) -> Option<String> {
|
||||
requested_model_name_candidates(requested_model_name, enable_model_directives).find_map(
|
||||
|requested_model_name| {
|
||||
let requested_model_name = requested_model_name.as_ref();
|
||||
resolve_global_model_name_by(rows, |row| {
|
||||
resolve_global_model_name_by(rows, reserved_global_model_name, |row| {
|
||||
row_has_available_provider_model(row, api_format, request_operation)
|
||||
&& row.global_model_name == requested_model_name
|
||||
})
|
||||
.or_else(|| {
|
||||
resolve_global_model_name_by(rows, |row| {
|
||||
resolve_global_model_name_by(rows, reserved_global_model_name, |row| {
|
||||
row_default_provider_model_name_available(row, api_format, request_operation)
|
||||
&& row.model_provider_model_name == requested_model_name
|
||||
})
|
||||
})
|
||||
.or_else(|| {
|
||||
resolve_global_model_name_by(rows, |row| {
|
||||
resolve_global_model_name_by(rows, reserved_global_model_name, |row| {
|
||||
row.model_provider_model_mappings
|
||||
.as_ref()
|
||||
.is_some_and(|mappings| {
|
||||
@@ -69,7 +101,7 @@ pub fn resolve_requested_global_model_name_with_model_directives_and_request_ope
|
||||
})
|
||||
})
|
||||
.or_else(|| {
|
||||
resolve_global_model_name_by(rows, |row| {
|
||||
resolve_global_model_name_by(rows, reserved_global_model_name, |row| {
|
||||
row_has_available_provider_model(row, api_format, request_operation)
|
||||
&& row.global_model_mappings.as_ref().is_some_and(|patterns| {
|
||||
patterns
|
||||
@@ -82,6 +114,13 @@ pub fn resolve_requested_global_model_name_with_model_directives_and_request_ope
|
||||
)
|
||||
}
|
||||
|
||||
fn reserved_global_model_allows_row(
|
||||
reserved_global_model_name: Option<&str>,
|
||||
row: &StoredMinimalCandidateSelectionRow,
|
||||
) -> bool {
|
||||
reserved_global_model_name.is_none_or(|reserved| row.global_model_name == reserved)
|
||||
}
|
||||
|
||||
pub fn row_supports_requested_model(
|
||||
row: &StoredMinimalCandidateSelectionRow,
|
||||
requested_model_name: &str,
|
||||
@@ -112,18 +151,39 @@ pub fn row_supports_requested_model_with_model_directives_and_request_operation(
|
||||
enable_model_directives: bool,
|
||||
request_operation: Option<&str>,
|
||||
) -> bool {
|
||||
requested_model_name_candidates(requested_model_name, enable_model_directives).any(
|
||||
|requested_model_name| {
|
||||
row_supports_requested_model_exact(
|
||||
row,
|
||||
requested_model_name.as_ref(),
|
||||
api_format,
|
||||
request_operation,
|
||||
)
|
||||
},
|
||||
row_supports_requested_model_with_reserved_global_model(
|
||||
row,
|
||||
requested_model_name,
|
||||
api_format,
|
||||
enable_model_directives,
|
||||
request_operation,
|
||||
None,
|
||||
)
|
||||
}
|
||||
|
||||
/// See [`resolve_requested_global_model_name_with_reserved_global_model`] for what
|
||||
/// `reserved_global_model_name` means.
|
||||
pub fn row_supports_requested_model_with_reserved_global_model(
|
||||
row: &StoredMinimalCandidateSelectionRow,
|
||||
requested_model_name: &str,
|
||||
api_format: &str,
|
||||
enable_model_directives: bool,
|
||||
request_operation: Option<&str>,
|
||||
reserved_global_model_name: Option<&str>,
|
||||
) -> bool {
|
||||
reserved_global_model_allows_row(reserved_global_model_name, row)
|
||||
&& requested_model_name_candidates(requested_model_name, enable_model_directives).any(
|
||||
|requested_model_name| {
|
||||
row_supports_requested_model_exact(
|
||||
row,
|
||||
requested_model_name.as_ref(),
|
||||
api_format,
|
||||
request_operation,
|
||||
)
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
fn row_supports_requested_model_exact(
|
||||
row: &StoredMinimalCandidateSelectionRow,
|
||||
requested_model_name: &str,
|
||||
@@ -152,13 +212,16 @@ fn row_supports_requested_model_exact(
|
||||
|
||||
fn resolve_global_model_name_by<F>(
|
||||
rows: &[StoredMinimalCandidateSelectionRow],
|
||||
reserved_global_model_name: Option<&str>,
|
||||
matches: F,
|
||||
) -> Option<String>
|
||||
where
|
||||
F: Fn(&StoredMinimalCandidateSelectionRow) -> bool,
|
||||
{
|
||||
let mut best_match = None::<&str>;
|
||||
for row in rows.iter().filter(|row| matches(row)) {
|
||||
for row in rows.iter().filter(|row| {
|
||||
reserved_global_model_allows_row(reserved_global_model_name, row) && matches(row)
|
||||
}) {
|
||||
let candidate = row.global_model_name.trim();
|
||||
if candidate.is_empty() {
|
||||
continue;
|
||||
@@ -561,8 +624,10 @@ mod tests {
|
||||
matches_model_mapping, resolve_provider_model_name,
|
||||
resolve_provider_model_name_with_model_directives,
|
||||
resolve_provider_model_name_with_model_directives_and_request_operation,
|
||||
resolve_requested_global_model_name_with_model_directives, row_supports_requested_model,
|
||||
row_supports_requested_model_with_model_directives,
|
||||
resolve_requested_global_model_name_with_model_directives,
|
||||
resolve_requested_global_model_name_with_reserved_global_model,
|
||||
row_supports_requested_model, row_supports_requested_model_with_model_directives,
|
||||
row_supports_requested_model_with_reserved_global_model,
|
||||
};
|
||||
use aether_data_contracts::repository::candidate_selection::{
|
||||
StoredMinimalCandidateSelectionRow, StoredProviderModelMapping,
|
||||
@@ -899,6 +964,111 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
fn cursor_alias_row() -> StoredMinimalCandidateSelectionRow {
|
||||
let mut row = sample_row("gemini-3.8-flash-cursor", "gemini-3.8-flash-cursor");
|
||||
row.endpoint_api_format = "claude:messages".to_string();
|
||||
row.model_provider_model_mappings = Some(vec![StoredProviderModelMapping {
|
||||
name: "gemini-3.8-flash".to_string(),
|
||||
priority: 1,
|
||||
api_formats: None,
|
||||
endpoint_ids: None,
|
||||
operations: None,
|
||||
}]);
|
||||
row
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn reserved_global_model_name_rejects_provider_alias_from_another_global_model() {
|
||||
let row = cursor_alias_row();
|
||||
|
||||
assert!(row_supports_requested_model_with_reserved_global_model(
|
||||
&row,
|
||||
"gemini-3.8-flash",
|
||||
"claude:messages",
|
||||
false,
|
||||
None,
|
||||
None,
|
||||
));
|
||||
assert!(!row_supports_requested_model_with_reserved_global_model(
|
||||
&row,
|
||||
"gemini-3.8-flash",
|
||||
"claude:messages",
|
||||
false,
|
||||
None,
|
||||
Some("gemini-3.8-flash"),
|
||||
));
|
||||
assert_eq!(
|
||||
resolve_requested_global_model_name_with_reserved_global_model(
|
||||
&[row],
|
||||
"gemini-3.8-flash",
|
||||
"claude:messages",
|
||||
false,
|
||||
None,
|
||||
Some("gemini-3.8-flash"),
|
||||
),
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn reserved_global_model_name_keeps_its_own_rows_addressable() {
|
||||
let row = cursor_alias_row();
|
||||
|
||||
assert!(row_supports_requested_model_with_reserved_global_model(
|
||||
&row,
|
||||
"gemini-3.8-flash-cursor",
|
||||
"claude:messages",
|
||||
false,
|
||||
None,
|
||||
Some("gemini-3.8-flash-cursor"),
|
||||
));
|
||||
assert_eq!(
|
||||
resolve_requested_global_model_name_with_reserved_global_model(
|
||||
&[row],
|
||||
"gemini-3.8-flash-cursor",
|
||||
"claude:messages",
|
||||
false,
|
||||
None,
|
||||
Some("gemini-3.8-flash-cursor"),
|
||||
)
|
||||
.as_deref(),
|
||||
Some("gemini-3.8-flash-cursor")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_alias_stays_addressable_when_it_is_not_a_global_model_name() {
|
||||
let mut row = sample_row("gpt-5", "gpt-5-upstream");
|
||||
row.model_provider_model_mappings = Some(vec![StoredProviderModelMapping {
|
||||
name: "gpt-5-alias".to_string(),
|
||||
priority: 1,
|
||||
api_formats: None,
|
||||
endpoint_ids: None,
|
||||
operations: None,
|
||||
}]);
|
||||
|
||||
assert!(row_supports_requested_model_with_reserved_global_model(
|
||||
&row,
|
||||
"gpt-5-alias",
|
||||
"openai:chat",
|
||||
false,
|
||||
None,
|
||||
None,
|
||||
));
|
||||
assert_eq!(
|
||||
resolve_requested_global_model_name_with_reserved_global_model(
|
||||
&[row],
|
||||
"gpt-5-alias",
|
||||
"openai:chat",
|
||||
false,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.as_deref(),
|
||||
Some("gpt-5")
|
||||
);
|
||||
}
|
||||
|
||||
fn sample_row(
|
||||
global_model_name: &str,
|
||||
model_provider_model_name: &str,
|
||||
|
||||
Reference in New Issue
Block a user