refactor: 将 global_priority 和 rate_multiplier 改为按 API 格式配置

- 新增 global_priority_by_format 字段,支持按 API 格式设置全局优先级
- 移除已废弃的 rate_multiplier 字段,统一使用 rate_multipliers
- 移除已废弃的 timeout 字段(providers 和 provider_endpoints 表)
- 更新调度器以支持按格式的优先级排序
- 同步更新前后端类型定义和 API 接口
- 添加数据库迁移脚本,自动迁移现有数据
This commit is contained in:
fawney19
2026-01-16 17:53:27 +08:00
parent 6c98816f9f
commit d5d74339dd
18 changed files with 247 additions and 136 deletions

View File

@@ -50,7 +50,7 @@ async def update_endpoint_key(
- `api_key`: 新的 API Key 原文
- `name`: Key 名称
- `note`: 备注
- `rate_multiplier`: 速率倍数
- `rate_multipliers`: 按 API 格式的成本倍率
- `internal_priority`: 内部优先级
- `rpm_limit`: RPM 限制(设置为 null 可切换到自适应模式)
- `allowed_models`: 允许的模型列表
@@ -82,8 +82,9 @@ async def get_keys_grouped_by_format(
- `name`: Key 名称
- `api_key_masked`: 脱敏后的 API Key
- `internal_priority`: 内部优先级
- `global_priority`: 全局优先级
- `rate_multiplier`: 速率倍数
- `global_priority_by_format`: 按 API 格式的全局优先级
- `format_priority`: 当前格式的优先级
- `rate_multipliers`: 按 API 格式的成本倍率
- `is_active`: 是否活跃
- `circuit_breaker_open`: 熔断器状态
- `provider_name`: Provider 名称
@@ -367,7 +368,6 @@ class AdminGetKeysGroupedByFormatAdapter(AdminApiAdapter):
Provider.is_active.is_(True),
)
.order_by(
ProviderAPIKey.global_priority.asc().nullslast(),
ProviderAPIKey.internal_priority.asc(),
)
.all()
@@ -427,8 +427,8 @@ class AdminGetKeysGroupedByFormatAdapter(AdminApiAdapter):
"name": key.name,
"api_key_masked": masked_key,
"internal_priority": key.internal_priority,
"global_priority": key.global_priority,
"rate_multiplier": key.rate_multiplier,
"global_priority_by_format": key.global_priority_by_format,
"rate_multipliers": key.rate_multipliers,
"is_active": key.is_active,
"provider_name": provider.name,
"api_formats": api_formats,
@@ -438,9 +438,10 @@ class AdminGetKeysGroupedByFormatAdapter(AdminApiAdapter):
"request_count": key.request_count,
}
# 将 Key 添加到每个支持的格式分组中,并附加格式特定的健康度数据
# 将 Key 添加到每个支持的格式分组中,并附加格式特定的数据
health_by_format = key.health_by_format or {}
circuit_by_format = key.circuit_breaker_by_format or {}
priority_by_format = key.global_priority_by_format or {}
provider_id = str(provider.id)
for api_format in api_formats:
if api_format not in grouped:
@@ -451,6 +452,8 @@ class AdminGetKeysGroupedByFormatAdapter(AdminApiAdapter):
format_key_info["endpoint_base_url"] = endpoint_base_url_map.get(
(provider_id, api_format)
)
# 添加格式特定的优先级
format_key_info["format_priority"] = priority_by_format.get(api_format)
# 添加格式特定的健康度数据
format_health = health_by_format.get(api_format, {})
format_circuit = circuit_by_format.get(api_format, {})
@@ -597,8 +600,7 @@ class AdminCreateProviderKeyAdapter(AdminApiAdapter):
api_key=encrypted_key,
name=self.key_data.name,
note=self.key_data.note,
rate_multiplier=self.key_data.rate_multiplier,
rate_multipliers=self.key_data.rate_multipliers, # 按 API 格式的成本倍率
rate_multipliers=self.key_data.rate_multipliers,
internal_priority=self.key_data.internal_priority,
rpm_limit=self.key_data.rpm_limit,
allowed_models=self.key_data.allowed_models if self.key_data.allowed_models else None,

