Merge branch 'fix/python314-upgrade'

# Conflicts:
#	src/api/handlers/base/base_handler.py
#	src/api/handlers/base/request_builder.py
#	src/models/endpoint_models.py
#	src/services/orchestration/candidate_resolver.py
#	src/services/orchestration/fallback_orchestrator.py
This commit is contained in:
fawney19
2026-01-30 12:59:52 +08:00
257 changed files with 4115 additions and 5236 deletions

View File

@@ -10,7 +10,6 @@
"""
from dataclasses import dataclass
from typing import List, Optional
from fastapi import APIRouter, Depends, HTTPException, Query, Request
from pydantic import BaseModel, Field, ValidationError
@@ -34,7 +33,7 @@ class EnableAdaptiveRequest(BaseModel):
"""启用自适应模式请求"""
enabled: bool = Field(..., description="是否启用自适应模式true=自适应false=固定限制)")
fixed_limit: Optional[int] = Field(
fixed_limit: int | None = Field(
None, ge=1, le=100, description="固定 RPM 限制(仅当 enabled=false 时生效1-100"
)
@@ -43,30 +42,30 @@ class AdaptiveStatsResponse(BaseModel):
"""自适应统计响应"""
adaptive_mode: bool = Field(..., description="是否为自适应模式rpm_limit=NULL")
rpm_limit: Optional[int] = Field(None, description="用户配置的固定限制NULL=自适应)")
effective_limit: Optional[int] = Field(
rpm_limit: int | None = Field(None, description="用户配置的固定限制NULL=自适应)")
effective_limit: int | None = Field(
None, description="当前有效限制(自适应使用学习值,固定使用配置值)"
)
learned_limit: Optional[int] = Field(None, description="学习到的 RPM 限制")
learned_limit: int | None = Field(None, description="学习到的 RPM 限制")
concurrent_429_count: int
rpm_429_count: int
last_429_at: Optional[str]
last_429_type: Optional[str]
last_429_at: str | None
last_429_type: str | None
adjustment_count: int
recent_adjustments: List[dict]
recent_adjustments: list[dict]
class KeyListItem(BaseModel):
"""Key 列表项"""
id: str
name: Optional[str]
name: str | None
provider_id: str
api_formats: List[str] = Field(default_factory=list)
api_formats: list[str] = Field(default_factory=list)
is_adaptive: bool = Field(..., description="是否为自适应模式rpm_limit=NULL")
rpm_limit: Optional[int] = Field(None, description="固定 RPM 限制NULL=自适应)")
effective_limit: Optional[int] = Field(None, description="当前有效限制")
learned_rpm_limit: Optional[int] = Field(None, description="学习到的 RPM 限制")
rpm_limit: int | None = Field(None, description="固定 RPM 限制NULL=自适应)")
effective_limit: int | None = Field(None, description="当前有效限制")
learned_rpm_limit: int | None = Field(None, description="学习到的 RPM 限制")
concurrent_429_count: int
rpm_429_count: int
@@ -76,12 +75,12 @@ class KeyListItem(BaseModel):
@router.get(
"/keys",
response_model=List[KeyListItem],
response_model=list[KeyListItem],
summary="获取所有启用自适应模式的Key",
)
async def list_adaptive_keys(
request: Request,
provider_id: Optional[str] = Query(None, description="按 Provider 过滤"),
provider_id: str | None = Query(None, description="按 Provider 过滤"),
db: Session = Depends(get_db),
):
"""
@@ -207,7 +206,7 @@ async def get_adaptive_summary(
@dataclass
class ListAdaptiveKeysAdapter(AdminApiAdapter):
provider_id: Optional[str] = None
provider_id: str | None = None
async def handle(self, context): # type: ignore[override]
# 自适应模式rpm_limit = NULL

View File

@@ -5,7 +5,6 @@
import os
from datetime import datetime, timedelta, timezone
from typing import Optional
from zoneinfo import ZoneInfo
from fastapi import APIRouter, Depends, HTTPException, Query, Request
@@ -25,7 +24,7 @@ from src.services.user.apikey import ApiKeyService
APP_TIMEZONE = ZoneInfo(os.getenv("APP_TIMEZONE", "Asia/Shanghai"))
def parse_expiry_date(date_str: Optional[str]) -> Optional[datetime]:
def parse_expiry_date(date_str: str | None) -> datetime | None:
"""解析过期日期字符串为 datetime 对象。
Args:
@@ -70,7 +69,7 @@ async def list_standalone_api_keys(
request: Request,
skip: int = Query(0, ge=0),
limit: int = Query(100, ge=1, le=500),
is_active: Optional[bool] = None,
is_active: bool | None = None,
db: Session = Depends(get_db),
):
"""
@@ -330,7 +329,7 @@ class AdminListStandaloneKeysAdapter(AdminApiAdapter):
self,
skip: int,
limit: int,
is_active: Optional[bool],
is_active: bool | None,
):
self.skip = skip
self.limit = limit

View File

