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:
ZheFox
2026-10-06 22:15:36 +08:00
committed by GitHub
7 changed files with 560 additions and 33 deletions
@@ -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,
+30 -2
View File
@@ -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();