refactor: 将模型别名(alias)统一重命名为模型映射(mapping)

- 重命名 API: alias-mapping-preview -> mapping-preview
- 重命名类型: AliasMatchedModel -> MappingMatchedModel 等
- 重命名字段: model_aliases -> model_mappings, alias_matched_model -> mapping_matched_model
- 重命名前端组件: ModelAliasesTab -> ModelMappingsTab
- 重命名验证函数: validate_model_aliases -> validate_model_mappings
- 同步更新相关测试用例
- 补充 ModelService 中新增/批量创建模型时的缓存失效逻辑
This commit is contained in:
fawney19
2026-01-14 19:35:41 +08:00
parent 24ece00e93
commit bec9c3a989
18 changed files with 328 additions and 279 deletions

View File

@@ -322,12 +322,12 @@ class AdminCreateGlobalModelAdapter(AdminApiAdapter):
async def handle(self, context): # type: ignore[override]
from src.core.exceptions import InvalidRequestException
from src.core.model_permissions import validate_and_extract_model_aliases
from src.core.model_permissions import validate_and_extract_model_mappings
# 验证 model_aliases如果有
is_valid, error, _ = validate_and_extract_model_aliases(self.payload.config)
# 验证 model_mappings如果有
is_valid, error, _ = validate_and_extract_model_mappings(self.payload.config)
if not is_valid:
raise InvalidRequestException(f"别名规则验证失败: {error}", "model_aliases")
raise InvalidRequestException(f"映射规则验证失败: {error}", "model_mappings")
# 将 TieredPricingConfig 转换为 dict
tiered_pricing_dict = self.payload.default_tiered_pricing.model_dump()
@@ -361,12 +361,12 @@ class AdminUpdateGlobalModelAdapter(AdminApiAdapter):
async def handle(self, context): # type: ignore[override]
from src.core.exceptions import InvalidRequestException
from src.core.model_permissions import validate_and_extract_model_aliases
from src.core.model_permissions import validate_and_extract_model_mappings
# 验证 model_aliases如果有
is_valid, error, _ = validate_and_extract_model_aliases(self.payload.config)
# 验证 model_mappings如果有
is_valid, error, _ = validate_and_extract_model_mappings(self.payload.config)
if not is_valid:
raise InvalidRequestException(f"别名规则验证失败: {error}", "model_aliases")
raise InvalidRequestException(f"映射规则验证失败: {error}", "model_mappings")
# 使用行级锁获取旧的 GlobalModel 信息,防止并发更新导致的竞态条件
# 设置 2 秒锁超时,允许短暂等待而非立即失败,提升并发操作的成功率

View File

@@ -405,6 +405,7 @@ class AdminCreateProviderModelAdapter(AdminApiAdapter):
try:
model = ModelService.create_model(db, self.provider_id, self.model_data)
logger.info(f"Model created: {model.provider_model_name} for provider {provider.name} by {context.user.username}")
# 缓存失效已在 ModelService.create_model 中处理
return ModelService.convert_to_response(model)
except Exception as exc:
raise InvalidRequestException(str(exc))
@@ -447,6 +448,7 @@ class AdminUpdateProviderModelAdapter(AdminApiAdapter):
try:
updated_model = ModelService.update_model(db, self.model_id, self.model_data)
logger.info(f"Model updated: {updated_model.provider_model_name} by {context.user.username}")
# 缓存失效已在 ModelService.update_model 中处理
return ModelService.convert_to_response(updated_model)
except Exception as exc:
raise InvalidRequestException(str(exc))
@@ -471,6 +473,7 @@ class AdminDeleteProviderModelAdapter(AdminApiAdapter):
try:
ModelService.delete_model(db, self.model_id)
logger.info(f"Model deleted: {model_name} by {context.user.username}")
# 缓存失效已在 ModelService.delete_model 中处理
return {"message": f"Model '{model_name}' deleted successfully"}
except Exception as exc:
raise InvalidRequestException(str(exc))
@@ -490,6 +493,7 @@ class AdminBatchCreateModelsAdapter(AdminApiAdapter):
try:
models = ModelService.batch_create_models(db, self.provider_id, self.models_data)
logger.info(f"Batch created {len(models)} models for provider {provider.name} by {context.user.username}")
# 缓存失效已在 ModelService.batch_create_models 中处理
return [ModelService.convert_to_response(model) for model in models]
except Exception as exc:
raise InvalidRequestException(str(exc))
@@ -633,6 +637,11 @@ class AdminBatchAssignModelsToProviderAdapter(AdminApiAdapter):
# 清除 /v1/models 列表缓存
if success:
# Provider 新增模型实现后,清除同进程的 ModelMapper 缓存,避免 TTL 内仍返回 None
from src.services.cache.invalidation import get_cache_invalidation_service
cache_service = get_cache_invalidation_service()
cache_service.on_model_changed(self.provider_id, success[0].get("global_model_id", ""))
await invalidate_models_list_cache()
return BatchAssignModelsToProviderResponse(success=success, errors=errors)

View File