@@ -5,7 +5,6 @@ Endpoint 健康监控 API
from collections import defaultdict
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from typing import Dict, List, Optional
from fastapi import APIRouter, Depends, Query, Request
from sqlalchemy import func
@@ -128,7 +127,7 @@ async def get_api_format_health_monitor(
async def get_key_health(
key_id: str,
request: Request,
api_format: Optional[str] = Query(None, description="API 格式(可选,如 CLAUDE、OPENAI"),
api_format: str | None = Query(None, description="API 格式(可选,如 CLAUDE、OPENAI"),
db: Session = Depends(get_db),
) -> HealthStatusResponse:
"""
@@ -161,7 +160,7 @@ async def get_key_health(
async def recover_key_health(
key_id: str,
request: Request,
api_format: Optional[str] = Query(None, description="API 格式(可选,不指定则恢复所有格式)"),
api_format: str | None = Query(None, description="API 格式(可选,不指定则恢复所有格式)"),
db: Session = Depends(get_db),
) -> dict:
"""
@@ -278,7 +277,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
)
# 构建所有格式的 provider_count 映射
all_formats: Dict[str, int] = {}
all_formats: dict[str, int] = {}
for api_format_enum, provider_count in active_formats:
api_format = (
api_format_enum.value if hasattr(api_format_enum, "value") else str(api_format_enum)
@@ -295,7 +294,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
)
.all()
)
endpoint_map: Dict[str, List[str]] = defaultdict(list)
endpoint_map: dict[str, list[str]] = defaultdict(list)
active_provider_formats: set[tuple[str, str]] = set()
for api_format_enum, endpoint_id, provider_id in endpoint_rows:
api_format = (
@@ -305,7 +304,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
active_provider_formats.add((str(provider_id), api_format))
# 1.2 统计每个 API 格式可用的活跃 Key 数量Key 属于 Provider通过 api_formats 关联格式)
key_counts: Dict[str, int] = {}
key_counts: dict[str, int] = {}
if active_provider_formats:
active_provider_keys = (
db.query(ProviderAPIKey.provider_id, ProviderAPIKey.api_formats)
@@ -342,7 +341,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
)
# 构建每个格式的状态统计
status_counts: Dict[str, Dict[str, int]] = {}
status_counts: dict[str, dict[str, int]] = {}
for api_format_enum, status, count in status_counts_query:
api_format = (
api_format_enum.value if hasattr(api_format_enum, "value") else str(api_format_enum)
@@ -370,7 +369,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
.all()
)
grouped_attempts: Dict[str, List[RequestCandidate]] = {}
grouped_attempts: dict[str, list[RequestCandidate]] = {}
for attempt, api_format_enum, provider_id in rows:
api_format = (
@@ -384,7 +383,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
grouped_attempts[api_format].append(attempt)
# 4. 为所有活跃格式生成监控数据(包括没有请求记录的)
monitors: List[ApiFormatHealthMonitor] = []
monitors: list[ApiFormatHealthMonitor] = []
for api_format in all_formats:
attempts = grouped_attempts.get(api_format, [])
# 获取窗口内的真实统计数据
@@ -399,7 +398,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
# 时间线按时间正序
attempts_sorted = list(reversed(attempts))
events: List[EndpointHealthEvent] = []
events: list[EndpointHealthEvent] = []
for attempt in attempts_sorted:
event_timestamp = attempt.finished_at or attempt.started_at or attempt.created_at
events.append(
@@ -462,7 +461,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
@dataclass
class AdminKeyHealthAdapter(AdminApiAdapter):
key_id: str
api_format: Optional[str] = None
api_format: str | None = None
async def handle(self, context): # type: ignore[override]
health_data = health_monitor.get_key_health(context.db, self.key_id, self.api_format)
@@ -500,7 +499,7 @@ class AdminKeyHealthAdapter(AdminApiAdapter):
@dataclass
class AdminRecoverKeyHealthAdapter(AdminApiAdapter):
key_id: str
api_format: Optional[str] = None
api_format: str | None = None
async def handle(self, context): # type: ignore[override]
db = context.db

View File

@@ -6,7 +6,6 @@ import json
import uuid
from dataclasses import dataclass
from datetime import datetime, timezone
from typing import Dict, List
from fastapi import APIRouter, Depends, Query, Request
from sqlalchemy.orm import Session
@@ -146,14 +145,14 @@ async def delete_endpoint_key(
# ========== Provider Keys API ==========
@router.get("/providers/{provider_id}/keys", response_model=List[EndpointAPIKeyResponse])
@router.get("/providers/{provider_id}/keys", response_model=list[EndpointAPIKeyResponse])
async def list_provider_keys(
provider_id: str,
request: Request,
skip: int = Query(0, ge=0, description="跳过的记录数"),
limit: int = Query(100, ge=1, le=1000, description="返回的最大记录数"),
db: Session = Depends(get_db),
) -> List[EndpointAPIKeyResponse]:
) -> list[EndpointAPIKeyResponse]:
"""
获取 Provider 的所有 Keys
@@ -503,12 +502,12 @@ class AdminGetKeysGroupedByFormatAdapter(AdminApiAdapter):
)
.all()
)
endpoint_base_url_map: Dict[tuple[str, str], str] = {}
endpoint_base_url_map: dict[tuple[str, str], str] = {}
for provider_id, api_format, base_url in endpoints:
fmt = api_format.value if hasattr(api_format, "value") else str(api_format)
endpoint_base_url_map[(str(provider_id), fmt)] = base_url
grouped: Dict[str, List[dict]] = {}
grouped: dict[str, list[dict]] = {}
for key, provider in keys:
api_formats = key.api_formats or []

View File

@@ -5,10 +5,9 @@ ProviderEndpoint CRUD 管理 API
import uuid
from dataclasses import dataclass
from datetime import datetime, timezone
from typing import List, Optional
from fastapi import APIRouter, Depends, Query, Request
from sqlalchemy import and_, func
from sqlalchemy import and_
from sqlalchemy.orm import Session
from sqlalchemy.orm.attributes import flag_modified
@@ -29,7 +28,7 @@ router = APIRouter(tags=["Endpoint Management"])
pipeline = ApiRequestPipeline()
def mask_proxy_password(proxy_config: Optional[dict]) -> Optional[dict]:
def mask_proxy_password(proxy_config: dict | None) -> dict | None:
"""对代理配置中的密码进行脱敏处理"""
if not proxy_config:
return None
@@ -39,14 +38,14 @@ def mask_proxy_password(proxy_config: Optional[dict]) -> Optional[dict]:
return masked
@router.get("/providers/{provider_id}/endpoints", response_model=List[ProviderEndpointResponse])
@router.get("/providers/{provider_id}/endpoints", response_model=list[ProviderEndpointResponse])
async def list_provider_endpoints(
provider_id: str,
request: Request,
skip: int = Query(0, ge=0, description="跳过的记录数"),
limit: int = Query(100, ge=1, le=1000, description="返回的最大记录数"),
db: Session = Depends(get_db),
) -> List[ProviderEndpointResponse]:
) -> list[ProviderEndpointResponse]:
"""
获取指定 Provider 的所有 Endpoints
@@ -245,7 +244,7 @@ class AdminListProviderEndpointsAdapter(AdminApiAdapter):
if is_active:
active_keys_map[fmt] = active_keys_map.get(fmt, 0) + 1
result: List[ProviderEndpointResponse] = []
result: list[ProviderEndpointResponse] = []
for endpoint in endpoints:
endpoint_format = (
endpoint.api_format

View File

@@ -1,7 +1,7 @@
"""LDAP配置管理API端点。"""
import re
from typing import Any, Dict, Optional
from typing import Any
from fastapi import APIRouter, Depends, Request
from pydantic import BaseModel, Field, ValidationError, field_validator
@@ -30,9 +30,9 @@ BCRYPT_HASH_PATTERN = re.compile(r"^\$2[aby]\$\d{2}\$.{53}$")
class LDAPConfigResponse(BaseModel):
"""LDAP配置响应不返回密码"""
server_url: Optional[str] = None
bind_dn: Optional[str] = None
base_dn: Optional[str] = None
server_url: str | None = None
bind_dn: str | None = None
base_dn: str | None = None
has_bind_password: bool = False
user_search_filter: str
username_attr: str
@@ -50,7 +50,7 @@ class LDAPConfigUpdate(BaseModel):
server_url: str = Field(..., min_length=1, max_length=255)
bind_dn: str = Field(..., min_length=1, max_length=255)
# 允许空字符串表示"清除密码";非空时自动 strip 并校验不能为空
bind_password: Optional[str] = Field(None, max_length=1024)
bind_password: str | None = Field(None, max_length=1024)
base_dn: str = Field(..., min_length=1, max_length=255)
user_search_filter: str = Field(default="(uid={username})", max_length=500)
username_attr: str = Field(default="uid", max_length=50)
@@ -63,7 +63,7 @@ class LDAPConfigUpdate(BaseModel):
@field_validator("bind_password")
@classmethod
def validate_bind_password(cls, v: Optional[str]) -> Optional[str]:
def validate_bind_password(cls, v: str | None) -> str | None:
if v is None or v == "":
return v
v = v.strip()
@@ -114,22 +114,22 @@ class LDAPTestResponse(BaseModel):
class LDAPConfigTest(BaseModel):
"""LDAP配置测试请求全部可选用于临时覆盖"""
server_url: Optional[str] = Field(None, min_length=1, max_length=255)
bind_dn: Optional[str] = Field(None, min_length=1, max_length=255)
bind_password: Optional[str] = Field(None, min_length=1)
base_dn: Optional[str] = Field(None, min_length=1, max_length=255)
user_search_filter: Optional[str] = Field(None, max_length=500)
username_attr: Optional[str] = Field(None, max_length=50)
email_attr: Optional[str] = Field(None, max_length=50)
display_name_attr: Optional[str] = Field(None, max_length=50)
is_enabled: Optional[bool] = None
is_exclusive: Optional[bool] = None
use_starttls: Optional[bool] = None
connect_timeout: Optional[int] = Field(None, ge=1, le=60)
server_url: str | None = Field(None, min_length=1, max_length=255)
bind_dn: str | None = Field(None, min_length=1, max_length=255)
bind_password: str | None = Field(None, min_length=1)
base_dn: str | None = Field(None, min_length=1, max_length=255)
user_search_filter: str | None = Field(None, max_length=500)
username_attr: str | None = Field(None, max_length=50)
email_attr: str | None = Field(None, max_length=50)
display_name_attr: str | None = Field(None, max_length=50)
is_enabled: bool | None = None
is_exclusive: bool | None = None
use_starttls: bool | None = None
connect_timeout: int | None = Field(None, ge=1, le=60)
@field_validator("user_search_filter")
@classmethod
def validate_search_filter(cls, v: Optional[str]) -> Optional[str]:
def validate_search_filter(cls, v: str | None) -> str | None:
if v is None:
return v
if "{username}" not in v:
@@ -263,7 +263,7 @@ async def test_ldap_connection(request: Request, db: Session = Depends(get_db))
class AdminGetLDAPConfigAdapter(AdminApiAdapter):
async def handle(self, context) -> Dict[str, Any]: # type: ignore[override]
async def handle(self, context) -> dict[str, Any]: # type: ignore[override]
db = context.db
config = db.query(LDAPConfig).first()
@@ -300,7 +300,7 @@ class AdminGetLDAPConfigAdapter(AdminApiAdapter):
class AdminUpdateLDAPConfigAdapter(AdminApiAdapter):
async def handle(self, context) -> Dict[str, str]: # type: ignore[override]
async def handle(self, context) -> dict[str, str]: # type: ignore[override]
db = context.db
payload = context.ensure_json_body()
@@ -421,7 +421,7 @@ class AdminUpdateLDAPConfigAdapter(AdminApiAdapter):
class AdminTestLDAPConnectionAdapter(AdminApiAdapter):
async def handle(self, context) -> Dict[str, Any]: # type: ignore[override]
async def handle(self, context) -> dict[str, Any]: # type: ignore[override]
from src.services.auth.ldap import LDAPService
db = context.db
@@ -442,7 +442,7 @@ class AdminTestLDAPConnectionAdapter(AdminApiAdapter):
raise InvalidRequestException(translate_pydantic_error(errors[0]))
raise InvalidRequestException("请求数据验证失败")
config_data: Dict[str, Any] = {}
config_data: dict[str, Any] = {}
if saved_config:
config_data = {

View File

@@ -1,7 +1,6 @@
"""管理员 Management Token 管理端点"""
from dataclasses import dataclass
from typing import Optional
from fastapi import APIRouter, Depends, HTTPException, Query, Request
from fastapi.responses import JSONResponse
@@ -46,8 +45,8 @@ class AdminManagementTokenApiAdapter(AdminApiAdapter):
@router.get("")
async def list_all_management_tokens(
request: Request,
user_id: Optional[str] = Query(None, description="筛选用户 ID"),
is_active: Optional[bool] = Query(None, description="筛选激活状态"),
user_id: str | None = Query(None, description="筛选用户 ID"),
is_active: bool | None = Query(None, description="筛选激活状态"),
skip: int = Query(0, ge=0),
limit: int = Query(50, ge=1, le=100),
db: Session = Depends(get_db),
@@ -174,8 +173,8 @@ class AdminListManagementTokensAdapter(AdminManagementTokenApiAdapter):
"""列出所有 Management Tokens"""
name: str = "admin_list_management_tokens"
user_id: Optional[str] = None
is_active: Optional[bool] = None
user_id: str | None = None
is_active: bool | None = None
skip: int = 0
limit: int = 50
@@ -197,7 +196,7 @@ class AdminListManagementTokensAdapter(AdminManagementTokenApiAdapter):
)
# 预加载用户信息
user_ids = list(set(t.user_id for t in tokens))
user_ids = list({t.user_id for t in tokens})
users = {u.id: u for u in context.db.query(User).filter(User.id.in_(user_ids)).all()}
for token in tokens:
token.user = users.get(token.user_id)

View File

@@ -5,7 +5,6 @@
"""
from dataclasses import dataclass
from typing import Dict, List
from fastapi import APIRouter, Depends, Request
from sqlalchemy.orm import Session, joinedload
@@ -64,12 +63,12 @@ class AdminGetModelCatalogAdapter(AdminApiAdapter):
db: Session = context.db
# 1. 获取所有活跃的 GlobalModel
global_models: List[GlobalModel] = (
global_models: list[GlobalModel] = (
db.query(GlobalModel).filter(GlobalModel.is_active == True).all()
)
# 2. 获取所有活跃的 Model 实现(包含 global_model 以便计算有效价格)
models: List[Model] = (
models: list[Model] = (
db.query(Model)
.options(joinedload(Model.provider), joinedload(Model.global_model))
.filter(Model.is_active == True)
@@ -77,17 +76,17 @@ class AdminGetModelCatalogAdapter(AdminApiAdapter):
)
# 按 GlobalModel ID 组织关联提供商
models_by_global_model: Dict[str, List[Model]] = {}
models_by_global_model: dict[str, list[Model]] = {}
for model in models:
if model.global_model_id:
models_by_global_model.setdefault(model.global_model_id, []).append(model)
# 3. 为每个 GlobalModel 构建 catalog item
catalog_items: List[ModelCatalogItem] = []
catalog_items: list[ModelCatalogItem] = []
for gm in global_models:
gm_id = gm.id
provider_entries: List[ModelCatalogProviderDetail] = []
provider_entries: list[ModelCatalogProviderDetail] = []
# 从 config JSON 读取能力标志
gm_config = gm.config or {}
capability_flags = {

View File

@@ -3,8 +3,7 @@ models.dev 外部模型数据代理
"""
import json
from typing import Any, Optional
from typing import Any
import httpx
from fastapi import APIRouter, Depends, HTTPException
@@ -43,7 +42,7 @@ OFFICIAL_PROVIDERS = {
}
async def _get_cached_data() -> Optional[dict[str, Any]]:
async def _get_cached_data() -> dict[str, Any] | None:
"""从 Redis 获取缓存数据"""
redis = await get_redis_client()
if redis is None:

View File

@@ -5,7 +5,6 @@ GlobalModel Admin API
"""
from dataclasses import dataclass
from typing import Optional
from fastapi import APIRouter, Depends, Query, Request
from sqlalchemy.orm import Session
@@ -37,8 +36,8 @@ async def list_global_models(
request: Request,
skip: int = Query(0, ge=0),
limit: int = Query(100, ge=1, le=1000),
is_active: Optional[bool] = Query(None),
search: Optional[str] = Query(None),
is_active: bool | None = Query(None),
search: str | None = Query(None),
db: Session = Depends(get_db),
) -> GlobalModelListResponse:
"""
@@ -254,8 +253,8 @@ class AdminListGlobalModelsAdapter(AdminApiAdapter):
skip: int
limit: int
is_active: Optional[bool]
search: Optional[str]
is_active: bool | None
search: str | None
async def handle(self, context): # type: ignore[override]
from sqlalchemy import func

View File

@@ -9,7 +9,6 @@ GlobalModel 请求链路预览 API
"""
from dataclasses import dataclass
from typing import Dict, List, Optional
from fastapi import APIRouter, Depends, Request
from pydantic import BaseModel, ConfigDict, Field
@@ -47,22 +46,22 @@ class RoutingKeyInfo(BaseModel):
name: str
masked_key: str = Field("", description="脱敏的 API Key")
internal_priority: int = Field(..., description="Key 内部优先级")
global_priority_by_format: Optional[Dict[str, int]] = Field(None, description="按 API 格式的全局优先级")
rpm_limit: Optional[int] = Field(None, description="RPM 限制null 表示自适应")
global_priority_by_format: dict[str, int] | None = Field(None, description="按 API 格式的全局优先级")
rpm_limit: int | None = Field(None, description="RPM 限制null 表示自适应")
is_adaptive: bool = Field(False, description="是否为自适应 RPM 模式")
effective_rpm: Optional[int] = Field(None, description="有效 RPM 限制")
effective_rpm: int | None = Field(None, description="有效 RPM 限制")
cache_ttl_minutes: int = Field(0, description="缓存 TTL分钟")
health_score: float = Field(1.0, description="健康度分数0-1 小数格式)")
is_active: bool
api_formats: List[str] = Field(default_factory=list, description="支持的 API 格式")
api_formats: list[str] = Field(default_factory=list, description="支持的 API 格式")
# 模型白名单
allowed_models: Optional[List[str]] = Field(None, description="允许的模型列表null 表示不限制")
allowed_models: list[str] | None = Field(None, description="允许的模型列表null 表示不限制")
# 熔断状态
circuit_breaker_open: bool = Field(False, description="熔断器是否打开")
circuit_breaker_formats: List[str] = Field(
circuit_breaker_formats: list[str] = Field(
default_factory=list, description="熔断的 API 格式列表"
)
next_probe_at: Optional[str] = Field(None, description="下次探测时间ISO格式")
next_probe_at: str | None = Field(None, description="下次探测时间ISO格式")
model_config = ConfigDict(from_attributes=True)
@@ -73,9 +72,9 @@ class RoutingEndpointInfo(BaseModel):
id: str
api_format: str
base_url: str
custom_path: Optional[str] = None
custom_path: str | None = None
is_active: bool
keys: List[RoutingKeyInfo] = Field(default_factory=list)
keys: list[RoutingKeyInfo] = Field(default_factory=list)
total_keys: int = 0
active_keys: int = 0
@@ -87,7 +86,7 @@ class RoutingModelMapping(BaseModel):
name: str = Field(..., description="映射名称")
priority: int = Field(..., description="优先级(数字越小优先级越高)")
api_formats: Optional[List[str]] = Field(None, description="作用域(适用的 API 格式)")
api_formats: list[str] | None = Field(None, description="作用域(适用的 API 格式)")
class RoutingProviderInfo(BaseModel):
@@ -97,18 +96,18 @@ class RoutingProviderInfo(BaseModel):
name: str
model_id: str = Field(..., description="Model IDGlobalModel 与 Provider 的关联记录 ID")
provider_priority: int = Field(..., description="提供商优先级(数字越小优先级越高)")
billing_type: Optional[str] = Field(None, description="计费类型")
monthly_quota_usd: Optional[float] = Field(None, description="月额度(美元)")
monthly_used_usd: Optional[float] = Field(None, description="已用额度(美元)")
billing_type: str | None = Field(None, description="计费类型")
monthly_quota_usd: float | None = Field(None, description="月额度(美元)")
monthly_used_usd: float | None = Field(None, description="已用额度(美元)")
is_active: bool
# 模型映射信息
provider_model_name: str = Field(..., description="提供商侧的模型名称")
model_mappings: List[RoutingModelMapping] = Field(
model_mappings: list[RoutingModelMapping] = Field(
default_factory=list, description="模型名称映射列表"
)
model_is_active: bool = Field(True, description="Model 是否活跃")
# Endpoint 和 Key 信息
endpoints: List[RoutingEndpointInfo] = Field(default_factory=list)
endpoints: list[RoutingEndpointInfo] = Field(default_factory=list)
total_endpoints: int = 0
active_endpoints: int = 0
@@ -123,7 +122,7 @@ class GlobalKeyWhitelistItem(BaseModel):
masked_key: str = Field(..., description="脱敏的 API Key")
provider_id: str = Field(..., description="Provider ID")
provider_name: str = Field(..., description="Provider 名称")
allowed_models: List[str] = Field(default_factory=list, description="Key 白名单模型列表")
allowed_models: list[str] = Field(default_factory=list, description="Key 白名单模型列表")
model_config = ConfigDict(from_attributes=True)
@@ -136,11 +135,11 @@ class ModelRoutingPreviewResponse(BaseModel):
display_name: str
is_active: bool
# GlobalModel 的模型映射(用于前端匹配 Key 白名单)
global_model_mappings: List[str] = Field(
global_model_mappings: list[str] = Field(
default_factory=list, description="GlobalModel 的模型映射规则(正则模式)"
)
# 链路信息
providers: List[RoutingProviderInfo] = Field(
providers: list[RoutingProviderInfo] = Field(
default_factory=list, description="按优先级排序的提供商列表"
)
total_providers: int = 0
@@ -149,7 +148,7 @@ class ModelRoutingPreviewResponse(BaseModel):
scheduling_mode: str = Field("cache_affinity", description="调度模式")
priority_mode: str = Field("provider", description="优先级模式")
# 全局 Key 白名单数据(供前端实时匹配,包含所有 Provider 的 Key
all_keys_whitelist: List[GlobalKeyWhitelistItem] = Field(
all_keys_whitelist: list[GlobalKeyWhitelistItem] = Field(
default_factory=list, description="所有 Provider 的 Key 白名单数据"
)
@@ -227,7 +226,7 @@ class AdminGetModelRoutingPreviewAdapter(AdminApiAdapter):
provider_ids = [m.provider_id for m in models if m.provider_id]
# 批量获取 Provider 的 Endpoints
endpoints_by_provider: Dict[str, List[ProviderEndpoint]] = {}
endpoints_by_provider: dict[str, list[ProviderEndpoint]] = {}
if provider_ids:
endpoints = (
db.query(ProviderEndpoint)
@@ -240,7 +239,7 @@ class AdminGetModelRoutingPreviewAdapter(AdminApiAdapter):
endpoints_by_provider[ep.provider_id].append(ep)
# 批量获取 Provider 的 Keys
keys_by_provider: Dict[str, List[ProviderAPIKey]] = {}
keys_by_provider: dict[str, list[ProviderAPIKey]] = {}
if provider_ids:
keys = (
db.query(ProviderAPIKey).filter(ProviderAPIKey.provider_id.in_(provider_ids)).all()
@@ -251,14 +250,14 @@ class AdminGetModelRoutingPreviewAdapter(AdminApiAdapter):
keys_by_provider[key.provider_id].append(key)
# 提取 GlobalModel 的 model_mappings用于 Key 白名单匹配)
global_model_mappings: List[str] = []
global_model_mappings: list[str] = []
if global_model.config and isinstance(global_model.config, dict):
mappings = global_model.config.get("model_mappings")
if isinstance(mappings, list):
global_model_mappings = [m for m in mappings if isinstance(m, str)]
# 构建 Provider 路由信息
provider_infos: List[RoutingProviderInfo] = []
provider_infos: list[RoutingProviderInfo] = []
for model in models:
provider = model.provider
if not provider:
@@ -281,7 +280,7 @@ class AdminGetModelRoutingPreviewAdapter(AdminApiAdapter):
provider_keys = keys_by_provider.get(provider.id, [])
# 按 api_format 组织 Keys
keys_by_endpoint: Dict[str, List[ProviderAPIKey]] = {}
keys_by_endpoint: dict[str, list[ProviderAPIKey]] = {}
for key in provider_keys:
# 每个 Key 可能支持多个 api_formats
for fmt in key.api_formats or []:
@@ -355,8 +354,8 @@ class AdminGetModelRoutingPreviewAdapter(AdminApiAdapter):
# 检查熔断状态
circuit_breaker_open = False
circuit_breaker_formats: List[str] = []
next_probe_at: Optional[str] = None
circuit_breaker_formats: list[str] = []
next_probe_at: str | None = None
if key.circuit_breaker_by_format:
for fmt, cb_state in key.circuit_breaker_by_format.items():
if isinstance(cb_state, dict) and cb_state.get("open"):
@@ -462,7 +461,7 @@ class AdminGetModelRoutingPreviewAdapter(AdminApiAdapter):
)
# 获取所有活跃 Provider 的 Key 白名单数据(供前端实时匹配)
all_keys_whitelist: List[GlobalKeyWhitelistItem] = []
all_keys_whitelist: list[GlobalKeyWhitelistItem] = []
crypto = CryptoService()
# 获取所有活跃的 Key带白名单使用 selectinload 避免 N+1 查询

View File

@@ -1,7 +1,8 @@
"""模块管理 API 端点"""
from __future__ import annotations
from dataclasses import dataclass
from typing import Any, Dict, Optional
from typing import Any
from fastapi import APIRouter, Depends, Request
from pydantic import BaseModel
@@ -28,18 +29,18 @@ class ModuleStatusResponse(BaseModel):
enabled: bool
active: bool
config_validated: bool
config_error: Optional[str]
config_error: str | None
display_name: str
description: str
category: str
admin_route: Optional[str]
admin_menu_icon: Optional[str]
admin_menu_group: Optional[str]
admin_route: str | None
admin_menu_icon: str | None
admin_menu_group: str | None
admin_menu_order: int
health: str
@classmethod
def from_status(cls, status: ModuleStatus) -> "ModuleStatusResponse":
def from_status(cls, status: ModuleStatus) -> ModuleStatusResponse:
return cls(
name=status.name,
available=status.available,
@@ -130,7 +131,7 @@ async def set_module_enabled(
class AdminGetAllModulesStatusAdapter(AdminApiAdapter):
"""获取所有模块状态"""
async def handle(self, context) -> Dict[str, Any]:
async def handle(self, context) -> dict[str, Any]:
registry = get_module_registry()
all_status = await registry.get_all_status_async(context.db)
@@ -146,7 +147,7 @@ class AdminGetModuleStatusAdapter(AdminApiAdapter):
module_name: str
async def handle(self, context) -> Dict[str, Any]:
async def handle(self, context) -> dict[str, Any]:
registry = get_module_registry()
status = await registry.get_module_status_async(self.module_name, context.db)
@@ -162,7 +163,7 @@ class AdminSetModuleEnabledAdapter(AdminApiAdapter):
module_name: str
async def handle(self, context) -> Dict[str, Any]:
async def handle(self, context) -> dict[str, Any]:
registry = get_module_registry()
# 检查模块是否存在

View File

@@ -2,7 +2,6 @@
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from typing import List, Optional
from fastapi import APIRouter, Depends, HTTPException, Query, Request
from sqlalchemy import func
@@ -33,8 +32,8 @@ pipeline = ApiRequestPipeline()
@router.get("/audit-logs")
async def get_audit_logs(
request: Request,
username: Optional[str] = Query(None, description="用户名筛选 (模糊匹配)"),
event_type: Optional[str] = Query(None, description="事件类型筛选"),
username: str | None = Query(None, description="用户名筛选 (模糊匹配)"),
event_type: str | None = Query(None, description="事件类型筛选"),
days: int = Query(7, description="查询天数"),
limit: int = Query(100, description="返回数量限制"),
offset: int = Query(0, description="偏移量"),
@@ -212,8 +211,8 @@ async def get_circuit_history(
@dataclass
class AdminGetAuditLogsAdapter(AdminApiAdapter):
username: Optional[str]
event_type: Optional[str]
username: str | None
event_type: str | None
days: int
limit: int
offset: int
@@ -497,8 +496,8 @@ class AdminCircuitHistoryAdapter(AdminApiAdapter):
return {"items": history, "count": len(history)}
def _get_health_recommendations(error_stats: dict, health_score: int) -> List[str]:
recommendations: List[str] = []
def _get_health_recommendations(error_stats: dict, health_score: int) -> list[str]:
recommendations: list[str] = []
if health_score < 50:
recommendations.append("系统健康状况严重,请立即检查错误日志")
if error_stats.get("total_errors", 0) > 100:

View File

@@ -5,7 +5,7 @@
"""
from dataclasses import dataclass
from typing import Any, Dict, List, Optional, Tuple
from typing import Any
from fastapi import APIRouter, Depends, HTTPException, Query, Request
from fastapi.responses import PlainTextResponse
@@ -13,7 +13,7 @@ from sqlalchemy.orm import Session
from src.api.base.admin_adapter import AdminApiAdapter
from src.api.base.context import ApiRequestContext
from src.api.base.pagination import PaginationMeta, build_pagination_payload, paginate_sequence
from src.api.base.pagination import build_pagination_payload, paginate_sequence
from src.api.base.pipeline import ApiRequestPipeline
from src.clients.redis_client import get_redis_client_sync
from src.core.crypto import crypto_service
@@ -28,7 +28,7 @@ router = APIRouter(prefix="/api/admin/monitoring/cache", tags=["Admin - Monitori
pipeline = ApiRequestPipeline()
def mask_api_key(api_key: Optional[str], prefix_len: int = 8, suffix_len: int = 4) -> Optional[str]:
def mask_api_key(api_key: str | None, prefix_len: int = 8, suffix_len: int = 4) -> str | None:
"""
脱敏 API Key显示前缀 + 星号 + 后缀
例如: sk-jhiId-xxxxxxxxxxxAABB -> sk-jhiId-********AABB
@@ -47,7 +47,7 @@ def mask_api_key(api_key: Optional[str], prefix_len: int = 8, suffix_len: int =
return f"{api_key[:prefix_len]}********{api_key[-suffix_len:]}"
def decrypt_and_mask(encrypted_key: Optional[str], prefix_len: int = 8) -> Optional[str]:
def decrypt_and_mask(encrypted_key: str | None, prefix_len: int = 8) -> str | None:
"""
解密 API Key 后脱敏显示
@@ -65,7 +65,7 @@ def decrypt_and_mask(encrypted_key: Optional[str], prefix_len: int = 8) -> Optio
return None
def resolve_user_identifier(db: Session, identifier: str) -> Optional[str]:
def resolve_user_identifier(db: Session, identifier: str) -> str | None:
"""
将用户标识符username/email/user_id/api_key_id解析为 user_id
@@ -181,7 +181,7 @@ async def get_user_affinity(
@router.get("/affinities")
async def list_affinities(
request: Request,
keyword: Optional[str] = None,
keyword: str | None = None,
limit: int = Query(100, ge=1, le=1000, description="返回数量限制"),
offset: int = Query(0, ge=0, description="偏移量"),
db: Session = Depends(get_db),
@@ -421,7 +421,7 @@ async def get_cache_metrics(
class AdminCacheStatsAdapter(AdminApiAdapter):
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
try:
redis_client = get_redis_client_sync()
# 读取系统配置,确保监控接口与编排器使用一致的模式
@@ -487,14 +487,14 @@ class AdminCacheMetricsAdapter(AdminApiAdapter):
logger.exception(f"导出缓存指标失败: {exc}")
raise HTTPException(status_code=500, detail=f"导出缓存指标失败: {exc}")
def _format_prometheus(self, stats: Dict[str, Any]) -> str:
def _format_prometheus(self, stats: dict[str, Any]) -> str:
"""
将 scheduler/affinity 指标转换为 Prometheus 文本格式。
"""
scheduler_metrics = stats.get("scheduler_metrics", {})
affinity_stats = stats.get("affinity_stats", {})
metric_map: List[Tuple[str, str, float]] = [
metric_map: list[tuple[str, str, float]] = [
(
"cache_scheduler_total_batches",
"Total batches pulled from provider list",
@@ -542,7 +542,7 @@ class AdminCacheMetricsAdapter(AdminApiAdapter):
),
]
affinity_map: List[Tuple[str, str, float]] = [
affinity_map: list[tuple[str, str, float]] = [
(
"cache_affinity_total",
"Total cache affinities stored",
@@ -596,7 +596,7 @@ class AdminCacheMetricsAdapter(AdminApiAdapter):
class AdminGetUserAffinityAdapter(AdminApiAdapter):
user_identifier: str
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
db = context.db
try:
user_id = resolve_user_identifier(db, self.user_identifier)
@@ -673,11 +673,11 @@ class AdminGetUserAffinityAdapter(AdminApiAdapter):
@dataclass
class AdminListAffinitiesAdapter(AdminApiAdapter):
keyword: Optional[str]
keyword: str | None
limit: int
offset: int
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
db = context.db
redis_client = get_redis_client_sync()
if not redis_client:
@@ -686,7 +686,7 @@ class AdminListAffinitiesAdapter(AdminApiAdapter):
affinity_mgr = await get_affinity_manager(redis_client)
matched_user_id = None
matched_api_key_id = None
raw_affinities: List[Dict[str, Any]] = []
raw_affinities: list[dict[str, Any]] = []
if self.keyword:
# 首先检查是否是 API Key IDaffinity_key
@@ -724,14 +724,14 @@ class AdminListAffinitiesAdapter(AdminApiAdapter):
}
# 批量查询用户 API Key 信息
user_api_key_map: Dict[str, ApiKey] = {}
user_api_key_map: dict[str, ApiKey] = {}
if affinity_keys:
user_api_keys = db.query(ApiKey).filter(ApiKey.id.in_(list(affinity_keys))).all()
user_api_key_map = {str(k.id): k for k in user_api_keys}
# 收集所有 user_id
user_ids = {str(k.user_id) for k in user_api_key_map.values()}
user_map: Dict[str, User] = {}
user_map: dict[str, User] = {}
if user_ids:
users = db.query(User).filter(User.id.in_(list(user_ids))).all()
user_map = {str(user.id): user for user in users}
@@ -771,7 +771,7 @@ class AdminListAffinitiesAdapter(AdminApiAdapter):
global_model_ids = {
item.get("model_name") for item in raw_affinities if item.get("model_name")
}
global_model_map: Dict[str, GlobalModel] = {}
global_model_map: dict[str, GlobalModel] = {}
if global_model_ids:
# model_name 可能是 UUID 格式的 global_model_id也可能是原始模型名称
global_models = db.query(GlobalModel).filter(
@@ -885,7 +885,7 @@ class AdminListAffinitiesAdapter(AdminApiAdapter):
class AdminClearUserCacheAdapter(AdminApiAdapter):
user_identifier: str
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
db = context.db
try:
redis_client = get_redis_client_sync()
@@ -995,7 +995,7 @@ class AdminClearSingleAffinityAdapter(AdminApiAdapter):
model_id: str
api_format: str
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
db = context.db
try:
redis_client = get_redis_client_sync()
@@ -1048,7 +1048,7 @@ class AdminClearSingleAffinityAdapter(AdminApiAdapter):
class AdminClearAllCacheAdapter(AdminApiAdapter):
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
try:
redis_client = get_redis_client_sync()
affinity_mgr = await get_affinity_manager(redis_client)
@@ -1068,7 +1068,7 @@ class AdminClearAllCacheAdapter(AdminApiAdapter):
class AdminClearProviderCacheAdapter(AdminApiAdapter):
provider_id: str
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
try:
redis_client = get_redis_client_sync()
affinity_mgr = await get_affinity_manager(redis_client)
@@ -1091,7 +1091,7 @@ class AdminClearProviderCacheAdapter(AdminApiAdapter):
class AdminCacheConfigAdapter(AdminApiAdapter):
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
from src.services.cache.affinity_manager import CacheAffinityManager
from src.config.constants import ConcurrencyDefaults
from src.services.rate_limit.adaptive_reservation import get_adaptive_reservation_manager
@@ -1260,7 +1260,7 @@ async def clear_provider_model_mapping_cache(
class AdminModelMappingCacheStatsAdapter(AdminApiAdapter):
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
import json
from src.clients.redis_client import get_redis_client
@@ -1510,7 +1510,7 @@ class AdminModelMappingCacheStatsAdapter(AdminApiAdapter):
class AdminClearAllModelMappingCacheAdapter(AdminApiAdapter):
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
from src.clients.redis_client import get_redis_client
try:
@@ -1552,7 +1552,7 @@ class AdminClearAllModelMappingCacheAdapter(AdminApiAdapter):
class AdminClearModelMappingCacheByNameAdapter(AdminApiAdapter):
model_name: str
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
from src.clients.redis_client import get_redis_client
try:
@@ -1599,7 +1599,7 @@ class AdminClearProviderModelMappingCacheAdapter(AdminApiAdapter):
provider_id: str
global_model_id: str
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
from src.clients.redis_client import get_redis_client
try:

View File

@@ -4,7 +4,6 @@
from dataclasses import dataclass
from datetime import datetime
from typing import List, Optional
from fastapi import APIRouter, Depends, HTTPException, Query, Request
from pydantic import BaseModel, ConfigDict
@@ -28,29 +27,29 @@ class CandidateResponse(BaseModel):
request_id: str
candidate_index: int
retry_index: int = 0 # 重试序号从0开始
provider_id: Optional[str] = None
provider_name: Optional[str] = None
provider_website: Optional[str] = None # Provider 官网
endpoint_id: Optional[str] = None
endpoint_name: Optional[str] = None # 端点显示名称api_format
key_id: Optional[str] = None
key_name: Optional[str] = None # 密钥名称
key_preview: Optional[str] = None # 密钥脱敏预览(如 sk-***abc
key_capabilities: Optional[dict] = None # Key 支持的能力
required_capabilities: Optional[dict] = None # 请求实际需要的能力标签
provider_id: str | None = None
provider_name: str | None = None
provider_website: str | None = None # Provider 官网
endpoint_id: str | None = None
endpoint_name: str | None = None # 端点显示名称api_format
key_id: str | None = None
key_name: str | None = None # 密钥名称
key_preview: str | None = None # 密钥脱敏预览(如 sk-***abc
key_capabilities: dict | None = None # Key 支持的能力
required_capabilities: dict | None = None # 请求实际需要的能力标签
status: str # 'pending', 'success', 'failed', 'skipped'
skip_reason: Optional[str] = None
skip_reason: str | None = None
is_cached: bool = False
# 执行结果字段
status_code: Optional[int] = None
error_type: Optional[str] = None
error_message: Optional[str] = None
latency_ms: Optional[int] = None
concurrent_requests: Optional[int] = None
extra_data: Optional[dict] = None
status_code: int | None = None
error_type: str | None = None
error_message: str | None = None
latency_ms: int | None = None
concurrent_requests: int | None = None
extra_data: dict | None = None
created_at: datetime
started_at: Optional[datetime] = None
finished_at: Optional[datetime] = None
started_at: datetime | None = None
finished_at: datetime | None = None
model_config = ConfigDict(from_attributes=True)
@@ -62,7 +61,7 @@ class RequestTraceResponse(BaseModel):
total_candidates: int
final_status: str # 'success', 'failed', 'cancelled', 'streaming', 'pending'
total_latency_ms: int
candidates: List[CandidateResponse]
candidates: list[CandidateResponse]
@router.get("/{request_id}", response_model=RequestTraceResponse)
@@ -253,7 +252,7 @@ class AdminGetRequestTraceAdapter(AdminApiAdapter):
key_preview_map[k.id] = "***"
# 构建 candidate 响应列表
candidate_responses: List[CandidateResponse] = []
candidate_responses: list[CandidateResponse] = []
for candidate in candidates:
provider_name = (
provider_map.get(candidate.provider_id) if candidate.provider_id else None

View File

@@ -9,7 +9,7 @@ Provider 操作 API 路由
"""
from dataclasses import asdict, is_dataclass
from typing import Any, Dict, List, Optional
from typing import Any
from fastapi import APIRouter, Depends, HTTPException
from pydantic import BaseModel, Field
@@ -18,9 +18,7 @@ from sqlalchemy.orm import Session
from src.database import get_db
from src.models.database import Provider, User
from src.services.provider_ops import (
ActionStatus,
ConnectorAuthType,
ConnectorStatus,
ProviderActionType,
ProviderOpsConfig,
ProviderOpsService,
@@ -40,46 +38,46 @@ class ArchitectureInfo(BaseModel):
architecture_id: str
display_name: str
description: str
supported_auth_types: List[Dict[str, str]]
supported_actions: List[Dict[str, Any]]
default_connector: Optional[str]
supported_auth_types: list[dict[str, str]]
supported_actions: list[dict[str, Any]]
default_connector: str | None
class ConnectorConfigRequest(BaseModel):
"""连接器配置请求"""
auth_type: str = Field(..., description="认证类型")
config: Dict[str, Any] = Field(default_factory=dict, description="连接器配置")
credentials: Dict[str, Any] = Field(default_factory=dict, description="凭据信息")
config: dict[str, Any] = Field(default_factory=dict, description="连接器配置")
credentials: dict[str, Any] = Field(default_factory=dict, description="凭据信息")
class ActionConfigRequest(BaseModel):
"""操作配置请求"""
enabled: bool = Field(True, description="是否启用")
config: Dict[str, Any] = Field(default_factory=dict, description="操作配置")
config: dict[str, Any] = Field(default_factory=dict, description="操作配置")
class SaveConfigRequest(BaseModel):
"""保存配置请求"""
architecture_id: str = Field("generic_api", description="架构 ID")
base_url: Optional[str] = Field(None, description="API 基础地址")
base_url: str | None = Field(None, description="API 基础地址")
connector: ConnectorConfigRequest
actions: Dict[str, ActionConfigRequest] = Field(default_factory=dict)
schedule: Dict[str, str] = Field(default_factory=dict, description="定时任务配置")
actions: dict[str, ActionConfigRequest] = Field(default_factory=dict)
schedule: dict[str, str] = Field(default_factory=dict, description="定时任务配置")
class ConnectRequest(BaseModel):
"""连接请求"""
credentials: Optional[Dict[str, Any]] = Field(None, description="凭据(可选,使用已保存的)")
credentials: dict[str, Any] | None = Field(None, description="凭据(可选,使用已保存的)")
class ExecuteActionRequest(BaseModel):
"""执行操作请求"""
config: Optional[Dict[str, Any]] = Field(None, description="操作配置(覆盖默认)")
config: dict[str, Any] | None = Field(None, description="操作配置(覆盖默认)")
class ConnectionStatusResponse(BaseModel):
@@ -87,9 +85,9 @@ class ConnectionStatusResponse(BaseModel):
status: str
auth_type: str
connected_at: Optional[str]
expires_at: Optional[str]
last_error: Optional[str]
connected_at: str | None
expires_at: str | None
last_error: str | None
class ActionResultResponse(BaseModel):
@@ -97,10 +95,10 @@ class ActionResultResponse(BaseModel):
status: str
action_type: str
data: Optional[Any]
message: Optional[str]
data: Any | None
message: str | None
executed_at: str
response_time_ms: Optional[int]
response_time_ms: int | None
cache_ttl_seconds: int
@@ -109,9 +107,9 @@ class ProviderOpsStatusResponse(BaseModel):
provider_id: str
is_configured: bool
architecture_id: Optional[str]
architecture_id: str | None
connection_status: ConnectionStatusResponse
enabled_actions: List[str]
enabled_actions: list[str]
class ProviderOpsConfigResponse(BaseModel):
@@ -119,17 +117,17 @@ class ProviderOpsConfigResponse(BaseModel):
provider_id: str
is_configured: bool
architecture_id: Optional[str] = None
base_url: Optional[str] = None
connector: Optional[Dict[str, Any]] = None # 脱敏后的连接器配置
architecture_id: str | None = None
base_url: str | None = None
connector: dict[str, Any] | None = None # 脱敏后的连接器配置
class VerifyAuthResponse(BaseModel):
"""验证认证响应"""
success: bool
message: Optional[str] = None
data: Optional[Dict[str, Any]] = None
message: str | None = None
data: dict[str, Any] | None = None
# ==================== Helper Functions ====================
@@ -147,7 +145,7 @@ def _serialize_data(data: Any) -> Any:
# ==================== Routes ====================
@router.get("/architectures", response_model=List[ArchitectureInfo])
@router.get("/architectures", response_model=list[ArchitectureInfo])
async def list_architectures(_: User = Depends(require_admin)):
"""获取所有可用的架构"""
registry = get_registry()
@@ -490,7 +488,7 @@ async def checkin(
@router.post("/batch/balance")
async def batch_query_balance(
provider_ids: Optional[List[str]] = None,
provider_ids: list[str] | None = None,
db: Session = Depends(get_db),
_: User = Depends(require_admin),
):

View File

@@ -4,7 +4,6 @@ Provider Query API 端点
"""
import asyncio
from typing import Optional
import httpx
from fastapi import APIRouter, Depends, HTTPException
@@ -40,7 +39,7 @@ class ModelsQueryRequest(BaseModel):
"""模型列表查询请求"""
provider_id: str
api_key_id: Optional[str] = None
api_key_id: str | None = None
force_refresh: bool = False # 强制刷新,跳过缓存
@@ -49,11 +48,11 @@ class TestModelRequest(BaseModel):
provider_id: str
model_name: str
api_key_id: Optional[str] = None
endpoint_id: Optional[str] = None # 指定使用的端点ID
api_key_id: str | None = None
endpoint_id: str | None = None # 指定使用的端点ID
stream: bool = False
message: Optional[str] = "你好"
api_format: Optional[str] = None # 指定使用的API格式如果不指定则使用端点的默认格式
message: str | None = "你好"
api_format: str | None = None # 指定使用的API格式如果不指定则使用端点的默认格式
# ============ API Endpoints ============

View File

@@ -3,7 +3,6 @@
"""
from datetime import datetime, timedelta, timezone
from typing import Optional
from fastapi import APIRouter, Depends, HTTPException, Request
from fastapi.responses import JSONResponse
@@ -25,11 +24,11 @@ pipeline = ApiRequestPipeline()
class ProviderBillingUpdate(BaseModel):
billing_type: ProviderBillingType
monthly_quota_usd: Optional[float] = None
monthly_quota_usd: float | None = None
quota_reset_day: int = Field(default=30, ge=1, le=365) # 重置周期(天数)
quota_last_reset_at: Optional[str] = None # 当前周期开始时间
quota_expires_at: Optional[str] = None
rpm_limit: Optional[int] = Field(default=None, ge=0)
quota_last_reset_at: str | None = None # 当前周期开始时间
quota_expires_at: str | None = None
rpm_limit: int | None = Field(default=None, ge=0)
provider_priority: int = Field(default=100, ge=0, le=200)
@@ -163,13 +162,12 @@ class AdminProviderBillingAdapter(AdminApiAdapter):
provider.quota_reset_day = config.quota_reset_day
provider.provider_priority = config.provider_priority
from dateutil import parser
from sqlalchemy import func
from src.models.database import Usage
if config.quota_last_reset_at:
new_reset_at = parser.parse(config.quota_last_reset_at)
new_reset_at = datetime.fromisoformat(config.quota_last_reset_at)
# 确保有时区信息,如果没有则假设为 UTC
if new_reset_at.tzinfo is None:
new_reset_at = new_reset_at.replace(tzinfo=timezone.utc)
@@ -188,7 +186,7 @@ class AdminProviderBillingAdapter(AdminApiAdapter):
logger.info(f"Synced usage for provider {provider.name}: ${period_usage:.4f} since {new_reset_at}")
if config.quota_expires_at:
expires_at = parser.parse(config.quota_expires_at)
expires_at = datetime.fromisoformat(config.quota_expires_at)
# 确保有时区信息,如果没有则假设为 UTC
if expires_at.tzinfo is None:
expires_at = expires_at.replace(tzinfo=timezone.utc)

View File

@@ -3,7 +3,7 @@ Provider 模型管理 API
"""
from dataclasses import dataclass
from typing import Any, Dict, List, Optional
from typing import Any
from fastapi import APIRouter, Depends, Request
from sqlalchemy.orm import Session, joinedload
@@ -40,15 +40,15 @@ router = APIRouter(tags=["Model Management"])
pipeline = ApiRequestPipeline()
@router.get("/{provider_id}/models", response_model=List[ModelResponse])
@router.get("/{provider_id}/models", response_model=list[ModelResponse])
async def list_provider_models(
provider_id: str,
request: Request,
is_active: Optional[bool] = None,
is_active: bool | None = None,
skip: int = 0,
limit: int = 100,
db: Session = Depends(get_db),
) -> List[ModelResponse]:
) -> list[ModelResponse]:
"""
获取提供商的所有模型
@@ -222,13 +222,13 @@ async def delete_provider_model(
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.post("/{provider_id}/models/batch", response_model=List[ModelResponse])
@router.post("/{provider_id}/models/batch", response_model=list[ModelResponse])
async def batch_create_provider_models(
provider_id: str,
models_data: List[ModelCreate],
models_data: list[ModelCreate],
request: Request,
db: Session = Depends(get_db),
) -> List[ModelResponse]:
) -> list[ModelResponse]:
"""
批量创建模型
@@ -375,7 +375,7 @@ async def import_models_from_upstream(
@dataclass
class AdminListProviderModelsAdapter(AdminApiAdapter):
provider_id: str
is_active: Optional[bool]
is_active: bool | None
skip: int
limit: int
@@ -482,7 +482,7 @@ class AdminDeleteProviderModelAdapter(AdminApiAdapter):
@dataclass
class AdminBatchCreateModelsAdapter(AdminApiAdapter):
provider_id: str
models_data: List[ModelCreate]
models_data: list[ModelCreate]
async def handle(self, context): # type: ignore[override]
db = context.db
@@ -525,7 +525,7 @@ class AdminGetProviderAvailableSourceModelsAdapter(AdminApiAdapter):
)
# 2. 构建以 GlobalModel 为主键的字典
global_models_dict: Dict[str, Dict[str, Any]] = {}
global_models_dict: dict[str, dict[str, Any]] = {}
for model in models:
global_model = model.global_model

View File

@@ -2,7 +2,6 @@
import asyncio
from datetime import datetime, timezone
from typing import Dict, List, Optional
from fastapi import APIRouter, Depends, Query, Request
from pydantic import BaseModel, ConfigDict, Field, ValidationError
@@ -48,7 +47,7 @@ class MappingMatchingGlobalModel(BaseModel):
global_model_name: str
display_name: str
is_active: bool
matched_models: List[MappingMatchedModel] = Field(
matched_models: list[MappingMatchedModel] = Field(
default_factory=list, description="匹配到的模型列表"
)
@@ -62,8 +61,8 @@ class MappingMatchingKey(BaseModel):
key_name: str
masked_key: str
is_active: bool
allowed_models: List[str] = Field(default_factory=list, description="Key 的模型白名单")
matching_global_models: List[MappingMatchingGlobalModel] = Field(
allowed_models: list[str] = Field(default_factory=list, description="Key 的模型白名单")
matching_global_models: list[MappingMatchingGlobalModel] = Field(
default_factory=list, description="匹配到的 GlobalModel 列表"
)
@@ -75,7 +74,7 @@ class ProviderMappingPreviewResponse(BaseModel):
provider_id: str
provider_name: str
keys: List[MappingMatchingKey] = Field(
keys: list[MappingMatchingKey] = Field(
default_factory=list, description="有白名单配置且匹配到映射的 Key 列表"
)
total_keys: int = Field(0, description="有匹配结果的 Key 数量")
@@ -95,7 +94,7 @@ async def list_providers(
request: Request,
skip: int = Query(0, ge=0),
limit: int = Query(100, ge=1, le=500),
is_active: Optional[bool] = None,
is_active: bool | None = None,
db: Session = Depends(get_db),
):
"""
@@ -209,7 +208,7 @@ async def delete_provider(provider_id: str, request: Request, db: Session = Depe
class AdminListProvidersAdapter(AdminApiAdapter):
def __init__(self, skip: int, limit: int, is_active: Optional[bool]):
def __init__(self, skip: int, limit: int, is_active: bool | None):
self.skip = skip
self.limit = limit
self.is_active = is_active
@@ -473,7 +472,7 @@ async def get_provider_mapping_preview(
pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode),
timeout=MAPPING_PREVIEW_TIMEOUT_SECONDS,
)
except asyncio.TimeoutError:
except TimeoutError:
logger.warning(f"映射预览超时: provider_id={provider_id}")
raise InvalidRequestException("映射预览超时,请简化配置或稍后重试")
@@ -565,7 +564,7 @@ class AdminGetProviderMappingPreviewAdapter(AdminApiAdapter):
truncated_models = total_models_with_mappings - MAPPING_PREVIEW_MAX_MODELS
# 构建有映射配置的 GlobalModel 映射
models_with_mappings: Dict[str, tuple] = {} # id -> (model_info, mappings)
models_with_mappings: dict[str, tuple] = {} # id -> (model_info, mappings)
for gm in global_models:
config = gm.config or {}
mappings = config.get("model_mappings", [])
@@ -585,7 +584,7 @@ class AdminGetProviderMappingPreviewAdapter(AdminApiAdapter):
truncated_models=0,
)
key_infos: List[MappingMatchingKey] = []
key_infos: list[MappingMatchingKey] = []
total_matches = 0
# 创建 CryptoService 实例
@@ -611,10 +610,10 @@ class AdminGetProviderMappingPreviewAdapter(AdminApiAdapter):
pass
# 查找匹配的 GlobalModel
matching_global_models: List[MappingMatchingGlobalModel] = []
matching_global_models: list[MappingMatchingGlobalModel] = []
for gm_id, (gm, mappings) in models_with_mappings.items():
matched_models: List[MappingMatchedModel] = []
matched_models: list[MappingMatchedModel] = []
for allowed_model in allowed_models_list:
for mapping_pattern in mappings:

View File

@@ -4,7 +4,6 @@ Provider 摘要与健康监控 API
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from typing import Dict, List
from fastapi import APIRouter, Depends, Query, Request
from sqlalchemy import case, func
@@ -35,11 +34,11 @@ router = APIRouter(tags=["Provider Summary"])
pipeline = ApiRequestPipeline()
@router.get("/summary", response_model=List[ProviderWithEndpointsSummary])
@router.get("/summary", response_model=list[ProviderWithEndpointsSummary])
async def get_providers_summary(
request: Request,
db: Session = Depends(get_db),
) -> List[ProviderWithEndpointsSummary]:
) -> list[ProviderWithEndpointsSummary]:
"""
获取所有提供商摘要信息
@@ -381,8 +380,8 @@ class AdminProviderHealthMonitorAdapter(AdminApiAdapter):
)
attempts = attempts_query.limit(limit_rows).all()
buffered_attempts: Dict[str, List[RequestCandidate]] = {eid: [] for eid in endpoint_ids}
counters: Dict[str, int] = {eid: 0 for eid in endpoint_ids}
buffered_attempts: dict[str, list[RequestCandidate]] = {eid: [] for eid in endpoint_ids}
counters: dict[str, int] = {eid: 0 for eid in endpoint_ids}
for attempt in attempts:
if not attempt.endpoint_id or attempt.endpoint_id not in buffered_attempts:
@@ -392,10 +391,10 @@ class AdminProviderHealthMonitorAdapter(AdminApiAdapter):
buffered_attempts[attempt.endpoint_id].append(attempt)
counters[attempt.endpoint_id] += 1
endpoint_monitors: List[EndpointHealthMonitor] = []
endpoint_monitors: list[EndpointHealthMonitor] = []
for endpoint in endpoints:
attempt_list = list(reversed(buffered_attempts.get(endpoint.id, [])))
events: List[EndpointHealthEvent] = []
events: list[EndpointHealthEvent] = []
for attempt in attempt_list:
event_timestamp = attempt.finished_at or attempt.started_at or attempt.created_at
events.append(

View File

@@ -4,8 +4,6 @@ IP 安全管理接口
提供 IP 黑白名单管理和速率限制统计
"""
from typing import List, Optional
from fastapi import APIRouter, Depends, HTTPException, Request, status
from pydantic import BaseModel, Field, ValidationError
from sqlalchemy.orm import Session
@@ -14,7 +12,6 @@ from src.api.base.adapter import ApiMode
from src.api.base.authenticated_adapter import AuthenticatedApiAdapter
from src.api.base.pipeline import ApiRequestPipeline
from src.core.exceptions import InvalidRequestException, translate_pydantic_error
from src.core.logger import logger
from src.database import get_db
from src.services.rate_limit.ip_limiter import IPRateLimiter
@@ -30,7 +27,7 @@ class AddIPToBlacklistRequest(BaseModel):
ip_address: str = Field(..., description="IP 地址")
reason: str = Field(..., min_length=1, max_length=200, description="加入黑名单的原因")
ttl: Optional[int] = Field(None, gt=0, description="过期时间None 表示永久")
ttl: int | None = Field(None, gt=0, description="过期时间None 表示永久")
class RemoveIPFromBlacklistRequest(BaseModel):

View File

@@ -1,10 +1,8 @@
"""系统设置API端点。"""
from __future__ import annotations
import copy
from dataclasses import dataclass
from typing import Optional
from fastapi import APIRouter, Depends, HTTPException, Request, status
from pydantic import ValidationError
@@ -649,9 +647,8 @@ class AdminSystemStatsAdapter(AdminApiAdapter):
class AdminTriggerCleanupAdapter(AdminApiAdapter):
async def handle(self, context): # type: ignore[override]
"""手动触发清理任务"""
from datetime import datetime, timedelta, timezone
from datetime import datetime, timezone
from sqlalchemy import func
from src.services.system.maintenance_scheduler import get_maintenance_scheduler

View File

@@ -3,7 +3,6 @@
from collections import defaultdict
from dataclasses import dataclass
from datetime import datetime
from typing import Optional
from fastapi import APIRouter, Depends, HTTPException, Query, Request
from sqlalchemy import func
@@ -34,8 +33,8 @@ pipeline = ApiRequestPipeline()
async def get_usage_aggregation(
request: Request,
group_by: str = Query(..., description="Aggregation dimension: model, user, provider, or api_format"),
start_date: Optional[datetime] = None,
end_date: Optional[datetime] = None,
start_date: datetime | None = None,
end_date: datetime | None = None,
limit: int = Query(20, ge=1, le=100),
db: Session = Depends(get_db),
):
@@ -75,8 +74,8 @@ async def get_usage_aggregation(
@router.get("/stats")
async def get_usage_stats(
request: Request,
start_date: Optional[datetime] = None,
end_date: Optional[datetime] = None,
start_date: datetime | None = None,
end_date: datetime | None = None,
db: Session = Depends(get_db),
):
"""
@@ -122,14 +121,14 @@ async def get_activity_heatmap(
@router.get("/records")
async def get_usage_records(
request: Request,
start_date: Optional[datetime] = None,
end_date: Optional[datetime] = None,
search: Optional[str] = None, # 通用搜索:用户名、密钥名、模型名、提供商名
user_id: Optional[str] = None,
username: Optional[str] = None,
model: Optional[str] = None,
provider: Optional[str] = None,
status: Optional[str] = None, # stream, standard, error
start_date: datetime | None = None,
end_date: datetime | None = None,
search: str | None = None, # 通用搜索:用户名、密钥名、模型名、提供商名
user_id: str | None = None,
username: str | None = None,
model: str | None = None,
provider: str | None = None,
status: str | None = None, # stream, standard, error
limit: int = Query(100, ge=1, le=500),
offset: int = Query(0, ge=0),
db: Session = Depends(get_db),
@@ -179,7 +178,7 @@ async def get_usage_records(
@router.get("/active")
async def get_active_requests(
request: Request,
ids: Optional[str] = Query(None, description="逗号分隔的请求 ID 列表,用于查询特定请求的状态"),
ids: str | None = Query(None, description="逗号分隔的请求 ID 列表,用于查询特定请求的状态"),
db: Session = Depends(get_db),
):
"""
@@ -259,7 +258,7 @@ async def get_usage_detail(
class AdminUsageStatsAdapter(AdminApiAdapter):
def __init__(self, start_date: Optional[datetime], end_date: Optional[datetime]):
def __init__(self, start_date: datetime | None, end_date: datetime | None):
self.start_date = start_date
self.end_date = end_date
@@ -339,7 +338,7 @@ class AdminActivityHeatmapAdapter(AdminApiAdapter):
class AdminUsageByModelAdapter(AdminApiAdapter):
def __init__(self, start_date: Optional[datetime], end_date: Optional[datetime], limit: int):
def __init__(self, start_date: datetime | None, end_date: datetime | None, limit: int):
self.start_date = start_date
self.end_date = end_date
self.limit = limit
@@ -386,7 +385,7 @@ class AdminUsageByModelAdapter(AdminApiAdapter):
class AdminUsageByUserAdapter(AdminApiAdapter):
def __init__(self, start_date: Optional[datetime], end_date: Optional[datetime], limit: int):
def __init__(self, start_date: datetime | None, end_date: datetime | None, limit: int):
self.start_date = start_date
self.end_date = end_date
self.limit = limit
@@ -436,7 +435,7 @@ class AdminUsageByUserAdapter(AdminApiAdapter):
class AdminUsageByProviderAdapter(AdminApiAdapter):
def __init__(self, start_date: Optional[datetime], end_date: Optional[datetime], limit: int):
def __init__(self, start_date: datetime | None, end_date: datetime | None, limit: int):
self.start_date = start_date
self.end_date = end_date
self.limit = limit
@@ -446,7 +445,7 @@ class AdminUsageByProviderAdapter(AdminApiAdapter):
# 从 request_candidates 表统计每个 Provider 的尝试次数和成功率
# 这样可以正确统计 Fallback 场景(一个请求可能尝试多个 Provider
from sqlalchemy import case, Integer
from sqlalchemy import case
attempt_query = db.query(
RequestCandidate.provider_id,
@@ -550,7 +549,7 @@ class AdminUsageByProviderAdapter(AdminApiAdapter):
class AdminUsageByApiFormatAdapter(AdminApiAdapter):
def __init__(self, start_date: Optional[datetime], end_date: Optional[datetime], limit: int):
def __init__(self, start_date: datetime | None, end_date: datetime | None, limit: int):
self.start_date = start_date
self.end_date = end_date
self.limit = limit
@@ -608,14 +607,14 @@ class AdminUsageByApiFormatAdapter(AdminApiAdapter):
class AdminUsageRecordsAdapter(AdminApiAdapter):
def __init__(
self,
start_date: Optional[datetime],
end_date: Optional[datetime],
search: Optional[str],
user_id: Optional[str],
username: Optional[str],
model: Optional[str],
provider: Optional[str],
status: Optional[str],
start_date: datetime | None,
end_date: datetime | None,
search: str | None,
user_id: str | None,
username: str | None,
model: str | None,
provider: str | None,
status: str | None,
limit: int,
offset: int,
):
@@ -744,7 +743,7 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
for req_id, candidates in request_candidates.items():
# 提取所有不同的 candidate_index
unique_candidates = set(c[0] for c in candidates)
unique_candidates = {c[0] for c in candidates}
# 如果有多个不同的 candidate_index说明发生了 FallbackProvider 切换)
fallback_map[req_id] = len(unique_candidates) > 1
@@ -877,7 +876,7 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
class AdminActiveRequestsAdapter(AdminApiAdapter):
"""轻量级活跃请求状态查询适配器"""
def __init__(self, ids: Optional[str]):
def __init__(self, ids: str | None):
self.ids = ids
async def handle(self, context): # type: ignore[override]
@@ -1033,8 +1032,8 @@ class AdminUsageDetailAdapter(AdminApiAdapter):
@router.get("/cache-affinity/ttl-analysis")
async def analyze_cache_affinity_ttl(
request: Request,
user_id: Optional[str] = Query(None, description="指定用户 ID"),
api_key_id: Optional[str] = Query(None, description="指定 API Key ID"),
user_id: str | None = Query(None, description="指定用户 ID"),
api_key_id: str | None = Query(None, description="指定 API Key ID"),
hours: int = Query(168, ge=1, le=720, description="分析最近多少小时的数据"),
db: Session = Depends(get_db),
):
@@ -1057,8 +1056,8 @@ async def analyze_cache_affinity_ttl(
@router.get("/cache-affinity/hit-analysis")
async def analyze_cache_hit(
request: Request,
user_id: Optional[str] = Query(None, description="指定用户 ID"),
api_key_id: Optional[str] = Query(None, description="指定 API Key ID"),
user_id: str | None = Query(None, description="指定用户 ID"),
api_key_id: str | None = Query(None, description="指定 API Key ID"),
hours: int = Query(168, ge=1, le=720, description="分析最近多少小时的数据"),
db: Session = Depends(get_db),
):
@@ -1080,8 +1079,8 @@ class CacheAffinityTTLAnalysisAdapter(AdminApiAdapter):
def __init__(
self,
user_id: Optional[str],
api_key_id: Optional[str],
user_id: str | None,
api_key_id: str | None,
hours: int,
):
self.user_id = user_id
@@ -1114,8 +1113,8 @@ class CacheHitAnalysisAdapter(AdminApiAdapter):
def __init__(
self,
user_id: Optional[str],
api_key_id: Optional[str],
user_id: str | None,
api_key_id: str | None,
hours: int,
):
self.user_id = user_id
@@ -1147,7 +1146,7 @@ async def get_interval_timeline(
request: Request,
hours: int = Query(24, ge=1, le=720, description="分析最近多少小时的数据"),
limit: int = Query(10000, ge=100, le=50000, description="最大返回数据点数量"),
user_id: Optional[str] = Query(None, description="指定用户 ID"),
user_id: str | None = Query(None, description="指定用户 ID"),
include_user_info: bool = Query(False, description="是否包含用户信息(用于管理员多用户视图)"),
db: Session = Depends(get_db),
):
@@ -1177,7 +1176,7 @@ class IntervalTimelineAdapter(AdminApiAdapter):
self,
hours: int,
limit: int,
user_id: Optional[str] = None,
user_id: str | None = None,
include_user_info: bool = False,
):
self.hours = hours

View File

@@ -1,7 +1,6 @@
"""用户管理 API 端点。"""
from datetime import datetime, timezone
from typing import Optional
from fastapi import APIRouter, Depends, HTTPException, Query, Request, status
from pydantic import ValidationError
@@ -48,8 +47,8 @@ async def list_users(
request: Request,
skip: int = Query(0, ge=0, description="跳过记录数"),
limit: int = Query(100, ge=1, le=1000, description="返回记录数"),
role: Optional[str] = Query(None, description="按角色筛选user/admin"),
is_active: Optional[bool] = Query(None, description="按状态筛选"),
role: str | None = Query(None, description="按角色筛选user/admin"),
is_active: bool | None = Query(None, description="按状态筛选"),
db: Session = Depends(get_db),
):
"""
@@ -136,7 +135,7 @@ async def reset_user_quota(user_id: str, request: Request, db: Session = Depends
async def get_user_api_keys(
user_id: str,
request: Request,
is_active: Optional[bool] = Query(None, description="按状态筛选"),
is_active: bool | None = Query(None, description="按状态筛选"),
db: Session = Depends(get_db),
):
"""
@@ -274,7 +273,7 @@ class AdminCreateUserAdapter(AdminApiAdapter):
class AdminListUsersAdapter(AdminApiAdapter):
def __init__(self, skip: int, limit: int, role: Optional[str], is_active: Optional[bool]):
def __init__(self, skip: int, limit: int, role: str | None, is_active: bool | None):
self.skip = skip
self.limit = limit
self.role = role
@@ -467,7 +466,7 @@ class AdminResetUserQuotaAdapter(AdminApiAdapter):
class AdminGetUserKeysAdapter(AdminApiAdapter):
"""获取用户的API Keys"""
def __init__(self, user_id: str, is_active: Optional[bool]):
def __init__(self, user_id: str, is_active: bool | None):
self.user_id = user_id
self.is_active = is_active

View File

@@ -1,9 +1,8 @@
"""公告系统 API 端点。"""
from dataclasses import dataclass
from typing import Optional
from fastapi import APIRouter, Depends, HTTPException, Query, Request, status
from fastapi import APIRouter, Depends, Query, Request
from pydantic import ValidationError
from sqlalchemy.orm import Session
@@ -12,7 +11,6 @@ from src.api.base.admin_adapter import AdminApiAdapter
from src.api.base.authenticated_adapter import AuthenticatedApiAdapter
from src.api.base.pipeline import ApiRequestPipeline
from src.core.exceptions import InvalidRequestException, translate_pydantic_error
from src.core.logger import logger
from src.database import get_db
from src.models.api import CreateAnnouncementRequest, UpdateAnnouncementRequest
from src.models.database import User
@@ -251,7 +249,7 @@ class AnnouncementOptionalAuthAdapter(ApiAdapter):
context.extra["optional_user"] = await self._resolve_optional_user(context)
return None
async def _resolve_optional_user(self, context) -> Optional[User]:
async def _resolve_optional_user(self, context) -> User | None:
if context.user:
return context.user
@@ -285,7 +283,7 @@ class AnnouncementOptionalAuthAdapter(ApiAdapter):
except Exception:
return None
def get_optional_user(self, context) -> Optional[User]:
def get_optional_user(self, context) -> User | None:
return context.extra.get("optional_user")

View File

@@ -2,10 +2,9 @@
认证相关API端点
"""
from typing import Optional, Tuple
from fastapi import APIRouter, Depends, HTTPException, Request, status
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
from fastapi.security import HTTPBearer
from pydantic import ValidationError
from sqlalchemy.orm import Session
@@ -42,7 +41,7 @@ from src.services.email import EmailSenderService, EmailVerificationService
from src.utils.request_utils import get_client_ip, get_user_agent
def validate_email_suffix(db: Session, email: str) -> Tuple[bool, Optional[str]]:
def validate_email_suffix(db: Session, email: str) -> tuple[bool, str | None]:
"""
验证邮箱后缀是否允许注册

View File

@@ -1,8 +1,6 @@
from __future__ import annotations
from abc import ABC, abstractmethod
from enum import Enum
from typing import Any, Dict, Optional
from typing import Any
from fastapi import Request, Response
@@ -23,7 +21,7 @@ class ApiAdapter(ABC):
name: str = "base"
mode: ApiMode = ApiMode.STANDARD
api_format: Optional[str] = None # 对应 Provider API 格式提示
api_format: str | None = None # 对应 Provider API 格式提示
audit_log_enabled: bool = True
audit_success_event = None
audit_failure_event = None
@@ -36,7 +34,7 @@ class ApiAdapter(ABC):
"""可选的授权钩子,默认允许通过。"""
return None
def extract_api_key(self, request: Request) -> Optional[str]:
def extract_api_key(self, request: Request) -> str | None:
"""
从请求中提取客户端 API 密钥。
@@ -55,17 +53,17 @@ class ApiAdapter(ABC):
context: ApiRequestContext,
*,
success: bool,
status_code: Optional[int],
error: Optional[str] = None,
) -> Dict[str, Any]:
status_code: int | None,
error: str | None = None,
) -> dict[str, Any]:
"""允许适配器在审计日志中追加自定义字段。"""
return {}
def detect_capability_requirements(
self,
headers: Dict[str, str],
request_body: Optional[Dict[str, Any]] = None,
) -> Dict[str, bool]:
headers: dict[str, str],
request_body: dict[str, Any] | None = None,
) -> dict[str, bool]:
"""
检测请求中隐含的能力需求(子类可覆盖)

View File

@@ -1,5 +1,3 @@
from __future__ import annotations
from fastapi import HTTPException
from src.models.database import UserRole

View File

@@ -1,10 +1,9 @@
from __future__ import annotations
import json
import time
import uuid
from dataclasses import dataclass, field
from typing import Any, Dict, Optional
from typing import Any
from fastapi import HTTPException, Request
from sqlalchemy.orm import Session
@@ -21,34 +20,34 @@ class ApiRequestContext:
request: Request
db: Session
user: Optional[User]
api_key: Optional[ApiKey]
user: User | None
api_key: ApiKey | None
request_id: str
start_time: float
client_ip: str
user_agent: str
original_headers: Dict[str, str]
query_params: Dict[str, str]
original_headers: dict[str, str]
query_params: dict[str, str]
raw_body: bytes | None = None
json_body: Optional[Dict[str, Any]] = None
quota_remaining: Optional[float] = None
json_body: dict[str, Any] | None = None
quota_remaining: float | None = None
mode: str = "standard" # standard / proxy
api_format_hint: Optional[str] = None
api_format_hint: str | None = None
# URL 路径参数(如 Gemini API 的 /v1beta/models/{model}:generateContent
path_params: Dict[str, Any] = field(default_factory=dict)
path_params: dict[str, Any] = field(default_factory=dict)
# Management Token用于管理 API 认证)
management_token: Optional[ManagementToken] = None
management_token: ManagementToken | None = None
# 供适配器扩展的状态存储
extra: Dict[str, Any] = field(default_factory=dict)
audit_metadata: Dict[str, Any] = field(default_factory=dict)
extra: dict[str, Any] = field(default_factory=dict)
audit_metadata: dict[str, Any] = field(default_factory=dict)
# 高频轮询端点日志抑制标志
quiet_logging: bool = False
def ensure_json_body(self) -> Dict[str, Any]:
def ensure_json_body(self) -> dict[str, Any]:
"""确保请求体已解析为JSON并返回。"""
if self.json_body is not None:
return self.json_body
@@ -70,7 +69,7 @@ class ApiRequestContext:
if value is not None:
self.audit_metadata[key] = value
def extend_audit_metadata(self, data: Dict[str, Any]) -> None:
def extend_audit_metadata(self, data: dict[str, Any]) -> None:
"""批量附加审计字段。"""
for key, value in data.items():
if value is not None:
@@ -81,13 +80,13 @@ class ApiRequestContext:
cls,
request: Request,
db: Session,
user: Optional[User],
api_key: Optional[ApiKey],
raw_body: Optional[bytes] = None,
user: User | None,
api_key: ApiKey | None,
raw_body: bytes | None = None,
mode: str = "standard",
api_format_hint: Optional[str] = None,
path_params: Optional[Dict[str, Any]] = None,
) -> "ApiRequestContext":
api_format_hint: str | None = None,
path_params: dict[str, Any] | None = None,
) -> ApiRequestContext:
"""创建上下文实例并提前读取必要的元数据。"""
request_id = getattr(request.state, "request_id", None) or str(uuid.uuid4())[:8]
setattr(request.state, "request_id", request_id)

View File

@@ -10,8 +10,9 @@
4. Key 的 allowed_models 允许该模型null = 允许所有)
"""
from __future__ import annotations
from dataclasses import asdict, dataclass
from typing import Any, Optional
from typing import Any
from sqlalchemy.orm import Session
@@ -27,7 +28,7 @@ _CACHE_KEY_PREFIX = "models:list"
_CACHE_TTL = CacheTTL.MODEL # 300 秒
def _get_cache_key(api_formats: list[str], client_format: Optional[str] = None) -> str:
def _get_cache_key(api_formats: list[str], client_format: str | None = None) -> str:
"""生成缓存 key"""
formats_str = ",".join(sorted(api_formats))
format_key = (client_format or "any").lower()
@@ -35,8 +36,8 @@ def _get_cache_key(api_formats: list[str], client_format: Optional[str] = None)
async def _get_cached_models(
api_formats: list[str], client_format: Optional[str] = None
) -> Optional[list["ModelInfo"]]:
api_formats: list[str], client_format: str | None = None
) -> list[ModelInfo] | None:
"""从缓存获取模型列表"""
cache_key = _get_cache_key(api_formats, client_format)
try:
@@ -51,8 +52,8 @@ async def _get_cached_models(
async def _set_cached_models(
api_formats: list[str],
models: list["ModelInfo"],
client_format: Optional[str] = None,
models: list[ModelInfo],
client_format: str | None = None,
) -> None:
"""将模型列表写入缓存"""
cache_key = _get_cache_key(api_formats, client_format)
@@ -87,8 +88,8 @@ class ModelInfo:
id: str # 模型 ID (GlobalModel.name 或 provider_model_name)
display_name: str
description: Optional[str]
created_at: Optional[str] # ISO 格式
description: str | None
created_at: str | None # ISO 格式
created_timestamp: int # Unix 时间戳
provider_name: str
provider_id: str = "" # Provider ID用于权限过滤
@@ -100,27 +101,27 @@ class ModelInfo:
image_generation: bool = False
structured_output: bool = False
# 规格参数
context_limit: Optional[int] = None
output_limit: Optional[int] = None
context_limit: int | None = None
output_limit: int | None = None
# 元信息
family: Optional[str] = None
knowledge_cutoff: Optional[str] = None
input_modalities: Optional[list[str]] = None
output_modalities: Optional[list[str]] = None
family: str | None = None
knowledge_cutoff: str | None = None
input_modalities: list[str] | None = None
output_modalities: list[str] | None = None
@dataclass
class AccessRestrictions:
"""API Key 或 User 的访问限制"""
allowed_providers: Optional[list[str]] = None # 允许的 Provider ID 列表
allowed_models: Optional[list[str]] = None # 允许的模型名称列表
allowed_api_formats: Optional[list[str]] = None # 允许的 API 格式列表
allowed_providers: list[str] | None = None # 允许的 Provider ID 列表
allowed_models: list[str] | None = None # 允许的模型名称列表
allowed_api_formats: list[str] | None = None # 允许的 API 格式列表
@classmethod
def from_api_key_and_user(
cls, api_key: Optional[ApiKey], user: Optional[User]
) -> "AccessRestrictions":
cls, api_key: ApiKey | None, user: User | None
) -> AccessRestrictions:
"""
从 API Key 和 User 合并访问限制
@@ -130,9 +131,9 @@ class AccessRestrictions:
- 如果 API Key 无限制但 User 有限制,使用 User 的限制
- 两者都无限制则返回空限制
"""
allowed_providers: Optional[list[str]] = None
allowed_models: Optional[list[str]] = None
allowed_api_formats: Optional[list[str]] = None
allowed_providers: list[str] | None = None
allowed_models: list[str] | None = None
allowed_api_formats: list[str] | None = None
# 优先使用 API Key 的限制
if api_key:
@@ -197,8 +198,8 @@ class AccessRestrictions:
def _normalize_api_formats(
api_formats: Optional[list[str]],
provider_to_formats: Optional[dict[str, set[str]]] = None,
api_formats: list[str] | None,
provider_to_formats: dict[str, set[str]] | None = None,
) -> list[str]:
"""规范化 API 格式列表(大写),必要时从 provider_to_formats 兜底"""
if api_formats:
@@ -212,7 +213,7 @@ def _normalize_api_formats(
def _get_provider_model_names_for_formats(
model: Model, usable_formats: Optional[set[str]] = None
model: Model, usable_formats: set[str] | None = None
) -> set[str]:
"""
获取模型在指定格式下支持的 Provider 模型名称集合
@@ -305,7 +306,7 @@ def get_compatible_provider_formats(
def get_available_provider_ids(
db: Session,
api_formats: list[str],
provider_to_formats: Optional[dict[str, set[str]]] = None,
provider_to_formats: dict[str, set[str]] | None = None,
) -> set[str]:
"""
返回有可用端点的 Provider IDs
@@ -334,7 +335,7 @@ def get_available_provider_ids(
def _get_available_model_ids_for_format(
db: Session,
api_formats: list[str],
provider_to_formats: Optional[dict[str, set[str]]] = None,
provider_to_formats: dict[str, set[str]] | None = None,
) -> set[str]:
"""
获取指定格式下真正可用的模型 ID 集合
@@ -410,7 +411,7 @@ def _get_available_model_ids_for_format(
return available_model_ids
def _extract_model_info(model: Any) -> Optional[ModelInfo]:
def _extract_model_info(model: Any) -> ModelInfo | None:
"""
从 Model 对象提取 ModelInfo
@@ -424,7 +425,7 @@ def _extract_model_info(model: Any) -> Optional[ModelInfo]:
model_id: str = global_model.name
display_name: str = global_model.display_name
created_at: Optional[str] = (
created_at: str | None = (
model.created_at.strftime("%Y-%m-%dT%H:%M:%SZ") if model.created_at else None
)
created_timestamp: int = int(model.created_at.timestamp()) if model.created_at else 0
@@ -433,7 +434,7 @@ def _extract_model_info(model: Any) -> Optional[ModelInfo]:
# 从 GlobalModel.config 提取配置信息
config: dict = global_model.config or {}
description: Optional[str] = config.get("description")
description: str | None = config.get("description")
return ModelInfo(
id=model_id,
@@ -464,10 +465,10 @@ def _extract_model_info(model: Any) -> Optional[ModelInfo]:
async def list_available_models(
db: Session,
available_provider_ids: set[str],
api_formats: Optional[list[str]] = None,
restrictions: Optional[AccessRestrictions] = None,
provider_to_formats: Optional[dict[str, set[str]]] = None,
client_format: Optional[str] = None,
api_formats: list[str] | None = None,
restrictions: AccessRestrictions | None = None,
provider_to_formats: dict[str, set[str]] | None = None,
client_format: str | None = None,
) -> list[ModelInfo]:
"""
获取可用模型列表(已去重,带缓存)
@@ -503,7 +504,7 @@ async def list_available_models(
return cached
# 如果提供了 api_formats获取真正可用的模型 ID
available_model_ids: Optional[set[str]] = None
available_model_ids: set[str] | None = None
if normalized_formats:
available_model_ids = _get_available_model_ids_for_format(
db, normalized_formats, provider_to_formats
@@ -551,10 +552,10 @@ def find_model_by_id(
db: Session,
model_id: str,
available_provider_ids: set[str],
api_formats: Optional[list[str]] = None,
restrictions: Optional[AccessRestrictions] = None,
provider_to_formats: Optional[dict[str, set[str]]] = None,
) -> Optional[ModelInfo]:
api_formats: list[str] | None = None,
restrictions: AccessRestrictions | None = None,
provider_to_formats: dict[str, set[str]] | None = None,
) -> ModelInfo | None:
"""
按 ID 查找模型(仅支持 GlobalModel.name
@@ -575,7 +576,7 @@ def find_model_by_id(
normalized_formats = _normalize_api_formats(api_formats, provider_to_formats)
# 如果提供了 api_formats获取真正可用的模型 ID
available_model_ids: Optional[set[str]] = None
available_model_ids: set[str] | None = None
if normalized_formats:
available_model_ids = _get_available_model_ids_for_format(
db, normalized_formats, provider_to_formats

View File

@@ -1,7 +1,6 @@
from __future__ import annotations
from dataclasses import asdict, dataclass
from typing import Any, List, Sequence, Tuple, TypeVar
from typing import Any, TypeVar
from collections.abc import Sequence
from sqlalchemy.orm import Query
@@ -19,7 +18,7 @@ class PaginationMeta:
return asdict(self)
def paginate_query(query: Query, limit: int, offset: int) -> Tuple[int, List[T]]:
def paginate_query(query: Query, limit: int, offset: int) -> tuple[int, list[T]]:
"""
对 SQLAlchemy 查询应用 limit/offset并返回总数与结果列表。
"""
@@ -30,7 +29,7 @@ def paginate_query(query: Query, limit: int, offset: int) -> Tuple[int, List[T]]
def paginate_sequence(
items: Sequence[T], limit: int, offset: int
) -> Tuple[List[T], PaginationMeta]:
) -> tuple[list[T], PaginationMeta]:
"""
对内存序列应用分页,返回切片和元数据。
"""
@@ -40,7 +39,7 @@ def paginate_sequence(
return sliced, meta
def build_pagination_payload(items: List[dict], meta: PaginationMeta, **extra: Any) -> dict:
def build_pagination_payload(items: list[dict], meta: PaginationMeta, **extra: Any) -> dict:
"""
构建标准分页响应 payload。
"""

View File

@@ -2,7 +2,7 @@ from __future__ import annotations
import time
from enum import Enum
from typing import TYPE_CHECKING, Any, Optional, Tuple
from typing import TYPE_CHECKING, Any
from fastapi import HTTPException, Request
from sqlalchemy.orm import Session
@@ -52,8 +52,8 @@ class ApiRequestPipeline:
db: Session,
*,
mode: ApiMode = ApiMode.STANDARD,
api_format_hint: Optional[str] = None,
path_params: Optional[dict[str, Any]] = None,
api_format_hint: str | None = None,
path_params: dict[str, Any] | None = None,
):
# 高频轮询端点抑制 debug 日志
is_quiet = http_request.url.path in QUIET_POLLING_PATHS
@@ -95,7 +95,7 @@ class ApiRequestPipeline:
)
if not is_quiet:
logger.debug("[Pipeline] Raw body读取完成 | size=%d bytes", len(raw_body) if raw_body is not None else 0)
except asyncio.TimeoutError:
except TimeoutError:
timeout_sec = int(config.request_body_timeout)
logger.error(f"读取请求体超时({timeout_sec}s),可能客户端未发送完整请求体")
raise HTTPException(
@@ -166,7 +166,7 @@ class ApiRequestPipeline:
def _authenticate_client(
self, request: Request, db: Session, adapter: ApiAdapter, *, quiet: bool = False
) -> Tuple[User, ApiKey]:
) -> tuple[User, ApiKey]:
if not quiet:
logger.debug("[Pipeline._authenticate_client] 开始")
# 使用 adapter 的 extract_api_key 方法,支持不同 API 格式的认证头
@@ -215,7 +215,7 @@ class ApiRequestPipeline:
async def _authenticate_admin(
self, request: Request, db: Session
) -> Tuple[User, Optional["ManagementToken"]]:
) -> tuple[User, ManagementToken | None]:
"""管理员认证,支持 JWT 和 Management Token 两种方式"""
from src.models.database import ManagementToken
from src.utils.request_utils import get_client_ip
@@ -278,7 +278,7 @@ class ApiRequestPipeline:
async def _authenticate_user(
self, request: Request, db: Session
) -> Tuple[User, Optional["ManagementToken"]]:
) -> tuple[User, ManagementToken | None]:
"""用户认证,支持 JWT 和 Management Token 两种方式"""
from src.models.database import ManagementToken
from src.utils.request_utils import get_client_ip
@@ -329,7 +329,7 @@ class ApiRequestPipeline:
async def _authenticate_management(
self, request: Request, db: Session
) -> Tuple[User, "ManagementToken"]:
) -> tuple[User, ManagementToken]:
"""Management Token 认证"""
from src.models.database import ManagementToken
from src.utils.request_utils import get_client_ip
@@ -362,7 +362,7 @@ class ApiRequestPipeline:
return user, management_token
def _calculate_quota_remaining(self, user: Optional[User]) -> Optional[float]:
def _calculate_quota_remaining(self, user: User | None) -> float | None:
if not user:
return None
if user.quota_usd is None or user.quota_usd < 0:
@@ -375,8 +375,8 @@ class ApiRequestPipeline:
adapter: ApiAdapter,
*,
success: bool,
status_code: Optional[int] = None,
error: Optional[str] = None,
status_code: int | None = None,
error: str | None = None,
) -> None:
"""记录审计事件
@@ -432,8 +432,8 @@ class ApiRequestPipeline:
adapter: ApiAdapter,
*,
success: bool,
status_code: Optional[int],
error: Optional[str],
status_code: int | None,
error: str | None,
) -> dict:
duration_ms = max((time.time() - context.start_time) * 1000, 0.0)
request = context.request

View File

@@ -2,7 +2,6 @@
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from typing import List
from fastapi import APIRouter, Depends, HTTPException, Query, Request
from sqlalchemy import and_, func
@@ -952,7 +951,7 @@ class DashboardDailyStatsAdapter(DashboardAdapter):
# 构建完整日期序列(使用业务时区日期)
current_date = start_date_local.date()
end_date_date = end_date_local.date()
formatted: List[dict] = []
formatted: list[dict] = []
while current_date <= end_date_date:
date_str = current_date.isoformat()
stat = stats_map.get(date_str)

View File

@@ -27,21 +27,20 @@
from __future__ import annotations
import asyncio
import time
from typing import (
TYPE_CHECKING,
Any,
Awaitable,
Callable,
Coroutine,
Dict,
Optional,
Protocol,
TypeVar,
runtime_checkable,
)
from collections.abc import Callable
from collections.abc import Awaitable, Coroutine
from fastapi import Request
from fastapi.responses import JSONResponse, StreamingResponse
from sqlalchemy.orm import Session
@@ -57,6 +56,9 @@ from src.services.usage.service import UsageService
if TYPE_CHECKING:
from src.api.handlers.base.stream_context import StreamContext
# Adapter 检测器类型:接受 headers 和可选的 request_body返回能力需求字典
type AdapterDetectorType = Callable[[dict[str, str], dict[str, Any] | None], dict[str, bool]]
class MessageTelemetry:
"""
@@ -105,29 +107,29 @@ class MessageTelemetry:
output_tokens: int,
response_time_ms: int,
status_code: int,
request_body: Dict[str, Any],
request_headers: Dict[str, Any],
request_body: dict[str, Any],
request_headers: dict[str, Any],
response_body: Any,
response_headers: Dict[str, Any],
client_response_headers: Optional[Dict[str, Any]] = None,
response_headers: dict[str, Any],
client_response_headers: dict[str, Any] | None = None,
cache_creation_tokens: int = 0,
cache_read_tokens: int = 0,
is_stream: bool = False,
provider_request_headers: Optional[Dict[str, Any]] = None,
provider_request_headers: dict[str, Any] | None = None,
# 时间指标
first_byte_time_ms: Optional[int] = None, # 首字时间/TTFB
first_byte_time_ms: int | None = None, # 首字时间/TTFB
# Provider 侧追踪信息(用于记录真实成本)
provider_id: Optional[str] = None,
provider_endpoint_id: Optional[str] = None,
provider_api_key_id: Optional[str] = None,
api_format: Optional[str] = None,
provider_id: str | None = None,
provider_endpoint_id: str | None = None,
provider_api_key_id: str | None = None,
api_format: str | None = None,
# 格式转换追踪
endpoint_api_format: Optional[str] = None, # 端点原生 API 格式
endpoint_api_format: str | None = None, # 端点原生 API 格式
has_format_conversion: bool = False, # 是否发生了格式转换
# 模型映射信息
target_model: Optional[str] = None,
target_model: str | None = None,
# Provider 响应元数据(如 Gemini 的 modelVersion
response_metadata: Optional[Dict[str, Any]] = None,
response_metadata: dict[str, Any] | None = None,
) -> float:
total_cost = await self.calculate_cost(
provider,
@@ -199,24 +201,24 @@ class MessageTelemetry:
response_time_ms: int,
status_code: int,
error_message: str,
request_body: Dict[str, Any],
request_headers: Dict[str, Any],
request_body: dict[str, Any],
request_headers: dict[str, Any],
is_stream: bool,
api_format: Optional[str] = None,
provider_request_headers: Optional[Dict[str, Any]] = None,
api_format: str | None = None,
provider_request_headers: dict[str, Any] | None = None,
# 预估 token 信息(来自 message_start 事件,用于中断请求的成本估算)
input_tokens: int = 0,
output_tokens: int = 0,
cache_creation_tokens: int = 0,
cache_read_tokens: int = 0,
response_body: Optional[Dict[str, Any]] = None,
response_headers: Optional[Dict[str, Any]] = None,
client_response_headers: Optional[Dict[str, Any]] = None,
response_body: dict[str, Any] | None = None,
response_headers: dict[str, Any] | None = None,
client_response_headers: dict[str, Any] | None = None,
# 格式转换追踪
endpoint_api_format: Optional[str] = None,
endpoint_api_format: str | None = None,
has_format_conversion: bool = False,
# 模型映射信息
target_model: Optional[str] = None,
target_model: str | None = None,
) -> None:
"""
记录失败请求
@@ -273,24 +275,24 @@ class MessageTelemetry:
provider: str,
model: str,
response_time_ms: int,
first_byte_time_ms: Optional[int],
first_byte_time_ms: int | None,
status_code: int,
request_body: Dict[str, Any],
request_headers: Dict[str, Any],
request_body: dict[str, Any],
request_headers: dict[str, Any],
is_stream: bool,
api_format: Optional[str] = None,
provider_request_headers: Optional[Dict[str, Any]] = None,
api_format: str | None = None,
provider_request_headers: dict[str, Any] | None = None,
input_tokens: int = 0,
output_tokens: int = 0,
cache_creation_tokens: int = 0,
cache_read_tokens: int = 0,
response_body: Optional[Dict[str, Any]] = None,
response_headers: Optional[Dict[str, Any]] = None,
client_response_headers: Optional[Dict[str, Any]] = None,
response_body: dict[str, Any] | None = None,
response_headers: dict[str, Any] | None = None,
client_response_headers: dict[str, Any] | None = None,
# 格式转换追踪
endpoint_api_format: Optional[str] = None,
endpoint_api_format: str | None = None,
has_format_conversion: bool = False,
target_model: Optional[str] = None,
target_model: str | None = None,
) -> None:
"""
记录客户端取消的请求
@@ -341,9 +343,9 @@ class MessageHandlerProtocol(Protocol):
self,
request: Any,
http_request: Request,
original_headers: Dict[str, str],
original_request_body: Dict[str, Any],
query_params: Optional[Dict[str, str]] = None,
original_headers: dict[str, str],
original_request_body: dict[str, Any],
query_params: dict[str, str] | None = None,
) -> StreamingResponse:
"""处理流式请求"""
...
@@ -352,9 +354,9 @@ class MessageHandlerProtocol(Protocol):
self,
request: Any,
http_request: Request,
original_headers: Dict[str, str],
original_request_body: Dict[str, Any],
query_params: Optional[Dict[str, str]] = None,
original_headers: dict[str, str],
original_request_body: dict[str, Any],
query_params: dict[str, str] | None = None,
) -> JSONResponse:
"""处理非流式请求"""
...
@@ -371,9 +373,6 @@ class BaseMessageHandler:
推荐使用 MessageHandlerProtocol 中定义的签名。
"""
# Adapter 检测器类型
AdapterDetectorType = Callable[[Dict[str, str], Optional[Dict[str, Any]]], Dict[str, bool]]
def __init__(
self,
*,
@@ -384,8 +383,8 @@ class BaseMessageHandler:
client_ip: str,
user_agent: str,
start_time: float,
allowed_api_formats: Optional[list[str]] = None,
adapter_detector: Optional[AdapterDetectorType] = None,
allowed_api_formats: list[str] | None = None,
adapter_detector: AdapterDetectorType | None = None,
) -> None:
self.db = db
self.user = user
@@ -408,9 +407,9 @@ class BaseMessageHandler:
def _resolve_capability_requirements(
self,
model_name: str,
request_headers: Optional[Dict[str, str]] = None,
request_body: Optional[Dict[str, Any]] = None,
) -> Dict[str, bool]:
request_headers: dict[str, str] | None = None,
request_body: dict[str, Any] | None = None,
) -> dict[str, bool]:
"""
解析请求的能力需求
@@ -442,12 +441,12 @@ class BaseMessageHandler:
async def _resolve_preferred_key_ids(
self,
model_name: str,
request_body: Optional[Dict[str, Any]] = None,
) -> Optional[list[str]]:
request_body: dict[str, Any] | None = None,
) -> list[str] | None:
"""可选的 Key 优先级解析钩子(默认不启用)。"""
return None
def get_api_format(self, provider_type: Optional[str] = None) -> APIFormat:
def get_api_format(self, provider_type: str | None = None) -> APIFormat:
"""根据 provider_type 解析 API 格式,未知类型默认 OPENAI"""
if provider_type:
result = resolve_api_format(provider_type, default=APIFormat.OPENAI)
@@ -456,17 +455,17 @@ class BaseMessageHandler:
def build_provider_payload(
self,
original_body: Dict[str, Any],
original_body: dict[str, Any],
*,
mapped_model: Optional[str] = None,
) -> Dict[str, Any]:
mapped_model: str | None = None,
) -> dict[str, Any]:
"""构建发送给 Provider 的请求体,替换 model 名称"""
payload = dict(original_body)
if mapped_model:
payload["model"] = mapped_model
return payload
def _update_usage_to_streaming(self, request_id: Optional[str] = None) -> None:
def _update_usage_to_streaming(self, request_id: str | None = None) -> None:
"""更新 Usage 状态为 streaming流式传输开始时调用
使用 asyncio 后台任务执行数据库更新,避免阻塞流式传输
@@ -500,7 +499,7 @@ class BaseMessageHandler:
# 创建后台任务,不阻塞当前流
asyncio.create_task(_do_update())
def _update_usage_to_streaming_with_ctx(self, ctx: "StreamContext") -> None:
def _update_usage_to_streaming_with_ctx(self, ctx: StreamContext) -> None:
"""更新 Usage 状态为 streaming同时更新 provider 相关信息
使用 asyncio 后台任务执行数据库更新,避免阻塞流式传输

View File

@@ -19,7 +19,7 @@ Chat Adapter 通用基类
import time
import traceback
from abc import abstractmethod
from typing import Any, Dict, Optional, Tuple, Type
from typing import Any
import httpx
from fastapi import HTTPException, Request
@@ -65,7 +65,7 @@ class ChatAdapterBase(ApiAdapter):
# 子类必须覆盖
FORMAT_ID: str = "UNKNOWN"
HANDLER_CLASS: Type[ChatHandlerBase]
HANDLER_CLASS: type[ChatHandlerBase]
# 适配器配置
name: str = "chat.base"
@@ -90,7 +90,7 @@ class ChatAdapterBase(ApiAdapter):
return base_url
@classmethod
def build_base_headers(cls, api_key: str) -> Dict[str, str]:
def build_base_headers(cls, api_key: str) -> dict[str, str]:
"""构建基础请求头,使用统一的 headers.py 实现"""
return build_adapter_base_headers(cls._get_api_format(), api_key)
@@ -101,13 +101,13 @@ class ChatAdapterBase(ApiAdapter):
@classmethod
def build_headers_with_extra(
cls, api_key: str, extra_headers: Optional[Dict[str, str]] = None
) -> Dict[str, str]:
cls, api_key: str, extra_headers: dict[str, str] | None = None
) -> dict[str, str]:
"""构建完整请求头(包含 extra_headers使用统一的 headers.py 实现"""
return build_adapter_headers(cls._get_api_format(), api_key, extra_headers)
@classmethod
def build_request_body(cls, request_data: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
def build_request_body(cls, request_data: dict[str, Any] | None = None) -> dict[str, Any]:
"""构建测试请求体,使用转换器注册表自动处理格式转换
Args:
@@ -120,11 +120,11 @@ class ChatAdapterBase(ApiAdapter):
return build_test_request_body(cls.FORMAT_ID, request_data)
def extract_api_key(self, request: Request) -> Optional[str]:
def extract_api_key(self, request: Request) -> str | None:
"""从请求中提取 API 密钥,使用统一的 headers.py 实现"""
return extract_client_api_key(dict(request.headers), self._get_api_format())
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
def __init__(self, allowed_api_formats: list[str] | None = None):
self.allowed_api_formats = allowed_api_formats or [self.FORMAT_ID]
async def handle(self, context: ApiRequestContext):
@@ -282,8 +282,8 @@ class ChatAdapterBase(ApiAdapter):
)
def _merge_path_params(
self, original_request_body: Dict[str, Any], path_params: Dict[str, Any]
) -> Dict[str, Any]:
self, original_request_body: dict[str, Any], path_params: dict[str, Any]
) -> dict[str, Any]:
"""
合并 URL 路径参数到请求体 - 子类可覆盖
@@ -316,7 +316,7 @@ class ChatAdapterBase(ApiAdapter):
"""
pass
def _extract_message_count(self, payload: Dict[str, Any], request_obj) -> int:
def _extract_message_count(self, payload: dict[str, Any], request_obj) -> int:
"""
提取消息数量 - 子类可覆盖
@@ -327,7 +327,7 @@ class ChatAdapterBase(ApiAdapter):
messages = request_obj.messages
return len(messages) if isinstance(messages, list) else 0
def _build_audit_metadata(self, payload: Dict[str, Any], request_obj) -> Dict[str, Any]:
def _build_audit_metadata(self, payload: dict[str, Any], request_obj) -> dict[str, Any]:
"""
构建审计日志元数据 - 子类可覆盖
"""
@@ -355,8 +355,8 @@ class ChatAdapterBase(ApiAdapter):
model: str,
stream: bool,
start_time: float,
original_headers: Dict[str, str],
original_request_body: Dict[str, Any],
original_headers: dict[str, str],
original_request_body: dict[str, Any],
client_ip: str,
request_id: str,
) -> JSONResponse:
@@ -426,8 +426,8 @@ class ChatAdapterBase(ApiAdapter):
model: str,
stream: bool,
start_time: float,
original_headers: Dict[str, str],
original_request_body: Dict[str, Any],
original_headers: dict[str, str],
original_request_body: dict[str, Any],
client_ip: str,
request_id: str,
) -> JSONResponse:
@@ -527,12 +527,12 @@ class ChatAdapterBase(ApiAdapter):
cache_read_input_tokens: int,
input_price_per_1m: float,
output_price_per_1m: float,
cache_creation_price_per_1m: Optional[float],
cache_read_price_per_1m: Optional[float],
price_per_request: Optional[float],
tiered_pricing: Optional[dict] = None,
cache_ttl_minutes: Optional[int] = None,
) -> Dict[str, Any]:
cache_creation_price_per_1m: float | None,
cache_read_price_per_1m: float | None,
price_per_request: float | None,
tiered_pricing: dict | None = None,
cache_ttl_minutes: int | None = None,
) -> dict[str, Any]:
"""
计算请求成本
@@ -597,8 +597,8 @@ class ChatAdapterBase(ApiAdapter):
client: httpx.AsyncClient,
base_url: str,
api_key: str,
extra_headers: Optional[Dict[str, str]] = None,
) -> Tuple[list, Optional[str]]:
extra_headers: dict[str, str] | None = None,
) -> tuple[list, str | None]:
"""
查询上游 API 支持的模型列表
@@ -626,16 +626,16 @@ class ChatAdapterBase(ApiAdapter):
client: httpx.AsyncClient,
base_url: str,
api_key: str,
request_data: Dict[str, Any],
extra_headers: Optional[Dict[str, str]] = None,
request_data: dict[str, Any],
extra_headers: dict[str, str] | None = None,
# 用量计算参数(现在强制记录)
db: Optional[Any] = None,
user: Optional[Any] = None,
provider_name: Optional[str] = None,
provider_id: Optional[str] = None,
api_key_id: Optional[str] = None,
model_name: Optional[str] = None,
) -> Dict[str, Any]:
db: Any | None = None,
user: Any | None = None,
provider_name: str | None = None,
provider_id: str | None = None,
api_key_id: str | None = None,
model_name: str | None = None,
) -> dict[str, Any]:
"""
测试模型连接性(非流式)
@@ -682,11 +682,11 @@ class ChatAdapterBase(ApiAdapter):
# Adapter 注册表 - 用于根据 API format 获取 Adapter 实例
# =========================================================================
_ADAPTER_REGISTRY: Dict[str, Type["ChatAdapterBase"]] = {}
_ADAPTER_REGISTRY: dict[str, type[ChatAdapterBase]] = {}
_ADAPTERS_LOADED = False
def register_adapter(adapter_class: Type["ChatAdapterBase"]) -> Type["ChatAdapterBase"]:
def register_adapter(adapter_class: type[ChatAdapterBase]) -> type[ChatAdapterBase]:
"""
注册 Adapter 类到注册表
@@ -731,7 +731,7 @@ def _ensure_adapters_loaded():
_ADAPTERS_LOADED = True
def get_adapter_class(api_format: str) -> Optional[Type["ChatAdapterBase"]]:
def get_adapter_class(api_format: str) -> type[ChatAdapterBase] | None:
"""
根据 API format 获取 Adapter 类
@@ -745,7 +745,7 @@ def get_adapter_class(api_format: str) -> Optional[Type["ChatAdapterBase"]]:
return _ADAPTER_REGISTRY.get(api_format.upper()) if api_format else None
def get_adapter_instance(api_format: str) -> Optional["ChatAdapterBase"]:
def get_adapter_instance(api_format: str) -> ChatAdapterBase | None:
"""
根据 API format 获取 Adapter 实例

View File

@@ -22,7 +22,10 @@ Chat Handler Base - Chat API 格式的通用基类
import asyncio
import json
from abc import ABC, abstractmethod
from typing import Any, AsyncGenerator, Awaitable, Callable, Dict, Optional, Union
from typing import Any
from collections.abc import Callable
from collections.abc import AsyncGenerator, Awaitable
import httpx
from fastapi import BackgroundTasks, Request
@@ -75,10 +78,10 @@ def _get_error_status_code(e: Exception, default: int = 400) -> int:
def _convert_error_response_best_effort(
error_response: Dict[str, Any],
error_response: dict[str, Any],
source_format: str,
target_format: str,
) -> Dict[str, Any]:
) -> dict[str, Any]:
"""
将上游错误响应 best-effort 转换为客户端格式。
@@ -97,7 +100,7 @@ def _convert_error_response_best_effort(
def _build_client_error_response_best_effort(
message: str,
target_format: str,
) -> Dict[str, Any]:
) -> dict[str, Any]:
"""
当无法解析上游错误 body 时构造一个目标格式的错误响应best-effort
"""
@@ -117,11 +120,11 @@ def _build_client_error_response_best_effort(
def _build_error_json_payload(
e: Union[ThinkingSignatureException, UpstreamClientException],
e: ThinkingSignatureException | UpstreamClientException,
client_format: str,
provider_format: str,
needs_conversion: bool = True,
) -> Dict[str, Any]:
) -> dict[str, Any]:
"""
构建错误 JSON 响应 payload公共逻辑
@@ -185,10 +188,10 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
client_ip: str,
user_agent: str,
start_time: float,
allowed_api_formats: Optional[list] = None,
adapter_detector: Optional[
Callable[[Dict[str, str], Optional[Dict[str, Any]]], Dict[str, bool]]
] = None,
allowed_api_formats: list | None = None,
adapter_detector: None | (
Callable[[dict[str, str], dict[str, Any] | None], dict[str, bool]]
) = None,
):
allowed = allowed_api_formats or [self.FORMAT_ID]
super().__init__(
@@ -202,7 +205,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
allowed_api_formats=allowed,
adapter_detector=adapter_detector,
)
self._parser: Optional[ResponseParser] = None
self._parser: ResponseParser | None = None
self._request_builder = PassthroughRequestBuilder()
@property
@@ -228,7 +231,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
pass
@abstractmethod
def _extract_usage(self, response: Dict) -> Dict[str, int]:
def _extract_usage(self, response: dict) -> dict[str, int]:
"""
从响应中提取 token 使用情况
@@ -241,7 +244,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
"""
pass
def _normalize_response(self, response: Dict) -> Dict:
def _normalize_response(self, response: dict) -> dict:
"""
规范化响应(可选覆盖)
@@ -257,8 +260,8 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
def extract_model_from_request(
self,
request_body: Dict[str, Any],
path_params: Optional[Dict[str, Any]] = None, # noqa: ARG002 - 子类使用
request_body: dict[str, Any],
path_params: dict[str, Any] | None = None, # noqa: ARG002 - 子类使用
) -> str:
"""
从请求中提取模型名 - 子类可覆盖
@@ -282,9 +285,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
def apply_mapped_model(
self,
request_body: Dict[str, Any],
request_body: dict[str, Any],
mapped_model: str, # noqa: ARG002 - 子类使用
) -> Dict[str, Any]:
) -> dict[str, Any]:
"""
将映射后的模型名应用到请求体
@@ -303,9 +306,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
def get_model_for_url(
self,
request_body: Dict[str, Any],
mapped_model: Optional[str],
) -> Optional[str]:
request_body: dict[str, Any],
mapped_model: str | None,
) -> str | None:
"""
获取用于 URL 路径的模型名
@@ -323,8 +326,8 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
def prepare_provider_request_body(
self,
request_body: Dict[str, Any],
) -> Dict[str, Any]:
request_body: dict[str, Any],
) -> dict[str, Any]:
"""
准备发送给 Provider 的请求体 - 子类可覆盖
@@ -341,9 +344,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
def _set_model_after_conversion(
self,
request_body: Dict[str, Any],
request_body: dict[str, Any],
provider_api_format: str,
mapped_model: Optional[str],
mapped_model: str | None,
fallback_model: str,
) -> None:
"""
@@ -372,7 +375,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
def _set_stream_after_conversion(
self,
request_body: Dict[str, Any],
request_body: dict[str, Any],
client_api_format: str,
provider_api_format: str,
is_stream: bool,
@@ -414,8 +417,8 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
self,
source_model: str,
provider_id: str,
api_format: Optional[str] = None,
) -> Optional[str]:
api_format: str | None = None,
) -> str | None:
"""
获取模型映射后的实际模型名
@@ -452,10 +455,10 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
self,
request: Any,
http_request: Request,
original_headers: Dict[str, Any],
original_request_body: Dict[str, Any],
query_params: Optional[Dict[str, str]] = None,
) -> Union[StreamingResponse, JSONResponse]:
original_headers: dict[str, Any],
original_request_body: dict[str, Any],
query_params: dict[str, str] | None = None,
) -> StreamingResponse | JSONResponse:
"""处理流式响应"""
logger.debug(f"开始流式响应处理 ({self.FORMAT_ID})")
@@ -466,7 +469,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
# 可变请求体容器:允许 orchestrator 在遇到 Thinking 签名错误时整流请求体后重试
# 结构: {"body": 实际请求体, "_rectified": 是否已整流, "_rectified_this_turn": 本轮是否整流}
request_body_ref: Dict[str, Any] = {"body": original_request_body}
request_body_ref: dict[str, Any] = {"body": original_request_body}
# 创建类型安全的流式上下文
ctx = StreamContext(model=model, api_format=api_format)
@@ -492,7 +495,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
endpoint: ProviderEndpoint,
key: ProviderAPIKey,
candidate: ProviderCandidate,
) -> AsyncGenerator[bytes, None]:
) -> AsyncGenerator[bytes]:
return await self._execute_stream_request(
ctx,
stream_processor,
@@ -615,12 +618,12 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
provider: Provider,
endpoint: ProviderEndpoint,
key: ProviderAPIKey,
original_request_body: Dict[str, Any],
original_headers: Dict[str, str],
query_params: Optional[Dict[str, str]] = None,
candidate: Optional[ProviderCandidate] = None,
is_disconnected: Optional[Callable[[], Awaitable[bool]]] = None,
) -> AsyncGenerator[bytes, None]:
original_request_body: dict[str, Any],
original_headers: dict[str, str],
query_params: dict[str, str] | None = None,
candidate: ProviderCandidate | None = None,
is_disconnected: Callable[[], Awaitable[bool]] | None = None,
) -> AsyncGenerator[bytes]:
"""执行流式请求并返回流生成器"""
# 重置上下文状态(重试时清除之前的数据)
ctx.reset_for_retry()
@@ -799,7 +802,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
ctx.error_message = "client_disconnected_during_prefetch"
raise
except asyncio.TimeoutError:
except TimeoutError:
# 整体请求超时(建立连接 + 获取首字节)
# 清理可能已建立的连接上下文
if response_ctx is not None:
@@ -856,8 +859,8 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
self,
ctx: StreamContext,
error: Exception,
original_headers: Dict[str, str],
original_request_body: Dict[str, Any],
original_headers: dict[str, str],
original_request_body: dict[str, Any],
) -> None:
"""记录流式请求失败"""
response_time_ms = self.elapsed_ms()
@@ -904,9 +907,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
self,
request: Any,
http_request: Request,
original_headers: Dict[str, Any],
original_request_body: Dict[str, Any],
query_params: Optional[Dict[str, str]] = None,
original_headers: dict[str, Any],
original_request_body: dict[str, Any],
query_params: dict[str, str] | None = None,
) -> JSONResponse:
"""处理非流式响应"""
logger.debug(f"开始非流式响应处理 ({self.FORMAT_ID})")
@@ -918,29 +921,29 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
# 可变请求体容器:允许 orchestrator 在遇到 Thinking 签名错误时整流请求体后重试
# 结构: {"body": 实际请求体, "_rectified": 是否已整流, "_rectified_this_turn": 本轮是否整流}
request_body_ref: Dict[str, Any] = {"body": original_request_body}
request_body_ref: dict[str, Any] = {"body": original_request_body}
# 用于跟踪的变量
provider_name: Optional[str] = None
response_json: Optional[Dict[str, Any]] = None
provider_name: str | None = None
response_json: dict[str, Any] | None = None
status_code = 200
response_headers: Dict[str, str] = {}
provider_request_headers: Dict[str, str] = {}
provider_request_body: Optional[Dict[str, Any]] = None
provider_api_format_for_error: Optional[str] = None
client_api_format_for_error: Optional[str] = None
response_headers: dict[str, str] = {}
provider_request_headers: dict[str, str] = {}
provider_request_body: dict[str, Any] | None = None
provider_api_format_for_error: str | None = None
client_api_format_for_error: str | None = None
needs_conversion_for_error: bool = False
provider_id: Optional[str] = None # Provider ID用于失败记录
endpoint_id: Optional[str] = None # Endpoint ID用于失败记录
key_id: Optional[str] = None # Key ID用于失败记录
mapped_model_result: Optional[str] = None # 映射后的目标模型名(用于 Usage 记录)
provider_id: str | None = None # Provider ID用于失败记录
endpoint_id: str | None = None # Endpoint ID用于失败记录
key_id: str | None = None # Key ID用于失败记录
mapped_model_result: str | None = None # 映射后的目标模型名(用于 Usage 记录)
async def sync_request_func(
provider: Provider,
endpoint: ProviderEndpoint,
key: ProviderAPIKey,
candidate: ProviderCandidate,
) -> Dict[str, Any]:
) -> dict[str, Any]:
nonlocal provider_name, response_json, status_code, response_headers
nonlocal provider_request_headers, provider_request_body, mapped_model_result
nonlocal provider_api_format_for_error, client_api_format_for_error, needs_conversion_for_error
@@ -1293,7 +1296,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
actual_request_body = provider_request_body or original_request_body
# 尝试从异常中提取响应头
error_response_headers: Dict[str, str] = {}
error_response_headers: dict[str, str] = {}
if isinstance(e, ProviderRateLimitException) and e.response_headers:
error_response_headers = e.response_headers
elif isinstance(e, httpx.HTTPStatusError) and hasattr(e, "response"):

