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

@@ -7,7 +7,7 @@ WORKDIR /app
COPY frontend/ ./frontend/ COPY frontend/ ./frontend/
RUN cd frontend && npm run build RUN cd frontend && npm run build
# ==================== 运行时镜像 ==================== # ==================== 运行时镜像 ====================
FROM python:3.12-slim FROM python:3.14-slim
WORKDIR /app WORKDIR /app
# 运行时依赖(无 gcc/nodejs/npm # 运行时依赖(无 gcc/nodejs/npm
RUN apt-get update && apt-get install -y \ RUN apt-get update && apt-get install -y \
@@ -17,7 +17,7 @@ RUN apt-get update && apt-get install -y \
curl \ curl \
&& rm -rf /var/lib/apt/lists/* && rm -rf /var/lib/apt/lists/*
# 从 base 镜像复制 Python 包 # 从 base 镜像复制 Python 包
COPY --from=builder /usr/local/lib/python3.12/site-packages /usr/local/lib/python3.12/site-packages COPY --from=builder /usr/local/lib/python3.14/site-packages /usr/local/lib/python3.14/site-packages
# 只复制需要的 Python 可执行文件 # 只复制需要的 Python 可执行文件
COPY --from=builder /usr/local/bin/gunicorn /usr/local/bin/ COPY --from=builder /usr/local/bin/gunicorn /usr/local/bin/
COPY --from=builder /usr/local/bin/uvicorn /usr/local/bin/ COPY --from=builder /usr/local/bin/uvicorn /usr/local/bin/

View File

@@ -10,7 +10,7 @@ COPY frontend/ ./frontend/
RUN cd frontend && npm run build RUN cd frontend && npm run build
# ==================== 运行时镜像 ==================== # ==================== 运行时镜像 ====================
FROM python:3.12-slim FROM python:3.14-slim
WORKDIR /app WORKDIR /app
@@ -24,7 +24,7 @@ RUN sed -i 's/deb.debian.org/mirrors.tuna.tsinghua.edu.cn/g' /etc/apt/sources.li
&& rm -rf /var/lib/apt/lists/* && rm -rf /var/lib/apt/lists/*
# 从 base 镜像复制 Python 包 # 从 base 镜像复制 Python 包
COPY --from=builder /usr/local/lib/python3.12/site-packages /usr/local/lib/python3.12/site-packages COPY --from=builder /usr/local/lib/python3.14/site-packages /usr/local/lib/python3.14/site-packages
# 只复制需要的 Python 可执行文件 # 只复制需要的 Python 可执行文件
COPY --from=builder /usr/local/bin/gunicorn /usr/local/bin/ COPY --from=builder /usr/local/bin/gunicorn /usr/local/bin/

View File

@@ -2,7 +2,7 @@
# 用于 GitHub Actions CI 构建(不使用国内镜像源) # 用于 GitHub Actions CI 构建(不使用国内镜像源)
# 构建命令: docker build -f Dockerfile.base -t aether-base:latest . # 构建命令: docker build -f Dockerfile.base -t aether-base:latest .
# 只在 pyproject.toml 或 frontend/package*.json 变化时需要重建 # 只在 pyproject.toml 或 frontend/package*.json 变化时需要重建
FROM python:3.12-slim FROM python:3.14-slim
WORKDIR /app WORKDIR /app

View File

@@ -1,7 +1,7 @@
# 构建镜像:编译环境 + 预编译的依赖(国内镜像源版本) # 构建镜像:编译环境 + 预编译的依赖(国内镜像源版本)
# 构建命令: docker build -f Dockerfile.base.local -t aether-base:latest . # 构建命令: docker build -f Dockerfile.base.local -t aether-base:latest .
# 只在 pyproject.toml 或 frontend/package*.json 变化时需要重建 # 只在 pyproject.toml 或 frontend/package*.json 变化时需要重建
FROM python:3.12-slim FROM python:3.14-slim
WORKDIR /app WORKDIR /app

View File

@@ -15,13 +15,11 @@ classifiers = [
"Intended Audience :: Developers", "Intended Audience :: Developers",
"License :: Other/Proprietary License", "License :: Other/Proprietary License",
"Programming Language :: Python :: 3", "Programming Language :: Python :: 3",
"Programming Language :: Python :: 3.8",
"Programming Language :: Python :: 3.9",
"Programming Language :: Python :: 3.10",
"Programming Language :: Python :: 3.11",
"Programming Language :: Python :: 3.12", "Programming Language :: Python :: 3.12",
"Programming Language :: Python :: 3.13",
"Programming Language :: Python :: 3.14",
] ]
requires-python = ">=3.9" requires-python = ">=3.12"
dependencies = [ dependencies = [
"fastapi[standard]>=0.115.11", "fastapi[standard]>=0.115.11",
"uvicorn>=0.34.0", "uvicorn>=0.34.0",
@@ -83,7 +81,7 @@ dev-dependencies = [
[tool.black] [tool.black]
line-length = 100 line-length = 100
target-version = ['py38'] target-version = ['py312']
[tool.isort] [tool.isort]
profile = "black" profile = "black"
@@ -99,7 +97,7 @@ source = "vcs"
version-file = "src/_version.py" version-file = "src/_version.py"
[tool.mypy] [tool.mypy]
python_version = "3.9" python_version = "3.12"
warn_return_any = true warn_return_any = true
warn_unused_configs = true warn_unused_configs = true
disallow_untyped_defs = true disallow_untyped_defs = true

View File

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

View File

@@ -5,7 +5,6 @@
import os import os
from datetime import datetime, timedelta, timezone from datetime import datetime, timedelta, timezone
from typing import Optional
from zoneinfo import ZoneInfo from zoneinfo import ZoneInfo
from fastapi import APIRouter, Depends, HTTPException, Query, Request 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")) 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 对象。 """解析过期日期字符串为 datetime 对象。
Args: Args:
@@ -70,7 +69,7 @@ async def list_standalone_api_keys(
request: Request, request: Request,
skip: int = Query(0, ge=0), skip: int = Query(0, ge=0),
limit: int = Query(100, ge=1, le=500), limit: int = Query(100, ge=1, le=500),
is_active: Optional[bool] = None, is_active: bool | None = None,
db: Session = Depends(get_db), db: Session = Depends(get_db),
): ):
""" """
@@ -330,7 +329,7 @@ class AdminListStandaloneKeysAdapter(AdminApiAdapter):
self, self,
skip: int, skip: int,
limit: int, limit: int,
is_active: Optional[bool], is_active: bool | None,
): ):
self.skip = skip self.skip = skip
self.limit = limit self.limit = limit

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

@@ -3,8 +3,7 @@ models.dev 外部模型数据代理
""" """
import json import json
from typing import Any
from typing import Any, Optional
import httpx import httpx
from fastapi import APIRouter, Depends, HTTPException 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 获取缓存数据"""
redis = await get_redis_client() redis = await get_redis_client()
if redis is None: if redis is None:

View File

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

View File

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

View File

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

View File

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

View File

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

View File

@@ -4,7 +4,6 @@
from dataclasses import dataclass from dataclasses import dataclass
from datetime import datetime from datetime import datetime
from typing import List, Optional
from fastapi import APIRouter, Depends, HTTPException, Query, Request from fastapi import APIRouter, Depends, HTTPException, Query, Request
from pydantic import BaseModel, ConfigDict from pydantic import BaseModel, ConfigDict
@@ -28,29 +27,29 @@ class CandidateResponse(BaseModel):
request_id: str request_id: str
candidate_index: int candidate_index: int
retry_index: int = 0 # 重试序号从0开始 retry_index: int = 0 # 重试序号从0开始
provider_id: Optional[str] = None provider_id: str | None = None
provider_name: Optional[str] = None provider_name: str | None = None
provider_website: Optional[str] = None # Provider 官网 provider_website: str | None = None # Provider 官网
endpoint_id: Optional[str] = None endpoint_id: str | None = None
endpoint_name: Optional[str] = None # 端点显示名称api_format endpoint_name: str | None = None # 端点显示名称api_format
key_id: Optional[str] = None key_id: str | None = None
key_name: Optional[str] = None # 密钥名称 key_name: str | None = None # 密钥名称
key_preview: Optional[str] = None # 密钥脱敏预览(如 sk-***abc key_preview: str | None = None # 密钥脱敏预览(如 sk-***abc
key_capabilities: Optional[dict] = None # Key 支持的能力 key_capabilities: dict | None = None # Key 支持的能力
required_capabilities: Optional[dict] = None # 请求实际需要的能力标签 required_capabilities: dict | None = None # 请求实际需要的能力标签
status: str # 'pending', 'success', 'failed', 'skipped' status: str # 'pending', 'success', 'failed', 'skipped'
skip_reason: Optional[str] = None skip_reason: str | None = None
is_cached: bool = False is_cached: bool = False
# 执行结果字段 # 执行结果字段
status_code: Optional[int] = None status_code: int | None = None
error_type: Optional[str] = None error_type: str | None = None
error_message: Optional[str] = None error_message: str | None = None
latency_ms: Optional[int] = None latency_ms: int | None = None
concurrent_requests: Optional[int] = None concurrent_requests: int | None = None
extra_data: Optional[dict] = None extra_data: dict | None = None
created_at: datetime created_at: datetime
started_at: Optional[datetime] = None started_at: datetime | None = None
finished_at: Optional[datetime] = None finished_at: datetime | None = None
model_config = ConfigDict(from_attributes=True) model_config = ConfigDict(from_attributes=True)
@@ -62,7 +61,7 @@ class RequestTraceResponse(BaseModel):
total_candidates: int total_candidates: int
final_status: str # 'success', 'failed', 'cancelled', 'streaming', 'pending' final_status: str # 'success', 'failed', 'cancelled', 'streaming', 'pending'
total_latency_ms: int total_latency_ms: int
candidates: List[CandidateResponse] candidates: list[CandidateResponse]
@router.get("/{request_id}", response_model=RequestTraceResponse) @router.get("/{request_id}", response_model=RequestTraceResponse)
@@ -253,7 +252,7 @@ class AdminGetRequestTraceAdapter(AdminApiAdapter):
key_preview_map[k.id] = "***" key_preview_map[k.id] = "***"
# 构建 candidate 响应列表 # 构建 candidate 响应列表
candidate_responses: List[CandidateResponse] = [] candidate_responses: list[CandidateResponse] = []
for candidate in candidates: for candidate in candidates:
provider_name = ( provider_name = (
provider_map.get(candidate.provider_id) if candidate.provider_id else None 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 dataclasses import asdict, is_dataclass
from typing import Any, Dict, List, Optional from typing import Any
from fastapi import APIRouter, Depends, HTTPException from fastapi import APIRouter, Depends, HTTPException
from pydantic import BaseModel, Field from pydantic import BaseModel, Field
@@ -18,9 +18,7 @@ from sqlalchemy.orm import Session
from src.database import get_db from src.database import get_db
from src.models.database import Provider, User from src.models.database import Provider, User
from src.services.provider_ops import ( from src.services.provider_ops import (
ActionStatus,
ConnectorAuthType, ConnectorAuthType,
ConnectorStatus,
ProviderActionType, ProviderActionType,
ProviderOpsConfig, ProviderOpsConfig,
ProviderOpsService, ProviderOpsService,
@@ -40,46 +38,46 @@ class ArchitectureInfo(BaseModel):
architecture_id: str architecture_id: str
display_name: str display_name: str
description: str description: str
supported_auth_types: List[Dict[str, str]] supported_auth_types: list[dict[str, str]]
supported_actions: List[Dict[str, Any]] supported_actions: list[dict[str, Any]]
default_connector: Optional[str] default_connector: str | None
class ConnectorConfigRequest(BaseModel): class ConnectorConfigRequest(BaseModel):
"""连接器配置请求""" """连接器配置请求"""
auth_type: str = Field(..., description="认证类型") auth_type: str = Field(..., description="认证类型")
config: 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="凭据信息") credentials: dict[str, Any] = Field(default_factory=dict, description="凭据信息")
class ActionConfigRequest(BaseModel): class ActionConfigRequest(BaseModel):
"""操作配置请求""" """操作配置请求"""
enabled: bool = Field(True, description="是否启用") 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): class SaveConfigRequest(BaseModel):
"""保存配置请求""" """保存配置请求"""
architecture_id: str = Field("generic_api", description="架构 ID") 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 connector: ConnectorConfigRequest
actions: Dict[str, ActionConfigRequest] = Field(default_factory=dict) actions: dict[str, ActionConfigRequest] = Field(default_factory=dict)
schedule: Dict[str, str] = Field(default_factory=dict, description="定时任务配置") schedule: dict[str, str] = Field(default_factory=dict, description="定时任务配置")
class ConnectRequest(BaseModel): class ConnectRequest(BaseModel):
"""连接请求""" """连接请求"""
credentials: Optional[Dict[str, Any]] = Field(None, description="凭据(可选,使用已保存的)") credentials: dict[str, Any] | None = Field(None, description="凭据(可选,使用已保存的)")
class ExecuteActionRequest(BaseModel): class ExecuteActionRequest(BaseModel):
"""执行操作请求""" """执行操作请求"""
config: Optional[Dict[str, Any]] = Field(None, description="操作配置(覆盖默认)") config: dict[str, Any] | None = Field(None, description="操作配置(覆盖默认)")
class ConnectionStatusResponse(BaseModel): class ConnectionStatusResponse(BaseModel):
@@ -87,9 +85,9 @@ class ConnectionStatusResponse(BaseModel):
status: str status: str
auth_type: str auth_type: str
connected_at: Optional[str] connected_at: str | None
expires_at: Optional[str] expires_at: str | None
last_error: Optional[str] last_error: str | None
class ActionResultResponse(BaseModel): class ActionResultResponse(BaseModel):
@@ -97,10 +95,10 @@ class ActionResultResponse(BaseModel):
status: str status: str
action_type: str action_type: str
data: Optional[Any] data: Any | None
message: Optional[str] message: str | None
executed_at: str executed_at: str
response_time_ms: Optional[int] response_time_ms: int | None
cache_ttl_seconds: int cache_ttl_seconds: int
@@ -109,9 +107,9 @@ class ProviderOpsStatusResponse(BaseModel):
provider_id: str provider_id: str
is_configured: bool is_configured: bool
architecture_id: Optional[str] architecture_id: str | None
connection_status: ConnectionStatusResponse connection_status: ConnectionStatusResponse
enabled_actions: List[str] enabled_actions: list[str]
class ProviderOpsConfigResponse(BaseModel): class ProviderOpsConfigResponse(BaseModel):
@@ -119,17 +117,17 @@ class ProviderOpsConfigResponse(BaseModel):
provider_id: str provider_id: str
is_configured: bool is_configured: bool
architecture_id: Optional[str] = None architecture_id: str | None = None
base_url: Optional[str] = None base_url: str | None = None
connector: Optional[Dict[str, Any]] = None # 脱敏后的连接器配置 connector: dict[str, Any] | None = None # 脱敏后的连接器配置
class VerifyAuthResponse(BaseModel): class VerifyAuthResponse(BaseModel):
"""验证认证响应""" """验证认证响应"""
success: bool success: bool
message: Optional[str] = None message: str | None = None
data: Optional[Dict[str, Any]] = None data: dict[str, Any] | None = None
# ==================== Helper Functions ==================== # ==================== Helper Functions ====================
@@ -147,7 +145,7 @@ def _serialize_data(data: Any) -> Any:
# ==================== Routes ==================== # ==================== Routes ====================
@router.get("/architectures", response_model=List[ArchitectureInfo]) @router.get("/architectures", response_model=list[ArchitectureInfo])
async def list_architectures(_: User = Depends(require_admin)): async def list_architectures(_: User = Depends(require_admin)):
"""获取所有可用的架构""" """获取所有可用的架构"""
registry = get_registry() registry = get_registry()
@@ -490,7 +488,7 @@ async def checkin(
@router.post("/batch/balance") @router.post("/batch/balance")
async def batch_query_balance( async def batch_query_balance(
provider_ids: Optional[List[str]] = None, provider_ids: list[str] | None = None,
db: Session = Depends(get_db), db: Session = Depends(get_db),
_: User = Depends(require_admin), _: User = Depends(require_admin),
): ):

View File

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

View File

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

View File

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

View File

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

View File

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

View File

@@ -4,8 +4,6 @@ IP 安全管理接口
提供 IP 黑白名单管理和速率限制统计 提供 IP 黑白名单管理和速率限制统计
""" """
from typing import List, Optional
from fastapi import APIRouter, Depends, HTTPException, Request, status from fastapi import APIRouter, Depends, HTTPException, Request, status
from pydantic import BaseModel, Field, ValidationError from pydantic import BaseModel, Field, ValidationError
from sqlalchemy.orm import Session 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.authenticated_adapter import AuthenticatedApiAdapter
from src.api.base.pipeline import ApiRequestPipeline from src.api.base.pipeline import ApiRequestPipeline
from src.core.exceptions import InvalidRequestException, translate_pydantic_error from src.core.exceptions import InvalidRequestException, translate_pydantic_error
from src.core.logger import logger
from src.database import get_db from src.database import get_db
from src.services.rate_limit.ip_limiter import IPRateLimiter from src.services.rate_limit.ip_limiter import IPRateLimiter
@@ -30,7 +27,7 @@ class AddIPToBlacklistRequest(BaseModel):
ip_address: str = Field(..., description="IP 地址") ip_address: str = Field(..., description="IP 地址")
reason: str = Field(..., min_length=1, max_length=200, description="加入黑名单的原因") 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): class RemoveIPFromBlacklistRequest(BaseModel):

View File

@@ -1,10 +1,8 @@
"""系统设置API端点。""" """系统设置API端点。"""
from __future__ import annotations
import copy import copy
from dataclasses import dataclass from dataclasses import dataclass
from typing import Optional
from fastapi import APIRouter, Depends, HTTPException, Request, status from fastapi import APIRouter, Depends, HTTPException, Request, status
from pydantic import ValidationError from pydantic import ValidationError
@@ -649,9 +647,8 @@ class AdminSystemStatsAdapter(AdminApiAdapter):
class AdminTriggerCleanupAdapter(AdminApiAdapter): class AdminTriggerCleanupAdapter(AdminApiAdapter):
async def handle(self, context): # type: ignore[override] 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 from src.services.system.maintenance_scheduler import get_maintenance_scheduler

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

@@ -1,7 +1,6 @@
from __future__ import annotations
from dataclasses import asdict, dataclass 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 from sqlalchemy.orm import Query
@@ -19,7 +18,7 @@ class PaginationMeta:
return asdict(self) 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并返回总数与结果列表。 对 SQLAlchemy 查询应用 limit/offset并返回总数与结果列表。
""" """
@@ -30,7 +29,7 @@ def paginate_query(query: Query, limit: int, offset: int) -> Tuple[int, List[T]]
def paginate_sequence( def paginate_sequence(
items: Sequence[T], limit: int, offset: int 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 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。 构建标准分页响应 payload。
""" """

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

@@ -8,7 +8,6 @@ StreamSmoother 使用这些提取器来处理不同格式的 SSE 事件。
import copy import copy
import json import json
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from typing import Optional
class ContentExtractor(ABC): class ContentExtractor(ABC):
@@ -20,7 +19,7 @@ class ContentExtractor(ABC):
""" """
@abstractmethod @abstractmethod
def extract_content(self, data: dict) -> Optional[str]: def extract_content(self, data: dict) -> str | None:
""" """
从 SSE 数据中提取可拆分的文本内容 从 SSE 数据中提取可拆分的文本内容
@@ -64,7 +63,7 @@ class OpenAIContentExtractor(ContentExtractor):
- 只在 delta 仅包含 role/content 时允许拆分,避免破坏 tool_calls 等结构 - 只在 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): if not isinstance(data, dict):
return None return None
@@ -115,7 +114,7 @@ class OpenAIContentExtractor(ContentExtractor):
new_choices.append(new_choice) new_choices.append(new_choice)
new_data["choices"] = new_choices 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): class ClaudeContentExtractor(ContentExtractor):
@@ -127,7 +126,7 @@ class ClaudeContentExtractor(ContentExtractor):
- 数据结构: delta.type=text_delta, delta.text - 数据结构: 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): if not isinstance(data, dict):
return None return None
@@ -165,9 +164,7 @@ class ClaudeContentExtractor(ContentExtractor):
# Claude 格式需要 event: 前缀 # Claude 格式需要 event: 前缀
event_name = event_type or "content_block_delta" event_name = event_type or "content_block_delta"
return f"event: {event_name}\ndata: {json.dumps(new_data, ensure_ascii=False)}\n\n".encode( return f"event: {event_name}\ndata: {json.dumps(new_data, ensure_ascii=False)}\n\n".encode()
"utf-8"
)
class GeminiContentExtractor(ContentExtractor): 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): if not isinstance(data, dict):
return None return None
@@ -226,7 +223,7 @@ class GeminiContentExtractor(ContentExtractor):
if "parts" in content and content["parts"]: if "parts" in content and content["parts"]:
content["parts"][0]["text"] = new_content 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: 协调整个流程 - EndpointCheckOrchestrator: 协调整个流程
""" """
from __future__ import annotations
from dataclasses import dataclass from dataclasses import dataclass
from typing import Any, AsyncIterator, Dict, Iterable, Optional, Union, List from typing import Any
from abc import ABC, abstractmethod from collections.abc import Iterable
import time import time
import uuid import uuid
import json import json
from functools import lru_cache
import asyncio import asyncio
from collections import defaultdict
import httpx 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 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) 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( def build_safe_headers(
base_headers: Dict[str, str], base_headers: dict[str, str],
extra_headers: Optional[Dict[str, str]], extra_headers: dict[str, str] | None,
protected_keys: Iterable[str], protected_keys: Iterable[str],
) -> Dict[str, str]: ) -> dict[str, str]:
""" """
合并 extra_headers但防止覆盖 protected_keys大小写不敏感 合并 extra_headers但防止覆盖 protected_keys大小写不敏感
""" """
@@ -60,16 +59,16 @@ async def run_endpoint_check(
*, *,
client: httpx.AsyncClient, # 保持兼容性,但内部不使用 client: httpx.AsyncClient, # 保持兼容性,但内部不使用
url: str, url: str,
headers: Dict[str, str], headers: dict[str, str],
json_body: Dict[str, Any], json_body: dict[str, Any],
api_format: str, api_format: str,
provider_name: Optional[str] = None, provider_name: str | None = None,
model_name: Optional[str] = None, model_name: str | None = None,
api_key_id: Optional[str] = None, api_key_id: str | None = None,
provider_id: Optional[str] = None, provider_id: str | None = None,
db: Optional[Any] = None, # Session对象需要时才导入 db: Any | None = None, # Session对象需要时才导入
user: Optional[Any] = None, # User对象 user: Any | None = None, # User对象
) -> Dict[str, Any]: ) -> dict[str, Any]:
""" """
执行端点检查(重构版本,使用新的架构): 执行端点检查(重构版本,使用新的架构):
- 使用新的架构类来分离关注点 - 使用新的架构类来分离关注点
@@ -123,21 +122,21 @@ async def _calculate_and_record_usage(
provider_id: str, provider_id: str,
api_key_id: str, api_key_id: str,
model_name: str, model_name: str,
request_data: Dict[str, Any], request_data: dict[str, Any],
response_data: Optional[Dict[str, Any]], response_data: dict[str, Any] | None,
request_id: str, request_id: str,
response_time_ms: int, response_time_ms: int,
request_headers: Dict[str, str], request_headers: dict[str, str],
response_headers: Optional[Dict[str, str]] = None, response_headers: dict[str, str] | None = None,
status_code: int = 0, status_code: int = 0,
error_message: Optional[str] = None, error_message: str | None = None,
# 新增支持直接传递token数据 # 新增支持直接传递token数据
input_tokens: Optional[int] = None, input_tokens: int | None = None,
output_tokens: Optional[int] = None, output_tokens: int | None = None,
cache_creation_input_tokens: Optional[int] = None, cache_creation_input_tokens: int | None = None,
cache_read_input_tokens: Optional[int] = None, cache_read_input_tokens: int | None = None,
api_format: Optional[str] = None, api_format: str | None = None,
) -> Dict[str, Any]: ) -> dict[str, Any]:
""" """
计算并记录用量数据(遗留函数) 计算并记录用量数据(遗留函数)
@@ -149,7 +148,7 @@ async def _calculate_and_record_usage(
""" """
from src.services.usage.service import UsageService from src.services.usage.service import UsageService
from src.services.request.candidate import RequestCandidateService 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对象不是用户API Key
provider_api_key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == api_key_id).first() 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计数信息 从响应中提取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计数方法简单估算 回退的Token计数方法简单估算
@@ -508,16 +507,16 @@ def _fallback_token_counting(request_data: Dict[str, Any], response_data: Option
class EndpointCheckRequest: class EndpointCheckRequest:
"""端点检查请求数据类""" """端点检查请求数据类"""
url: str url: str
headers: Dict[str, str] headers: dict[str, str]
json_body: Dict[str, Any] json_body: dict[str, Any]
api_format: str api_format: str
provider_name: Optional[str] = None provider_name: str | None = None
model_name: Optional[str] = None model_name: str | None = None
api_key_id: Optional[str] = None api_key_id: str | None = None
provider_id: Optional[str] = None provider_id: str | None = None
db: Optional[Any] = None db: Any | None = None
user: Optional[Any] = None user: Any | None = None
request_id: Optional[str] = None request_id: str | None = None
timeout: float = 30.0 timeout: float = 30.0
@@ -525,12 +524,12 @@ class EndpointCheckRequest:
class EndpointCheckResult: class EndpointCheckResult:
"""端点检查结果数据类""" """端点检查结果数据类"""
status_code: int status_code: int
headers: Dict[str, str] headers: dict[str, str]
response_time_ms: int response_time_ms: int
request_id: str request_id: str
response_data: Optional[Dict[str, Any]] = None response_data: dict[str, Any] | None = None
error_message: Optional[str] = None error_message: str | None = None
usage_data: Optional[Dict[str, Any]] = None usage_data: dict[str, Any] | None = None
class HttpRequestExecutor: class HttpRequestExecutor:
@@ -613,7 +612,7 @@ class UsageCalculator:
return _extract_tokens_from_response(api_identifier, result.response_data) return _extract_tokens_from_response(api_identifier, result.response_data)
@staticmethod @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计数方法简单估算"""
# 估算输入Token # 估算输入Token
messages = request_data.get("messages", request_data.get("contents", [])) 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): def __init__(self, batch_size: int = 10, flush_interval: float = 2.0):
self.batch_size = batch_size self.batch_size = batch_size
self.flush_interval = flush_interval self.flush_interval = flush_interval
self.pending_records: List[Dict[str, Any]] = [] self.pending_records: list[dict[str, Any]] = []
self._flush_task: Optional[asyncio.Task] = None self._flush_task: asyncio.Task | None = None
self._lock = asyncio.Lock() self._lock = asyncio.Lock()
self._running = True 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: async with self._lock:
self.pending_records.append(usage_data) 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: def get_batch_recorder() -> AsyncBatchUsageRecorder:
"""获取全局批处理器实例""" """获取全局批处理器实例"""
@@ -756,7 +755,7 @@ def get_batch_recorder() -> AsyncBatchUsageRecorder:
class EndpointCheckError(Exception): 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) super().__init__(message)
self.message = message self.message = message
self.error_type = error_type self.error_type = error_type
@@ -765,22 +764,22 @@ class EndpointCheckError(Exception):
class NetworkError(EndpointCheckError): 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) super().__init__(message, "network_error", 0, details)
class AuthenticationError(EndpointCheckError): 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) super().__init__(message, "authentication_error", 401, details)
class RateLimitError(EndpointCheckError): 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) super().__init__(message, "rate_limit_error", 429, details)
class UpstreamError(EndpointCheckError): 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) super().__init__(message, "upstream_error", status_code, details)
@@ -982,7 +981,7 @@ class EndpointCheckConfig:
retry_on_timeouts: bool = True retry_on_timeouts: bool = True
@classmethod @classmethod
def from_env(cls) -> 'EndpointCheckConfig': def from_env(cls) -> EndpointCheckConfig:
"""从环境变量创建配置""" """从环境变量创建配置"""
import os import os
@@ -1004,7 +1003,7 @@ class EndpointCheckConfig:
) )
@classmethod @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)}) return cls(**{k: v for k, v in config_dict.items() if hasattr(cls, k)})
@@ -1012,7 +1011,7 @@ class EndpointCheckConfig:
class ConfigurableEndpointChecker: class ConfigurableEndpointChecker:
"""可配置的端点检查器""" """可配置的端点检查器"""
def __init__(self, config: Optional[EndpointCheckConfig] = None): def __init__(self, config: EndpointCheckConfig | None = None):
self.config = config or EndpointCheckConfig() self.config = config or EndpointCheckConfig()
self.executor = HttpRequestExecutor(timeout=self.config.timeout) self.executor = HttpRequestExecutor(timeout=self.config.timeout)
self.usage_calculator = UsageCalculator() 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 global _global_configured_checker
if _global_configured_checker is None or config is not None: 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: class EndpointCheckOrchestrator:
"""端点检查协调器 - 协调整个流程""" """端点检查协调器 - 协调整个流程"""
def __init__(self, executor: Optional[HttpRequestExecutor] = None, def __init__(self, executor: HttpRequestExecutor | None = None,
usage_calculator: Optional[UsageCalculator] = None): usage_calculator: UsageCalculator | None = None):
self.executor = executor or HttpRequestExecutor() self.executor = executor or HttpRequestExecutor()
self.usage_calculator = usage_calculator or UsageCalculator() self.usage_calculator = usage_calculator or UsageCalculator()

View File

@@ -6,7 +6,7 @@
""" """
import re import re
from typing import Any, Dict, Optional, Tuple, Type from typing import Any
from src.api.handlers.base.response_parser import ( from src.api.handlers.base.response_parser import (
ParsedChunk, 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 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 但在响应体中包含错误) 检查响应中是否存在嵌套错误(某些代理服务返回 HTTP 200 但在响应体中包含错误)
@@ -62,7 +62,7 @@ def _check_nested_error(response: Dict[str, Any]) -> Tuple[bool, Optional[Dict[s
return False, None 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.name = "OPENAI"
self.api_format = "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(): if not line or not line.strip():
return None return None
@@ -186,7 +186,7 @@ class OpenAIResponseParser(ResponseParser):
return chunk 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( result = ParsedResponse(
raw_response=response, raw_response=response,
status_code=status_code, status_code=status_code,
@@ -217,7 +217,7 @@ class OpenAIResponseParser(ResponseParser):
return result 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 {} usage = response.get("usage") or {}
return { return {
"input_tokens": usage.get("prompt_tokens", 0), "input_tokens": usage.get("prompt_tokens", 0),
@@ -226,7 +226,7 @@ class OpenAIResponseParser(ResponseParser):
"cache_read_tokens": 0, "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", []) choices = response.get("choices", [])
if choices: if choices:
message = choices[0].get("message", {}) message = choices[0].get("message", {})
@@ -235,7 +235,7 @@ class OpenAIResponseParser(ResponseParser):
return content return content
return "" 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) is_error, _ = _check_nested_error(response)
return is_error return is_error
@@ -259,7 +259,7 @@ class ClaudeResponseParser(ResponseParser):
self.name = "CLAUDE" self.name = "CLAUDE"
self.api_format = "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(): if not line or not line.strip():
return None return None
@@ -324,7 +324,7 @@ class ClaudeResponseParser(ResponseParser):
return chunk 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( result = ParsedResponse(
raw_response=response, raw_response=response,
status_code=status_code, status_code=status_code,
@@ -358,7 +358,7 @@ class ClaudeResponseParser(ResponseParser):
return result 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 路径下 # 对于 message_start 事件usage 在 message.usage 路径下
# 对于其他响应usage 在顶层 # 对于其他响应usage 在顶层
usage = response.get("usage") or {} usage = response.get("usage") or {}
@@ -372,7 +372,7 @@ class ClaudeResponseParser(ResponseParser):
"cache_read_tokens": usage.get("cache_read_input_tokens", 0), "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", []) content = response.get("content", [])
if isinstance(content, list): if isinstance(content, list):
text_parts = [] text_parts = []
@@ -382,7 +382,7 @@ class ClaudeResponseParser(ResponseParser):
return "".join(text_parts) return "".join(text_parts)
return "" 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) is_error, _ = _check_nested_error(response)
return is_error return is_error
@@ -406,7 +406,7 @@ class GeminiResponseParser(ResponseParser):
self.name = "GEMINI" self.name = "GEMINI"
self.api_format = "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 行 解析 Gemini SSE 行
@@ -473,7 +473,7 @@ class GeminiResponseParser(ResponseParser):
return chunk 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( result = ParsedResponse(
raw_response=response, raw_response=response,
status_code=status_code, status_code=status_code,
@@ -509,7 +509,7 @@ class GeminiResponseParser(ResponseParser):
return result 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 使用量 从 Gemini 响应中提取 token 使用量
@@ -531,7 +531,7 @@ class GeminiResponseParser(ResponseParser):
"cache_read_tokens": usage.get("cached_tokens", 0), "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", []) candidates = response.get("candidates", [])
if candidates: if candidates:
content = candidates[0].get("content", {}) content = candidates[0].get("content", {})
@@ -543,7 +543,7 @@ class GeminiResponseParser(ResponseParser):
return "".join(text_parts) return "".join(text_parts)
return "" 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": ClaudeResponseParser,
"CLAUDE_CLI": ClaudeCliResponseParser, "CLAUDE_CLI": ClaudeCliResponseParser,
"OPENAI": OpenAIResponseParser, "OPENAI": OpenAIResponseParser,

View File

@@ -11,12 +11,11 @@
payload, headers = builder.build(original_body, original_headers, endpoint, key) payload, headers = builder.build(original_body, original_headers, endpoint, key)
""" """
from __future__ import annotations
import json import json
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from dataclasses import dataclass 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.crypto import crypto_service
from src.core.api_format import HeaderBuilder, UPSTREAM_DROP_HEADERS from src.core.api_format import HeaderBuilder, UPSTREAM_DROP_HEADERS
@@ -37,9 +36,9 @@ class ProviderAuthInfo:
auth_header: str auth_header: str
auth_value: str auth_value: str
# 解密后的认证配置(用于 URL 构建等场景,避免重复解密) # 解密后的认证配置(用于 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) 元组""" """返回 (auth_header, auth_value) 元组"""
return (self.auth_header, self.auth_value) return (self.auth_header, self.auth_value)
@@ -48,7 +47,7 @@ class ProviderAuthInfo:
# ============================================================================== # ==============================================================================
# 兼容别名:历史代码使用 SENSITIVE_HEADERS 命名 # 兼容别名:历史代码使用 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 格式) # 标准测试请求体OpenAI 格式)
# 用于 check_endpoint 等测试场景,使用简单安全的消息内容避免触发安全过滤 # 用于 check_endpoint 等测试场景,使用简单安全的消息内容避免触发安全过滤
DEFAULT_TEST_REQUEST: Dict[str, Any] = { DEFAULT_TEST_REQUEST: dict[str, Any] = {
"messages": [{"role": "user", "content": "Hi"}], "messages": [{"role": "user", "content": "Hi"}],
"max_tokens": 5, "max_tokens": 5,
"temperature": 0, "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则合并到默认测试请求中 如果传入 request_data则合并到默认测试请求中
@@ -85,8 +84,8 @@ def get_test_request_data(request_data: Optional[Dict[str, Any]] = None) -> Dict
def build_test_request_body( def build_test_request_body(
format_id: str, format_id: str,
request_data: Optional[Dict[str, Any]] = None, request_data: dict[str, Any] | None = None,
) -> Dict[str, Any]: ) -> dict[str, Any]:
"""构建测试请求体,自动处理格式转换 """构建测试请求体,自动处理格式转换
使用格式转换注册表将 OpenAI 格式的测试请求转换为目标格式。 使用格式转换注册表将 OpenAI 格式的测试请求转换为目标格式。
@@ -127,39 +126,39 @@ class RequestBuilder(ABC):
@abstractmethod @abstractmethod
def build_payload( def build_payload(
self, 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, is_stream: bool = False,
) -> Dict[str, Any]: ) -> dict[str, Any]:
"""构建请求体""" """构建请求体"""
pass pass
@abstractmethod @abstractmethod
def build_headers( def build_headers(
self, self,
original_headers: Dict[str, str], original_headers: dict[str, str],
endpoint: Any, endpoint: Any,
key: Any, key: Any,
*, *,
extra_headers: Optional[Dict[str, str]] = None, extra_headers: dict[str, str] | None = None,
pre_computed_auth: Optional[Tuple[str, str]] = None, pre_computed_auth: tuple[str, str] | None = None,
) -> Dict[str, str]: ) -> dict[str, str]:
"""构建请求头""" """构建请求头"""
pass pass
def build( def build(
self, self,
original_body: Dict[str, Any], original_body: dict[str, Any],
original_headers: Dict[str, str], original_headers: dict[str, str],
endpoint: Any, endpoint: Any,
key: Any, key: Any,
*, *,
mapped_model: Optional[str] = None, mapped_model: str | None = None,
is_stream: bool = False, is_stream: bool = False,
extra_headers: Optional[Dict[str, str]] = None, extra_headers: dict[str, str] | None = None,
pre_computed_auth: Optional[Tuple[str, str]] = None, pre_computed_auth: tuple[str, str] | None = None,
) -> Tuple[Dict[str, Any], Dict[str, str]]: ) -> tuple[dict[str, Any], dict[str, str]]:
""" """
构建完整的请求(请求体 + 请求头) 构建完整的请求(请求体 + 请求头)
@@ -202,11 +201,11 @@ class PassthroughRequestBuilder(RequestBuilder):
def build_payload( def build_payload(
self, 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 - 保留原始值,不自动添加 is_stream: bool = False, # noqa: ARG002 - 保留原始值,不自动添加
) -> Dict[str, Any]: ) -> dict[str, Any]:
""" """
透传请求体 - 原样复制,不做任何修改 透传请求体 - 原样复制,不做任何修改
@@ -218,13 +217,13 @@ class PassthroughRequestBuilder(RequestBuilder):
def build_headers( def build_headers(
self, self,
original_headers: Dict[str, str], original_headers: dict[str, str],
endpoint: Any, endpoint: Any,
key: Any, key: Any,
*, *,
extra_headers: Optional[Dict[str, str]] = None, extra_headers: dict[str, str] | None = None,
pre_computed_auth: Optional[Tuple[str, str]] = None, pre_computed_auth: tuple[str, str] | None = None,
) -> Dict[str, str]: ) -> dict[str, str]:
""" """
透传请求头 - 清理敏感头部(黑名单),透传其他所有头部 透传请求头 - 清理敏感头部(黑名单),透传其他所有头部
@@ -289,11 +288,11 @@ class PassthroughRequestBuilder(RequestBuilder):
def build_passthrough_request( def build_passthrough_request(
original_body: Dict[str, Any], original_body: dict[str, Any],
original_headers: Dict[str, str], original_headers: dict[str, str],
endpoint: Any, endpoint: Any,
key: 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 abc import ABC, abstractmethod
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import Any, Dict, Optional from typing import Any
@dataclass @dataclass
@@ -13,14 +13,14 @@ class ParsedChunk:
# 原始数据 # 原始数据
raw_line: str raw_line: str
event_type: Optional[str] = None event_type: str | None = None
data: Optional[Dict[str, Any]] = None data: dict[str, Any] | None = None
# 提取的内容 # 提取的内容
text_delta: str = "" text_delta: str = ""
is_done: bool = False is_done: bool = False
is_error: bool = False is_error: bool = False
error_message: Optional[str] = None error_message: str | None = None
# 使用量信息(通常在最后一个 chunk 中) # 使用量信息(通常在最后一个 chunk 中)
input_tokens: int = 0 input_tokens: int = 0
@@ -29,7 +29,7 @@ class ParsedChunk:
cache_read_tokens: int = 0 cache_read_tokens: int = 0
# 响应 ID # 响应 ID
response_id: Optional[str] = None response_id: str | None = None
@dataclass @dataclass
@@ -48,21 +48,21 @@ class StreamStats:
# 内容 # 内容
collected_text: str = "" collected_text: str = ""
response_id: Optional[str] = None response_id: str | None = None
# 状态 # 状态
has_completion: bool = False has_completion: bool = False
status_code: int = 200 status_code: int = 200
error_message: Optional[str] = None error_message: str | None = None
# Provider 信息 # Provider 信息
provider_name: Optional[str] = None provider_name: str | None = None
endpoint_id: Optional[str] = None endpoint_id: str | None = None
key_id: Optional[str] = None key_id: str | None = None
# 响应头和完整响应 # 响应头和完整响应
response_headers: Dict[str, str] = field(default_factory=dict) response_headers: dict[str, str] = field(default_factory=dict)
final_response: Optional[Dict[str, Any]] = None final_response: dict[str, Any] | None = None
@dataclass @dataclass
@@ -70,12 +70,12 @@ class ParsedResponse:
"""解析后的非流式响应""" """解析后的非流式响应"""
# 原始响应 # 原始响应
raw_response: Dict[str, Any] raw_response: dict[str, Any]
status_code: int status_code: int
# 提取的内容 # 提取的内容
text_content: str = "" text_content: str = ""
response_id: Optional[str] = None response_id: str | None = None
# 使用量 # 使用量
input_tokens: int = 0 input_tokens: int = 0
@@ -85,10 +85,10 @@ class ParsedResponse:
# 错误信息 # 错误信息
is_error: bool = False is_error: bool = False
error_type: Optional[str] = None error_type: str | None = None
error_message: Optional[str] = None error_message: str | None = None
# 从响应体解析出的嵌套状态码(当 HTTP 200 但响应体含错误时使用) # 从响应体解析出的嵌套状态码(当 HTTP 200 但响应体含错误时使用)
embedded_status_code: Optional[int] = None embedded_status_code: int | None = None
class ResponseParser(ABC): class ResponseParser(ABC):
@@ -106,7 +106,7 @@ class ResponseParser(ABC):
api_format: str = "UNKNOWN" api_format: str = "UNKNOWN"
@abstractmethod @abstractmethod
def parse_sse_line(self, line: str, stats: StreamStats) -> Optional[ParsedChunk]: def parse_sse_line(self, line: str, stats: StreamStats) -> ParsedChunk | None:
""" """
解析单行 SSE 数据 解析单行 SSE 数据
@@ -120,7 +120,7 @@ class ResponseParser(ABC):
pass pass
@abstractmethod @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 pass
@abstractmethod @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 使用量 从响应中提取 token 使用量
@@ -147,7 +147,7 @@ class ResponseParser(ABC):
pass pass
@abstractmethod @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 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 import time
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any, Dict, List, Optional from typing import TYPE_CHECKING, Any
if TYPE_CHECKING: if TYPE_CHECKING:
from src.core.api_format.conversion.stream_state import StreamState from src.core.api_format.conversion.stream_state import StreamState
@@ -35,16 +37,16 @@ class StreamContext:
api_key_id: int = 0 api_key_id: int = 0
# Provider 信息(在请求执行时填充) # Provider 信息(在请求执行时填充)
provider_name: Optional[str] = None provider_name: str | None = None
provider_id: Optional[str] = None provider_id: str | None = None
endpoint_id: Optional[str] = None endpoint_id: str | None = None
key_id: Optional[str] = None key_id: str | None = None
attempt_id: Optional[str] = None attempt_id: str | None = None
attempt_synced: bool = False 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 统计 # Token 统计
input_tokens: int = 0 input_tokens: int = 0
@@ -53,33 +55,33 @@ class StreamContext:
cache_creation_tokens: int = 0 cache_creation_tokens: int = 0
# 响应内容 # 响应内容
_collected_text_parts: List[str] = field(default_factory=list, repr=False) _collected_text_parts: list[str] = field(default_factory=list, repr=False)
response_id: Optional[str] = None response_id: str | None = None
final_usage: Optional[Dict[str, Any]] = None final_usage: dict[str, Any] | None = None
final_response: Optional[Dict[str, Any]] = 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) start_time: float = field(default_factory=time.time)
# 响应状态 # 响应状态
status_code: int = 200 status_code: int = 200
error_message: Optional[str] = None # 客户端友好的错误消息 error_message: str | None = None # 客户端友好的错误消息
upstream_response: Optional[str] = None # 原始 Provider 响应(用于请求链路追踪) upstream_response: str | None = None # 原始 Provider 响应(用于请求链路追踪)
has_completion: bool = False has_completion: bool = False
# 请求/响应数据 # 请求/响应数据
response_headers: Dict[str, str] = field(default_factory=dict) # 提供商响应头 response_headers: dict[str, str] = field(default_factory=dict) # 提供商响应头
client_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_headers: dict[str, str] = field(default_factory=dict)
provider_request_body: Optional[Dict[str, Any]] = None provider_request_body: dict[str, Any] | None = None
# 格式转换信息CLI handler 需要) # 格式转换信息CLI handler 需要)
client_api_format: str = "" client_api_format: str = ""
needs_conversion: bool = False # 是否需要跨格式转换(由 handler 层设置) needs_conversion: bool = False # 是否需要跨格式转换(由 handler 层设置)
# Provider 响应元数据CLI handler 需要) # Provider 响应元数据CLI handler 需要)
response_metadata: Dict[str, Any] = field(default_factory=dict) response_metadata: dict[str, Any] = field(default_factory=dict)
# 整流标记Thinking Rectifier # 整流标记Thinking Rectifier
rectified: bool = False # 请求是否经过整流(移除 thinking 块后重试) rectified: bool = False # 请求是否经过整流(移除 thinking 块后重试)
@@ -87,10 +89,10 @@ class StreamContext:
# 流式处理统计 # 流式处理统计
data_count: int = 0 data_count: int = 0
chunk_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 追踪) # 流式格式转换状态(跨 chunk 追踪)
stream_conversion_state: Optional["StreamState"] = None stream_conversion_state: StreamState | None = None
def reset_for_retry(self) -> None: def reset_for_retry(self) -> None:
""" """
@@ -138,7 +140,7 @@ class StreamContext:
provider_id: str, provider_id: str,
endpoint_id: str, endpoint_id: str,
key_id: str, key_id: str,
provider_api_format: Optional[str] = None, provider_api_format: str | None = None,
) -> None: ) -> None:
"""更新 Provider 信息""" """更新 Provider 信息"""
self.provider_name = provider_name self.provider_name = provider_name
@@ -149,10 +151,10 @@ class StreamContext:
def update_usage( def update_usage(
self, self,
input_tokens: Optional[int] = None, input_tokens: int | None = None,
output_tokens: Optional[int] = None, output_tokens: int | None = None,
cached_tokens: Optional[int] = None, cached_tokens: int | None = None,
cache_creation_tokens: Optional[int] = None, cache_creation_tokens: int | None = None,
) -> None: ) -> None:
""" """
更新 Token 使用统计 更新 Token 使用统计
@@ -194,7 +196,7 @@ class StreamContext:
self, self,
status_code: int, status_code: int,
error_message: str, error_message: str,
upstream_response: Optional[str] = None, upstream_response: str | None = None,
) -> None: ) -> None:
""" """
标记请求失败 标记请求失败
@@ -230,7 +232,7 @@ class StreamContext:
"""检查是否因客户端断开连接而结束""" """检查是否因客户端断开连接而结束"""
return self.status_code == 499 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 codecs
import json import json
from dataclasses import dataclass 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 import httpx
@@ -65,10 +68,10 @@ class StreamProcessor:
self, self,
request_id: str, request_id: str,
default_parser: ResponseParser, default_parser: ResponseParser,
on_streaming_start: Optional[Callable[[], None]] = None, on_streaming_start: Callable[[], None] | None = None,
*, *,
collect_text: bool = False, collect_text: bool = False,
smoothing_config: Optional[StreamSmoothingConfig] = None, smoothing_config: StreamSmoothingConfig | None = None,
): ):
""" """
初始化流处理器 初始化流处理器
@@ -105,7 +108,7 @@ class StreamProcessor:
def handle_sse_event( def handle_sse_event(
self, self,
ctx: StreamContext, ctx: StreamContext,
event_name: Optional[str], event_name: str | None,
data_str: str, data_str: str,
*, *,
skip_record: bool = False, skip_record: bool = False,
@@ -363,7 +366,7 @@ class StreamProcessor:
): ):
# 重新抛出可重试的 Provider 异常,触发故障转移 # 重新抛出可重试的 Provider 异常,触发故障转移
raise raise
except (OSError, IOError) as e: except OSError as e:
# 网络 I/O 异常:记录警告,可能需要重试 # 网络 I/O 异常:记录警告,可能需要重试
logger.warning(f" [{self.request_id}] 预读流时发生网络异常: {type(e).__name__}: {e}") logger.warning(f" [{self.request_id}] 预读流时发生网络异常: {type(e).__name__}: {e}")
except Exception as e: except Exception as e:
@@ -382,10 +385,10 @@ class StreamProcessor:
byte_iterator: Any, byte_iterator: Any,
response_ctx: Any, response_ctx: Any,
http_client: httpx.AsyncClient, http_client: httpx.AsyncClient,
prefetched_chunks: Optional[list] = None, prefetched_chunks: list | None = None,
*, *,
start_time: Optional[float] = None, start_time: float | None = None,
) -> AsyncGenerator[bytes, None]: ) -> AsyncGenerator[bytes]:
""" """
创建响应流生成器 创建响应流生成器
@@ -547,9 +550,7 @@ class StreamProcessor:
logger.warning(f"[{self.request_id}] 流式格式转换失败: {conv_err}") logger.warning(f"[{self.request_id}] 流式格式转换失败: {conv_err}")
payload = _build_stream_error_payload("响应格式转换失败,请稍后重试") payload = _build_stream_error_payload("响应格式转换失败,请稍后重试")
error_bytes = ( error_bytes = (
f"data: {json.dumps(payload, ensure_ascii=False)}\n\n".encode( f"data: {json.dumps(payload, ensure_ascii=False)}\n\n".encode()
"utf-8"
)
) )
done_bytes = ( done_bytes = (
b"data: [DONE]\n\n" if client_format.startswith("OPENAI") else b"" b"data: [DONE]\n\n" if client_format.startswith("OPENAI") else b""
@@ -570,7 +571,7 @@ class StreamProcessor:
# 统一使用 SSE 格式输出Gemini streamGenerateContent 也使用 SSE # 统一使用 SSE 格式输出Gemini streamGenerateContent 也使用 SSE
# 参考: https://ai.google.dev/api/generate-content # 参考: https://ai.google.dev/api/generate-content
out.append( 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 return out
@@ -769,9 +770,9 @@ class StreamProcessor:
async def create_monitored_stream( async def create_monitored_stream(
self, self,
ctx: StreamContext, ctx: StreamContext,
stream_generator: AsyncGenerator[bytes, None], stream_generator: AsyncGenerator[bytes],
is_disconnected: Callable[[], Any], is_disconnected: Callable[[], Any],
) -> AsyncGenerator[bytes, None]: ) -> AsyncGenerator[bytes]:
""" """
创建带监控的流生成器 创建带监控的流生成器
@@ -833,8 +834,8 @@ class StreamProcessor:
async def create_smoothed_stream( async def create_smoothed_stream(
self, self,
stream_generator: AsyncGenerator[bytes, None], stream_generator: AsyncGenerator[bytes],
) -> AsyncGenerator[bytes, None]: ) -> AsyncGenerator[bytes]:
""" """
创建平滑输出的流生成器 创建平滑输出的流生成器
@@ -933,7 +934,7 @@ class StreamProcessor:
if buffer: if buffer:
yield 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: if format_name not in self._extractors:
extractor = get_extractor(format_name) extractor = get_extractor(format_name)
@@ -943,7 +944,7 @@ class StreamProcessor:
def _detect_format_and_extract( def _detect_format_and_extract(
self, data: dict self, data: dict
) -> tuple[Optional[str], Optional[ContentExtractor]]: ) -> tuple[str | None, ContentExtractor | None]:
""" """
检测数据格式并提取内容 检测数据格式并提取内容
@@ -998,10 +999,10 @@ class StreamProcessor:
async def create_smoothed_stream( async def create_smoothed_stream(
stream_generator: AsyncGenerator[bytes, None], stream_generator: AsyncGenerator[bytes],
chunk_size: int = 20, chunk_size: int = 20,
delay_ms: int = 8, delay_ms: int = 8,
) -> AsyncGenerator[bytes, None]: ) -> AsyncGenerator[bytes]:
""" """
独立的平滑流生成函数 独立的平滑流生成函数
@@ -1032,7 +1033,7 @@ class _LightweightSmoother:
self.delay_ms = delay_ms self.delay_ms = delay_ms
self._extractors: dict[str, ContentExtractor] = {} 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: if format_name not in self._extractors:
extractor = get_extractor(format_name) extractor = get_extractor(format_name)
if extractor: if extractor:
@@ -1041,7 +1042,7 @@ class _LightweightSmoother:
def _detect_format_and_extract( def _detect_format_and_extract(
self, data: dict self, data: dict
) -> tuple[Optional[str], Optional[ContentExtractor]]: ) -> tuple[str | None, ContentExtractor | None]:
for format_name in get_extractor_formats(): for format_name in get_extractor_formats():
extractor = self._get_extractor(format_name) extractor = self._get_extractor(format_name)
if extractor: 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)] return [content[i : i + self.chunk_size] for i in range(0, text_length, self.chunk_size)]
async def smooth( async def smooth(
self, stream_generator: AsyncGenerator[bytes, None] self, stream_generator: AsyncGenerator[bytes]
) -> AsyncGenerator[bytes, None]: ) -> AsyncGenerator[bytes]:
buffer = b"" buffer = b""
is_first_content = True is_first_content = True

View File

@@ -9,7 +9,7 @@
import asyncio import asyncio
import time import time
from typing import Any, Dict, Optional from typing import Any
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
@@ -58,8 +58,8 @@ class StreamTelemetryRecorder:
async def record_stream_stats( async def record_stream_stats(
self, self,
ctx: StreamContext, ctx: StreamContext,
original_headers: Dict[str, str], original_headers: dict[str, str],
original_request_body: Dict[str, Any], original_request_body: dict[str, Any],
start_time: float, start_time: float,
) -> None: ) -> None:
""" """
@@ -144,9 +144,9 @@ class StreamTelemetryRecorder:
self, self,
writer: TelemetryWriter, writer: TelemetryWriter,
ctx: StreamContext, ctx: StreamContext,
original_headers: Dict[str, str], original_headers: dict[str, str],
actual_request_body: Dict[str, Any], actual_request_body: dict[str, Any],
response_body: Optional[Dict[str, Any]], response_body: dict[str, Any] | None,
response_time_ms: int, response_time_ms: int,
) -> None: ) -> None:
"""记录成功的请求""" """记录成功的请求"""
@@ -193,9 +193,9 @@ class StreamTelemetryRecorder:
self, self,
writer: TelemetryWriter, writer: TelemetryWriter,
ctx: StreamContext, ctx: StreamContext,
original_headers: Dict[str, str], original_headers: dict[str, str],
actual_request_body: Dict[str, Any], actual_request_body: dict[str, Any],
response_body: Optional[Dict[str, Any]], response_body: dict[str, Any] | None,
response_time_ms: int, response_time_ms: int,
) -> None: ) -> None:
"""记录失败的请求""" """记录失败的请求"""
@@ -236,9 +236,9 @@ class StreamTelemetryRecorder:
self, self,
writer: TelemetryWriter, writer: TelemetryWriter,
ctx: StreamContext, ctx: StreamContext,
original_headers: Dict[str, str], original_headers: dict[str, str],
actual_request_body: Dict[str, Any], actual_request_body: dict[str, Any],
response_body: Optional[Dict[str, Any]], response_body: dict[str, Any] | None,
response_time_ms: int, response_time_ms: int,
) -> None: ) -> None:
"""记录客户端取消的请求""" """记录客户端取消的请求"""
@@ -285,7 +285,7 @@ class StreamTelemetryRecorder:
from src.services.request.candidate import RequestCandidateService from src.services.request.candidate import RequestCandidateService
extra_data: Dict[str, Any] = { extra_data: dict[str, Any] = {
"stream_completed": ctx.is_success(), "stream_completed": ctx.is_success(),
"data_count": ctx.data_count, "data_count": ctx.data_count,
} }
@@ -358,7 +358,7 @@ class StreamTelemetryRecorder:
status: str, status: str,
response_time_ms: int, response_time_ms: int,
status_code: int = 200, status_code: int = 200,
error_message: Optional[str] = None, error_message: str | None = None,
) -> None: ) -> None:
"""直接更新 Usage 表的状态字段""" """直接更新 Usage 表的状态字段"""
try: try:
@@ -378,7 +378,7 @@ class StreamTelemetryRecorder:
async def _get_telemetry_writer( async def _get_telemetry_writer(
self, bg_db: Session, ctx: StreamContext, response_time_ms: int 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: if config.usage_queue_enabled and self.user_id and self.api_key_id:
return QueueTelemetryWriter( return QueueTelemetryWriter(
request_id=self.request_id, request_id=self.request_id,
@@ -400,9 +400,9 @@ class StreamTelemetryRecorder:
self, self,
writer: TelemetryWriter, writer: TelemetryWriter,
ctx: StreamContext, ctx: StreamContext,
original_headers: Dict[str, str], original_headers: dict[str, str],
actual_request_body: Dict[str, Any], actual_request_body: dict[str, Any],
response_body: Optional[Dict[str, Any]], response_body: dict[str, Any] | None,
response_time_ms: int, response_time_ms: int,
) -> None: ) -> None:
"""根据上下文状态分发到对应的记录方法""" """根据上下文状态分发到对应的记录方法"""
@@ -430,7 +430,7 @@ class StreamTelemetryRecorder:
return "cancelled" return "cancelled"
return "failed" 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() 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() 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 from __future__ import annotations
import json 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.exceptions import EmbeddedErrorException, ProviderNotAvailableException
from src.core.api_format import filter_response_headers 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 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 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兼容三种格式 提取缓存创建 tokens兼容三种格式
@@ -99,7 +100,7 @@ def extract_cache_creation_tokens(usage: Dict[str, Any]) -> int:
return old_format 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推荐响应头用于减少代理缓冲带来的卡顿/成段输出。 构建 SSEtext/event-stream推荐响应头用于减少代理缓冲带来的卡顿/成段输出。
@@ -107,7 +108,7 @@ def build_sse_headers(extra_headers: Optional[Dict[str, str]] = None) -> Dict[st
- Cache-Control: no-transform 可避免部分代理对流做压缩/改写导致缓冲 - Cache-Control: no-transform 可避免部分代理对流做压缩/改写导致缓冲
- X-Accel-Buffering: no 可显式提示 Nginx 关闭缓冲(即使全局已关闭也无害) - X-Accel-Buffering: no 可显式提示 Nginx 关闭缓冲(即使全局已关闭也无害)
""" """
headers: Dict[str, str] = { headers: dict[str, str] = {
"Cache-Control": "no-cache, no-transform", "Cache-Control": "no-cache, no-transform",
"X-Accel-Buffering": "no", "X-Accel-Buffering": "no",
} }
@@ -116,7 +117,7 @@ def build_sse_headers(extra_headers: Optional[Dict[str, str]] = None) -> Dict[st
return headers 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, parser: Any,
request_id: str, request_id: str,
provider_name: str, provider_name: str,
endpoint_id: Optional[str], endpoint_id: str | None,
base_url: Optional[str], base_url: str | None,
) -> None: ) -> None:
""" """
检查预读的响应是否为非 SSE 格式的错误响应HTML 或纯 JSON 错误) 检查预读的响应是否为非 SSE 格式的错误响应HTML 或纯 JSON 错误)

View File

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

View File

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

View File

@@ -4,10 +4,9 @@ Claude SSE 流解析器
解析 Claude Messages API 的 Server-Sent Events 流。 解析 Claude Messages API 的 Server-Sent Events 流。
""" """
from __future__ import annotations
import json import json
from typing import Any, Dict, List, Optional from typing import Any
from src.api.handlers.base.utils import extract_cache_creation_tokens from src.api.handlers.base.utils import extract_cache_creation_tokens
@@ -43,7 +42,7 @@ class ClaudeStreamParser:
DELTA_TEXT = "text_delta" DELTA_TEXT = "text_delta"
DELTA_INPUT_JSON = "input_json_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 数据块 解析 SSE 数据块
@@ -58,10 +57,10 @@ class ClaudeStreamParser:
else: else:
text = chunk text = chunk
events: List[Dict[str, Any]] = [] events: list[dict[str, Any]] = []
lines = text.strip().split("\n") lines = text.strip().split("\n")
current_event_type: Optional[str] = None current_event_type: str | None = None
for line in lines: for line in lines:
line = line.strip() line = line.strip()
@@ -96,7 +95,7 @@ class ClaudeStreamParser:
return events return events
def parse_line(self, line: str) -> Optional[Dict[str, Any]]: def parse_line(self, line: str) -> dict[str, Any] | None:
""" """
解析单行 SSE 数据 解析单行 SSE 数据
@@ -117,7 +116,7 @@ class ClaudeStreamParser:
except json.JSONDecodeError: except json.JSONDecodeError:
return None 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") event_type = event.get("type")
return event_type in (self.EVENT_MESSAGE_STOP, "__done__") 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 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") event_type = event.get("type")
return str(event_type) if event_type is not None else None 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 事件中提取文本增量 从 content_block_delta 事件中提取文本增量
@@ -175,7 +174,7 @@ class ClaudeStreamParser:
return None 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 使用量 从事件中提取 token 使用量
@@ -212,7 +211,7 @@ class ClaudeStreamParser:
return None 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 从 message_start 事件中提取消息 ID
@@ -229,7 +228,7 @@ class ClaudeStreamParser:
msg_id = message.get("id") msg_id = message.get("id")
return str(msg_id) if msg_id is not None else None 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 事件中提取停止原因 从 message_delta 事件中提取停止原因

View File

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

View File

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

View File

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

View File

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

View File

@@ -15,7 +15,7 @@ Gemini streamGenerateContent 常见两种返回(与上游/代理实现有关
""" """
import json import json
from typing import Any, Dict, List, Optional, Union from typing import Any
class GeminiStreamParser: class GeminiStreamParser:
@@ -43,7 +43,7 @@ class GeminiStreamParser:
self._in_array = False self._in_array = False
self._brace_depth = 0 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: else:
text = chunk text = chunk
events: List[Dict[str, Any]] = [] events: list[dict[str, Any]] = []
for char in text: for char in text:
if char == "[" and not self._in_array: if char == "[" and not self._in_array:
@@ -97,7 +97,7 @@ class GeminiStreamParser:
return events return events
def parse_line(self, line: str) -> Optional[Dict[str, Any]]: def parse_line(self, line: str) -> dict[str, Any] | None:
""" """
解析单行 JSON 数据 解析单行 JSON 数据
@@ -118,7 +118,7 @@ class GeminiStreamParser:
except json.JSONDecodeError: except json.JSONDecodeError:
return None 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 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 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 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 str(reason) if reason is not None else None
return 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 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 使用量 从事件中提取 token 使用量
@@ -280,7 +280,7 @@ class GeminiStreamParser:
"cached_tokens": usage_metadata.get("cachedContentTokenCount", 0), "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") version = event.get("modelVersion")
return str(version) if version is not None else None 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 格式的请求。 继承 CliAdapterBase处理 Gemini CLI 格式的请求。
""" """
from typing import Any, Dict, Optional, Tuple, Type from typing import Any
import httpx import httpx
from fastapi import Request from fastapi import Request
@@ -29,16 +29,16 @@ class GeminiCliAdapter(CliAdapterBase):
name = "gemini.cli" name = "gemini.cli"
@property @property
def HANDLER_CLASS(self) -> Type[CliMessageHandlerBase]: def HANDLER_CLASS(self) -> type[CliMessageHandlerBase]:
"""延迟导入 Handler 类避免循环依赖""" """延迟导入 Handler 类避免循环依赖"""
from src.api.handlers.gemini_cli.handler import GeminiCliMessageHandler from src.api.handlers.gemini_cli.handler import GeminiCliMessageHandler
return 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"]) 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 两种方式 从请求中提取 API 密钥 - Gemini CLI 支持 header 和 query 两种方式
@@ -53,8 +53,8 @@ class GeminiCliAdapter(CliAdapterBase):
) )
def _merge_path_params( def _merge_path_params(
self, original_request_body: Dict[str, Any], path_params: Dict[str, Any] # noqa: ARG002 self, original_request_body: dict[str, Any], path_params: dict[str, Any] # noqa: ARG002
) -> Dict[str, Any]: ) -> dict[str, Any]:
""" """
合并 URL 路径参数到请求体 - Gemini CLI 特化版本 合并 URL 路径参数到请求体 - Gemini CLI 特化版本
@@ -74,23 +74,23 @@ class GeminiCliAdapter(CliAdapterBase):
# Gemini: 不合并任何 path_params 到请求体 # Gemini: 不合并任何 path_params 到请求体
return original_request_body.copy() 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 字段""" """Gemini CLI 使用 contents 字段"""
contents = payload.get("contents", []) contents = payload.get("contents", [])
return len(contents) if isinstance(contents, list) else 0 return len(contents) if isinstance(contents, list) else 0
def _build_audit_metadata( def _build_audit_metadata(
self, self,
payload: Dict[str, Any], payload: dict[str, Any],
path_params: Optional[Dict[str, Any]] = None, path_params: dict[str, Any] | None = None,
) -> Dict[str, Any]: ) -> dict[str, Any]:
"""Gemini CLI 特定的审计元数据""" """Gemini CLI 特定的审计元数据"""
# 从 path_params 获取 modelGemini 请求体不含 model # 从 path_params 获取 modelGemini 请求体不含 model
model = path_params.get("model", "unknown") if path_params else "unknown" model = path_params.get("model", "unknown") if path_params else "unknown"
contents = payload.get("contents", []) contents = payload.get("contents", [])
generation_config = payload.get("generation_config", {}) or {} generation_config = payload.get("generation_config", {}) or {}
role_counts: Dict[str, int] = {} role_counts: dict[str, int] = {}
for content in contents: for content in contents:
role = content.get("role", "unknown") if isinstance(content, dict) else "unknown" role = content.get("role", "unknown") if isinstance(content, dict) else "unknown"
role_counts[role] = role_counts.get(role, 0) + 1 role_counts[role] = role_counts.get(role, 0) + 1
@@ -120,8 +120,8 @@ class GeminiCliAdapter(CliAdapterBase):
client: httpx.AsyncClient, client: httpx.AsyncClient,
base_url: str, base_url: str,
api_key: str, api_key: str,
extra_headers: Optional[Dict[str, str]] = None, extra_headers: dict[str, str] | None = None,
) -> Tuple[list, Optional[str]]: ) -> tuple[list, str | None]:
"""查询 Gemini API 支持的模型列表(带 CLI User-Agent""" """查询 Gemini API 支持的模型列表(带 CLI User-Agent"""
# 复用 GeminiChatAdapter 的实现,添加 CLI User-Agent # 复用 GeminiChatAdapter 的实现,添加 CLI User-Agent
cli_headers = {"User-Agent": config.internal_user_agent_gemini_cli} cli_headers = {"User-Agent": config.internal_user_agent_gemini_cli}
@@ -136,7 +136,7 @@ class GeminiCliAdapter(CliAdapterBase):
return models, error return models, error
@classmethod @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""" """构建Gemini CLI API端点URL"""
effective_model_name = model_name or request_data.get("model", "") effective_model_name = model_name or request_data.get("model", "")
if not effective_model_name: if not effective_model_name:
@@ -152,12 +152,12 @@ class GeminiCliAdapter(CliAdapterBase):
# build_request_body 使用基类实现,通过 format_conversion_registry 自动转换 OPENAI -> GEMINI_CLI # build_request_body 使用基类实现,通过 format_conversion_registry 自动转换 OPENAI -> GEMINI_CLI
@classmethod @classmethod
def get_cli_user_agent(cls) -> Optional[str]: def get_cli_user_agent(cls) -> str | None:
"""获取Gemini CLI User-Agent""" """获取Gemini CLI User-Agent"""
return config.internal_user_agent_gemini_cli return config.internal_user_agent_gemini_cli
@classmethod @classmethod
def get_cli_extra_headers(cls) -> Dict[str, str]: def get_cli_extra_headers(cls) -> dict[str, str]:
"""获取Gemini CLI额外请求头包含 x-app: cli 标识""" """获取Gemini CLI额外请求头包含 x-app: cli 标识"""
headers = super().get_cli_extra_headers() headers = super().get_cli_extra_headers()
headers["x-app"] = "cli" # 标识 CLI 模式,让上游使用正确的 adapter headers["x-app"] = "cli" # 标识 CLI 模式,让上游使用正确的 adapter

View File

@@ -4,7 +4,7 @@ Gemini CLI Message Handler - 基于通用 CLI Handler 基类的实现
继承 CliMessageHandlerBase处理 Gemini CLI API 格式的请求。 继承 CliMessageHandlerBase处理 Gemini CLI API 格式的请求。
""" """
from typing import Any, Dict, Optional from typing import Any
from src.api.handlers.base.cli_handler_base import ( from src.api.handlers.base.cli_handler_base import (
CliMessageHandlerBase, CliMessageHandlerBase,
@@ -34,8 +34,8 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
def extract_model_from_request( def extract_model_from_request(
self, self,
request_body: Dict[str, Any], # noqa: ARG002 - 基类签名要求 request_body: dict[str, Any], # noqa: ARG002 - 基类签名要求
path_params: Optional[Dict[str, Any]] = None, path_params: dict[str, Any] | None = None,
) -> str: ) -> str:
""" """
从请求中提取模型名 - Gemini 格式实现 从请求中提取模型名 - Gemini 格式实现
@@ -57,8 +57,8 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
def prepare_provider_request_body( def prepare_provider_request_body(
self, self,
request_body: Dict[str, Any], request_body: dict[str, Any],
) -> Dict[str, Any]: ) -> dict[str, Any]:
""" """
准备发送给 Gemini API 的请求体 - 移除 model 字段 准备发送给 Gemini API 的请求体 - 移除 model 字段
@@ -77,9 +77,9 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
def get_model_for_url( def get_model_for_url(
self, self,
request_body: Dict[str, Any], request_body: dict[str, Any],
mapped_model: Optional[str], mapped_model: str | None,
) -> Optional[str]: ) -> str | None:
""" """
Gemini 需要将 model 放入 URL 路径中 Gemini 需要将 model 放入 URL 路径中
@@ -93,7 +93,7 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
# 优先使用映射后的模型名,否则使用请求体中的 # 优先使用映射后的模型名,否则使用请求体中的
return mapped_model or request_body.get("model") 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 使用情况 从 Gemini 事件中提取 token 使用情况
@@ -126,7 +126,7 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
self, self,
ctx: StreamContext, ctx: StreamContext,
_event_type: str, _event_type: str,
data: Dict[str, Any], data: dict[str, Any],
) -> None: ) -> None:
""" """
处理 Gemini CLI 格式的流式事件 处理 Gemini CLI 格式的流式事件
@@ -190,8 +190,8 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
def _extract_response_metadata( def _extract_response_metadata(
self, self,
response: Dict[str, Any], response: dict[str, Any],
) -> Dict[str, Any]: ) -> dict[str, Any]:
""" """
从 Gemini 响应中提取元数据 从 Gemini 响应中提取元数据
@@ -203,7 +203,7 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
Returns: Returns:
包含 model_version 的元数据字典 包含 model_version 的元数据字典
""" """
metadata: Dict[str, Any] = {} metadata: dict[str, Any] = {}
model_version = response.get("modelVersion") model_version = response.get("modelVersion")
if model_version: if model_version:
metadata["model_version"] = 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 格式请求。 处理 /v1/chat/completions 端点的 OpenAI Chat 格式请求。
""" """
from typing import Any, Dict, Optional, Tuple, Type from typing import Any
import httpx import httpx
from fastapi.responses import JSONResponse from fastapi.responses import JSONResponse
@@ -28,13 +28,13 @@ class OpenAIChatAdapter(ChatAdapterBase):
name = "openai.chat" name = "openai.chat"
@property @property
def HANDLER_CLASS(self) -> Type[ChatHandlerBase]: def HANDLER_CLASS(self) -> type[ChatHandlerBase]:
"""延迟导入 Handler 类避免循环依赖""" """延迟导入 Handler 类避免循环依赖"""
from src.api.handlers.openai.handler import OpenAIChatHandler from src.api.handlers.openai.handler import OpenAIChatHandler
return 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"]) super().__init__(allowed_api_formats or ["OPENAI"])
def _validate_request_body(self, original_request_body: dict, path_params: dict = None): 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"), 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 特定的审计元数据""" """构建 OpenAI Chat 特定的审计元数据"""
role_counts = {} role_counts = {}
for message in request_obj.messages: for message in request_obj.messages:
@@ -105,8 +105,8 @@ class OpenAIChatAdapter(ChatAdapterBase):
client: httpx.AsyncClient, client: httpx.AsyncClient,
base_url: str, base_url: str,
api_key: str, api_key: str,
extra_headers: Optional[Dict[str, str]] = None, extra_headers: dict[str, str] | None = None,
) -> Tuple[list, Optional[str]]: ) -> tuple[list, str | None]:
"""查询 OpenAI 兼容 API 支持的模型列表""" """查询 OpenAI 兼容 API 支持的模型列表"""
headers = cls.build_headers_with_extra(api_key, extra_headers) headers = cls.build_headers_with_extra(api_key, extra_headers)

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

@@ -1,12 +1,11 @@
"""公开模块状态 API供登录页等使用""" """公开模块状态 API供登录页等使用"""
from typing import List
from fastapi import APIRouter, Depends from fastapi import APIRouter, Depends
from pydantic import BaseModel from pydantic import BaseModel
from sqlalchemy.orm import Session 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 from src.database import get_db
router = APIRouter(prefix="/api/modules", tags=["Modules"]) router = APIRouter(prefix="/api/modules", tags=["Modules"])
@@ -20,7 +19,7 @@ class AuthModuleInfo(BaseModel):
active: bool 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)): 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 datetime import datetime, timezone
from typing import Any, Dict, Optional from typing import Any
import httpx import httpx
from fastapi import APIRouter, Depends, HTTPException, Query, Request 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: if value is None:
return default return default
@@ -39,9 +39,9 @@ def _serialize_provider(
provider: Provider, provider: Provider,
include_models: bool, include_models: bool,
include_endpoints: bool, include_endpoints: bool,
) -> Dict[str, Any]: ) -> dict[str, Any]:
"""序列化 Provider 对象""" """序列化 Provider 对象"""
provider_data: Dict[str, Any] = { provider_data: dict[str, Any] = {
"id": provider.id, "id": provider.id,
"name": provider.name, "name": provider.name,
"is_active": provider.is_active, "is_active": provider.is_active,
@@ -81,7 +81,7 @@ def _serialize_provider(
return provider_data 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 优先级选择)""" """选择 Provider按 provider_priority 优先级选择)"""
query = db.query(Provider).filter(Provider.is_active == True) query = db.query(Provider).filter(Provider.is_active == True)
if provider_name: 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 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: try:
redis = await get_redis_client() redis = await get_redis_client()
if redis: if redis:
@@ -245,9 +245,9 @@ async def provider_detail(
async def test_connection( async def test_connection(
request: Request, request: Request,
db: Session = Depends(get_db), db: Session = Depends(get_db),
provider: Optional[str] = Query(None), provider: str | None = Query(None),
model: str = Query("claude-3-haiku-20240307"), model: str = Query("claude-3-haiku-20240307"),
api_format: Optional[str] = Query(None), api_format: str | None = Query(None),
): ):
"""测试 Provider 连接""" """测试 Provider 连接"""
selected_provider = _select_provider(db, provider) selected_provider = _select_provider(db, provider)

View File

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

View File

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

View File

@@ -8,11 +8,13 @@
3. 连接池复用Keep-alive 连接减少 TCP 握手开销 3. 连接池复用Keep-alive 连接减少 TCP 握手开销
""" """
from __future__ import annotations
import asyncio import asyncio
import hashlib import hashlib
import time import time
from contextlib import asynccontextmanager from contextlib import asynccontextmanager
from typing import Any, Dict, Optional, Tuple from typing import Any
from urllib.parse import quote, urlparse from urllib.parse import quote, urlparse
import httpx import httpx
@@ -26,7 +28,7 @@ _proxy_clients_lock = asyncio.Lock()
_default_client_lock = asyncio.Lock() _default_client_lock = asyncio.Lock()
def _compute_proxy_cache_key(proxy_config: Optional[Dict[str, Any]]) -> str: def _compute_proxy_cache_key(proxy_config: dict[str, Any] | None) -> str:
""" """
计算代理配置的缓存键 计算代理配置的缓存键
@@ -48,7 +50,7 @@ def _compute_proxy_cache_key(proxy_config: Optional[Dict[str, Any]]) -> str:
return f"proxy:{hashlib.md5(proxy_url.encode()).hexdigest()[:16]}" return f"proxy:{hashlib.md5(proxy_url.encode()).hexdigest()[:16]}"
def build_proxy_url(proxy_config: Dict[str, Any]) -> Optional[str]: def build_proxy_url(proxy_config: dict[str, Any]) -> str | None:
""" """
根据代理配置构建完整的代理 URL 根据代理配置构建完整的代理 URL
@@ -103,11 +105,11 @@ class HTTPClientPool:
3. LRU 淘汰:代理客户端超过上限时淘汰最久未使用的 3. LRU 淘汰:代理客户端超过上限时淘汰最久未使用的
""" """
_instance: Optional["HTTPClientPool"] = None _instance: HTTPClientPool | None = None
_default_client: Optional[httpx.AsyncClient] = None _default_client: httpx.AsyncClient | None = None
_clients: Dict[str, httpx.AsyncClient] = {} _clients: dict[str, httpx.AsyncClient] = {}
# 代理客户端缓存:{cache_key: (client, last_used_time)} # 代理客户端缓存:{cache_key: (client, last_used_time)}
_proxy_clients: Dict[str, Tuple[httpx.AsyncClient, float]] = {} _proxy_clients: dict[str, tuple[httpx.AsyncClient, float]] = {}
# 代理客户端缓存上限(避免内存泄漏) # 代理客户端缓存上限(避免内存泄漏)
_max_proxy_clients: int = 50 _max_proxy_clients: int = 50
@@ -242,7 +244,7 @@ class HTTPClientPool:
@classmethod @classmethod
async def get_proxy_client( async def get_proxy_client(
cls, cls,
proxy_config: Optional[Dict[str, Any]] = None, proxy_config: dict[str, Any] | None = None,
) -> httpx.AsyncClient: ) -> httpx.AsyncClient:
""" """
获取代理客户端(带缓存复用) 获取代理客户端(带缓存复用)
@@ -280,7 +282,7 @@ class HTTPClientPool:
await cls._evict_lru_proxy_client() await cls._evict_lru_proxy_client()
# 创建新客户端(使用默认超时,请求时可覆盖) # 创建新客户端(使用默认超时,请求时可覆盖)
client_config: Dict[str, Any] = { client_config: dict[str, Any] = {
"http2": False, "http2": False,
"verify": get_ssl_context(), "verify": get_ssl_context(),
"follow_redirects": True, "follow_redirects": True,
@@ -370,8 +372,8 @@ class HTTPClientPool:
@classmethod @classmethod
def create_client_with_proxy( def create_client_with_proxy(
cls, cls,
proxy_config: Optional[Dict[str, Any]] = None, proxy_config: dict[str, Any] | None = None,
timeout: Optional[httpx.Timeout] = None, timeout: httpx.Timeout | None = None,
**kwargs: Any, **kwargs: Any,
) -> httpx.AsyncClient: ) -> httpx.AsyncClient:
""" """
@@ -387,7 +389,7 @@ class HTTPClientPool:
Returns: Returns:
配置好的 httpx.AsyncClient 实例(调用者需要负责关闭) 配置好的 httpx.AsyncClient 实例(调用者需要负责关闭)
""" """
client_config: Dict[str, Any] = { client_config: dict[str, Any] = {
"http2": False, "http2": False,
"verify": get_ssl_context(), "verify": get_ssl_context(),
"follow_redirects": True, "follow_redirects": True,
@@ -413,7 +415,7 @@ class HTTPClientPool:
return httpx.AsyncClient(**client_config) return httpx.AsyncClient(**client_config)
@classmethod @classmethod
def get_pool_stats(cls) -> Dict[str, Any]: def get_pool_stats(cls) -> dict[str, Any]:
"""获取连接池统计信息""" """获取连接池统计信息"""
return { return {
"default_client_active": cls._default_client is not None, "default_client_active": cls._default_client is not None,

View File

@@ -9,10 +9,11 @@
- 调用方可以根据状态决定降级策略 - 调用方可以根据状态决定降级策略
""" """
from __future__ import annotations
import os import os
import time import time
from enum import Enum from enum import Enum
from typing import Optional
import redis.asyncio as aioredis import redis.asyncio as aioredis
from src.core.logger import logger from src.core.logger import logger
@@ -35,8 +36,8 @@ class RedisClientManager:
提供 Redis 连接管理、熔断器保护和状态监控。 提供 Redis 连接管理、熔断器保护和状态监控。
""" """
_instance: Optional["RedisClientManager"] = None _instance: RedisClientManager | None = None
_redis: Optional[aioredis.Redis] = None _redis: aioredis.Redis | None = None
def __new__(cls): def __new__(cls):
"""单例模式""" """单例模式"""
@@ -50,11 +51,11 @@ class RedisClientManager:
return return
self._initialized = True self._initialized = True
self._circuit_open_until: Optional[float] = None self._circuit_open_until: float | None = None
self._consecutive_failures: int = 0 self._consecutive_failures: int = 0
self._circuit_threshold = int(os.getenv("REDIS_CIRCUIT_BREAKER_THRESHOLD", "3")) self._circuit_threshold = int(os.getenv("REDIS_CIRCUIT_BREAKER_THRESHOLD", "3"))
self._circuit_reset_seconds = int(os.getenv("REDIS_CIRCUIT_BREAKER_RESET_SECONDS", "60")) self._circuit_reset_seconds = int(os.getenv("REDIS_CIRCUIT_BREAKER_RESET_SECONDS", "60"))
self._last_error: Optional[str] = None # 记录最后一次错误 self._last_error: str | None = None # 记录最后一次错误
def get_state(self) -> RedisState: def get_state(self) -> RedisState:
""" """
@@ -100,7 +101,7 @@ class RedisClientManager:
self._consecutive_failures = 0 self._consecutive_failures = 0
self._last_error = None self._last_error = None
async def initialize(self, require_redis: bool = False) -> Optional[aioredis.Redis]: async def initialize(self, require_redis: bool = False) -> aioredis.Redis | None:
""" """
初始化Redis连接 初始化Redis连接
@@ -236,7 +237,7 @@ class RedisClientManager:
self._redis = None self._redis = None
logger.info("全局Redis客户端已关闭") logger.info("全局Redis客户端已关闭")
def get_client(self) -> Optional[aioredis.Redis]: def get_client(self) -> aioredis.Redis | None:
""" """
获取Redis客户端非异步 获取Redis客户端非异步
@@ -249,10 +250,10 @@ class RedisClientManager:
# 全局单例 # 全局单例
_redis_manager: Optional[RedisClientManager] = None _redis_manager: RedisClientManager | None = None
async def get_redis_client(require_redis: bool = False) -> Optional[aioredis.Redis]: async def get_redis_client(require_redis: bool = False) -> aioredis.Redis | None:
""" """
获取全局Redis客户端 获取全局Redis客户端
@@ -277,7 +278,7 @@ async def get_redis_client(require_redis: bool = False) -> Optional[aioredis.Red
return _redis_manager.get_client() return _redis_manager.get_client()
def get_redis_client_sync() -> Optional[aioredis.Redis]: def get_redis_client_sync() -> aioredis.Redis | None:
""" """
同步获取Redis客户端不会初始化 同步获取Redis客户端不会初始化

View File

@@ -12,8 +12,9 @@
from __future__ import annotations from __future__ import annotations
import logging import logging
from typing import TYPE_CHECKING, Optional, Tuple from typing import TYPE_CHECKING
if TYPE_CHECKING: if TYPE_CHECKING:
from src.core.api_format.conversion.registry import FormatConversionRegistry from src.core.api_format.conversion.registry import FormatConversionRegistry
@@ -26,11 +27,11 @@ logger = logging.getLogger(__name__)
def is_format_compatible( def is_format_compatible(
client_format: str, client_format: str,
endpoint_api_format: str, endpoint_api_format: str,
endpoint_format_acceptance_config: Optional[dict], endpoint_format_acceptance_config: dict | None,
is_stream: bool, is_stream: bool,
global_conversion_enabled: bool, global_conversion_enabled: bool,
registry: Optional["FormatConversionRegistry"] = None, registry: FormatConversionRegistry | None = None,
) -> Tuple[bool, bool, Optional[str]]: ) -> tuple[bool, bool, str | None]:
""" """
检查端点是否兼容客户端格式 检查端点是否兼容客户端格式

View File

@@ -4,7 +4,6 @@
用于严格模式转换失败时抛出,让编排器可以尝试下一个候选。 用于严格模式转换失败时抛出,让编排器可以尝试下一个候选。
""" """
from __future__ import annotations
class FormatConversionError(Exception): class FormatConversionError(Exception):

View File

@@ -9,13 +9,11 @@
这些应复用 `src/core/api_format/metadata.py`API_FORMAT_DEFINITIONS作为单一事实来源。 这些应复用 `src/core/api_format/metadata.py`API_FORMAT_DEFINITIONS作为单一事实来源。
""" """
from __future__ import annotations
from typing import Dict, Set
# 角色映射仅作为辅助system/tool 的具体落点以 Normalizer 规则为准) # 角色映射仅作为辅助system/tool 的具体落点以 Normalizer 规则为准)
ROLE_MAPPINGS: Dict[str, Dict[str, str]] = { ROLE_MAPPINGS: dict[str, dict[str, str]] = {
"OPENAI": { "OPENAI": {
"user": "user", "user": "user",
"assistant": "assistant", "assistant": "assistant",
@@ -29,7 +27,7 @@ ROLE_MAPPINGS: Dict[str, Dict[str, str]] = {
# 停止原因映射internal -> provider未知值使用 UNKNOWN 并写入 extra/raw # 停止原因映射internal -> provider未知值使用 UNKNOWN 并写入 extra/raw
STOP_REASON_MAPPINGS: Dict[str, Dict[str, str]] = { STOP_REASON_MAPPINGS: dict[str, dict[str, str]] = {
"CLAUDE": { "CLAUDE": {
"end_turn": "end_turn", "end_turn": "end_turn",
"max_tokens": "max_tokens", "max_tokens": "max_tokens",
@@ -60,7 +58,7 @@ STOP_REASON_MAPPINGS: Dict[str, Dict[str, str]] = {
# 使用量字段映射provider usage field -> internal UsageInfo field # 使用量字段映射provider usage field -> internal UsageInfo field
USAGE_FIELD_MAPPINGS: Dict[str, Dict[str, str]] = { USAGE_FIELD_MAPPINGS: dict[str, dict[str, str]] = {
"CLAUDE": { "CLAUDE": {
"input_tokens": "input_tokens", "input_tokens": "input_tokens",
"output_tokens": "output_tokens", "output_tokens": "output_tokens",
@@ -82,7 +80,7 @@ USAGE_FIELD_MAPPINGS: Dict[str, Dict[str, str]] = {
# 错误类型映射provider -> internal ErrorType.value # 错误类型映射provider -> internal ErrorType.value
ERROR_TYPE_MAPPINGS: Dict[str, Dict[str, str]] = { ERROR_TYPE_MAPPINGS: dict[str, dict[str, str]] = {
"CLAUDE": { "CLAUDE": {
"invalid_request_error": "invalid_request", "invalid_request_error": "invalid_request",
"authentication_error": "authentication", "authentication_error": "authentication",
@@ -116,7 +114,7 @@ ERROR_TYPE_MAPPINGS: Dict[str, Dict[str, str]] = {
# 可重试的错误类型internal ErrorType.value # 可重试的错误类型internal ErrorType.value
RETRYABLE_ERROR_TYPES: Set[str] = { RETRYABLE_ERROR_TYPES: set[str] = {
"rate_limit", "rate_limit",
"overloaded", "overloaded",
"server_error", "server_error",

View File

@@ -10,11 +10,10 @@
- 兼容优先UnknownBlock 在内部保留,但默认在输出阶段丢弃(可观测、可随时调整策略) - 兼容优先UnknownBlock 在内部保留,但默认在输出阶段丢弃(可观测、可随时调整策略)
""" """
from __future__ import annotations
from dataclasses import dataclass, field from dataclasses import dataclass, field
from enum import Enum from enum import Enum
from typing import Any, Dict, FrozenSet, List, Optional, Union from typing import Any
class Role(str, Enum): class Role(str, Enum):
@@ -65,7 +64,7 @@ class TextBlock:
type: ContentType = field(default=ContentType.TEXT, init=False) type: ContentType = field(default=ContentType.TEXT, init=False)
text: str = "" text: str = ""
extra: Dict[str, Any] = field(default_factory=dict) extra: dict[str, Any] = field(default_factory=dict)
@dataclass @dataclass
@@ -74,11 +73,11 @@ class ImageBlock:
type: ContentType = field(default=ContentType.IMAGE, init=False) type: ContentType = field(default=ContentType.IMAGE, init=False)
# base64 编码的图片数据(二选一) # base64 编码的图片数据(二选一)
data: Optional[str] = None data: str | None = None
media_type: Optional[str] = None media_type: str | None = None
# 或者 URL 引用 # 或者 URL 引用
url: Optional[str] = None url: str | None = None
extra: Dict[str, Any] = field(default_factory=dict) extra: dict[str, Any] = field(default_factory=dict)
@dataclass @dataclass
@@ -88,8 +87,8 @@ class ToolUseBlock:
type: ContentType = field(default=ContentType.TOOL_USE, init=False) type: ContentType = field(default=ContentType.TOOL_USE, init=False)
tool_id: str = "" tool_id: str = ""
tool_name: str = "" tool_name: str = ""
tool_input: Dict[str, Any] = field(default_factory=dict) tool_input: dict[str, Any] = field(default_factory=dict)
extra: Dict[str, Any] = field(default_factory=dict) extra: dict[str, Any] = field(default_factory=dict)
@dataclass @dataclass
@@ -100,9 +99,9 @@ class ToolResultBlock:
tool_use_id: str = "" # 对应的 ToolUseBlock.tool_id tool_use_id: str = "" # 对应的 ToolUseBlock.tool_id
# 工具输出可能是纯文本,也可能是结构化 JSONGemini functionResponse 等) # 工具输出可能是纯文本,也可能是结构化 JSONGemini functionResponse 等)
output: Any = None output: Any = None
content_text: Optional[str] = None content_text: str | None = None
is_error: bool = False is_error: bool = False
extra: Dict[str, Any] = field(default_factory=dict) extra: dict[str, Any] = field(default_factory=dict)
@dataclass @dataclass
@@ -111,11 +110,11 @@ class UnknownBlock:
type: ContentType = field(default=ContentType.UNKNOWN, init=False) type: ContentType = field(default=ContentType.UNKNOWN, init=False)
raw_type: str = "" # 原始的类型字符串(各格式不一致) raw_type: str = "" # 原始的类型字符串(各格式不一致)
payload: Dict[str, Any] = field(default_factory=dict) # 原始结构(尽量保持) payload: dict[str, Any] = field(default_factory=dict) # 原始结构(尽量保持)
extra: Dict[str, Any] = field(default_factory=dict) extra: dict[str, Any] = field(default_factory=dict)
ContentBlock = Union[TextBlock, ImageBlock, ToolUseBlock, ToolResultBlock, UnknownBlock] ContentBlock = TextBlock | ImageBlock | ToolUseBlock | ToolResultBlock | UnknownBlock
@dataclass @dataclass
@@ -123,8 +122,8 @@ class InternalMessage:
"""统一的消息表示""" """统一的消息表示"""
role: Role role: Role
content: List[ContentBlock] # 统一使用列表,纯文本用单个 TextBlock content: list[ContentBlock] # 统一使用列表,纯文本用单个 TextBlock
extra: Dict[str, Any] = field(default_factory=dict) extra: dict[str, Any] = field(default_factory=dict)
@dataclass @dataclass
@@ -132,9 +131,9 @@ class ToolDefinition:
"""统一的工具定义""" """统一的工具定义"""
name: str name: str
description: Optional[str] = None description: str | None = None
parameters: Optional[Dict[str, Any]] = None # JSON Schema parameters: dict[str, Any] | None = None # JSON Schema
extra: Dict[str, Any] = field(default_factory=dict) extra: dict[str, Any] = field(default_factory=dict)
class ToolChoiceType(str, Enum): class ToolChoiceType(str, Enum):
@@ -149,8 +148,8 @@ class ToolChoice:
"""统一的工具选择""" """统一的工具选择"""
type: ToolChoiceType type: ToolChoiceType
tool_name: Optional[str] = None tool_name: str | None = None
extra: Dict[str, Any] = field(default_factory=dict) extra: dict[str, Any] = field(default_factory=dict)
@dataclass @dataclass
@@ -159,7 +158,7 @@ class InstructionSegment:
role: Role # 仅允许 Role.SYSTEM / Role.DEVELOPER role: Role # 仅允许 Role.SYSTEM / Role.DEVELOPER
text: str = "" text: str = ""
extra: Dict[str, Any] = field(default_factory=dict) extra: dict[str, Any] = field(default_factory=dict)
@dataclass @dataclass
@@ -167,25 +166,25 @@ class InternalRequest:
"""统一的请求表示""" """统一的请求表示"""
model: str model: str
messages: List[InternalMessage] messages: list[InternalMessage]
# 指令层:保留 system/developer 结构与顺序 # 指令层:保留 system/developer 结构与顺序
instructions: List[InstructionSegment] = field(default_factory=list) instructions: list[InstructionSegment] = field(default_factory=list)
# 兼容字段instructions 的 join 文本(无 role 标签),用于 Claude/Gemini 这类仅接受字符串 system 的格式 # 兼容字段instructions 的 join 文本(无 role 标签),用于 Claude/Gemini 这类仅接受字符串 system 的格式
system: Optional[str] = None system: str | None = None
max_tokens: Optional[int] = None max_tokens: int | None = None
temperature: Optional[float] = None temperature: float | None = None
top_p: Optional[float] = None top_p: float | None = None
top_k: Optional[int] = None top_k: int | None = None
stop_sequences: Optional[List[str]] = None stop_sequences: list[str] | None = None
stream: bool = False stream: bool = False
tools: Optional[List[ToolDefinition]] = None tools: list[ToolDefinition] | None = None
tool_choice: Optional[ToolChoice] = None # auto/none/required 或指定 tool_name tool_choice: ToolChoice | None = None # auto/none/required 或指定 tool_name
extra: Dict[str, Any] = field(default_factory=dict) # 未识别字段透传 extra: dict[str, Any] = field(default_factory=dict) # 未识别字段透传
def to_debug_dict(self) -> Dict[str, Any]: def to_debug_dict(self) -> dict[str, Any]:
"""用于日志和调试的简化表示""" """用于日志和调试的简化表示"""
return { return {
"model": self.model, "model": self.model,
@@ -208,7 +207,7 @@ class UsageInfo:
total_tokens: int = 0 total_tokens: int = 0
cache_read_tokens: int = 0 cache_read_tokens: int = 0
cache_write_tokens: int = 0 cache_write_tokens: int = 0
extra: Dict[str, Any] = field(default_factory=dict) extra: dict[str, Any] = field(default_factory=dict)
@dataclass @dataclass
@@ -217,12 +216,12 @@ class InternalResponse:
id: str id: str
model: str model: str
content: List[ContentBlock] content: list[ContentBlock]
stop_reason: Optional[StopReason] = None stop_reason: StopReason | None = None
usage: Optional[UsageInfo] = None usage: UsageInfo | None = None
extra: Dict[str, Any] = field(default_factory=dict) extra: dict[str, Any] = field(default_factory=dict)
def to_debug_dict(self) -> Dict[str, Any]: def to_debug_dict(self) -> dict[str, Any]:
"""用于日志和调试的简化表示""" """用于日志和调试的简化表示"""
usage = None usage = None
if self.usage: if self.usage:
@@ -246,12 +245,12 @@ class InternalError:
type: ErrorType type: ErrorType
message: str message: str
code: Optional[str] = None # 原始错误码 code: str | None = None # 原始错误码
param: Optional[str] = None # 导致错误的参数 param: str | None = None # 导致错误的参数
retryable: bool = False # 是否可重试 retryable: bool = False # 是否可重试
extra: Dict[str, Any] = field(default_factory=dict) extra: dict[str, Any] = field(default_factory=dict)
def to_debug_dict(self) -> Dict[str, Any]: def to_debug_dict(self) -> dict[str, Any]:
"""用于日志和调试""" """用于日志和调试"""
return { return {
"type": self.type.value, "type": self.type.value,
@@ -269,7 +268,7 @@ class FormatCapabilities:
supports_error_conversion: bool = True supports_error_conversion: bool = True
supports_tools: bool = True supports_tools: bool = True
supports_images: bool = False supports_images: bool = False
supported_features: FrozenSet[str] = field(default_factory=frozenset) supported_features: frozenset[str] = field(default_factory=frozenset)
__all__ = [ __all__ = [

View File

@@ -5,10 +5,9 @@
再从 internal 输出到目标格式。 再从 internal 输出到目标格式。
""" """
from __future__ import annotations
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from typing import Any, Dict, List, Optional from typing import Any
from .internal import FormatCapabilities, InternalError, InternalRequest, InternalResponse from .internal import FormatCapabilities, InternalError, InternalRequest, InternalResponse
from .stream_events import InternalStreamEvent from .stream_events import InternalStreamEvent
@@ -24,19 +23,19 @@ class FormatNormalizer(ABC):
# ============ 请求转换 ============ # ============ 请求转换 ============
@abstractmethod @abstractmethod
def request_to_internal(self, request: Dict[str, Any]) -> InternalRequest: def request_to_internal(self, request: dict[str, Any]) -> InternalRequest:
"""将格式特定请求转换为内部表示""" """将格式特定请求转换为内部表示"""
raise NotImplementedError raise NotImplementedError
@abstractmethod @abstractmethod
def request_from_internal(self, internal: InternalRequest) -> Dict[str, Any]: def request_from_internal(self, internal: InternalRequest) -> dict[str, Any]:
"""将内部表示转换为格式特定请求""" """将内部表示转换为格式特定请求"""
raise NotImplementedError raise NotImplementedError
# ============ 响应转换 ============ # ============ 响应转换 ============
@abstractmethod @abstractmethod
def response_to_internal(self, response: Dict[str, Any]) -> InternalResponse: def response_to_internal(self, response: dict[str, Any]) -> InternalResponse:
"""将格式特定响应转换为内部表示""" """将格式特定响应转换为内部表示"""
raise NotImplementedError raise NotImplementedError
@@ -45,8 +44,8 @@ class FormatNormalizer(ABC):
self, self,
internal: InternalResponse, internal: InternalResponse,
*, *,
requested_model: Optional[str] = None, requested_model: str | None = None,
) -> Dict[str, Any]: ) -> dict[str, Any]:
"""将内部表示转换为格式特定响应 """将内部表示转换为格式特定响应
Args: Args:
@@ -61,9 +60,9 @@ class FormatNormalizer(ABC):
def stream_chunk_to_internal( def stream_chunk_to_internal(
self, self,
chunk: Dict[str, Any], chunk: dict[str, Any],
state: StreamState, state: StreamState,
) -> List[InternalStreamEvent]: ) -> list[InternalStreamEvent]:
"""将格式特定流式块转换为内部事件""" """将格式特定流式块转换为内部事件"""
raise NotImplementedError raise NotImplementedError
@@ -71,21 +70,21 @@ class FormatNormalizer(ABC):
self, self,
event: InternalStreamEvent, event: InternalStreamEvent,
state: StreamState, state: StreamState,
) -> List[Dict[str, Any]]: ) -> list[dict[str, Any]]:
"""将内部事件转换为格式特定流式块""" """将内部事件转换为格式特定流式块"""
raise NotImplementedError raise NotImplementedError
# ============ 错误转换(可选) ============ # ============ 错误转换(可选) ============
def is_error_response(self, response: Dict[str, Any]) -> bool: def is_error_response(self, response: dict[str, Any]) -> bool:
"""基于 body 的兜底判断(不可靠),子类可覆盖""" """基于 body 的兜底判断(不可靠),子类可覆盖"""
return False return False
def error_to_internal(self, error_response: Dict[str, Any]) -> InternalError: def error_to_internal(self, error_response: dict[str, Any]) -> InternalError:
"""将格式特定错误转换为内部表示""" """将格式特定错误转换为内部表示"""
raise NotImplementedError raise NotImplementedError
def error_from_internal(self, internal: InternalError) -> Dict[str, Any]: def error_from_internal(self, internal: InternalError) -> dict[str, Any]:
"""将内部错误表示转换为格式特定错误""" """将内部错误表示转换为格式特定错误"""
raise NotImplementedError raise NotImplementedError

View File

@@ -6,7 +6,6 @@ Normalizers
本目录在 Phase 1 仅创建结构;具体实现将在 Phase 2+ 补齐。 本目录在 Phase 1 仅创建结构;具体实现将在 Phase 2+ 补齐。
""" """
from __future__ import annotations
__all__: list[str] = [] __all__: list[str] = []

View File

@@ -7,10 +7,9 @@ Claude Messages API Normalizer
- 可选Claude error <-> InternalError - 可选Claude error <-> InternalError
""" """
from __future__ import annotations
import json import json
from typing import Any, Dict, List, Optional, Tuple from typing import Any
from src.core.api_format.conversion.field_mappings import ( from src.core.api_format.conversion.field_mappings import (
ERROR_TYPE_MAPPINGS, ERROR_TYPE_MAPPINGS,
@@ -63,7 +62,7 @@ class ClaudeNormalizer(FormatNormalizer):
supports_images=True, supports_images=True,
) )
_CLAUDE_STOP_TO_INTERNAL: Dict[str, StopReason] = { _CLAUDE_STOP_TO_INTERNAL: dict[str, StopReason] = {
"end_turn": StopReason.END_TURN, "end_turn": StopReason.END_TURN,
"max_tokens": StopReason.MAX_TOKENS, "max_tokens": StopReason.MAX_TOKENS,
"stop_sequence": StopReason.STOP_SEQUENCE, "stop_sequence": StopReason.STOP_SEQUENCE,
@@ -73,7 +72,7 @@ class ClaudeNormalizer(FormatNormalizer):
"content_filtered": StopReason.CONTENT_FILTERED, "content_filtered": StopReason.CONTENT_FILTERED,
} }
_ERROR_TYPE_TO_CLAUDE: Dict[ErrorType, str] = { _ERROR_TYPE_TO_CLAUDE: dict[ErrorType, str] = {
ErrorType.INVALID_REQUEST: "invalid_request_error", ErrorType.INVALID_REQUEST: "invalid_request_error",
ErrorType.AUTHENTICATION: "authentication_error", ErrorType.AUTHENTICATION: "authentication_error",
ErrorType.PERMISSION_DENIED: "permission_error", ErrorType.PERMISSION_DENIED: "permission_error",
@@ -90,11 +89,11 @@ class ClaudeNormalizer(FormatNormalizer):
# Requests # Requests
# ========================= # =========================
def request_to_internal(self, request: Dict[str, Any]) -> InternalRequest: def request_to_internal(self, request: dict[str, Any]) -> InternalRequest:
model = str(request.get("model") or "") model = str(request.get("model") or "")
dropped: Dict[str, int] = {} dropped: dict[str, int] = {}
instructions: List[InstructionSegment] = [] instructions: list[InstructionSegment] = []
# 顶层 system 先进入 instructions保持确定性优先级 # 顶层 system 先进入 instructions保持确定性优先级
sys_value = request.get("system") sys_value = request.get("system")
@@ -103,7 +102,7 @@ class ClaudeNormalizer(FormatNormalizer):
if sys_text: if sys_text:
instructions.append(InstructionSegment(role=Role.SYSTEM, text=sys_text)) instructions.append(InstructionSegment(role=Role.SYSTEM, text=sys_text))
messages: List[InternalMessage] = [] messages: list[InternalMessage] = []
for msg in request.get("messages") or []: for msg in request.get("messages") or []:
if not isinstance(msg, dict): if not isinstance(msg, dict):
dropped["claude_message_non_dict"] = dropped.get("claude_message_non_dict", 0) + 1 dropped["claude_message_non_dict"] = dropped.get("claude_message_non_dict", 0) + 1
@@ -156,15 +155,15 @@ class ClaudeNormalizer(FormatNormalizer):
return internal return internal
def request_from_internal(self, internal: InternalRequest) -> Dict[str, Any]: def request_from_internal(self, internal: InternalRequest) -> dict[str, Any]:
system_text = internal.system or self._join_instructions(internal.instructions) system_text = internal.system or self._join_instructions(internal.instructions)
# Claude Messages API: messages[] 仅允许 user/assistant且需要交替这里做最小修复 # Claude Messages API: messages[] 仅允许 user/assistant且需要交替这里做最小修复
fixed_messages = self._coerce_claude_message_sequence(internal.messages) fixed_messages = self._coerce_claude_message_sequence(internal.messages)
out_messages: List[Dict[str, Any]] = [self._internal_message_to_claude(m) for m in fixed_messages] out_messages: list[dict[str, Any]] = [self._internal_message_to_claude(m) for m in fixed_messages]
result: Dict[str, Any] = { result: dict[str, Any] = {
"model": internal.model, "model": internal.model,
"messages": out_messages, "messages": out_messages,
"max_tokens": internal.max_tokens if internal.max_tokens is not None else 4096, "max_tokens": internal.max_tokens if internal.max_tokens is not None else 4096,
@@ -210,20 +209,20 @@ class ClaudeNormalizer(FormatNormalizer):
# Responses # Responses
# ========================= # =========================
def response_to_internal(self, response: Dict[str, Any]) -> InternalResponse: def response_to_internal(self, response: dict[str, Any]) -> InternalResponse:
rid = str(response.get("id") or "") rid = str(response.get("id") or "")
model = str(response.get("model") or "") model = str(response.get("model") or "")
blocks, dropped = self._claude_content_to_blocks(response.get("content")) blocks, dropped = self._claude_content_to_blocks(response.get("content"))
raw_stop = response.get("stop_reason") raw_stop = response.get("stop_reason")
stop_reason: Optional[StopReason] = None stop_reason: StopReason | None = None
if raw_stop is not None: if raw_stop is not None:
stop_reason = self._CLAUDE_STOP_TO_INTERNAL.get(str(raw_stop), StopReason.UNKNOWN) stop_reason = self._CLAUDE_STOP_TO_INTERNAL.get(str(raw_stop), StopReason.UNKNOWN)
usage_info = self._claude_usage_to_internal(response.get("usage")) usage_info = self._claude_usage_to_internal(response.get("usage"))
extra: Dict[str, Any] = {} extra: dict[str, Any] = {}
if raw_stop is not None: if raw_stop is not None:
extra.setdefault("raw", {})["stop_reason"] = raw_stop extra.setdefault("raw", {})["stop_reason"] = raw_stop
@@ -245,13 +244,13 @@ class ClaudeNormalizer(FormatNormalizer):
self, self,
internal: InternalResponse, internal: InternalResponse,
*, *,
requested_model: Optional[str] = None, requested_model: str | None = None,
) -> Dict[str, Any]: ) -> dict[str, Any]:
cid = internal.id or "unknown" cid = internal.id or "unknown"
if not cid.startswith("msg_"): if not cid.startswith("msg_"):
cid = f"msg_{cid}" cid = f"msg_{cid}"
content: List[Dict[str, Any]] = [] content: list[dict[str, Any]] = []
for b in internal.content: for b in internal.content:
if isinstance(b, TextBlock): if isinstance(b, TextBlock):
if b.text: if b.text:
@@ -288,7 +287,7 @@ class ClaudeNormalizer(FormatNormalizer):
if internal.stop_reason is not None: if internal.stop_reason is not None:
stop_reason = STOP_REASON_MAPPINGS.get("CLAUDE", {}).get(internal.stop_reason.value, "end_turn") stop_reason = STOP_REASON_MAPPINGS.get("CLAUDE", {}).get(internal.stop_reason.value, "end_turn")
usage: Dict[str, Any] = {"input_tokens": 0, "output_tokens": 0} usage: dict[str, Any] = {"input_tokens": 0, "output_tokens": 0}
if internal.usage: if internal.usage:
usage = { usage = {
"input_tokens": int(internal.usage.input_tokens), "input_tokens": int(internal.usage.input_tokens),
@@ -319,11 +318,11 @@ class ClaudeNormalizer(FormatNormalizer):
def stream_chunk_to_internal( def stream_chunk_to_internal(
self, self,
chunk: Dict[str, Any], chunk: dict[str, Any],
state: StreamState, state: StreamState,
) -> List[InternalStreamEvent]: ) -> list[InternalStreamEvent]:
ss = state.substate(self.FORMAT_ID) ss = state.substate(self.FORMAT_ID)
events: List[InternalStreamEvent] = [] events: list[InternalStreamEvent] = []
event_type = chunk.get("type") event_type = chunk.get("type")
if event_type is None: if event_type is None:
@@ -335,7 +334,7 @@ class ClaudeNormalizer(FormatNormalizer):
if event_type == "message_start": if event_type == "message_start":
message_raw = chunk.get("message") message_raw = chunk.get("message")
message: Dict[str, Any] = message_raw if isinstance(message_raw, dict) else {} message: dict[str, Any] = message_raw if isinstance(message_raw, dict) else {}
msg_id = str(message.get("id") or "") msg_id = str(message.get("id") or "")
# 保留初始化时设置的 model客户端请求的模型仅在空时用上游值 # 保留初始化时设置的 model客户端请求的模型仅在空时用上游值
model = state.model or str(message.get("model") or "") model = state.model or str(message.get("model") or "")
@@ -350,7 +349,7 @@ class ClaudeNormalizer(FormatNormalizer):
if event_type == "content_block_start": if event_type == "content_block_start":
index = int(chunk.get("index") or 0) index = int(chunk.get("index") or 0)
block_raw = chunk.get("content_block") block_raw = chunk.get("content_block")
block: Dict[str, Any] = block_raw if isinstance(block_raw, dict) else {} block: dict[str, Any] = block_raw if isinstance(block_raw, dict) else {}
btype = str(block.get("type") or "unknown") btype = str(block.get("type") or "unknown")
if btype == "text": if btype == "text":
@@ -385,7 +384,7 @@ class ClaudeNormalizer(FormatNormalizer):
if event_type == "content_block_delta": if event_type == "content_block_delta":
index = int(chunk.get("index") or 0) index = int(chunk.get("index") or 0)
delta_raw = chunk.get("delta") delta_raw = chunk.get("delta")
delta: Dict[str, Any] = delta_raw if isinstance(delta_raw, dict) else {} delta: dict[str, Any] = delta_raw if isinstance(delta_raw, dict) else {}
dtype = str(delta.get("type") or "unknown") dtype = str(delta.get("type") or "unknown")
if dtype == "text_delta": if dtype == "text_delta":
@@ -415,7 +414,7 @@ class ClaudeNormalizer(FormatNormalizer):
if event_type == "message_delta": if event_type == "message_delta":
delta_raw2 = chunk.get("delta") delta_raw2 = chunk.get("delta")
delta2: Dict[str, Any] = delta_raw2 if isinstance(delta_raw2, dict) else {} delta2: dict[str, Any] = delta_raw2 if isinstance(delta_raw2, dict) else {}
raw_stop = delta2.get("stop_reason") raw_stop = delta2.get("stop_reason")
if raw_stop is not None: if raw_stop is not None:
ss["stop_reason"] = str(raw_stop) ss["stop_reason"] = str(raw_stop)
@@ -426,7 +425,7 @@ class ClaudeNormalizer(FormatNormalizer):
if event_type == "message_stop": if event_type == "message_stop":
raw_stop = ss.get("stop_reason") raw_stop = ss.get("stop_reason")
stop_reason: Optional[StopReason] = None stop_reason: StopReason | None = None
if raw_stop is not None: if raw_stop is not None:
stop_reason = self._CLAUDE_STOP_TO_INTERNAL.get(str(raw_stop), StopReason.UNKNOWN) stop_reason = self._CLAUDE_STOP_TO_INTERNAL.get(str(raw_stop), StopReason.UNKNOWN)
usage_info = self._claude_usage_to_internal(ss.get("usage")) usage_info = self._claude_usage_to_internal(ss.get("usage"))
@@ -444,9 +443,9 @@ class ClaudeNormalizer(FormatNormalizer):
self, self,
event: InternalStreamEvent, event: InternalStreamEvent,
state: StreamState, state: StreamState,
) -> List[Dict[str, Any]]: ) -> list[dict[str, Any]]:
ss = state.substate(self.FORMAT_ID) ss = state.substate(self.FORMAT_ID)
out: List[Dict[str, Any]] = [] out: list[dict[str, Any]] = []
if isinstance(event, MessageStartEvent): if isinstance(event, MessageStartEvent):
state.message_id = event.message_id or state.message_id state.message_id = event.message_id or state.message_id
@@ -454,7 +453,7 @@ class ClaudeNormalizer(FormatNormalizer):
if not state.model: if not state.model:
state.model = event.model or "" state.model = event.model or ""
ss.setdefault("block_index_to_tool_id", {}) ss.setdefault("block_index_to_tool_id", {})
message_obj: Dict[str, Any] = { message_obj: dict[str, Any] = {
"id": state.message_id or "msg_stream", "id": state.message_id or "msg_stream",
"type": "message", "type": "message",
"role": "assistant", "role": "assistant",
@@ -531,7 +530,7 @@ class ClaudeNormalizer(FormatNormalizer):
if event.stop_reason is not None: if event.stop_reason is not None:
stop_reason = STOP_REASON_MAPPINGS.get("CLAUDE", {}).get(event.stop_reason.value, "end_turn") stop_reason = STOP_REASON_MAPPINGS.get("CLAUDE", {}).get(event.stop_reason.value, "end_turn")
msg_delta: Dict[str, Any] = { msg_delta: dict[str, Any] = {
"type": "message_delta", "type": "message_delta",
"delta": {"stop_reason": stop_reason}, "delta": {"stop_reason": stop_reason},
} }
@@ -553,15 +552,15 @@ class ClaudeNormalizer(FormatNormalizer):
# Error conversion # Error conversion
# ========================= # =========================
def is_error_response(self, response: Dict[str, Any]) -> bool: def is_error_response(self, response: dict[str, Any]) -> bool:
if not isinstance(response, dict): if not isinstance(response, dict):
return False return False
if response.get("type") == "error": if response.get("type") == "error":
return True return True
return "error" in response return "error" in response
def error_to_internal(self, error_response: Dict[str, Any]) -> InternalError: def error_to_internal(self, error_response: dict[str, Any]) -> InternalError:
err: Dict[str, Any] = {} err: dict[str, Any] = {}
if isinstance(error_response, dict): if isinstance(error_response, dict):
err_raw = error_response.get("error") err_raw = error_response.get("error")
err = err_raw if isinstance(err_raw, dict) else {} err = err_raw if isinstance(err_raw, dict) else {}
@@ -580,9 +579,9 @@ class ClaudeNormalizer(FormatNormalizer):
extra={"claude": {"error": err}, "raw": {"type": raw_type}}, extra={"claude": {"error": err}, "raw": {"type": raw_type}},
) )
def error_from_internal(self, internal: InternalError) -> Dict[str, Any]: def error_from_internal(self, internal: InternalError) -> dict[str, Any]:
type_str = self._ERROR_TYPE_TO_CLAUDE.get(internal.type, "api_error") type_str = self._ERROR_TYPE_TO_CLAUDE.get(internal.type, "api_error")
payload: Dict[str, Any] = {"type": type_str, "message": internal.message} payload: dict[str, Any] = {"type": type_str, "message": internal.message}
if internal.param is not None: if internal.param is not None:
payload["param"] = internal.param payload["param"] = internal.param
if internal.code is not None: if internal.code is not None:
@@ -593,8 +592,8 @@ class ClaudeNormalizer(FormatNormalizer):
# Helpers # Helpers
# ========================= # =========================
def _claude_message_to_internal(self, msg: Dict[str, Any]) -> Tuple[Optional[InternalMessage], Dict[str, int]]: def _claude_message_to_internal(self, msg: dict[str, Any]) -> tuple[InternalMessage | None, dict[str, int]]:
dropped: Dict[str, int] = {} dropped: dict[str, int] = {}
role_raw = str(msg.get("role") or "unknown") role_raw = str(msg.get("role") or "unknown")
if role_raw == "user": if role_raw == "user":
@@ -616,8 +615,8 @@ class ClaudeNormalizer(FormatNormalizer):
dropped, dropped,
) )
def _claude_content_to_blocks(self, content: Any) -> Tuple[List[ContentBlock], Dict[str, int]]: def _claude_content_to_blocks(self, content: Any) -> tuple[list[ContentBlock], dict[str, int]]:
dropped: Dict[str, int] = {} dropped: dict[str, int] = {}
if content is None: if content is None:
return [], dropped return [], dropped
if isinstance(content, str): if isinstance(content, str):
@@ -626,7 +625,7 @@ class ClaudeNormalizer(FormatNormalizer):
dropped["claude_content_non_list"] = dropped.get("claude_content_non_list", 0) + 1 dropped["claude_content_non_list"] = dropped.get("claude_content_non_list", 0) + 1
return [], dropped return [], dropped
blocks: List[ContentBlock] = [] blocks: list[ContentBlock] = []
for block in content: for block in content:
if not isinstance(block, dict): if not isinstance(block, dict):
dropped["claude_block_non_dict"] = dropped.get("claude_block_non_dict", 0) + 1 dropped["claude_block_non_dict"] = dropped.get("claude_block_non_dict", 0) + 1
@@ -641,7 +640,7 @@ class ClaudeNormalizer(FormatNormalizer):
if btype == "image": if btype == "image":
src_raw = block.get("source") src_raw = block.get("source")
src: Dict[str, Any] = src_raw if isinstance(src_raw, dict) else {} src: dict[str, Any] = src_raw if isinstance(src_raw, dict) else {}
stype = src.get("type") stype = src.get("type")
if stype == "base64": if stype == "base64":
data = src.get("data") data = src.get("data")
@@ -687,7 +686,7 @@ class ClaudeNormalizer(FormatNormalizer):
tool_use_id: str, tool_use_id: str,
raw_content: Any, raw_content: Any,
is_error: bool, is_error: bool,
raw_block: Dict[str, Any], raw_block: dict[str, Any],
) -> ToolResultBlock: ) -> ToolResultBlock:
if raw_content is None: if raw_content is None:
return ToolResultBlock( return ToolResultBlock(
@@ -723,7 +722,7 @@ class ClaudeNormalizer(FormatNormalizer):
) )
if isinstance(raw_content, list): if isinstance(raw_content, list):
text_parts: List[str] = [] text_parts: list[str] = []
for part in raw_content: for part in raw_content:
if isinstance(part, dict) and part.get("type") == "text": if isinstance(part, dict) and part.get("type") == "text":
text = part.get("text") text = part.get("text")
@@ -747,15 +746,15 @@ class ClaudeNormalizer(FormatNormalizer):
extra={"claude": raw_block}, extra={"claude": raw_block},
) )
def _collapse_claude_system(self, system_value: Any) -> Tuple[Optional[str], Dict[str, int]]: def _collapse_claude_system(self, system_value: Any) -> tuple[str | None, dict[str, int]]:
dropped: Dict[str, int] = {} dropped: dict[str, int] = {}
if system_value is None: if system_value is None:
return None, dropped return None, dropped
if isinstance(system_value, str): if isinstance(system_value, str):
return (system_value or None), dropped return (system_value or None), dropped
if isinstance(system_value, list): if isinstance(system_value, list):
texts: List[str] = [] texts: list[str] = []
for item in system_value: for item in system_value:
if not isinstance(item, dict): if not isinstance(item, dict):
dropped["claude_system_item_non_dict"] = dropped.get("claude_system_item_non_dict", 0) + 1 dropped["claude_system_item_non_dict"] = dropped.get("claude_system_item_non_dict", 0) + 1
@@ -773,16 +772,16 @@ class ClaudeNormalizer(FormatNormalizer):
dropped["claude_system_unsupported"] = dropped.get("claude_system_unsupported", 0) + 1 dropped["claude_system_unsupported"] = dropped.get("claude_system_unsupported", 0) + 1
return None, dropped return None, dropped
def _join_instructions(self, instructions: List[InstructionSegment]) -> Optional[str]: def _join_instructions(self, instructions: list[InstructionSegment]) -> str | None:
parts = [seg.text for seg in instructions if seg.text] parts = [seg.text for seg in instructions if seg.text]
joined = "\n\n".join(parts) joined = "\n\n".join(parts)
return joined or None return joined or None
def _claude_tools_to_internal(self, tools: Any) -> Optional[List[ToolDefinition]]: def _claude_tools_to_internal(self, tools: Any) -> list[ToolDefinition] | None:
if not tools or not isinstance(tools, list): if not tools or not isinstance(tools, list):
return None return None
out: List[ToolDefinition] = [] out: list[ToolDefinition] = []
for tool in tools: for tool in tools:
if not isinstance(tool, dict): if not isinstance(tool, dict):
continue continue
@@ -799,7 +798,7 @@ class ClaudeNormalizer(FormatNormalizer):
) )
return out or None return out or None
def _claude_tool_choice_to_internal(self, tool_choice: Any) -> Optional[ToolChoice]: def _claude_tool_choice_to_internal(self, tool_choice: Any) -> ToolChoice | None:
if tool_choice is None: if tool_choice is None:
return None return None
if not isinstance(tool_choice, dict): if not isinstance(tool_choice, dict):
@@ -818,7 +817,7 @@ class ClaudeNormalizer(FormatNormalizer):
return ToolChoice(type=ToolChoiceType.AUTO, extra={"claude": tool_choice}) return ToolChoice(type=ToolChoiceType.AUTO, extra={"claude": tool_choice})
def _tool_choice_to_claude(self, tool_choice: ToolChoice) -> Dict[str, Any]: def _tool_choice_to_claude(self, tool_choice: ToolChoice) -> dict[str, Any]:
if tool_choice.type == ToolChoiceType.NONE: if tool_choice.type == ToolChoiceType.NONE:
return {"type": "none"} return {"type": "none"}
if tool_choice.type == ToolChoiceType.AUTO: if tool_choice.type == ToolChoiceType.AUTO:
@@ -829,11 +828,11 @@ class ClaudeNormalizer(FormatNormalizer):
return {"type": "tool_use", "name": tool_choice.tool_name or ""} return {"type": "tool_use", "name": tool_choice.tool_name or ""}
return {"type": "auto"} return {"type": "auto"}
def _internal_message_to_claude(self, msg: InternalMessage) -> Dict[str, Any]: def _internal_message_to_claude(self, msg: InternalMessage) -> dict[str, Any]:
role = "user" if msg.role == Role.USER else "assistant" role = "user" if msg.role == Role.USER else "assistant"
blocks: List[Dict[str, Any]] = [] blocks: list[dict[str, Any]] = []
text_parts: List[str] = [] text_parts: list[str] = []
for b in msg.content: for b in msg.content:
if isinstance(b, UnknownBlock): if isinstance(b, UnknownBlock):
@@ -902,8 +901,8 @@ class ClaudeNormalizer(FormatNormalizer):
return {"role": role, "content": blocks} return {"role": role, "content": blocks}
def _coerce_claude_message_sequence(self, messages: List[InternalMessage]) -> List[InternalMessage]: def _coerce_claude_message_sequence(self, messages: list[InternalMessage]) -> list[InternalMessage]:
normalized: List[InternalMessage] = [] normalized: list[InternalMessage] = []
for m in messages: for m in messages:
role = m.role role = m.role
if role not in (Role.USER, Role.ASSISTANT): if role not in (Role.USER, Role.ASSISTANT):
@@ -916,7 +915,7 @@ class ClaudeNormalizer(FormatNormalizer):
if normalized[0].role != Role.USER: if normalized[0].role != Role.USER:
normalized = [InternalMessage(role=Role.USER, content=[])] + normalized normalized = [InternalMessage(role=Role.USER, content=[])] + normalized
merged: List[InternalMessage] = [] merged: list[InternalMessage] = []
for m in normalized: for m in normalized:
if merged and merged[-1].role == m.role: if merged and merged[-1].role == m.role:
merged[-1].content.extend(m.content) merged[-1].content.extend(m.content)
@@ -925,12 +924,12 @@ class ClaudeNormalizer(FormatNormalizer):
return merged return merged
def _claude_usage_to_internal(self, usage: Any) -> Optional[UsageInfo]: def _claude_usage_to_internal(self, usage: Any) -> UsageInfo | None:
if not isinstance(usage, dict): if not isinstance(usage, dict):
return None return None
mapping = USAGE_FIELD_MAPPINGS.get("CLAUDE", {}) mapping = USAGE_FIELD_MAPPINGS.get("CLAUDE", {})
fields: Dict[str, int] = {} fields: dict[str, int] = {}
extra = self._extract_extra(usage, set(mapping.keys())) extra = self._extract_extra(usage, set(mapping.keys()))
for provider_key, internal_key in mapping.items(): for provider_key, internal_key in mapping.items():
@@ -952,8 +951,8 @@ class ClaudeNormalizer(FormatNormalizer):
extra={"claude": extra} if extra else {}, extra={"claude": extra} if extra else {},
) )
def _usage_to_claude(self, usage: UsageInfo) -> Dict[str, Any]: def _usage_to_claude(self, usage: UsageInfo) -> dict[str, Any]:
result: Dict[str, Any] = { result: dict[str, Any] = {
"input_tokens": int(usage.input_tokens), "input_tokens": int(usage.input_tokens),
"output_tokens": int(usage.output_tokens), "output_tokens": int(usage.output_tokens),
} }
@@ -969,7 +968,7 @@ class ClaudeNormalizer(FormatNormalizer):
except ValueError: except ValueError:
return ErrorType.UNKNOWN return ErrorType.UNKNOWN
def _optional_int(self, value: Any) -> Optional[int]: def _optional_int(self, value: Any) -> int | None:
if value is None: if value is None:
return None return None
try: try:
@@ -977,7 +976,7 @@ class ClaudeNormalizer(FormatNormalizer):
except (TypeError, ValueError): except (TypeError, ValueError):
return None return None
def _optional_float(self, value: Any) -> Optional[float]: def _optional_float(self, value: Any) -> float | None:
if value is None: if value is None:
return None return None
try: try:
@@ -985,7 +984,7 @@ class ClaudeNormalizer(FormatNormalizer):
except (TypeError, ValueError): except (TypeError, ValueError):
return None return None
def _coerce_str_list(self, value: Any) -> Optional[List[str]]: def _coerce_str_list(self, value: Any) -> list[str] | None:
if value is None: if value is None:
return None return None
if isinstance(value, str): if isinstance(value, str):
@@ -994,10 +993,10 @@ class ClaudeNormalizer(FormatNormalizer):
return [str(x) for x in value if x is not None] return [str(x) for x in value if x is not None]
return None return None
def _extract_extra(self, payload: Dict[str, Any], known_keys: set[str]) -> Dict[str, Any]: def _extract_extra(self, payload: dict[str, Any], known_keys: set[str]) -> dict[str, Any]:
return {k: v for k, v in payload.items() if k not in known_keys} return {k: v for k, v in payload.items() if k not in known_keys}
def _merge_dropped(self, target: Dict[str, int], source: Dict[str, int]) -> None: def _merge_dropped(self, target: dict[str, int], source: dict[str, int]) -> None:
for k, v in source.items(): for k, v in source.items():
target[k] = target.get(k, 0) + int(v) target[k] = target.get(k, 0) + int(v)

View File

@@ -7,7 +7,6 @@ CLAUDE_CLI 的请求/响应 body 与 CLAUDE 一致Anthropic Messages API
如需 CLI 特殊处理,可覆盖 request_from_internal / request_to_internal 等方法。 如需 CLI 特殊处理,可覆盖 request_from_internal / request_to_internal 等方法。
""" """
from __future__ import annotations
from src.core.api_format.conversion.normalizers.claude import ClaudeNormalizer from src.core.api_format.conversion.normalizers.claude import ClaudeNormalizer

View File

@@ -11,11 +11,9 @@ Gemini (GenerateContent / streamGenerateContent) Normalizer
- 响应/流式通常为 camelCasecandidates/finishReason/usageMetadata/modelVersion - 响应/流式通常为 camelCasecandidates/finishReason/usageMetadata/modelVersion
""" """
from __future__ import annotations
import json import json
import time from typing import Any
from typing import Any, Dict, List, Optional, Tuple
from src.core.api_format.conversion.field_mappings import ( from src.core.api_format.conversion.field_mappings import (
ERROR_TYPE_MAPPINGS, ERROR_TYPE_MAPPINGS,
@@ -68,7 +66,7 @@ class GeminiNormalizer(FormatNormalizer):
supports_images=True, supports_images=True,
) )
_FINISH_REASON_TO_STOP: Dict[str, StopReason] = { _FINISH_REASON_TO_STOP: dict[str, StopReason] = {
"STOP": StopReason.END_TURN, "STOP": StopReason.END_TURN,
"MAX_TOKENS": StopReason.MAX_TOKENS, "MAX_TOKENS": StopReason.MAX_TOKENS,
"SAFETY": StopReason.CONTENT_FILTERED, "SAFETY": StopReason.CONTENT_FILTERED,
@@ -77,7 +75,7 @@ class GeminiNormalizer(FormatNormalizer):
"OTHER": StopReason.UNKNOWN, "OTHER": StopReason.UNKNOWN,
} }
_ERROR_TYPE_TO_GEMINI_STATUS: Dict[ErrorType, str] = { _ERROR_TYPE_TO_GEMINI_STATUS: dict[ErrorType, str] = {
ErrorType.INVALID_REQUEST: "INVALID_ARGUMENT", ErrorType.INVALID_REQUEST: "INVALID_ARGUMENT",
ErrorType.AUTHENTICATION: "UNAUTHENTICATED", ErrorType.AUTHENTICATION: "UNAUTHENTICATED",
ErrorType.PERMISSION_DENIED: "PERMISSION_DENIED", ErrorType.PERMISSION_DENIED: "PERMISSION_DENIED",
@@ -94,11 +92,11 @@ class GeminiNormalizer(FormatNormalizer):
# Requests # Requests
# ========================= # =========================
def request_to_internal(self, request: Dict[str, Any]) -> InternalRequest: def request_to_internal(self, request: dict[str, Any]) -> InternalRequest:
model = str(request.get("model") or "") model = str(request.get("model") or "")
dropped: Dict[str, int] = {} dropped: dict[str, int] = {}
instructions: List[InstructionSegment] = [] instructions: list[InstructionSegment] = []
system_text, sys_dropped = self._collapse_system_instruction( system_text, sys_dropped = self._collapse_system_instruction(
request.get("system_instruction") request.get("system_instruction")
if "system_instruction" in request if "system_instruction" in request
@@ -108,7 +106,7 @@ class GeminiNormalizer(FormatNormalizer):
if system_text: if system_text:
instructions.append(InstructionSegment(role=Role.SYSTEM, text=system_text)) instructions.append(InstructionSegment(role=Role.SYSTEM, text=system_text))
messages: List[InternalMessage] = [] messages: list[InternalMessage] = []
contents = request.get("contents") or [] contents = request.get("contents") or []
if isinstance(contents, list): if isinstance(contents, list):
for content in contents: for content in contents:
@@ -152,7 +150,7 @@ class GeminiNormalizer(FormatNormalizer):
) )
# 构建 extra保留原始 gemini 字段 # 构建 extra保留原始 gemini 字段
extra: Dict[str, Any] = {"gemini": self._extract_extra(request, {"contents"})} extra: dict[str, Any] = {"gemini": self._extract_extra(request, {"contents"})}
# 保留 generationConfig 中的特殊字段responseModalities, thinkingConfig 等) # 保留 generationConfig 中的特殊字段responseModalities, thinkingConfig 等)
# 这些字段在 _get_generation_config 中已提取,需要单独存储以便转换时使用 # 这些字段在 _get_generation_config 中已提取,需要单独存储以便转换时使用
@@ -160,7 +158,7 @@ class GeminiNormalizer(FormatNormalizer):
response_modalities = generation_config.get("response_modalities") response_modalities = generation_config.get("response_modalities")
thinking_config = generation_config.get("thinking_config") thinking_config = generation_config.get("thinking_config")
if response_modalities or thinking_config: if response_modalities or thinking_config:
google_extra: Dict[str, Any] = {} google_extra: dict[str, Any] = {}
if response_modalities: if response_modalities:
google_extra["response_modalities"] = response_modalities google_extra["response_modalities"] = response_modalities
if thinking_config: if thinking_config:
@@ -188,7 +186,7 @@ class GeminiNormalizer(FormatNormalizer):
return internal return internal
def request_from_internal(self, internal: InternalRequest) -> Dict[str, Any]: def request_from_internal(self, internal: InternalRequest) -> dict[str, Any]:
system_text = internal.system or self._join_instructions(internal.instructions) system_text = internal.system or self._join_instructions(internal.instructions)
# tools/tool_choice # tools/tool_choice
@@ -212,7 +210,7 @@ class GeminiNormalizer(FormatNormalizer):
if internal.tool_choice: if internal.tool_choice:
tool_config = self._tool_choice_to_gemini_tool_config(internal.tool_choice) tool_config = self._tool_choice_to_gemini_tool_config(internal.tool_choice)
generation_config: Dict[str, Any] = {} generation_config: dict[str, Any] = {}
if internal.max_tokens is not None: if internal.max_tokens is not None:
generation_config["max_output_tokens"] = internal.max_tokens generation_config["max_output_tokens"] = internal.max_tokens
if internal.temperature is not None: if internal.temperature is not None:
@@ -231,7 +229,7 @@ class GeminiNormalizer(FormatNormalizer):
thinking_config = google_extra.get("thinking_config") thinking_config = google_extra.get("thinking_config")
if isinstance(thinking_config, dict): if isinstance(thinking_config, dict):
# snake_case -> camelCase 转换 # snake_case -> camelCase 转换
gemini_thinking: Dict[str, Any] = {} gemini_thinking: dict[str, Any] = {}
if "thinking_budget" in thinking_config: if "thinking_budget" in thinking_config:
gemini_thinking["thinkingBudget"] = thinking_config["thinking_budget"] gemini_thinking["thinkingBudget"] = thinking_config["thinking_budget"]
if "include_thoughts" in thinking_config: if "include_thoughts" in thinking_config:
@@ -265,11 +263,11 @@ class GeminiNormalizer(FormatNormalizer):
if "thinking_config" in orig_gc and "thinkingConfig" not in generation_config: if "thinking_config" in orig_gc and "thinkingConfig" not in generation_config:
generation_config["thinkingConfig"] = orig_gc["thinking_config"] generation_config["thinkingConfig"] = orig_gc["thinking_config"]
contents: List[Dict[str, Any]] = [] contents: list[dict[str, Any]] = []
for msg in internal.messages: for msg in internal.messages:
contents.append(self._internal_message_to_content(msg)) contents.append(self._internal_message_to_content(msg))
result: Dict[str, Any] = { result: dict[str, Any] = {
"contents": contents, "contents": contents,
} }
@@ -295,7 +293,7 @@ class GeminiNormalizer(FormatNormalizer):
# Responses # Responses
# ========================= # =========================
def response_to_internal(self, response: Dict[str, Any]) -> InternalResponse: def response_to_internal(self, response: dict[str, Any]) -> InternalResponse:
rid = str(response.get("id") or "") rid = str(response.get("id") or "")
model = str(response.get("modelVersion") or response.get("model") or "") model = str(response.get("modelVersion") or response.get("model") or "")
@@ -315,7 +313,7 @@ class GeminiNormalizer(FormatNormalizer):
usage_info = self._usage_metadata_to_internal(response.get("usageMetadata")) usage_info = self._usage_metadata_to_internal(response.get("usageMetadata"))
extra: Dict[str, Any] = {} extra: dict[str, Any] = {}
if finish_reason is not None: if finish_reason is not None:
extra.setdefault("raw", {})["finishReason"] = finish_reason extra.setdefault("raw", {})["finishReason"] = finish_reason
@@ -337,9 +335,9 @@ class GeminiNormalizer(FormatNormalizer):
self, self,
internal: InternalResponse, internal: InternalResponse,
*, *,
requested_model: Optional[str] = None, requested_model: str | None = None,
) -> Dict[str, Any]: ) -> dict[str, Any]:
parts: List[Dict[str, Any]] = [] parts: list[dict[str, Any]] = []
for b in internal.content: for b in internal.content:
if isinstance(b, TextBlock): if isinstance(b, TextBlock):
if b.text: if b.text:
@@ -374,7 +372,7 @@ class GeminiNormalizer(FormatNormalizer):
if internal.stop_reason is not None: if internal.stop_reason is not None:
finish_reason = STOP_REASON_MAPPINGS.get("GEMINI", {}).get(internal.stop_reason.value, "OTHER") finish_reason = STOP_REASON_MAPPINGS.get("GEMINI", {}).get(internal.stop_reason.value, "OTHER")
usage_metadata: Dict[str, Any] = {} usage_metadata: dict[str, Any] = {}
if internal.usage: if internal.usage:
usage_metadata = { usage_metadata = {
"promptTokenCount": int(internal.usage.input_tokens), "promptTokenCount": int(internal.usage.input_tokens),
@@ -384,7 +382,7 @@ class GeminiNormalizer(FormatNormalizer):
if internal.usage.cache_read_tokens: if internal.usage.cache_read_tokens:
usage_metadata["cachedContentTokenCount"] = int(internal.usage.cache_read_tokens) usage_metadata["cachedContentTokenCount"] = int(internal.usage.cache_read_tokens)
candidate: Dict[str, Any] = { candidate: dict[str, Any] = {
"content": {"parts": parts, "role": "model"}, "content": {"parts": parts, "role": "model"},
"index": 0, "index": 0,
} }
@@ -394,7 +392,7 @@ class GeminiNormalizer(FormatNormalizer):
# 优先使用用户请求的原始模型名,回退到上游返回的模型名 # 优先使用用户请求的原始模型名,回退到上游返回的模型名
model_name = requested_model if requested_model else (internal.model or "gemini") model_name = requested_model if requested_model else (internal.model or "gemini")
out: Dict[str, Any] = { out: dict[str, Any] = {
"candidates": [candidate], "candidates": [candidate],
"modelVersion": model_name, "modelVersion": model_name,
} }
@@ -412,9 +410,9 @@ class GeminiNormalizer(FormatNormalizer):
# Streaming # Streaming
# ========================= # =========================
def stream_chunk_to_internal(self, chunk: Dict[str, Any], state: StreamState) -> List[InternalStreamEvent]: def stream_chunk_to_internal(self, chunk: dict[str, Any], state: StreamState) -> list[InternalStreamEvent]:
ss = state.substate(self.FORMAT_ID) ss = state.substate(self.FORMAT_ID)
events: List[InternalStreamEvent] = [] events: list[InternalStreamEvent] = []
if not ss.get("message_started"): if not ss.get("message_started"):
# 保留初始化时设置的 model客户端请求的模型仅在空时用上游值 # 保留初始化时设置的 model客户端请求的模型仅在空时用上游值
@@ -544,11 +542,11 @@ class GeminiNormalizer(FormatNormalizer):
self, self,
event: InternalStreamEvent, event: InternalStreamEvent,
state: StreamState, state: StreamState,
) -> List[Dict[str, Any]]: ) -> list[dict[str, Any]]:
ss = state.substate(self.FORMAT_ID) ss = state.substate(self.FORMAT_ID)
out: List[Dict[str, Any]] = [] out: list[dict[str, Any]] = []
def base_chunk(parts: List[Dict[str, Any]]) -> Dict[str, Any]: def base_chunk(parts: list[dict[str, Any]]) -> dict[str, Any]:
return { return {
"candidates": [ "candidates": [
{ {
@@ -622,7 +620,7 @@ class GeminiNormalizer(FormatNormalizer):
name = str(entry.get("name") or "") name = str(entry.get("name") or "")
raw_json = str(entry.get("json") or "") raw_json = str(entry.get("json") or "")
args: Dict[str, Any] = {} args: dict[str, Any] = {}
if raw_json: if raw_json:
try: try:
parsed = json.loads(raw_json) parsed = json.loads(raw_json)
@@ -639,7 +637,7 @@ class GeminiNormalizer(FormatNormalizer):
if event.stop_reason is not None: if event.stop_reason is not None:
finish_reason = STOP_REASON_MAPPINGS.get("GEMINI", {}).get(event.stop_reason.value, "OTHER") finish_reason = STOP_REASON_MAPPINGS.get("GEMINI", {}).get(event.stop_reason.value, "OTHER")
chunk: Dict[str, Any] = base_chunk([]) chunk: dict[str, Any] = base_chunk([])
if finish_reason is not None: if finish_reason is not None:
chunk["candidates"][0]["finishReason"] = finish_reason chunk["candidates"][0]["finishReason"] = finish_reason
@@ -665,10 +663,10 @@ class GeminiNormalizer(FormatNormalizer):
# Error conversion # Error conversion
# ========================= # =========================
def is_error_response(self, response: Dict[str, Any]) -> bool: def is_error_response(self, response: dict[str, Any]) -> bool:
return isinstance(response, dict) and "error" in response return isinstance(response, dict) and "error" in response
def error_to_internal(self, error_response: Dict[str, Any]) -> InternalError: def error_to_internal(self, error_response: dict[str, Any]) -> InternalError:
err = error_response.get("error") if isinstance(error_response, dict) else None err = error_response.get("error") if isinstance(error_response, dict) else None
err = err if isinstance(err, dict) else {} err = err if isinstance(err, dict) else {}
@@ -691,9 +689,9 @@ class GeminiNormalizer(FormatNormalizer):
extra={"gemini": {"error": err}, "raw": {"status": raw_status}}, extra={"gemini": {"error": err}, "raw": {"status": raw_status}},
) )
def error_from_internal(self, internal: InternalError) -> Dict[str, Any]: def error_from_internal(self, internal: InternalError) -> dict[str, Any]:
status = self._ERROR_TYPE_TO_GEMINI_STATUS.get(internal.type, "INTERNAL") status = self._ERROR_TYPE_TO_GEMINI_STATUS.get(internal.type, "INTERNAL")
payload: Dict[str, Any] = { payload: dict[str, Any] = {
"code": 400 if internal.type == ErrorType.INVALID_REQUEST else 500, "code": 400 if internal.type == ErrorType.INVALID_REQUEST else 500,
"message": internal.message, "message": internal.message,
"status": status, "status": status,
@@ -704,8 +702,8 @@ class GeminiNormalizer(FormatNormalizer):
# Helpers # Helpers
# ========================= # =========================
def _content_to_internal_message(self, content: Dict[str, Any]) -> Tuple[Optional[InternalMessage], Dict[str, int]]: def _content_to_internal_message(self, content: dict[str, Any]) -> tuple[InternalMessage | None, dict[str, int]]:
dropped: Dict[str, int] = {} dropped: dict[str, int] = {}
role_raw = str(content.get("role") or "user") role_raw = str(content.get("role") or "user")
if role_raw == "model": if role_raw == "model":
@@ -727,15 +725,15 @@ class GeminiNormalizer(FormatNormalizer):
dropped, dropped,
) )
def _parts_to_blocks(self, parts: Any) -> Tuple[List[ContentBlock], Dict[str, int]]: def _parts_to_blocks(self, parts: Any) -> tuple[list[ContentBlock], dict[str, int]]:
dropped: Dict[str, int] = {} dropped: dict[str, int] = {}
if parts is None: if parts is None:
return [], dropped return [], dropped
if not isinstance(parts, list): if not isinstance(parts, list):
dropped["gemini_parts_non_list"] = dropped.get("gemini_parts_non_list", 0) + 1 dropped["gemini_parts_non_list"] = dropped.get("gemini_parts_non_list", 0) + 1
return [], dropped return [], dropped
blocks: List[ContentBlock] = [] blocks: list[ContentBlock] = []
for part in parts: for part in parts:
if not isinstance(part, dict): if not isinstance(part, dict):
dropped["gemini_part_non_dict"] = dropped.get("gemini_part_non_dict", 0) + 1 dropped["gemini_part_non_dict"] = dropped.get("gemini_part_non_dict", 0) + 1
@@ -785,7 +783,7 @@ class GeminiNormalizer(FormatNormalizer):
name = str(func_resp.get("name") or "") name = str(func_resp.get("name") or "")
response = func_resp.get("response") response = func_resp.get("response")
output: Any = None output: Any = None
content_text: Optional[str] = None content_text: str | None = None
# 兼容历史response 常见结构为 {"result": ...} # 兼容历史response 常见结构为 {"result": ...}
if isinstance(response, dict) and "result" in response: if isinstance(response, dict) and "result" in response:
@@ -815,10 +813,10 @@ class GeminiNormalizer(FormatNormalizer):
return blocks, dropped return blocks, dropped
def _internal_message_to_content(self, msg: InternalMessage) -> Dict[str, Any]: def _internal_message_to_content(self, msg: InternalMessage) -> dict[str, Any]:
role = "model" if msg.role == Role.ASSISTANT else "user" role = "model" if msg.role == Role.ASSISTANT else "user"
parts: List[Dict[str, Any]] = [] parts: list[dict[str, Any]] = []
for b in msg.content: for b in msg.content:
if isinstance(b, UnknownBlock): if isinstance(b, UnknownBlock):
continue continue
@@ -861,8 +859,8 @@ class GeminiNormalizer(FormatNormalizer):
return {"role": role, "parts": parts} return {"role": role, "parts": parts}
def _collapse_system_instruction(self, system_instruction: Any) -> Tuple[Optional[str], Dict[str, int]]: def _collapse_system_instruction(self, system_instruction: Any) -> tuple[str | None, dict[str, int]]:
dropped: Dict[str, int] = {} dropped: dict[str, int] = {}
if system_instruction is None: if system_instruction is None:
return None, dropped return None, dropped
@@ -870,7 +868,7 @@ class GeminiNormalizer(FormatNormalizer):
if isinstance(system_instruction, dict): if isinstance(system_instruction, dict):
parts = system_instruction.get("parts") parts = system_instruction.get("parts")
if isinstance(parts, list): if isinstance(parts, list):
texts: List[str] = [] texts: list[str] = []
for part in parts: for part in parts:
if isinstance(part, dict) and "text" in part and part.get("text"): if isinstance(part, dict) and "text" in part and part.get("text"):
texts.append(str(part.get("text"))) texts.append(str(part.get("text")))
@@ -880,7 +878,7 @@ class GeminiNormalizer(FormatNormalizer):
dropped["gemini_system_instruction_unsupported"] = dropped.get("gemini_system_instruction_unsupported", 0) + 1 dropped["gemini_system_instruction_unsupported"] = dropped.get("gemini_system_instruction_unsupported", 0) + 1
return None, dropped return None, dropped
def _get_generation_config(self, request: Dict[str, Any]) -> Dict[str, Any]: def _get_generation_config(self, request: dict[str, Any]) -> dict[str, Any]:
# 兼容 snake_case 与 camelCase # 兼容 snake_case 与 camelCase
gc = request.get("generation_config") if "generation_config" in request else request.get("generationConfig") gc = request.get("generation_config") if "generation_config" in request else request.get("generationConfig")
if not isinstance(gc, dict): if not isinstance(gc, dict):
@@ -893,7 +891,7 @@ class GeminiNormalizer(FormatNormalizer):
return gc.get(k) return gc.get(k)
return None return None
normalized: Dict[str, Any] = {} normalized: dict[str, Any] = {}
normalized["max_output_tokens"] = pick("max_output_tokens", "maxOutputTokens") normalized["max_output_tokens"] = pick("max_output_tokens", "maxOutputTokens")
normalized["temperature"] = pick("temperature") normalized["temperature"] = pick("temperature")
normalized["top_p"] = pick("top_p", "topP") normalized["top_p"] = pick("top_p", "topP")
@@ -912,11 +910,11 @@ class GeminiNormalizer(FormatNormalizer):
return {k: v for k, v in normalized.items() if v is not None} return {k: v for k, v in normalized.items() if v is not None}
def _gemini_tools_to_internal(self, tools: Any) -> Optional[List[ToolDefinition]]: def _gemini_tools_to_internal(self, tools: Any) -> list[ToolDefinition] | None:
if not tools or not isinstance(tools, list): if not tools or not isinstance(tools, list):
return None return None
out: List[ToolDefinition] = [] out: list[ToolDefinition] = []
for tool in tools: for tool in tools:
if not isinstance(tool, dict): if not isinstance(tool, dict):
continue continue
@@ -945,7 +943,7 @@ class GeminiNormalizer(FormatNormalizer):
return out or None return out or None
def _gemini_tool_config_to_tool_choice(self, tool_config: Any) -> Optional[ToolChoice]: def _gemini_tool_config_to_tool_choice(self, tool_config: Any) -> ToolChoice | None:
if tool_config is None: if tool_config is None:
return None return None
if not isinstance(tool_config, dict): if not isinstance(tool_config, dict):
@@ -972,9 +970,9 @@ class GeminiNormalizer(FormatNormalizer):
return ToolChoice(type=ToolChoiceType.AUTO, extra={"gemini": tool_config}) return ToolChoice(type=ToolChoiceType.AUTO, extra={"gemini": tool_config})
def _tool_choice_to_gemini_tool_config(self, tool_choice: ToolChoice) -> Dict[str, Any]: def _tool_choice_to_gemini_tool_config(self, tool_choice: ToolChoice) -> dict[str, Any]:
mode = "AUTO" mode = "AUTO"
cfg: Dict[str, Any] = {} cfg: dict[str, Any] = {}
if tool_choice.type == ToolChoiceType.NONE: if tool_choice.type == ToolChoiceType.NONE:
mode = "NONE" mode = "NONE"
@@ -987,12 +985,12 @@ class GeminiNormalizer(FormatNormalizer):
cfg["mode"] = mode cfg["mode"] = mode
return {"function_calling_config": cfg} return {"function_calling_config": cfg}
def _usage_metadata_to_internal(self, usage_metadata: Any) -> Optional[UsageInfo]: def _usage_metadata_to_internal(self, usage_metadata: Any) -> UsageInfo | None:
if not isinstance(usage_metadata, dict): if not isinstance(usage_metadata, dict):
return None return None
mapping = USAGE_FIELD_MAPPINGS.get("GEMINI", {}) mapping = USAGE_FIELD_MAPPINGS.get("GEMINI", {})
fields: Dict[str, int] = {} fields: dict[str, int] = {}
extra = self._extract_extra(usage_metadata, set(mapping.keys())) extra = self._extract_extra(usage_metadata, set(mapping.keys()))
# promptTokenCount/candidatesTokenCount/totalTokenCount/cachedContentTokenCount # promptTokenCount/candidatesTokenCount/totalTokenCount/cachedContentTokenCount
@@ -1023,7 +1021,7 @@ class GeminiNormalizer(FormatNormalizer):
extra={"gemini": extra} if extra else {}, extra={"gemini": extra} if extra else {},
) )
def _join_instructions(self, instructions: List[InstructionSegment]) -> Optional[str]: def _join_instructions(self, instructions: list[InstructionSegment]) -> str | None:
parts = [seg.text for seg in instructions if seg.text] parts = [seg.text for seg in instructions if seg.text]
joined = "\n\n".join(parts) joined = "\n\n".join(parts)
return joined or None return joined or None
@@ -1034,7 +1032,7 @@ class GeminiNormalizer(FormatNormalizer):
except ValueError: except ValueError:
return ErrorType.UNKNOWN return ErrorType.UNKNOWN
def _optional_int(self, value: Any) -> Optional[int]: def _optional_int(self, value: Any) -> int | None:
if value is None: if value is None:
return None return None
try: try:
@@ -1042,7 +1040,7 @@ class GeminiNormalizer(FormatNormalizer):
except (TypeError, ValueError): except (TypeError, ValueError):
return None return None
def _optional_float(self, value: Any) -> Optional[float]: def _optional_float(self, value: Any) -> float | None:
if value is None: if value is None:
return None return None
try: try:
@@ -1050,7 +1048,7 @@ class GeminiNormalizer(FormatNormalizer):
except (TypeError, ValueError): except (TypeError, ValueError):
return None return None
def _coerce_str_list(self, value: Any) -> Optional[List[str]]: def _coerce_str_list(self, value: Any) -> list[str] | None:
if value is None: if value is None:
return None return None
if isinstance(value, str): if isinstance(value, str):
@@ -1059,10 +1057,10 @@ class GeminiNormalizer(FormatNormalizer):
return [str(x) for x in value if x is not None] return [str(x) for x in value if x is not None]
return None return None
def _extract_extra(self, payload: Dict[str, Any], known_keys: set[str]) -> Dict[str, Any]: def _extract_extra(self, payload: dict[str, Any], known_keys: set[str]) -> dict[str, Any]:
return {k: v for k, v in payload.items() if k not in known_keys} return {k: v for k, v in payload.items() if k not in known_keys}
def _merge_dropped(self, target: Dict[str, int], source: Dict[str, int]) -> None: def _merge_dropped(self, target: dict[str, int], source: dict[str, int]) -> None:
for k, v in source.items(): for k, v in source.items():
target[k] = target.get(k, 0) + int(v) target[k] = target.get(k, 0) + int(v)

View File

@@ -7,7 +7,6 @@ GEMINI_CLI 的请求/响应 body 与 GEMINI 一致Google Gemini API
如需 CLI 特殊处理,可覆盖 request_from_internal / request_to_internal 等方法。 如需 CLI 特殊处理,可覆盖 request_from_internal / request_to_internal 等方法。
""" """
from __future__ import annotations
from src.core.api_format.conversion.normalizers.gemini import GeminiNormalizer from src.core.api_format.conversion.normalizers.gemini import GeminiNormalizer

View File

@@ -7,11 +7,10 @@ OpenAI Chat Completions Normalizer
- 可选OpenAI error <-> InternalError - 可选OpenAI error <-> InternalError
""" """
from __future__ import annotations
import json import json
import time import time
from typing import Any, Dict, List, Optional, Tuple, Union from typing import Any
from src.core.logger import logger from src.core.logger import logger
from src.core.api_format.conversion.field_mappings import ( from src.core.api_format.conversion.field_mappings import (
@@ -49,7 +48,6 @@ from src.core.api_format.conversion.stream_events import (
InternalStreamEvent, InternalStreamEvent,
MessageStartEvent, MessageStartEvent,
MessageStopEvent, MessageStopEvent,
StreamEventType,
ToolCallDeltaEvent, ToolCallDeltaEvent,
) )
from src.core.api_format.conversion.stream_state import StreamState from src.core.api_format.conversion.stream_state import StreamState
@@ -65,7 +63,7 @@ class OpenAINormalizer(FormatNormalizer):
) )
# finish_reason -> StopReason # finish_reason -> StopReason
_FINISH_REASON_TO_STOP: Dict[str, StopReason] = { _FINISH_REASON_TO_STOP: dict[str, StopReason] = {
"stop": StopReason.END_TURN, "stop": StopReason.END_TURN,
"length": StopReason.MAX_TOKENS, "length": StopReason.MAX_TOKENS,
"tool_calls": StopReason.TOOL_USE, "tool_calls": StopReason.TOOL_USE,
@@ -74,7 +72,7 @@ class OpenAINormalizer(FormatNormalizer):
} }
# StopReason -> finish_reason # StopReason -> finish_reason
_STOP_TO_FINISH_REASON: Dict[StopReason, str] = { _STOP_TO_FINISH_REASON: dict[StopReason, str] = {
StopReason.END_TURN: "stop", StopReason.END_TURN: "stop",
StopReason.MAX_TOKENS: "length", StopReason.MAX_TOKENS: "length",
StopReason.STOP_SEQUENCE: "stop", StopReason.STOP_SEQUENCE: "stop",
@@ -84,7 +82,7 @@ class OpenAINormalizer(FormatNormalizer):
} }
# InternalError.type -> OpenAI error.type最佳努力 # InternalError.type -> OpenAI error.type最佳努力
_ERROR_TYPE_TO_OPENAI: Dict[ErrorType, str] = { _ERROR_TYPE_TO_OPENAI: dict[ErrorType, str] = {
ErrorType.INVALID_REQUEST: "invalid_request_error", ErrorType.INVALID_REQUEST: "invalid_request_error",
ErrorType.AUTHENTICATION: "invalid_api_key", ErrorType.AUTHENTICATION: "invalid_api_key",
ErrorType.PERMISSION_DENIED: "invalid_request_error", ErrorType.PERMISSION_DENIED: "invalid_request_error",
@@ -101,13 +99,13 @@ class OpenAINormalizer(FormatNormalizer):
# Requests # Requests
# ========================= # =========================
def request_to_internal(self, request: Dict[str, Any]) -> InternalRequest: def request_to_internal(self, request: dict[str, Any]) -> InternalRequest:
model = str(request.get("model") or "") model = str(request.get("model") or "")
dropped: Dict[str, int] = {} dropped: dict[str, int] = {}
instructions: List[InstructionSegment] = [] instructions: list[InstructionSegment] = []
messages: List[InternalMessage] = [] messages: list[InternalMessage] = []
for msg in request.get("messages") or []: for msg in request.get("messages") or []:
if not isinstance(msg, dict): if not isinstance(msg, dict):
@@ -146,7 +144,7 @@ class OpenAINormalizer(FormatNormalizer):
) )
# 构建 extra保留未识别字段 # 构建 extra保留未识别字段
extra: Dict[str, Any] = {"openai": self._extract_extra(request, {"messages"})} extra: dict[str, Any] = {"openai": self._extract_extra(request, {"messages"})}
# 处理 extra_body.google (用于 Gemini 特定功能透传,如 thinkingConfig, responseModalities) # 处理 extra_body.google (用于 Gemini 特定功能透传,如 thinkingConfig, responseModalities)
extra_body = request.get("extra_body") extra_body = request.get("extra_body")
@@ -175,8 +173,8 @@ class OpenAINormalizer(FormatNormalizer):
return internal return internal
def request_from_internal(self, internal: InternalRequest) -> Dict[str, Any]: def request_from_internal(self, internal: InternalRequest) -> dict[str, Any]:
out_messages: List[Dict[str, Any]] = [] out_messages: list[dict[str, Any]] = []
if internal.instructions: if internal.instructions:
for seg in internal.instructions: for seg in internal.instructions:
@@ -189,7 +187,7 @@ class OpenAINormalizer(FormatNormalizer):
for msg in internal.messages: for msg in internal.messages:
out_messages.extend(self._internal_message_to_openai_messages(msg)) out_messages.extend(self._internal_message_to_openai_messages(msg))
result: Dict[str, Any] = { result: dict[str, Any] = {
"model": internal.model, "model": internal.model,
"messages": out_messages, "messages": out_messages,
} }
@@ -233,11 +231,11 @@ class OpenAINormalizer(FormatNormalizer):
# Responses # Responses
# ========================= # =========================
def response_to_internal(self, response: Dict[str, Any]) -> InternalResponse: def response_to_internal(self, response: dict[str, Any]) -> InternalResponse:
rid = str(response.get("id") or "") rid = str(response.get("id") or "")
model = str(response.get("model") or "") model = str(response.get("model") or "")
extra: Dict[str, Any] = {} extra: dict[str, Any] = {}
choices = response.get("choices") or [] choices = response.get("choices") or []
if isinstance(choices, list) and len(choices) > 1: if isinstance(choices, list) and len(choices) > 1:
@@ -285,12 +283,12 @@ class OpenAINormalizer(FormatNormalizer):
self, self,
internal: InternalResponse, internal: InternalResponse,
*, *,
requested_model: Optional[str] = None, requested_model: str | None = None,
) -> Dict[str, Any]: ) -> dict[str, Any]:
# OpenAI Chat Completions response envelope # OpenAI Chat Completions response envelope
# 优先使用用户请求的原始模型名,回退到上游返回的模型名 # 优先使用用户请求的原始模型名,回退到上游返回的模型名
model_name = requested_model if requested_model else internal.model model_name = requested_model if requested_model else internal.model
out: Dict[str, Any] = { out: dict[str, Any] = {
"id": internal.id or "chatcmpl-unknown", "id": internal.id or "chatcmpl-unknown",
"object": "chat.completion", "object": "chat.completion",
"created": int(time.time()), "created": int(time.time()),
@@ -298,7 +296,7 @@ class OpenAINormalizer(FormatNormalizer):
"choices": [], "choices": [],
} }
message: Dict[str, Any] = {"role": "assistant"} message: dict[str, Any] = {"role": "assistant"}
content_blocks, tool_blocks = self._split_blocks(internal.content) content_blocks, tool_blocks = self._split_blocks(internal.content)
content_value = self._blocks_to_openai_content(content_blocks) content_value = self._blocks_to_openai_content(content_blocks)
@@ -336,9 +334,9 @@ class OpenAINormalizer(FormatNormalizer):
# Streaming # Streaming
# ========================= # =========================
def stream_chunk_to_internal(self, chunk: Dict[str, Any], state: StreamState) -> List[InternalStreamEvent]: def stream_chunk_to_internal(self, chunk: dict[str, Any], state: StreamState) -> list[InternalStreamEvent]:
ss = state.substate(self.FORMAT_ID) ss = state.substate(self.FORMAT_ID)
events: List[InternalStreamEvent] = [] events: list[InternalStreamEvent] = []
# OpenAI streaming error通常是单个 {"error": {...}} # OpenAI streaming error通常是单个 {"error": {...}}
if isinstance(chunk, dict) and "error" in chunk: if isinstance(chunk, dict) and "error" in chunk:
@@ -437,11 +435,11 @@ class OpenAINormalizer(FormatNormalizer):
self, self,
event: InternalStreamEvent, event: InternalStreamEvent,
state: StreamState, state: StreamState,
) -> List[Dict[str, Any]]: ) -> list[dict[str, Any]]:
ss = state.substate(self.FORMAT_ID) ss = state.substate(self.FORMAT_ID)
out: List[Dict[str, Any]] = [] out: list[dict[str, Any]] = []
def base_chunk(delta: Dict[str, Any], finish_reason: Optional[str] = None) -> Dict[str, Any]: def base_chunk(delta: dict[str, Any], finish_reason: str | None = None) -> dict[str, Any]:
return { return {
"id": state.message_id or "chatcmpl-stream", "id": state.message_id or "chatcmpl-stream",
"object": "chat.completion.chunk", "object": "chat.completion.chunk",
@@ -570,10 +568,10 @@ class OpenAINormalizer(FormatNormalizer):
# Error conversion # Error conversion
# ========================= # =========================
def is_error_response(self, response: Dict[str, Any]) -> bool: def is_error_response(self, response: dict[str, Any]) -> bool:
return isinstance(response, dict) and "error" in response return isinstance(response, dict) and "error" in response
def error_to_internal(self, error_response: Dict[str, Any]) -> InternalError: def error_to_internal(self, error_response: dict[str, Any]) -> InternalError:
err = error_response.get("error") if isinstance(error_response, dict) else None err = error_response.get("error") if isinstance(error_response, dict) else None
err = err if isinstance(err, dict) else {} err = err if isinstance(err, dict) else {}
@@ -592,9 +590,9 @@ class OpenAINormalizer(FormatNormalizer):
extra={"openai": {"error": err}, "raw": {"type": raw_type}}, extra={"openai": {"error": err}, "raw": {"type": raw_type}},
) )
def error_from_internal(self, internal: InternalError) -> Dict[str, Any]: def error_from_internal(self, internal: InternalError) -> dict[str, Any]:
type_str = self._ERROR_TYPE_TO_OPENAI.get(internal.type, "server_error") type_str = self._ERROR_TYPE_TO_OPENAI.get(internal.type, "server_error")
payload: Dict[str, Any] = { payload: dict[str, Any] = {
"message": internal.message, "message": internal.message,
"type": type_str, "type": type_str,
} }
@@ -608,8 +606,8 @@ class OpenAINormalizer(FormatNormalizer):
# Helpers # Helpers
# ========================= # =========================
def _openai_message_to_internal(self, msg: Dict[str, Any]) -> Tuple[Optional[InternalMessage], Dict[str, int]]: def _openai_message_to_internal(self, msg: dict[str, Any]) -> tuple[InternalMessage | None, dict[str, int]]:
dropped: Dict[str, int] = {} dropped: dict[str, int] = {}
role_raw = str(msg.get("role") or "unknown") role_raw = str(msg.get("role") or "unknown")
role = self._role_from_openai(role_raw) role = self._role_from_openai(role_raw)
@@ -656,8 +654,8 @@ class OpenAINormalizer(FormatNormalizer):
dropped, dropped,
) )
def _openai_content_to_blocks(self, content: Any) -> Tuple[List[ContentBlock], Dict[str, int]]: def _openai_content_to_blocks(self, content: Any) -> tuple[list[ContentBlock], dict[str, int]]:
dropped: Dict[str, int] = {} dropped: dict[str, int] = {}
if content is None: if content is None:
return [], dropped return [], dropped
@@ -667,7 +665,7 @@ class OpenAINormalizer(FormatNormalizer):
dropped["openai_content_non_list"] = dropped.get("openai_content_non_list", 0) + 1 dropped["openai_content_non_list"] = dropped.get("openai_content_non_list", 0) + 1
return [], dropped return [], dropped
blocks: List[ContentBlock] = [] blocks: list[ContentBlock] = []
for part in content: for part in content:
if not isinstance(part, dict): if not isinstance(part, dict):
dropped["openai_content_part_non_dict"] = dropped.get("openai_content_part_non_dict", 0) + 1 dropped["openai_content_part_non_dict"] = dropped.get("openai_content_part_non_dict", 0) + 1
@@ -698,21 +696,21 @@ class OpenAINormalizer(FormatNormalizer):
return blocks, dropped return blocks, dropped
def _collapse_openai_text(self, content: Any) -> Tuple[str, Dict[str, int]]: def _collapse_openai_text(self, content: Any) -> tuple[str, dict[str, int]]:
blocks, dropped = self._openai_content_to_blocks(content) blocks, dropped = self._openai_content_to_blocks(content)
text_parts = [b.text for b in blocks if isinstance(b, TextBlock) and b.text] text_parts = [b.text for b in blocks if isinstance(b, TextBlock) and b.text]
return ("\n\n".join(text_parts), dropped) return ("\n\n".join(text_parts), dropped)
def _join_instructions(self, instructions: List[InstructionSegment]) -> Optional[str]: def _join_instructions(self, instructions: list[InstructionSegment]) -> str | None:
parts = [seg.text for seg in instructions if seg.text] parts = [seg.text for seg in instructions if seg.text]
joined = "\n\n".join(parts) joined = "\n\n".join(parts)
return joined or None return joined or None
def _openai_tools_to_internal(self, tools: Any) -> Optional[List[ToolDefinition]]: def _openai_tools_to_internal(self, tools: Any) -> list[ToolDefinition] | None:
if not tools or not isinstance(tools, list): if not tools or not isinstance(tools, list):
return None return None
out: List[ToolDefinition] = [] out: list[ToolDefinition] = []
for tool in tools: for tool in tools:
if not isinstance(tool, dict): if not isinstance(tool, dict):
continue continue
@@ -720,7 +718,7 @@ class OpenAINormalizer(FormatNormalizer):
continue continue
function_raw = tool.get("function") function_raw = tool.get("function")
function: Dict[str, Any] = function_raw if isinstance(function_raw, dict) else {} function: dict[str, Any] = function_raw if isinstance(function_raw, dict) else {}
name = str(function.get("name") or "") name = str(function.get("name") or "")
if not name: if not name:
continue continue
@@ -740,7 +738,7 @@ class OpenAINormalizer(FormatNormalizer):
return out or None return out or None
def _openai_tool_choice_to_internal(self, tool_choice: Any) -> Optional[ToolChoice]: def _openai_tool_choice_to_internal(self, tool_choice: Any) -> ToolChoice | None:
if tool_choice is None: if tool_choice is None:
return None return None
@@ -756,13 +754,13 @@ class OpenAINormalizer(FormatNormalizer):
if isinstance(tool_choice, dict) and tool_choice.get("type") == "function": if isinstance(tool_choice, dict) and tool_choice.get("type") == "function":
fn_raw = tool_choice.get("function") fn_raw = tool_choice.get("function")
fn: Dict[str, Any] = fn_raw if isinstance(fn_raw, dict) else {} fn: dict[str, Any] = fn_raw if isinstance(fn_raw, dict) else {}
name = str(fn.get("name") or "") name = str(fn.get("name") or "")
return ToolChoice(type=ToolChoiceType.TOOL, tool_name=name, extra={"openai": tool_choice}) return ToolChoice(type=ToolChoiceType.TOOL, tool_name=name, extra={"openai": tool_choice})
return ToolChoice(type=ToolChoiceType.AUTO, extra={"openai": tool_choice}) return ToolChoice(type=ToolChoiceType.AUTO, extra={"openai": tool_choice})
def _tool_choice_to_openai(self, tool_choice: ToolChoice) -> Union[str, Dict[str, Any]]: def _tool_choice_to_openai(self, tool_choice: ToolChoice) -> str | dict[str, Any]:
if tool_choice.type == ToolChoiceType.NONE: if tool_choice.type == ToolChoiceType.NONE:
return "none" return "none"
if tool_choice.type == ToolChoiceType.AUTO: if tool_choice.type == ToolChoiceType.AUTO:
@@ -773,8 +771,8 @@ class OpenAINormalizer(FormatNormalizer):
return {"type": "function", "function": {"name": tool_choice.tool_name or ""}} return {"type": "function", "function": {"name": tool_choice.tool_name or ""}}
return "auto" return "auto"
def _openai_tool_call_to_block(self, tool_call: Any) -> Tuple[Optional[ToolUseBlock], Dict[str, int]]: def _openai_tool_call_to_block(self, tool_call: Any) -> tuple[ToolUseBlock | None, dict[str, int]]:
dropped: Dict[str, int] = {} dropped: dict[str, int] = {}
if not isinstance(tool_call, dict): if not isinstance(tool_call, dict):
dropped["openai_tool_call_non_dict"] = dropped.get("openai_tool_call_non_dict", 0) + 1 dropped["openai_tool_call_non_dict"] = dropped.get("openai_tool_call_non_dict", 0) + 1
return None, dropped return None, dropped
@@ -785,12 +783,12 @@ class OpenAINormalizer(FormatNormalizer):
return None, dropped return None, dropped
fn_raw = tool_call.get("function") fn_raw = tool_call.get("function")
fn: Dict[str, Any] = fn_raw if isinstance(fn_raw, dict) else {} fn: dict[str, Any] = fn_raw if isinstance(fn_raw, dict) else {}
name = str(fn.get("name") or "") name = str(fn.get("name") or "")
args_str = str(fn.get("arguments") or "") args_str = str(fn.get("arguments") or "")
tool_id = str(tool_call.get("id") or "") tool_id = str(tool_call.get("id") or "")
tool_input: Dict[str, Any] tool_input: dict[str, Any]
if args_str: if args_str:
try: try:
parsed = json.loads(args_str) parsed = json.loads(args_str)
@@ -810,15 +808,15 @@ class OpenAINormalizer(FormatNormalizer):
dropped, dropped,
) )
def _legacy_function_call_to_block(self, func_call: Dict[str, Any]) -> Tuple[Optional[ToolUseBlock], Dict[str, int]]: def _legacy_function_call_to_block(self, func_call: dict[str, Any]) -> tuple[ToolUseBlock | None, dict[str, int]]:
dropped: Dict[str, int] = {} dropped: dict[str, int] = {}
name = str(func_call.get("name") or "") name = str(func_call.get("name") or "")
args_str = str(func_call.get("arguments") or "") args_str = str(func_call.get("arguments") or "")
if not name: if not name:
dropped["openai_function_call_missing_name"] = dropped.get("openai_function_call_missing_name", 0) + 1 dropped["openai_function_call_missing_name"] = dropped.get("openai_function_call_missing_name", 0) + 1
return None, dropped return None, dropped
tool_input: Dict[str, Any] tool_input: dict[str, Any]
if args_str: if args_str:
try: try:
parsed = json.loads(args_str) parsed = json.loads(args_str)
@@ -840,10 +838,10 @@ class OpenAINormalizer(FormatNormalizer):
def _openai_tool_result_message_to_block( def _openai_tool_result_message_to_block(
self, self,
msg: Dict[str, Any], msg: dict[str, Any],
tool_call_id: str, tool_call_id: str,
) -> Tuple[Optional[ToolResultBlock], Dict[str, int]]: ) -> tuple[ToolResultBlock | None, dict[str, int]]:
dropped: Dict[str, int] = {} dropped: dict[str, int] = {}
content = msg.get("content") content = msg.get("content")
if content is None: if content is None:
return ToolResultBlock(tool_use_id=tool_call_id, output=None, content_text=None, extra={"openai": msg}), dropped return ToolResultBlock(tool_use_id=tool_call_id, output=None, content_text=None, extra={"openai": msg}), dropped
@@ -887,7 +885,7 @@ class OpenAINormalizer(FormatNormalizer):
dropped, dropped,
) )
def _openai_usage_to_internal(self, usage: Any) -> Optional[UsageInfo]: def _openai_usage_to_internal(self, usage: Any) -> UsageInfo | None:
if not isinstance(usage, dict): if not isinstance(usage, dict):
return None return None
@@ -908,9 +906,9 @@ class OpenAINormalizer(FormatNormalizer):
extra={"openai": extra} if extra else {}, extra={"openai": extra} if extra else {},
) )
def _blocks_to_openai_content(self, blocks: List[ContentBlock]) -> Optional[Union[str, List[Dict[str, Any]]]]: def _blocks_to_openai_content(self, blocks: list[ContentBlock]) -> str | list[dict[str, Any]] | None:
parts: List[Dict[str, Any]] = [] parts: list[dict[str, Any]] = []
text_parts: List[str] = [] text_parts: list[str] = []
for b in blocks: for b in blocks:
if isinstance(b, TextBlock): if isinstance(b, TextBlock):
@@ -948,9 +946,9 @@ class OpenAINormalizer(FormatNormalizer):
# OpenAI content 可以是空字符串;但作为响应 message.content 通常允许为 ""/None。 # OpenAI content 可以是空字符串;但作为响应 message.content 通常允许为 ""/None。
return "" return ""
def _split_blocks(self, blocks: List[ContentBlock]) -> Tuple[List[ContentBlock], List[ToolUseBlock]]: def _split_blocks(self, blocks: list[ContentBlock]) -> tuple[list[ContentBlock], list[ToolUseBlock]]:
content_blocks: List[ContentBlock] = [] content_blocks: list[ContentBlock] = []
tool_blocks: List[ToolUseBlock] = [] tool_blocks: list[ToolUseBlock] = []
for b in blocks: for b in blocks:
if isinstance(b, ToolUseBlock): if isinstance(b, ToolUseBlock):
tool_blocks.append(b) tool_blocks.append(b)
@@ -963,7 +961,7 @@ class OpenAINormalizer(FormatNormalizer):
content_blocks.append(b) content_blocks.append(b)
return content_blocks, tool_blocks return content_blocks, tool_blocks
def _internal_message_to_openai_messages(self, msg: InternalMessage) -> List[Dict[str, Any]]: def _internal_message_to_openai_messages(self, msg: InternalMessage) -> list[dict[str, Any]]:
if msg.role == Role.USER: if msg.role == Role.USER:
return self._user_message_to_openai(msg) return self._user_message_to_openai(msg)
if msg.role == Role.ASSISTANT: if msg.role == Role.ASSISTANT:
@@ -974,9 +972,9 @@ class OpenAINormalizer(FormatNormalizer):
return [{"role": "tool", "content": content_value or ""}] return [{"role": "tool", "content": content_value or ""}]
return [{"role": "user", "content": ""}] return [{"role": "user", "content": ""}]
def _user_message_to_openai(self, msg: InternalMessage) -> List[Dict[str, Any]]: def _user_message_to_openai(self, msg: InternalMessage) -> list[dict[str, Any]]:
out: List[Dict[str, Any]] = [] out: list[dict[str, Any]] = []
pending: List[ContentBlock] = [] pending: list[ContentBlock] = []
def flush_user() -> None: def flush_user() -> None:
nonlocal pending nonlocal pending
@@ -1008,9 +1006,9 @@ class OpenAINormalizer(FormatNormalizer):
return out return out
def _assistant_message_to_openai(self, msg: InternalMessage) -> Dict[str, Any]: def _assistant_message_to_openai(self, msg: InternalMessage) -> dict[str, Any]:
content_blocks: List[ContentBlock] = [] content_blocks: list[ContentBlock] = []
tool_blocks: List[ToolUseBlock] = [] tool_blocks: list[ToolUseBlock] = []
for b in msg.content: for b in msg.content:
if isinstance(b, ToolUseBlock): if isinstance(b, ToolUseBlock):
@@ -1022,7 +1020,7 @@ class OpenAINormalizer(FormatNormalizer):
continue continue
content_blocks.append(b) content_blocks.append(b)
out: Dict[str, Any] = {"role": "assistant"} out: dict[str, Any] = {"role": "assistant"}
content_value = self._blocks_to_openai_content(content_blocks) content_value = self._blocks_to_openai_content(content_blocks)
out["content"] = content_value if content_value is not None else "" out["content"] = content_value if content_value is not None else ""
@@ -1031,7 +1029,7 @@ class OpenAINormalizer(FormatNormalizer):
return out return out
def _tool_result_block_to_openai_message(self, block: ToolResultBlock) -> Dict[str, Any]: def _tool_result_block_to_openai_message(self, block: ToolResultBlock) -> dict[str, Any]:
content: str content: str
if block.content_text is not None: if block.content_text is not None:
content = block.content_text content = block.content_text
@@ -1048,7 +1046,7 @@ class OpenAINormalizer(FormatNormalizer):
"content": content, "content": content,
} }
def _tool_use_block_to_openai_call(self, block: ToolUseBlock, index: int) -> Dict[str, Any]: def _tool_use_block_to_openai_call(self, block: ToolUseBlock, index: int) -> dict[str, Any]:
return { return {
"index": index, "index": index,
"id": block.tool_id or f"call_{index}", "id": block.tool_id or f"call_{index}",
@@ -1085,7 +1083,7 @@ class OpenAINormalizer(FormatNormalizer):
except ValueError: except ValueError:
return ErrorType.UNKNOWN return ErrorType.UNKNOWN
def _optional_int(self, value: Any) -> Optional[int]: def _optional_int(self, value: Any) -> int | None:
if value is None: if value is None:
return None return None
try: try:
@@ -1093,7 +1091,7 @@ class OpenAINormalizer(FormatNormalizer):
except (TypeError, ValueError): except (TypeError, ValueError):
return None return None
def _optional_float(self, value: Any) -> Optional[float]: def _optional_float(self, value: Any) -> float | None:
if value is None: if value is None:
return None return None
try: try:
@@ -1101,7 +1099,7 @@ class OpenAINormalizer(FormatNormalizer):
except (TypeError, ValueError): except (TypeError, ValueError):
return None return None
def _coerce_str_list(self, value: Any) -> Optional[List[str]]: def _coerce_str_list(self, value: Any) -> list[str] | None:
if value is None: if value is None:
return None return None
if isinstance(value, str): if isinstance(value, str):
@@ -1110,14 +1108,14 @@ class OpenAINormalizer(FormatNormalizer):
return [str(x) for x in value if x is not None] return [str(x) for x in value if x is not None]
return None return None
def _extract_extra(self, payload: Dict[str, Any], known_keys: set[str]) -> Dict[str, Any]: def _extract_extra(self, payload: dict[str, Any], known_keys: set[str]) -> dict[str, Any]:
return {k: v for k, v in payload.items() if k not in known_keys} return {k: v for k, v in payload.items() if k not in known_keys}
def _merge_dropped(self, target: Dict[str, int], source: Dict[str, int]) -> None: def _merge_dropped(self, target: dict[str, int], source: dict[str, int]) -> None:
for k, v in source.items(): for k, v in source.items():
target[k] = target.get(k, 0) + int(v) target[k] = target.get(k, 0) + int(v)
def _ensure_tool_block_index(self, ss: Dict[str, Any], tool_key: str) -> int: def _ensure_tool_block_index(self, ss: dict[str, Any], tool_key: str) -> int:
mapping = ss.get("tool_id_to_block_index") mapping = ss.get("tool_id_to_block_index")
if not isinstance(mapping, dict): if not isinstance(mapping, dict):
mapping = {} mapping = {}
@@ -1131,7 +1129,7 @@ class OpenAINormalizer(FormatNormalizer):
ss["next_block_index"] = next_idx + 1 ss["next_block_index"] = next_idx + 1
return next_idx return next_idx
def _ensure_tool_call_index(self, ss: Dict[str, Any], tool_id: str) -> int: def _ensure_tool_call_index(self, ss: dict[str, Any], tool_id: str) -> int:
mapping = ss.get("tool_id_to_index") mapping = ss.get("tool_id_to_index")
if not isinstance(mapping, dict): if not isinstance(mapping, dict):
mapping = {} mapping = {}

View File

@@ -10,11 +10,10 @@ OpenAI CLI / Responses Normalizer (OPENAI_CLI)
- 未识别的字段会进入 extra/raw未知内容块保留在 internal但默认输出阶段会丢弃。 - 未识别的字段会进入 extra/raw未知内容块保留在 internal但默认输出阶段会丢弃。
""" """
from __future__ import annotations
import json import json
import time import time
from typing import Any, Dict, List, Optional, Tuple, Union from typing import Any
from src.core.api_format.conversion.field_mappings import ( from src.core.api_format.conversion.field_mappings import (
ERROR_TYPE_MAPPINGS, ERROR_TYPE_MAPPINGS,
@@ -65,7 +64,7 @@ class OpenAICliNormalizer(FormatNormalizer):
supports_images=True, supports_images=True,
) )
_ERROR_TYPE_TO_OPENAI: Dict[ErrorType, str] = { _ERROR_TYPE_TO_OPENAI: dict[ErrorType, str] = {
ErrorType.INVALID_REQUEST: "invalid_request_error", ErrorType.INVALID_REQUEST: "invalid_request_error",
ErrorType.AUTHENTICATION: "invalid_api_key", ErrorType.AUTHENTICATION: "invalid_api_key",
ErrorType.PERMISSION_DENIED: "invalid_request_error", ErrorType.PERMISSION_DENIED: "invalid_request_error",
@@ -82,12 +81,12 @@ class OpenAICliNormalizer(FormatNormalizer):
# Requests # Requests
# ========================= # =========================
def request_to_internal(self, request: Dict[str, Any]) -> InternalRequest: def request_to_internal(self, request: dict[str, Any]) -> InternalRequest:
model = str(request.get("model") or "") model = str(request.get("model") or "")
instructions_text = request.get("instructions") instructions_text = request.get("instructions")
instructions: List[InstructionSegment] = [] instructions: list[InstructionSegment] = []
system_text: Optional[str] = None system_text: str | None = None
if isinstance(instructions_text, str) and instructions_text.strip(): if isinstance(instructions_text, str) and instructions_text.strip():
system_text = instructions_text system_text = instructions_text
instructions.append(InstructionSegment(role=Role.SYSTEM, text=instructions_text)) instructions.append(InstructionSegment(role=Role.SYSTEM, text=instructions_text))
@@ -118,8 +117,8 @@ class OpenAICliNormalizer(FormatNormalizer):
return internal return internal
def request_from_internal(self, internal: InternalRequest) -> Dict[str, Any]: def request_from_internal(self, internal: InternalRequest) -> dict[str, Any]:
result: Dict[str, Any] = { result: dict[str, Any] = {
"model": internal.model, "model": internal.model,
"input": self._internal_messages_to_input(internal.messages), "input": self._internal_messages_to_input(internal.messages),
} }
@@ -164,7 +163,7 @@ class OpenAICliNormalizer(FormatNormalizer):
# Responses # Responses
# ========================= # =========================
def response_to_internal(self, response: Dict[str, Any]) -> InternalResponse: def response_to_internal(self, response: dict[str, Any]) -> InternalResponse:
payload = self._unwrap_response_object(response) payload = self._unwrap_response_object(response)
rid = str(payload.get("id") or "") rid = str(payload.get("id") or "")
@@ -191,8 +190,8 @@ class OpenAICliNormalizer(FormatNormalizer):
self, self,
internal: InternalResponse, internal: InternalResponse,
*, *,
requested_model: Optional[str] = None, requested_model: str | None = None,
) -> Dict[str, Any]: ) -> dict[str, Any]:
text = self._collapse_internal_text(internal.content) text = self._collapse_internal_text(internal.content)
output_message = { output_message = {
@@ -203,7 +202,7 @@ class OpenAICliNormalizer(FormatNormalizer):
} }
usage = internal.usage or UsageInfo() usage = internal.usage or UsageInfo()
usage_obj: Dict[str, Any] = { usage_obj: dict[str, Any] = {
"input_tokens": usage.input_tokens, "input_tokens": usage.input_tokens,
"output_tokens": usage.output_tokens, "output_tokens": usage.output_tokens,
"total_tokens": usage.total_tokens or (usage.input_tokens + usage.output_tokens), "total_tokens": usage.total_tokens or (usage.input_tokens + usage.output_tokens),
@@ -228,11 +227,11 @@ class OpenAICliNormalizer(FormatNormalizer):
def stream_chunk_to_internal( def stream_chunk_to_internal(
self, self,
chunk: Dict[str, Any], chunk: dict[str, Any],
state: StreamState, state: StreamState,
) -> List[InternalStreamEvent]: ) -> list[InternalStreamEvent]:
ss = state.substate(self.FORMAT_ID) ss = state.substate(self.FORMAT_ID)
events: List[InternalStreamEvent] = [] events: list[InternalStreamEvent] = []
# 统一错误结构(最佳努力) # 统一错误结构(最佳努力)
if isinstance(chunk, dict) and "error" in chunk: if isinstance(chunk, dict) and "error" in chunk:
@@ -392,11 +391,11 @@ class OpenAICliNormalizer(FormatNormalizer):
self, self,
event: InternalStreamEvent, event: InternalStreamEvent,
state: StreamState, state: StreamState,
) -> List[Dict[str, Any]]: ) -> list[dict[str, Any]]:
ss = state.substate(self.FORMAT_ID) ss = state.substate(self.FORMAT_ID)
out: List[Dict[str, Any]] = [] out: list[dict[str, Any]] = []
def event_block(payload: Dict[str, Any]) -> Dict[str, Any]: def event_block(payload: dict[str, Any]) -> dict[str, Any]:
# OpenAI Responses SSE 的 payload 通常自带 type 字段;这里强制保证 # OpenAI Responses SSE 的 payload 通常自带 type 字段;这里强制保证
return payload return payload
@@ -463,10 +462,10 @@ class OpenAICliNormalizer(FormatNormalizer):
# Error conversion # Error conversion
# ========================= # =========================
def is_error_response(self, response: Dict[str, Any]) -> bool: def is_error_response(self, response: dict[str, Any]) -> bool:
return isinstance(response, dict) and "error" in response return isinstance(response, dict) and "error" in response
def error_to_internal(self, error_response: Dict[str, Any]) -> InternalError: def error_to_internal(self, error_response: dict[str, Any]) -> InternalError:
err = error_response.get("error") if isinstance(error_response, dict) else None err = error_response.get("error") if isinstance(error_response, dict) else None
err = err if isinstance(err, dict) else {} err = err if isinstance(err, dict) else {}
@@ -484,9 +483,9 @@ class OpenAICliNormalizer(FormatNormalizer):
extra={"openai_cli": {"error": err}, "raw": {"type": raw_type}}, extra={"openai_cli": {"error": err}, "raw": {"type": raw_type}},
) )
def error_from_internal(self, internal: InternalError) -> Dict[str, Any]: def error_from_internal(self, internal: InternalError) -> dict[str, Any]:
type_str = self._ERROR_TYPE_TO_OPENAI.get(internal.type, "server_error") type_str = self._ERROR_TYPE_TO_OPENAI.get(internal.type, "server_error")
payload: Dict[str, Any] = {"type": type_str, "message": internal.message} payload: dict[str, Any] = {"type": type_str, "message": internal.message}
if internal.code is not None: if internal.code is not None:
payload["code"] = internal.code payload["code"] = internal.code
if internal.param is not None: if internal.param is not None:
@@ -497,7 +496,7 @@ class OpenAICliNormalizer(FormatNormalizer):
# Helpers # Helpers
# ========================= # =========================
def _unwrap_response_object(self, response: Dict[str, Any]) -> Dict[str, Any]: def _unwrap_response_object(self, response: dict[str, Any]) -> dict[str, Any]:
if not isinstance(response, dict): if not isinstance(response, dict):
return {} return {}
resp_inner = response.get("response") resp_inner = response.get("response")
@@ -506,8 +505,8 @@ class OpenAICliNormalizer(FormatNormalizer):
return resp_inner return resp_inner
return response return response
def _extract_output_text_blocks(self, payload: Dict[str, Any]) -> Tuple[List[ContentBlock], Dict[str, Any]]: def _extract_output_text_blocks(self, payload: dict[str, Any]) -> tuple[list[ContentBlock], dict[str, Any]]:
text_parts: List[str] = [] text_parts: list[str] = []
output = payload.get("output") output = payload.get("output")
if isinstance(output, list): if isinstance(output, list):
@@ -532,12 +531,12 @@ class OpenAICliNormalizer(FormatNormalizer):
if not text_parts and isinstance(payload.get("output_text"), str): if not text_parts and isinstance(payload.get("output_text"), str):
text_parts.append(payload.get("output_text") or "") text_parts.append(payload.get("output_text") or "")
blocks: List[ContentBlock] = [] blocks: list[ContentBlock] = []
text = "".join(text_parts) text = "".join(text_parts)
if text: if text:
blocks.append(TextBlock(text=text)) blocks.append(TextBlock(text=text))
extra: Dict[str, Any] = {"raw": {"openai_cli_output": output}} if output is not None else {} extra: dict[str, Any] = {"raw": {"openai_cli_output": output}} if output is not None else {}
return blocks, extra return blocks, extra
def _usage_to_internal(self, usage: Any) -> UsageInfo: def _usage_to_internal(self, usage: Any) -> UsageInfo:
@@ -553,14 +552,14 @@ class OpenAICliNormalizer(FormatNormalizer):
extra={"openai_cli": {"usage": usage}}, extra={"openai_cli": {"usage": usage}},
) )
def _collapse_internal_text(self, blocks: List[ContentBlock]) -> str: def _collapse_internal_text(self, blocks: list[ContentBlock]) -> str:
parts: List[str] = [] parts: list[str] = []
for block in blocks: for block in blocks:
if isinstance(block, TextBlock) and block.text: if isinstance(block, TextBlock) and block.text:
parts.append(block.text) parts.append(block.text)
return "".join(parts) return "".join(parts)
def _input_to_internal_messages(self, input_data: Any) -> List[InternalMessage]: def _input_to_internal_messages(self, input_data: Any) -> list[InternalMessage]:
if input_data is None: if input_data is None:
return [] return []
@@ -575,7 +574,7 @@ class OpenAICliNormalizer(FormatNormalizer):
if not isinstance(input_data, list): if not isinstance(input_data, list):
return [InternalMessage(role=Role.USER, content=[UnknownBlock(raw_type="input", payload={"input": input_data})])] return [InternalMessage(role=Role.USER, content=[UnknownBlock(raw_type="input", payload={"input": input_data})])]
messages: List[InternalMessage] = [] messages: list[InternalMessage] = []
for item in input_data: for item in input_data:
if not isinstance(item, dict): if not isinstance(item, dict):
continue continue
@@ -624,7 +623,7 @@ class OpenAICliNormalizer(FormatNormalizer):
# reasoning -> assistant 消息,提取 summary 作为文本 # reasoning -> assistant 消息,提取 summary 作为文本
if item_type == "reasoning": if item_type == "reasoning":
summary_parts: List[str] = [] summary_parts: list[str] = []
summary = item.get("summary") summary = item.get("summary")
if isinstance(summary, list): if isinstance(summary, list):
for s in summary: for s in summary:
@@ -638,7 +637,7 @@ class OpenAICliNormalizer(FormatNormalizer):
summary_parts.append(summary) summary_parts.append(summary)
# 如果有 summary 文本,创建一个 UnknownBlock 保留原始结构 # 如果有 summary 文本,创建一个 UnknownBlock 保留原始结构
reasoning_blocks: List[ContentBlock] = [] reasoning_blocks: list[ContentBlock] = []
if summary_parts: if summary_parts:
# 保留 reasoning 的 summary 作为 UnknownBlock便于输出时决策 # 保留 reasoning 的 summary 作为 UnknownBlock便于输出时决策
reasoning_blocks.append(UnknownBlock( reasoning_blocks.append(UnknownBlock(
@@ -662,7 +661,7 @@ class OpenAICliNormalizer(FormatNormalizer):
return messages return messages
def _responses_content_to_blocks(self, content: Any) -> List[ContentBlock]: def _responses_content_to_blocks(self, content: Any) -> list[ContentBlock]:
if content is None: if content is None:
return [] return []
if isinstance(content, str): if isinstance(content, str):
@@ -673,7 +672,7 @@ class OpenAICliNormalizer(FormatNormalizer):
if not isinstance(content, list): if not isinstance(content, list):
return [UnknownBlock(raw_type="content", payload={"content": content})] return [UnknownBlock(raw_type="content", payload={"content": content})]
blocks: List[ContentBlock] = [] blocks: list[ContentBlock] = []
for part in content: for part in content:
if isinstance(part, str): if isinstance(part, str):
if part: if part:
@@ -690,8 +689,8 @@ class OpenAICliNormalizer(FormatNormalizer):
blocks.append(UnknownBlock(raw_type=ptype or "unknown", payload=part)) blocks.append(UnknownBlock(raw_type=ptype or "unknown", payload=part))
return blocks return blocks
def _internal_messages_to_input(self, messages: List[InternalMessage]) -> List[Dict[str, Any]]: def _internal_messages_to_input(self, messages: list[InternalMessage]) -> list[dict[str, Any]]:
out: List[Dict[str, Any]] = [] out: list[dict[str, Any]] = []
for msg in messages: for msg in messages:
# ToolUseBlock -> function_call # ToolUseBlock -> function_call
for block in msg.content: for block in msg.content:
@@ -729,7 +728,7 @@ class OpenAICliNormalizer(FormatNormalizer):
# 普通 messageTextBlock # 普通 messageTextBlock
role = self._role_to_openai(msg.role) role = self._role_to_openai(msg.role)
content_items: List[Dict[str, Any]] = [] content_items: list[dict[str, Any]] = []
has_text = False has_text = False
for block in msg.content: for block in msg.content:
@@ -748,10 +747,10 @@ class OpenAICliNormalizer(FormatNormalizer):
return out return out
def _tools_to_internal(self, tools: Any) -> Optional[List[ToolDefinition]]: def _tools_to_internal(self, tools: Any) -> list[ToolDefinition] | None:
if not isinstance(tools, list): if not isinstance(tools, list):
return None return None
out: List[ToolDefinition] = [] out: list[ToolDefinition] = []
for tool in tools: for tool in tools:
if not isinstance(tool, dict): if not isinstance(tool, dict):
continue continue
@@ -783,7 +782,7 @@ class OpenAICliNormalizer(FormatNormalizer):
) )
return out or None return out or None
def _tool_choice_to_internal(self, tool_choice: Any) -> Optional[ToolChoice]: def _tool_choice_to_internal(self, tool_choice: Any) -> ToolChoice | None:
if tool_choice is None: if tool_choice is None:
return None return None
if isinstance(tool_choice, str): if isinstance(tool_choice, str):
@@ -802,7 +801,7 @@ class OpenAICliNormalizer(FormatNormalizer):
return ToolChoice(type=ToolChoiceType.AUTO, extra={"raw": tool_choice}) return ToolChoice(type=ToolChoiceType.AUTO, extra={"raw": tool_choice})
def _tool_choice_to_openai(self, tool_choice: ToolChoice) -> Union[str, Dict[str, Any]]: def _tool_choice_to_openai(self, tool_choice: ToolChoice) -> str | dict[str, Any]:
if tool_choice.type == ToolChoiceType.NONE: if tool_choice.type == ToolChoiceType.NONE:
return "none" return "none"
if tool_choice.type == ToolChoiceType.AUTO: if tool_choice.type == ToolChoiceType.AUTO:
@@ -840,7 +839,7 @@ class OpenAICliNormalizer(FormatNormalizer):
return "tool" return "tool"
return "user" return "user"
def _optional_int(self, value: Any) -> Optional[int]: def _optional_int(self, value: Any) -> int | None:
if value is None: if value is None:
return None return None
try: try:
@@ -848,7 +847,7 @@ class OpenAICliNormalizer(FormatNormalizer):
except (TypeError, ValueError): except (TypeError, ValueError):
return None return None
def _optional_float(self, value: Any) -> Optional[float]: def _optional_float(self, value: Any) -> float | None:
if value is None: if value is None:
return None return None
try: try:
@@ -856,13 +855,13 @@ class OpenAICliNormalizer(FormatNormalizer):
except (TypeError, ValueError): except (TypeError, ValueError):
return None return None
def _coerce_str_list(self, value: Any) -> Optional[List[str]]: def _coerce_str_list(self, value: Any) -> list[str] | None:
if value is None: if value is None:
return None return None
if isinstance(value, str): if isinstance(value, str):
return [value] return [value]
if isinstance(value, list): if isinstance(value, list):
out: List[str] = [] out: list[str] = []
for item in value: for item in value:
if item is None: if item is None:
continue continue
@@ -870,14 +869,14 @@ class OpenAICliNormalizer(FormatNormalizer):
return out return out
return [str(value)] return [str(value)]
def _extract_extra(self, payload: Dict[str, Any], keep_keys: set[str]) -> Dict[str, Any]: def _extract_extra(self, payload: dict[str, Any], keep_keys: set[str]) -> dict[str, Any]:
if not isinstance(payload, dict): if not isinstance(payload, dict):
return {} return {}
return {k: v for k, v in payload.items() if k not in keep_keys} return {k: v for k, v in payload.items() if k not in keep_keys}
def _join_instructions(self, internal: InternalRequest) -> str: def _join_instructions(self, internal: InternalRequest) -> str:
if internal.instructions: if internal.instructions:
parts: List[str] = [] parts: list[str] = []
for seg in internal.instructions: for seg in internal.instructions:
if seg.text: if seg.text:
parts.append(seg.text) parts.append(seg.text)

View File

@@ -9,12 +9,12 @@ source -> internal -> target
- 转换失败将抛出 `FormatConversionError`(不再静默回退)。 - 转换失败将抛出 `FormatConversionError`(不再静默回退)。
""" """
from __future__ import annotations
import threading import threading
import time import time
from contextlib import contextmanager from contextlib import contextmanager
from typing import Any, Dict, Generator, List, Optional from typing import Any
from collections.abc import Generator
from src.core.logger import logger from src.core.logger import logger
from src.core.metrics import format_conversion_duration_seconds, format_conversion_total from src.core.metrics import format_conversion_duration_seconds, format_conversion_total
@@ -29,7 +29,7 @@ def _track_conversion_metrics(
direction: str, direction: str,
source: str, source: str,
target: str, target: str,
) -> Generator[None, None, None]: ) -> Generator[None]:
start = time.perf_counter() start = time.perf_counter()
try: try:
yield yield
@@ -47,13 +47,13 @@ class FormatConversionRegistry:
"""基于 Normalizer 的格式转换注册表""" """基于 Normalizer 的格式转换注册表"""
def __init__(self) -> None: def __init__(self) -> None:
self._normalizers: Dict[str, FormatNormalizer] = {} self._normalizers: dict[str, FormatNormalizer] = {}
def register(self, normalizer: FormatNormalizer) -> None: def register(self, normalizer: FormatNormalizer) -> None:
self._normalizers[str(normalizer.FORMAT_ID).upper()] = normalizer self._normalizers[str(normalizer.FORMAT_ID).upper()] = normalizer
logger.info(f"[FormatConversionRegistry] 注册 normalizer: {normalizer.FORMAT_ID}") logger.info(f"[FormatConversionRegistry] 注册 normalizer: {normalizer.FORMAT_ID}")
def get_normalizer(self, format_id: str) -> Optional[FormatNormalizer]: def get_normalizer(self, format_id: str) -> FormatNormalizer | None:
return self._normalizers.get(str(format_id).upper()) return self._normalizers.get(str(format_id).upper())
def _require_normalizer(self, format_id: str) -> FormatNormalizer: def _require_normalizer(self, format_id: str) -> FormatNormalizer:
@@ -66,10 +66,10 @@ class FormatConversionRegistry:
def convert_request( def convert_request(
self, self,
request: Dict[str, Any], request: dict[str, Any],
source_format: str, source_format: str,
target_format: str, target_format: str,
) -> Dict[str, Any]: ) -> dict[str, Any]:
if str(source_format).upper() == str(target_format).upper(): if str(source_format).upper() == str(target_format).upper():
return request return request
@@ -85,12 +85,12 @@ class FormatConversionRegistry:
def convert_response( def convert_response(
self, self,
response: Dict[str, Any], response: dict[str, Any],
source_format: str, source_format: str,
target_format: str, target_format: str,
*, *,
requested_model: Optional[str] = None, requested_model: str | None = None,
) -> Dict[str, Any]: ) -> dict[str, Any]:
"""转换响应格式 """转换响应格式
Args: Args:
@@ -124,10 +124,10 @@ class FormatConversionRegistry:
def convert_error_response( def convert_error_response(
self, self,
error_response: Dict[str, Any], error_response: dict[str, Any],
source_format: str, source_format: str,
target_format: str, target_format: str,
) -> Dict[str, Any]: ) -> dict[str, Any]:
if str(source_format).upper() == str(target_format).upper(): if str(source_format).upper() == str(target_format).upper():
return error_response return error_response
@@ -152,11 +152,11 @@ class FormatConversionRegistry:
def convert_stream_chunk( def convert_stream_chunk(
self, self,
chunk: Dict[str, Any], chunk: dict[str, Any],
source_format: str, source_format: str,
target_format: str, target_format: str,
state: Optional[StreamState] = None, state: StreamState | None = None,
) -> List[Dict[str, Any]]: ) -> list[dict[str, Any]]:
if str(source_format).upper() == str(target_format).upper(): if str(source_format).upper() == str(target_format).upper():
return [chunk] return [chunk]
@@ -182,7 +182,7 @@ class FormatConversionRegistry:
with _track_conversion_metrics("stream", str(source_format).upper(), str(target_format).upper()): with _track_conversion_metrics("stream", str(source_format).upper(), str(target_format).upper()):
try: try:
events = src.stream_chunk_to_internal(chunk, state) events = src.stream_chunk_to_internal(chunk, state)
out: List[Dict[str, Any]] = [] out: list[dict[str, Any]] = []
for event in events: for event in events:
out.extend(tgt.stream_event_from_internal(event, state)) out.extend(tgt.stream_event_from_internal(event, state))
return out return out
@@ -226,10 +226,10 @@ class FormatConversionRegistry:
return self.can_convert_stream(format_a, format_b) and self.can_convert_stream(format_b, format_a) return self.can_convert_stream(format_a, format_b) and self.can_convert_stream(format_b, format_a)
return True return True
def list_normalizers(self) -> List[str]: def list_normalizers(self) -> list[str]:
return sorted(self._normalizers.keys()) return sorted(self._normalizers.keys())
def get_supported_targets(self, source_format: str) -> List[str]: def get_supported_targets(self, source_format: str) -> list[str]:
src = str(source_format).upper() src = str(source_format).upper()
if src not in self._normalizers: if src not in self._normalizers:
return [] return []

View File

@@ -4,11 +4,10 @@
用于把 OpenAI/Claude/Gemini 的流式协议映射为统一事件序列,再由目标格式 Normalizer 输出。 用于把 OpenAI/Claude/Gemini 的流式协议映射为统一事件序列,再由目标格式 Normalizer 输出。
""" """
from __future__ import annotations
from dataclasses import dataclass, field from dataclasses import dataclass, field
from enum import Enum from enum import Enum
from typing import Any, Dict, Optional, Union from typing import Any
from .internal import ContentType, InternalError, StopReason, UsageInfo from .internal import ContentType, InternalError, StopReason, UsageInfo
@@ -34,8 +33,8 @@ class MessageStartEvent:
type: StreamEventType = field(default=StreamEventType.MESSAGE_START, init=False) type: StreamEventType = field(default=StreamEventType.MESSAGE_START, init=False)
message_id: str = "" message_id: str = ""
model: str = "" model: str = ""
usage: Optional[UsageInfo] = None # Claude 流式响应的 message_start 可能包含 usage usage: UsageInfo | None = None # Claude 流式响应的 message_start 可能包含 usage
extra: Dict[str, Any] = field(default_factory=dict) extra: dict[str, Any] = field(default_factory=dict)
@dataclass @dataclass
@@ -46,9 +45,9 @@ class ContentBlockStartEvent:
block_index: int = 0 block_index: int = 0
block_type: ContentType = ContentType.TEXT block_type: ContentType = ContentType.TEXT
# 工具调用时使用TOOL_USE block # 工具调用时使用TOOL_USE block
tool_id: Optional[str] = None tool_id: str | None = None
tool_name: Optional[str] = None tool_name: str | None = None
extra: Dict[str, Any] = field(default_factory=dict) extra: dict[str, Any] = field(default_factory=dict)
@dataclass @dataclass
@@ -58,7 +57,7 @@ class ContentDeltaEvent:
type: StreamEventType = field(default=StreamEventType.CONTENT_DELTA, init=False) type: StreamEventType = field(default=StreamEventType.CONTENT_DELTA, init=False)
block_index: int = 0 block_index: int = 0
text_delta: str = "" text_delta: str = ""
extra: Dict[str, Any] = field(default_factory=dict) extra: dict[str, Any] = field(default_factory=dict)
@dataclass @dataclass
@@ -69,7 +68,7 @@ class ToolCallDeltaEvent:
block_index: int = 0 block_index: int = 0
tool_id: str = "" tool_id: str = ""
input_delta: str = "" # JSON 字符串片段 input_delta: str = "" # JSON 字符串片段
extra: Dict[str, Any] = field(default_factory=dict) extra: dict[str, Any] = field(default_factory=dict)
@dataclass @dataclass
@@ -78,7 +77,7 @@ class ContentBlockStopEvent:
type: StreamEventType = field(default=StreamEventType.CONTENT_BLOCK_STOP, init=False) type: StreamEventType = field(default=StreamEventType.CONTENT_BLOCK_STOP, init=False)
block_index: int = 0 block_index: int = 0
extra: Dict[str, Any] = field(default_factory=dict) extra: dict[str, Any] = field(default_factory=dict)
@dataclass @dataclass
@@ -86,9 +85,9 @@ class MessageStopEvent:
"""消息结束事件""" """消息结束事件"""
type: StreamEventType = field(default=StreamEventType.MESSAGE_STOP, init=False) type: StreamEventType = field(default=StreamEventType.MESSAGE_STOP, init=False)
stop_reason: Optional[StopReason] = None stop_reason: StopReason | None = None
usage: Optional[UsageInfo] = None usage: UsageInfo | None = None
extra: Dict[str, Any] = field(default_factory=dict) extra: dict[str, Any] = field(default_factory=dict)
@dataclass @dataclass
@@ -97,7 +96,7 @@ class UsageEvent:
type: StreamEventType = field(default=StreamEventType.USAGE, init=False) type: StreamEventType = field(default=StreamEventType.USAGE, init=False)
usage: UsageInfo = field(default_factory=UsageInfo) usage: UsageInfo = field(default_factory=UsageInfo)
extra: Dict[str, Any] = field(default_factory=dict) extra: dict[str, Any] = field(default_factory=dict)
@dataclass @dataclass
@@ -106,7 +105,7 @@ class ErrorEvent:
type: StreamEventType = field(default=StreamEventType.ERROR, init=False) type: StreamEventType = field(default=StreamEventType.ERROR, init=False)
error: InternalError error: InternalError
extra: Dict[str, Any] = field(default_factory=dict) extra: dict[str, Any] = field(default_factory=dict)
@dataclass @dataclass
@@ -115,21 +114,21 @@ class UnknownStreamEvent:
type: StreamEventType = field(default=StreamEventType.UNKNOWN, init=False) type: StreamEventType = field(default=StreamEventType.UNKNOWN, init=False)
raw_type: str = "" raw_type: str = ""
payload: Dict[str, Any] = field(default_factory=dict) payload: dict[str, Any] = field(default_factory=dict)
extra: Dict[str, Any] = field(default_factory=dict) extra: dict[str, Any] = field(default_factory=dict)
InternalStreamEvent = Union[ InternalStreamEvent = (
MessageStartEvent, MessageStartEvent
ContentBlockStartEvent, | ContentBlockStartEvent
ContentDeltaEvent, | ContentDeltaEvent
ToolCallDeltaEvent, | ToolCallDeltaEvent
ContentBlockStopEvent, | ContentBlockStopEvent
MessageStopEvent, | MessageStopEvent
UsageEvent, | UsageEvent
ErrorEvent, | ErrorEvent
UnknownStreamEvent, | UnknownStreamEvent
] )
__all__ = [ __all__ = [

View File

@@ -5,10 +5,9 @@
每个 Normalizer 通过 `substate(format_id)` 获取自己的隔离状态字典。 每个 Normalizer 通过 `substate(format_id)` 获取自己的隔离状态字典。
""" """
from __future__ import annotations
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import Any, Dict from typing import Any
@dataclass @dataclass
@@ -26,12 +25,12 @@ class StreamState:
message_id: str = "" message_id: str = ""
# Registry/调用层的通用扩展信息(与具体格式无关) # Registry/调用层的通用扩展信息(与具体格式无关)
extra: Dict[str, Any] = field(default_factory=dict) extra: dict[str, Any] = field(default_factory=dict)
# 各 Normalizer 的隔离状态key: FORMAT_ID # 各 Normalizer 的隔离状态key: FORMAT_ID
by_format: Dict[str, Dict[str, Any]] = field(default_factory=dict) by_format: dict[str, dict[str, Any]] = field(default_factory=dict)
def substate(self, format_id: str) -> Dict[str, Any]: def substate(self, format_id: str) -> dict[str, Any]:
"""获取指定格式的隔离子状态""" """获取指定格式的隔离子状态"""
key = str(format_id).upper() key = str(format_id).upper()
return self.by_format.setdefault(key, {}) return self.by_format.setdefault(key, {})

View File

@@ -6,7 +6,8 @@ API 格式检测
from __future__ import annotations from __future__ import annotations
from typing import TYPE_CHECKING, Dict, Optional, Tuple
from typing import TYPE_CHECKING
if TYPE_CHECKING: if TYPE_CHECKING:
from starlette.requests import Request from starlette.requests import Request
@@ -16,10 +17,10 @@ from src.core.api_format.metadata import API_FORMAT_DEFINITIONS, ApiFormatDefini
def _extract_api_key_by_definition( def _extract_api_key_by_definition(
headers: Dict[str, str], headers: dict[str, str],
query_params: Optional[Dict[str, str]], query_params: dict[str, str] | None,
definition: ApiFormatDefinition, definition: ApiFormatDefinition,
) -> Tuple[Optional[str], str]: ) -> tuple[str | None, str]:
""" """
根据格式定义从请求中提取 API Key 根据格式定义从请求中提取 API Key
@@ -64,9 +65,9 @@ def _extract_api_key_by_definition(
def detect_format_from_request( def detect_format_from_request(
headers: Dict[str, str], headers: dict[str, str],
query_params: Optional[Dict[str, str]] = None, query_params: dict[str, str] | None = None,
) -> Tuple[APIFormat, Optional[str], str]: ) -> tuple[APIFormat, str | None, str]:
""" """
从请求头检测 API 格式和 API Key 从请求头检测 API 格式和 API Key
@@ -107,8 +108,8 @@ def detect_format_from_request(
def detect_format_and_key_from_starlette( def detect_format_and_key_from_starlette(
request: "Request", request: Request,
) -> Tuple[str, Optional[str], str]: ) -> tuple[str, str | None, str]:
""" """
从 Starlette Request 对象检测 API 格式和 API Key 从 Starlette Request 对象检测 API 格式和 API Key
@@ -135,7 +136,7 @@ def detect_format_and_key_from_starlette(
def detect_format_from_response( def detect_format_from_response(
response_data: dict, response_data: dict,
) -> Optional[APIFormat]: ) -> APIFormat | None:
""" """
从响应内容检测 API 格式 从响应内容检测 API 格式

View File

@@ -12,7 +12,8 @@
from __future__ import annotations from __future__ import annotations
from typing import AbstractSet, Any, Dict, FrozenSet, Optional, Set from collections.abc import Set as AbstractSet
from typing import Any
from src.core.api_format.enums import APIFormat from src.core.api_format.enums import APIFormat
from src.core.api_format.metadata import get_auth_config, get_extra_headers, get_protected_keys from src.core.api_format.metadata import get_auth_config, get_extra_headers, get_protected_keys
@@ -23,7 +24,7 @@ from src.core.api_format.metadata import get_auth_config, get_extra_headers, get
# ============================================================================= # =============================================================================
# 转发给上游时需要剔除的头部(系统管理 + 认证替换) # 转发给上游时需要剔除的头部(系统管理 + 认证替换)
UPSTREAM_DROP_HEADERS: FrozenSet[str] = frozenset( UPSTREAM_DROP_HEADERS: frozenset[str] = frozenset(
{ {
# 认证头 - 会被替换为 Provider 的认证 # 认证头 - 会被替换为 Provider 的认证
"authorization", "authorization",
@@ -41,7 +42,7 @@ UPSTREAM_DROP_HEADERS: FrozenSet[str] = frozenset(
# 最小必脱敏集合(编译时常量,用于快速路径) # 最小必脱敏集合(编译时常量,用于快速路径)
# 完整脱敏应使用 SystemConfigService.get_sensitive_headers() # 完整脱敏应使用 SystemConfigService.get_sensitive_headers()
CORE_REDACT_HEADERS: FrozenSet[str] = frozenset( CORE_REDACT_HEADERS: frozenset[str] = frozenset(
{ {
"authorization", "authorization",
"x-api-key", "x-api-key",
@@ -50,7 +51,7 @@ CORE_REDACT_HEADERS: FrozenSet[str] = frozenset(
) )
# Hop-by-hop 头部 (RFC 7230) # Hop-by-hop 头部 (RFC 7230)
HOP_BY_HOP_HEADERS: FrozenSet[str] = frozenset( HOP_BY_HOP_HEADERS: frozenset[str] = frozenset(
{ {
"connection", "connection",
"keep-alive", "keep-alive",
@@ -64,7 +65,7 @@ HOP_BY_HOP_HEADERS: FrozenSet[str] = frozenset(
) )
# 响应时需要过滤的头部body-dependent + hop-by-hop # 响应时需要过滤的头部body-dependent + hop-by-hop
RESPONSE_DROP_HEADERS: FrozenSet[str] = ( RESPONSE_DROP_HEADERS: frozenset[str] = (
frozenset( frozenset(
{ {
"content-length", "content-length",
@@ -82,7 +83,7 @@ RESPONSE_DROP_HEADERS: FrozenSet[str] = (
# ============================================================================= # =============================================================================
def normalize_headers(headers: Dict[str, str]) -> Dict[str, str]: def normalize_headers(headers: dict[str, str]) -> dict[str, str]:
""" """
将请求头 key 统一为小写 将请求头 key 统一为小写
@@ -92,7 +93,7 @@ def normalize_headers(headers: Dict[str, str]) -> Dict[str, str]:
return {k.lower(): v for k, v in headers.items()} return {k.lower(): v for k, v in headers.items()}
def get_header_value(headers: Dict[str, str], key: str, default: str = "") -> str: def get_header_value(headers: dict[str, str], key: str, default: str = "") -> str:
""" """
大小写不敏感地获取请求头值 大小写不敏感地获取请求头值
@@ -117,7 +118,7 @@ def get_header_value(headers: Dict[str, str], key: str, default: str = "") -> st
# ============================================================================= # =============================================================================
def extract_client_api_key(headers: Dict[str, str], api_format: APIFormat) -> Optional[str]: def extract_client_api_key(headers: dict[str, str], api_format: APIFormat) -> str | None:
""" """
从客户端请求头提取 API Key 从客户端请求头提取 API Key
@@ -147,10 +148,10 @@ def extract_client_api_key(headers: Dict[str, str], api_format: APIFormat) -> Op
def extract_client_api_key_with_query( def extract_client_api_key_with_query(
headers: Dict[str, str], headers: dict[str, str],
query_params: Optional[Dict[str, str]], query_params: dict[str, str] | None,
api_format: APIFormat, api_format: APIFormat,
) -> Optional[str]: ) -> str | None:
""" """
从客户端请求头或 URL 参数提取 API Key 从客户端请求头或 URL 参数提取 API Key
@@ -184,10 +185,10 @@ def extract_client_api_key_with_query(
def detect_capabilities( def detect_capabilities(
headers: Dict[str, str], headers: dict[str, str],
api_format: APIFormat, api_format: APIFormat,
request_body: Optional[Dict[str, Any]] = None, # noqa: ARG001 - 预留给部分格式使用 request_body: dict[str, Any] | None = None, # noqa: ARG001 - 预留给部分格式使用
) -> Dict[str, bool]: ) -> dict[str, bool]:
""" """
从请求头检测能力需求 从请求头检测能力需求
@@ -203,7 +204,7 @@ def detect_capabilities(
能力需求字典,如 {"context_1m": True} 能力需求字典,如 {"context_1m": True}
""" """
requirements: Dict[str, bool] = {} requirements: dict[str, bool] = {}
if api_format in (APIFormat.CLAUDE, APIFormat.CLAUDE_CLI): if api_format in (APIFormat.CLAUDE, APIFormat.CLAUDE_CLI):
beta_header = get_header_value(headers, "anthropic-beta") beta_header = get_header_value(headers, "anthropic-beta")
@@ -228,20 +229,20 @@ class HeaderBuilder:
def __init__(self) -> None: def __init__(self) -> None:
# key: (original_case_key, value) # key: (original_case_key, value)
self._headers: Dict[str, tuple[str, str]] = {} self._headers: dict[str, tuple[str, str]] = {}
def add(self, key: str, value: str) -> "HeaderBuilder": def add(self, key: str, value: str) -> HeaderBuilder:
"""添加单个头部(会覆盖同名头部)""" """添加单个头部(会覆盖同名头部)"""
self._headers[key.lower()] = (key, value) self._headers[key.lower()] = (key, value)
return self return self
def add_many(self, headers: Dict[str, str]) -> "HeaderBuilder": def add_many(self, headers: dict[str, str]) -> HeaderBuilder:
"""批量添加头部""" """批量添加头部"""
for k, v in headers.items(): for k, v in headers.items():
self.add(k, v) self.add(k, v)
return self return self
def add_protected(self, headers: Dict[str, str], protected_keys: AbstractSet[str]) -> "HeaderBuilder": def add_protected(self, headers: dict[str, str], protected_keys: AbstractSet[str]) -> HeaderBuilder:
""" """
添加头部但保护指定的 key 不被覆盖 添加头部但保护指定的 key 不被覆盖
@@ -253,13 +254,13 @@ class HeaderBuilder:
self.add(k, v) self.add(k, v)
return self return self
def remove(self, keys: FrozenSet[str]) -> "HeaderBuilder": def remove(self, keys: frozenset[str]) -> HeaderBuilder:
"""移除指定的头部""" """移除指定的头部"""
for k in keys: for k in keys:
self._headers.pop(k.lower(), None) self._headers.pop(k.lower(), None)
return self return self
def rename(self, from_key: str, to_key: str) -> "HeaderBuilder": def rename(self, from_key: str, to_key: str) -> HeaderBuilder:
""" """
重命名头部(保留原值) 重命名头部(保留原值)
@@ -273,9 +274,9 @@ class HeaderBuilder:
def apply_rules( def apply_rules(
self, self,
rules: list[Dict[str, Any]], rules: list[dict[str, Any]],
protected_keys: Optional[AbstractSet[str]] = None, protected_keys: AbstractSet[str] | None = None,
) -> "HeaderBuilder": ) -> HeaderBuilder:
""" """
应用请求头规则 应用请求头规则
@@ -314,20 +315,20 @@ class HeaderBuilder:
return self return self
def build(self) -> Dict[str, str]: def build(self) -> dict[str, str]:
"""构建最终的头部字典""" """构建最终的头部字典"""
return {original_key: value for original_key, value in self._headers.values()} return {original_key: value for original_key, value in self._headers.values()}
def build_upstream_headers( def build_upstream_headers(
original_headers: Dict[str, str], original_headers: dict[str, str],
api_format: APIFormat, api_format: APIFormat,
provider_api_key: str, provider_api_key: str,
*, *,
endpoint_headers: Optional[Dict[str, str]] = None, endpoint_headers: dict[str, str] | None = None,
extra_headers: Optional[Dict[str, str]] = None, extra_headers: dict[str, str] | None = None,
drop_headers: Optional[FrozenSet[str]] = None, drop_headers: frozenset[str] | None = None,
) -> Dict[str, str]: ) -> dict[str, str]:
""" """
构建发送给上游 Provider 的请求头 构建发送给上游 Provider 的请求头
@@ -386,10 +387,10 @@ def build_upstream_headers(
def merge_headers_with_protection( def merge_headers_with_protection(
base_headers: Dict[str, str], base_headers: dict[str, str],
extra_headers: Optional[Dict[str, str]], extra_headers: dict[str, str] | None,
protected_keys: FrozenSet[str] | Set[str], protected_keys: frozenset[str] | set[str],
) -> Dict[str, str]: ) -> dict[str, str]:
""" """
合并头部但保护指定的 key 不被覆盖 合并头部但保护指定的 key 不被覆盖
@@ -418,9 +419,9 @@ def merge_headers_with_protection(
def filter_response_headers( def filter_response_headers(
headers: Optional[Dict[str, str]], headers: dict[str, str] | None,
drop_headers: Optional[FrozenSet[str]] = None, drop_headers: frozenset[str] | None = None,
) -> Dict[str, str]: ) -> dict[str, str]:
""" """
过滤上游响应头中不应透传给客户端的字段 过滤上游响应头中不应透传给客户端的字段
@@ -446,9 +447,9 @@ def filter_response_headers(
def redact_headers_for_log( def redact_headers_for_log(
headers: Dict[str, str], headers: dict[str, str],
redact_keys: Optional[FrozenSet[str]] = None, redact_keys: frozenset[str] | None = None,
) -> Dict[str, str]: ) -> dict[str, str]:
""" """
将敏感头部值替换为 *** 用于日志记录 将敏感头部值替换为 *** 用于日志记录
@@ -487,7 +488,7 @@ def build_adapter_base_headers(
api_key: str, api_key: str,
*, *,
include_extra: bool = True, include_extra: bool = True,
) -> Dict[str, str]: ) -> dict[str, str]:
""" """
根据 API 格式构建基础请求头 根据 API 格式构建基础请求头
@@ -504,7 +505,7 @@ def build_adapter_base_headers(
auth_header, auth_type = get_auth_config(api_format) auth_header, auth_type = get_auth_config(api_format)
auth_value = f"Bearer {api_key}" if auth_type == "bearer" else api_key auth_value = f"Bearer {api_key}" if auth_type == "bearer" else api_key
headers: Dict[str, str] = { headers: dict[str, str] = {
auth_header: auth_value, auth_header: auth_value,
"Content-Type": "application/json", "Content-Type": "application/json",
} }
@@ -520,8 +521,8 @@ def build_adapter_base_headers(
def build_adapter_headers( def build_adapter_headers(
api_format: APIFormat, api_format: APIFormat,
api_key: str, api_key: str,
extra_headers: Optional[Dict[str, str]] = None, extra_headers: dict[str, str] | None = None,
) -> Dict[str, str]: ) -> dict[str, str]:
""" """
构建完整的 Adapter 请求头 构建完整的 Adapter 请求头
@@ -565,8 +566,8 @@ def get_adapter_protected_keys(api_format: APIFormat) -> tuple[str, ...]:
def extract_set_headers_from_rules( def extract_set_headers_from_rules(
header_rules: Optional[list[Dict[str, Any]]], header_rules: list[dict[str, Any]] | None,
) -> Optional[Dict[str, str]]: ) -> dict[str, str] | None:
""" """
从 header_rules 中提取 set 操作生成的头部字典 从 header_rules 中提取 set 操作生成的头部字典
@@ -582,7 +583,7 @@ def extract_set_headers_from_rules(
if not header_rules: if not header_rules:
return None return None
headers: Dict[str, str] = {} headers: dict[str, str] = {}
for rule in header_rules: for rule in header_rules:
if rule.get("action") == "set": if rule.get("action") == "set":
key = rule.get("key", "") key = rule.get("key", "")
@@ -593,7 +594,7 @@ def extract_set_headers_from_rules(
return headers if headers else None return headers if headers else None
def get_extra_headers_from_endpoint(endpoint: Any) -> Optional[Dict[str, str]]: def get_extra_headers_from_endpoint(endpoint: Any) -> dict[str, str] | None:
""" """
从 endpoint 提取额外请求头 从 endpoint 提取额外请求头

View File

@@ -13,13 +13,12 @@ API 格式元数据定义
definition = get_api_format_definition(APIFormat.CLAUDE) definition = get_api_format_definition(APIFormat.CLAUDE)
""" """
from __future__ import annotations
import re import re
from dataclasses import dataclass, field from dataclasses import dataclass, field
from functools import lru_cache from functools import lru_cache
from types import MappingProxyType from types import MappingProxyType
from typing import Dict, Iterable, List, Mapping, MutableMapping, Optional, Sequence, Union from collections.abc import Iterable, Mapping, MutableMapping, Sequence
from .enums import APIFormat from .enums import APIFormat
@@ -64,7 +63,7 @@ class ApiFormatDefinition:
yield normalized yield normalized
_DEFINITIONS: Dict[APIFormat, ApiFormatDefinition] = { _DEFINITIONS: dict[APIFormat, ApiFormatDefinition] = {
APIFormat.CLAUDE: ApiFormatDefinition( APIFormat.CLAUDE: ApiFormatDefinition(
api_format=APIFormat.CLAUDE, api_format=APIFormat.CLAUDE,
aliases=("claude", "anthropic", "claude_compatible"), aliases=("claude", "anthropic", "claude_compatible"),
@@ -151,12 +150,12 @@ def get_api_format_definition(api_format: APIFormat) -> ApiFormatDefinition:
return API_FORMAT_DEFINITIONS[api_format] return API_FORMAT_DEFINITIONS[api_format]
def list_api_format_definitions() -> List[ApiFormatDefinition]: def list_api_format_definitions() -> list[ApiFormatDefinition]:
"""返回所有定义的浅拷贝列表,供遍历使用。""" """返回所有定义的浅拷贝列表,供遍历使用。"""
return list(API_FORMAT_DEFINITIONS.values()) return list(API_FORMAT_DEFINITIONS.values())
def build_alias_lookup() -> Dict[str, APIFormat]: def build_alias_lookup() -> dict[str, APIFormat]:
""" """
构建 alias -> APIFormat 的查找表。 构建 alias -> APIFormat 的查找表。
每次调用都会返回新的 dict避免可变全局引发并发问题。 每次调用都会返回新的 dict避免可变全局引发并发问题。
@@ -237,7 +236,7 @@ def get_protected_keys(api_format: APIFormat) -> frozenset[str]:
return frozenset({"authorization", "content-type"}) return frozenset({"authorization", "content-type"})
def get_data_format_id(api_format: Union[str, APIFormat]) -> str: def get_data_format_id(api_format: str | APIFormat) -> str:
""" """
获取格式的数据格式标识。 获取格式的数据格式标识。
@@ -264,7 +263,7 @@ def get_data_format_id(api_format: Union[str, APIFormat]) -> str:
return api_format.value.lower() return api_format.value.lower()
def can_passthrough(client_format: Union[str, APIFormat], endpoint_format: Union[str, APIFormat]) -> bool: def can_passthrough(client_format: str | APIFormat, endpoint_format: str | APIFormat) -> bool:
""" """
判断两个格式之间是否可以透传(不需要数据转换)。 判断两个格式之间是否可以透传(不需要数据转换)。
@@ -294,12 +293,12 @@ def can_passthrough(client_format: Union[str, APIFormat], endpoint_format: Union
@lru_cache(maxsize=1) @lru_cache(maxsize=1)
def _alias_lookup_cache() -> Dict[str, APIFormat]: def _alias_lookup_cache() -> dict[str, APIFormat]:
"""缓存 alias -> APIFormat 查找表,减少重复构建。""" """缓存 alias -> APIFormat 查找表,减少重复构建。"""
return build_alias_lookup() return build_alias_lookup()
def resolve_api_format_alias(value: str) -> Optional[APIFormat]: def resolve_api_format_alias(value: str) -> APIFormat | None:
"""根据别名查找 APIFormat找不到时返回 None。""" """根据别名查找 APIFormat找不到时返回 None。"""
if not value: if not value:
return None return None
@@ -310,9 +309,9 @@ def resolve_api_format_alias(value: str) -> Optional[APIFormat]:
def resolve_api_format( def resolve_api_format(
value: Union[str, APIFormat, None], value: str | APIFormat | None,
default: Optional[APIFormat] = None, default: APIFormat | None = None,
) -> Optional[APIFormat]: ) -> APIFormat | None:
""" """
将任意字符串/枚举值解析为 APIFormat。 将任意字符串/枚举值解析为 APIFormat。

View File

@@ -6,13 +6,13 @@ API 格式工具函数
from __future__ import annotations from __future__ import annotations
from typing import TYPE_CHECKING, Optional, Union from typing import TYPE_CHECKING
if TYPE_CHECKING: if TYPE_CHECKING:
from src.core.api_format.enums import APIFormat from src.core.api_format.enums import APIFormat
def is_cli_format(format_id: Union[str, "APIFormat", None]) -> bool: def is_cli_format(format_id: str | APIFormat | None) -> bool:
""" """
判断是否为 CLI 透传格式 判断是否为 CLI 透传格式
@@ -40,7 +40,7 @@ def is_cli_format(format_id: Union[str, "APIFormat", None]) -> bool:
return str(format_id).upper().endswith("_CLI") return str(format_id).upper().endswith("_CLI")
def get_base_format(format_id: Union[str, "APIFormat", None]) -> Optional[str]: def get_base_format(format_id: str | APIFormat | None) -> str | None:
""" """
获取基础格式(去除 _CLI 后缀) 获取基础格式(去除 _CLI 后缀)
@@ -66,7 +66,7 @@ def get_base_format(format_id: Union[str, "APIFormat", None]) -> Optional[str]:
return format_str return format_str
def normalize_format(format_id: Union[str, "APIFormat", None]) -> Optional[str]: def normalize_format(format_id: str | APIFormat | None) -> str | None:
""" """
规范化格式标识符 规范化格式标识符
@@ -84,8 +84,8 @@ def normalize_format(format_id: Union[str, "APIFormat", None]) -> Optional[str]:
def is_same_format( def is_same_format(
format1: Union[str, "APIFormat", None], format1: str | APIFormat | None,
format2: Union[str, "APIFormat", None], format2: str | APIFormat | None,
) -> bool: ) -> bool:
""" """
判断两个格式是否相同 判断两个格式是否相同
@@ -95,7 +95,7 @@ def is_same_format(
return normalize_format(format1) == normalize_format(format2) return normalize_format(format1) == normalize_format(format2)
def is_convertible_format(format_id: Union[str, "APIFormat", None]) -> bool: def is_convertible_format(format_id: str | APIFormat | None) -> bool:
""" """
判断是否为可转换格式 判断是否为可转换格式

View File

@@ -8,7 +8,6 @@
""" """
import asyncio import asyncio
from typing import Set
from src.core.logger import logger from src.core.logger import logger
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
@@ -23,7 +22,7 @@ class BatchCommitter:
interval_seconds: 批量提交间隔(秒) interval_seconds: 批量提交间隔(秒)
""" """
self.interval_seconds = interval_seconds self.interval_seconds = interval_seconds
self._pending_sessions: Set[Session] = set() self._pending_sessions: set[Session] = set()
self._lock = asyncio.Lock() self._lock = asyncio.Lock()
self._task = None self._task = None

Some files were not shown because too many files have changed in this diff Show More