feat: 实现 GlobalModel 别名匹配系统

主要更改:
- GlobalModel 支持 model_aliases 配置,允许使用正则表达式定义别名规则
- Provider Key 的 allowed_models 现在可以通过别名规则匹配 GlobalModel
- 新增 ModelAliasesTab 组件用于管理模型别名配置
- Provider 详情页新增别名映射预览功能,展示 Key 白名单与 GlobalModel 别名的匹配关系
- 路由预览 API 返回 Key 的 allowed_models 信息

安全特性:
- 使用 regex 库的原生超时保护(100ms)防止 ReDoS 攻击
- 别名规则数量限制(50 条/模型)和长度限制(200 字符)
- 别名映射预览 API 添加超时保护和结果截断

其他改进:
- GlobalModel 更新/删除时使用行级锁防止并发竞态
- 缓存失效逻辑优化,支持异步清理和正则缓存清空
- 路由 Tab 布局重构,使用 flexbox 替代绝对定位
This commit is contained in:
fawney19
2026-01-13 16:04:15 +08:00
parent 9fea71a70c
commit 85decd7487
21 changed files with 3845 additions and 2308 deletions

View File

@@ -321,6 +321,14 @@ class AdminCreateGlobalModelAdapter(AdminApiAdapter):
payload: GlobalModelCreate
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
# 验证 model_aliases如果有
is_valid, error, _ = validate_and_extract_model_aliases(self.payload.config)
if not is_valid:
raise InvalidRequestException(f"别名规则验证失败: {error}", "model_aliases")
# 将 TieredPricingConfig 转换为 dict
tiered_pricing_dict = self.payload.default_tiered_pricing.model_dump()
@@ -352,6 +360,40 @@ class AdminUpdateGlobalModelAdapter(AdminApiAdapter):
payload: GlobalModelUpdate
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
# 验证 model_aliases如果有
is_valid, error, _ = validate_and_extract_model_aliases(self.payload.config)
if not is_valid:
raise InvalidRequestException(f"别名规则验证失败: {error}", "model_aliases")
# 使用行级锁获取旧的 GlobalModel 信息,防止并发更新导致的竞态条件
# 设置 2 秒锁超时,允许短暂等待而非立即失败,提升并发操作的成功率
from sqlalchemy import text
from sqlalchemy.exc import OperationalError
from src.models.database import GlobalModel
try:
# 设置会话级别的锁超时(仅影响当前事务)
context.db.execute(text("SET LOCAL lock_timeout = '2s'"))
old_global_model = (
context.db.query(GlobalModel)
.filter(GlobalModel.id == self.global_model_id)
.with_for_update()
.first()
)
except OperationalError as e:
# 锁超时或锁冲突时返回友好的错误提示
error_msg = str(e).lower()
if "lock" in error_msg or "timeout" in error_msg:
raise InvalidRequestException("该模型正在被其他操作更新,请稍后重试")
raise
old_model_name = old_global_model.name if old_global_model else None
new_model_name = self.payload.name if self.payload.name else old_model_name
# 执行更新(此时仍持有行锁)
global_model = GlobalModelService.update_global_model(
db=context.db,
global_model_id=self.global_model_id,
@@ -360,11 +402,18 @@ class AdminUpdateGlobalModelAdapter(AdminApiAdapter):
logger.info(f"GlobalModel 已更新: id={global_model.id} name={global_model.name}")
# 失效相关缓存
# 更新成功后才失效缓存(避免回滚时缓存已被清除的竞态问题)
# 注意:此时事务已提交(由 pipeline 管理),数据已持久化
from src.services.cache.invalidation import get_cache_invalidation_service
cache_service = get_cache_invalidation_service()
cache_service.on_global_model_changed(global_model.name)
# 同步清理新旧两个名称的缓存(防止名称变更时的竞态)
if old_model_name:
cache_service.on_global_model_changed(old_model_name, self.global_model_id)
if new_model_name and new_model_name != old_model_name:
cache_service.on_global_model_changed(new_model_name, self.global_model_id)
# 异步失效更多缓存
await cache_service.on_global_model_changed_async(global_model.name, global_model.id)
return GlobalModelResponse.model_validate(global_model)
@@ -376,24 +425,44 @@ class AdminDeleteGlobalModelAdapter(AdminApiAdapter):
global_model_id: str
async def handle(self, context): # type: ignore[override]
# 获取 GlobalModel 信息(用于失效缓存)
# 使用行级锁获取 GlobalModel 信息,防止并发操作导致的竞态条件
# 设置 2 秒锁超时,允许短暂等待而非立即失败
from sqlalchemy import text
from sqlalchemy.exc import OperationalError
from src.core.exceptions import InvalidRequestException
from src.models.database import GlobalModel
global_model = (
context.db.query(GlobalModel).filter(GlobalModel.id == self.global_model_id).first()
)
try:
# 设置会话级别的锁超时(仅影响当前事务)
context.db.execute(text("SET LOCAL lock_timeout = '2s'"))
global_model = (
context.db.query(GlobalModel)
.filter(GlobalModel.id == self.global_model_id)
.with_for_update()
.first()
)
except OperationalError as e:
# 锁超时或锁冲突时返回友好的错误提示
error_msg = str(e).lower()
if "lock" in error_msg or "timeout" in error_msg:
raise InvalidRequestException("该模型正在被其他操作处理,请稍后重试")
raise
model_name = global_model.name if global_model else None
model_id = global_model.id if global_model else self.global_model_id
# 执行删除(此时仍持有行锁)
GlobalModelService.delete_global_model(context.db, self.global_model_id)
logger.info(f"GlobalModel 已删除: id={self.global_model_id}")
# 失效相关缓存
if model_name:
from src.services.cache.invalidation import get_cache_invalidation_service
# 删除成功后才失效缓存(避免回滚时缓存已被清除的竞态问题)
from src.services.cache.invalidation import get_cache_invalidation_service
cache_service = get_cache_invalidation_service()
cache_service.on_global_model_changed(model_name)
cache_service = get_cache_invalidation_service()
if model_name:
cache_service.on_global_model_changed(model_name, model_id)
await cache_service.on_global_model_changed_async(model_name, model_id)
return None
@@ -413,7 +482,9 @@ class AdminBatchAssignToProvidersAdapter(AdminApiAdapter):
create_models=self.payload.create_models,
)
logger.info(f"批量为 Provider 添加 GlobalModel: global_model_id={self.global_model_id} success={len(result['success'])} errors={len(result['errors'])}")
logger.info(
f"批量为 Provider 添加 GlobalModel: global_model_id={self.global_model_id} success={len(result['success'])} errors={len(result['errors'])}"
)
return BatchAssignToProvidersResponse(**result)