View File

@@ -17,7 +17,7 @@ CLI Adapter 通用基类
import time
import traceback
from typing import Any, Dict, Optional, Tuple, Type
from typing import Any
import httpx
from fastapi import HTTPException, Request
@@ -63,7 +63,7 @@ class CliAdapterBase(ApiAdapter):
# 子类必须覆盖
FORMAT_ID: str = "UNKNOWN"
HANDLER_CLASS: Type[CliMessageHandlerBase]
HANDLER_CLASS: type[CliMessageHandlerBase]
# 适配器配置
name: str = "cli.base"
@@ -72,7 +72,7 @@ class CliAdapterBase(ApiAdapter):
# 计费模板配置(子类可覆盖,如 "claude", "openai", "gemini"
BILLING_TEMPLATE: str = "claude"
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
def __init__(self, allowed_api_formats: list[str] | None = None):
self.allowed_api_formats = allowed_api_formats or [self.FORMAT_ID]
# =========================================================================
@@ -87,7 +87,7 @@ class CliAdapterBase(ApiAdapter):
except KeyError:
return APIFormat.OPENAI
def extract_api_key(self, request: Request) -> Optional[str]:
def extract_api_key(self, request: Request) -> str | None:
"""
从请求中提取 API 密钥
@@ -96,7 +96,7 @@ class CliAdapterBase(ApiAdapter):
return extract_client_api_key(dict(request.headers), self._get_api_format())
@classmethod
def build_base_headers(cls, api_key: str) -> Dict[str, str]:
def build_base_headers(cls, api_key: str) -> dict[str, str]:
"""
构建 CLI API 认证头
@@ -106,8 +106,8 @@ class CliAdapterBase(ApiAdapter):
@classmethod
def build_headers_with_extra(
cls, api_key: str, extra_headers: Optional[Dict[str, str]] = None
) -> Dict[str, str]:
cls, api_key: str, extra_headers: dict[str, str] | None = None
) -> dict[str, str]:
"""
构建带额外头部的完整请求头
@@ -260,8 +260,8 @@ class CliAdapterBase(ApiAdapter):
)
def _merge_path_params(
self, original_request_body: Dict[str, Any], path_params: Dict[str, Any]
) -> Dict[str, Any]:
self, original_request_body: dict[str, Any], path_params: dict[str, Any]
) -> dict[str, Any]:
"""
合并 URL 路径参数到请求体 - 子类可覆盖
@@ -280,7 +280,7 @@ class CliAdapterBase(ApiAdapter):
merged[key] = value
return merged
def _extract_message_count(self, payload: Dict[str, Any]) -> int:
def _extract_message_count(self, payload: dict[str, Any]) -> int:
"""
提取消息数量 - 子类可覆盖
@@ -297,9 +297,9 @@ class CliAdapterBase(ApiAdapter):
def _build_audit_metadata(
self,
payload: Dict[str, Any],
path_params: Optional[Dict[str, Any]] = None,
) -> Dict[str, Any]:
payload: dict[str, Any],
path_params: dict[str, Any] | None = None,
) -> dict[str, Any]:
"""
构建审计日志元数据 - 子类可覆盖
@@ -338,8 +338,8 @@ class CliAdapterBase(ApiAdapter):
model: str,
stream: bool,
start_time: float,
original_headers: Dict[str, str],
original_request_body: Dict[str, Any],
original_headers: dict[str, str],
original_request_body: dict[str, Any],
client_ip: str,
request_id: str,
) -> JSONResponse:
@@ -409,8 +409,8 @@ class CliAdapterBase(ApiAdapter):
model: str,
stream: bool,
start_time: float,
original_headers: Dict[str, str],
original_request_body: Dict[str, Any],
original_headers: dict[str, str],
original_request_body: dict[str, Any],
client_ip: str,
request_id: str,
) -> JSONResponse:
@@ -507,12 +507,12 @@ class CliAdapterBase(ApiAdapter):
cache_read_input_tokens: int,
input_price_per_1m: float,
output_price_per_1m: float,
cache_creation_price_per_1m: Optional[float],
cache_read_price_per_1m: Optional[float],
price_per_request: Optional[float],
tiered_pricing: Optional[dict] = None,
cache_ttl_minutes: Optional[int] = None,
) -> Dict[str, Any]:
cache_creation_price_per_1m: float | None,
cache_read_price_per_1m: float | None,
price_per_request: float | None,
tiered_pricing: dict | None = None,
cache_ttl_minutes: int | None = None,
) -> dict[str, Any]:
"""
计算请求成本
@@ -567,8 +567,8 @@ class CliAdapterBase(ApiAdapter):
client: httpx.AsyncClient,
base_url: str,
api_key: str,
extra_headers: Optional[Dict[str, str]] = None,
) -> Tuple[list, Optional[str]]:
extra_headers: dict[str, str] | None = None,
) -> tuple[list, str | None]:
"""
查询上游 API 支持的模型列表
@@ -596,16 +596,16 @@ class CliAdapterBase(ApiAdapter):
client: httpx.AsyncClient,
base_url: str,
api_key: str,
request_data: Dict[str, Any],
extra_headers: Optional[Dict[str, str]] = None,
request_data: dict[str, Any],
extra_headers: dict[str, str] | None = None,
# 用量计算参数
db: Optional[Any] = None,
user: Optional[Any] = None,
provider_name: Optional[str] = None,
provider_id: Optional[str] = None,
api_key_id: Optional[str] = None,
model_name: Optional[str] = None,
) -> Dict[str, Any]:
db: Any | None = None,
user: Any | None = None,
provider_name: str | None = None,
provider_id: str | None = None,
api_key_id: str | None = None,
model_name: str | None = None,
) -> dict[str, Any]:
"""
测试模型连接性(非流式)
@@ -669,7 +669,7 @@ class CliAdapterBase(ApiAdapter):
# =========================================================================
@classmethod
def build_endpoint_url(cls, base_url: str, request_data: Dict[str, Any], model_name: Optional[str] = None) -> str:
def build_endpoint_url(cls, base_url: str, request_data: dict[str, Any], model_name: str | None = None) -> str:
"""
构建CLI API端点URL - 子类应覆盖
@@ -684,7 +684,7 @@ class CliAdapterBase(ApiAdapter):
raise NotImplementedError(f"{cls.FORMAT_ID} adapter must implement build_endpoint_url")
@classmethod
def build_request_body(cls, request_data: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
def build_request_body(cls, request_data: dict[str, Any] | None = None) -> dict[str, Any]:
"""构建测试请求体,使用转换器注册表自动处理格式转换
Args:
@@ -698,7 +698,7 @@ class CliAdapterBase(ApiAdapter):
return build_test_request_body(cls.FORMAT_ID, request_data)
@classmethod
def get_cli_user_agent(cls) -> Optional[str]:
def get_cli_user_agent(cls) -> str | None:
"""
获取CLI User-Agent - 子类可覆盖
@@ -708,7 +708,7 @@ class CliAdapterBase(ApiAdapter):
return None
@classmethod
def get_cli_extra_headers(cls) -> Dict[str, str]:
def get_cli_extra_headers(cls) -> dict[str, str]:
"""
获取CLI额外请求头 - 子类可覆盖
@@ -718,7 +718,7 @@ class CliAdapterBase(ApiAdapter):
Returns:
额外请求头字典
"""
headers: Dict[str, str] = {}
headers: dict[str, str] = {}
cli_user_agent = cls.get_cli_user_agent()
if cli_user_agent:
headers["User-Agent"] = cli_user_agent
@@ -728,11 +728,11 @@ class CliAdapterBase(ApiAdapter):
# CLI Adapter 注册表 - 用于根据 API format 获取 CLI Adapter 实例
# =========================================================================
_CLI_ADAPTER_REGISTRY: Dict[str, Type["CliAdapterBase"]] = {}
_CLI_ADAPTER_REGISTRY: dict[str, type[CliAdapterBase]] = {}
_CLI_ADAPTERS_LOADED = False
def register_cli_adapter(adapter_class: Type["CliAdapterBase"]) -> Type["CliAdapterBase"]:
def register_cli_adapter(adapter_class: type[CliAdapterBase]) -> type[CliAdapterBase]:
"""
注册 CLI Adapter 类到注册表
@@ -771,13 +771,13 @@ def _ensure_cli_adapters_loaded():
_CLI_ADAPTERS_LOADED = True
def get_cli_adapter_class(api_format: str) -> Optional[Type["CliAdapterBase"]]:
def get_cli_adapter_class(api_format: str) -> type[CliAdapterBase] | None:
"""根据 API format 获取 CLI Adapter 类"""
_ensure_cli_adapters_loaded()
return _CLI_ADAPTER_REGISTRY.get(api_format.upper()) if api_format else None
def get_cli_adapter_instance(api_format: str) -> Optional["CliAdapterBase"]:
def get_cli_adapter_instance(api_format: str) -> CliAdapterBase | None:
"""根据 API format 获取 CLI Adapter 实例"""
adapter_class = get_cli_adapter_class(api_format)
if adapter_class:

View File

@@ -10,6 +10,8 @@ CLI Message Handler 通用基类
3. 简化新格式接入 - 只需实现 ResponseParser 和少量钩子方法
"""
from __future__ import annotations
import asyncio
import codecs
import json
@@ -17,14 +19,11 @@ import time
from typing import (
TYPE_CHECKING,
Any,
AsyncGenerator,
Callable,
Dict,
List,
Optional,
Tuple,
)
from collections.abc import Callable
from collections.abc import AsyncGenerator
import httpx
from fastapi import BackgroundTasks, Request
from fastapi.responses import JSONResponse, StreamingResponse
@@ -45,7 +44,6 @@ from src.api.handlers.base.request_builder import PassthroughRequestBuilder, get
# 直接从具体模块导入,避免循环依赖
from src.api.handlers.base.response_parser import (
ResponseParser,
StreamStats,
)
from src.api.handlers.base.stream_context import StreamContext
from src.api.handlers.base.utils import (
@@ -86,7 +84,7 @@ from src.utils.timeout import read_first_chunk_with_ttfb_timeout
# ==============================================================================
def _parse_sse_data_line(line: str) -> Tuple[Optional[Any], str]:
def _parse_sse_data_line(line: str) -> tuple[Any | None, str]:
"""
解析标准 SSE data 行
@@ -108,7 +106,7 @@ def _parse_sse_data_line(line: str) -> Tuple[Optional[Any], str]:
return None, "invalid"
def _parse_sse_event_data_line(line: str) -> Tuple[Optional[Any], str]:
def _parse_sse_event_data_line(line: str) -> tuple[Any | None, str]:
"""
解析 event + data 同行格式(如 "event: xxx data: {...}"
@@ -126,7 +124,7 @@ def _parse_sse_event_data_line(line: str) -> Tuple[Optional[Any], str]:
return None, "invalid"
def _parse_gemini_json_array_line(line: str) -> Tuple[Optional[Any], str]:
def _parse_gemini_json_array_line(line: str) -> tuple[Any | None, str]:
"""
解析 Gemini JSON-array 格式的裸 JSON 行
@@ -151,9 +149,9 @@ def _parse_gemini_json_array_line(line: str) -> Tuple[Optional[Any], str]:
def _format_converted_events_to_sse(
converted_events: List[Dict[str, Any]],
converted_events: list[dict[str, Any]],
client_format: str,
) -> List[str]:
) -> list[str]:
"""
将转换后的事件格式化为 SSE 行
@@ -164,7 +162,7 @@ def _format_converted_events_to_sse(
Returns:
SSE 行列表(每个元素是完整的 SSE 事件,包含尾部空行)
"""
result: List[str] = []
result: list[str] = []
needs_event_line = client_format.upper() in ("CLAUDE", "CLAUDE_CLI")
for evt in converted_events:
@@ -213,10 +211,10 @@ class CliMessageHandlerBase(BaseMessageHandler):
client_ip: str,
user_agent: str,
start_time: float,
allowed_api_formats: Optional[list] = None,
adapter_detector: Optional[
Callable[[Dict[str, str], Optional[Dict[str, Any]]], Dict[str, bool]]
] = None,
allowed_api_formats: list | None = None,
adapter_detector: None | (
Callable[[dict[str, str], dict[str, Any] | None], dict[str, bool]]
) = None,
):
allowed = allowed_api_formats or [self.FORMAT_ID]
super().__init__(
@@ -230,7 +228,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
allowed_api_formats=allowed,
adapter_detector=adapter_detector,
)
self._parser: Optional[ResponseParser] = None
self._parser: ResponseParser | None = None
self._request_builder = PassthroughRequestBuilder()
@property
@@ -253,7 +251,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
self,
source_model: str,
provider_id: str,
) -> Optional[str]:
) -> str | None:
"""
获取模型映射后的实际模型名
@@ -296,8 +294,8 @@ class CliMessageHandlerBase(BaseMessageHandler):
def extract_model_from_request(
self,
request_body: Dict[str, Any],
path_params: Optional[Dict[str, Any]] = None, # noqa: ARG002 - 子类使用
request_body: dict[str, Any],
path_params: dict[str, Any] | None = None, # noqa: ARG002 - 子类使用
) -> str:
"""
从请求中提取模型名 - 子类可覆盖
@@ -321,9 +319,9 @@ class CliMessageHandlerBase(BaseMessageHandler):
def apply_mapped_model(
self,
request_body: Dict[str, Any],
request_body: dict[str, Any],
mapped_model: str, # noqa: ARG002 - 子类使用
) -> Dict[str, Any]:
) -> dict[str, Any]:
"""
将映射后的模型名应用到请求体
@@ -342,8 +340,8 @@ class CliMessageHandlerBase(BaseMessageHandler):
def prepare_provider_request_body(
self,
request_body: Dict[str, Any],
) -> Dict[str, Any]:
request_body: dict[str, Any],
) -> dict[str, Any]:
"""
准备发送给 Provider 的请求体 - 子类可覆盖
@@ -359,7 +357,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
return request_body
@staticmethod
def _get_format_metadata(format_id: str) -> Optional["ApiFormatDefinition"]:
def _get_format_metadata(format_id: str) -> ApiFormatDefinition | None:
"""获取格式元数据(解析失败返回 None"""
from src.core.api_format import APIFormat
from src.core.api_format.metadata import API_FORMAT_DEFINITIONS
@@ -372,10 +370,10 @@ class CliMessageHandlerBase(BaseMessageHandler):
def _finalize_converted_request(
self,
request_body: Dict[str, Any],
request_body: dict[str, Any],
client_api_format: str,
provider_api_format: str,
mapped_model: Optional[str],
mapped_model: str | None,
fallback_model: str,
is_stream: bool,
) -> None:
@@ -418,13 +416,13 @@ class CliMessageHandlerBase(BaseMessageHandler):
def _convert_request_for_cross_format(
self,
request_body: Dict[str, Any],
request_body: dict[str, Any],
client_api_format: str,
provider_api_format: str,
mapped_model: Optional[str],
mapped_model: str | None,
fallback_model: str,
is_stream: bool,
) -> Tuple[Dict[str, Any], str]:
) -> tuple[dict[str, Any], str]:
"""
跨格式请求转换的公共逻辑
@@ -465,9 +463,9 @@ class CliMessageHandlerBase(BaseMessageHandler):
def get_model_for_url(
self,
request_body: Dict[str, Any],
mapped_model: Optional[str],
) -> Optional[str]:
request_body: dict[str, Any],
mapped_model: str | None,
) -> str | None:
"""
获取用于 URL 路径的模型名
@@ -485,8 +483,8 @@ class CliMessageHandlerBase(BaseMessageHandler):
def _extract_response_metadata(
self,
response: Dict[str, Any],
) -> Dict[str, Any]:
response: dict[str, Any],
) -> dict[str, Any]:
"""
从响应中提取 Provider 特有的元数据 - 子类可覆盖
@@ -503,11 +501,11 @@ class CliMessageHandlerBase(BaseMessageHandler):
async def process_stream(
self,
original_request_body: Dict[str, Any],
original_headers: Dict[str, str],
query_params: Optional[Dict[str, str]] = None,
path_params: Optional[Dict[str, Any]] = None,
http_request: Optional[Request] = None,
original_request_body: dict[str, Any],
original_headers: dict[str, str],
query_params: dict[str, str] | None = None,
path_params: dict[str, Any] | None = None,
http_request: Request | None = None,
) -> StreamingResponse:
"""
处理流式请求
@@ -529,7 +527,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
# 可变请求体容器:允许 orchestrator 在遇到 Thinking 签名错误时整流请求体后重试
# 结构: {"body": 实际请求体, "_rectified": 是否已整流, "_rectified_this_turn": 本轮是否整流}
request_body_ref: Dict[str, Any] = {"body": original_request_body}
request_body_ref: dict[str, Any] = {"body": original_request_body}
# 使用子类实现的方法提取 model不同 API 格式的 model 位置不同)
# 注意:使用 original_request_body因为整流只修改 messages不影响 model 字段
@@ -550,7 +548,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
endpoint: ProviderEndpoint,
key: ProviderAPIKey,
candidate: ProviderCandidate,
) -> AsyncGenerator[bytes, None]:
) -> AsyncGenerator[bytes]:
return await self._execute_stream_request(
ctx,
provider,
@@ -653,12 +651,12 @@ class CliMessageHandlerBase(BaseMessageHandler):
provider: Provider,
endpoint: ProviderEndpoint,
key: ProviderAPIKey,
original_request_body: Dict[str, Any],
original_headers: Dict[str, str],
query_params: Optional[Dict[str, str]] = None,
candidate: Optional[ProviderCandidate] = None,
http_request: Optional[Request] = None,
) -> AsyncGenerator[bytes, None]:
original_request_body: dict[str, Any],
original_headers: dict[str, str],
query_params: dict[str, str] | None = None,
candidate: ProviderCandidate | None = None,
http_request: Request | None = None,
) -> AsyncGenerator[bytes]:
"""执行流式请求并返回流生成器"""
# 重置上下文状态(重试时清除之前的数据,避免累积)
ctx.parsed_chunks = []
@@ -824,7 +822,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
else:
await asyncio.wait_for(_connect_and_prefetch(), timeout=request_timeout)
except asyncio.TimeoutError:
except TimeoutError:
# 整体请求超时(建立连接 + 获取首字节)
# 清理可能已建立的连接上下文
if response_ctx is not None:
@@ -898,7 +896,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
stream_response: httpx.Response,
response_ctx: Any,
http_client: httpx.AsyncClient,
) -> AsyncGenerator[bytes, None]:
) -> AsyncGenerator[bytes]:
"""创建响应流生成器(使用字节流)"""
try:
sse_parser = SSEEventParser()
@@ -961,9 +959,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
},
}
self._mark_first_output(ctx, output_state)
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode(
"utf-8"
)
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
return # 结束生成器
# 格式转换或直接透传
@@ -1015,7 +1011,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
"message": ctx.error_message,
},
}
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode("utf-8")
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
else:
logger.debug("流式数据转发完成")
# 为 OpenAI 客户端补齐 [DONE] 标记(非 CLI 格式)
@@ -1040,7 +1036,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
"message": ctx.error_message,
},
}
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode("utf-8")
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
except httpx.RemoteProtocolError:
if ctx.data_count > 0:
error_event = {
@@ -1050,7 +1046,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
"message": "上游连接意外关闭,部分响应已成功传输",
},
}
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode("utf-8")
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
else:
raise
finally:
@@ -1241,7 +1237,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
except (EmbeddedErrorException, ProviderTimeoutException, ProviderNotAvailableException):
# 重新抛出可重试的 Provider 异常,触发故障转移
raise
except (OSError, IOError) as e:
except OSError as e:
# 网络 I/O 异常:记录警告,可能需要重试
logger.warning(f" [{self.request_id}] 预读流时发生网络异常: {type(e).__name__}: {e}")
except Exception as e:
@@ -1261,7 +1257,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
response_ctx: Any,
http_client: httpx.AsyncClient,
prefetched_chunks: list,
) -> AsyncGenerator[bytes, None]:
) -> AsyncGenerator[bytes]:
"""创建响应流生成器(带预读数据,使用字节流)"""
try:
sse_parser = SSEEventParser()
@@ -1382,9 +1378,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
},
}
self._mark_first_output(ctx, output_state)
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode(
"utf-8"
)
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
return
# 格式转换或直接透传
@@ -1439,7 +1433,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
"message": ctx.error_message,
},
}
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode("utf-8")
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
else:
logger.debug("流式数据转发完成")
# 为 OpenAI 客户端补齐 [DONE] 标记(非 CLI 格式)
@@ -1463,7 +1457,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
"message": ctx.error_message,
},
}
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode("utf-8")
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
except httpx.RemoteProtocolError:
if ctx.data_count > 0:
error_event = {
@@ -1473,7 +1467,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
"message": "上游连接意外关闭,部分响应已成功传输",
},
}
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode("utf-8")
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
else:
raise
finally:
@@ -1489,7 +1483,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
def _handle_sse_event(
self,
ctx: StreamContext,
event_name: Optional[str],
event_name: str | None,
data_str: str,
record_chunk: bool = False,
) -> None:
@@ -1538,7 +1532,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
self,
ctx: StreamContext,
event_type: str,
data: Dict[str, Any],
data: dict[str, Any],
) -> None:
"""
处理解析后的事件数据 - 子类应覆盖此方法
@@ -1612,7 +1606,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
def _record_converted_chunks(
self,
ctx: StreamContext,
converted_events: List[Dict[str, Any]],
converted_events: list[dict[str, Any]],
) -> None:
"""
记录转换后的 chunk 数据到 parsed_chunks并更新统计信息
@@ -1656,7 +1650,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
def _extract_usage_from_converted_event(
self,
ctx: StreamContext,
evt: Dict[str, Any],
evt: dict[str, Any],
event_type: str,
) -> None:
"""
@@ -1672,7 +1666,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
evt: 转换后的事件
event_type: 事件类型
"""
usage: Optional[Dict[str, Any]] = None
usage: dict[str, Any] | None = None
# Claude 格式: message_delta 或 message_start
if event_type == "message_delta":
@@ -1737,9 +1731,9 @@ class CliMessageHandlerBase(BaseMessageHandler):
async def _create_monitored_stream(
self,
ctx: StreamContext,
stream_generator: AsyncGenerator[bytes, None],
http_request: Optional[Request] = None,
) -> AsyncGenerator[bytes, None]:
stream_generator: AsyncGenerator[bytes],
http_request: Request | None = None,
) -> AsyncGenerator[bytes]:
"""
创建带监控的流生成器
@@ -1833,8 +1827,8 @@ class CliMessageHandlerBase(BaseMessageHandler):
async def _record_stream_stats(
self,
ctx: StreamContext,
original_headers: Dict[str, str],
original_request_body: Dict[str, Any],
original_headers: dict[str, str],
original_request_body: dict[str, Any],
) -> None:
"""在流完成后记录统计信息"""
try:
@@ -1996,7 +1990,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
from src.services.request.candidate import RequestCandidateService
# 计算候选自身的 TTFB
candidate_first_byte_time_ms: Optional[int] = None
candidate_first_byte_time_ms: int | None = None
if ctx.first_byte_time_ms is not None:
candidate_first_byte_time_ms = (
RequestCandidateService.calculate_candidate_ttfb(
@@ -2061,8 +2055,8 @@ class CliMessageHandlerBase(BaseMessageHandler):
self,
ctx: StreamContext,
error: Exception,
original_headers: Dict[str, str],
original_request_body: Dict[str, Any],
original_headers: dict[str, str],
original_request_body: dict[str, Any],
) -> None:
"""记录流式请求失败"""
# 使用 self.start_time 作为时间基准,与首字时间保持一致
@@ -2111,10 +2105,10 @@ class CliMessageHandlerBase(BaseMessageHandler):
async def process_sync(
self,
original_request_body: Dict[str, Any],
original_headers: Dict[str, str],
query_params: Optional[Dict[str, str]] = None,
path_params: Optional[Dict[str, Any]] = None,
original_request_body: dict[str, Any],
original_headers: dict[str, str],
query_params: dict[str, str] | None = None,
path_params: dict[str, Any] | None = None,
) -> JSONResponse:
"""
处理非流式请求
@@ -2142,19 +2136,19 @@ class CliMessageHandlerBase(BaseMessageHandler):
endpoint_id = None # Endpoint ID用于失败记录
key_id = None # Key ID用于失败记录
mapped_model_result = None # 映射后的目标模型名(用于 Usage 记录)
response_metadata_result: Dict[str, Any] = {} # Provider 响应元数据
response_metadata_result: dict[str, Any] = {} # Provider 响应元数据
needs_conversion = False # 是否需要格式转换(由 candidate 决定)
# 可变请求体容器:允许 orchestrator 在遇到 Thinking 签名错误时整流请求体后重试
# 结构: {"body": 实际请求体, "_rectified": 是否已整流, "_rectified_this_turn": 本轮是否整流}
request_body_ref: Dict[str, Any] = {"body": original_request_body}
request_body_ref: dict[str, Any] = {"body": original_request_body}
async def sync_request_func(
provider: Provider,
endpoint: ProviderEndpoint,
key: ProviderAPIKey,
candidate: ProviderCandidate,
) -> Dict[str, Any]:
) -> 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, needs_conversion
provider_name = str(provider.name)
provider_api_format = str(endpoint.api_format) if endpoint.api_format else ""
@@ -2470,7 +2464,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
actual_request_body = provider_request_body or original_request_body
# 尝试从异常中提取响应头
error_response_headers: Dict[str, str] = {}
error_response_headers: dict[str, str] = {}
if isinstance(e, ProviderRateLimitException) and e.response_headers:
error_response_headers = e.response_headers
elif isinstance(e, httpx.HTTPStatusError) and hasattr(e, "response"):
@@ -2581,7 +2575,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
)
return True
def _mark_first_output(self, ctx: StreamContext, state: Dict[str, bool]) -> None:
def _mark_first_output(self, ctx: StreamContext, state: dict[str, bool]) -> None:
"""
标记首次输出:记录 TTFB 并更新 streaming 状态
@@ -2605,7 +2599,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
ctx: StreamContext,
line: str,
events: list, # noqa: ARG002 - 预留给上下文感知转换
) -> Tuple[List[str], List[Dict[str, Any]]]:
) -> tuple[list[str], list[dict[str, Any]]]:
"""
将 SSE 行从 Provider 格式转换为客户端格式
@@ -2690,7 +2684,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
def _parse_sse_line_to_json(
self, line: str, provider_format: str
) -> Tuple[Optional[Any], str]:
) -> tuple[Any | None, str]:
"""
解析 SSE 行为 JSON 对象

View File

@@ -8,7 +8,6 @@ StreamSmoother 使用这些提取器来处理不同格式的 SSE 事件。
import copy
import json
from abc import ABC, abstractmethod
from typing import Optional
class ContentExtractor(ABC):
@@ -20,7 +19,7 @@ class ContentExtractor(ABC):
"""
@abstractmethod
def extract_content(self, data: dict) -> Optional[str]:
def extract_content(self, data: dict) -> str | None:
"""
从 SSE 数据中提取可拆分的文本内容
@@ -64,7 +63,7 @@ class OpenAIContentExtractor(ContentExtractor):
- 只在 delta 仅包含 role/content 时允许拆分,避免破坏 tool_calls 等结构
"""
def extract_content(self, data: dict) -> Optional[str]:
def extract_content(self, data: dict) -> str | None:
if not isinstance(data, dict):
return None
@@ -115,7 +114,7 @@ class OpenAIContentExtractor(ContentExtractor):
new_choices.append(new_choice)
new_data["choices"] = new_choices
return f"data: {json.dumps(new_data, ensure_ascii=False)}\n\n".encode("utf-8")
return f"data: {json.dumps(new_data, ensure_ascii=False)}\n\n".encode()
class ClaudeContentExtractor(ContentExtractor):
@@ -127,7 +126,7 @@ class ClaudeContentExtractor(ContentExtractor):
- 数据结构: delta.type=text_delta, delta.text
"""
def extract_content(self, data: dict) -> Optional[str]:
def extract_content(self, data: dict) -> str | None:
if not isinstance(data, dict):
return None
@@ -165,9 +164,7 @@ class ClaudeContentExtractor(ContentExtractor):
# Claude 格式需要 event: 前缀
event_name = event_type or "content_block_delta"
return f"event: {event_name}\ndata: {json.dumps(new_data, ensure_ascii=False)}\n\n".encode(
"utf-8"
)
return f"event: {event_name}\ndata: {json.dumps(new_data, ensure_ascii=False)}\n\n".encode()
class GeminiContentExtractor(ContentExtractor):
@@ -179,7 +176,7 @@ class GeminiContentExtractor(ContentExtractor):
- 只有纯文本块才拆分
"""
def extract_content(self, data: dict) -> Optional[str]:
def extract_content(self, data: dict) -> str | None:
if not isinstance(data, dict):
return None
@@ -226,7 +223,7 @@ class GeminiContentExtractor(ContentExtractor):
if "parts" in content and content["parts"]:
content["parts"][0]["text"] = new_content
return f"data: {json.dumps(new_data, ensure_ascii=False)}\n\n".encode("utf-8")
return f"data: {json.dumps(new_data, ensure_ascii=False)}\n\n".encode()
# 提取器注册表
@@ -237,7 +234,7 @@ _EXTRACTORS: dict[str, type[ContentExtractor]] = {
}
def get_extractor(format_name: str) -> Optional[ContentExtractor]:
def get_extractor(format_name: str) -> ContentExtractor | None:
"""
根据格式名获取对应的内容提取器实例

View File

@@ -14,15 +14,14 @@
- EndpointCheckOrchestrator: 协调整个流程
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import Any, AsyncIterator, Dict, Iterable, Optional, Union, List
from abc import ABC, abstractmethod
from typing import Any
from collections.abc import Iterable
import time
import uuid
import json
from functools import lru_cache
import asyncio
from collections import defaultdict
import httpx
@@ -31,7 +30,7 @@ from src.core.api_format import CORE_REDACT_HEADERS, merge_headers_with_protecti
from src.utils.ssl_utils import get_ssl_context
def _redact_headers(headers: Dict[str, str]) -> Dict[str, str]:
def _redact_headers(headers: dict[str, str]) -> dict[str, str]:
return redact_headers_for_log(headers, CORE_REDACT_HEADERS)
@@ -46,10 +45,10 @@ def _truncate_repr(value: Any, limit: int = 1200) -> str:
def build_safe_headers(
base_headers: Dict[str, str],
extra_headers: Optional[Dict[str, str]],
base_headers: dict[str, str],
extra_headers: dict[str, str] | None,
protected_keys: Iterable[str],
) -> Dict[str, str]:
) -> dict[str, str]:
"""
合并 extra_headers但防止覆盖 protected_keys大小写不敏感
"""
@@ -60,16 +59,16 @@ async def run_endpoint_check(
*,
client: httpx.AsyncClient, # 保持兼容性,但内部不使用
url: str,
headers: Dict[str, str],
json_body: Dict[str, Any],
headers: dict[str, str],
json_body: dict[str, Any],
api_format: str,
provider_name: Optional[str] = None,
model_name: Optional[str] = None,
api_key_id: Optional[str] = None,
provider_id: Optional[str] = None,
db: Optional[Any] = None, # Session对象需要时才导入
user: Optional[Any] = None, # User对象
) -> Dict[str, Any]:
provider_name: str | None = None,
model_name: str | None = None,
api_key_id: str | None = None,
provider_id: str | None = None,
db: Any | None = None, # Session对象需要时才导入
user: Any | None = None, # User对象
) -> dict[str, Any]:
"""
执行端点检查(重构版本,使用新的架构):
- 使用新的架构类来分离关注点
@@ -123,21 +122,21 @@ async def _calculate_and_record_usage(
provider_id: str,
api_key_id: str,
model_name: str,
request_data: Dict[str, Any],
response_data: Optional[Dict[str, Any]],
request_data: dict[str, Any],
response_data: dict[str, Any] | None,
request_id: str,
response_time_ms: int,
request_headers: Dict[str, str],
response_headers: Optional[Dict[str, str]] = None,
request_headers: dict[str, str],
response_headers: dict[str, str] | None = None,
status_code: int = 0,
error_message: Optional[str] = None,
error_message: str | None = None,
# 新增支持直接传递token数据
input_tokens: Optional[int] = None,
output_tokens: Optional[int] = None,
cache_creation_input_tokens: Optional[int] = None,
cache_read_input_tokens: Optional[int] = None,
api_format: Optional[str] = None,
) -> Dict[str, Any]:
input_tokens: int | None = None,
output_tokens: int | None = None,
cache_creation_input_tokens: int | None = None,
cache_read_input_tokens: int | None = None,
api_format: str | None = None,
) -> dict[str, Any]:
"""
计算并记录用量数据(遗留函数)
@@ -149,7 +148,7 @@ async def _calculate_and_record_usage(
"""
from src.services.usage.service import UsageService
from src.services.request.candidate import RequestCandidateService
from src.models.database import ApiKey, ProviderAPIKey, ProviderEndpoint
from src.models.database import ApiKey, ProviderAPIKey
# 获取Provider API Key对象不是用户API Key
provider_api_key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == api_key_id).first()
@@ -360,7 +359,7 @@ async def _calculate_and_record_usage(
}
def _extract_tokens_from_response(api_identifier: str, response_data: Optional[Dict[str, Any]]) -> tuple[int, int, int, int]:
def _extract_tokens_from_response(api_identifier: str, response_data: dict[str, Any] | None) -> tuple[int, int, int, int]:
"""
从响应中提取Token计数信息
@@ -446,7 +445,7 @@ def _extract_tokens_from_response(api_identifier: str, response_data: Optional[D
def _fallback_token_counting(request_data: Dict[str, Any], response_data: Optional[Dict[str, Any]]) -> tuple[int, int, int, int]:
def _fallback_token_counting(request_data: dict[str, Any], response_data: dict[str, Any] | None) -> tuple[int, int, int, int]:
"""
回退的Token计数方法简单估算
@@ -508,16 +507,16 @@ def _fallback_token_counting(request_data: Dict[str, Any], response_data: Option
class EndpointCheckRequest:
"""端点检查请求数据类"""
url: str
headers: Dict[str, str]
json_body: Dict[str, Any]
headers: dict[str, str]
json_body: dict[str, Any]
api_format: str
provider_name: Optional[str] = None
model_name: Optional[str] = None
api_key_id: Optional[str] = None
provider_id: Optional[str] = None
db: Optional[Any] = None
user: Optional[Any] = None
request_id: Optional[str] = None
provider_name: str | None = None
model_name: str | None = None
api_key_id: str | None = None
provider_id: str | None = None
db: Any | None = None
user: Any | None = None
request_id: str | None = None
timeout: float = 30.0
@@ -525,12 +524,12 @@ class EndpointCheckRequest:
class EndpointCheckResult:
"""端点检查结果数据类"""
status_code: int
headers: Dict[str, str]
headers: dict[str, str]
response_time_ms: int
request_id: str
response_data: Optional[Dict[str, Any]] = None
error_message: Optional[str] = None
usage_data: Optional[Dict[str, Any]] = None
response_data: dict[str, Any] | None = None
error_message: str | None = None
usage_data: dict[str, Any] | None = None
class HttpRequestExecutor:
@@ -613,7 +612,7 @@ class UsageCalculator:
return _extract_tokens_from_response(api_identifier, result.response_data)
@staticmethod
def _fallback_token_counting(request_data: Dict[str, Any], response_data: Optional[Dict[str, Any]]) -> tuple[int, int, int, int]:
def _fallback_token_counting(request_data: dict[str, Any], response_data: dict[str, Any] | None) -> tuple[int, int, int, int]:
"""回退的Token计数方法简单估算"""
# 估算输入Token
messages = request_data.get("messages", request_data.get("contents", []))
@@ -665,12 +664,12 @@ class AsyncBatchUsageRecorder:
def __init__(self, batch_size: int = 10, flush_interval: float = 2.0):
self.batch_size = batch_size
self.flush_interval = flush_interval
self.pending_records: List[Dict[str, Any]] = []
self._flush_task: Optional[asyncio.Task] = None
self.pending_records: list[dict[str, Any]] = []
self._flush_task: asyncio.Task | None = None
self._lock = asyncio.Lock()
self._running = True
async def add_record(self, usage_data: Dict[str, Any]) -> None:
async def add_record(self, usage_data: dict[str, Any]) -> None:
"""添加用量记录到批处理队列"""
async with self._lock:
self.pending_records.append(usage_data)
@@ -740,7 +739,7 @@ class AsyncBatchUsageRecorder:
# 全局批处理器实例(单例)
_global_batch_recorder: Optional[AsyncBatchUsageRecorder] = None
_global_batch_recorder: AsyncBatchUsageRecorder | None = None
def get_batch_recorder() -> AsyncBatchUsageRecorder:
"""获取全局批处理器实例"""
@@ -756,7 +755,7 @@ def get_batch_recorder() -> AsyncBatchUsageRecorder:
class EndpointCheckError(Exception):
"""端点检查错误基类"""
def __init__(self, message: str, error_type: str, status_code: int = 500, details: Optional[Dict[str, Any]] = None):
def __init__(self, message: str, error_type: str, status_code: int = 500, details: dict[str, Any] | None = None):
super().__init__(message)
self.message = message
self.error_type = error_type
@@ -765,22 +764,22 @@ class EndpointCheckError(Exception):
class NetworkError(EndpointCheckError):
"""网络请求错误"""
def __init__(self, message: str, details: Optional[Dict[str, Any]] = None):
def __init__(self, message: str, details: dict[str, Any] | None = None):
super().__init__(message, "network_error", 0, details)
class AuthenticationError(EndpointCheckError):
"""认证错误"""
def __init__(self, message: str, details: Optional[Dict[str, Any]] = None):
def __init__(self, message: str, details: dict[str, Any] | None = None):
super().__init__(message, "authentication_error", 401, details)
class RateLimitError(EndpointCheckError):
"""速率限制错误"""
def __init__(self, message: str, details: Optional[Dict[str, Any]] = None):
def __init__(self, message: str, details: dict[str, Any] | None = None):
super().__init__(message, "rate_limit_error", 429, details)
class UpstreamError(EndpointCheckError):
"""上游服务错误"""
def __init__(self, message: str, status_code: int, details: Optional[Dict[str, Any]] = None):
def __init__(self, message: str, status_code: int, details: dict[str, Any] | None = None):
super().__init__(message, "upstream_error", status_code, details)
@@ -982,7 +981,7 @@ class EndpointCheckConfig:
retry_on_timeouts: bool = True
@classmethod
def from_env(cls) -> 'EndpointCheckConfig':
def from_env(cls) -> EndpointCheckConfig:
"""从环境变量创建配置"""
import os
@@ -1004,7 +1003,7 @@ class EndpointCheckConfig:
)
@classmethod
def from_dict(cls, config_dict: Dict[str, Any]) -> 'EndpointCheckConfig':
def from_dict(cls, config_dict: dict[str, Any]) -> EndpointCheckConfig:
"""从字典创建配置"""
return cls(**{k: v for k, v in config_dict.items() if hasattr(cls, k)})
@@ -1012,7 +1011,7 @@ class EndpointCheckConfig:
class ConfigurableEndpointChecker:
"""可配置的端点检查器"""
def __init__(self, config: Optional[EndpointCheckConfig] = None):
def __init__(self, config: EndpointCheckConfig | None = None):
self.config = config or EndpointCheckConfig()
self.executor = HttpRequestExecutor(timeout=self.config.timeout)
self.usage_calculator = UsageCalculator()
@@ -1171,9 +1170,9 @@ class ConfigurableEndpointChecker:
# 全局配置检查器实例
_global_configured_checker: Optional[ConfigurableEndpointChecker] = None
_global_configured_checker: ConfigurableEndpointChecker | None = None
def get_configured_checker(config: Optional[EndpointCheckConfig] = None) -> ConfigurableEndpointChecker:
def get_configured_checker(config: EndpointCheckConfig | None = None) -> ConfigurableEndpointChecker:
"""获取全局配置检查器实例"""
global _global_configured_checker
if _global_configured_checker is None or config is not None:
@@ -1186,8 +1185,8 @@ def get_configured_checker(config: Optional[EndpointCheckConfig] = None) -> Conf
class EndpointCheckOrchestrator:
"""端点检查协调器 - 协调整个流程"""
def __init__(self, executor: Optional[HttpRequestExecutor] = None,
usage_calculator: Optional[UsageCalculator] = None):
def __init__(self, executor: HttpRequestExecutor | None = None,
usage_calculator: UsageCalculator | None = None):
self.executor = executor or HttpRequestExecutor()
self.usage_calculator = usage_calculator or UsageCalculator()

View File

@@ -6,7 +6,7 @@
"""
import re
from typing import Any, Dict, Optional, Tuple, Type
from typing import Any
from src.api.handlers.base.response_parser import (
ParsedChunk,
@@ -20,7 +20,7 @@ from src.api.handlers.base.utils import extract_cache_creation_tokens
from src.core.api_format import is_cli_format
def _check_nested_error(response: Dict[str, Any]) -> Tuple[bool, Optional[Dict[str, Any]]]:
def _check_nested_error(response: dict[str, Any]) -> tuple[bool, dict[str, Any] | None]:
"""
检查响应中是否存在嵌套错误(某些代理服务返回 HTTP 200 但在响应体中包含错误)
@@ -62,7 +62,7 @@ def _check_nested_error(response: Dict[str, Any]) -> Tuple[bool, Optional[Dict[s
return False, None
def _extract_embedded_status_code(error_info: Optional[Dict[str, Any]]) -> Optional[int]:
def _extract_embedded_status_code(error_info: dict[str, Any] | None) -> int | None:
"""
从错误信息中提取嵌套的状态码
@@ -137,7 +137,7 @@ class OpenAIResponseParser(ResponseParser):
self.name = "OPENAI"
self.api_format = "OPENAI"
def parse_sse_line(self, line: str, stats: StreamStats) -> Optional[ParsedChunk]:
def parse_sse_line(self, line: str, stats: StreamStats) -> ParsedChunk | None:
if not line or not line.strip():
return None
@@ -186,7 +186,7 @@ class OpenAIResponseParser(ResponseParser):
return chunk
def parse_response(self, response: Dict[str, Any], status_code: int) -> ParsedResponse:
def parse_response(self, response: dict[str, Any], status_code: int) -> ParsedResponse:
result = ParsedResponse(
raw_response=response,
status_code=status_code,
@@ -217,7 +217,7 @@ class OpenAIResponseParser(ResponseParser):
return result
def extract_usage_from_response(self, response: Dict[str, Any]) -> Dict[str, int]:
def extract_usage_from_response(self, response: dict[str, Any]) -> dict[str, int]:
usage = response.get("usage") or {}
return {
"input_tokens": usage.get("prompt_tokens", 0),
@@ -226,7 +226,7 @@ class OpenAIResponseParser(ResponseParser):
"cache_read_tokens": 0,
}
def extract_text_content(self, response: Dict[str, Any]) -> str:
def extract_text_content(self, response: dict[str, Any]) -> str:
choices = response.get("choices", [])
if choices:
message = choices[0].get("message", {})
@@ -235,7 +235,7 @@ class OpenAIResponseParser(ResponseParser):
return content
return ""
def is_error_response(self, response: Dict[str, Any]) -> bool:
def is_error_response(self, response: dict[str, Any]) -> bool:
is_error, _ = _check_nested_error(response)
return is_error
@@ -259,7 +259,7 @@ class ClaudeResponseParser(ResponseParser):
self.name = "CLAUDE"
self.api_format = "CLAUDE"
def parse_sse_line(self, line: str, stats: StreamStats) -> Optional[ParsedChunk]:
def parse_sse_line(self, line: str, stats: StreamStats) -> ParsedChunk | None:
if not line or not line.strip():
return None
@@ -324,7 +324,7 @@ class ClaudeResponseParser(ResponseParser):
return chunk
def parse_response(self, response: Dict[str, Any], status_code: int) -> ParsedResponse:
def parse_response(self, response: dict[str, Any], status_code: int) -> ParsedResponse:
result = ParsedResponse(
raw_response=response,
status_code=status_code,
@@ -358,7 +358,7 @@ class ClaudeResponseParser(ResponseParser):
return result
def extract_usage_from_response(self, response: Dict[str, Any]) -> Dict[str, int]:
def extract_usage_from_response(self, response: dict[str, Any]) -> dict[str, int]:
# 对于 message_start 事件usage 在 message.usage 路径下
# 对于其他响应usage 在顶层
usage = response.get("usage") or {}
@@ -372,7 +372,7 @@ class ClaudeResponseParser(ResponseParser):
"cache_read_tokens": usage.get("cache_read_input_tokens", 0),
}
def extract_text_content(self, response: Dict[str, Any]) -> str:
def extract_text_content(self, response: dict[str, Any]) -> str:
content = response.get("content", [])
if isinstance(content, list):
text_parts = []
@@ -382,7 +382,7 @@ class ClaudeResponseParser(ResponseParser):
return "".join(text_parts)
return ""
def is_error_response(self, response: Dict[str, Any]) -> bool:
def is_error_response(self, response: dict[str, Any]) -> bool:
is_error, _ = _check_nested_error(response)
return is_error
@@ -406,7 +406,7 @@ class GeminiResponseParser(ResponseParser):
self.name = "GEMINI"
self.api_format = "GEMINI"
def parse_sse_line(self, line: str, stats: StreamStats) -> Optional[ParsedChunk]:
def parse_sse_line(self, line: str, stats: StreamStats) -> ParsedChunk | None:
"""
解析 Gemini SSE 行
@@ -473,7 +473,7 @@ class GeminiResponseParser(ResponseParser):
return chunk
def parse_response(self, response: Dict[str, Any], status_code: int) -> ParsedResponse:
def parse_response(self, response: dict[str, Any], status_code: int) -> ParsedResponse:
result = ParsedResponse(
raw_response=response,
status_code=status_code,
@@ -509,7 +509,7 @@ class GeminiResponseParser(ResponseParser):
return result
def extract_usage_from_response(self, response: Dict[str, Any]) -> Dict[str, int]:
def extract_usage_from_response(self, response: dict[str, Any]) -> dict[str, int]:
"""
从 Gemini 响应中提取 token 使用量
@@ -531,7 +531,7 @@ class GeminiResponseParser(ResponseParser):
"cache_read_tokens": usage.get("cached_tokens", 0),
}
def extract_text_content(self, response: Dict[str, Any]) -> str:
def extract_text_content(self, response: dict[str, Any]) -> str:
candidates = response.get("candidates", [])
if candidates:
content = candidates[0].get("content", {})
@@ -543,7 +543,7 @@ class GeminiResponseParser(ResponseParser):
return "".join(text_parts)
return ""
def is_error_response(self, response: Dict[str, Any]) -> bool:
def is_error_response(self, response: dict[str, Any]) -> bool:
"""
判断响应是否为错误响应
@@ -562,7 +562,7 @@ class GeminiCliResponseParser(GeminiResponseParser):
# 解析器注册表
_PARSERS: Dict[str, Type[ResponseParser]] = {
_PARSERS: dict[str, type[ResponseParser]] = {
"CLAUDE": ClaudeResponseParser,
"CLAUDE_CLI": ClaudeCliResponseParser,
"OPENAI": OpenAIResponseParser,

View File

@@ -11,12 +11,11 @@
payload, headers = builder.build(original_body, original_headers, endpoint, key)
"""
from __future__ import annotations
import json
from abc import ABC, abstractmethod
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, Dict, FrozenSet, Optional, Tuple
from typing import TYPE_CHECKING, Any
from src.core.crypto import crypto_service
from src.core.api_format import HeaderBuilder, UPSTREAM_DROP_HEADERS
@@ -37,9 +36,9 @@ class ProviderAuthInfo:
auth_header: str
auth_value: str
# 解密后的认证配置(用于 URL 构建等场景,避免重复解密)
decrypted_auth_config: Optional[Dict[str, Any]] = None
decrypted_auth_config: dict[str, Any] | None = None
def as_tuple(self) -> Tuple[str, str]:
def as_tuple(self) -> tuple[str, str]:
"""返回 (auth_header, auth_value) 元组"""
return (self.auth_header, self.auth_value)
@@ -48,7 +47,7 @@ class ProviderAuthInfo:
# ==============================================================================
# 兼容别名:历史代码使用 SENSITIVE_HEADERS 命名
SENSITIVE_HEADERS: FrozenSet[str] = UPSTREAM_DROP_HEADERS
SENSITIVE_HEADERS: frozenset[str] = UPSTREAM_DROP_HEADERS
# ==============================================================================
@@ -57,14 +56,14 @@ SENSITIVE_HEADERS: FrozenSet[str] = UPSTREAM_DROP_HEADERS
# 标准测试请求体OpenAI 格式)
# 用于 check_endpoint 等测试场景,使用简单安全的消息内容避免触发安全过滤
DEFAULT_TEST_REQUEST: Dict[str, Any] = {
DEFAULT_TEST_REQUEST: dict[str, Any] = {
"messages": [{"role": "user", "content": "Hi"}],
"max_tokens": 5,
"temperature": 0,
}
def get_test_request_data(request_data: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
def get_test_request_data(request_data: dict[str, Any] | None = None) -> dict[str, Any]:
"""获取测试请求数据
如果传入 request_data则合并到默认测试请求中
@@ -85,8 +84,8 @@ def get_test_request_data(request_data: Optional[Dict[str, Any]] = None) -> Dict
def build_test_request_body(
format_id: str,
request_data: Optional[Dict[str, Any]] = None,
) -> Dict[str, Any]:
request_data: dict[str, Any] | None = None,
) -> dict[str, Any]:
"""构建测试请求体,自动处理格式转换
使用格式转换注册表将 OpenAI 格式的测试请求转换为目标格式。
@@ -127,39 +126,39 @@ class RequestBuilder(ABC):
@abstractmethod
def build_payload(
self,
original_body: Dict[str, Any],
original_body: dict[str, Any],
*,
mapped_model: Optional[str] = None,
mapped_model: str | None = None,
is_stream: bool = False,
) -> Dict[str, Any]:
) -> dict[str, Any]:
"""构建请求体"""
pass
@abstractmethod
def build_headers(
self,
original_headers: Dict[str, str],
original_headers: dict[str, str],
endpoint: Any,
key: Any,
*,
extra_headers: Optional[Dict[str, str]] = None,
pre_computed_auth: Optional[Tuple[str, str]] = None,
) -> Dict[str, str]:
extra_headers: dict[str, str] | None = None,
pre_computed_auth: tuple[str, str] | None = None,
) -> dict[str, str]:
"""构建请求头"""
pass
def build(
self,
original_body: Dict[str, Any],
original_headers: Dict[str, str],
original_body: dict[str, Any],
original_headers: dict[str, str],
endpoint: Any,
key: Any,
*,
mapped_model: Optional[str] = None,
mapped_model: str | None = None,
is_stream: bool = False,
extra_headers: Optional[Dict[str, str]] = None,
pre_computed_auth: Optional[Tuple[str, str]] = None,
) -> Tuple[Dict[str, Any], Dict[str, str]]:
extra_headers: dict[str, str] | None = None,
pre_computed_auth: tuple[str, str] | None = None,
) -> tuple[dict[str, Any], dict[str, str]]:
"""
构建完整的请求(请求体 + 请求头)
@@ -202,11 +201,11 @@ class PassthroughRequestBuilder(RequestBuilder):
def build_payload(
self,
original_body: Dict[str, Any],
original_body: dict[str, Any],
*,
mapped_model: Optional[str] = None, # noqa: ARG002 - 由 apply_mapped_model 处理
mapped_model: str | None = None, # noqa: ARG002 - 由 apply_mapped_model 处理
is_stream: bool = False, # noqa: ARG002 - 保留原始值,不自动添加
) -> Dict[str, Any]:
) -> dict[str, Any]:
"""
透传请求体 - 原样复制,不做任何修改
@@ -218,13 +217,13 @@ class PassthroughRequestBuilder(RequestBuilder):
def build_headers(
self,
original_headers: Dict[str, str],
original_headers: dict[str, str],
endpoint: Any,
key: Any,
*,
extra_headers: Optional[Dict[str, str]] = None,
pre_computed_auth: Optional[Tuple[str, str]] = None,
) -> Dict[str, str]:
extra_headers: dict[str, str] | None = None,
pre_computed_auth: tuple[str, str] | None = None,
) -> dict[str, str]:
"""
透传请求头 - 清理敏感头部(黑名单),透传其他所有头部
@@ -289,11 +288,11 @@ class PassthroughRequestBuilder(RequestBuilder):
def build_passthrough_request(
original_body: Dict[str, Any],
original_headers: Dict[str, str],
original_body: dict[str, Any],
original_headers: dict[str, str],
endpoint: Any,
key: Any,
) -> Tuple[Dict[str, Any], Dict[str, str]]:
) -> tuple[dict[str, Any], dict[str, str]]:
"""
构建透传模式的请求

View File

@@ -4,7 +4,7 @@
from abc import ABC, abstractmethod
from dataclasses import dataclass, field
from typing import Any, Dict, Optional
from typing import Any
@dataclass
@@ -13,14 +13,14 @@ class ParsedChunk:
# 原始数据
raw_line: str
event_type: Optional[str] = None
data: Optional[Dict[str, Any]] = None
event_type: str | None = None
data: dict[str, Any] | None = None
# 提取的内容
text_delta: str = ""
is_done: bool = False
is_error: bool = False
error_message: Optional[str] = None
error_message: str | None = None
# 使用量信息(通常在最后一个 chunk 中)
input_tokens: int = 0
@@ -29,7 +29,7 @@ class ParsedChunk:
cache_read_tokens: int = 0
# 响应 ID
response_id: Optional[str] = None
response_id: str | None = None
@dataclass
@@ -48,21 +48,21 @@ class StreamStats:
# 内容
collected_text: str = ""
response_id: Optional[str] = None
response_id: str | None = None
# 状态
has_completion: bool = False
status_code: int = 200
error_message: Optional[str] = None
error_message: str | None = None
# Provider 信息
provider_name: Optional[str] = None
endpoint_id: Optional[str] = None
key_id: Optional[str] = None
provider_name: str | None = None
endpoint_id: str | None = None
key_id: str | None = None
# 响应头和完整响应
response_headers: Dict[str, str] = field(default_factory=dict)
final_response: Optional[Dict[str, Any]] = None
response_headers: dict[str, str] = field(default_factory=dict)
final_response: dict[str, Any] | None = None
@dataclass
@@ -70,12 +70,12 @@ class ParsedResponse:
"""解析后的非流式响应"""
# 原始响应
raw_response: Dict[str, Any]
raw_response: dict[str, Any]
status_code: int
# 提取的内容
text_content: str = ""
response_id: Optional[str] = None
response_id: str | None = None
# 使用量
input_tokens: int = 0
@@ -85,10 +85,10 @@ class ParsedResponse:
# 错误信息
is_error: bool = False
error_type: Optional[str] = None
error_message: Optional[str] = None
error_type: str | None = None
error_message: str | None = None
# 从响应体解析出的嵌套状态码(当 HTTP 200 但响应体含错误时使用)
embedded_status_code: Optional[int] = None
embedded_status_code: int | None = None
class ResponseParser(ABC):
@@ -106,7 +106,7 @@ class ResponseParser(ABC):
api_format: str = "UNKNOWN"
@abstractmethod
def parse_sse_line(self, line: str, stats: StreamStats) -> Optional[ParsedChunk]:
def parse_sse_line(self, line: str, stats: StreamStats) -> ParsedChunk | None:
"""
解析单行 SSE 数据
@@ -120,7 +120,7 @@ class ResponseParser(ABC):
pass
@abstractmethod
def parse_response(self, response: Dict[str, Any], status_code: int) -> ParsedResponse:
def parse_response(self, response: dict[str, Any], status_code: int) -> ParsedResponse:
"""
解析非流式响应
@@ -134,7 +134,7 @@ class ResponseParser(ABC):
pass
@abstractmethod
def extract_usage_from_response(self, response: Dict[str, Any]) -> Dict[str, int]:
def extract_usage_from_response(self, response: dict[str, Any]) -> dict[str, int]:
"""
从响应中提取 token 使用量
@@ -147,7 +147,7 @@ class ResponseParser(ABC):
pass
@abstractmethod
def extract_text_content(self, response: Dict[str, Any]) -> str:
def extract_text_content(self, response: dict[str, Any]) -> str:
"""
从响应中提取文本内容
@@ -159,7 +159,7 @@ class ResponseParser(ABC):
"""
pass
def is_error_response(self, response: Dict[str, Any]) -> bool:
def is_error_response(self, response: dict[str, Any]) -> bool:
"""
判断响应是否为错误响应

View File

@@ -8,9 +8,11 @@
- 请求/响应数据
"""
from __future__ import annotations
import time
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any, Dict, List, Optional
from typing import TYPE_CHECKING, Any
if TYPE_CHECKING:
from src.core.api_format.conversion.stream_state import StreamState
@@ -35,16 +37,16 @@ class StreamContext:
api_key_id: int = 0
# Provider 信息(在请求执行时填充)
provider_name: Optional[str] = None
provider_id: Optional[str] = None
endpoint_id: Optional[str] = None
key_id: Optional[str] = None
attempt_id: Optional[str] = None
provider_name: str | None = None
provider_id: str | None = None
endpoint_id: str | None = None
key_id: str | None = None
attempt_id: str | None = None
attempt_synced: bool = False
provider_api_format: Optional[str] = None # Provider 的响应格式
provider_api_format: str | None = None # Provider 的响应格式
# 模型映射
mapped_model: Optional[str] = None
mapped_model: str | None = None
# Token 统计
input_tokens: int = 0
@@ -53,33 +55,33 @@ class StreamContext:
cache_creation_tokens: int = 0
# 响应内容
_collected_text_parts: List[str] = field(default_factory=list, repr=False)
response_id: Optional[str] = None
final_usage: Optional[Dict[str, Any]] = None
final_response: Optional[Dict[str, Any]] = None
_collected_text_parts: list[str] = field(default_factory=list, repr=False)
response_id: str | None = None
final_usage: dict[str, Any] | None = None
final_response: dict[str, Any] | None = None
# 时间指标
first_byte_time_ms: Optional[int] = None # 首字时间 (TTFB - Time To First Byte)
first_byte_time_ms: int | None = None # 首字时间 (TTFB - Time To First Byte)
start_time: float = field(default_factory=time.time)
# 响应状态
status_code: int = 200
error_message: Optional[str] = None # 客户端友好的错误消息
upstream_response: Optional[str] = None # 原始 Provider 响应(用于请求链路追踪)
error_message: str | None = None # 客户端友好的错误消息
upstream_response: str | None = None # 原始 Provider 响应(用于请求链路追踪)
has_completion: bool = False
# 请求/响应数据
response_headers: Dict[str, str] = field(default_factory=dict) # 提供商响应头
client_response_headers: Dict[str, str] = field(default_factory=dict) # 返回给客户端的响应头
provider_request_headers: Dict[str, str] = field(default_factory=dict)
provider_request_body: Optional[Dict[str, Any]] = None
response_headers: dict[str, str] = field(default_factory=dict) # 提供商响应头
client_response_headers: dict[str, str] = field(default_factory=dict) # 返回给客户端的响应头
provider_request_headers: dict[str, str] = field(default_factory=dict)
provider_request_body: dict[str, Any] | None = None
# 格式转换信息CLI handler 需要)
client_api_format: str = ""
needs_conversion: bool = False # 是否需要跨格式转换(由 handler 层设置)
# Provider 响应元数据CLI handler 需要)
response_metadata: Dict[str, Any] = field(default_factory=dict)
response_metadata: dict[str, Any] = field(default_factory=dict)
# 整流标记Thinking Rectifier
rectified: bool = False # 请求是否经过整流(移除 thinking 块后重试)
@@ -87,10 +89,10 @@ class StreamContext:
# 流式处理统计
data_count: int = 0
chunk_count: int = 0
parsed_chunks: List[Dict[str, Any]] = field(default_factory=list)
parsed_chunks: list[dict[str, Any]] = field(default_factory=list)
# 流式格式转换状态(跨 chunk 追踪)
stream_conversion_state: Optional["StreamState"] = None
stream_conversion_state: StreamState | None = None
def reset_for_retry(self) -> None:
"""
@@ -138,7 +140,7 @@ class StreamContext:
provider_id: str,
endpoint_id: str,
key_id: str,
provider_api_format: Optional[str] = None,
provider_api_format: str | None = None,
) -> None:
"""更新 Provider 信息"""
self.provider_name = provider_name
@@ -149,10 +151,10 @@ class StreamContext:
def update_usage(
self,
input_tokens: Optional[int] = None,
output_tokens: Optional[int] = None,
cached_tokens: Optional[int] = None,
cache_creation_tokens: Optional[int] = None,
input_tokens: int | None = None,
output_tokens: int | None = None,
cached_tokens: int | None = None,
cache_creation_tokens: int | None = None,
) -> None:
"""
更新 Token 使用统计
@@ -194,7 +196,7 @@ class StreamContext:
self,
status_code: int,
error_message: str,
upstream_response: Optional[str] = None,
upstream_response: str | None = None,
) -> None:
"""
标记请求失败
@@ -230,7 +232,7 @@ class StreamContext:
"""检查是否因客户端断开连接而结束"""
return self.status_code == 499
def build_response_body(self, response_time_ms: int) -> Dict[str, Any]:
def build_response_body(self, response_time_ms: int) -> dict[str, Any]:
"""
构建响应体元数据

View File

@@ -13,7 +13,10 @@ import asyncio
import codecs
import json
from dataclasses import dataclass
from typing import Any, AsyncGenerator, Callable, Optional
from typing import Any
from collections.abc import Callable
from collections.abc import AsyncGenerator
import httpx
@@ -65,10 +68,10 @@ class StreamProcessor:
self,
request_id: str,
default_parser: ResponseParser,
on_streaming_start: Optional[Callable[[], None]] = None,
on_streaming_start: Callable[[], None] | None = None,
*,
collect_text: bool = False,
smoothing_config: Optional[StreamSmoothingConfig] = None,
smoothing_config: StreamSmoothingConfig | None = None,
):
"""
初始化流处理器
@@ -105,7 +108,7 @@ class StreamProcessor:
def handle_sse_event(
self,
ctx: StreamContext,
event_name: Optional[str],
event_name: str | None,
data_str: str,
*,
skip_record: bool = False,
@@ -363,7 +366,7 @@ class StreamProcessor:
):
# 重新抛出可重试的 Provider 异常,触发故障转移
raise
except (OSError, IOError) as e:
except OSError as e:
# 网络 I/O 异常:记录警告,可能需要重试
logger.warning(f" [{self.request_id}] 预读流时发生网络异常: {type(e).__name__}: {e}")
except Exception as e:
@@ -382,10 +385,10 @@ class StreamProcessor:
byte_iterator: Any,
response_ctx: Any,
http_client: httpx.AsyncClient,
prefetched_chunks: Optional[list] = None,
prefetched_chunks: list | None = None,
*,
start_time: Optional[float] = None,
) -> AsyncGenerator[bytes, None]:
start_time: float | None = None,
) -> AsyncGenerator[bytes]:
"""
创建响应流生成器
@@ -547,9 +550,7 @@ class StreamProcessor:
logger.warning(f"[{self.request_id}] 流式格式转换失败: {conv_err}")
payload = _build_stream_error_payload("响应格式转换失败,请稍后重试")
error_bytes = (
f"data: {json.dumps(payload, ensure_ascii=False)}\n\n".encode(
"utf-8"
)
f"data: {json.dumps(payload, ensure_ascii=False)}\n\n".encode()
)
done_bytes = (
b"data: [DONE]\n\n" if client_format.startswith("OPENAI") else b""
@@ -570,7 +571,7 @@ class StreamProcessor:
# 统一使用 SSE 格式输出Gemini streamGenerateContent 也使用 SSE
# 参考: https://ai.google.dev/api/generate-content
out.append(
f"data: {json.dumps(evt, ensure_ascii=False)}\n\n".encode("utf-8")
f"data: {json.dumps(evt, ensure_ascii=False)}\n\n".encode()
)
return out
@@ -769,9 +770,9 @@ class StreamProcessor:
async def create_monitored_stream(
self,
ctx: StreamContext,
stream_generator: AsyncGenerator[bytes, None],
stream_generator: AsyncGenerator[bytes],
is_disconnected: Callable[[], Any],
) -> AsyncGenerator[bytes, None]:
) -> AsyncGenerator[bytes]:
"""
创建带监控的流生成器
@@ -833,8 +834,8 @@ class StreamProcessor:
async def create_smoothed_stream(
self,
stream_generator: AsyncGenerator[bytes, None],
) -> AsyncGenerator[bytes, None]:
stream_generator: AsyncGenerator[bytes],
) -> AsyncGenerator[bytes]:
"""
创建平滑输出的流生成器
@@ -933,7 +934,7 @@ class StreamProcessor:
if buffer:
yield buffer
def _get_extractor(self, format_name: str) -> Optional[ContentExtractor]:
def _get_extractor(self, format_name: str) -> ContentExtractor | None:
"""获取或创建格式对应的提取器(带缓存)"""
if format_name not in self._extractors:
extractor = get_extractor(format_name)
@@ -943,7 +944,7 @@ class StreamProcessor:
def _detect_format_and_extract(
self, data: dict
) -> tuple[Optional[str], Optional[ContentExtractor]]:
) -> tuple[str | None, ContentExtractor | None]:
"""
检测数据格式并提取内容
@@ -998,10 +999,10 @@ class StreamProcessor:
async def create_smoothed_stream(
stream_generator: AsyncGenerator[bytes, None],
stream_generator: AsyncGenerator[bytes],
chunk_size: int = 20,
delay_ms: int = 8,
) -> AsyncGenerator[bytes, None]:
) -> AsyncGenerator[bytes]:
"""
独立的平滑流生成函数
@@ -1032,7 +1033,7 @@ class _LightweightSmoother:
self.delay_ms = delay_ms
self._extractors: dict[str, ContentExtractor] = {}
def _get_extractor(self, format_name: str) -> Optional[ContentExtractor]:
def _get_extractor(self, format_name: str) -> ContentExtractor | None:
if format_name not in self._extractors:
extractor = get_extractor(format_name)
if extractor:
@@ -1041,7 +1042,7 @@ class _LightweightSmoother:
def _detect_format_and_extract(
self, data: dict
) -> tuple[Optional[str], Optional[ContentExtractor]]:
) -> tuple[str | None, ContentExtractor | None]:
for format_name in get_extractor_formats():
extractor = self._get_extractor(format_name)
if extractor:
@@ -1060,8 +1061,8 @@ class _LightweightSmoother:
return [content[i : i + self.chunk_size] for i in range(0, text_length, self.chunk_size)]
async def smooth(
self, stream_generator: AsyncGenerator[bytes, None]
) -> AsyncGenerator[bytes, None]:
self, stream_generator: AsyncGenerator[bytes]
) -> AsyncGenerator[bytes]:
buffer = b""
is_first_content = True

View File

@@ -9,7 +9,7 @@
import asyncio
import time
from typing import Any, Dict, Optional
from typing import Any
from sqlalchemy.orm import Session
@@ -58,8 +58,8 @@ class StreamTelemetryRecorder:
async def record_stream_stats(
self,
ctx: StreamContext,
original_headers: Dict[str, str],
original_request_body: Dict[str, Any],
original_headers: dict[str, str],
original_request_body: dict[str, Any],
start_time: float,
) -> None:
"""
@@ -144,9 +144,9 @@ class StreamTelemetryRecorder:
self,
writer: TelemetryWriter,
ctx: StreamContext,
original_headers: Dict[str, str],
actual_request_body: Dict[str, Any],
response_body: Optional[Dict[str, Any]],
original_headers: dict[str, str],
actual_request_body: dict[str, Any],
response_body: dict[str, Any] | None,
response_time_ms: int,
) -> None:
"""记录成功的请求"""
@@ -193,9 +193,9 @@ class StreamTelemetryRecorder:
self,
writer: TelemetryWriter,
ctx: StreamContext,
original_headers: Dict[str, str],
actual_request_body: Dict[str, Any],
response_body: Optional[Dict[str, Any]],
original_headers: dict[str, str],
actual_request_body: dict[str, Any],
response_body: dict[str, Any] | None,
response_time_ms: int,
) -> None:
"""记录失败的请求"""
@@ -236,9 +236,9 @@ class StreamTelemetryRecorder:
self,
writer: TelemetryWriter,
ctx: StreamContext,
original_headers: Dict[str, str],
actual_request_body: Dict[str, Any],
response_body: Optional[Dict[str, Any]],
original_headers: dict[str, str],
actual_request_body: dict[str, Any],
response_body: dict[str, Any] | None,
response_time_ms: int,
) -> None:
"""记录客户端取消的请求"""
@@ -285,7 +285,7 @@ class StreamTelemetryRecorder:
from src.services.request.candidate import RequestCandidateService
extra_data: Dict[str, Any] = {
extra_data: dict[str, Any] = {
"stream_completed": ctx.is_success(),
"data_count": ctx.data_count,
}
@@ -358,7 +358,7 @@ class StreamTelemetryRecorder:
status: str,
response_time_ms: int,
status_code: int = 200,
error_message: Optional[str] = None,
error_message: str | None = None,
) -> None:
"""直接更新 Usage 表的状态字段"""
try:
@@ -378,7 +378,7 @@ class StreamTelemetryRecorder:
async def _get_telemetry_writer(
self, bg_db: Session, ctx: StreamContext, response_time_ms: int
) -> Optional[TelemetryWriter]:
) -> TelemetryWriter | None:
if config.usage_queue_enabled and self.user_id and self.api_key_id:
return QueueTelemetryWriter(
request_id=self.request_id,
@@ -400,9 +400,9 @@ class StreamTelemetryRecorder:
self,
writer: TelemetryWriter,
ctx: StreamContext,
original_headers: Dict[str, str],
actual_request_body: Dict[str, Any],
response_body: Optional[Dict[str, Any]],
original_headers: dict[str, str],
actual_request_body: dict[str, Any],
response_body: dict[str, Any] | None,
response_time_ms: int,
) -> None:
"""根据上下文状态分发到对应的记录方法"""
@@ -430,7 +430,7 @@ class StreamTelemetryRecorder:
return "cancelled"
return "failed"
def _build_db_writer(self, bg_db: Session) -> Optional[DbTelemetryWriter]:
def _build_db_writer(self, bg_db: Session) -> DbTelemetryWriter | None:
user = bg_db.query(User).filter(User.id == self.user_id).first()
api_key_obj = bg_db.query(ApiKey).filter(ApiKey.id == self.api_key_id).first()

