mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 09:20:22 +08:00
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:
@@ -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 秒锁超时,允许短暂等待而非立即失败,提升并发操作的成功率
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user