Merge pull request #874 from hkxiaoyao/fix/model-mapping-endpoint-scope

fix(scheduler): confine model mappings to endpoint and api-format scope
This commit is contained in:
ZheFox
2026-10-09 14:24:36 +08:00
committed by GitHub
6 changed files with 244 additions and 51 deletions
@@ -19,7 +19,9 @@ use super::super::{
list_selectable_candidates_for_required_capability_without_requested_model,
list_selectable_candidates_for_required_capability_without_requested_model_with_auth_limit_signal,
};
use super::support::{sample_auth_snapshot, sample_provider, sample_row};
use super::support::{
sample_auth_snapshot, sample_provider, sample_row, sample_row_without_model_mappings,
};
#[tokio::test]
async fn compatible_required_capability_prefers_matching_keys_without_hard_filtering() {
@@ -77,7 +79,7 @@ async fn compatible_required_capability_prefers_matching_keys_without_hard_filte
#[tokio::test]
async fn exclusive_required_capability_keeps_hard_filtering_only_matching_keys() {
let mut incompatible = sample_row();
let mut incompatible = sample_row_without_model_mappings();
incompatible.provider_id = "provider-a".to_string();
incompatible.provider_name = "provider-a".to_string();
incompatible.endpoint_id = "endpoint-a".to_string();
@@ -89,7 +91,7 @@ async fn exclusive_required_capability_keeps_hard_filtering_only_matching_keys()
incompatible.global_model_name = "gemini-2.5-pro".to_string();
incompatible.key_capabilities = Some(serde_json::json!({}));
let mut compatible = sample_row();
let mut compatible = sample_row_without_model_mappings();
compatible.provider_id = "provider-b".to_string();
compatible.provider_name = "provider-b".to_string();
compatible.endpoint_id = "endpoint-b".to_string();
@@ -133,7 +135,7 @@ async fn exclusive_required_capability_keeps_hard_filtering_only_matching_keys()
#[tokio::test]
async fn required_capability_without_model_uses_session_scoped_affinity() {
let mut fallback = sample_row();
let mut fallback = sample_row_without_model_mappings();
fallback.provider_id = "provider-a".to_string();
fallback.provider_name = "provider-a".to_string();
fallback.endpoint_id = "endpoint-a".to_string();
@@ -35,7 +35,10 @@ use super::super::selection::{
collect_selectable_candidates_with_skip_reasons as collect_selectable_candidates_with_skip_reasons_impl,
is_exact_all_skipped_by_auth_limit, select_minimal_candidate as select_candidate_impl,
};
use super::support::{sample_auth_snapshot, sample_key, sample_provider, sample_row};
use super::support::{
sample_auth_snapshot, sample_key, sample_provider, sample_row,
sample_row_without_model_mappings,
};
async fn state_with_routing_default_policy(
data_state: GatewayDataState,
@@ -2225,7 +2228,7 @@ async fn keeps_codex_candidate_selectable_when_oauth_token_is_expired() {
#[tokio::test]
async fn keeps_refreshable_kiro_candidate_selectable_with_runtime_oauth_invalid_marker() {
let mut row = sample_row();
let mut row = sample_row_without_model_mappings();
row.provider_id = "provider-kiro".to_string();
row.provider_name = "kiro".to_string();
row.provider_type = "kiro".to_string();
@@ -2286,7 +2289,7 @@ async fn keeps_refreshable_kiro_candidate_selectable_with_runtime_oauth_invalid_
#[tokio::test]
async fn keeps_refreshable_kiro_candidate_selectable_when_oauth_token_expired() {
let mut row = sample_row();
let mut row = sample_row_without_model_mappings();
row.provider_id = "provider-kiro".to_string();
row.provider_name = "kiro".to_string();
row.provider_type = "kiro".to_string();
@@ -2346,7 +2349,7 @@ async fn keeps_refreshable_kiro_candidate_selectable_when_oauth_token_expired()
#[tokio::test]
async fn keeps_kiro_candidate_selectable_after_refresh_token_failure_until_access_token_expiry() {
let mut row = sample_row();
let mut row = sample_row_without_model_mappings();
row.provider_id = "provider-kiro".to_string();
row.provider_name = "kiro".to_string();
row.provider_type = "kiro".to_string();
@@ -2411,7 +2414,7 @@ async fn keeps_kiro_candidate_selectable_after_refresh_token_failure_until_acces
#[tokio::test]
async fn skips_kiro_candidate_after_refresh_token_failure_and_access_token_expiry() {
let mut row = sample_row();
let mut row = sample_row_without_model_mappings();
row.provider_id = "provider-kiro".to_string();
row.provider_name = "kiro".to_string();
row.provider_type = "kiro".to_string();
@@ -2477,7 +2480,7 @@ async fn skips_kiro_candidate_after_refresh_token_failure_and_access_token_expir
#[tokio::test]
async fn skips_refreshable_kiro_candidate_when_oauth_marker_is_account_block() {
let mut row = sample_row();
let mut row = sample_row_without_model_mappings();
row.provider_id = "provider-kiro".to_string();
row.provider_name = "kiro".to_string();
row.provider_type = "kiro".to_string();
@@ -2620,7 +2623,7 @@ async fn keeps_codex_candidate_selectable_when_exhausted_account_flag_is_disable
#[tokio::test]
async fn skips_kiro_candidate_when_account_quota_is_exhausted_and_pool_flag_enabled() {
let mut first = sample_row();
let mut first = sample_row_without_model_mappings();
first.provider_id = "provider-kiro".to_string();
first.provider_name = "kiro".to_string();
first.provider_type = "kiro".to_string();
@@ -2632,7 +2635,7 @@ async fn skips_kiro_candidate_when_account_quota_is_exhausted_and_pool_flag_enab
first.key_api_formats = Some(vec!["claude:messages".to_string()]);
first.key_global_priority_by_format = Some(serde_json::json!({"claude:messages": 1}));
let mut second = sample_row();
let mut second = sample_row_without_model_mappings();
second.provider_id = "provider-openai".to_string();
second.provider_name = "openai".to_string();
second.endpoint_id = "endpoint-openai".to_string();
@@ -56,6 +56,14 @@ pub(super) fn sample_row() -> StoredMinimalCandidateSelectionRow {
}
}
/// 共享夹具默认把上游名映射限定在 openai 格式;需要其它端点/格式的用例
/// 用本函数移除映射约束(无映射的行在所有端点按默认上游名可用)。
pub(super) fn sample_row_without_model_mappings() -> StoredMinimalCandidateSelectionRow {
let mut row = sample_row();
row.model_provider_model_mappings = None;
row
}
pub(super) fn sample_provider(
id: &str,
concurrent_limit: Option<i32>,
@@ -1050,17 +1050,6 @@ fn requested_model_selection_sql() -> String {
)
)
)
OR NOT EXISTS (
SELECT 1
FROM jsonb_array_elements(
CASE
WHEN jsonb_typeof(m.provider_model_mappings) = 'array'
THEN m.provider_model_mappings
ELSE '[]'::jsonb
END
) AS mapping(value)
WHERE mapping.value ->> 'name' = m.provider_model_name
)
)
)
OR (
@@ -1068,17 +1057,6 @@ fn requested_model_selection_sql() -> String {
AND (
m.provider_model_mappings IS NULL
OR jsonb_typeof(m.provider_model_mappings) <> 'array'
OR NOT EXISTS (
SELECT 1
FROM jsonb_array_elements(
CASE
WHEN jsonb_typeof(m.provider_model_mappings) = 'array'
THEN m.provider_model_mappings
ELSE '[]'::jsonb
END
) AS mapping(value)
WHERE mapping.value ->> 'name' = m.provider_model_name
)
OR EXISTS (
SELECT 1
FROM jsonb_array_elements(
@@ -1108,6 +1086,47 @@ fn requested_model_selection_sql() -> String {
)
)
)
OR (
NOT EXISTS (
SELECT 1
FROM jsonb_array_elements(
CASE
WHEN jsonb_typeof(m.provider_model_mappings) = 'array'
THEN m.provider_model_mappings
ELSE '[]'::jsonb
END
) AS mapping(value)
WHERE mapping.value ->> 'name' = m.provider_model_name
)
AND EXISTS (
SELECT 1
FROM jsonb_array_elements(
CASE
WHEN jsonb_typeof(m.provider_model_mappings) = 'array'
THEN m.provider_model_mappings
ELSE '[]'::jsonb
END
) AS mapping(value)
WHERE (
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 __AETHER_PROVIDER_MODEL_MAPPING_API_FORMAT_MATCH__
)
)
AND (
mapping.value -> 'endpoint_ids' IS NULL
OR jsonb_typeof(mapping.value -> 'endpoint_ids') <> 'array'
OR EXISTS (
SELECT 1
FROM jsonb_array_elements_text(mapping.value -> 'endpoint_ids') AS endpoint(value)
WHERE endpoint.value = pe.id
)
)
)
)
)
)
OR (
@@ -1706,7 +1725,7 @@ mod tests {
let sql = requested_model_selection_sql();
let compatibility = PROVIDER_MODEL_MAPPING_API_FORMAT_MATCH_SQL;
assert_eq!(sql.matches(compatibility).count(), 3);
assert_eq!(sql.matches(compatibility).count(), 4);
assert!(!sql.contains(PROVIDER_MODEL_MAPPING_API_FORMAT_MATCH_MARKER));
assert!(compatibility.contains("LOWER(BTRIM(p.provider_type)) = 'codex'"));
assert!(compatibility.contains("LOWER($4) = 'codex:live'"));
@@ -1730,6 +1749,21 @@ mod tests {
}
}
#[test]
fn requested_model_sql_confines_default_name_to_mapping_scope() {
let sql = requested_model_selection_sql();
// 默认名(provider_model_name)的可用性由映射的端点/格式范围决定:
// 分支 A(全局名匹配)不再单独放行“无默认名映射”的行;
// 分支 B(provider 名匹配)要求默认名映射命中,或映射范围内至少有一条映射命中。
assert_eq!(sql.matches("->> 'name' = m.provider_model_name").count(), 2);
assert_eq!(
sql.matches(PROVIDER_MODEL_MAPPING_API_FORMAT_MATCH_SQL)
.count(),
4
);
}
#[test]
fn candidate_selection_sql_allows_grok_oauth_chat_auth() {
let requested_model_sql = requested_model_selection_sql();
@@ -284,7 +284,12 @@ fn row_default_provider_model_name_available(
return true;
}
}
!has_explicit_default_mapping
if has_explicit_default_mapping {
return false;
}
// 与调度核心保持一致:映射一旦限定了端点或 API 格式范围,
// 默认模型名就只在范围内可用;范围之外整行不再匹配。
row_mapping_matches_scope(row, api_format)
}
fn row_mapping_matches_scope(row: &StoredMinimalCandidateSelectionRow, api_format: &str) -> bool {
@@ -691,6 +696,45 @@ mod tests {
assert_eq!(rows[0].endpoint_id, "endpoint-openai");
}
#[tokio::test]
async fn requested_model_filter_confines_rows_to_mapping_endpoint_scope() {
// 场景:映射只把上游名绑定到 openai:chat 端点;
// 其他端点(claude:messages)不应回退到默认名,整行不可匹配。
let mut selected = sample_row("provider-1", "openai:chat", "deepseek-flash", 10);
selected.endpoint_id = "endpoint-chat".to_string();
selected.model_provider_model_name = "deepseek-flash".to_string();
selected.model_provider_model_mappings = Some(vec![StoredProviderModelMapping {
name: "gpt-6-luna".to_string(),
priority: 1,
api_formats: None,
endpoint_ids: Some(vec!["endpoint-chat".to_string()]),
operations: None,
}]);
let mut scoped_out = selected.clone();
scoped_out.provider_id = "provider-2".to_string();
scoped_out.endpoint_id = "endpoint-claude".to_string();
scoped_out.endpoint_api_format = "claude:messages".to_string();
scoped_out.key_id = "key-provider-2".to_string();
scoped_out.key_api_formats = Some(vec!["claude:messages".to_string()]);
let repository =
InMemoryMinimalCandidateSelectionReadRepository::seed(vec![scoped_out, selected]);
let rows = repository
.list_for_exact_api_format_and_requested_model("claude:messages", "deepseek-flash")
.await
.expect("list should succeed");
assert!(rows.is_empty());
let rows = repository
.list_for_exact_api_format_and_requested_model("openai:chat", "deepseek-flash")
.await
.expect("list should succeed");
assert_eq!(rows.len(), 1);
assert_eq!(rows[0].endpoint_id, "endpoint-chat");
}
#[tokio::test]
async fn requested_model_page_returns_requested_slice_only() {
let mut rows = Vec::new();
+117 -15
View File
@@ -393,14 +393,20 @@ fn row_default_provider_model_name_available(
return true;
}
}
!has_explicit_default_mapping
if has_explicit_default_mapping {
return false;
}
// 映射一旦限定了端点或 API 格式范围,默认模型名就只在范围内可用;
// 范围之外整行不再参与候选,避免把默认名发到不认识它的上游端点。
mappings
.iter()
.any(|mapping| mapping_endpoint_and_api_format_scope_matches(mapping, row, api_format))
}
fn mapping_scope_matches(
fn mapping_endpoint_and_api_format_scope_matches(
mapping: &StoredProviderModelMapping,
row: &StoredMinimalCandidateSelectionRow,
api_format: &str,
request_operation: Option<&str>,
) -> bool {
let api_format_matches_scope = mapping.api_formats.as_ref().is_none_or(|api_formats| {
api_formats.iter().any(|value| {
@@ -411,24 +417,29 @@ fn mapping_scope_matches(
return false;
}
let endpoint_matches_scope = mapping.endpoint_ids.as_ref().is_none_or(|endpoint_ids| {
mapping.endpoint_ids.as_ref().is_none_or(|endpoint_ids| {
endpoint_ids
.iter()
.any(|endpoint_id| endpoint_id == &row.endpoint_id)
});
if !endpoint_matches_scope {
return false;
}
mapping.operations.as_ref().is_none_or(|operations| {
request_operation.is_some_and(|request_operation| {
operations
.iter()
.any(|operation| operation.eq_ignore_ascii_case(request_operation))
})
})
}
fn mapping_scope_matches(
mapping: &StoredProviderModelMapping,
row: &StoredMinimalCandidateSelectionRow,
api_format: &str,
request_operation: Option<&str>,
) -> bool {
mapping_endpoint_and_api_format_scope_matches(mapping, row, api_format)
&& mapping.operations.as_ref().is_none_or(|operations| {
request_operation.is_some_and(|request_operation| {
operations
.iter()
.any(|operation| operation.eq_ignore_ascii_case(request_operation))
})
})
}
fn mapping_operation_scope_rank(mapping: &StoredProviderModelMapping) -> u8 {
u8::from(mapping.operations.is_some())
}
@@ -1010,6 +1021,55 @@ mod tests {
);
}
#[test]
fn endpoint_scoped_mapping_confines_row_availability_to_scoped_endpoints() {
// 场景:某行只把上游名 gpt-6-luna 映射到 openai:chat 端点;
// 其他端点(如 claude:messages)不应回退到默认名,整行不可用。
let mut row = sample_row("deepseek-flash", "deepseek-flash");
row.endpoint_id = "endpoint-chat".to_string();
row.endpoint_api_format = "openai:chat".to_string();
row.model_provider_model_mappings = Some(vec![StoredProviderModelMapping {
name: "gpt-6-luna".to_string(),
priority: 1,
api_formats: None,
endpoint_ids: Some(vec!["endpoint-chat".to_string()]),
operations: None,
}]);
assert!(row_supports_requested_model(
&row,
"deepseek-flash",
"openai:chat"
));
assert_eq!(
resolve_provider_model_name(&row, "deepseek-flash", "openai:chat")
.map(|resolved| resolved.0),
Some("gpt-6-luna".to_string())
);
let mut claude_row = row.clone();
claude_row.endpoint_id = "endpoint-claude".to_string();
claude_row.endpoint_api_format = "claude:messages".to_string();
assert!(!row_supports_requested_model(
&claude_row,
"deepseek-flash",
"claude:messages"
));
assert!(
resolve_provider_model_name(&claude_row, "deepseek-flash", "claude:messages").is_none()
);
assert_eq!(
resolve_requested_global_model_name_with_model_directives(
&[claude_row],
"deepseek-flash",
"claude:messages",
false,
),
None
);
}
#[test]
fn reserved_global_model_name_keeps_its_own_rows_addressable() {
let row = cursor_alias_row();
@@ -1069,6 +1129,48 @@ mod tests {
);
}
#[test]
fn operation_scoped_mapping_keeps_default_name_available() {
// 只限定“适用请求”(如 compact)的映射不应把整行从其他请求里排除掉。
let mut row = sample_row("gpt-5.6-sol", "gpt-5.6-sol");
row.endpoint_api_format = "openai:responses".to_string();
row.model_provider_model_mappings = Some(vec![StoredProviderModelMapping {
name: "gpt-5.6-terra".to_string(),
priority: 1,
api_formats: Some(vec!["openai:responses".to_string()]),
endpoint_ids: None,
operations: Some(vec!["compact".to_string()]),
}]);
assert!(row_supports_requested_model(
&row,
"gpt-5.6-sol",
"openai:responses"
));
assert_eq!(
resolve_provider_model_name_with_model_directives_and_request_operation(
&row,
"gpt-5.6-sol",
"openai:responses",
false,
None,
)
.map(|resolved| resolved.0),
Some("gpt-5.6-sol".to_string())
);
assert_eq!(
resolve_provider_model_name_with_model_directives_and_request_operation(
&row,
"gpt-5.6-sol",
"openai:responses",
false,
Some("compact"),
)
.map(|resolved| resolved.0),
Some("gpt-5.6-terra".to_string())
);
}
fn sample_row(
global_model_name: &str,
model_provider_model_name: &str,