View File

@@ -4,8 +4,9 @@ Handler 基础工具函数
from __future__ import annotations
import json
from typing import TYPE_CHECKING, Any, Dict, Optional
from typing import TYPE_CHECKING, Any
from src.core.exceptions import EmbeddedErrorException, ProviderNotAvailableException
from src.core.api_format import filter_response_headers
@@ -15,7 +16,7 @@ if TYPE_CHECKING:
from src.core.api_format.conversion.registry import FormatConversionRegistry
def get_format_converter_registry() -> "FormatConversionRegistry":
def get_format_converter_registry() -> FormatConversionRegistry:
"""
获取格式转换注册表(线程安全)
@@ -31,7 +32,7 @@ def get_format_converter_registry() -> "FormatConversionRegistry":
return format_conversion_registry
def extract_cache_creation_tokens(usage: Dict[str, Any]) -> int:
def extract_cache_creation_tokens(usage: dict[str, Any]) -> int:
"""
提取缓存创建 tokens兼容三种格式
@@ -99,7 +100,7 @@ def extract_cache_creation_tokens(usage: Dict[str, Any]) -> int:
return old_format
def build_sse_headers(extra_headers: Optional[Dict[str, str]] = None) -> Dict[str, str]:
def build_sse_headers(extra_headers: dict[str, str] | None = None) -> dict[str, str]:
"""
构建 SSEtext/event-stream推荐响应头用于减少代理缓冲带来的卡顿/成段输出。
@@ -107,7 +108,7 @@ def build_sse_headers(extra_headers: Optional[Dict[str, str]] = None) -> Dict[st
- Cache-Control: no-transform 可避免部分代理对流做压缩/改写导致缓冲
- X-Accel-Buffering: no 可显式提示 Nginx 关闭缓冲(即使全局已关闭也无害)
"""
headers: Dict[str, str] = {
headers: dict[str, str] = {
"Cache-Control": "no-cache, no-transform",
"X-Accel-Buffering": "no",
}
@@ -116,7 +117,7 @@ def build_sse_headers(extra_headers: Optional[Dict[str, str]] = None) -> Dict[st
return headers
def filter_proxy_response_headers(headers: Optional[Dict[str, str]]) -> Dict[str, str]:
def filter_proxy_response_headers(headers: dict[str, str] | None) -> dict[str, str]:
"""
过滤上游响应头中不应透传给客户端的字段。
@@ -148,8 +149,8 @@ def check_prefetched_response_error(
parser: Any,
request_id: str,
provider_name: str,
endpoint_id: Optional[str],
base_url: Optional[str],
endpoint_id: str | None,
base_url: str | None,
) -> None:
"""
检查预读的响应是否为非 SSE 格式的错误响应HTML 或纯 JSON 错误)

View File

@@ -4,7 +4,7 @@ Claude Chat Adapter - 基于 ChatAdapterBase 的 Claude Chat API 适配器
处理 /v1/messages 端点的 Claude Chat 格式请求。
"""
from typing import Any, Dict, Optional, Tuple, Type
from typing import Any
import httpx
from fastapi import HTTPException, Request
@@ -25,9 +25,9 @@ class ClaudeCapabilityDetector:
@staticmethod
def detect_from_headers(
headers: Dict[str, str],
request_body: Optional[Dict[str, Any]] = None,
) -> Dict[str, bool]:
headers: dict[str, str],
request_body: dict[str, Any] | None = None,
) -> dict[str, bool]:
"""
从 Claude 请求头检测能力需求
@@ -38,7 +38,7 @@ class ClaudeCapabilityDetector:
headers: 请求头字典
request_body: 请求体Claude 不使用,保留用于接口统一)
"""
requirements: Dict[str, bool] = {}
requirements: dict[str, bool] = {}
# 使用统一的大小写不敏感获取
beta_header = get_header_value(headers, "anthropic-beta")
@@ -61,21 +61,21 @@ class ClaudeChatAdapter(ChatAdapterBase):
name = "claude.chat"
@property
def HANDLER_CLASS(self) -> Type[ChatHandlerBase]:
def HANDLER_CLASS(self) -> type[ChatHandlerBase]:
"""延迟导入 Handler 类避免循环依赖"""
from src.api.handlers.claude.handler import ClaudeChatHandler
return ClaudeChatHandler
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
def __init__(self, allowed_api_formats: list[str] | None = None):
super().__init__(allowed_api_formats or ["CLAUDE"])
logger.info(f"[{self.name}] 初始化Chat模式适配器 | API格式: {self.allowed_api_formats}")
def detect_capability_requirements(
self,
headers: Dict[str, str],
request_body: Optional[Dict[str, Any]] = None,
) -> Dict[str, bool]:
headers: dict[str, str],
request_body: dict[str, Any] | None = None,
) -> dict[str, bool]:
"""检测 Claude 请求中隐含的能力需求"""
return ClaudeCapabilityDetector.detect_from_headers(headers)
@@ -124,7 +124,7 @@ class ClaudeChatAdapter(ChatAdapterBase):
)
return request
def _build_audit_metadata(self, _payload: Dict[str, Any], request_obj) -> Dict[str, Any]:
def _build_audit_metadata(self, _payload: dict[str, Any], request_obj) -> dict[str, Any]:
"""构建 Claude Chat 特定的审计元数据"""
role_counts: dict[str, int] = {}
for message in request_obj.messages:
@@ -153,8 +153,8 @@ class ClaudeChatAdapter(ChatAdapterBase):
client: httpx.AsyncClient,
base_url: str,
api_key: str,
extra_headers: Optional[Dict[str, str]] = None,
) -> Tuple[list, Optional[str]]:
extra_headers: dict[str, str] | None = None,
) -> tuple[list, str | None]:
"""查询 Claude API 支持的模型列表"""
headers = cls.build_headers_with_extra(api_key, extra_headers)
@@ -201,7 +201,7 @@ class ClaudeChatAdapter(ChatAdapterBase):
# build_request_body 使用基类实现,通过 format_conversion_registry 自动转换 OPENAI -> CLAUDE
def build_claude_adapter(x_app_header: Optional[str]):
def build_claude_adapter(x_app_header: str | None):
"""根据 x-app 头部构造 Chat 或 Claude Code 适配器。"""
if x_app_header and x_app_header.lower() == "cli":
from src.api.handlers.claude_cli.adapter import ClaudeCliAdapter
@@ -216,7 +216,7 @@ class ClaudeTokenCountAdapter(ApiAdapter):
name = "claude.token_count"
mode = ApiMode.STANDARD
def extract_api_key(self, request: Request) -> Optional[str]:
def extract_api_key(self, request: Request) -> str | None:
"""从请求中提取 API 密钥 (x-api-key 或 Authorization: Bearer)"""
# 优先检查 x-api-key
api_key = request.headers.get("x-api-key")

View File

@@ -5,7 +5,7 @@ Claude Chat Handler - 基于通用 Chat Handler 基类的简化实现
代码量从原来的 ~1470 行减少到 ~120 行。
"""
from typing import Any, Dict, Optional
from typing import Any
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
from src.api.handlers.base.utils import extract_cache_creation_tokens
@@ -25,8 +25,8 @@ class ClaudeChatHandler(ChatHandlerBase):
def extract_model_from_request(
self,
request_body: Dict[str, Any],
path_params: Optional[Dict[str, Any]] = None, # noqa: ARG002
request_body: dict[str, Any],
path_params: dict[str, Any] | None = None, # noqa: ARG002
) -> str:
"""
从请求中提取模型名 - Claude 格式实现
@@ -45,9 +45,9 @@ class ClaudeChatHandler(ChatHandlerBase):
def apply_mapped_model(
self,
request_body: Dict[str, Any],
request_body: dict[str, Any],
mapped_model: str,
) -> Dict[str, Any]:
) -> dict[str, Any]:
"""
将映射后的模型名应用到请求体
@@ -90,7 +90,7 @@ class ClaudeChatHandler(ChatHandlerBase):
return request
def _extract_usage(self, response: Dict) -> Dict[str, int]:
def _extract_usage(self, response: dict) -> dict[str, int]:
"""
从 Claude 响应中提取 token 使用情况
@@ -108,7 +108,7 @@ class ClaudeChatHandler(ChatHandlerBase):
"cache_read_input_tokens": usage.get("cache_read_input_tokens", 0),
}
def _normalize_response(self, response: Dict[str, Any]) -> Dict[str, Any]:
def _normalize_response(self, response: dict[str, Any]) -> dict[str, Any]:
"""
规范化 Claude 响应

