fix(model-resolution): 修复模型解析优先级和别名匹配范围

- 调整 ModelCacheService 解析优先级:优先直接匹配 GlobalModel.name,
  避免被 provider_model_name 误导到错误的 GlobalModel
- 为 check_model_allowed_with_aliases 新增 candidate_models 参数,
  限制别名匹配只能落到 Provider 实际配置的模型名上
- 对 allowed_set 排序以确保别名匹配结果的确定性
- 调整 HTTP_READ_TIMEOUT 默认值从 300s 改为 60s
- 优化 ModelMappingTab 组件布局,将正则表达式提取到 Key 级别展示
- 补充 .env.example 超时配置文档
This commit is contained in:
fawney19
2026-01-14 17:29:04 +08:00
parent 97d66a8d0c
commit 0ed053b02f
8 changed files with 294 additions and 67 deletions

View File

@@ -141,7 +141,7 @@ class Config:
# HTTP 请求超时配置(秒)
self.http_connect_timeout = float(os.getenv("HTTP_CONNECT_TIMEOUT", "10.0"))
self.http_read_timeout = float(os.getenv("HTTP_READ_TIMEOUT", "300.0"))
self.http_read_timeout = float(os.getenv("HTTP_READ_TIMEOUT", "60.0"))
self.http_write_timeout = float(os.getenv("HTTP_WRITE_TIMEOUT", "60.0"))
self.http_pool_timeout = float(os.getenv("HTTP_POOL_TIMEOUT", "10.0"))

View File