View File

@@ -47,7 +47,7 @@ class RoutingKeyInfo(BaseModel):
name: str
masked_key: str = Field("", description="脱敏的 API Key")
internal_priority: int = Field(..., description="Key 内部优先级")
global_priority: Optional[int] = Field(None, description="全局 Key 优先级")
global_priority_by_format: Optional[Dict[str, int]] = Field(None, description="按 API 格式的全局优先级")
rpm_limit: Optional[int] = Field(None, description="RPM 限制null 表示自适应")
is_adaptive: bool = Field(False, description="是否为自适应 RPM 模式")
effective_rpm: Optional[int] = Field(None, description="有效 RPM 限制")
@@ -293,8 +293,14 @@ class AdminGetModelRoutingPreviewAdapter(AdminApiAdapter):
for ep in provider_endpoints:
# 获取该 Endpoint 格式对应的 Keys
ep_keys = keys_by_endpoint.get(ep.api_format or "", [])
# 按优先级排序
ep_keys.sort(key=lambda k: (k.global_priority or 999, k.internal_priority or 0))
# 按优先级排序(使用当前格式的全局优先级)
api_format = ep.api_format or ""
def get_key_priority(k: ProviderAPIKey) -> tuple[int, int]:
format_priority = 999
if k.global_priority_by_format and api_format in k.global_priority_by_format:
format_priority = k.global_priority_by_format[api_format]
return (format_priority, k.internal_priority or 0)
ep_keys.sort(key=get_key_priority)
key_infos = []
for key in ep_keys:
@@ -353,7 +359,7 @@ class AdminGetModelRoutingPreviewAdapter(AdminApiAdapter):
name=key.name or "",
masked_key=masked_key,
internal_priority=key.internal_priority or 0,
global_priority=key.global_priority,
global_priority_by_format=key.global_priority_by_format,
rpm_limit=key.rpm_limit,
is_adaptive=is_adaptive,
effective_rpm=effective_rpm,

View File