View File

@@ -4,10 +4,9 @@ Claude SSE 流解析器
解析 Claude Messages API 的 Server-Sent Events 流。
"""
from __future__ import annotations
import json
from typing import Any, Dict, List, Optional
from typing import Any
from src.api.handlers.base.utils import extract_cache_creation_tokens
@@ -43,7 +42,7 @@ class ClaudeStreamParser:
DELTA_TEXT = "text_delta"
DELTA_INPUT_JSON = "input_json_delta"
def parse_chunk(self, chunk: bytes | str) -> List[Dict[str, Any]]:
def parse_chunk(self, chunk: bytes | str) -> list[dict[str, Any]]:
"""
解析 SSE 数据块
@@ -58,10 +57,10 @@ class ClaudeStreamParser:
else:
text = chunk
events: List[Dict[str, Any]] = []
events: list[dict[str, Any]] = []
lines = text.strip().split("\n")
current_event_type: Optional[str] = None
current_event_type: str | None = None
for line in lines:
line = line.strip()
@@ -96,7 +95,7 @@ class ClaudeStreamParser:
return events
def parse_line(self, line: str) -> Optional[Dict[str, Any]]:
def parse_line(self, line: str) -> dict[str, Any] | None:
"""
解析单行 SSE 数据
@@ -117,7 +116,7 @@ class ClaudeStreamParser:
except json.JSONDecodeError:
return None
def is_done_event(self, event: Dict[str, Any]) -> bool:
def is_done_event(self, event: dict[str, Any]) -> bool:
"""
判断是否为结束事件
@@ -130,7 +129,7 @@ class ClaudeStreamParser:
event_type = event.get("type")
return event_type in (self.EVENT_MESSAGE_STOP, "__done__")
def is_error_event(self, event: Dict[str, Any]) -> bool:
def is_error_event(self, event: dict[str, Any]) -> bool:
"""
判断是否为错误事件
@@ -142,7 +141,7 @@ class ClaudeStreamParser:
"""
return event.get("type") == self.EVENT_ERROR
def get_event_type(self, event: Dict[str, Any]) -> Optional[str]:
def get_event_type(self, event: dict[str, Any]) -> str | None:
"""
获取事件类型
@@ -155,7 +154,7 @@ class ClaudeStreamParser:
event_type = event.get("type")
return str(event_type) if event_type is not None else None
def extract_text_delta(self, event: Dict[str, Any]) -> Optional[str]:
def extract_text_delta(self, event: dict[str, Any]) -> str | None:
"""
从 content_block_delta 事件中提取文本增量
@@ -175,7 +174,7 @@ class ClaudeStreamParser:
return None
def extract_usage(self, event: Dict[str, Any]) -> Optional[Dict[str, int]]:
def extract_usage(self, event: dict[str, Any]) -> dict[str, int] | None:
"""
从事件中提取 token 使用量
@@ -212,7 +211,7 @@ class ClaudeStreamParser:
return None
def extract_message_id(self, event: Dict[str, Any]) -> Optional[str]:
def extract_message_id(self, event: dict[str, Any]) -> str | None:
"""
从 message_start 事件中提取消息 ID
@@ -229,7 +228,7 @@ class ClaudeStreamParser:
msg_id = message.get("id")
return str(msg_id) if msg_id is not None else None
def extract_stop_reason(self, event: Dict[str, Any]) -> Optional[str]:
def extract_stop_reason(self, event: dict[str, Any]) -> str | None:
"""
从 message_delta 事件中提取停止原因

View File

@@ -4,7 +4,7 @@ Claude CLI Adapter - 基于通用 CLI Adapter 基类的简化实现
继承 CliAdapterBase只需配置 FORMAT_ID 和 HANDLER_CLASS。
"""
from typing import Any, Dict, Optional, Tuple, Type
from typing import Any
import httpx
@@ -27,20 +27,20 @@ class ClaudeCliAdapter(CliAdapterBase):
name = "claude.cli"
@property
def HANDLER_CLASS(self) -> Type[CliMessageHandlerBase]:
def HANDLER_CLASS(self) -> type[CliMessageHandlerBase]:
"""延迟导入 Handler 类避免循环依赖"""
from src.api.handlers.claude_cli.handler import ClaudeCliMessageHandler
return ClaudeCliMessageHandler
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
def __init__(self, allowed_api_formats: list[str] | None = None):
super().__init__(allowed_api_formats or ["CLAUDE_CLI"])
def detect_capability_requirements(
self,
headers: Dict[str, str],
request_body: Optional[Dict[str, Any]] = None,
) -> Dict[str, bool]:
headers: dict[str, str],
request_body: dict[str, Any] | None = None,
) -> dict[str, bool]:
"""检测 Claude CLI 请求中隐含的能力需求"""
return ClaudeCapabilityDetector.detect_from_headers(headers)
@@ -61,16 +61,16 @@ class ClaudeCliAdapter(CliAdapterBase):
"""
return input_tokens + cache_creation_input_tokens + cache_read_input_tokens
def _extract_message_count(self, payload: Dict[str, Any]) -> int:
def _extract_message_count(self, payload: dict[str, Any]) -> int:
"""Claude CLI 使用 messages 字段"""
messages = payload.get("messages", [])
return len(messages) if isinstance(messages, list) else 0
def _build_audit_metadata(
self,
payload: Dict[str, Any],
path_params: Optional[Dict[str, Any]] = None, # noqa: ARG002
) -> Dict[str, Any]:
payload: dict[str, Any],
path_params: dict[str, Any] | None = None, # noqa: ARG002
) -> dict[str, Any]:
"""Claude CLI 特定的审计元数据"""
model = payload.get("model", "unknown")
stream = payload.get("stream", False)
@@ -104,8 +104,8 @@ class ClaudeCliAdapter(CliAdapterBase):
client: httpx.AsyncClient,
base_url: str,
api_key: str,
extra_headers: Optional[Dict[str, str]] = None,
) -> Tuple[list, Optional[str]]:
extra_headers: dict[str, str] | None = None,
) -> tuple[list, str | None]:
"""查询 Claude API 支持的模型列表(带 CLI User-Agent"""
# 复用 ClaudeChatAdapter 的实现,添加 CLI User-Agent
cli_headers = {"User-Agent": config.internal_user_agent_claude_cli}
@@ -120,7 +120,7 @@ class ClaudeCliAdapter(CliAdapterBase):
return models, error
@classmethod
def build_endpoint_url(cls, base_url: str, request_data: Dict[str, Any], model_name: Optional[str] = None) -> str:
def build_endpoint_url(cls, base_url: str, request_data: dict[str, Any], model_name: str | None = None) -> str:
"""构建Claude CLI API端点URL"""
base_url = base_url.rstrip("/")
if base_url.endswith("/v1"):
@@ -131,12 +131,12 @@ class ClaudeCliAdapter(CliAdapterBase):
# build_request_body 使用基类实现,通过 format_conversion_registry 自动转换 OPENAI -> CLAUDE_CLI
@classmethod
def get_cli_user_agent(cls) -> Optional[str]:
def get_cli_user_agent(cls) -> str | None:
"""获取Claude CLI User-Agent"""
return config.internal_user_agent_claude_cli
@classmethod
def get_cli_extra_headers(cls) -> Dict[str, str]:
def get_cli_extra_headers(cls) -> dict[str, str]:
"""获取Claude CLI额外请求头包含 x-app: cli 标识"""
headers = super().get_cli_extra_headers()
headers["x-app"] = "cli" # 标识 CLI 模式,让上游使用正确的认证方式

View File

@@ -4,7 +4,7 @@ Claude CLI Message Handler - 基于通用 CLI Handler 基类的简化实现
继承 CliMessageHandlerBase只需覆盖格式特定的配置和事件处理逻辑。
"""
from typing import Any, Dict, Optional
from typing import Any
from src.api.handlers.base.cli_handler_base import (
CliMessageHandlerBase,
@@ -33,8 +33,8 @@ class ClaudeCliMessageHandler(CliMessageHandlerBase):
def extract_model_from_request(
self,
request_body: Dict[str, Any],
path_params: Optional[Dict[str, Any]] = None, # noqa: ARG002
request_body: dict[str, Any],
path_params: dict[str, Any] | None = None, # noqa: ARG002
) -> str:
"""
从请求中提取模型名 - Claude 格式实现
@@ -53,9 +53,9 @@ class ClaudeCliMessageHandler(CliMessageHandlerBase):
def apply_mapped_model(
self,
request_body: Dict[str, Any],
request_body: dict[str, Any],
mapped_model: str,
) -> Dict[str, Any]:
) -> dict[str, Any]:
"""
Claude API 的 model 在请求体顶级
@@ -74,7 +74,7 @@ class ClaudeCliMessageHandler(CliMessageHandlerBase):
self,
ctx: StreamContext,
event_type: str,
data: Dict[str, Any],
data: dict[str, Any],
) -> None:
"""
处理 Claude CLI 格式的 SSE 事件
@@ -142,8 +142,8 @@ class ClaudeCliMessageHandler(CliMessageHandlerBase):
def _extract_response_metadata(
self,
response: Dict[str, Any],
) -> Dict[str, Any]:
response: dict[str, Any],
) -> dict[str, Any]:
"""
从 Claude 响应中提取元数据
@@ -155,7 +155,7 @@ class ClaudeCliMessageHandler(CliMessageHandlerBase):
Returns:
提取的元数据字典
"""
metadata: Dict[str, Any] = {}
metadata: dict[str, Any] = {}
# 提取模型名称(实际使用的模型)
if "model" in response:

View File

@@ -4,7 +4,7 @@ Gemini Chat Adapter
处理 Gemini API 格式的请求适配
"""
from typing import Any, Dict, Optional, Tuple, Type
from typing import Any
import httpx
from fastapi import HTTPException, Request
@@ -33,17 +33,17 @@ class GeminiChatAdapter(ChatAdapterBase):
name = "gemini.chat"
@property
def HANDLER_CLASS(self) -> Type[ChatHandlerBase]:
def HANDLER_CLASS(self) -> type[ChatHandlerBase]:
"""延迟导入 Handler 类避免循环依赖"""
from src.api.handlers.gemini.handler import GeminiChatHandler
return GeminiChatHandler
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
def __init__(self, allowed_api_formats: list[str] | None = None):
super().__init__(allowed_api_formats or ["GEMINI"])
logger.info(f"[{self.name}] 初始化 Gemini Chat 适配器 | API格式: {self.allowed_api_formats}")
def extract_api_key(self, request: Request) -> Optional[str]:
def extract_api_key(self, request: Request) -> str | None:
"""
从请求中提取 API 密钥 - Gemini 支持 header 和 query 两种方式
@@ -68,8 +68,8 @@ class GeminiChatAdapter(ChatAdapterBase):
return {}
def _merge_path_params(
self, original_request_body: Dict[str, Any], path_params: Dict[str, Any] # noqa: ARG002
) -> Dict[str, Any]:
self, original_request_body: dict[str, Any], path_params: dict[str, Any] # noqa: ARG002
) -> dict[str, Any]:
"""
合并 URL 路径参数到请求体 - Gemini 特化版本
@@ -122,14 +122,14 @@ class GeminiChatAdapter(ChatAdapterBase):
request.stream = is_stream
return request
def _extract_message_count(self, payload: Dict[str, Any], request_obj) -> int:
def _extract_message_count(self, payload: dict[str, Any], request_obj) -> int:
"""提取消息数量"""
contents = payload.get("contents", [])
if hasattr(request_obj, "contents"):
contents = request_obj.contents
return len(contents) if isinstance(contents, list) else 0
def _build_audit_metadata(self, payload: Dict[str, Any], request_obj) -> Dict[str, Any]:
def _build_audit_metadata(self, payload: dict[str, Any], request_obj) -> dict[str, Any]:
"""构建 Gemini Chat 特定的审计元数据"""
role_counts: dict[str, int] = {}
@@ -182,8 +182,8 @@ class GeminiChatAdapter(ChatAdapterBase):
client: httpx.AsyncClient,
base_url: str,
api_key: str,
extra_headers: Optional[Dict[str, str]] = None,
) -> Tuple[list, Optional[str]]:
extra_headers: dict[str, str] | None = None,
) -> tuple[list, str | None]:
"""查询 Gemini API 支持的模型列表"""
# Gemini 使用 URL 参数传递 key不需要 headers 中的认证
base_url_clean = base_url.rstrip("/")
@@ -192,7 +192,7 @@ class GeminiChatAdapter(ChatAdapterBase):
else:
models_url = f"{base_url_clean}/v1beta/models?key={api_key}"
headers: Dict[str, str] = {}
headers: dict[str, str] = {}
if extra_headers:
headers.update(extra_headers)
@@ -242,16 +242,16 @@ class GeminiChatAdapter(ChatAdapterBase):
client: httpx.AsyncClient,
base_url: str,
api_key: str,
request_data: Dict[str, Any],
extra_headers: Optional[Dict[str, str]] = None,
request_data: dict[str, Any],
extra_headers: dict[str, str] | None = None,
# 用量计算参数
db: Optional[Any] = None,
user: Optional[Any] = None,
provider_name: Optional[str] = None,
provider_id: Optional[str] = None,
api_key_id: Optional[str] = None,
model_name: Optional[str] = None,
) -> Dict[str, Any]:
db: Any | None = None,
user: Any | None = None,
provider_name: str | None = None,
provider_id: str | None = None,
api_key_id: str | None = None,
model_name: str | None = None,
) -> dict[str, Any]:
"""测试 Gemini API 模型连接性(非流式)"""
# Gemini需要从request_data或model_name参数获取model名称
effective_model_name = model_name or request_data.get("model", "")

