mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-09 02:47:45 +08:00
fix(routing): keep global model names out of provider alias reach
A provider model can be addressed by its upstream name or by any of its `provider_model_mappings` entries, and neither name is published in the model catalog, which lists global model names only. Resolution ran per API format, so an alias could win a format the real global model had no provider in: a `claude:messages` client asking for `gemini-3.8-flash` landed on the provider that merely renames its own `gemini-3.8-flash-cursor` model to `gemini-3.8-flash` on the way upstream, and the separate global model stopped distinguishing the two routes. Treat global model names as a reserved namespace instead: when the request names an active global model, only rows bound to it may serve it, whatever API format they sit in. Rows in hand answer that question for free whenever one of them is bound to a global model of that name, so the lookup stays off the path ordinary requests take. Authorization resolves the same way, so an API key's allowed models cannot be satisfied through a resolution candidate planning will no longer make. A request naming a model that is not a global model keeps every matching rule, so addressing a provider variant by its upstream name still works. Co-Authored-By: Claude Opus 5 <[email protected]>
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