fix(mapping): 对齐全局模型映射的正则匹配行为与范围 (#296)

- scheduler_core: `matches_model_mapping` 改为大小写不敏感且整串匹配,并补充单测
- global model routing 预览:
  - Key 过滤增加 `allowed_models + model_mappings` 校验
  - `all_keys_whitelist` 改为收集全站活跃 Provider 的活跃 Key 白名单
- provider mapping-preview:
  - 优先使用 admin 全量 GlobalModel(含非激活)参与映射
  - admin 数据为空时回退 public 模型,保持兼容
- public models 匹配逻辑统一复用 scheduler_core 实现,避免行为分叉
- 更新网关测试,覆盖未关联 Provider 的 Key 也进入 whitelist 的场景
This commit is contained in:
AAEE86
2026-04-14 21:58:52 +08:00
committed by GitHub
parent 47bf1d04a1
commit fb31928e44
5 changed files with 210 additions and 36 deletions
@@ -7,7 +7,9 @@ use aether_data_contracts::repository::global_models::{
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
};
use aether_scheduler_core::{is_provider_key_circuit_open, provider_key_health_score};
use aether_scheduler_core::{
is_provider_key_circuit_open, matches_model_mapping, provider_key_health_score,
};
use serde_json::json;
use std::collections::BTreeMap;
use uuid::Uuid;
@@ -121,6 +123,13 @@ pub(crate) async fn build_admin_global_model_routing_payload(
.unwrap_or_default()
.into_iter()
.filter(|key| provider_catalog_key_supports_format(key, &endpoint.api_format))
.filter(|key| {
key_allowed_models_match_global_model_for_routing(
key.allowed_models.as_ref(),
&global_model.name,
&global_model_mappings,
)
})
.collect::<Vec<_>>();
endpoint_keys.sort_by(|left, right| {
left.internal_priority
@@ -170,14 +179,6 @@ pub(crate) async fn build_admin_global_model_routing_payload(
"circuit_breaker_formats": circuit_breaker_formats,
"next_probe_at": next_probe_at,
});
all_keys_whitelist.push(json!({
"key_id": &key.id,
"key_name": &key.name,
"masked_key": state.masked_catalog_api_key(key),
"provider_id": &provider.id,
"provider_name": &provider.name,
"allowed_models": json_string_list(key.allowed_models.as_ref()),
}));
payload
})
.collect::<Vec<_>>();
@@ -218,6 +219,53 @@ pub(crate) async fn build_admin_global_model_routing_payload(
"active_endpoints": active_endpoints,
}));
}
// 与 Python 逻辑对齐:供前端实时匹配的白名单数据来自“全站活跃 Provider 的活跃 Key”
// (仅保留配置了非空 allowed_models 的 Key),而不是仅当前 GlobalModel 关联 Provider。
let active_providers = state
.list_provider_catalog_providers(true)
.await
.ok()
.unwrap_or_default();
let active_provider_ids = active_providers
.iter()
.map(|provider| provider.id.clone())
.collect::<Vec<_>>();
let active_provider_name_by_id = active_providers
.into_iter()
.map(|provider| (provider.id, provider.name))
.collect::<BTreeMap<_, _>>();
let active_keys = if active_provider_ids.is_empty() {
Vec::new()
} else {
state
.list_provider_catalog_keys_by_provider_ids(&active_provider_ids)
.await
.ok()
.unwrap_or_default()
};
for key in active_keys {
if !key.is_active {
continue;
}
let allowed_models = json_string_list(key.allowed_models.as_ref());
if allowed_models.is_empty() {
continue;
}
let provider_name = active_provider_name_by_id
.get(&key.provider_id)
.cloned()
.unwrap_or_default();
all_keys_whitelist.push(json!({
"key_id": key.id,
"key_name": key.name,
"masked_key": state.masked_catalog_api_key(&key),
"provider_id": key.provider_id,
"provider_name": provider_name,
"allowed_models": allowed_models,
}));
}
providers_payload.sort_by(|left, right| {
left.get("provider_priority")
.and_then(serde_json::Value::as_i64)
@@ -257,6 +305,35 @@ pub(crate) async fn build_admin_global_model_routing_payload(
}))
}
fn key_allowed_models_match_global_model_for_routing(
raw_allowed_models: Option<&serde_json::Value>,
global_model_name: &str,
global_model_mappings: &[String],
) -> bool {
// 兼容 Python 预览逻辑:None/[] 视为“不限制”,在链路预览中保留该 Key。
let allowed_models = json_string_list(raw_allowed_models);
if raw_allowed_models.is_none() || allowed_models.is_empty() {
return true;
}
if allowed_models
.iter()
.any(|value| value == global_model_name)
{
return true;
}
for allowed_model in &allowed_models {
for pattern in global_model_mappings {
if matches_model_mapping(pattern, allowed_model) {
return true;
}
}
}
false
}
pub(crate) async fn build_admin_assign_global_model_to_providers_payload(
state: &AdminAppState<'_>,
global_model_id: &str,