View File

@@ -4,7 +4,7 @@ Gemini Chat Handler
处理 Gemini API 格式的请求
"""
from typing import Any, Dict, Optional
from typing import Any
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
@@ -76,8 +76,8 @@ class GeminiChatHandler(ChatHandlerBase):
def extract_model_from_request(
self,
request_body: Dict[str, Any],
path_params: Optional[Dict[str, Any]] = None,
request_body: dict[str, Any],
path_params: dict[str, Any] | None = None,
) -> str:
"""
从请求中提取模型名 - Gemini Chat 格式实现
@@ -126,7 +126,7 @@ class GeminiChatHandler(ChatHandlerBase):
return request
def _extract_usage(self, response: Dict) -> Dict[str, int]:
def _extract_usage(self, response: dict) -> dict[str, int]:
"""
从 Gemini 响应中提取 token 使用情况
@@ -151,7 +151,7 @@ class GeminiChatHandler(ChatHandlerBase):
"cache_read_input_tokens": usage.get("cached_tokens", 0),
}
def _normalize_response(self, response: Dict) -> Dict:
def _normalize_response(self, response: dict) -> dict:
"""
规范化 Gemini 响应

View File

@@ -15,7 +15,7 @@ Gemini streamGenerateContent 常见两种返回(与上游/代理实现有关
"""
import json
from typing import Any, Dict, List, Optional, Union
from typing import Any
class GeminiStreamParser:
@@ -43,7 +43,7 @@ class GeminiStreamParser:
self._in_array = False
self._brace_depth = 0
def parse_chunk(self, chunk: Union[bytes, str]) -> List[Dict[str, Any]]:
def parse_chunk(self, chunk: bytes | str) -> list[dict[str, Any]]:
"""
解析流式数据块
@@ -58,7 +58,7 @@ class GeminiStreamParser:
else:
text = chunk
events: List[Dict[str, Any]] = []
events: list[dict[str, Any]] = []
for char in text:
if char == "[" and not self._in_array:
@@ -97,7 +97,7 @@ class GeminiStreamParser:
return events
def parse_line(self, line: str) -> Optional[Dict[str, Any]]:
def parse_line(self, line: str) -> dict[str, Any] | None:
"""
解析单行 JSON 数据
@@ -118,7 +118,7 @@ class GeminiStreamParser:
except json.JSONDecodeError:
return None
def is_done_event(self, event: Dict[str, Any]) -> bool:
def is_done_event(self, event: dict[str, Any]) -> bool:
"""
判断是否为结束事件
@@ -143,7 +143,7 @@ class GeminiStreamParser:
return False
def is_error_event(self, event: Dict[str, Any]) -> bool:
def is_error_event(self, event: dict[str, Any]) -> bool:
"""
判断是否为错误事件
@@ -171,7 +171,7 @@ class GeminiStreamParser:
return False
def extract_error_info(self, event: Dict[str, Any]) -> Optional[Dict[str, Any]]:
def extract_error_info(self, event: dict[str, Any]) -> dict[str, Any] | None:
"""
从事件中提取错误信息
@@ -208,7 +208,7 @@ class GeminiStreamParser:
return None
def get_finish_reason(self, event: Dict[str, Any]) -> Optional[str]:
def get_finish_reason(self, event: dict[str, Any]) -> str | None:
"""
获取结束原因
@@ -224,7 +224,7 @@ class GeminiStreamParser:
return str(reason) if reason is not None else None
return None
def extract_text_delta(self, event: Dict[str, Any]) -> Optional[str]:
def extract_text_delta(self, event: dict[str, Any]) -> str | None:
"""
从响应中提取文本内容
@@ -248,7 +248,7 @@ class GeminiStreamParser:
return "".join(text_parts) if text_parts else None
def extract_usage(self, event: Dict[str, Any]) -> Optional[Dict[str, int]]:
def extract_usage(self, event: dict[str, Any]) -> dict[str, int] | None:
"""
从事件中提取 token 使用量
@@ -280,7 +280,7 @@ class GeminiStreamParser:
"cached_tokens": usage_metadata.get("cachedContentTokenCount", 0),
}
def extract_model_version(self, event: Dict[str, Any]) -> Optional[str]:
def extract_model_version(self, event: dict[str, Any]) -> str | None:
"""
从响应中提取模型版本
@@ -293,7 +293,7 @@ class GeminiStreamParser:
version = event.get("modelVersion")
return str(version) if version is not None else None
def extract_safety_ratings(self, event: Dict[str, Any]) -> Optional[List[Dict[str, Any]]]:
def extract_safety_ratings(self, event: dict[str, Any]) -> list[dict[str, Any]] | None:
"""
从响应中提取安全评级

View File

@@ -4,7 +4,7 @@ Gemini CLI Adapter - 基于通用 CLI Adapter 基类的实现
继承 CliAdapterBase处理 Gemini CLI 格式的请求。
"""
from typing import Any, Dict, Optional, Tuple, Type
from typing import Any
import httpx
from fastapi import Request
@@ -29,16 +29,16 @@ class GeminiCliAdapter(CliAdapterBase):
name = "gemini.cli"
@property
def HANDLER_CLASS(self) -> Type[CliMessageHandlerBase]:
def HANDLER_CLASS(self) -> type[CliMessageHandlerBase]:
"""延迟导入 Handler 类避免循环依赖"""
from src.api.handlers.gemini_cli.handler import GeminiCliMessageHandler
return GeminiCliMessageHandler
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
def __init__(self, allowed_api_formats: list[str] | None = None):
super().__init__(allowed_api_formats or ["GEMINI_CLI"])
def extract_api_key(self, request: Request) -> Optional[str]:
def extract_api_key(self, request: Request) -> str | None:
"""
从请求中提取 API 密钥 - Gemini CLI 支持 header 和 query 两种方式
@@ -53,8 +53,8 @@ class GeminiCliAdapter(CliAdapterBase):
)
def _merge_path_params(
self, original_request_body: Dict[str, Any], path_params: Dict[str, Any] # noqa: ARG002
) -> Dict[str, Any]:
self, original_request_body: dict[str, Any], path_params: dict[str, Any] # noqa: ARG002
) -> dict[str, Any]:
"""
合并 URL 路径参数到请求体 - Gemini CLI 特化版本
@@ -74,23 +74,23 @@ class GeminiCliAdapter(CliAdapterBase):
# Gemini: 不合并任何 path_params 到请求体
return original_request_body.copy()
def _extract_message_count(self, payload: Dict[str, Any]) -> int:
def _extract_message_count(self, payload: dict[str, Any]) -> int:
"""Gemini CLI 使用 contents 字段"""
contents = payload.get("contents", [])
return len(contents) if isinstance(contents, list) else 0
def _build_audit_metadata(
self,
payload: Dict[str, Any],
path_params: Optional[Dict[str, Any]] = None,
) -> Dict[str, Any]:
payload: dict[str, Any],
path_params: dict[str, Any] | None = None,
) -> dict[str, Any]:
"""Gemini CLI 特定的审计元数据"""
# 从 path_params 获取 modelGemini 请求体不含 model
model = path_params.get("model", "unknown") if path_params else "unknown"
contents = payload.get("contents", [])
generation_config = payload.get("generation_config", {}) or {}
role_counts: Dict[str, int] = {}
role_counts: dict[str, int] = {}
for content in contents:
role = content.get("role", "unknown") if isinstance(content, dict) else "unknown"
role_counts[role] = role_counts.get(role, 0) + 1
@@ -120,8 +120,8 @@ class GeminiCliAdapter(CliAdapterBase):
client: httpx.AsyncClient,
base_url: str,
api_key: str,
extra_headers: Optional[Dict[str, str]] = None,
) -> Tuple[list, Optional[str]]:
extra_headers: dict[str, str] | None = None,
) -> tuple[list, str | None]:
"""查询 Gemini API 支持的模型列表(带 CLI User-Agent"""
# 复用 GeminiChatAdapter 的实现,添加 CLI User-Agent
cli_headers = {"User-Agent": config.internal_user_agent_gemini_cli}
@@ -136,7 +136,7 @@ class GeminiCliAdapter(CliAdapterBase):
return models, error
@classmethod
def build_endpoint_url(cls, base_url: str, request_data: Dict[str, Any], model_name: Optional[str] = None) -> str:
def build_endpoint_url(cls, base_url: str, request_data: dict[str, Any], model_name: str | None = None) -> str:
"""构建Gemini CLI API端点URL"""
effective_model_name = model_name or request_data.get("model", "")
if not effective_model_name:
@@ -152,12 +152,12 @@ class GeminiCliAdapter(CliAdapterBase):
# build_request_body 使用基类实现,通过 format_conversion_registry 自动转换 OPENAI -> GEMINI_CLI
@classmethod
def get_cli_user_agent(cls) -> Optional[str]:
def get_cli_user_agent(cls) -> str | None:
"""获取Gemini CLI User-Agent"""
return config.internal_user_agent_gemini_cli
@classmethod
def get_cli_extra_headers(cls) -> Dict[str, str]:
def get_cli_extra_headers(cls) -> dict[str, str]:
"""获取Gemini CLI额外请求头包含 x-app: cli 标识"""
headers = super().get_cli_extra_headers()
headers["x-app"] = "cli" # 标识 CLI 模式,让上游使用正确的 adapter