@@ -213,12 +213,11 @@ async def list_affinities(
- `provider_id`: Provider ID
- `provider_name`: Provider 显示名称
- `endpoint_id`: Endpoint ID
- `endpoint_api_format`: Endpoint API 格式
- `endpoint_url`: Endpoint 基础 URL
- `key_id`: Key ID
- `key_name`: Key 名称
- `key_prefix`: 脱敏后的 Provider Key
- `rate_multiplier`: 速率倍数
- `rate_multipliers`: 按 API 格式的成本倍率
- `global_model_id`: GlobalModel ID
- `model_name`: 模型名称
- `model_display_name`: 模型显示名称
@@ -821,14 +820,11 @@ class AdminListAffinitiesAdapter(AdminApiAdapter):
"provider_id": provider_id,
"provider_name": provider.name if provider else None,
"endpoint_id": endpoint_id,
"endpoint_api_format": (
endpoint.api_format if endpoint and endpoint.api_format else None
),
"endpoint_url": endpoint.base_url if endpoint else None,
"key_id": key_id,
"key_name": key.name if key else None,
"key_prefix": provider_key_masked,
"rate_multiplier": key.rate_multiplier if key else 1.0,
"rate_multipliers": key.rate_multipliers if key else None,
"global_model_id": affinity.get("model_name"), # 原始的 global_model_id
"model_name": (
global_model_map.get(affinity.get("model_name")).name

View File

@@ -803,10 +803,9 @@ class AdminExportConfigAdapter(AdminApiAdapter):
"name": key.name,
"note": key.note,
"api_formats": key.api_formats or [],
"rate_multiplier": key.rate_multiplier,
"rate_multipliers": key.rate_multipliers,
"internal_priority": key.internal_priority,
"global_priority": key.global_priority,
"global_priority_by_format": key.global_priority_by_format,
"rpm_limit": key.rpm_limit,
"allowed_models": key.allowed_models,
"capabilities": key.capabilities,
@@ -1159,10 +1158,9 @@ class AdminImportConfigAdapter(AdminApiAdapter):
api_key=encrypted_key,
name=key_data.get("name") or "Imported Key",
note=key_data.get("note"),
rate_multiplier=key_data.get("rate_multiplier", 1.0),
rate_multipliers=key_data.get("rate_multipliers"),
internal_priority=key_data.get("internal_priority", 50),
global_priority=key_data.get("global_priority"),
global_priority_by_format=key_data.get("global_priority_by_format"),
rpm_limit=key_data.get("rpm_limit"),
allowed_models=key_data.get("allowed_models"),
capabilities=key_data.get("capabilities"),

View File

@@ -551,11 +551,7 @@ class Provider(Base):
# 限制
concurrent_limit = Column(Integer, nullable=True) # 并发请求限制
# 请求配置(从 Endpoint 迁移,作为全局默认值)
# [已废弃] timeout 字段不再使用,超时由环境变量控制:
# - 非流式请求: HTTP_REQUEST_TIMEOUT默认 300 秒)
# - 流式首字节: STREAM_FIRST_BYTE_TIMEOUT默认 30 秒)
timeout = Column(Integer, default=300, nullable=True) # [已废弃] 请求超时(秒)
# 请求配置
max_retries = Column(Integer, default=2, nullable=True) # 最大重试次数
proxy = Column(JSONB, nullable=True) # 代理配置: {url, username, password, enabled}
@@ -603,7 +599,6 @@ class ProviderEndpoint(Base):
# 请求配置
header_rules = Column(JSON, nullable=True) # 请求头规则 [{action, key, value, from, to}]
timeout = Column(Integer, default=300) # [已废弃] 超时(秒),由环境变量控制
max_retries = Column(Integer, default=2) # 最大重试次数
# 状态
@@ -1004,11 +999,6 @@ class ProviderAPIKey(Base):
note = Column(String(500), nullable=True) # 备注说明(可选)
# 成本计算
# [DEPRECATED] rate_multiplier 已废弃,请使用 rate_multipliers
# 将在未来版本中移除,目前仅作为 rate_multipliers 未配置时的回退值
rate_multiplier = Column(
Float, default=1.0, nullable=False
) # [DEPRECATED] 默认成本倍率,请使用 rate_multipliers
rate_multipliers = Column(
JSON, nullable=True
) # 按 API 格式的成本倍率 {"CLAUDE_CLI": 1.0, "OPENAI_CLI": 0.8}
@@ -1017,9 +1007,9 @@ class ProviderAPIKey(Base):
internal_priority = Column(
Integer, default=50
) # Endpoint 内部优先级(用于提供商优先模式,同 Endpoint 内 Keys 的排序,同优先级参与负载均衡)
global_priority = Column(
Integer, nullable=True
) # 全局 Key 优先级(用于全局 Key 优先模式,跨 Provider 的 Key 排序NULL=未配置使用默认排序)
global_priority_by_format = Column(
JSON, nullable=True
) # 按 API 格式的全局优先级 {"CLAUDE": 1, "CLAUDE_CLI": 2}
# RPM 限制配置(自适应学习)
# rpm_limit 决定 RPM 控制模式:

View File

@@ -152,10 +152,6 @@ class EndpointAPIKeyCreate(BaseModel):
name: str = Field(..., min_length=1, max_length=100, description="密钥名称(必填,用于识别)")
# 成本计算
# [DEPRECATED] rate_multiplier 已废弃,请使用 rate_multipliers
rate_multiplier: float = Field(
default=1.0, ge=0.01, description="[DEPRECATED] 默认成本倍率,已废弃,请使用 rate_multipliers"
)
rate_multipliers: Optional[Dict[str, float]] = Field(
default=None, description="按 API 格式的成本倍率,如 {'CLAUDE_CLI': 1.0, 'OPENAI_CLI': 0.8}"
)
@@ -294,16 +290,14 @@ class EndpointAPIKeyUpdate(BaseModel):
default=None, min_length=3, max_length=500, description="API Key将自动加密"
)
name: Optional[str] = Field(default=None, min_length=1, max_length=100, description="密钥名称")
# [DEPRECATED] rate_multiplier 已废弃,请使用 rate_multipliers
rate_multiplier: Optional[float] = Field(default=None, ge=0.01, description="[DEPRECATED] 默认成本倍率,已废弃")
rate_multipliers: Optional[Dict[str, float]] = Field(
default=None, description="按 API 格式的成本倍率,如 {'CLAUDE_CLI': 1.0, 'OPENAI_CLI': 0.8}"
)
internal_priority: Optional[int] = Field(
default=None, description="Key 内部优先级(提供商优先模式,数字越小越优先)"
)
global_priority: Optional[int] = Field(
default=None, description="全局 Key 优先级(全局 Key 优先模式,数字越小越优先)"
global_priority_by_format: Optional[Dict[str, int]] = Field(
default=None, description="按 API 格式的全局优先级,如 {'CLAUDE': 1, 'CLAUDE_CLI': 2}"
)
# rpm_limit: 使用特殊标记区分"未提供"和"设置为 null自适应模式"
# - 不提供字段:不更新
@@ -421,15 +415,15 @@ class EndpointAPIKeyResponse(BaseModel):
name: str = Field(..., description="密钥名称")
# 成本计算
# [DEPRECATED] rate_multiplier 已废弃,请使用 rate_multipliers
rate_multiplier: float = Field(default=1.0, description="[DEPRECATED] 默认成本倍率,已废弃")
rate_multipliers: Optional[Dict[str, float]] = Field(
default=None, description="按 API 格式的成本倍率,如 {'CLAUDE_CLI': 1.0, 'OPENAI_CLI': 0.8}"
)
# 优先级和限制
internal_priority: int = Field(default=50, description="Endpoint 内部优先级")
global_priority: Optional[int] = Field(default=None, description="全局 Key 优先级")
global_priority_by_format: Optional[Dict[str, int]] = Field(
default=None, description="按 API 格式的全局优先级"
)
rpm_limit: Optional[int] = None
allowed_models: Optional[List[str]] = None
capabilities: Optional[Dict[str, bool]] = Field(default=None, description="Key 能力标签")
@@ -591,7 +585,6 @@ class ProviderUpdateRequest(BaseModel):
quota_reset_day: Optional[int] = Field(None, ge=1, le=31, description="配额重置日1-31")
quota_expires_at: Optional[datetime] = Field(None, description="配额过期时间")
# 请求配置(从 Endpoint 迁移)
timeout: Optional[int] = Field(None, ge=1, le=600, description="请求超时(秒)")
max_retries: Optional[int] = Field(None, ge=0, le=10, description="最大重试次数")
proxy: Optional[Dict[str, Any]] = Field(None, description="代理配置")

View File

@@ -670,7 +670,7 @@ class CacheAwareScheduler:
)
# 3. 应用优先级模式排序
candidates = self._apply_priority_mode_sort(candidates, affinity_key)
candidates = self._apply_priority_mode_sort(candidates, affinity_key, target_format.value)
# 更新指标
self._metrics["total_candidates"] += len(candidates)
@@ -693,7 +693,7 @@ class CacheAwareScheduler:
)
elif self.scheduling_mode == self.SCHEDULING_MODE_LOAD_BALANCE:
# 负载均衡模式:忽略缓存,同优先级内随机轮换
candidates = self._apply_load_balance(candidates)
candidates = self._apply_load_balance(candidates, target_format.value)
for candidate in candidates:
candidate.is_cached = False
else:
@@ -1188,50 +1188,57 @@ class CacheAwareScheduler:
logger.debug(f"[CacheAwareScheduler] 切换调度模式为: {self.scheduling_mode}")
def _apply_priority_mode_sort(
self, candidates: List[ProviderCandidate], affinity_key: Optional[str] = None
self, candidates: List[ProviderCandidate], affinity_key: Optional[str] = None,
api_format: Optional[str] = None
) -> List[ProviderCandidate]:
"""
根据优先级模式对候选列表排序(数字越小越优先)
- provider: 提供商优先模式,保持原有顺序(按 Provider.provider_priority -> Key.internal_priority 排序,已由查询保证)
Key.internal_priority 表示 Endpoint 内部优先级,同优先级内通过哈希分散负载均衡
- global_key: 全局 Key 优先模式,按 Key.global_priority 升序排序(数字小的优先)
global_priority 的优先NULL 的排后面
global_priority 内通过哈希分散实现负载均衡
- global_key: 全局 Key 优先模式,按 Key.global_priority_by_format 升序排序(数字小的优先)
优先级的优先NULL 的排后面
优先级内通过哈希分散实现负载均衡
"""
if not candidates:
return candidates
if self.priority_mode == self.PRIORITY_MODE_GLOBAL_KEY:
# 全局 Key 优先模式:按 global_priority 分组,同组内哈希分散负载均衡
return self._sort_by_global_priority_with_hash(candidates, affinity_key)
return self._sort_by_global_priority_with_hash(candidates, affinity_key, api_format)
# 提供商优先模式保持原有顺序provider_priority 排序已经由查询保证)
return candidates
def _sort_by_global_priority_with_hash(
self, candidates: List[ProviderCandidate], affinity_key: Optional[str] = None
self, candidates: List[ProviderCandidate], affinity_key: Optional[str] = None,
api_format: Optional[str] = None
) -> List[ProviderCandidate]:
"""
按 global_priority 分组排序,同优先级内通过哈希分散实现负载均衡
按 global_priority_by_format 分组排序,同优先级内通过哈希分散实现负载均衡
排序逻辑:
1. 按 global_priority 分组数字小的优先NULL 排后面)
2. 同 global_priority 组内,使用 affinity_key 哈希分散
1. 按 global_priority_by_format[api_format] 分组数字小的优先NULL 排后面)
2. 同优先级组内,使用 affinity_key 哈希分散
3. 确保同一用户请求稳定选择同一个 Key缓存亲和性
"""
import hashlib
from collections import defaultdict
# 按 global_priority 分组
def get_priority(candidate: ProviderCandidate) -> int:
"""获取候选的优先级"""
if not candidate.key:
return 999999
priority_by_format = candidate.key.global_priority_by_format or {}
if api_format and api_format in priority_by_format:
return priority_by_format[api_format]
return 999999 # NULL 排在后面
# 按优先级分组
priority_groups: Dict[int, List[ProviderCandidate]] = defaultdict(list)
for candidate in candidates:
global_priority = (
candidate.key.global_priority
if candidate.key and candidate.key.global_priority is not None
else 999999 # NULL 排在后面
)
priority_groups[global_priority].append(candidate)
priority = get_priority(candidate)
priority_groups[priority].append(candidate)
result = []
for priority in sorted(priority_groups.keys()): # 数字小的优先级高
@@ -1263,13 +1270,13 @@ class CacheAwareScheduler:
return result
def _apply_load_balance(
self, candidates: List[ProviderCandidate]
self, candidates: List[ProviderCandidate], api_format: Optional[str] = None
) -> List[ProviderCandidate]:
"""
负载均衡模式:同优先级内随机轮换
排序逻辑:
1. 按优先级分组provider_priority, internal_priority 或 global_priority
1. 按优先级分组provider_priority, internal_priority 或 global_priority_by_format
2. 同优先级组内随机打乱
3. 不考虑缓存亲和性
"""
@@ -1282,14 +1289,14 @@ class CacheAwareScheduler:
# 根据优先级模式选择分组方式
if self.priority_mode == self.PRIORITY_MODE_GLOBAL_KEY:
# 全局 Key 优先模式:按 global_priority 分组
# 全局 Key 优先模式:按格式特定优先级分组
for candidate in candidates:
global_priority = (
candidate.key.global_priority
if candidate.key and candidate.key.global_priority is not None
else 999999
)
priority_groups[(global_priority,)].append(candidate)
priority = 999999
if candidate.key:
priority_by_format = candidate.key.global_priority_by_format or {}
if api_format and api_format in priority_by_format:
priority = priority_by_format[api_format]
priority_groups[(priority,)].append(candidate)
else:
# 提供商优先模式:按 (provider_priority, internal_priority) 分组
for candidate in candidates:

View File

@@ -27,22 +27,15 @@ class ProviderCacheService:
@staticmethod
def compute_rate_multiplier(
rate_multiplier: Optional[float],
rate_multipliers: Optional[dict],
api_format: Optional[str] = None,
) -> float:
"""
计算 rate_multiplier 的纯函数(无数据库/缓存依赖)
优先返回指定 API 格式的倍率,如果没有则返回默认倍率
规则:
- 如果指定了 api_format 且 rate_multipliers 存在:
- 如果 rate_multipliers[api_format] 存在,返回它
- 否则返回 1.0rate_multipliers 存在但该格式未配置)
- 否则返回 rate_multiplier 或 1.0
返回指定 API 格式的倍率,如果没有则返回 1.0
Args:
rate_multiplier: 默认倍率
rate_multipliers: 按 API 格式的倍率配置字典
api_format: API 格式(可选),如 "CLAUDE""OPENAI"
@@ -53,12 +46,7 @@ class ProviderCacheService:
format_upper = api_format.upper()
if format_upper in rate_multipliers:
return float(rate_multipliers[format_upper])
else:
# rate_multipliers 存在但该格式未配置,使用默认值 1.0
return 1.0
else:
# rate_multipliers 不存在或未指定 api_format回退到默认倍率
return rate_multiplier or 1.0
return 1.0
@staticmethod
async def get_provider_api_key_rate_multiplier(
@@ -92,7 +80,7 @@ class ProviderCacheService:
# 2. 缓存未命中,查询数据库
provider_key = (
db.query(ProviderAPIKey.rate_multiplier, ProviderAPIKey.rate_multipliers)
db.query(ProviderAPIKey.rate_multipliers)
.filter(ProviderAPIKey.id == provider_api_key_id)
.first()
)
@@ -100,7 +88,7 @@ class ProviderCacheService:
# 3. 计算倍率并写入缓存
if provider_key:
rate_multiplier = ProviderCacheService.compute_rate_multiplier(
provider_key.rate_multiplier, provider_key.rate_multipliers, api_format
provider_key.rate_multipliers, api_format
)
await CacheService.set(

View File

@@ -1567,7 +1567,7 @@ class UsageService:
from src.services.cache.provider_cache import ProviderCacheService
provider_key = (
db.query(ProviderAPIKey.rate_multiplier, ProviderAPIKey.rate_multipliers)
db.query(ProviderAPIKey.rate_multipliers)
.filter(ProviderAPIKey.id == provider_api_key_id)
.first()
)
@@ -1576,7 +1576,7 @@ class UsageService:
return None
return ProviderCacheService.compute_rate_multiplier(
provider_key.rate_multiplier, provider_key.rate_multipliers, api_format
provider_key.rate_multipliers, api_format
)
@classmethod