View File

@@ -17,6 +17,7 @@ from sqlalchemy.orm import Session, selectinload
from src.api.base.admin_adapter import AdminApiAdapter
from src.api.base.pipeline import ApiRequestPipeline
from src.core.model_permissions import parse_allowed_models_to_list
from src.database import get_db
from src.models.database import (
GlobalModel,
@@ -49,6 +50,8 @@ class RoutingKeyInfo(BaseModel):
health_score: float = Field(100.0, description="健康度分数")
is_active: bool
api_formats: List[str] = Field(default_factory=list, description="支持的 API 格式")
# 模型白名单
allowed_models: Optional[List[str]] = Field(None, description="允许的模型列表null 表示不限制")
# 熔断状态
circuit_breaker_open: bool = Field(False, description="熔断器是否打开")
circuit_breaker_formats: List[str] = Field(default_factory=list, description="熔断的 API 格式列表")
@@ -299,6 +302,21 @@ class AdminGetModelRoutingPreviewAdapter(AdminApiAdapter):
circuit_breaker_open = True
circuit_breaker_formats.append(fmt)
# 解析 allowed_models
# 语义说明:
# - None: 不限制(允许所有模型)
# - {}: 空字典 = 不限制normalize_allowed_models 返回 None
# - []: 空列表 = 拒绝所有模型
# - {"CLAUDE": []}: 指定格式空列表 = 该格式拒绝所有
raw_allowed_models = key.allowed_models
if raw_allowed_models is None:
allowed_models_list = None
elif isinstance(raw_allowed_models, dict) and not raw_allowed_models:
# 空 dict {} 在语义上等价于不限制
allowed_models_list = None
else:
allowed_models_list = parse_allowed_models_to_list(raw_allowed_models)
key_infos.append(
RoutingKeyInfo(
id=key.id or "",
@@ -313,6 +331,7 @@ class AdminGetModelRoutingPreviewAdapter(AdminApiAdapter):
health_score=health_score,
is_active=bool(key.is_active),
api_formats=key.api_formats or [],
allowed_models=allowed_models_list,
circuit_breaker_open=circuit_breaker_open,
circuit_breaker_formats=circuit_breaker_formats,
)

View File

@@ -1,10 +1,11 @@
"""管理员 Provider 管理路由。"""
import asyncio
from datetime import datetime, timezone
from typing import Optional
from typing import Dict, List, Optional
from fastapi import APIRouter, Depends, Query, Request
from pydantic import ValidationError
from pydantic import BaseModel, ConfigDict, Field, ValidationError
from sqlalchemy.orm import Session
from src.api.base.admin_adapter import AdminApiAdapter
@@ -12,15 +13,81 @@ from src.api.base.pipeline import ApiRequestPipeline
from src.core.enums import ProviderBillingType
from src.core.exceptions import InvalidRequestException, NotFoundException
from src.core.logger import logger
from src.core.model_permissions import match_model_with_pattern, parse_allowed_models_to_list
from src.database import get_db
from src.models.admin_requests import CreateProviderRequest, UpdateProviderRequest
from src.models.database import Provider
from src.models.database import GlobalModel, Provider, ProviderAPIKey
from src.services.cache.provider_cache import ProviderCacheService
router = APIRouter(tags=["Provider CRUD"])
pipeline = ApiRequestPipeline()
# 别名映射预览配置(管理后台功能,限制宽松)
ALIAS_PREVIEW_MAX_KEYS = 200
ALIAS_PREVIEW_MAX_MODELS = 500
ALIAS_PREVIEW_TIMEOUT_SECONDS = 10.0
# ========== Response Models ==========
class AliasMatchedModel(BaseModel):
"""匹配到的模型名称"""
allowed_model: str = Field(..., description="Key 白名单中匹配到的模型名")
alias_pattern: str = Field(..., description="匹配的别名规则")
class AliasMatchingGlobalModel(BaseModel):
"""有别名匹配的 GlobalModel"""
global_model_id: str
global_model_name: str
display_name: str
is_active: bool
matched_models: List[AliasMatchedModel] = Field(
default_factory=list, description="匹配到的模型列表"
)
model_config = ConfigDict(from_attributes=True)
class AliasMatchingKey(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(
default_factory=list, description="匹配到的 GlobalModel 列表"
)
model_config = ConfigDict(from_attributes=True)
class ProviderAliasMappingPreviewResponse(BaseModel):
"""Provider 别名映射预览响应"""
provider_id: str
provider_name: str
keys: List[AliasMatchingKey] = Field(
default_factory=list, description="有白名单配置且匹配到别名的 Key 列表"
)
total_keys: int = Field(0, description="有匹配结果的 Key 数量")
total_matches: int = Field(
0, description="匹配到的 GlobalModel 数量(同一 GlobalModel 被多个 Key 匹配会重复计数)"
)
# 截断提示字段
truncated: bool = Field(False, description="是否因限制而截断结果")
truncated_keys: int = Field(0, description="被截断的 Key 数量")
truncated_models: int = Field(0, description="被截断的 GlobalModel 数量")
model_config = ConfigDict(from_attributes=True)
@router.get("/")
async def list_providers(
request: Request,
@@ -292,7 +359,9 @@ class AdminUpdateProviderAdapter(AdminApiAdapter):
setattr(provider, field, ProviderBillingType(value))
elif field == "proxy" and value is not None:
# proxy 需要转换为 dict如果是 Pydantic 模型)
setattr(provider, field, value if isinstance(value, dict) else value.model_dump())
setattr(
provider, field, value if isinstance(value, dict) else value.model_dump()
)
else:
setattr(provider, field, value)
@@ -345,3 +414,232 @@ class AdminDeleteProviderAdapter(AdminApiAdapter):
db.delete(provider)
db.commit()
return {"message": "提供商已删除"}
@router.get(
"/{provider_id}/alias-mapping-preview",
response_model=ProviderAliasMappingPreviewResponse,
)
async def get_provider_alias_mapping_preview(
request: Request,
provider_id: str,
db: Session = Depends(get_db),
) -> ProviderAliasMappingPreviewResponse:
"""
获取 Provider 别名映射预览
查看该 Provider 的 Key 白名单能够被哪些 GlobalModel 的别名规则匹配。
**路径参数**:
- `provider_id`: Provider ID
**返回字段**:
- `provider_id`: Provider ID
- `provider_name`: Provider 名称
- `keys`: 有白名单配置的 Key 列表,每个包含:
- `key_id`: Key ID
- `key_name`: Key 名称
- `masked_key`: 脱敏的 Key
- `allowed_models`: Key 的白名单模型列表
- `matching_global_models`: 匹配到的 GlobalModel 列表
- `total_keys`: 有白名单配置的 Key 总数
- `total_matches`: 匹配到的 GlobalModel 总数
"""
adapter = AdminGetProviderAliasMappingPreviewAdapter(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,
)
except asyncio.TimeoutError:
logger.warning(f"别名映射预览超时: provider_id={provider_id}")
raise InvalidRequestException("别名映射预览超时,请简化配置或稍后重试")
class AdminGetProviderAliasMappingPreviewAdapter(AdminApiAdapter):
"""获取 Provider 别名映射预览"""
def __init__(self, provider_id: str):
self.provider_id = provider_id
async def handle(self, context) -> ProviderAliasMappingPreviewResponse: # type: ignore[override]
db = context.db
# 获取 Provider
provider = db.query(Provider).filter(Provider.id == self.provider_id).first()
if not provider:
raise NotFoundException("提供商不存在", "provider")
# 统计截断情况
truncated_keys = 0
truncated_models = 0
# 获取该 Provider 有白名单配置的 Key 总数(用于截断统计)
from sqlalchemy import func
total_keys_with_allowed_models = (
db.query(func.count(ProviderAPIKey.id))
.filter(
ProviderAPIKey.provider_id == self.provider_id,
ProviderAPIKey.allowed_models.isnot(None),
)
.scalar()
or 0
)
# 获取该 Provider 有白名单配置的 Key只查询需要的字段
keys = (
db.query(
ProviderAPIKey.id,
ProviderAPIKey.name,
ProviderAPIKey.api_key,
ProviderAPIKey.is_active,
ProviderAPIKey.allowed_models,
)
.filter(
ProviderAPIKey.provider_id == self.provider_id,
ProviderAPIKey.allowed_models.isnot(None),
)
.limit(ALIAS_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
# 获取有 model_aliases 配置的 GlobalModel 总数(用于截断统计)
total_models_with_aliases = (
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,
)
.scalar()
or 0
)
# 只查询有 model_aliases 配置的 GlobalModel使用 SQLAlchemy JSONB 操作符)
global_models = (
db.query(
GlobalModel.id,
GlobalModel.name,
GlobalModel.display_name,
GlobalModel.is_active,
GlobalModel.config,
)
.filter(
GlobalModel.config.isnot(None),
GlobalModel.config["model_aliases"].isnot(None),
func.jsonb_array_length(GlobalModel.config["model_aliases"]) > 0,
)
.limit(ALIAS_PREVIEW_MAX_MODELS)
.all()
)
# 计算被截断的 GlobalModel 数量
if total_models_with_aliases > ALIAS_PREVIEW_MAX_MODELS:
truncated_models = total_models_with_aliases - ALIAS_PREVIEW_MAX_MODELS
# 构建有别名配置的 GlobalModel 映射
models_with_aliases: Dict[str, tuple] = {} # id -> (model_info, aliases)
for gm in global_models:
config = gm.config or {}
aliases = config.get("model_aliases", [])
if aliases:
models_with_aliases[gm.id] = (gm, aliases)
# 如果没有任何带别名的 GlobalModel直接返回空结果
if not models_with_aliases:
return ProviderAliasMappingPreviewResponse(
provider_id=provider.id,
provider_name=provider.name,
keys=[],
total_keys=0,
total_matches=0,
truncated=False,
truncated_keys=0,
truncated_models=0,
)
key_infos: List[AliasMatchingKey] = []
total_matches = 0
# 创建 CryptoService 实例
from src.core.crypto import CryptoService
crypto = CryptoService()
for key in keys:
allowed_models_list = parse_allowed_models_to_list(key.allowed_models)
if not allowed_models_list:
continue
# 生成脱敏 Key
masked_key = "***"
if key.api_key:
try:
decrypted_key = crypto.decrypt(key.api_key, silent=True)
if len(decrypted_key) > 8:
masked_key = f"{decrypted_key[:4]}***{decrypted_key[-4:]}"
else:
masked_key = f"{decrypted_key[:2]}***"
except Exception:
pass
# 查找匹配的 GlobalModel
matching_global_models: List[AliasMatchingGlobalModel] = []
for gm_id, (gm, aliases) in models_with_aliases.items():
matched_models: List[AliasMatchedModel] = []
for allowed_model in allowed_models_list:
for alias_pattern in aliases:
if match_model_with_pattern(alias_pattern, allowed_model):
matched_models.append(
AliasMatchedModel(
allowed_model=allowed_model,
alias_pattern=alias_pattern,
)
)
break # 一个 allowed_model 只需匹配一个别名
if matched_models:
matching_global_models.append(
AliasMatchingGlobalModel(
global_model_id=gm.id,
global_model_name=gm.name,
display_name=gm.display_name,
is_active=bool(gm.is_active),
matched_models=matched_models,
)
)
total_matches += 1
if matching_global_models:
key_infos.append(
AliasMatchingKey(
key_id=key.id or "",
key_name=key.name or "",
masked_key=masked_key,
is_active=bool(key.is_active),
allowed_models=allowed_models_list,
matching_global_models=matching_global_models,
)
)
is_truncated = truncated_keys > 0 or truncated_models > 0
return ProviderAliasMappingPreviewResponse(
provider_id=provider.id,
provider_name=provider.name,
keys=key_infos,
total_keys=len(key_infos),
total_matches=total_matches,
truncated=is_truncated,
truncated_keys=truncated_keys,
truncated_models=truncated_models,
)

View File

@@ -53,6 +53,7 @@ from src.models.database import (
ProviderEndpoint,
User,
)
from src.services.cache.aware_scheduler import ProviderCandidate
from src.services.provider.transport import build_provider_url
@@ -312,6 +313,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
provider: Provider,
endpoint: ProviderEndpoint,
key: ProviderAPIKey,
candidate: ProviderCandidate,
) -> AsyncGenerator[bytes, None]:
return await self._execute_stream_request(
ctx,
@@ -322,6 +324,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
original_request_body,
original_headers,
query_params,
candidate,
)
try:
@@ -411,6 +414,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
original_request_body: Dict[str, Any],
original_headers: Dict[str, str],
query_params: Optional[Dict[str, str]] = None,
candidate: Optional[ProviderCandidate] = None,
) -> AsyncGenerator[bytes, None]:
"""执行流式请求并返回流生成器"""
# 重置上下文状态(重试时清除之前的数据)
@@ -425,11 +429,13 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
provider_api_format=str(endpoint.api_format) if endpoint.api_format else None,
)
# 获取模型映射
mapped_model = await self._get_mapped_model(
source_model=ctx.model,
provider_id=str(provider.id),
)
# 获取模型映射(优先使用别名匹配到的模型,其次是 Provider 级别的映射)
mapped_model = candidate.alias_matched_model if candidate else None
if not mapped_model:
mapped_model = await self._get_mapped_model(
source_model=ctx.model,
provider_id=str(provider.id),
)
# 应用模型映射到请求体
if mapped_model:
@@ -650,17 +656,20 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
provider: Provider,
endpoint: ProviderEndpoint,
key: ProviderAPIKey,
candidate: ProviderCandidate,
) -> Dict[str, Any]:
nonlocal provider_name, response_json, status_code, response_headers
nonlocal provider_request_headers, provider_request_body, mapped_model_result
provider_name = str(provider.name)
# 获取模型映射
mapped_model = await self._get_mapped_model(
source_model=model,
provider_id=str(provider.id),
)
# 获取模型映射(优先使用别名匹配到的模型,其次是 Provider 级别的映射)
mapped_model = candidate.alias_matched_model if candidate else None
if not mapped_model:
mapped_model = await self._get_mapped_model(
source_model=model,
provider_id=str(provider.id),
)
# 应用模型映射
if mapped_model:

View File

@@ -64,6 +64,7 @@ from src.models.database import (
ProviderEndpoint,
User,
)
from src.services.cache.aware_scheduler import ProviderCandidate
from src.services.provider.transport import build_provider_url
from src.utils.sse_parser import SSEEventParser
from src.utils.timeout import read_first_chunk_with_ttfb_timeout
@@ -317,6 +318,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
provider: Provider,
endpoint: ProviderEndpoint,
key: ProviderAPIKey,
candidate: ProviderCandidate,
) -> AsyncGenerator[bytes, None]:
return await self._execute_stream_request(
ctx,
@@ -326,6 +328,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
original_request_body,
original_headers,
query_params,
candidate,
)
try:
@@ -405,6 +408,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
original_request_body: Dict[str, Any],
original_headers: Dict[str, str],
query_params: Optional[Dict[str, str]] = None,
candidate: Optional[ProviderCandidate] = None,
) -> AsyncGenerator[bytes, None]:
"""执行流式请求并返回流生成器"""
# 重置上下文状态(重试时清除之前的数据,避免累积)
@@ -432,11 +436,13 @@ 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 中设置
# 获取模型映射(映射名称 → 实际模型名
mapped_model = await self._get_mapped_model(
source_model=ctx.model,
provider_id=str(provider.id),
)
# 获取模型映射(优先使用别名匹配到的模型,其次是 Provider 级别的映射)
mapped_model = candidate.alias_matched_model if candidate else None
if not mapped_model:
mapped_model = await self._get_mapped_model(
source_model=ctx.model,
provider_id=str(provider.id),
)
# 应用模型映射到请求体(子类可覆盖此方法处理不同格式)
if mapped_model:
@@ -1247,14 +1253,29 @@ class CliMessageHandlerBase(BaseMessageHandler):
stream_generator: AsyncGenerator[bytes, None],
) -> AsyncGenerator[bytes, None]:
"""创建带监控的流生成器"""
import time as time_module
last_chunk_time = time_module.time()
chunk_count = 0
try:
async for chunk in stream_generator:
last_chunk_time = time_module.time()
chunk_count += 1
yield chunk
except asyncio.CancelledError:
# 计算距离上次收到 chunk 的时间
time_since_last_chunk = time_module.time() - last_chunk_time
# 如果响应已完成,不标记为失败
if not ctx.has_completion:
ctx.status_code = 499
ctx.error_message = "Client disconnected"
logger.warning(
f"ID:{ctx.request_id} | Stream cancelled: "
f"chunks={chunk_count}, "
f"has_completion={ctx.has_completion}, "
f"time_since_last_chunk={time_since_last_chunk:.2f}s, "
f"output_tokens={ctx.output_tokens}"
)
raise
except httpx.TimeoutException as e:
ctx.status_code = 504
@@ -1536,16 +1557,19 @@ class CliMessageHandlerBase(BaseMessageHandler):
provider: Provider,
endpoint: ProviderEndpoint,
key: ProviderAPIKey,
candidate: ProviderCandidate,
) -> Dict[str, Any]:
nonlocal provider_name, response_json, status_code, response_headers, provider_api_format, provider_request_headers, provider_request_body, mapped_model_result, response_metadata_result
provider_name = str(provider.name)
provider_api_format = str(endpoint.api_format) if endpoint.api_format else ""
# 获取模型映射(映射名称 → 实际模型名
mapped_model = await self._get_mapped_model(
source_model=model,
provider_id=str(provider.id),
)
# 获取模型映射(优先使用别名匹配到的模型,其次是 Provider 级别的映射)
mapped_model = candidate.alias_matched_model if candidate else None
if not mapped_model:
mapped_model = await self._get_mapped_model(
source_model=model,
provider_id=str(provider.id),
)
# 应用模型映射到请求体(子类可覆盖此方法处理不同格式)
if mapped_model:

View File

@@ -268,7 +268,7 @@ async def test_connection(
orchestrator = FallbackOrchestrator(db, redis_client)
# 定义请求函数
async def test_request_func(_prov, endpoint, key):
async def test_request_func(_prov, endpoint, key, _candidate):
request_builder = PassthroughRequestBuilder()
provider_payload, provider_headers = request_builder.build(
payload, {}, endpoint, key, is_stream=False