View File

@@ -4,7 +4,7 @@ Gemini CLI Message Handler - 基于通用 CLI Handler 基类的实现
继承 CliMessageHandlerBase处理 Gemini CLI API 格式的请求。
"""
from typing import Any, Dict, Optional
from typing import Any
from src.api.handlers.base.cli_handler_base import (
CliMessageHandlerBase,
@@ -34,8 +34,8 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
def extract_model_from_request(
self,
request_body: Dict[str, Any], # noqa: ARG002 - 基类签名要求
path_params: Optional[Dict[str, Any]] = None,
request_body: dict[str, Any], # noqa: ARG002 - 基类签名要求
path_params: dict[str, Any] | None = None,
) -> str:
"""
从请求中提取模型名 - Gemini 格式实现
@@ -57,8 +57,8 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
def prepare_provider_request_body(
self,
request_body: Dict[str, Any],
) -> Dict[str, Any]:
request_body: dict[str, Any],
) -> dict[str, Any]:
"""
准备发送给 Gemini API 的请求体 - 移除 model 字段
@@ -77,9 +77,9 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
def get_model_for_url(
self,
request_body: Dict[str, Any],
mapped_model: Optional[str],
) -> Optional[str]:
request_body: dict[str, Any],
mapped_model: str | None,
) -> str | None:
"""
Gemini 需要将 model 放入 URL 路径中
@@ -93,7 +93,7 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
# 优先使用映射后的模型名,否则使用请求体中的
return mapped_model or request_body.get("model")
def _extract_usage_from_event(self, event: Dict[str, Any]) -> Dict[str, int]:
def _extract_usage_from_event(self, event: dict[str, Any]) -> dict[str, int]:
"""
从 Gemini 事件中提取 token 使用情况
@@ -126,7 +126,7 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
self,
ctx: StreamContext,
_event_type: str,
data: Dict[str, Any],
data: dict[str, Any],
) -> None:
"""
处理 Gemini CLI 格式的流式事件
@@ -190,8 +190,8 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
def _extract_response_metadata(
self,
response: Dict[str, Any],
) -> Dict[str, Any]:
response: dict[str, Any],
) -> dict[str, Any]:
"""
从 Gemini 响应中提取元数据
@@ -203,7 +203,7 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
Returns:
包含 model_version 的元数据字典
"""
metadata: Dict[str, Any] = {}
metadata: dict[str, Any] = {}
model_version = response.get("modelVersion")
if model_version:
metadata["model_version"] = model_version

View File

@@ -4,7 +4,7 @@ OpenAI Chat Adapter - 基于 ChatAdapterBase 的 OpenAI Chat API 适配器
处理 /v1/chat/completions 端点的 OpenAI Chat 格式请求。
"""
from typing import Any, Dict, Optional, Tuple, Type
from typing import Any
import httpx
from fastapi.responses import JSONResponse
@@ -28,13 +28,13 @@ class OpenAIChatAdapter(ChatAdapterBase):
name = "openai.chat"
@property
def HANDLER_CLASS(self) -> Type[ChatHandlerBase]:
def HANDLER_CLASS(self) -> type[ChatHandlerBase]:
"""延迟导入 Handler 类避免循环依赖"""
from src.api.handlers.openai.handler import OpenAIChatHandler
return OpenAIChatHandler
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
def __init__(self, allowed_api_formats: list[str] | None = None):
super().__init__(allowed_api_formats or ["OPENAI"])
def _validate_request_body(self, original_request_body: dict, path_params: dict = None):
@@ -66,7 +66,7 @@ class OpenAIChatAdapter(ChatAdapterBase):
max_tokens=original_request_body.get("max_tokens"),
)
def _build_audit_metadata(self, payload: Dict[str, Any], request_obj) -> Dict[str, Any]:
def _build_audit_metadata(self, payload: dict[str, Any], request_obj) -> dict[str, Any]:
"""构建 OpenAI Chat 特定的审计元数据"""
role_counts = {}
for message in request_obj.messages:
@@ -105,8 +105,8 @@ class OpenAIChatAdapter(ChatAdapterBase):
client: httpx.AsyncClient,
base_url: str,
api_key: str,
extra_headers: Optional[Dict[str, str]] = None,
) -> Tuple[list, Optional[str]]:
extra_headers: dict[str, str] | None = None,
) -> tuple[list, str | None]:
"""查询 OpenAI 兼容 API 支持的模型列表"""
headers = cls.build_headers_with_extra(api_key, extra_headers)

View File

@@ -5,7 +5,7 @@ OpenAI Chat Handler - 基于通用 Chat Handler 基类的简化实现
代码量从原来的 ~1315 行减少到 ~100 行。
"""
from typing import Any, Dict, Optional
from typing import Any
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
@@ -24,8 +24,8 @@ class OpenAIChatHandler(ChatHandlerBase):
def extract_model_from_request(
self,
request_body: Dict[str, Any],
path_params: Optional[Dict[str, Any]] = None, # noqa: ARG002
request_body: dict[str, Any],
path_params: dict[str, Any] | None = None, # noqa: ARG002
) -> str:
"""
从请求中提取模型名 - OpenAI 格式实现
@@ -44,9 +44,9 @@ class OpenAIChatHandler(ChatHandlerBase):
def apply_mapped_model(
self,
request_body: Dict[str, Any],
request_body: dict[str, Any],
mapped_model: str,
) -> Dict[str, Any]:
) -> dict[str, Any]:
"""
将映射后的模型名应用到请求体
@@ -89,7 +89,7 @@ class OpenAIChatHandler(ChatHandlerBase):
return request
def _extract_usage(self, response: Dict) -> Dict[str, int]:
def _extract_usage(self, response: dict) -> dict[str, int]:
"""
从 OpenAI 响应中提取 token 使用情况
@@ -106,7 +106,7 @@ class OpenAIChatHandler(ChatHandlerBase):
"cache_read_input_tokens": 0,
}
def _normalize_response(self, response: Dict) -> Dict:
def _normalize_response(self, response: dict) -> dict:
"""
规范化 OpenAI 响应

View File