@@ -23,58 +23,58 @@ router = APIRouter(tags=["Provider CRUD"])
pipeline = ApiRequestPipeline()
# 别名映射预览配置(管理后台功能,限制宽松)
ALIAS_PREVIEW_MAX_KEYS = 200
ALIAS_PREVIEW_MAX_MODELS = 500
ALIAS_PREVIEW_TIMEOUT_SECONDS = 10.0
# 映射预览配置(管理后台功能,限制宽松)
MAPPING_PREVIEW_MAX_KEYS = 200
MAPPING_PREVIEW_MAX_MODELS = 500
MAPPING_PREVIEW_TIMEOUT_SECONDS = 10.0
# ========== Response Models ==========
class AliasMatchedModel(BaseModel):
class MappingMatchedModel(BaseModel):
"""匹配到的模型名称"""
allowed_model: str = Field(..., description="Key 白名单中匹配到的模型名")
alias_pattern: str = Field(..., description="匹配的别名规则")
mapping_pattern: str = Field(..., description="匹配的映射规则")
class AliasMatchingGlobalModel(BaseModel):
"""别名匹配的 GlobalModel"""
class MappingMatchingGlobalModel(BaseModel):
"""映射匹配的 GlobalModel"""
global_model_id: str
global_model_name: str
display_name: str
is_active: bool
matched_models: List[AliasMatchedModel] = Field(
matched_models: List[MappingMatchedModel] = Field(
default_factory=list, description="匹配到的模型列表"
)
model_config = ConfigDict(from_attributes=True)
class AliasMatchingKey(BaseModel):
"""别名匹配的 Key"""
class MappingMatchingKey(BaseModel):
"""映射匹配的 Key"""
key_id: str
key_name: str
masked_key: str
is_active: bool
allowed_models: List[str] = Field(default_factory=list, description="Key 的模型白名单")
matching_global_models: List[AliasMatchingGlobalModel] = Field(
matching_global_models: List[MappingMatchingGlobalModel] = Field(
default_factory=list, description="匹配到的 GlobalModel 列表"
)
model_config = ConfigDict(from_attributes=True)
class ProviderAliasMappingPreviewResponse(BaseModel):
"""Provider 别名映射预览响应"""
class ProviderMappingPreviewResponse(BaseModel):
"""Provider 映射预览响应"""
provider_id: str
provider_name: str
keys: List[AliasMatchingKey] = Field(
default_factory=list, description="有白名单配置且匹配到别名的 Key 列表"
keys: List[MappingMatchingKey] = Field(
default_factory=list, description="有白名单配置且匹配到映射的 Key 列表"
)
total_keys: int = Field(0, description="有匹配结果的 Key 数量")
total_matches: int = Field(
@@ -417,18 +417,18 @@ class AdminDeleteProviderAdapter(AdminApiAdapter):
@router.get(
"/{provider_id}/alias-mapping-preview",
response_model=ProviderAliasMappingPreviewResponse,
"/{provider_id}/mapping-preview",
response_model=ProviderMappingPreviewResponse,
)
async def get_provider_alias_mapping_preview(
async def get_provider_mapping_preview(
request: Request,
provider_id: str,
db: Session = Depends(get_db),
) -> ProviderAliasMappingPreviewResponse:
) -> ProviderMappingPreviewResponse:
"""
获取 Provider 别名映射预览
获取 Provider 映射预览
查看该 Provider 的 Key 白名单能够被哪些 GlobalModel 的别名规则匹配。
查看该 Provider 的 Key 白名单能够被哪些 GlobalModel 的映射规则匹配。
**路径参数**:
- `provider_id`: Provider ID
@@ -445,26 +445,26 @@ async def get_provider_alias_mapping_preview(
- `total_keys`: 有白名单配置的 Key 总数
- `total_matches`: 匹配到的 GlobalModel 总数
"""
adapter = AdminGetProviderAliasMappingPreviewAdapter(provider_id=provider_id)
adapter = AdminGetProviderMappingPreviewAdapter(provider_id=provider_id)
# 添加超时保护,防止复杂匹配导致的 DoS
try:
return await asyncio.wait_for(
pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode),
timeout=ALIAS_PREVIEW_TIMEOUT_SECONDS,
timeout=MAPPING_PREVIEW_TIMEOUT_SECONDS,
)
except asyncio.TimeoutError:
logger.warning(f"别名映射预览超时: provider_id={provider_id}")
raise InvalidRequestException("别名映射预览超时,请简化配置或稍后重试")
logger.warning(f"映射预览超时: provider_id={provider_id}")
raise InvalidRequestException("映射预览超时,请简化配置或稍后重试")
class AdminGetProviderAliasMappingPreviewAdapter(AdminApiAdapter):
"""获取 Provider 别名映射预览"""
class AdminGetProviderMappingPreviewAdapter(AdminApiAdapter):
"""获取 Provider 映射预览"""
def __init__(self, provider_id: str):
self.provider_id = provider_id
async def handle(self, context) -> ProviderAliasMappingPreviewResponse: # type: ignore[override]
async def handle(self, context) -> ProviderMappingPreviewResponse: # type: ignore[override]
db = context.db
# 获取 Provider
@@ -502,27 +502,27 @@ class AdminGetProviderAliasMappingPreviewAdapter(AdminApiAdapter):
ProviderAPIKey.provider_id == self.provider_id,
ProviderAPIKey.allowed_models.isnot(None),
)
.limit(ALIAS_PREVIEW_MAX_KEYS)
.limit(MAPPING_PREVIEW_MAX_KEYS)
.all()
)
# 计算被截断的 Key 数量
if total_keys_with_allowed_models > ALIAS_PREVIEW_MAX_KEYS:
truncated_keys = total_keys_with_allowed_models - ALIAS_PREVIEW_MAX_KEYS
if total_keys_with_allowed_models > MAPPING_PREVIEW_MAX_KEYS:
truncated_keys = total_keys_with_allowed_models - MAPPING_PREVIEW_MAX_KEYS
# 获取有 model_aliases 配置的 GlobalModel 总数(用于截断统计)
total_models_with_aliases = (
# 获取有 model_mappings 配置的 GlobalModel 总数(用于截断统计)
total_models_with_mappings = (
db.query(func.count(GlobalModel.id))
.filter(
GlobalModel.config.isnot(None),
GlobalModel.config["model_aliases"].isnot(None),
func.jsonb_array_length(GlobalModel.config["model_aliases"]) > 0,
GlobalModel.config["model_mappings"].isnot(None),
func.jsonb_array_length(GlobalModel.config["model_mappings"]) > 0,
)
.scalar()
or 0
)
# 只查询有 model_aliases 配置的 GlobalModel使用 SQLAlchemy JSONB 操作符)
# 只查询有 model_mappings 配置的 GlobalModel使用 SQLAlchemy JSONB 操作符)
global_models = (
db.query(
GlobalModel.id,
@@ -533,28 +533,28 @@ class AdminGetProviderAliasMappingPreviewAdapter(AdminApiAdapter):
)
.filter(
GlobalModel.config.isnot(None),
GlobalModel.config["model_aliases"].isnot(None),
func.jsonb_array_length(GlobalModel.config["model_aliases"]) > 0,
GlobalModel.config["model_mappings"].isnot(None),
func.jsonb_array_length(GlobalModel.config["model_mappings"]) > 0,
)
.limit(ALIAS_PREVIEW_MAX_MODELS)
.limit(MAPPING_PREVIEW_MAX_MODELS)
.all()
)
# 计算被截断的 GlobalModel 数量
if total_models_with_aliases > ALIAS_PREVIEW_MAX_MODELS:
truncated_models = total_models_with_aliases - ALIAS_PREVIEW_MAX_MODELS
if total_models_with_mappings > MAPPING_PREVIEW_MAX_MODELS:
truncated_models = total_models_with_mappings - MAPPING_PREVIEW_MAX_MODELS
# 构建有别名配置的 GlobalModel 映射
models_with_aliases: Dict[str, tuple] = {} # id -> (model_info, aliases)
# 构建有映射配置的 GlobalModel 映射
models_with_mappings: Dict[str, tuple] = {} # id -> (model_info, mappings)
for gm in global_models:
config = gm.config or {}
aliases = config.get("model_aliases", [])
if aliases:
models_with_aliases[gm.id] = (gm, aliases)
mappings = config.get("model_mappings", [])
if mappings:
models_with_mappings[gm.id] = (gm, mappings)
# 如果没有任何带别名的 GlobalModel直接返回空结果
if not models_with_aliases:
return ProviderAliasMappingPreviewResponse(
# 如果没有任何带映射的 GlobalModel直接返回空结果
if not models_with_mappings:
return ProviderMappingPreviewResponse(
provider_id=provider.id,
provider_name=provider.name,
keys=[],
@@ -565,7 +565,7 @@ class AdminGetProviderAliasMappingPreviewAdapter(AdminApiAdapter):
truncated_models=0,
)
key_infos: List[AliasMatchingKey] = []
key_infos: List[MappingMatchingKey] = []
total_matches = 0
# 创建 CryptoService 实例
@@ -591,25 +591,25 @@ class AdminGetProviderAliasMappingPreviewAdapter(AdminApiAdapter):
pass
# 查找匹配的 GlobalModel
matching_global_models: List[AliasMatchingGlobalModel] = []
matching_global_models: List[MappingMatchingGlobalModel] = []
for gm_id, (gm, aliases) in models_with_aliases.items():
matched_models: List[AliasMatchedModel] = []
for gm_id, (gm, mappings) in models_with_mappings.items():
matched_models: List[MappingMatchedModel] = []
for allowed_model in allowed_models_list:
for alias_pattern in aliases:
if match_model_with_pattern(alias_pattern, allowed_model):
for mapping_pattern in mappings:
if match_model_with_pattern(mapping_pattern, allowed_model):
matched_models.append(
AliasMatchedModel(
MappingMatchedModel(
allowed_model=allowed_model,
alias_pattern=alias_pattern,
mapping_pattern=mapping_pattern,
)
)
break # 一个 allowed_model 只需匹配一个别名
break # 一个 allowed_model 只需匹配一个映射
if matched_models:
matching_global_models.append(
AliasMatchingGlobalModel(
MappingMatchingGlobalModel(
global_model_id=gm.id,
global_model_name=gm.name,
display_name=gm.display_name,
@@ -621,7 +621,7 @@ class AdminGetProviderAliasMappingPreviewAdapter(AdminApiAdapter):
if matching_global_models:
key_infos.append(
AliasMatchingKey(
MappingMatchingKey(
key_id=key.id or "",
key_name=key.name or "",
masked_key=masked_key,
@@ -633,7 +633,7 @@ class AdminGetProviderAliasMappingPreviewAdapter(AdminApiAdapter):
is_truncated = truncated_keys > 0 or truncated_models > 0
return ProviderAliasMappingPreviewResponse(
return ProviderMappingPreviewResponse(
provider_id=provider.id,
provider_name=provider.name,
keys=key_infos,

View File

@@ -429,8 +429,8 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
provider_api_format=str(endpoint.api_format) if endpoint.api_format else None,
)
# 获取模型映射(优先使用别名匹配到的模型,其次是 Provider 级别的映射)
mapped_model = candidate.alias_matched_model if candidate else None
# 获取模型映射(优先使用映射匹配到的模型,其次是 Provider 级别的映射)
mapped_model = candidate.mapping_matched_model if candidate else None
if not mapped_model:
mapped_model = await self._get_mapped_model(
source_model=ctx.model,
@@ -663,8 +663,8 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
provider_name = str(provider.name)
# 获取模型映射(优先使用别名匹配到的模型,其次是 Provider 级别的映射)
mapped_model = candidate.alias_matched_model if candidate else None
# 获取模型映射(优先使用映射匹配到的模型,其次是 Provider 级别的映射)
mapped_model = candidate.mapping_matched_model if candidate else None
if not mapped_model:
mapped_model = await self._get_mapped_model(
source_model=model,

View File

@@ -437,8 +437,8 @@ class CliMessageHandlerBase(BaseMessageHandler):
ctx.provider_api_format = str(endpoint.api_format) if endpoint.api_format else ""
ctx.client_api_format = ctx.api_format # 已在 process_stream 中设置
# 获取模型映射(优先使用别名匹配到的模型,其次是 Provider 级别的映射)
mapped_model = candidate.alias_matched_model if candidate else None
# 获取模型映射(优先使用映射匹配到的模型,其次是 Provider 级别的映射)
mapped_model = candidate.mapping_matched_model if candidate else None
if not mapped_model:
mapped_model = await self._get_mapped_model(
source_model=ctx.model,
@@ -1570,8 +1570,8 @@ class CliMessageHandlerBase(BaseMessageHandler):
provider_name = str(provider.name)
provider_api_format = str(endpoint.api_format) if endpoint.api_format else ""
# 获取模型映射(优先使用别名匹配到的模型,其次是 Provider 级别的映射)
mapped_model = candidate.alias_matched_model if candidate else None
# 获取模型映射(优先使用映射匹配到的模型,其次是 Provider 级别的映射)
mapped_model = candidate.mapping_matched_model if candidate else None
if not mapped_model:
mapped_model = await self._get_mapped_model(
source_model=model,

View File

@@ -4,9 +4,9 @@
allowed_models 格式: ["claude-sonnet-4", "gpt-4o"]
使用 None/null 表示不限制(允许所有模型)
支持模型别名匹配:
- GlobalModel.config.model_aliases 定义别名模式
- 别名模式支持正则表达式语法
支持模型映射匹配:
- GlobalModel.config.model_mappings 定义映射模式
- 映射模式支持正则表达式语法
- 例如claude-haiku-.* 可匹配 claude-haiku-4.5, claude-haiku-last
- 使用 regex 库的原生超时保护100ms防止 ReDoS
"""
@@ -19,10 +19,10 @@ import regex
from src.core.logger import logger
# 别名规则限制
MAX_ALIASES_PER_MODEL = 50
MAX_ALIAS_LENGTH = 200
MAX_MODEL_NAME_LENGTH = 200 # 与 MAX_ALIAS_LENGTH 保持一致
# 映射规则限制
MAX_MAPPINGS_PER_MODEL = 50
MAX_MAPPING_LENGTH = 200
MAX_MODEL_NAME_LENGTH = 200 # 与 MAX_MAPPING_LENGTH 保持一致
REGEX_MATCH_TIMEOUT_MS = 100 # 正则匹配超时(毫秒)
# 类型别名
@@ -154,9 +154,9 @@ def parse_allowed_models_to_list(allowed_models: AllowedModels) -> List[str]:
return list(allowed_models)
def validate_alias_pattern(pattern: str) -> Tuple[bool, Optional[str]]:
def validate_mapping_pattern(pattern: str) -> Tuple[bool, Optional[str]]:
"""
验证别名模式是否安全
验证映射模式是否安全
Args:
pattern: 待验证的正则模式
@@ -165,10 +165,10 @@ def validate_alias_pattern(pattern: str) -> Tuple[bool, Optional[str]]:
(is_valid, error_message)
"""
if not pattern or not pattern.strip():
return False, "别名规则不能为空"
return False, "映射规则不能为空"
if len(pattern) > MAX_ALIAS_LENGTH:
return False, f"别名规则过长 (最大 {MAX_ALIAS_LENGTH} 字符)"
if len(pattern) > MAX_MAPPING_LENGTH:
return False, f"映射规则过长 (最大 {MAX_MAPPING_LENGTH} 字符)"
# 尝试编译验证语法
try:
@@ -179,35 +179,35 @@ def validate_alias_pattern(pattern: str) -> Tuple[bool, Optional[str]]:
return True, None
def validate_model_aliases(aliases: Optional[List[str]]) -> Tuple[bool, Optional[str]]:
def validate_model_mappings(mappings: Optional[List[str]]) -> Tuple[bool, Optional[str]]:
"""
验证别名列表是否合法
验证映射列表是否合法
Args:
aliases: 别名列表
mappings: 映射列表
Returns:
(is_valid, error_message)
"""
if not aliases:
if not mappings:
return True, None
if len(aliases) > MAX_ALIASES_PER_MODEL:
return False, f"别名规则数量超限 (最大 {MAX_ALIASES_PER_MODEL} 条)"
if len(mappings) > MAX_MAPPINGS_PER_MODEL:
return False, f"映射规则数量超限 (最大 {MAX_MAPPINGS_PER_MODEL} 条)"
for i, alias in enumerate(aliases):
is_valid, error = validate_alias_pattern(alias)
for i, mapping in enumerate(mappings):
is_valid, error = validate_mapping_pattern(mapping)
if not is_valid:
return False, f"{i + 1} 条规则无效: {error}"
return True, None
def validate_and_extract_model_aliases(
def validate_and_extract_model_mappings(
config: Optional[dict],
) -> Tuple[bool, Optional[str], Optional[List[str]]]:
"""
从 config 中验证并提取 model_aliases
从 config 中验证并提取 model_mappings
用于 GlobalModel 创建/更新时的统一验证
@@ -215,34 +215,34 @@ def validate_and_extract_model_aliases(
config: GlobalModel 的 config 字典
Returns:
(is_valid, error_message, aliases):
(is_valid, error_message, mappings):
- is_valid: 验证是否通过
- error_message: 错误信息(验证失败时)
- aliases: 提取的别名列表(验证成功时)
- mappings: 提取的映射列表(验证成功时)
"""
if not config or "model_aliases" not in config:
if not config or "model_mappings" not in config:
return True, None, None
aliases = config.get("model_aliases")
mappings = config.get("model_mappings")
# 允许显式设置为 None表示清除别名
if aliases is None:
# 允许显式设置为 None表示清除映射
if mappings is None:
return True, None, None
# 类型验证:必须是列表
if not isinstance(aliases, list):
return False, "model_aliases 必须是数组类型", None
if not isinstance(mappings, list):
return False, "model_mappings 必须是数组类型", None
# 元素类型验证:必须是字符串
if not all(isinstance(a, str) for a in aliases):
return False, "model_aliases 数组元素必须是字符串", None
if not all(isinstance(m, str) for m in mappings):
return False, "model_mappings 数组元素必须是字符串", None
# 业务规则验证
is_valid, error = validate_model_aliases(aliases)
is_valid, error = validate_model_mappings(mappings)
if not is_valid:
return False, error, None
return True, None, aliases
return True, None, mappings
@lru_cache(maxsize=2000)
@@ -267,7 +267,7 @@ def clear_regex_cache() -> None:
"""
清空正则缓存
在 GlobalModel 别名更新时调用此函数以确保缓存一致性
在 GlobalModel 映射更新时调用此函数以确保缓存一致性
"""
_compile_pattern_cached.cache_clear()
logger.debug("[RegexCache] 缓存已清空")
@@ -310,7 +310,7 @@ def _match_with_timeout(
def match_model_with_pattern(pattern: str, model_name: str) -> bool:
"""
检查模型名是否匹配别名模式(支持正则表达式)
检查模型名是否匹配映射模式(支持正则表达式)
安全特性:
- 长度限制检查
@@ -318,7 +318,7 @@ def match_model_with_pattern(pattern: str, model_name: str) -> bool:
- 正则匹配超时保护100ms使用 regex 库原生超时)
Args:
pattern: 别名模式,支持正则表达式语法
pattern: 映射模式,支持正则表达式语法
model_name: 被检查的模型名(来自 Key 的 allowed_models
Returns:
@@ -334,7 +334,7 @@ def match_model_with_pattern(pattern: str, model_name: str) -> bool:
return True
# 长度检查
if len(pattern) > MAX_ALIAS_LENGTH or len(model_name) > MAX_MODEL_NAME_LENGTH:
if len(pattern) > MAX_MAPPING_LENGTH or len(model_name) > MAX_MODEL_NAME_LENGTH:
return False
# 使用缓存的编译结果
@@ -347,45 +347,45 @@ def match_model_with_pattern(pattern: str, model_name: str) -> bool:
return result is True
def check_model_allowed_with_aliases(
def check_model_allowed_with_mappings(
model_name: str,
allowed_models: AllowedModels,
resolved_model_name: Optional[str] = None,
model_aliases: Optional[List[str]] = None,
model_mappings: Optional[List[str]] = None,
candidate_models: Optional[set[str]] = None,
) -> tuple[bool, Optional[str]]:
"""
检查模型是否被允许(支持别名通配符匹配)
检查模型是否被允许(支持映射通配符匹配)
匹配优先级:
1. 精确匹配 model_name用户请求的模型名
2. 精确匹配 resolved_model_nameGlobalModel.name
3. 遍历 model_aliases检查每个别名是否匹配 allowed_models 中的任一项
3. 遍历 model_mappings检查每个映射是否匹配 allowed_models 中的任一项
别名匹配顺序说明:
映射匹配顺序说明:
- 按 allowed_models 集合的迭代顺序遍历(通常为字母顺序,因为内部使用 set
- 对于每个 allowed_model按 model_aliases 数组顺序依次尝试匹配
- 对于每个 allowed_model按 model_mappings 数组顺序依次尝试匹配
- 返回第一个成功匹配的 allowed_model
- 如需确定性行为,请确保 model_aliases 中的规则从最具体到最通用排序
- 如需确定性行为,请确保 model_mappings 中的规则从最具体到最通用排序
Args:
model_name: 请求的模型名称
allowed_models: 允许的模型配置(来自 Provider Key
resolved_model_name: 解析后的 GlobalModel.name
model_aliases: GlobalModel 的别名列表(来自 config.model_aliases
candidate_models: 可选的候选模型集合(用于限制别名匹配只能落到这些模型名上)
model_mappings: GlobalModel 的映射列表(来自 config.model_mappings
candidate_models: 可选的候选模型集合(用于限制映射匹配只能落到这些模型名上)
Returns:
(is_allowed, matched_model_name):
- is_allowed: 是否允许使用该模型
- matched_model_name: 通过别名匹配到的模型名(仅别名匹配时有值,精确匹配时为 None
- matched_model_name: 通过映射匹配到的模型名(仅映射匹配时有值,精确匹配时为 None
"""
# 先尝试精确匹配(使用原有逻辑)
if check_model_allowed(model_name, allowed_models, resolved_model_name):
return True, None
# 如果精确匹配失败且有别名配置,尝试别名匹配
if not model_aliases:
# 如果精确匹配失败且有映射配置,尝试映射匹配
if not model_mappings:
return False, None
# 获取 allowed_models 的集合
@@ -398,18 +398,18 @@ 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_models 中的每个模型名,检查是否有映射能匹配
# 注意:为了避免 set 迭代顺序带来的非确定性,这里对 allowed_set 做排序
# 返回第一个匹配的模型名,匹配顺序由 allowed_models 排序结果和 model_aliases 数组顺序共同决定
# 返回第一个匹配的模型名,匹配顺序由 allowed_models 排序结果和 model_mappings 数组顺序共同决定
for allowed_model in sorted(allowed_set):
for alias_pattern in model_aliases:
if match_model_with_pattern(alias_pattern, allowed_model):
for mapping_pattern in model_mappings:
if match_model_with_pattern(mapping_pattern, allowed_model):
# 返回匹配到的模型名,用于实际请求
return True, allowed_model

View File

@@ -77,7 +77,7 @@ class ProviderCandidate:
is_cached: bool = False
is_skipped: bool = False # 是否被跳过
skip_reason: Optional[str] = None # 跳过原因
alias_matched_model: Optional[str] = None # 通过别名匹配到的模型名(用于实际请求)
mapping_matched_model: Optional[str] = None # 通过映射匹配到的模型名(用于实际请求)
@dataclass
@@ -580,7 +580,7 @@ class CacheAwareScheduler:
target_format = normalize_api_format(api_format)
# 0. 解析 model_name 到 GlobalModel支持直接匹配和映射名匹配使用 ModelCacheService
global_model = await ModelCacheService.resolve_global_model_by_name_or_alias(db, model_name)
global_model = await ModelCacheService.resolve_global_model_by_name_or_mapping(db, model_name)
if not global_model:
logger.warning(f"GlobalModel not found: {model_name}")
@@ -591,8 +591,8 @@ class CacheAwareScheduler:
requested_model_name = model_name
resolved_model_name = str(global_model.name)
# 提取模型别名(用于 Provider Key 的 allowed_models 匹配)
model_aliases: List[str] = (global_model.config or {}).get("model_aliases", [])
# 提取模型映射(用于 Provider Key 的 allowed_models 匹配)
model_mappings: List[str] = (global_model.config or {}).get("model_mappings", [])
# 获取合并后的访问限制ApiKey + User
restrictions = self._get_effective_restrictions(user_api_key)
@@ -660,7 +660,7 @@ class CacheAwareScheduler:
target_format=target_format,
model_name=requested_model_name,
resolved_model_name=resolved_model_name,
model_aliases=model_aliases,
model_mappings=model_mappings,
affinity_key=affinity_key,
max_candidates=max_candidates,
is_stream=is_stream,
@@ -774,7 +774,7 @@ class CacheAwareScheduler:
- provider_model_names: Provider 侧可用的模型名称集合(主名称 + 映射名称,按 api_format 过滤)
"""
# 使用 ModelCacheService 解析模型名称(支持映射名)
global_model = await ModelCacheService.resolve_global_model_by_name_or_alias(db, model_name)
global_model = await ModelCacheService.resolve_global_model_by_name_or_mapping(db, model_name)
if not global_model:
# 完全未找到匹配
@@ -883,7 +883,7 @@ class CacheAwareScheduler:
model_name: str,
capability_requirements: Optional[Dict[str, bool]] = None,
resolved_model_name: Optional[str] = None,
model_aliases: Optional[List[str]] = None,
model_mappings: Optional[List[str]] = None,
candidate_models: Optional[set[str]] = None,
) -> Tuple[bool, Optional[str], Optional[str]]:
"""
@@ -897,14 +897,14 @@ class CacheAwareScheduler:
model_name: 模型名称
capability_requirements: 能力需求(可选)
resolved_model_name: 解析后的 GlobalModel.name可选
model_aliases: GlobalModel 的别名列表(用于通配符匹配)
candidate_models: Provider 侧可用的模型名称集合(用于限制别名匹配范围)
model_mappings: GlobalModel 的映射列表(用于通配符匹配)
candidate_models: Provider 侧可用的模型名称集合(用于限制映射匹配范围)
Returns:
(is_available, skip_reason, alias_matched_model)
(is_available, skip_reason, mapping_matched_model)
- is_available: Key 是否可用
- skip_reason: 不可用时的原因
- alias_matched_model: 通过别名匹配到的模型名(用于实际请求)
- mapping_matched_model: 通过映射匹配到的模型名(用于实际请求)
"""
# 检查熔断器状态(使用详细状态方法获取更丰富的跳过原因,按 API 格式)
is_available, circuit_reason = health_monitor.get_circuit_breaker_status(
@@ -915,33 +915,33 @@ class CacheAwareScheduler:
# 模型权限检查:使用 allowed_models 白名单
# None = 允许所有模型,[] = 拒绝所有模型,["a","b"] = 只允许指定模型
# 支持通配符别名匹配(通过 model_aliases
# 支持通配符映射匹配(通过 model_mappings
from src.core.model_permissions import (
check_model_allowed_with_aliases,
check_model_allowed_with_mappings,
get_allowed_models_preview,
)
try:
is_allowed, alias_matched_model = check_model_allowed_with_aliases(
is_allowed, mapping_matched_model = check_model_allowed_with_mappings(
model_name=model_name,
allowed_models=key.allowed_models,
resolved_model_name=resolved_model_name,
model_aliases=model_aliases,
model_mappings=model_mappings,
candidate_models=candidate_models,
)
except TimeoutError:
# 正则匹配超时(可能是 ReDoS 攻击或复杂模式)
logger.warning(f"别名匹配超时: key_id={key.id}, model={model_name}")
return False, "别名匹配超时,请简化配置", None
logger.warning(f"映射匹配超时: key_id={key.id}, model={model_name}")
return False, "映射匹配超时,请简化配置", None
except re.error as e:
# 正则语法错误(配置问题)
logger.warning(f"别名规则无效: key_id={key.id}, model={model_name}, error={e}")
return False, f"别名规则无效: {str(e)}", None
logger.warning(f"映射规则无效: key_id={key.id}, model={model_name}, error={e}")
return False, f"映射规则无效: {str(e)}", None
except Exception as e:
# 其他未知异常
logger.error(f"别名匹配异常: key_id={key.id}, model={model_name}, error={e}", exc_info=True)
logger.error(f"映射匹配异常: key_id={key.id}, model={model_name}, error={e}", exc_info=True)
# 异常时保守处理:不允许使用该 Key
return False, "别名匹配失败", None
return False, "映射匹配失败", None
if not is_allowed:
return False, f"模型权限不匹配(允许: {get_allowed_models_preview(key.allowed_models)})", None
@@ -957,7 +957,7 @@ class CacheAwareScheduler:
if not is_match:
return False, skip_reason, None
return True, None, alias_matched_model
return True, None, mapping_matched_model
async def _build_candidates(
self,
@@ -967,7 +967,7 @@ class CacheAwareScheduler:
model_name: str,
affinity_key: Optional[str],
resolved_model_name: Optional[str] = None,
model_aliases: Optional[List[str]] = None,
model_mappings: Optional[List[str]] = None,
max_candidates: Optional[int] = None,
is_stream: bool = False,
capability_requirements: Optional[Dict[str, bool]] = None,
@@ -984,7 +984,7 @@ class CacheAwareScheduler:
model_name: 模型名称(用户请求的名称,可能是映射名)
affinity_key: 亲和性标识符通常为API Key ID
resolved_model_name: 解析后的 GlobalModel.name用于 Key.allowed_models 校验)
model_aliases: GlobalModel 的别名列表(用于 Key.allowed_models 通配符匹配)
model_mappings: GlobalModel 的映射列表(用于 Key.allowed_models 通配符匹配)
max_candidates: 最大候选数
is_stream: 是否是流式请求,如果为 True 则过滤不支持流式的 Provider
capability_requirements: 能力需求(可选)
@@ -1047,13 +1047,13 @@ class CacheAwareScheduler:
for key in keys:
# Key 级别的能力检查
is_available, skip_reason, alias_matched_model = self._check_key_availability(
is_available, skip_reason, mapping_matched_model = self._check_key_availability(
key,
target_format_str,
model_name,
capability_requirements,
resolved_model_name=resolved_model_name,
model_aliases=model_aliases,
model_mappings=model_mappings,
candidate_models=provider_model_names,
)
@@ -1063,7 +1063,7 @@ class CacheAwareScheduler:
key=key,
is_skipped=not is_available,
skip_reason=skip_reason,
alias_matched_model=alias_matched_model,
mapping_matched_model=mapping_matched_model,
)
candidates.append(candidate)

View File

@@ -15,7 +15,7 @@ Model 映射缓存服务 - 减少模型查询
使用示例
--------
global_model = await ModelCacheService.resolve_global_model_by_name_or_alias(db, "gpt-4")
global_model = await ModelCacheService.resolve_global_model_by_name_or_mapping(db, "gpt-4")
"""
import time
@@ -250,7 +250,7 @@ class ModelCacheService:
logger.debug(f"GlobalModel 缓存已清除: {global_model_id}")
@staticmethod
async def resolve_global_model_by_name_or_alias(
async def resolve_global_model_by_name_or_mapping(
db: Session, model_name: str
) -> Optional[GlobalModel]:
"""

View File

@@ -102,7 +102,7 @@ class ModelMapperMiddleware:
mapping = None
# 步骤 1: 解析 GlobalModel支持映射名
global_model = await ModelCacheService.resolve_global_model_by_name_or_alias(
global_model = await ModelCacheService.resolve_global_model_by_name_or_mapping(
self.db, source_model
)

View File

@@ -77,6 +77,24 @@ class ModelService:
logger.info(f"创建模型成功: provider={provider.name}, model={model.provider_model_name}, global_model_id={model.global_model_id}")
# 清除 Redis 缓存(异步执行,不阻塞返回)
# 重要:新增模型可能需要清除 resolver 的 NOT_FOUND 负缓存global_model:resolve:*
# 否则请求链路在 TTL 内可能无法立刻解析到新模型。
asyncio.create_task(
ModelCacheService.invalidate_model_cache(
model_id=model.id,
provider_id=model.provider_id,
global_model_id=model.global_model_id,
provider_model_name=model.provider_model_name,
provider_model_mappings=model.provider_model_mappings,
)
)
# 清除内存缓存ModelMapperMiddleware 实例)
if model.provider_id and model.global_model_id:
cache_service = get_cache_invalidation_service()
cache_service.on_model_changed(model.provider_id, model.global_model_id)
# 清除 /v1/models 列表缓存
asyncio.create_task(invalidate_models_list_cache())
@@ -154,6 +172,7 @@ class ModelService:
raise NotFoundException(f"模型 {model_id} 不存在")
# 保存旧的映射,用于清除缓存
old_global_model_id = model.global_model_id
old_provider_model_name = model.provider_model_name
old_provider_model_mappings = model.provider_model_mappings
@@ -179,14 +198,17 @@ class ModelService:
ModelCacheService.invalidate_model_cache(
model_id=model.id,
provider_id=model.provider_id,
global_model_id=model.global_model_id,
global_model_id=old_global_model_id,
provider_model_name=old_provider_model_name,
provider_model_mappings=old_provider_model_mappings,
)
)
# 再清除新的映射缓存(如果有变化)
if (model.provider_model_name != old_provider_model_name or
model.provider_model_mappings != old_provider_model_mappings):
# 再清除新的映射缓存(如果有变化,包括 global_model_id 变更
if (
model.provider_model_name != old_provider_model_name
or model.provider_model_mappings != old_provider_model_mappings
or model.global_model_id != old_global_model_id
):
asyncio.create_task(
ModelCacheService.invalidate_model_cache(
model_id=model.id,
@@ -354,6 +376,7 @@ class ModelService:
provider_id=provider_id,
global_model_id=model_data.global_model_id,
provider_model_name=model_data.provider_model_name,
provider_model_mappings=model_data.provider_model_mappings,
price_per_request=model_data.price_per_request,
tiered_pricing=model_data.tiered_pricing,
supports_vision=model_data.supports_vision,
@@ -373,6 +396,23 @@ class ModelService:
db.refresh(model)
logger.info(f"批量创建 {len(created_models)} 个模型成功")
# 清除 Redis 缓存(异步执行,不阻塞返回)
# 逐个清除 resolver 的映射缓存,避免 NOT_FOUND 负缓存阻塞新模型生效。
for model in created_models:
asyncio.create_task(
ModelCacheService.invalidate_model_cache(
model_id=model.id,
provider_id=model.provider_id,
global_model_id=model.global_model_id,
provider_model_name=model.provider_model_name,
provider_model_mappings=model.provider_model_mappings,
)
)
# 清除内存缓存ModelMapperMiddleware 实例)
cache_service = get_cache_invalidation_service()
cache_service.on_model_changed(provider_id, created_models[0].global_model_id)
# 清除 /v1/models 列表缓存
asyncio.create_task(invalidate_models_list_cache())
except IntegrityError as e: