mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 17:30:23 +08:00
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:
62
src/services/cache/aware_scheduler.py
vendored
62
src/services/cache/aware_scheduler.py
vendored
@@ -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(
|
||||
|
||||
41
src/services/cache/model_cache.py
vendored
41
src/services/cache/model_cache.py
vendored
@@ -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"
|
||||
# 未找到匹配,缓存负结果
|
||||
|
||||
Reference in New Issue
Block a user