@@ -4,10 +4,9 @@ OpenAI SSE 流解析器
解析 OpenAI Chat Completions API 的 Server-Sent Events 流。
"""
from __future__ import annotations
import json
from typing import Any, Dict, List, Optional
from typing import Any
class OpenAIStreamParser:
@@ -23,7 +22,7 @@ class OpenAIStreamParser:
- 流结束时发送 data: [DONE]
"""
def parse_chunk(self, chunk: bytes | str) -> List[Dict[str, Any]]:
def parse_chunk(self, chunk: bytes | str) -> list[dict[str, Any]]:
"""
解析 SSE 数据块
@@ -38,7 +37,7 @@ class OpenAIStreamParser:
else:
text = chunk
chunks: List[Dict[str, Any]] = []
chunks: list[dict[str, Any]] = []
lines = text.strip().split("\n")
for line in lines:
@@ -64,7 +63,7 @@ class OpenAIStreamParser:
return chunks
def parse_line(self, line: str) -> Optional[Dict[str, Any]]:
def parse_line(self, line: str) -> dict[str, Any] | None:
"""
解析单行 SSE 数据
@@ -85,7 +84,7 @@ class OpenAIStreamParser:
except json.JSONDecodeError:
return None
def is_done_chunk(self, chunk: Dict[str, Any]) -> bool:
def is_done_chunk(self, chunk: dict[str, Any]) -> bool:
"""
判断是否为结束 chunk
@@ -107,7 +106,7 @@ class OpenAIStreamParser:
return False
def get_finish_reason(self, chunk: Dict[str, Any]) -> Optional[str]:
def get_finish_reason(self, chunk: dict[str, Any]) -> str | None:
"""
获取结束原因
@@ -123,7 +122,7 @@ class OpenAIStreamParser:
return str(reason) if reason is not None else None
return None
def extract_text_delta(self, chunk: Dict[str, Any]) -> Optional[str]:
def extract_text_delta(self, chunk: dict[str, Any]) -> str | None:
"""
从 chunk 中提取文本增量
@@ -145,7 +144,7 @@ class OpenAIStreamParser:
return None
def extract_tool_calls_delta(self, chunk: Dict[str, Any]) -> Optional[List[Dict[str, Any]]]:
def extract_tool_calls_delta(self, chunk: dict[str, Any]) -> list[dict[str, Any]] | None:
"""
从 chunk 中提取工具调用增量
@@ -165,7 +164,7 @@ class OpenAIStreamParser:
return tool_calls
return None
def extract_role(self, chunk: Dict[str, Any]) -> Optional[str]:
def extract_role(self, chunk: dict[str, Any]) -> str | None:
"""
从 chunk 中提取角色

View File

@@ -4,7 +4,7 @@ OpenAI CLI Adapter - 基于通用 CLI Adapter 基类的简化实现
继承 CliAdapterBase只需配置 FORMAT_ID 和 HANDLER_CLASS。
"""
from typing import Any, Dict, Optional, Tuple, Type
from typing import Any
import httpx
@@ -27,13 +27,13 @@ class OpenAICliAdapter(CliAdapterBase):
name = "openai.cli"
@property
def HANDLER_CLASS(self) -> Type[CliMessageHandlerBase]:
def HANDLER_CLASS(self) -> type[CliMessageHandlerBase]:
"""延迟导入 Handler 类避免循环依赖"""
from src.api.handlers.openai_cli.handler import OpenAICliMessageHandler
return OpenAICliMessageHandler
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
def __init__(self, allowed_api_formats: list[str] | None = None):
super().__init__(allowed_api_formats or ["OPENAI_CLI"])
# =========================================================================
@@ -46,8 +46,8 @@ class OpenAICliAdapter(CliAdapterBase):
client: httpx.AsyncClient,
base_url: str,
api_key: str,
extra_headers: Optional[Dict[str, str]] = None,
) -> Tuple[list, Optional[str]]:
extra_headers: dict[str, str] | None = None,
) -> tuple[list, str | None]:
"""查询 OpenAI 兼容 API 支持的模型列表(带 CLI User-Agent"""
# 复用 OpenAIChatAdapter 的实现,添加 CLI User-Agent
cli_headers = {"User-Agent": config.internal_user_agent_openai_cli}
@@ -62,7 +62,7 @@ class OpenAICliAdapter(CliAdapterBase):
return models, error
@classmethod
def build_endpoint_url(cls, base_url: str, request_data: Dict[str, Any], model_name: Optional[str] = None) -> str:
def build_endpoint_url(cls, base_url: str, request_data: dict[str, Any], model_name: str | None = None) -> str:
"""构建OpenAI CLI API端点URL"""
base_url = base_url.rstrip("/")
if base_url.endswith("/v1"):
@@ -74,7 +74,7 @@ class OpenAICliAdapter(CliAdapterBase):
# OPENAI -> OPENAI_CLI 无转换器,会直接透传原始请求
@classmethod
def get_cli_user_agent(cls) -> Optional[str]:
def get_cli_user_agent(cls) -> str | None:
"""获取OpenAI CLI User-Agent"""
return config.internal_user_agent_openai_cli

View File

@@ -5,7 +5,7 @@ OpenAI CLI Message Handler - 基于通用 CLI Handler 基类的简化实现
代码量从原来的 900+ 行减少到 ~100 行。
"""
from typing import Any, Dict, Optional
from typing import Any
from src.api.handlers.base.cli_handler_base import (
CliMessageHandlerBase,
@@ -32,8 +32,8 @@ class OpenAICliMessageHandler(CliMessageHandlerBase):
def extract_model_from_request(
self,
request_body: Dict[str, Any],
path_params: Optional[Dict[str, Any]] = None, # noqa: ARG002
request_body: dict[str, Any],
path_params: dict[str, Any] | None = None, # noqa: ARG002
) -> str:
"""
从请求中提取模型名 - OpenAI 格式实现
@@ -52,9 +52,9 @@ class OpenAICliMessageHandler(CliMessageHandlerBase):
def apply_mapped_model(
self,
request_body: Dict[str, Any],
request_body: dict[str, Any],
mapped_model: str,
) -> Dict[str, Any]:
) -> dict[str, Any]:
"""
OpenAI CLI (Responses API) 的 model 在请求体顶级字段。
@@ -73,7 +73,7 @@ class OpenAICliMessageHandler(CliMessageHandlerBase):
self,
ctx: StreamContext,
event_type: str,
data: Dict[str, Any],
data: dict[str, Any],
) -> None:
"""
处理 OpenAI CLI 格式的 SSE 事件
@@ -144,8 +144,8 @@ class OpenAICliMessageHandler(CliMessageHandlerBase):
def _extract_response_metadata(
self,
response: Dict[str, Any],
) -> Dict[str, Any]:
response: dict[str, Any],
) -> dict[str, Any]:
"""
从 OpenAI 响应中提取元数据
@@ -157,7 +157,7 @@ class OpenAICliMessageHandler(CliMessageHandlerBase):
Returns:
提取的元数据字典
"""
metadata: Dict[str, Any] = {}
metadata: dict[str, Any] = {}
# 提取模型名称(实际使用的模型)
if "model" in response:

View File

@@ -2,7 +2,6 @@
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from typing import Optional
from fastapi import APIRouter, Depends, HTTPException, Query, Request
from sqlalchemy.orm import Session
@@ -22,7 +21,7 @@ pipeline = ApiRequestPipeline()
@router.get("/my-audit-logs")
async def get_my_audit_logs(
request: Request,
event_type: Optional[str] = Query(None, description="事件类型筛选"),
event_type: str | None = Query(None, description="事件类型筛选"),
days: int = Query(30, description="查询天数"),
limit: int = Query(50, description="返回数量限制"),
offset: int = Query(0, ge=0, description="偏移量"),
@@ -86,7 +85,7 @@ class AuthenticatedApiAdapter(ApiAdapter):
@dataclass
class UserAuditLogsAdapter(AuthenticatedApiAdapter):
event_type: Optional[str]
event_type: str | None
days: int
limit: int
offset: int

View File

@@ -1,8 +1,7 @@
"""OAuth 管理端点(管理员)。"""
from __future__ import annotations
from typing import Any, Dict, List, Optional
from typing import Any
from fastapi import APIRouter, Depends, Request
from pydantic import BaseModel, Field, ValidationError
@@ -27,24 +26,24 @@ class SupportedOAuthType(BaseModel):
default_authorization_url: str
default_token_url: str
default_userinfo_url: str
default_scopes: List[str]
default_scopes: list[str]
class OAuthProviderUpsertRequest(BaseModel):
display_name: str = Field(..., min_length=1, max_length=100)
client_id: str = Field(..., min_length=1, max_length=255)
client_secret: Optional[str] = Field(None, max_length=2048)
client_secret: str | None = Field(None, max_length=2048)
authorization_url_override: Optional[str] = Field(None, max_length=500)
token_url_override: Optional[str] = Field(None, max_length=500)
userinfo_url_override: Optional[str] = Field(None, max_length=500)
scopes: Optional[List[str]] = None
authorization_url_override: str | None = Field(None, max_length=500)
token_url_override: str | None = Field(None, max_length=500)
userinfo_url_override: str | None = Field(None, max_length=500)
scopes: list[str] | None = None
redirect_uri: str = Field(..., min_length=1, max_length=500)
frontend_callback_url: str = Field(..., min_length=1, max_length=500)
attribute_mapping: Optional[Dict[str, Any]] = None
extra_config: Optional[Dict[str, Any]] = None
attribute_mapping: dict[str, Any] | None = None
extra_config: dict[str, Any] | None = None
is_enabled: bool = False
force: bool = False
@@ -55,14 +54,14 @@ class OAuthProviderAdminResponse(BaseModel):
display_name: str
client_id: str
has_secret: bool
authorization_url_override: Optional[str] = None
token_url_override: Optional[str] = None
userinfo_url_override: Optional[str] = None
scopes: Optional[List[str]] = None
authorization_url_override: str | None = None
token_url_override: str | None = None
userinfo_url_override: str | None = None
scopes: list[str] | None = None
redirect_uri: str
frontend_callback_url: str
attribute_mapping: Optional[Dict[str, Any]] = None
extra_config: Optional[Dict[str, Any]] = None
attribute_mapping: dict[str, Any] | None = None
extra_config: dict[str, Any] | None = None
is_enabled: bool
@@ -77,19 +76,19 @@ class OAuthProviderTestRequest(BaseModel):
"""测试请求,使用表单数据而非数据库配置"""
client_id: str = Field(..., min_length=1)
client_secret: Optional[str] = None
authorization_url_override: Optional[str] = None
token_url_override: Optional[str] = None
client_secret: str | None = None
authorization_url_override: str | None = None
token_url_override: str | None = None
redirect_uri: str = Field(..., min_length=1)
@router.get("/supported-types", response_model=List[SupportedOAuthType])
@router.get("/supported-types", response_model=list[SupportedOAuthType])
async def get_supported_types(request: Request, db: Session = Depends(get_db)) -> Any:
adapter = GetSupportedTypesAdapter()
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.get("/providers", response_model=List[OAuthProviderAdminResponse])
@router.get("/providers", response_model=list[OAuthProviderAdminResponse])
async def list_provider_configs(request: Request, db: Session = Depends(get_db)) -> Any:
adapter = ListOAuthProviderConfigsAdapter()
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)

View File

@@ -1,8 +1,7 @@
"""OAuth 公开端点(无需登录)。"""
from __future__ import annotations
from typing import Any, Optional
from typing import Any
from fastapi import APIRouter, Depends, Query, status
from sqlalchemy.orm import Session
@@ -38,10 +37,10 @@ async def oauth_authorize(provider_type: str, db: Session = Depends(get_db)) ->
async def oauth_callback(
provider_type: str,
db: Session = Depends(get_db),
code: Optional[str] = Query(None),
state: Optional[str] = Query(None),
error: Optional[str] = Query(None),
error_description: Optional[str] = Query(None),
code: str | None = Query(None),
state: str | None = Query(None),
error: str | None = Query(None),
error_description: str | None = Query(None),
) -> RedirectResponse:
"""
OAuth 回调端点。

View File

@@ -1,8 +1,7 @@
"""OAuth 用户端点(需登录)。"""
from __future__ import annotations
from typing import Any, Optional, cast
from typing import Any, cast
from fastapi import APIRouter, Depends, HTTPException, Request, status
from sqlalchemy.orm import Session
@@ -53,7 +52,7 @@ async def bind_oauth_provider(
provider_type: str,
request: Request,
db: Session = Depends(get_db),
bind_token: Optional[str] = None,
bind_token: str | None = None,
) -> RedirectResponse:
"""发起 OAuth 绑定流程,支持通过 bind_token 参数进行安全认证"""
adapter = BindOAuthProviderAdapter(provider_type=provider_type, bind_token=bind_token)
@@ -115,10 +114,10 @@ class BindOAuthProviderAdapter(AuthenticatedApiAdapter):
2. bind_token 参数 (浏览器跳转场景)
"""
def __init__(self, provider_type: str, bind_token: Optional[str] = None):
def __init__(self, provider_type: str, bind_token: str | None = None):
self.provider_type = provider_type
self.bind_token = bind_token
self._user_from_bind_token: Optional[User] = None
self._user_from_bind_token: User | None = None
@property
def mode(self) -> ApiMode: # type: ignore[override]
@@ -136,7 +135,7 @@ class BindOAuthProviderAdapter(AuthenticatedApiAdapter):
raise HTTPException(status_code=401, detail="未登录")
async def handle(self, context: ApiRequestContext) -> RedirectResponse: # type: ignore[override]
user: Optional[User] = context.user
user: User | None = context.user
# 如果使用 bind_token验证并获取用户
if self.bind_token:

View File

@@ -6,10 +6,9 @@
from collections import defaultdict
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from typing import Dict, List, Optional
from fastapi import APIRouter, Depends, Query, Request
from sqlalchemy import and_, func, or_
from sqlalchemy import and_, or_
from sqlalchemy.orm import Session, joinedload
from src.api.base.adapter import ApiAdapter, ApiMode
@@ -41,10 +40,10 @@ router = APIRouter(prefix="/api/public", tags=["System Catalog"])
pipeline = ApiRequestPipeline()
@router.get("/providers", response_model=List[PublicProviderResponse])
@router.get("/providers", response_model=list[PublicProviderResponse])
async def get_public_providers(
request: Request,
is_active: Optional[bool] = Query(None, description="过滤活跃状态"),
is_active: bool | None = Query(None, description="过滤活跃状态"),
skip: int = Query(0, description="跳过记录数"),
limit: int = Query(100, description="返回记录数限制"),
db: Session = Depends(get_db),
@@ -77,11 +76,11 @@ async def get_public_providers(
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=ApiMode.PUBLIC)
@router.get("/models", response_model=List[PublicModelResponse])
@router.get("/models", response_model=list[PublicModelResponse])
async def get_public_models(
request: Request,
provider_id: Optional[str] = Query(None, description="提供商ID过滤"),
is_active: Optional[bool] = Query(None, description="过滤活跃状态"),
provider_id: str | None = Query(None, description="提供商ID过滤"),
is_active: bool | None = Query(None, description="过滤活跃状态"),
skip: int = Query(0, description="跳过记录数"),
limit: int = Query(100, description="返回记录数限制"),
db: Session = Depends(get_db),
@@ -145,7 +144,7 @@ async def get_public_stats(request: Request, db: Session = Depends(get_db)):
async def search_models(
request: Request,
q: str = Query(..., description="搜索关键词"),
provider_id: Optional[int] = Query(None, description="提供商ID过滤"),
provider_id: int | None = Query(None, description="提供商ID过滤"),
limit: int = Query(20, description="返回记录数限制"),
db: Session = Depends(get_db),
):
@@ -234,8 +233,8 @@ async def get_public_global_models(
request: Request,
skip: int = Query(0, ge=0, description="跳过记录数"),
limit: int = Query(100, ge=1, le=1000, description="返回记录数限制"),
is_active: Optional[bool] = Query(None, description="过滤活跃状态"),
search: Optional[str] = Query(None, description="搜索关键词"),
is_active: bool | None = Query(None, description="过滤活跃状态"),
search: str | None = Query(None, description="搜索关键词"),
db: Session = Depends(get_db),
):
"""
@@ -283,7 +282,7 @@ class PublicApiAdapter(ApiAdapter):
@dataclass
class PublicProvidersAdapter(PublicApiAdapter):
is_active: Optional[bool]
is_active: bool | None
skip: int
limit: int
@@ -338,8 +337,8 @@ class PublicProvidersAdapter(PublicApiAdapter):
@dataclass
class PublicModelsAdapter(PublicApiAdapter):
provider_id: Optional[str]
is_active: Optional[bool]
provider_id: str | None
is_active: bool | None
skip: int
limit: int
@@ -426,7 +425,7 @@ class PublicStatsAdapter(PublicApiAdapter):
@dataclass
class PublicSearchModelsAdapter(PublicApiAdapter):
query: str
provider_id: Optional[int]
provider_id: int | None
limit: int
async def handle(self, context): # type: ignore[override]
@@ -508,7 +507,7 @@ class PublicApiFormatHealthMonitorAdapter(PublicApiAdapter):
.all()
)
all_formats: List[str] = []
all_formats: list[str] = []
for (api_format_enum,) in active_formats:
api_format = (
api_format_enum.value if hasattr(api_format_enum, "value") else str(api_format_enum)
@@ -525,7 +524,7 @@ class PublicApiFormatHealthMonitorAdapter(PublicApiAdapter):
)
.all()
)
endpoint_map: Dict[str, List[str]] = defaultdict(list)
endpoint_map: dict[str, list[str]] = defaultdict(list)
for api_format_enum, endpoint_id in endpoint_rows:
api_format = (
api_format_enum.value if hasattr(api_format_enum, "value") else str(api_format_enum)
@@ -551,7 +550,7 @@ class PublicApiFormatHealthMonitorAdapter(PublicApiAdapter):
.all()
)
grouped_candidates: Dict[str, List[RequestCandidate]] = {}
grouped_candidates: dict[str, list[RequestCandidate]] = {}
for candidate, api_format_enum in rows:
api_format = (
@@ -564,7 +563,7 @@ class PublicApiFormatHealthMonitorAdapter(PublicApiAdapter):
grouped_candidates[api_format].append(candidate)
# 3. 为所有活跃格式生成监控数据
monitors: List[PublicApiFormatHealthMonitor] = []
monitors: list[PublicApiFormatHealthMonitor] = []
for api_format in all_formats:
candidates = grouped_candidates.get(api_format, [])
@@ -579,7 +578,7 @@ class PublicApiFormatHealthMonitorAdapter(PublicApiAdapter):
success_rate = success_count / actual_completed if actual_completed > 0 else 1.0
# 转换为公开版事件列表(不含敏感信息如 provider_id, key_id
events: List[PublicHealthEvent] = []
events: list[PublicHealthEvent] = []
for c in candidates:
event_time = c.finished_at or c.started_at or c.created_at
events.append(
@@ -649,8 +648,8 @@ class PublicGlobalModelsAdapter(PublicApiAdapter):
skip: int
limit: int
is_active: Optional[bool]
search: Optional[str]
is_active: bool | None
search: str | None
async def handle(self, context): # type: ignore[override]
db = context.db

View File

@@ -7,8 +7,6 @@
- Authorization: Bearer (bearer) -> OpenAI 格式
"""
from typing import Optional, Tuple, Union
from fastapi import APIRouter, Depends, Query, Request
from fastapi.responses import JSONResponse
from sqlalchemy.orm import Session
@@ -35,7 +33,6 @@ from src.core.logger import logger
from src.database import get_db
from src.models.database import ApiKey, User
from src.services.auth.service import AuthService
from src.services.system.config import SystemConfigService
router = APIRouter(tags=["System Catalog"])
@@ -54,7 +51,7 @@ _ALL_CHAT_FORMATS = [
def _extract_api_key_from_request(
request: Request, definition: ApiFormatDefinition
) -> Optional[str]:
) -> str | None:
"""根据格式定义从请求中提取 API Key"""
auth_header = definition.auth_header.lower()
auth_type = definition.auth_type
@@ -76,7 +73,7 @@ def _extract_api_key_from_request(
return header_value
def _detect_api_format_and_key(request: Request) -> Tuple[str, Optional[str]]:
def _detect_api_format_and_key(request: Request) -> tuple[str, str | None]:
"""
根据请求头检测 API 格式并提取 API Key
@@ -163,7 +160,7 @@ def _build_empty_list_response(api_format: str) -> dict:
def _filter_formats_by_restrictions(
formats: list[str], restrictions: AccessRestrictions, api_format: str
) -> Tuple[list[str], Optional[dict]]:
) -> tuple[list[str], dict | None]:
"""
根据访问限制过滤 API 格式
@@ -182,7 +179,7 @@ def _filter_formats_by_restrictions(
return filtered, None
def _authenticate(db: Session, api_key: Optional[str]) -> Tuple[Optional[User], Optional[ApiKey]]:
def _authenticate(db: Session, api_key: str | None) -> tuple[User | None, ApiKey | None]:
"""
认证 API Key
@@ -248,8 +245,8 @@ def _build_auth_error_response(api_format: str) -> JSONResponse:
def _build_claude_list_response(
models: list[ModelInfo],
before_id: Optional[str],
after_id: Optional[str],
before_id: str | None,
after_id: str | None,
limit: int,
) -> dict:
"""构建 Claude 格式的列表响应"""
@@ -309,7 +306,7 @@ def _build_openai_list_response(models: list[ModelInfo]) -> dict:
def _build_gemini_list_response(
models: list[ModelInfo],
page_size: int,
page_token: Optional[str],
page_token: str | None,
) -> dict:
"""构建 Gemini 格式的列表响应"""
# 处理分页
@@ -435,14 +432,14 @@ def _build_404_response(model_id: str, api_format: str) -> JSONResponse:
async def list_models(
request: Request,
# Claude 分页参数
before_id: Optional[str] = Query(None, description="返回此 ID 之前的结果 (Claude)"),
after_id: Optional[str] = Query(None, description="返回此 ID 之后的结果 (Claude)"),
before_id: str | None = Query(None, description="返回此 ID 之前的结果 (Claude)"),
after_id: str | None = Query(None, description="返回此 ID 之后的结果 (Claude)"),
limit: int = Query(20, ge=1, le=1000, description="返回数量限制 (Claude)"),
# Gemini 分页参数
page_size: int = Query(50, alias="pageSize", ge=1, le=1000, description="每页数量 (Gemini)"),
page_token: Optional[str] = Query(None, alias="pageToken", description="分页 token (Gemini)"),
page_token: str | None = Query(None, alias="pageToken", description="分页 token (Gemini)"),
db: Session = Depends(get_db),
) -> Union[dict, JSONResponse]:
) -> dict | JSONResponse:
"""
列出可用模型(统一端点)
@@ -556,7 +553,7 @@ async def retrieve_model(
model_id: str,
request: Request,
db: Session = Depends(get_db),
) -> Union[dict, JSONResponse]:
) -> dict | JSONResponse:
"""
获取单个模型详情(统一端点)
@@ -658,9 +655,9 @@ async def retrieve_model(
async def list_models_gemini(
request: Request,
page_size: int = Query(50, alias="pageSize", ge=1, le=1000),
page_token: Optional[str] = Query(None, alias="pageToken"),
page_token: str | None = Query(None, alias="pageToken"),
db: Session = Depends(get_db),
) -> Union[dict, JSONResponse]:
) -> dict | JSONResponse:
"""
列出可用模型Gemini v1beta 专用端点)
@@ -741,7 +738,7 @@ async def get_model_gemini(
request: Request,
model_name: str,
db: Session = Depends(get_db),
) -> Union[dict, JSONResponse]:
) -> dict | JSONResponse:
"""
获取单个模型详情Gemini v1beta 专用端点)

View File

@@ -1,12 +1,11 @@
"""公开模块状态 API供登录页等使用"""
from typing import List
from fastapi import APIRouter, Depends
from pydantic import BaseModel
from sqlalchemy.orm import Session
from src.core.modules import ModuleCategory, get_module_registry
from src.core.modules import get_module_registry
from src.database import get_db
router = APIRouter(prefix="/api/modules", tags=["Modules"])
@@ -20,7 +19,7 @@ class AuthModuleInfo(BaseModel):
active: bool
@router.get("/auth-status", response_model=List[AuthModuleInfo])
@router.get("/auth-status", response_model=list[AuthModuleInfo])
async def get_auth_modules_status(db: Session = Depends(get_db)):
"""
获取认证模块状态(公开接口)

View File

@@ -5,7 +5,7 @@ System Catalog / 健康检查相关端点
"""
from datetime import datetime, timezone
from typing import Any, Dict, Optional
from typing import Any
import httpx
from fastapi import APIRouter, Depends, HTTPException, Query, Request
@@ -28,7 +28,7 @@ router = APIRouter(tags=["System Catalog"])
# ============== 辅助函数 ==============
def _as_bool(value: Optional[str], default: bool) -> bool:
def _as_bool(value: str | None, default: bool) -> bool:
"""将字符串转换为布尔值"""
if value is None:
return default
@@ -39,9 +39,9 @@ def _serialize_provider(
provider: Provider,
include_models: bool,
include_endpoints: bool,
) -> Dict[str, Any]:
) -> dict[str, Any]:
"""序列化 Provider 对象"""
provider_data: Dict[str, Any] = {
provider_data: dict[str, Any] = {
"id": provider.id,
"name": provider.name,
"is_active": provider.is_active,
@@ -81,7 +81,7 @@ def _serialize_provider(
return provider_data
def _select_provider(db: Session, provider_name: Optional[str]) -> Optional[Provider]:
def _select_provider(db: Session, provider_name: str | None) -> Provider | None:
"""选择 Provider按 provider_priority 优先级选择)"""
query = db.query(Provider).filter(Provider.is_active == True)
if provider_name:
@@ -104,7 +104,7 @@ async def service_health(db: Session = Depends(get_db)):
)
active_models = db.query(func.count(Model.id)).filter(Model.is_active == True).scalar() or 0
redis_info: Dict[str, Any] = {"status": "unknown"}
redis_info: dict[str, Any] = {"status": "unknown"}
try:
redis = await get_redis_client()
if redis:
@@ -245,9 +245,9 @@ async def provider_detail(
async def test_connection(
request: Request,
db: Session = Depends(get_db),
provider: Optional[str] = Query(None),
provider: str | None = Query(None),
model: str = Query("claude-3-haiku-20240307"),
api_format: Optional[str] = Query(None),
api_format: str | None = Query(None),
):
"""测试 Provider 连接"""
selected_provider = _select_provider(db, provider)

View File

@@ -2,7 +2,6 @@
from dataclasses import dataclass
from datetime import datetime
from typing import Optional
from fastapi import APIRouter, Depends, HTTPException, Query, Request
from fastapi.responses import JSONResponse
@@ -55,13 +54,13 @@ class CreateManagementTokenRequest(BaseModel):
"""创建 Management Token 请求"""
name: str = Field(..., min_length=1, max_length=100, description="Token 名称")
description: Optional[str] = Field(None, max_length=500, description="描述")
allowed_ips: Optional[list[str]] = Field(None, description="IP 白名单")
expires_at: Optional[datetime] = Field(None, description="过期时间")
description: str | None = Field(None, max_length=500, description="描述")
allowed_ips: list[str] | None = Field(None, description="IP 白名单")
expires_at: datetime | None = Field(None, description="过期时间")
@field_validator("allowed_ips")
@classmethod
def validate_allowed_ips(cls, v: Optional[list[str]]) -> Optional[list[str]]:
def validate_allowed_ips(cls, v: list[str] | None) -> list[str] | None:
return validate_ip_list(v)
@field_validator("expires_at", mode="before")
@@ -81,10 +80,10 @@ class UpdateManagementTokenRequest(BaseModel):
model_config = {"extra": "allow"} # 允许额外字段以便检测哪些字段被显式提供
name: Optional[str] = Field(None, min_length=1, max_length=100)
description: Optional[str] = Field(None, max_length=500)
allowed_ips: Optional[list[str]] = None
expires_at: Optional[datetime] = None
name: str | None = Field(None, min_length=1, max_length=100)
description: str | None = Field(None, max_length=500)
allowed_ips: list[str] | None = None
expires_at: datetime | None = None
# 用于追踪哪些字段被显式提供(包括显式设为 null 的情况)
_provided_fields: set[str] = set()
@@ -101,7 +100,7 @@ class UpdateManagementTokenRequest(BaseModel):
@field_validator("allowed_ips")
@classmethod
def validate_allowed_ips(cls, v: Optional[list[str]]) -> Optional[list[str]]:
def validate_allowed_ips(cls, v: list[str] | None) -> list[str] | None:
# 如果是 None表示要清空直接返回
if v is None:
return None
@@ -122,7 +121,7 @@ class UpdateManagementTokenRequest(BaseModel):
@router.get("")
async def list_my_management_tokens(
request: Request,
is_active: Optional[bool] = Query(None, description="筛选激活状态"),
is_active: bool | None = Query(None, description="筛选激活状态"),
skip: int = Query(0, ge=0),
limit: int = Query(50, ge=1, le=100),
db: Session = Depends(get_db),
@@ -347,7 +346,7 @@ class ListMyManagementTokensAdapter(ManagementTokenApiAdapter):
"""列出用户的 Management Tokens"""
name: str = "list_my_management_tokens"
is_active: Optional[bool] = None
is_active: bool | None = None
skip: int = 0
limit: int = 50

View File

@@ -2,14 +2,12 @@
from dataclasses import dataclass
from datetime import datetime, timezone
from typing import Optional
from fastapi import APIRouter, Depends, HTTPException, Query, Request
from pydantic import ValidationError
from sqlalchemy import and_, func
from sqlalchemy.orm import Session
from src.api.base.adapter import ApiAdapter, ApiMode
from src.api.base.authenticated_adapter import AuthenticatedApiAdapter
from src.api.base.pipeline import ApiRequestPipeline
from src.core.crypto import crypto_service
@@ -170,9 +168,9 @@ async def toggle_my_api_key(key_id: str, request: Request, db: Session = Depends
@router.get("/usage")
async def get_my_usage(
request: Request,
start_date: Optional[datetime] = Query(None, description="开始时间ISO 格式)"),
end_date: Optional[datetime] = Query(None, description="结束时间ISO 格式)"),
search: Optional[str] = Query(None, description="搜索关键词(密钥名、模型名)"),
start_date: datetime | None = Query(None, description="开始时间ISO 格式)"),
end_date: datetime | None = Query(None, description="结束时间ISO 格式)"),
search: str | None = Query(None, description="搜索关键词(密钥名、模型名)"),
limit: int = Query(100, ge=1, le=200, description="每页记录数默认100最大200"),
offset: int = Query(0, ge=0, le=2000, description="偏移量用于分页最大2000"),
db: Session = Depends(get_db),
@@ -200,7 +198,7 @@ async def get_my_usage(
@router.get("/usage/active")
async def get_my_active_requests(
request: Request,
ids: Optional[str] = Query(None, description="请求 ID 列表,逗号分隔"),
ids: str | None = Query(None, description="请求 ID 列表,逗号分隔"),
db: Session = Depends(get_db),
):
"""
@@ -268,7 +266,7 @@ async def list_available_models(
request: Request,
skip: int = Query(0, ge=0, description="跳过记录数"),
limit: int = Query(100, ge=1, le=1000, description="返回记录数限制"),
search: Optional[str] = Query(None, description="搜索关键词"),
search: str | None = Query(None, description="搜索关键词"),
db: Session = Depends(get_db),
):
"""
@@ -721,9 +719,9 @@ class ToggleMyApiKeyAdapter(AuthenticatedApiAdapter):
class GetUsageAdapter(AuthenticatedApiAdapter):
"""获取用户使用统计的适配器"""
start_date: Optional[datetime]
end_date: Optional[datetime]
search: Optional[str] = None
start_date: datetime | None
end_date: datetime | None
search: str | None = None
limit: int = 100
offset: int = 0
@@ -983,7 +981,7 @@ class GetUsageAdapter(AuthenticatedApiAdapter):
class GetActiveRequestsAdapter(AuthenticatedApiAdapter):
"""轻量级活跃请求状态查询适配器(用于用户端轮询)"""
ids: Optional[str] = None
ids: str | None = None
async def handle(self, context): # type: ignore[override]
from src.services.usage import UsageService
@@ -1045,7 +1043,7 @@ class ListAvailableModelsAdapter(AuthenticatedApiAdapter):
skip: int
limit: int
search: Optional[str]
search: str | None
async def handle(self, context): # type: ignore[override]
from sqlalchemy import or_
@@ -1220,7 +1218,6 @@ class ListAvailableProvidersAdapter(AuthenticatedApiAdapter):
async def handle(self, context): # type: ignore[override]
from sqlalchemy.orm import selectinload
from src.models.database import ProviderEndpoint
db = context.db