mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-03 01:40:21 +08:00
refactor: 全局适配 ApiFamily/EndpointKind 结构化标识体系
将新的 (ApiFamily, EndpointKind) / `family:kind` 签名体系应用到整个代码库: - API Handlers: 所有 adapter/handler 使用新的签名格式 - Services: provider, model, usage, cache, auth 等服务层适配 - Database: ProviderEndpoint 新增 api_family/endpoint_kind 字段 - Frontend: Provider 管理、Usage 表格等组件适配 - Tests: 更新所有相关测试用例
This commit is contained in:
@@ -12,7 +12,6 @@ from typing import Any
|
||||
|
||||
from pydantic import BaseModel, Field, field_validator
|
||||
|
||||
from src.core.api_format import APIFormat
|
||||
from src.core.enums import ProviderBillingType
|
||||
|
||||
|
||||
@@ -71,8 +70,20 @@ class CreateProviderRequest(BaseModel):
|
||||
|
||||
# 检查 SQL 注入关键字(不区分大小写)
|
||||
sql_keywords = [
|
||||
"SELECT", "INSERT", "UPDATE", "DELETE", "DROP", "CREATE",
|
||||
"ALTER", "TRUNCATE", "UNION", "EXEC", "EXECUTE", "--", "/*", "*/"
|
||||
"SELECT",
|
||||
"INSERT",
|
||||
"UPDATE",
|
||||
"DELETE",
|
||||
"DROP",
|
||||
"CREATE",
|
||||
"ALTER",
|
||||
"TRUNCATE",
|
||||
"UNION",
|
||||
"EXEC",
|
||||
"EXECUTE",
|
||||
"--",
|
||||
"/*",
|
||||
"*/",
|
||||
]
|
||||
v_upper = v.upper()
|
||||
for keyword in sql_keywords:
|
||||
@@ -80,6 +91,7 @@ class CreateProviderRequest(BaseModel):
|
||||
raise ValueError(f"名称包含非法关键字: {keyword}")
|
||||
|
||||
return v
|
||||
|
||||
billing_type: str | None = Field(
|
||||
ProviderBillingType.PAY_AS_YOU_GO.value, description="计费类型"
|
||||
)
|
||||
@@ -87,15 +99,21 @@ class CreateProviderRequest(BaseModel):
|
||||
quota_reset_day: int | None = Field(30, ge=1, le=365, description="配额重置周期(天数)")
|
||||
quota_last_reset_at: datetime | None = Field(None, description="当前周期开始时间")
|
||||
quota_expires_at: datetime | None = Field(None, description="配额过期时间")
|
||||
provider_priority: int | None = Field(100, ge=0, le=1000, description="提供商优先级(数字越小越优先)")
|
||||
provider_priority: int | None = Field(
|
||||
100, ge=0, le=1000, description="提供商优先级(数字越小越优先)"
|
||||
)
|
||||
is_active: bool | None = Field(True, description="是否启用")
|
||||
concurrent_limit: int | None = Field(None, ge=0, description="并发限制")
|
||||
# 请求配置(从 Endpoint 迁移)
|
||||
max_retries: int | None = Field(2, ge=0, le=10, description="最大重试次数")
|
||||
proxy: ProxyConfig | None = Field(None, description="代理配置")
|
||||
# 超时配置(秒),为空时使用全局配置
|
||||
stream_first_byte_timeout: float | None = Field(None, ge=1, le=300, description="流式请求首字节超时(秒)")
|
||||
request_timeout: float | None = Field(None, ge=1, le=600, description="非流式请求整体超时(秒)")
|
||||
stream_first_byte_timeout: float | None = Field(
|
||||
None, ge=1, le=300, description="流式请求首字节超时(秒)"
|
||||
)
|
||||
request_timeout: float | None = Field(
|
||||
None, ge=1, le=600, description="非流式请求整体超时(秒)"
|
||||
)
|
||||
config: dict[str, Any] | None = Field(None, description="其他配置")
|
||||
|
||||
@field_validator("name", "description")
|
||||
@@ -167,8 +185,12 @@ class UpdateProviderRequest(BaseModel):
|
||||
max_retries: int | None = Field(None, ge=0, le=10, description="最大重试次数")
|
||||
proxy: ProxyConfig | None = Field(None, description="代理配置")
|
||||
# 超时配置(秒),为空时使用全局配置
|
||||
stream_first_byte_timeout: float | None = Field(None, ge=1, le=300, description="流式请求首字节超时(秒)")
|
||||
request_timeout: float | None = Field(None, ge=1, le=600, description="非流式请求整体超时(秒)")
|
||||
stream_first_byte_timeout: float | None = Field(
|
||||
None, ge=1, le=300, description="流式请求首字节超时(秒)"
|
||||
)
|
||||
request_timeout: float | None = Field(
|
||||
None, ge=1, le=600, description="非流式请求整体超时(秒)"
|
||||
)
|
||||
config: dict[str, Any] | None = None
|
||||
|
||||
# 复用相同的验证器
|
||||
@@ -187,7 +209,9 @@ class CreateEndpointRequest(BaseModel):
|
||||
provider_id: str = Field(..., description="Provider ID")
|
||||
name: str = Field(..., min_length=1, max_length=100, description="Endpoint 名称")
|
||||
base_url: str = Field(..., min_length=1, max_length=500, description="API 基础 URL")
|
||||
api_format: str = Field(..., description="API 格式(CLAUDE 或 OPENAI)")
|
||||
api_format: str = Field(
|
||||
..., description="Endpoint signature(如 openai:chat, claude:cli, gemini:video)"
|
||||
)
|
||||
custom_path: str | None = Field(None, max_length=200, description="自定义路径")
|
||||
priority: int | None = Field(100, ge=0, le=1000, description="优先级")
|
||||
is_active: bool | None = Field(True, description="是否启用")
|
||||
@@ -216,12 +240,14 @@ class CreateEndpointRequest(BaseModel):
|
||||
@classmethod
|
||||
def validate_api_format(cls, v: str) -> str:
|
||||
"""验证 API 格式"""
|
||||
try:
|
||||
APIFormat(v)
|
||||
return v
|
||||
except ValueError:
|
||||
valid_formats = [f.value for f in APIFormat]
|
||||
raise ValueError(f"无效的 API 格式,有效值为: {', '.join(valid_formats)}")
|
||||
from src.core.api_format import list_endpoint_definitions, resolve_endpoint_definition
|
||||
from src.core.api_format.signature import normalize_signature_key
|
||||
|
||||
normalized = normalize_signature_key(v)
|
||||
if resolve_endpoint_definition(normalized) is None:
|
||||
valid_formats = [d.signature_key for d in list_endpoint_definitions()]
|
||||
raise ValueError(f"无效的 api_format,有效值为: {', '.join(valid_formats)}")
|
||||
return normalized
|
||||
|
||||
@field_validator("custom_path")
|
||||
@classmethod
|
||||
@@ -307,7 +333,9 @@ class UpdateUserRequest(BaseModel):
|
||||
|
||||
username: str | None = Field(None, min_length=1, max_length=50)
|
||||
email: str | None = Field(None, max_length=100)
|
||||
password: str | None = Field(None, min_length=6, max_length=128, description="新密码(留空保持不变)")
|
||||
password: str | None = Field(
|
||||
None, min_length=6, max_length=128, description="新密码(留空保持不变)"
|
||||
)
|
||||
quota_usd: float | None = Field(None, ge=0)
|
||||
is_active: bool | None = None
|
||||
role: str | None = None
|
||||
|
||||
@@ -244,9 +244,15 @@ class CreateUserRequest(BaseModel):
|
||||
quota_usd: float | None = Field(default=None, description="USD配额,null表示使用系统默认配额")
|
||||
unlimited: bool = Field(default=False, description="是否无限配额")
|
||||
# 访问限制字段
|
||||
allowed_providers: list[str] | None = Field(default=None, description="允许使用的提供商ID列表,null表示无限制")
|
||||
allowed_api_formats: list[str] | None = Field(default=None, description="允许使用的API格式列表,null表示无限制")
|
||||
allowed_models: list[str] | None = Field(default=None, description="允许使用的模型名称列表,null表示无限制")
|
||||
allowed_providers: list[str] | None = Field(
|
||||
default=None, description="允许使用的提供商ID列表,null表示无限制"
|
||||
)
|
||||
allowed_api_formats: list[str] | None = Field(
|
||||
default=None, description="允许使用的API格式列表,null表示无限制"
|
||||
)
|
||||
allowed_models: list[str] | None = Field(
|
||||
default=None, description="允许使用的模型名称列表,null表示无限制"
|
||||
)
|
||||
|
||||
@field_validator("quota_usd", mode="before")
|
||||
@classmethod
|
||||
@@ -285,6 +291,30 @@ class CreateUserRequest(BaseModel):
|
||||
raise ValueError("用户名只能包含字母、数字、下划线、连字符和点号")
|
||||
return v
|
||||
|
||||
@field_validator("allowed_api_formats")
|
||||
@classmethod
|
||||
def validate_allowed_api_formats(cls, v: list[str] | None) -> list[str] | None:
|
||||
"""校验并规范化 allowed_api_formats(endpoint signature: family:kind)。"""
|
||||
if v is None:
|
||||
return None
|
||||
from src.core.api_format import list_endpoint_definitions, resolve_endpoint_definition
|
||||
from src.core.api_format.signature import normalize_signature_key
|
||||
|
||||
allowed = [d.signature_key for d in list_endpoint_definitions()]
|
||||
out: list[str] = []
|
||||
seen: set[str] = set()
|
||||
for fmt in v:
|
||||
if not fmt:
|
||||
continue
|
||||
norm = normalize_signature_key(fmt)
|
||||
if resolve_endpoint_definition(norm) is None:
|
||||
raise ValueError(f"allowed_api_formats 必须是以下之一: {allowed},当前值: {fmt}")
|
||||
if norm in seen:
|
||||
continue
|
||||
seen.add(norm)
|
||||
out.append(norm)
|
||||
return out
|
||||
|
||||
@classmethod
|
||||
@field_validator("password")
|
||||
def validate_password(cls, v: Any) -> Any:
|
||||
@@ -313,6 +343,12 @@ class UpdateUserRequest(BaseModel):
|
||||
quota_usd: float | None = None
|
||||
is_active: bool | None = None
|
||||
|
||||
@field_validator("allowed_api_formats")
|
||||
@classmethod
|
||||
def validate_allowed_api_formats(cls, v: list[str] | None) -> list[str] | None:
|
||||
# 与 CreateUserRequest 保持一致
|
||||
return CreateUserRequest.validate_allowed_api_formats(v)
|
||||
|
||||
@field_validator("quota_usd", mode="before")
|
||||
@classmethod
|
||||
def validate_quota_usd(cls, v: Any) -> Any:
|
||||
@@ -344,6 +380,12 @@ class CreateApiKeyRequest(BaseModel):
|
||||
False, description="过期后是否自动删除(True=物理删除,False=仅禁用)"
|
||||
)
|
||||
|
||||
@field_validator("allowed_api_formats")
|
||||
@classmethod
|
||||
def validate_allowed_api_formats(cls, v: list[str] | None) -> list[str] | None:
|
||||
# 与 CreateUserRequest 保持一致
|
||||
return CreateUserRequest.validate_allowed_api_formats(v)
|
||||
|
||||
|
||||
class UserResponse(BaseModel):
|
||||
"""用户响应"""
|
||||
@@ -408,8 +450,12 @@ class ProviderCreate(BaseModel):
|
||||
is_active: bool = Field(False, description="是否启用(默认false,需要配置API密钥后才能启用)")
|
||||
|
||||
# 超时配置(秒),为空时使用全局配置
|
||||
stream_first_byte_timeout: float | None = Field(None, ge=1, le=300, description="流式请求首字节超时(秒)")
|
||||
request_timeout: float | None = Field(None, ge=1, le=600, description="非流式请求整体超时(秒)")
|
||||
stream_first_byte_timeout: float | None = Field(
|
||||
None, ge=1, le=300, description="流式请求首字节超时(秒)"
|
||||
)
|
||||
request_timeout: float | None = Field(
|
||||
None, ge=1, le=600, description="非流式请求整体超时(秒)"
|
||||
)
|
||||
|
||||
|
||||
class ProviderUpdate(BaseModel):
|
||||
@@ -430,8 +476,12 @@ class ProviderUpdate(BaseModel):
|
||||
is_active: bool | None = None
|
||||
|
||||
# 超时配置(秒),为空时使用全局配置
|
||||
stream_first_byte_timeout: float | None = Field(None, ge=1, le=300, description="流式请求首字节超时(秒)")
|
||||
request_timeout: float | None = Field(None, ge=1, le=600, description="非流式请求整体超时(秒)")
|
||||
stream_first_byte_timeout: float | None = Field(
|
||||
None, ge=1, le=300, description="流式请求首字节超时(秒)"
|
||||
)
|
||||
request_timeout: float | None = Field(
|
||||
None, ge=1, le=600, description="非流式请求整体超时(秒)"
|
||||
)
|
||||
|
||||
|
||||
class ProviderResponse(BaseModel):
|
||||
|
||||
@@ -707,7 +707,11 @@ class ProviderEndpoint(Base):
|
||||
provider_id = Column(String(36), ForeignKey("providers.id", ondelete="CASCADE"), nullable=False)
|
||||
|
||||
# API 格式和配置
|
||||
api_format = Column(String(50), nullable=False) # 存储 APIFormat 枚举值的字符串
|
||||
# 新模式:存储 endpoint signature key(family:kind),如 "openai:chat"
|
||||
api_format = Column(String(50), nullable=False)
|
||||
# 新架构字段(Phase 1/3):用于将 api_format 拆分为结构化维度
|
||||
api_family = Column(String(50), nullable=True) # openai/claude/gemini
|
||||
endpoint_kind = Column(String(50), nullable=True) # chat/cli/video/...
|
||||
base_url = Column(String(500), nullable=False)
|
||||
|
||||
# 请求配置
|
||||
@@ -754,6 +758,7 @@ class ProviderEndpoint(Base):
|
||||
__table_args__ = (
|
||||
UniqueConstraint("provider_id", "api_format", name="uq_provider_api_format"),
|
||||
Index("idx_endpoint_format_active", "api_format", "is_active"),
|
||||
Index("idx_provider_family_kind", "provider_id", "api_family", "endpoint_kind"),
|
||||
)
|
||||
|
||||
|
||||
@@ -1021,7 +1026,7 @@ class Model(Base):
|
||||
|
||||
Args:
|
||||
affinity_key: 用于哈希分散的亲和键(如用户 API Key 哈希),确保同一用户稳定选择同一映射
|
||||
api_format: 当前请求的 API 格式(如 CLAUDE、OPENAI 等),用于过滤适用的映射
|
||||
api_format: 当前请求的 endpoint signature(如 "openai:chat"),用于过滤适用的映射
|
||||
"""
|
||||
import hashlib
|
||||
|
||||
@@ -1044,8 +1049,11 @@ class Model(Base):
|
||||
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
|
||||
if isinstance(mapping_api_formats, list):
|
||||
target = str(api_format).strip().lower()
|
||||
allowed = {str(fmt).strip().lower() for fmt in mapping_api_formats if fmt}
|
||||
if target not in allowed:
|
||||
continue
|
||||
|
||||
raw_priority = raw.get("priority", 1)
|
||||
try:
|
||||
@@ -1228,7 +1236,7 @@ class ProviderAPIKey(Base):
|
||||
|
||||
# API 格式支持列表(核心字段)
|
||||
# None 表示支持所有格式(兼容历史数据),空列表 [] 表示不支持任何格式
|
||||
api_formats = Column(JSON, nullable=True, default=list) # ["CLAUDE", "CLAUDE_CLI"]
|
||||
api_formats = Column(JSON, nullable=True, default=list) # ["claude:chat", "claude:cli"]
|
||||
|
||||
# 认证类型
|
||||
# - "api_key": 标准 API Key 认证(默认)
|
||||
@@ -1252,7 +1260,7 @@ class ProviderAPIKey(Base):
|
||||
# 成本计算
|
||||
rate_multipliers = Column(
|
||||
JSON, nullable=True
|
||||
) # 按 API 格式的成本倍率 {"CLAUDE_CLI": 1.0, "OPENAI_CLI": 0.8}
|
||||
) # 按 endpoint signature 的成本倍率 {"claude:cli": 1.0, "openai:cli": 0.8}
|
||||
|
||||
# 优先级配置 (数字越小越优先)
|
||||
internal_priority = Column(
|
||||
@@ -1260,7 +1268,7 @@ class ProviderAPIKey(Base):
|
||||
) # Endpoint 内部优先级(用于提供商优先模式,同 Endpoint 内 Keys 的排序,同优先级参与负载均衡)
|
||||
global_priority_by_format = Column(
|
||||
JSON, nullable=True
|
||||
) # 按 API 格式的全局优先级 {"CLAUDE": 1, "CLAUDE_CLI": 2}
|
||||
) # 按 endpoint signature 的全局优先级 {"claude:chat": 1, "claude:cli": 2}
|
||||
|
||||
# RPM 限制配置(自适应学习)
|
||||
# rpm_limit 决定 RPM 控制模式:
|
||||
@@ -1289,8 +1297,8 @@ class ProviderAPIKey(Base):
|
||||
) # 利用率采样窗口 [{"ts": timestamp, "util": 0.8}, ...]
|
||||
last_probe_increase_at = Column(DateTime(timezone=True), nullable=True) # 上次探测性扩容时间
|
||||
|
||||
# 健康度追踪(按 API 格式存储)
|
||||
# 结构: {"CLAUDE": {"health_score": 1.0, "consecutive_failures": 0, "last_failure_at": null, "request_results_window": []}, ...}
|
||||
# 健康度追踪(按 endpoint signature 存储)
|
||||
# 结构: {"claude:chat": {"health_score": 1.0, "consecutive_failures": 0, ...}, ...}
|
||||
health_by_format = Column(JSON, nullable=True, default=dict)
|
||||
|
||||
# 缓存与熔断配置
|
||||
@@ -1301,8 +1309,8 @@ class ProviderAPIKey(Base):
|
||||
Integer, default=32, nullable=False
|
||||
) # 最大探测间隔(分钟),默认32分钟(硬上限)
|
||||
|
||||
# 熔断器状态(按 API 格式存储)
|
||||
# 结构: {"CLAUDE": {"open": false, "open_at": null, "next_probe_at": null, "half_open_until": null, "half_open_successes": 0, "half_open_failures": 0}, ...}
|
||||
# 熔断器状态(按 endpoint signature 存储)
|
||||
# 结构: {"claude:chat": {"open": false, "open_at": null, ...}, ...}
|
||||
circuit_breaker_by_format = Column(JSON, nullable=True, default=dict)
|
||||
|
||||
# 使用统计
|
||||
|
||||
@@ -12,7 +12,6 @@ from pydantic import BaseModel, ConfigDict, Field, field_validator
|
||||
|
||||
from src.models.admin_requests import ProxyConfig
|
||||
|
||||
|
||||
# ========== Header Rule 类型定义 ==========
|
||||
# 请求头规则支持三种操作:
|
||||
# - set: 设置/覆盖请求头 {"action": "set", "key": "X-Custom", "value": "val"}
|
||||
@@ -29,7 +28,12 @@ class ProviderEndpointCreate(BaseModel):
|
||||
"""创建 Endpoint 请求"""
|
||||
|
||||
provider_id: str = Field(..., description="Provider ID")
|
||||
api_format: str = Field(..., description="API 格式 (CLAUDE, OPENAI, CLAUDE_CLI, OPENAI_CLI)")
|
||||
api_format: str = Field(
|
||||
...,
|
||||
description=(
|
||||
"Endpoint signature(例如: claude:chat/claude:cli, openai:chat/openai:cli/openai:video, gemini:chat/gemini:cli/gemini:video)"
|
||||
),
|
||||
)
|
||||
base_url: str = Field(..., min_length=1, max_length=500, description="API 基础 URL")
|
||||
custom_path: str | None = Field(default=None, max_length=200, description="自定义请求路径")
|
||||
|
||||
@@ -57,13 +61,14 @@ class ProviderEndpointCreate(BaseModel):
|
||||
@classmethod
|
||||
def validate_api_format(cls, v: str) -> str:
|
||||
"""验证 API 格式"""
|
||||
from src.core.api_format import APIFormat
|
||||
from src.core.api_format import list_endpoint_definitions, resolve_endpoint_definition
|
||||
from src.core.api_format.signature import normalize_signature_key
|
||||
|
||||
allowed = [fmt.value for fmt in APIFormat]
|
||||
v_upper = v.upper()
|
||||
if v_upper not in allowed:
|
||||
raise ValueError(f"API 格式必须是 {allowed} 之一")
|
||||
return v_upper
|
||||
normalized = normalize_signature_key(v)
|
||||
if resolve_endpoint_definition(normalized) is None:
|
||||
allowed = [d.signature_key for d in list_endpoint_definitions()]
|
||||
raise ValueError(f"api_format 必须是以下之一: {allowed}")
|
||||
return normalized
|
||||
|
||||
@field_validator("base_url")
|
||||
@classmethod
|
||||
@@ -125,9 +130,7 @@ class ProviderEndpointResponse(BaseModel):
|
||||
custom_path: str | None = None
|
||||
|
||||
# 请求头配置
|
||||
header_rules: list[HeaderRule] | None = Field(
|
||||
default=None, description="请求头规则列表"
|
||||
)
|
||||
header_rules: list[HeaderRule] | None = Field(default=None, description="请求头规则列表")
|
||||
|
||||
max_retries: int
|
||||
|
||||
@@ -165,23 +168,25 @@ class EndpointAPIKeyCreate(BaseModel):
|
||||
|
||||
provider_id: str | None = Field(default=None, description="Provider ID(从 URL 获取)")
|
||||
api_formats: list[str] | None = Field(
|
||||
default=None, min_length=1, description="支持的 API 格式列表(必填,路由层校验)"
|
||||
default=None, min_length=1, description="支持的 endpoint signature 列表(必填,路由层校验)"
|
||||
)
|
||||
|
||||
api_key: str = Field(default="", max_length=500, description="API Key(标准认证时必填,将自动加密)")
|
||||
api_key: str = Field(
|
||||
default="", max_length=500, description="API Key(标准认证时必填,将自动加密)"
|
||||
)
|
||||
auth_type: Literal["api_key", "vertex_ai"] = Field(
|
||||
default="api_key",
|
||||
description="认证类型:api_key(标准 API Key)或 vertex_ai(Vertex AI Service Account)"
|
||||
description="认证类型:api_key(标准 API Key)或 vertex_ai(Vertex AI Service Account)",
|
||||
)
|
||||
auth_config: dict[str, Any] | None = Field(
|
||||
default=None,
|
||||
description="认证配置(JSON):vertex_ai 时存储完整 Service Account JSON"
|
||||
default=None, description="认证配置(JSON):vertex_ai 时存储完整 Service Account JSON"
|
||||
)
|
||||
name: str = Field(..., min_length=1, max_length=100, description="密钥名称(必填,用于识别)")
|
||||
|
||||
# 成本计算
|
||||
rate_multipliers: dict[str, float] | None = Field(
|
||||
default=None, description="按 API 格式的成本倍率,如 {'CLAUDE_CLI': 1.0, 'OPENAI_CLI': 0.8}"
|
||||
default=None,
|
||||
description="按 endpoint signature 的成本倍率,如 {'claude:cli': 1.0, 'openai:cli': 0.8}",
|
||||
)
|
||||
|
||||
# 优先级和限制(数字越小越优先)
|
||||
@@ -236,19 +241,20 @@ class EndpointAPIKeyCreate(BaseModel):
|
||||
if v is None:
|
||||
return v
|
||||
|
||||
from src.core.api_format import APIFormat
|
||||
from src.core.api_format import list_endpoint_definitions, resolve_endpoint_definition
|
||||
from src.core.api_format.signature import normalize_signature_key
|
||||
|
||||
allowed = [fmt.value for fmt in APIFormat]
|
||||
validated = []
|
||||
seen = set()
|
||||
allowed = [d.signature_key for d in list_endpoint_definitions()]
|
||||
validated: list[str] = []
|
||||
seen: set[str] = set()
|
||||
for fmt in v:
|
||||
fmt_upper = fmt.upper()
|
||||
if fmt_upper not in allowed:
|
||||
raise ValueError(f"API 格式必须是 {allowed} 之一,当前值: {fmt}")
|
||||
if fmt_upper in seen:
|
||||
normalized = normalize_signature_key(fmt)
|
||||
if resolve_endpoint_definition(normalized) is None:
|
||||
raise ValueError(f"api_formats 必须是以下之一: {allowed},当前值: {fmt}")
|
||||
if normalized in seen:
|
||||
continue # 静默去重
|
||||
seen.add(fmt_upper)
|
||||
validated.append(fmt_upper)
|
||||
seen.add(normalized)
|
||||
validated.append(normalized)
|
||||
return validated
|
||||
|
||||
@field_validator("allowed_models")
|
||||
@@ -323,25 +329,29 @@ class EndpointAPIKeyUpdate(BaseModel):
|
||||
)
|
||||
|
||||
api_key: str | None = Field(
|
||||
default=None, min_length=3, max_length=500, description="API Key(标准认证时使用,将自动加密)"
|
||||
default=None,
|
||||
min_length=3,
|
||||
max_length=500,
|
||||
description="API Key(标准认证时使用,将自动加密)",
|
||||
)
|
||||
auth_type: Literal["api_key", "vertex_ai"] | None = Field(
|
||||
default=None,
|
||||
description="认证类型:api_key(标准 API Key)或 vertex_ai(Vertex AI Service Account)"
|
||||
description="认证类型:api_key(标准 API Key)或 vertex_ai(Vertex AI Service Account)",
|
||||
)
|
||||
auth_config: dict[str, Any] | None = Field(
|
||||
default=None,
|
||||
description="认证配置(JSON):vertex_ai 时存储完整 Service Account JSON"
|
||||
default=None, description="认证配置(JSON):vertex_ai 时存储完整 Service Account JSON"
|
||||
)
|
||||
name: str | None = Field(default=None, min_length=1, max_length=100, description="密钥名称")
|
||||
rate_multipliers: dict[str, float] | None = Field(
|
||||
default=None, description="按 API 格式的成本倍率,如 {'CLAUDE_CLI': 1.0, 'OPENAI_CLI': 0.8}"
|
||||
default=None,
|
||||
description="按 endpoint signature 的成本倍率,如 {'claude:cli': 1.0, 'openai:cli': 0.8}",
|
||||
)
|
||||
internal_priority: int | None = Field(
|
||||
default=None, description="Key 内部优先级(提供商优先模式,数字越小越优先)"
|
||||
)
|
||||
global_priority_by_format: dict[str, int] | None = Field(
|
||||
default=None, description="按 API 格式的全局优先级,如 {'CLAUDE': 1, 'CLAUDE_CLI': 2}"
|
||||
default=None,
|
||||
description="按 endpoint signature 的全局优先级,如 {'claude:chat': 1, 'claude:cli': 2}",
|
||||
)
|
||||
# rpm_limit: 使用特殊标记区分"未提供"和"设置为 null(自适应模式)"
|
||||
# - 不提供字段:不更新
|
||||
@@ -365,9 +375,7 @@ class EndpointAPIKeyUpdate(BaseModel):
|
||||
)
|
||||
is_active: bool | None = Field(default=None, description="是否启用")
|
||||
note: str | None = Field(default=None, max_length=500, description="备注说明")
|
||||
auto_fetch_models: bool | None = Field(
|
||||
default=None, description="是否启用自动获取模型"
|
||||
)
|
||||
auto_fetch_models: bool | None = Field(default=None, description="是否启用自动获取模型")
|
||||
locked_models: list[str] | None = Field(
|
||||
default=None, description="被锁定的模型列表(刷新时不会被删除)"
|
||||
)
|
||||
@@ -386,20 +394,7 @@ class EndpointAPIKeyUpdate(BaseModel):
|
||||
if v is None:
|
||||
return v
|
||||
|
||||
from src.core.api_format import APIFormat
|
||||
|
||||
allowed = [fmt.value for fmt in APIFormat]
|
||||
validated = []
|
||||
seen = set()
|
||||
for fmt in v:
|
||||
fmt_upper = fmt.upper()
|
||||
if fmt_upper not in allowed:
|
||||
raise ValueError(f"API 格式必须是 {allowed} 之一,当前值: {fmt}")
|
||||
if fmt_upper in seen:
|
||||
continue # 静默去重
|
||||
seen.add(fmt_upper)
|
||||
validated.append(fmt_upper)
|
||||
return validated
|
||||
return EndpointAPIKeyCreate.validate_api_formats(v)
|
||||
|
||||
@field_validator("allowed_models")
|
||||
@classmethod
|
||||
@@ -458,7 +453,9 @@ class EndpointAPIKeyResponse(BaseModel):
|
||||
id: str
|
||||
|
||||
provider_id: str = Field(..., description="Provider ID")
|
||||
api_formats: list[str] = Field(default=[], description="支持的 API 格式列表")
|
||||
api_formats: list[str] = Field(
|
||||
default=[], description="支持的 endpoint signature 列表(如 openai:chat, claude:cli)"
|
||||
)
|
||||
|
||||
# Key 信息(脱敏)
|
||||
api_key_masked: str = Field(..., description="脱敏后的 Key")
|
||||
@@ -469,13 +466,14 @@ class EndpointAPIKeyResponse(BaseModel):
|
||||
|
||||
# 成本计算
|
||||
rate_multipliers: dict[str, float] | None = Field(
|
||||
default=None, description="按 API 格式的成本倍率,如 {'CLAUDE_CLI': 1.0, 'OPENAI_CLI': 0.8}"
|
||||
default=None,
|
||||
description="按 endpoint signature 的成本倍率,如 {'claude:cli': 1.0, 'openai:cli': 0.8}",
|
||||
)
|
||||
|
||||
# 优先级和限制
|
||||
internal_priority: int = Field(default=50, description="Endpoint 内部优先级")
|
||||
global_priority_by_format: dict[str, int] | None = Field(
|
||||
default=None, description="按 API 格式的全局优先级"
|
||||
default=None, description="按 endpoint signature 的全局优先级"
|
||||
)
|
||||
rpm_limit: int | None = None
|
||||
allowed_models: list[str] | None = None
|
||||
@@ -485,12 +483,12 @@ class EndpointAPIKeyResponse(BaseModel):
|
||||
cache_ttl_minutes: int = Field(default=5, description="缓存 TTL(分钟),0=禁用")
|
||||
max_probe_interval_minutes: int = Field(default=32, description="熔断探测间隔(分钟)")
|
||||
|
||||
# 按格式的健康度数据
|
||||
# 按 endpoint signature 的健康度数据
|
||||
health_by_format: dict[str, Any] | None = Field(
|
||||
default=None, description="按 API 格式存储的健康度数据"
|
||||
default=None, description="按 endpoint signature 存储的健康度数据"
|
||||
)
|
||||
circuit_breaker_by_format: dict[str, Any] | None = Field(
|
||||
default=None, description="按 API 格式存储的熔断器状态"
|
||||
default=None, description="按 endpoint signature 存储的熔断器状态"
|
||||
)
|
||||
|
||||
# 聚合字段(从 health_by_format 计算,用于列表显示)
|
||||
@@ -648,8 +646,12 @@ class ProviderUpdateRequest(BaseModel):
|
||||
max_retries: int | None = Field(None, ge=0, le=10, description="最大重试次数")
|
||||
proxy: dict[str, Any] | None = Field(None, description="代理配置")
|
||||
# 超时配置(秒),为空时使用全局配置
|
||||
stream_first_byte_timeout: float | None = Field(None, ge=1, le=300, description="流式请求首字节超时(秒)")
|
||||
request_timeout: float | None = Field(None, ge=1, le=600, description="非流式请求整体超时(秒)")
|
||||
stream_first_byte_timeout: float | None = Field(
|
||||
None, ge=1, le=300, description="流式请求首字节超时(秒)"
|
||||
)
|
||||
request_timeout: float | None = Field(
|
||||
None, ge=1, le=600, description="非流式请求整体超时(秒)"
|
||||
)
|
||||
|
||||
|
||||
class ProviderWithEndpointsSummary(BaseModel):
|
||||
@@ -679,7 +681,9 @@ class ProviderWithEndpointsSummary(BaseModel):
|
||||
max_retries: int | None = Field(default=2, description="最大重试次数")
|
||||
proxy: dict[str, Any] | None = Field(default=None, description="代理配置")
|
||||
# 超时配置(秒),为空时使用全局配置
|
||||
stream_first_byte_timeout: float | None = Field(default=None, description="流式请求首字节超时(秒)")
|
||||
stream_first_byte_timeout: float | None = Field(
|
||||
default=None, description="流式请求首字节超时(秒)"
|
||||
)
|
||||
request_timeout: float | None = Field(default=None, description="非流式请求整体超时(秒)")
|
||||
|
||||
# Endpoint 统计
|
||||
@@ -780,9 +784,7 @@ class ApiFormatHealthMonitor(BaseModel):
|
||||
time_range_start: datetime | None = Field(
|
||||
default=None, description="时间线所覆盖区间的开始时间"
|
||||
)
|
||||
time_range_end: datetime | None = Field(
|
||||
default=None, description="时间线所覆盖区间的结束时间"
|
||||
)
|
||||
time_range_end: datetime | None = Field(default=None, description="时间线所覆盖区间的结束时间")
|
||||
|
||||
|
||||
class ApiFormatHealthMonitorResponse(BaseModel):
|
||||
|
||||
@@ -3,12 +3,12 @@ Pydantic 数据模型(阶段一统一模型管理)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, model_validator
|
||||
|
||||
|
||||
# ========== 阶梯计费相关模型 ==========
|
||||
|
||||
|
||||
@@ -16,25 +16,23 @@ class CacheTTLPricing(BaseModel):
|
||||
"""缓存时长定价配置"""
|
||||
|
||||
ttl_minutes: int = Field(..., ge=1, description="缓存时长(分钟)")
|
||||
cache_creation_price_per_1m: float = Field(..., ge=0, description="该时长的缓存创建价格/M tokens")
|
||||
cache_creation_price_per_1m: float = Field(
|
||||
..., ge=0, description="该时长的缓存创建价格/M tokens"
|
||||
)
|
||||
|
||||
|
||||
class PricingTier(BaseModel):
|
||||
"""单个价格阶梯配置"""
|
||||
|
||||
up_to: int | None = Field(
|
||||
None,
|
||||
ge=1,
|
||||
description="阶梯上限(tokens),null 表示无上限(最后一个阶梯)"
|
||||
None, ge=1, description="阶梯上限(tokens),null 表示无上限(最后一个阶梯)"
|
||||
)
|
||||
input_price_per_1m: float = Field(..., ge=0, description="输入价格/M tokens")
|
||||
output_price_per_1m: float = Field(..., ge=0, description="输出价格/M tokens")
|
||||
cache_creation_price_per_1m: float | None = Field(
|
||||
None, ge=0, description="缓存创建价格/M tokens"
|
||||
)
|
||||
cache_read_price_per_1m: float | None = Field(
|
||||
None, ge=0, description="缓存读取价格/M tokens"
|
||||
)
|
||||
cache_read_price_per_1m: float | None = Field(None, ge=0, description="缓存读取价格/M tokens")
|
||||
cache_ttl_pricing: list[CacheTTLPricing] | None = Field(
|
||||
None, description="按缓存时长分价格(可选)"
|
||||
)
|
||||
@@ -44,9 +42,7 @@ class TieredPricingConfig(BaseModel):
|
||||
"""阶梯计费配置"""
|
||||
|
||||
tiers: list[PricingTier] = Field(
|
||||
...,
|
||||
min_length=1,
|
||||
description="价格阶梯列表,按 up_to 升序排列"
|
||||
..., min_length=1, description="价格阶梯列表,按 up_to 升序排列"
|
||||
)
|
||||
|
||||
@model_validator(mode="after")
|
||||
@@ -78,9 +74,7 @@ class TieredPricingConfig(BaseModel):
|
||||
prev_ttl = 0
|
||||
for ttl_pricing in tier.cache_ttl_pricing:
|
||||
if ttl_pricing.ttl_minutes <= prev_ttl:
|
||||
raise ValueError(
|
||||
f"cache_ttl_pricing 必须按 ttl_minutes 升序排列"
|
||||
)
|
||||
raise ValueError(f"cache_ttl_pricing 必须按 ttl_minutes 升序排列")
|
||||
prev_ttl = ttl_pricing.ttl_minutes
|
||||
|
||||
# 最后一个阶梯必须是无上限的
|
||||
@@ -195,13 +189,10 @@ class GlobalModelCreate(BaseModel):
|
||||
..., description="阶梯计费配置(固定价格用单阶梯表示)"
|
||||
)
|
||||
# Key 能力配置 - 模型支持的能力列表(如 ["cache_1h", "context_1m"])
|
||||
supported_capabilities: list[str] | None = Field(
|
||||
None, description="支持的 Key 能力列表"
|
||||
)
|
||||
supported_capabilities: list[str] | None = Field(None, description="支持的 Key 能力列表")
|
||||
# 模型配置(JSON格式)- 包含能力、规格、元信息等
|
||||
config: dict[str, Any] | None = Field(
|
||||
None,
|
||||
description="模型配置(streaming, vision, context_limit, description 等)"
|
||||
None, description="模型配置(streaming, vision, context_limit, description 等)"
|
||||
)
|
||||
is_active: bool | None = Field(True, description="是否激活")
|
||||
|
||||
@@ -214,17 +205,12 @@ class GlobalModelUpdate(BaseModel):
|
||||
# 按次计费配置
|
||||
default_price_per_request: float | None = Field(None, ge=0, description="每次请求固定费用")
|
||||
# 阶梯计费配置
|
||||
default_tiered_pricing: TieredPricingConfig | None = Field(
|
||||
None, description="阶梯计费配置"
|
||||
)
|
||||
default_tiered_pricing: TieredPricingConfig | None = Field(None, description="阶梯计费配置")
|
||||
# Key 能力配置 - 模型支持的能力列表(如 ["cache_1h", "context_1m"])
|
||||
supported_capabilities: list[str] | None = Field(
|
||||
None, description="支持的 Key 能力列表"
|
||||
)
|
||||
supported_capabilities: list[str] | None = Field(None, description="支持的 Key 能力列表")
|
||||
# 模型配置(JSON格式)- 包含能力、规格、元信息等
|
||||
config: dict[str, Any] | None = Field(
|
||||
None,
|
||||
description="模型配置(streaming, vision, context_limit, description 等)"
|
||||
None, description="模型配置(streaming, vision, context_limit, description 等)"
|
||||
)
|
||||
|
||||
|
||||
@@ -247,8 +233,7 @@ class GlobalModelResponse(BaseModel):
|
||||
)
|
||||
# 模型配置(JSON格式)
|
||||
config: dict[str, Any] | None = Field(
|
||||
default=None,
|
||||
description="模型配置(streaming, vision, context_limit, description 等)"
|
||||
default=None, description="模型配置(streaming, vision, context_limit, description 等)"
|
||||
)
|
||||
# 统计数据(可选)
|
||||
provider_count: int | None = Field(default=0, description="支持的 Provider 数量")
|
||||
@@ -315,12 +300,10 @@ class ImportFromUpstreamRequest(BaseModel):
|
||||
# 价格覆盖配置(应用于所有导入的模型)
|
||||
tiered_pricing: dict | None = Field(
|
||||
None,
|
||||
description="阶梯计费配置(可选),格式: {tiers: [{up_to, input_price_per_1m, output_price_per_1m, ...}]}"
|
||||
description="阶梯计费配置(可选),格式: {tiers: [{up_to, input_price_per_1m, output_price_per_1m, ...}]}",
|
||||
)
|
||||
price_per_request: float | None = Field(
|
||||
None,
|
||||
ge=0,
|
||||
description="按次计费价格(可选,单位:美元)"
|
||||
None, ge=0, description="按次计费价格(可选,单位:美元)"
|
||||
)
|
||||
|
||||
|
||||
@@ -331,7 +314,9 @@ class ImportFromUpstreamSuccessItem(BaseModel):
|
||||
provider_model_id: str = Field(..., description="Provider Model ID")
|
||||
global_model_id: str | None = Field("", description="GlobalModel ID(如果已关联)")
|
||||
global_model_name: str | None = Field("", description="GlobalModel 名称(如果已关联)")
|
||||
created_global_model: bool = Field(False, description="是否新创建了 GlobalModel(始终为 false)")
|
||||
created_global_model: bool = Field(
|
||||
False, description="是否新创建了 GlobalModel(始终为 false)"
|
||||
)
|
||||
|
||||
|
||||
class ImportFromUpstreamErrorItem(BaseModel):
|
||||
|
||||
Reference in New Issue
Block a user