@@ -16,7 +16,7 @@
import re
from functools import lru_cache
from typing import Dict, List, Optional, Set, Tuple, Union
from typing import Dict, List, Optional, Tuple, Union
import regex
@@ -35,7 +35,7 @@ AllowedModels = Optional[Union[List[str], Dict[str, List[str]]]]
def normalize_allowed_models(
allowed_models: AllowedModels,
api_format: Optional[str] = None,
) -> Optional[Set[str]]:
) -> Optional[set[str]]:
"""
将 allowed_models 规范化为模型名称集合
@@ -45,7 +45,7 @@ def normalize_allowed_models(
Returns:
- None: 不限制(允许所有模型)
- Set[str]: 允许的模型名称集合(可能为空集,表示拒绝所有)
- set[str]: 允许的模型名称集合(可能为空集,表示拒绝所有)
"""
if allowed_models is None:
return None
@@ -58,7 +58,7 @@ def normalize_allowed_models(
if isinstance(allowed_models, dict):
if api_format is None:
# 没有指定格式,合并所有格式的模型
all_models: Set[str] = set()
all_models: set[str] = set()
for models in allowed_models.values():
if isinstance(models, list):
all_models.update(models)
@@ -153,7 +153,7 @@ def merge_allowed_models(
# 任一为字典模式:按 API 格式分别取交集,避免把 dict 合并成 list 导致权限过宽
from src.core.enums import APIFormat
def merge_sets(a: Optional[Set[str]], b: Optional[Set[str]]) -> Optional[Set[str]]:
def merge_sets(a: Optional[set[str]], b: Optional[set[str]]) -> Optional[set[str]]:
# None 表示不限制:交集规则下等价于“只受另一方限制”
if a is None:
return b
@@ -163,7 +163,7 @@ def merge_allowed_models(
known_formats = [fmt.value for fmt in APIFormat]
per_format: Dict[str, Optional[Set[str]]] = {}
per_format: Dict[str, Optional[set[str]]] = {}
for fmt in known_formats:
s1 = normalize_allowed_models(allowed_models_1, api_format=fmt)
s2 = normalize_allowed_models(allowed_models_2, api_format=fmt)
@@ -215,7 +215,7 @@ def get_allowed_models_preview(
if allowed_models is None:
return "(不限制)"
all_models: Set[str] = set()
all_models: set[str] = set()
if isinstance(allowed_models, list):
all_models = set(allowed_models)
@@ -295,7 +295,7 @@ def convert_to_simple_mode(allowed_models: AllowedModels) -> Optional[List[str]]
return allowed_models
if isinstance(allowed_models, dict):
all_models: Set[str] = set()
all_models: set[str] = set()
for models in allowed_models.values():
if isinstance(models, list):
all_models.update(models)
@@ -325,7 +325,7 @@ def parse_allowed_models_to_list(allowed_models: AllowedModels) -> List[str]:
return allowed_models
if isinstance(allowed_models, dict):
all_models: Set[str] = set()
all_models: set[str] = set()
for models in allowed_models.values():
if isinstance(models, list):
all_models.update(models)
@@ -533,6 +533,7 @@ def check_model_allowed_with_aliases(
api_format: Optional[str] = None,
resolved_model_name: Optional[str] = None,
model_aliases: Optional[List[str]] = None,
candidate_models: Optional[set[str]] = None,
) -> tuple[bool, Optional[str]]:
"""
检查模型是否被允许(支持别名通配符匹配)
@@ -554,6 +555,7 @@ def check_model_allowed_with_aliases(
api_format: 当前请求的 API 格式
resolved_model_name: 解析后的 GlobalModel.name
model_aliases: GlobalModel 的别名列表(来自 config.model_aliases
candidate_models: 可选的候选模型集合(用于限制别名匹配只能落到这些模型名上)
Returns:
(is_allowed, matched_model_name):
@@ -578,9 +580,16 @@ def check_model_allowed_with_aliases(
# 空集合 = 拒绝所有
return False, None
# 如果提供了候选集合,只允许在候选集合中进行别名匹配
if candidate_models is not None:
allowed_set = allowed_set & candidate_models
if len(allowed_set) == 0:
return False, None
# 遍历 allowed_models 中的每个模型名,检查是否有别名能匹配
# 注意:返回第一个匹配的模型名,匹配顺序由 allowed_set 迭代顺序和 model_aliases 数组顺序决定
for allowed_model in allowed_set:
# 注意:为了避免 set 迭代顺序带来的非确定性,这里对 allowed_set 做排序
# 返回第一个匹配的模型名,匹配顺序由 allowed_models 排序结果和 model_aliases 数组顺序共同决定
for allowed_model in sorted(allowed_set):
for alias_pattern in model_aliases:
if match_model_with_pattern(alias_pattern, allowed_model):
# 返回匹配到的模型名,用于实际请求

View File

@@ -746,9 +746,10 @@ class CacheAwareScheduler:
db: Session,
provider: Provider,
model_name: str,
api_format: Optional[str] = None,
is_stream: bool = False,
capability_requirements: Optional[Dict[str, bool]] = None,
) -> Tuple[bool, Optional[str], Optional[List[str]]]:
) -> Tuple[bool, Optional[str], Optional[List[str]], Optional[set[str]]]:
"""
检查 Provider 是否支持指定模型(可选检查流式支持和能力需求)
@@ -768,20 +769,24 @@ class CacheAwareScheduler:
capability_requirements: 能力需求(可选),用于检查模型是否支持所需能力
Returns:
(is_supported, skip_reason, supported_capabilities) - 是否支持、跳过原因、模型支持的能力列表
(is_supported, skip_reason, supported_capabilities, provider_model_names)
- is_supported: 是否支持
- skip_reason: 跳过原因
- supported_capabilities: 模型支持的能力列表
- provider_model_names: Provider 侧可用的模型名称集合(主名称 + 映射名称,按 api_format 过滤)
"""
# 使用 ModelCacheService 解析模型名称(支持映射名)
global_model = await ModelCacheService.resolve_global_model_by_name_or_alias(db, model_name)
if not global_model:
# 完全未找到匹配
return False, "模型不存在或 Provider 未配置此模型", None
return False, "模型不存在或 Provider 未配置此模型", None, None
# 找到 GlobalModel 后,检查当前 Provider 是否支持
is_supported, skip_reason, caps = await self._check_model_support_for_global_model(
db, provider, global_model, model_name, is_stream, capability_requirements
is_supported, skip_reason, caps, provider_model_names = await self._check_model_support_for_global_model(
db, provider, global_model, model_name, api_format, is_stream, capability_requirements
)
return is_supported, skip_reason, caps
return is_supported, skip_reason, caps, provider_model_names
async def _check_model_support_for_global_model(
self,
@@ -789,9 +794,10 @@ class CacheAwareScheduler:
provider: Provider,
global_model: "GlobalModel",
model_name: str,
api_format: Optional[str] = None,
is_stream: bool = False,
capability_requirements: Optional[Dict[str, bool]] = None,
) -> Tuple[bool, Optional[str], Optional[List[str]]]:
) -> Tuple[bool, Optional[str], Optional[List[str]], Optional[set[str]]]:
"""
检查 Provider 是否支持指定的 GlobalModel
@@ -804,7 +810,7 @@ class CacheAwareScheduler:
capability_requirements: 能力需求
Returns:
(is_supported, skip_reason, supported_capabilities)
(is_supported, skip_reason, supported_capabilities, provider_model_names)
"""
# 确保 global_model 附加到当前 Session
# 注意:从缓存重建的对象是 transient 状态,不能使用 load=False
@@ -833,7 +839,7 @@ class CacheAwareScheduler:
if is_stream:
supports_streaming = model.get_effective_supports_streaming()
if not supports_streaming:
return False, f"模型 {model_name} 在此 Provider 不支持流式", None
return False, f"模型 {model_name} 在此 Provider 不支持流式", None, None
# 检查模型是否支持所需的能力(在 Provider 级别检查,而不是 Key 级别)
# 只有当 model_supported_capabilities 非空时才进行检查
@@ -845,11 +851,32 @@ class CacheAwareScheduler:
False,
f"模型 {model_name} 不支持能力: {cap_name}",
list(model_supported_capabilities),
None,
)
return True, None, list(model_supported_capabilities)
provider_model_names: set[str] = {model.provider_model_name}
raw_mappings = model.provider_model_mappings
if isinstance(raw_mappings, list):
for raw in raw_mappings:
if not isinstance(raw, dict):
continue
name = raw.get("name")
if not isinstance(name, str) or not name.strip():
continue
return False, "Provider 未实现此模型", None
mapping_api_formats = raw.get("api_formats")
if api_format and mapping_api_formats:
if (
isinstance(mapping_api_formats, list)
and api_format not in mapping_api_formats
):
continue
provider_model_names.add(name.strip())
return True, None, list(model_supported_capabilities), provider_model_names
return False, "Provider 未实现此模型", None, None
def _check_key_availability(
self,
@@ -859,6 +886,7 @@ class CacheAwareScheduler:
capability_requirements: Optional[Dict[str, bool]] = None,
resolved_model_name: Optional[str] = None,
model_aliases: Optional[List[str]] = None,
candidate_models: Optional[set[str]] = None,
) -> Tuple[bool, Optional[str], Optional[str]]:
"""
检查 API Key 的可用性
@@ -872,6 +900,7 @@ class CacheAwareScheduler:
capability_requirements: 能力需求(可选)
resolved_model_name: 解析后的 GlobalModel.name可选
model_aliases: GlobalModel 的别名列表(用于通配符匹配)
candidate_models: Provider 侧可用的模型名称集合(用于限制别名匹配范围)
Returns:
(is_available, skip_reason, alias_matched_model)
@@ -901,6 +930,7 @@ class CacheAwareScheduler:
api_format=api_format,
resolved_model_name=resolved_model_name,
model_aliases=model_aliases,
candidate_models=candidate_models,
)
except TimeoutError:
# 正则匹配超时(可能是 ReDoS 攻击或复杂模式)
@@ -970,8 +1000,13 @@ class CacheAwareScheduler:
for provider in providers:
# 检查模型支持(同时检查流式支持和模型能力需求)
supports_model, skip_reason, _model_caps = await self._check_model_support(
db, provider, model_name, is_stream, capability_requirements
supports_model, skip_reason, _model_caps, provider_model_names = await self._check_model_support(
db,
provider,
model_name,
api_format=target_format_str,
is_stream=is_stream,
capability_requirements=capability_requirements,
)
if not supports_model:
logger.debug(f"Provider {provider.name} 不支持模型 {model_name}: {skip_reason}")
@@ -1022,6 +1057,7 @@ class CacheAwareScheduler:
capability_requirements,
resolved_model_name=resolved_model_name,
model_aliases=model_aliases,
candidate_models=provider_model_names,
)
candidate = ProviderCandidate(

View File

@@ -258,8 +258,8 @@ class ModelCacheService:
查找顺序:
1. 检查缓存
2. 通过 provider_model_name 匹配(查询 Model 表)
3. 直接匹配 GlobalModel.name(兜底
2. 直接匹配 GlobalModel.name
3. 通过 provider_model_name 匹配(查询 Model 表
注意:此方法不使用 provider_model_mappings 进行全局解析。
provider_model_mappings 是 Provider 级别的映射配置,只在特定 Provider 上下文中生效,
@@ -301,7 +301,25 @@ class ModelCacheService:
logger.debug(f"GlobalModel 缓存命中(映射解析): {normalized_name}")
return ModelCacheService._dict_to_global_model(cached_data)
# 2. 通过 provider_model_name 匹配(不考虑 provider_model_mappings
# 2. 直接通过 GlobalModel.name 匹配(优先级最高
# 说明:如果存在同名 GlobalModel应优先解析为 GlobalModel 本身,
# 避免被某个 Provider 的 provider_model_name 误导导致解析到错误的 GlobalModel。
global_model = (
db.query(GlobalModel)
.filter(GlobalModel.name == normalized_name, GlobalModel.is_active == True)
.first()
)
if global_model:
resolution_method = "direct_match"
global_model_dict = ModelCacheService._global_model_to_dict(global_model)
await CacheService.set(
cache_key, global_model_dict, ttl_seconds=ModelCacheService.CACHE_TTL
)
logger.debug(f"GlobalModel 已缓存(映射解析-直接匹配): {normalized_name}")
return global_model
# 3. 通过 provider_model_name 匹配(不考虑 provider_model_mappings
# 重要provider_model_mappings 是 Provider 级别的映射配置,只在特定 Provider 上下文中生效
# 全局解析不应该受到某个 Provider 映射配置的影响
# 例如Provider A 把 "haiku" 映射到 "sonnet",不应该影响 Provider B 的 "haiku" 解析
@@ -357,23 +375,6 @@ class ModelCacheService:
)
return result_global_model
# 3. 如果通过 provider 映射没找到,最后尝试直接通过 GlobalModel.name 查找
global_model = (
db.query(GlobalModel)
.filter(GlobalModel.name == normalized_name, GlobalModel.is_active == True)
.first()
)
if global_model:
resolution_method = "direct_match"
# 缓存结果
global_model_dict = ModelCacheService._global_model_to_dict(global_model)
await CacheService.set(
cache_key, global_model_dict, ttl_seconds=ModelCacheService.CACHE_TTL
)
logger.debug(f"GlobalModel 已缓存(映射解析-直接匹配): {normalized_name}")
return global_model
# 4. 完全未找到
resolution_method = "not_found"
# 未找到匹配,缓存负结果