mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-03 01:40:21 +08:00
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:
@@ -7,7 +7,7 @@ WORKDIR /app
|
||||
COPY frontend/ ./frontend/
|
||||
RUN cd frontend && npm run build
|
||||
# ==================== 运行时镜像 ====================
|
||||
FROM python:3.12-slim
|
||||
FROM python:3.14-slim
|
||||
WORKDIR /app
|
||||
# 运行时依赖(无 gcc/nodejs/npm)
|
||||
RUN apt-get update && apt-get install -y \
|
||||
@@ -17,7 +17,7 @@ RUN apt-get update && apt-get install -y \
|
||||
curl \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
# 从 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 可执行文件
|
||||
COPY --from=builder /usr/local/bin/gunicorn /usr/local/bin/
|
||||
COPY --from=builder /usr/local/bin/uvicorn /usr/local/bin/
|
||||
|
||||
@@ -10,7 +10,7 @@ COPY frontend/ ./frontend/
|
||||
RUN cd frontend && npm run build
|
||||
|
||||
# ==================== 运行时镜像 ====================
|
||||
FROM python:3.12-slim
|
||||
FROM python:3.14-slim
|
||||
|
||||
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/*
|
||||
|
||||
# 从 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 可执行文件
|
||||
COPY --from=builder /usr/local/bin/gunicorn /usr/local/bin/
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
# 用于 GitHub Actions CI 构建(不使用国内镜像源)
|
||||
# 构建命令: docker build -f Dockerfile.base -t aether-base:latest .
|
||||
# 只在 pyproject.toml 或 frontend/package*.json 变化时需要重建
|
||||
FROM python:3.12-slim
|
||||
FROM python:3.14-slim
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
# 构建镜像:编译环境 + 预编译的依赖(国内镜像源版本)
|
||||
# 构建命令: docker build -f Dockerfile.base.local -t aether-base:latest .
|
||||
# 只在 pyproject.toml 或 frontend/package*.json 变化时需要重建
|
||||
FROM python:3.12-slim
|
||||
FROM python:3.14-slim
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
|
||||
@@ -15,13 +15,11 @@ classifiers = [
|
||||
"Intended Audience :: Developers",
|
||||
"License :: Other/Proprietary License",
|
||||
"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.13",
|
||||
"Programming Language :: Python :: 3.14",
|
||||
]
|
||||
requires-python = ">=3.9"
|
||||
requires-python = ">=3.12"
|
||||
dependencies = [
|
||||
"fastapi[standard]>=0.115.11",
|
||||
"uvicorn>=0.34.0",
|
||||
@@ -83,7 +81,7 @@ dev-dependencies = [
|
||||
|
||||
[tool.black]
|
||||
line-length = 100
|
||||
target-version = ['py38']
|
||||
target-version = ['py312']
|
||||
|
||||
[tool.isort]
|
||||
profile = "black"
|
||||
@@ -99,7 +97,7 @@ source = "vcs"
|
||||
version-file = "src/_version.py"
|
||||
|
||||
[tool.mypy]
|
||||
python_version = "3.9"
|
||||
python_version = "3.12"
|
||||
warn_return_any = true
|
||||
warn_unused_configs = true
|
||||
disallow_untyped_defs = true
|
||||
|
||||
@@ -10,7 +10,6 @@
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from pydantic import BaseModel, Field, ValidationError
|
||||
@@ -34,7 +33,7 @@ class EnableAdaptiveRequest(BaseModel):
|
||||
"""启用自适应模式请求"""
|
||||
|
||||
enabled: bool = Field(..., description="是否启用自适应模式(true=自适应,false=固定限制)")
|
||||
fixed_limit: Optional[int] = Field(
|
||||
fixed_limit: int | None = Field(
|
||||
None, ge=1, le=100, description="固定 RPM 限制(仅当 enabled=false 时生效,1-100)"
|
||||
)
|
||||
|
||||
@@ -43,30 +42,30 @@ class AdaptiveStatsResponse(BaseModel):
|
||||
"""自适应统计响应"""
|
||||
|
||||
adaptive_mode: bool = Field(..., description="是否为自适应模式(rpm_limit=NULL)")
|
||||
rpm_limit: Optional[int] = Field(None, description="用户配置的固定限制(NULL=自适应)")
|
||||
effective_limit: Optional[int] = Field(
|
||||
rpm_limit: int | None = Field(None, description="用户配置的固定限制(NULL=自适应)")
|
||||
effective_limit: int | None = Field(
|
||||
None, description="当前有效限制(自适应使用学习值,固定使用配置值)"
|
||||
)
|
||||
learned_limit: Optional[int] = Field(None, description="学习到的 RPM 限制")
|
||||
learned_limit: int | None = Field(None, description="学习到的 RPM 限制")
|
||||
concurrent_429_count: int
|
||||
rpm_429_count: int
|
||||
last_429_at: Optional[str]
|
||||
last_429_type: Optional[str]
|
||||
last_429_at: str | None
|
||||
last_429_type: str | None
|
||||
adjustment_count: int
|
||||
recent_adjustments: List[dict]
|
||||
recent_adjustments: list[dict]
|
||||
|
||||
|
||||
class KeyListItem(BaseModel):
|
||||
"""Key 列表项"""
|
||||
|
||||
id: str
|
||||
name: Optional[str]
|
||||
name: str | None
|
||||
provider_id: str
|
||||
api_formats: List[str] = Field(default_factory=list)
|
||||
api_formats: list[str] = Field(default_factory=list)
|
||||
is_adaptive: bool = Field(..., description="是否为自适应模式(rpm_limit=NULL)")
|
||||
rpm_limit: Optional[int] = Field(None, description="固定 RPM 限制(NULL=自适应)")
|
||||
effective_limit: Optional[int] = Field(None, description="当前有效限制")
|
||||
learned_rpm_limit: Optional[int] = Field(None, description="学习到的 RPM 限制")
|
||||
rpm_limit: int | None = Field(None, description="固定 RPM 限制(NULL=自适应)")
|
||||
effective_limit: int | None = Field(None, description="当前有效限制")
|
||||
learned_rpm_limit: int | None = Field(None, description="学习到的 RPM 限制")
|
||||
concurrent_429_count: int
|
||||
rpm_429_count: int
|
||||
|
||||
@@ -76,12 +75,12 @@ class KeyListItem(BaseModel):
|
||||
|
||||
@router.get(
|
||||
"/keys",
|
||||
response_model=List[KeyListItem],
|
||||
response_model=list[KeyListItem],
|
||||
summary="获取所有启用自适应模式的Key",
|
||||
)
|
||||
async def list_adaptive_keys(
|
||||
request: Request,
|
||||
provider_id: Optional[str] = Query(None, description="按 Provider 过滤"),
|
||||
provider_id: str | None = Query(None, description="按 Provider 过滤"),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
"""
|
||||
@@ -207,7 +206,7 @@ async def get_adaptive_summary(
|
||||
|
||||
@dataclass
|
||||
class ListAdaptiveKeysAdapter(AdminApiAdapter):
|
||||
provider_id: Optional[str] = None
|
||||
provider_id: str | None = None
|
||||
|
||||
async def handle(self, context): # type: ignore[override]
|
||||
# 自适应模式:rpm_limit = NULL
|
||||
|
||||
@@ -5,7 +5,6 @@
|
||||
|
||||
import os
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Optional
|
||||
from zoneinfo import ZoneInfo
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
@@ -25,7 +24,7 @@ from src.services.user.apikey import ApiKeyService
|
||||
APP_TIMEZONE = ZoneInfo(os.getenv("APP_TIMEZONE", "Asia/Shanghai"))
|
||||
|
||||
|
||||
def parse_expiry_date(date_str: Optional[str]) -> Optional[datetime]:
|
||||
def parse_expiry_date(date_str: str | None) -> datetime | None:
|
||||
"""解析过期日期字符串为 datetime 对象。
|
||||
|
||||
Args:
|
||||
@@ -70,7 +69,7 @@ async def list_standalone_api_keys(
|
||||
request: Request,
|
||||
skip: int = Query(0, ge=0),
|
||||
limit: int = Query(100, ge=1, le=500),
|
||||
is_active: Optional[bool] = None,
|
||||
is_active: bool | None = None,
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
"""
|
||||
@@ -330,7 +329,7 @@ class AdminListStandaloneKeysAdapter(AdminApiAdapter):
|
||||
self,
|
||||
skip: int,
|
||||
limit: int,
|
||||
is_active: Optional[bool],
|
||||
is_active: bool | None,
|
||||
):
|
||||
self.skip = skip
|
||||
self.limit = limit
|
||||
|
||||
@@ -5,7 +5,6 @@ Endpoint 健康监控 API
|
||||
from collections import defaultdict
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, Request
|
||||
from sqlalchemy import func
|
||||
@@ -128,7 +127,7 @@ async def get_api_format_health_monitor(
|
||||
async def get_key_health(
|
||||
key_id: str,
|
||||
request: Request,
|
||||
api_format: Optional[str] = Query(None, description="API 格式(可选,如 CLAUDE、OPENAI)"),
|
||||
api_format: str | None = Query(None, description="API 格式(可选,如 CLAUDE、OPENAI)"),
|
||||
db: Session = Depends(get_db),
|
||||
) -> HealthStatusResponse:
|
||||
"""
|
||||
@@ -161,7 +160,7 @@ async def get_key_health(
|
||||
async def recover_key_health(
|
||||
key_id: str,
|
||||
request: Request,
|
||||
api_format: Optional[str] = Query(None, description="API 格式(可选,不指定则恢复所有格式)"),
|
||||
api_format: str | None = Query(None, description="API 格式(可选,不指定则恢复所有格式)"),
|
||||
db: Session = Depends(get_db),
|
||||
) -> dict:
|
||||
"""
|
||||
@@ -278,7 +277,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
||||
)
|
||||
|
||||
# 构建所有格式的 provider_count 映射
|
||||
all_formats: Dict[str, int] = {}
|
||||
all_formats: dict[str, int] = {}
|
||||
for api_format_enum, provider_count in active_formats:
|
||||
api_format = (
|
||||
api_format_enum.value if hasattr(api_format_enum, "value") else str(api_format_enum)
|
||||
@@ -295,7 +294,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
||||
)
|
||||
.all()
|
||||
)
|
||||
endpoint_map: Dict[str, List[str]] = defaultdict(list)
|
||||
endpoint_map: dict[str, list[str]] = defaultdict(list)
|
||||
active_provider_formats: set[tuple[str, str]] = set()
|
||||
for api_format_enum, endpoint_id, provider_id in endpoint_rows:
|
||||
api_format = (
|
||||
@@ -305,7 +304,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
||||
active_provider_formats.add((str(provider_id), api_format))
|
||||
|
||||
# 1.2 统计每个 API 格式可用的活跃 Key 数量(Key 属于 Provider,通过 api_formats 关联格式)
|
||||
key_counts: Dict[str, int] = {}
|
||||
key_counts: dict[str, int] = {}
|
||||
if active_provider_formats:
|
||||
active_provider_keys = (
|
||||
db.query(ProviderAPIKey.provider_id, ProviderAPIKey.api_formats)
|
||||
@@ -342,7 +341,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
||||
)
|
||||
|
||||
# 构建每个格式的状态统计
|
||||
status_counts: Dict[str, Dict[str, int]] = {}
|
||||
status_counts: dict[str, dict[str, int]] = {}
|
||||
for api_format_enum, status, count in status_counts_query:
|
||||
api_format = (
|
||||
api_format_enum.value if hasattr(api_format_enum, "value") else str(api_format_enum)
|
||||
@@ -370,7 +369,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
||||
.all()
|
||||
)
|
||||
|
||||
grouped_attempts: Dict[str, List[RequestCandidate]] = {}
|
||||
grouped_attempts: dict[str, list[RequestCandidate]] = {}
|
||||
|
||||
for attempt, api_format_enum, provider_id in rows:
|
||||
api_format = (
|
||||
@@ -384,7 +383,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
||||
grouped_attempts[api_format].append(attempt)
|
||||
|
||||
# 4. 为所有活跃格式生成监控数据(包括没有请求记录的)
|
||||
monitors: List[ApiFormatHealthMonitor] = []
|
||||
monitors: list[ApiFormatHealthMonitor] = []
|
||||
for api_format in all_formats:
|
||||
attempts = grouped_attempts.get(api_format, [])
|
||||
# 获取窗口内的真实统计数据
|
||||
@@ -399,7 +398,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
||||
|
||||
# 时间线按时间正序
|
||||
attempts_sorted = list(reversed(attempts))
|
||||
events: List[EndpointHealthEvent] = []
|
||||
events: list[EndpointHealthEvent] = []
|
||||
for attempt in attempts_sorted:
|
||||
event_timestamp = attempt.finished_at or attempt.started_at or attempt.created_at
|
||||
events.append(
|
||||
@@ -462,7 +461,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
||||
@dataclass
|
||||
class AdminKeyHealthAdapter(AdminApiAdapter):
|
||||
key_id: str
|
||||
api_format: Optional[str] = None
|
||||
api_format: str | None = None
|
||||
|
||||
async def handle(self, context): # type: ignore[override]
|
||||
health_data = health_monitor.get_key_health(context.db, self.key_id, self.api_format)
|
||||
@@ -500,7 +499,7 @@ class AdminKeyHealthAdapter(AdminApiAdapter):
|
||||
@dataclass
|
||||
class AdminRecoverKeyHealthAdapter(AdminApiAdapter):
|
||||
key_id: str
|
||||
api_format: Optional[str] = None
|
||||
api_format: str | None = None
|
||||
|
||||
async def handle(self, context): # type: ignore[override]
|
||||
db = context.db
|
||||
|
||||
@@ -6,7 +6,6 @@ import json
|
||||
import uuid
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from typing import Dict, List
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, Request
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -146,14 +145,14 @@ async def delete_endpoint_key(
|
||||
# ========== Provider Keys API ==========
|
||||
|
||||
|
||||
@router.get("/providers/{provider_id}/keys", response_model=List[EndpointAPIKeyResponse])
|
||||
@router.get("/providers/{provider_id}/keys", response_model=list[EndpointAPIKeyResponse])
|
||||
async def list_provider_keys(
|
||||
provider_id: str,
|
||||
request: Request,
|
||||
skip: int = Query(0, ge=0, description="跳过的记录数"),
|
||||
limit: int = Query(100, ge=1, le=1000, description="返回的最大记录数"),
|
||||
db: Session = Depends(get_db),
|
||||
) -> List[EndpointAPIKeyResponse]:
|
||||
) -> list[EndpointAPIKeyResponse]:
|
||||
"""
|
||||
获取 Provider 的所有 Keys
|
||||
|
||||
@@ -503,12 +502,12 @@ class AdminGetKeysGroupedByFormatAdapter(AdminApiAdapter):
|
||||
)
|
||||
.all()
|
||||
)
|
||||
endpoint_base_url_map: Dict[tuple[str, str], str] = {}
|
||||
endpoint_base_url_map: dict[tuple[str, str], str] = {}
|
||||
for provider_id, api_format, base_url in endpoints:
|
||||
fmt = api_format.value if hasattr(api_format, "value") else str(api_format)
|
||||
endpoint_base_url_map[(str(provider_id), fmt)] = base_url
|
||||
|
||||
grouped: Dict[str, List[dict]] = {}
|
||||
grouped: dict[str, list[dict]] = {}
|
||||
for key, provider in keys:
|
||||
api_formats = key.api_formats or []
|
||||
|
||||
|
||||
@@ -5,10 +5,9 @@ ProviderEndpoint CRUD 管理 API
|
||||
import uuid
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from typing import List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, Request
|
||||
from sqlalchemy import and_, func
|
||||
from sqlalchemy import and_
|
||||
from sqlalchemy.orm import Session
|
||||
from sqlalchemy.orm.attributes import flag_modified
|
||||
|
||||
@@ -29,7 +28,7 @@ router = APIRouter(tags=["Endpoint Management"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
|
||||
|
||||
def mask_proxy_password(proxy_config: Optional[dict]) -> Optional[dict]:
|
||||
def mask_proxy_password(proxy_config: dict | None) -> dict | None:
|
||||
"""对代理配置中的密码进行脱敏处理"""
|
||||
if not proxy_config:
|
||||
return None
|
||||
@@ -39,14 +38,14 @@ def mask_proxy_password(proxy_config: Optional[dict]) -> Optional[dict]:
|
||||
return masked
|
||||
|
||||
|
||||
@router.get("/providers/{provider_id}/endpoints", response_model=List[ProviderEndpointResponse])
|
||||
@router.get("/providers/{provider_id}/endpoints", response_model=list[ProviderEndpointResponse])
|
||||
async def list_provider_endpoints(
|
||||
provider_id: str,
|
||||
request: Request,
|
||||
skip: int = Query(0, ge=0, description="跳过的记录数"),
|
||||
limit: int = Query(100, ge=1, le=1000, description="返回的最大记录数"),
|
||||
db: Session = Depends(get_db),
|
||||
) -> List[ProviderEndpointResponse]:
|
||||
) -> list[ProviderEndpointResponse]:
|
||||
"""
|
||||
获取指定 Provider 的所有 Endpoints
|
||||
|
||||
@@ -245,7 +244,7 @@ class AdminListProviderEndpointsAdapter(AdminApiAdapter):
|
||||
if is_active:
|
||||
active_keys_map[fmt] = active_keys_map.get(fmt, 0) + 1
|
||||
|
||||
result: List[ProviderEndpointResponse] = []
|
||||
result: list[ProviderEndpointResponse] = []
|
||||
for endpoint in endpoints:
|
||||
endpoint_format = (
|
||||
endpoint.api_format
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
"""LDAP配置管理API端点。"""
|
||||
|
||||
import re
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, Request
|
||||
from pydantic import BaseModel, Field, ValidationError, field_validator
|
||||
@@ -30,9 +30,9 @@ BCRYPT_HASH_PATTERN = re.compile(r"^\$2[aby]\$\d{2}\$.{53}$")
|
||||
class LDAPConfigResponse(BaseModel):
|
||||
"""LDAP配置响应(不返回密码)"""
|
||||
|
||||
server_url: Optional[str] = None
|
||||
bind_dn: Optional[str] = None
|
||||
base_dn: Optional[str] = None
|
||||
server_url: str | None = None
|
||||
bind_dn: str | None = None
|
||||
base_dn: str | None = None
|
||||
has_bind_password: bool = False
|
||||
user_search_filter: str
|
||||
username_attr: str
|
||||
@@ -50,7 +50,7 @@ class LDAPConfigUpdate(BaseModel):
|
||||
server_url: str = Field(..., min_length=1, max_length=255)
|
||||
bind_dn: str = Field(..., min_length=1, max_length=255)
|
||||
# 允许空字符串表示"清除密码";非空时自动 strip 并校验不能为空
|
||||
bind_password: Optional[str] = Field(None, max_length=1024)
|
||||
bind_password: str | None = Field(None, max_length=1024)
|
||||
base_dn: str = Field(..., min_length=1, max_length=255)
|
||||
user_search_filter: str = Field(default="(uid={username})", max_length=500)
|
||||
username_attr: str = Field(default="uid", max_length=50)
|
||||
@@ -63,7 +63,7 @@ class LDAPConfigUpdate(BaseModel):
|
||||
|
||||
@field_validator("bind_password")
|
||||
@classmethod
|
||||
def validate_bind_password(cls, v: Optional[str]) -> Optional[str]:
|
||||
def validate_bind_password(cls, v: str | None) -> str | None:
|
||||
if v is None or v == "":
|
||||
return v
|
||||
v = v.strip()
|
||||
@@ -114,22 +114,22 @@ class LDAPTestResponse(BaseModel):
|
||||
class LDAPConfigTest(BaseModel):
|
||||
"""LDAP配置测试请求(全部可选,用于临时覆盖)"""
|
||||
|
||||
server_url: Optional[str] = Field(None, min_length=1, max_length=255)
|
||||
bind_dn: Optional[str] = Field(None, min_length=1, max_length=255)
|
||||
bind_password: Optional[str] = Field(None, min_length=1)
|
||||
base_dn: Optional[str] = Field(None, min_length=1, max_length=255)
|
||||
user_search_filter: Optional[str] = Field(None, max_length=500)
|
||||
username_attr: Optional[str] = Field(None, max_length=50)
|
||||
email_attr: Optional[str] = Field(None, max_length=50)
|
||||
display_name_attr: Optional[str] = Field(None, max_length=50)
|
||||
is_enabled: Optional[bool] = None
|
||||
is_exclusive: Optional[bool] = None
|
||||
use_starttls: Optional[bool] = None
|
||||
connect_timeout: Optional[int] = Field(None, ge=1, le=60)
|
||||
server_url: str | None = Field(None, min_length=1, max_length=255)
|
||||
bind_dn: str | None = Field(None, min_length=1, max_length=255)
|
||||
bind_password: str | None = Field(None, min_length=1)
|
||||
base_dn: str | None = Field(None, min_length=1, max_length=255)
|
||||
user_search_filter: str | None = Field(None, max_length=500)
|
||||
username_attr: str | None = Field(None, max_length=50)
|
||||
email_attr: str | None = Field(None, max_length=50)
|
||||
display_name_attr: str | None = Field(None, max_length=50)
|
||||
is_enabled: bool | None = None
|
||||
is_exclusive: bool | None = None
|
||||
use_starttls: bool | None = None
|
||||
connect_timeout: int | None = Field(None, ge=1, le=60)
|
||||
|
||||
@field_validator("user_search_filter")
|
||||
@classmethod
|
||||
def validate_search_filter(cls, v: Optional[str]) -> Optional[str]:
|
||||
def validate_search_filter(cls, v: str | None) -> str | None:
|
||||
if v is None:
|
||||
return v
|
||||
if "{username}" not in v:
|
||||
@@ -263,7 +263,7 @@ async def test_ldap_connection(request: Request, db: Session = Depends(get_db))
|
||||
|
||||
|
||||
class AdminGetLDAPConfigAdapter(AdminApiAdapter):
|
||||
async def handle(self, context) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context) -> dict[str, Any]: # type: ignore[override]
|
||||
db = context.db
|
||||
config = db.query(LDAPConfig).first()
|
||||
|
||||
@@ -300,7 +300,7 @@ class AdminGetLDAPConfigAdapter(AdminApiAdapter):
|
||||
|
||||
|
||||
class AdminUpdateLDAPConfigAdapter(AdminApiAdapter):
|
||||
async def handle(self, context) -> Dict[str, str]: # type: ignore[override]
|
||||
async def handle(self, context) -> dict[str, str]: # type: ignore[override]
|
||||
db = context.db
|
||||
payload = context.ensure_json_body()
|
||||
|
||||
@@ -421,7 +421,7 @@ class AdminUpdateLDAPConfigAdapter(AdminApiAdapter):
|
||||
|
||||
|
||||
class AdminTestLDAPConnectionAdapter(AdminApiAdapter):
|
||||
async def handle(self, context) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context) -> dict[str, Any]: # type: ignore[override]
|
||||
from src.services.auth.ldap import LDAPService
|
||||
|
||||
db = context.db
|
||||
@@ -442,7 +442,7 @@ class AdminTestLDAPConnectionAdapter(AdminApiAdapter):
|
||||
raise InvalidRequestException(translate_pydantic_error(errors[0]))
|
||||
raise InvalidRequestException("请求数据验证失败")
|
||||
|
||||
config_data: Dict[str, Any] = {}
|
||||
config_data: dict[str, Any] = {}
|
||||
|
||||
if saved_config:
|
||||
config_data = {
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
"""管理员 Management Token 管理端点"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from fastapi.responses import JSONResponse
|
||||
@@ -46,8 +45,8 @@ class AdminManagementTokenApiAdapter(AdminApiAdapter):
|
||||
@router.get("")
|
||||
async def list_all_management_tokens(
|
||||
request: Request,
|
||||
user_id: Optional[str] = Query(None, description="筛选用户 ID"),
|
||||
is_active: Optional[bool] = Query(None, description="筛选激活状态"),
|
||||
user_id: str | None = Query(None, description="筛选用户 ID"),
|
||||
is_active: bool | None = Query(None, description="筛选激活状态"),
|
||||
skip: int = Query(0, ge=0),
|
||||
limit: int = Query(50, ge=1, le=100),
|
||||
db: Session = Depends(get_db),
|
||||
@@ -174,8 +173,8 @@ class AdminListManagementTokensAdapter(AdminManagementTokenApiAdapter):
|
||||
"""列出所有 Management Tokens"""
|
||||
|
||||
name: str = "admin_list_management_tokens"
|
||||
user_id: Optional[str] = None
|
||||
is_active: Optional[bool] = None
|
||||
user_id: str | None = None
|
||||
is_active: bool | None = None
|
||||
skip: int = 0
|
||||
limit: int = 50
|
||||
|
||||
@@ -197,7 +196,7 @@ class AdminListManagementTokensAdapter(AdminManagementTokenApiAdapter):
|
||||
)
|
||||
|
||||
# 预加载用户信息
|
||||
user_ids = list(set(t.user_id for t in tokens))
|
||||
user_ids = list({t.user_id for t in tokens})
|
||||
users = {u.id: u for u in context.db.query(User).filter(User.id.in_(user_ids)).all()}
|
||||
for token in tokens:
|
||||
token.user = users.get(token.user_id)
|
||||
|
||||
@@ -5,7 +5,6 @@
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Dict, List
|
||||
|
||||
from fastapi import APIRouter, Depends, Request
|
||||
from sqlalchemy.orm import Session, joinedload
|
||||
@@ -64,12 +63,12 @@ class AdminGetModelCatalogAdapter(AdminApiAdapter):
|
||||
db: Session = context.db
|
||||
|
||||
# 1. 获取所有活跃的 GlobalModel
|
||||
global_models: List[GlobalModel] = (
|
||||
global_models: list[GlobalModel] = (
|
||||
db.query(GlobalModel).filter(GlobalModel.is_active == True).all()
|
||||
)
|
||||
|
||||
# 2. 获取所有活跃的 Model 实现(包含 global_model 以便计算有效价格)
|
||||
models: List[Model] = (
|
||||
models: list[Model] = (
|
||||
db.query(Model)
|
||||
.options(joinedload(Model.provider), joinedload(Model.global_model))
|
||||
.filter(Model.is_active == True)
|
||||
@@ -77,17 +76,17 @@ class AdminGetModelCatalogAdapter(AdminApiAdapter):
|
||||
)
|
||||
|
||||
# 按 GlobalModel ID 组织关联提供商
|
||||
models_by_global_model: Dict[str, List[Model]] = {}
|
||||
models_by_global_model: dict[str, list[Model]] = {}
|
||||
for model in models:
|
||||
if model.global_model_id:
|
||||
models_by_global_model.setdefault(model.global_model_id, []).append(model)
|
||||
|
||||
# 3. 为每个 GlobalModel 构建 catalog item
|
||||
catalog_items: List[ModelCatalogItem] = []
|
||||
catalog_items: list[ModelCatalogItem] = []
|
||||
|
||||
for gm in global_models:
|
||||
gm_id = gm.id
|
||||
provider_entries: List[ModelCatalogProviderDetail] = []
|
||||
provider_entries: list[ModelCatalogProviderDetail] = []
|
||||
# 从 config JSON 读取能力标志
|
||||
gm_config = gm.config or {}
|
||||
capability_flags = {
|
||||
|
||||
@@ -3,8 +3,7 @@ models.dev 外部模型数据代理
|
||||
"""
|
||||
|
||||
import json
|
||||
|
||||
from typing import Any, Optional
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
@@ -43,7 +42,7 @@ OFFICIAL_PROVIDERS = {
|
||||
}
|
||||
|
||||
|
||||
async def _get_cached_data() -> Optional[dict[str, Any]]:
|
||||
async def _get_cached_data() -> dict[str, Any] | None:
|
||||
"""从 Redis 获取缓存数据"""
|
||||
redis = await get_redis_client()
|
||||
if redis is None:
|
||||
|
||||
@@ -5,7 +5,6 @@ GlobalModel Admin API
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, Request
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -37,8 +36,8 @@ async def list_global_models(
|
||||
request: Request,
|
||||
skip: int = Query(0, ge=0),
|
||||
limit: int = Query(100, ge=1, le=1000),
|
||||
is_active: Optional[bool] = Query(None),
|
||||
search: Optional[str] = Query(None),
|
||||
is_active: bool | None = Query(None),
|
||||
search: str | None = Query(None),
|
||||
db: Session = Depends(get_db),
|
||||
) -> GlobalModelListResponse:
|
||||
"""
|
||||
@@ -254,8 +253,8 @@ class AdminListGlobalModelsAdapter(AdminApiAdapter):
|
||||
|
||||
skip: int
|
||||
limit: int
|
||||
is_active: Optional[bool]
|
||||
search: Optional[str]
|
||||
is_active: bool | None
|
||||
search: str | None
|
||||
|
||||
async def handle(self, context): # type: ignore[override]
|
||||
from sqlalchemy import func
|
||||
|
||||
@@ -9,7 +9,6 @@ GlobalModel 请求链路预览 API
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, Request
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
@@ -47,22 +46,22 @@ class RoutingKeyInfo(BaseModel):
|
||||
name: str
|
||||
masked_key: str = Field("", description="脱敏的 API Key")
|
||||
internal_priority: int = Field(..., description="Key 内部优先级")
|
||||
global_priority_by_format: Optional[Dict[str, int]] = Field(None, description="按 API 格式的全局优先级")
|
||||
rpm_limit: Optional[int] = Field(None, description="RPM 限制,null 表示自适应")
|
||||
global_priority_by_format: dict[str, int] | None = Field(None, description="按 API 格式的全局优先级")
|
||||
rpm_limit: int | None = Field(None, description="RPM 限制,null 表示自适应")
|
||||
is_adaptive: bool = Field(False, description="是否为自适应 RPM 模式")
|
||||
effective_rpm: Optional[int] = Field(None, description="有效 RPM 限制")
|
||||
effective_rpm: int | None = Field(None, description="有效 RPM 限制")
|
||||
cache_ttl_minutes: int = Field(0, description="缓存 TTL(分钟)")
|
||||
health_score: float = Field(1.0, description="健康度分数(0-1 小数格式)")
|
||||
is_active: bool
|
||||
api_formats: List[str] = Field(default_factory=list, description="支持的 API 格式")
|
||||
api_formats: list[str] = Field(default_factory=list, description="支持的 API 格式")
|
||||
# 模型白名单
|
||||
allowed_models: Optional[List[str]] = Field(None, description="允许的模型列表,null 表示不限制")
|
||||
allowed_models: list[str] | None = Field(None, description="允许的模型列表,null 表示不限制")
|
||||
# 熔断状态
|
||||
circuit_breaker_open: bool = Field(False, description="熔断器是否打开")
|
||||
circuit_breaker_formats: List[str] = Field(
|
||||
circuit_breaker_formats: list[str] = Field(
|
||||
default_factory=list, description="熔断的 API 格式列表"
|
||||
)
|
||||
next_probe_at: Optional[str] = Field(None, description="下次探测时间(ISO格式)")
|
||||
next_probe_at: str | None = Field(None, description="下次探测时间(ISO格式)")
|
||||
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
@@ -73,9 +72,9 @@ class RoutingEndpointInfo(BaseModel):
|
||||
id: str
|
||||
api_format: str
|
||||
base_url: str
|
||||
custom_path: Optional[str] = None
|
||||
custom_path: str | None = None
|
||||
is_active: bool
|
||||
keys: List[RoutingKeyInfo] = Field(default_factory=list)
|
||||
keys: list[RoutingKeyInfo] = Field(default_factory=list)
|
||||
total_keys: int = 0
|
||||
active_keys: int = 0
|
||||
|
||||
@@ -87,7 +86,7 @@ class RoutingModelMapping(BaseModel):
|
||||
|
||||
name: str = Field(..., description="映射名称")
|
||||
priority: int = Field(..., description="优先级(数字越小优先级越高)")
|
||||
api_formats: Optional[List[str]] = Field(None, description="作用域(适用的 API 格式)")
|
||||
api_formats: list[str] | None = Field(None, description="作用域(适用的 API 格式)")
|
||||
|
||||
|
||||
class RoutingProviderInfo(BaseModel):
|
||||
@@ -97,18 +96,18 @@ class RoutingProviderInfo(BaseModel):
|
||||
name: str
|
||||
model_id: str = Field(..., description="Model ID(GlobalModel 与 Provider 的关联记录 ID)")
|
||||
provider_priority: int = Field(..., description="提供商优先级(数字越小优先级越高)")
|
||||
billing_type: Optional[str] = Field(None, description="计费类型")
|
||||
monthly_quota_usd: Optional[float] = Field(None, description="月额度(美元)")
|
||||
monthly_used_usd: Optional[float] = Field(None, description="已用额度(美元)")
|
||||
billing_type: str | None = Field(None, description="计费类型")
|
||||
monthly_quota_usd: float | None = Field(None, description="月额度(美元)")
|
||||
monthly_used_usd: float | None = Field(None, description="已用额度(美元)")
|
||||
is_active: bool
|
||||
# 模型映射信息
|
||||
provider_model_name: str = Field(..., description="提供商侧的模型名称")
|
||||
model_mappings: List[RoutingModelMapping] = Field(
|
||||
model_mappings: list[RoutingModelMapping] = Field(
|
||||
default_factory=list, description="模型名称映射列表"
|
||||
)
|
||||
model_is_active: bool = Field(True, description="Model 是否活跃")
|
||||
# Endpoint 和 Key 信息
|
||||
endpoints: List[RoutingEndpointInfo] = Field(default_factory=list)
|
||||
endpoints: list[RoutingEndpointInfo] = Field(default_factory=list)
|
||||
total_endpoints: int = 0
|
||||
active_endpoints: int = 0
|
||||
|
||||
@@ -123,7 +122,7 @@ class GlobalKeyWhitelistItem(BaseModel):
|
||||
masked_key: str = Field(..., description="脱敏的 API Key")
|
||||
provider_id: str = Field(..., description="Provider ID")
|
||||
provider_name: str = Field(..., description="Provider 名称")
|
||||
allowed_models: List[str] = Field(default_factory=list, description="Key 白名单模型列表")
|
||||
allowed_models: list[str] = Field(default_factory=list, description="Key 白名单模型列表")
|
||||
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
@@ -136,11 +135,11 @@ class ModelRoutingPreviewResponse(BaseModel):
|
||||
display_name: str
|
||||
is_active: bool
|
||||
# GlobalModel 的模型映射(用于前端匹配 Key 白名单)
|
||||
global_model_mappings: List[str] = Field(
|
||||
global_model_mappings: list[str] = Field(
|
||||
default_factory=list, description="GlobalModel 的模型映射规则(正则模式)"
|
||||
)
|
||||
# 链路信息
|
||||
providers: List[RoutingProviderInfo] = Field(
|
||||
providers: list[RoutingProviderInfo] = Field(
|
||||
default_factory=list, description="按优先级排序的提供商列表"
|
||||
)
|
||||
total_providers: int = 0
|
||||
@@ -149,7 +148,7 @@ class ModelRoutingPreviewResponse(BaseModel):
|
||||
scheduling_mode: str = Field("cache_affinity", description="调度模式")
|
||||
priority_mode: str = Field("provider", description="优先级模式")
|
||||
# 全局 Key 白名单数据(供前端实时匹配,包含所有 Provider 的 Key)
|
||||
all_keys_whitelist: List[GlobalKeyWhitelistItem] = Field(
|
||||
all_keys_whitelist: list[GlobalKeyWhitelistItem] = Field(
|
||||
default_factory=list, description="所有 Provider 的 Key 白名单数据"
|
||||
)
|
||||
|
||||
@@ -227,7 +226,7 @@ class AdminGetModelRoutingPreviewAdapter(AdminApiAdapter):
|
||||
provider_ids = [m.provider_id for m in models if m.provider_id]
|
||||
|
||||
# 批量获取 Provider 的 Endpoints
|
||||
endpoints_by_provider: Dict[str, List[ProviderEndpoint]] = {}
|
||||
endpoints_by_provider: dict[str, list[ProviderEndpoint]] = {}
|
||||
if provider_ids:
|
||||
endpoints = (
|
||||
db.query(ProviderEndpoint)
|
||||
@@ -240,7 +239,7 @@ class AdminGetModelRoutingPreviewAdapter(AdminApiAdapter):
|
||||
endpoints_by_provider[ep.provider_id].append(ep)
|
||||
|
||||
# 批量获取 Provider 的 Keys
|
||||
keys_by_provider: Dict[str, List[ProviderAPIKey]] = {}
|
||||
keys_by_provider: dict[str, list[ProviderAPIKey]] = {}
|
||||
if provider_ids:
|
||||
keys = (
|
||||
db.query(ProviderAPIKey).filter(ProviderAPIKey.provider_id.in_(provider_ids)).all()
|
||||
@@ -251,14 +250,14 @@ class AdminGetModelRoutingPreviewAdapter(AdminApiAdapter):
|
||||
keys_by_provider[key.provider_id].append(key)
|
||||
|
||||
# 提取 GlobalModel 的 model_mappings(用于 Key 白名单匹配)
|
||||
global_model_mappings: List[str] = []
|
||||
global_model_mappings: list[str] = []
|
||||
if global_model.config and isinstance(global_model.config, dict):
|
||||
mappings = global_model.config.get("model_mappings")
|
||||
if isinstance(mappings, list):
|
||||
global_model_mappings = [m for m in mappings if isinstance(m, str)]
|
||||
|
||||
# 构建 Provider 路由信息
|
||||
provider_infos: List[RoutingProviderInfo] = []
|
||||
provider_infos: list[RoutingProviderInfo] = []
|
||||
for model in models:
|
||||
provider = model.provider
|
||||
if not provider:
|
||||
@@ -281,7 +280,7 @@ class AdminGetModelRoutingPreviewAdapter(AdminApiAdapter):
|
||||
provider_keys = keys_by_provider.get(provider.id, [])
|
||||
|
||||
# 按 api_format 组织 Keys
|
||||
keys_by_endpoint: Dict[str, List[ProviderAPIKey]] = {}
|
||||
keys_by_endpoint: dict[str, list[ProviderAPIKey]] = {}
|
||||
for key in provider_keys:
|
||||
# 每个 Key 可能支持多个 api_formats
|
||||
for fmt in key.api_formats or []:
|
||||
@@ -355,8 +354,8 @@ class AdminGetModelRoutingPreviewAdapter(AdminApiAdapter):
|
||||
|
||||
# 检查熔断状态
|
||||
circuit_breaker_open = False
|
||||
circuit_breaker_formats: List[str] = []
|
||||
next_probe_at: Optional[str] = None
|
||||
circuit_breaker_formats: list[str] = []
|
||||
next_probe_at: str | None = None
|
||||
if key.circuit_breaker_by_format:
|
||||
for fmt, cb_state in key.circuit_breaker_by_format.items():
|
||||
if isinstance(cb_state, dict) and cb_state.get("open"):
|
||||
@@ -462,7 +461,7 @@ class AdminGetModelRoutingPreviewAdapter(AdminApiAdapter):
|
||||
)
|
||||
|
||||
# 获取所有活跃 Provider 的 Key 白名单数据(供前端实时匹配)
|
||||
all_keys_whitelist: List[GlobalKeyWhitelistItem] = []
|
||||
all_keys_whitelist: list[GlobalKeyWhitelistItem] = []
|
||||
crypto = CryptoService()
|
||||
|
||||
# 获取所有活跃的 Key(带白名单),使用 selectinload 避免 N+1 查询
|
||||
|
||||
@@ -1,7 +1,8 @@
|
||||
"""模块管理 API 端点"""
|
||||
|
||||
from __future__ import annotations
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, Request
|
||||
from pydantic import BaseModel
|
||||
@@ -28,18 +29,18 @@ class ModuleStatusResponse(BaseModel):
|
||||
enabled: bool
|
||||
active: bool
|
||||
config_validated: bool
|
||||
config_error: Optional[str]
|
||||
config_error: str | None
|
||||
display_name: str
|
||||
description: str
|
||||
category: str
|
||||
admin_route: Optional[str]
|
||||
admin_menu_icon: Optional[str]
|
||||
admin_menu_group: Optional[str]
|
||||
admin_route: str | None
|
||||
admin_menu_icon: str | None
|
||||
admin_menu_group: str | None
|
||||
admin_menu_order: int
|
||||
health: str
|
||||
|
||||
@classmethod
|
||||
def from_status(cls, status: ModuleStatus) -> "ModuleStatusResponse":
|
||||
def from_status(cls, status: ModuleStatus) -> ModuleStatusResponse:
|
||||
return cls(
|
||||
name=status.name,
|
||||
available=status.available,
|
||||
@@ -130,7 +131,7 @@ async def set_module_enabled(
|
||||
class AdminGetAllModulesStatusAdapter(AdminApiAdapter):
|
||||
"""获取所有模块状态"""
|
||||
|
||||
async def handle(self, context) -> Dict[str, Any]:
|
||||
async def handle(self, context) -> dict[str, Any]:
|
||||
registry = get_module_registry()
|
||||
all_status = await registry.get_all_status_async(context.db)
|
||||
|
||||
@@ -146,7 +147,7 @@ class AdminGetModuleStatusAdapter(AdminApiAdapter):
|
||||
|
||||
module_name: str
|
||||
|
||||
async def handle(self, context) -> Dict[str, Any]:
|
||||
async def handle(self, context) -> dict[str, Any]:
|
||||
registry = get_module_registry()
|
||||
status = await registry.get_module_status_async(self.module_name, context.db)
|
||||
|
||||
@@ -162,7 +163,7 @@ class AdminSetModuleEnabledAdapter(AdminApiAdapter):
|
||||
|
||||
module_name: str
|
||||
|
||||
async def handle(self, context) -> Dict[str, Any]:
|
||||
async def handle(self, context) -> dict[str, Any]:
|
||||
registry = get_module_registry()
|
||||
|
||||
# 检查模块是否存在
|
||||
|
||||
@@ -2,7 +2,6 @@
|
||||
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from sqlalchemy import func
|
||||
@@ -33,8 +32,8 @@ pipeline = ApiRequestPipeline()
|
||||
@router.get("/audit-logs")
|
||||
async def get_audit_logs(
|
||||
request: Request,
|
||||
username: Optional[str] = Query(None, description="用户名筛选 (模糊匹配)"),
|
||||
event_type: Optional[str] = Query(None, description="事件类型筛选"),
|
||||
username: str | None = Query(None, description="用户名筛选 (模糊匹配)"),
|
||||
event_type: str | None = Query(None, description="事件类型筛选"),
|
||||
days: int = Query(7, description="查询天数"),
|
||||
limit: int = Query(100, description="返回数量限制"),
|
||||
offset: int = Query(0, description="偏移量"),
|
||||
@@ -212,8 +211,8 @@ async def get_circuit_history(
|
||||
|
||||
@dataclass
|
||||
class AdminGetAuditLogsAdapter(AdminApiAdapter):
|
||||
username: Optional[str]
|
||||
event_type: Optional[str]
|
||||
username: str | None
|
||||
event_type: str | None
|
||||
days: int
|
||||
limit: int
|
||||
offset: int
|
||||
@@ -497,8 +496,8 @@ class AdminCircuitHistoryAdapter(AdminApiAdapter):
|
||||
return {"items": history, "count": len(history)}
|
||||
|
||||
|
||||
def _get_health_recommendations(error_stats: dict, health_score: int) -> List[str]:
|
||||
recommendations: List[str] = []
|
||||
def _get_health_recommendations(error_stats: dict, health_score: int) -> list[str]:
|
||||
recommendations: list[str] = []
|
||||
if health_score < 50:
|
||||
recommendations.append("系统健康状况严重,请立即检查错误日志")
|
||||
if error_stats.get("total_errors", 0) > 100:
|
||||
|
||||
@@ -5,7 +5,7 @@
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from fastapi.responses import PlainTextResponse
|
||||
@@ -13,7 +13,7 @@ from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pagination import PaginationMeta, build_pagination_payload, paginate_sequence
|
||||
from src.api.base.pagination import build_pagination_payload, paginate_sequence
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.clients.redis_client import get_redis_client_sync
|
||||
from src.core.crypto import crypto_service
|
||||
@@ -28,7 +28,7 @@ router = APIRouter(prefix="/api/admin/monitoring/cache", tags=["Admin - Monitori
|
||||
pipeline = ApiRequestPipeline()
|
||||
|
||||
|
||||
def mask_api_key(api_key: Optional[str], prefix_len: int = 8, suffix_len: int = 4) -> Optional[str]:
|
||||
def mask_api_key(api_key: str | None, prefix_len: int = 8, suffix_len: int = 4) -> str | None:
|
||||
"""
|
||||
脱敏 API Key,显示前缀 + 星号 + 后缀
|
||||
例如: sk-jhiId-xxxxxxxxxxxAABB -> sk-jhiId-********AABB
|
||||
@@ -47,7 +47,7 @@ def mask_api_key(api_key: Optional[str], prefix_len: int = 8, suffix_len: int =
|
||||
return f"{api_key[:prefix_len]}********{api_key[-suffix_len:]}"
|
||||
|
||||
|
||||
def decrypt_and_mask(encrypted_key: Optional[str], prefix_len: int = 8) -> Optional[str]:
|
||||
def decrypt_and_mask(encrypted_key: str | None, prefix_len: int = 8) -> str | None:
|
||||
"""
|
||||
解密 API Key 后脱敏显示
|
||||
|
||||
@@ -65,7 +65,7 @@ def decrypt_and_mask(encrypted_key: Optional[str], prefix_len: int = 8) -> Optio
|
||||
return None
|
||||
|
||||
|
||||
def resolve_user_identifier(db: Session, identifier: str) -> Optional[str]:
|
||||
def resolve_user_identifier(db: Session, identifier: str) -> str | None:
|
||||
"""
|
||||
将用户标识符(username/email/user_id/api_key_id)解析为 user_id
|
||||
|
||||
@@ -181,7 +181,7 @@ async def get_user_affinity(
|
||||
@router.get("/affinities")
|
||||
async def list_affinities(
|
||||
request: Request,
|
||||
keyword: Optional[str] = None,
|
||||
keyword: str | None = None,
|
||||
limit: int = Query(100, ge=1, le=1000, description="返回数量限制"),
|
||||
offset: int = Query(0, ge=0, description="偏移量"),
|
||||
db: Session = Depends(get_db),
|
||||
@@ -421,7 +421,7 @@ async def get_cache_metrics(
|
||||
|
||||
|
||||
class AdminCacheStatsAdapter(AdminApiAdapter):
|
||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
try:
|
||||
redis_client = get_redis_client_sync()
|
||||
# 读取系统配置,确保监控接口与编排器使用一致的模式
|
||||
@@ -487,14 +487,14 @@ class AdminCacheMetricsAdapter(AdminApiAdapter):
|
||||
logger.exception(f"导出缓存指标失败: {exc}")
|
||||
raise HTTPException(status_code=500, detail=f"导出缓存指标失败: {exc}")
|
||||
|
||||
def _format_prometheus(self, stats: Dict[str, Any]) -> str:
|
||||
def _format_prometheus(self, stats: dict[str, Any]) -> str:
|
||||
"""
|
||||
将 scheduler/affinity 指标转换为 Prometheus 文本格式。
|
||||
"""
|
||||
scheduler_metrics = stats.get("scheduler_metrics", {})
|
||||
affinity_stats = stats.get("affinity_stats", {})
|
||||
|
||||
metric_map: List[Tuple[str, str, float]] = [
|
||||
metric_map: list[tuple[str, str, float]] = [
|
||||
(
|
||||
"cache_scheduler_total_batches",
|
||||
"Total batches pulled from provider list",
|
||||
@@ -542,7 +542,7 @@ class AdminCacheMetricsAdapter(AdminApiAdapter):
|
||||
),
|
||||
]
|
||||
|
||||
affinity_map: List[Tuple[str, str, float]] = [
|
||||
affinity_map: list[tuple[str, str, float]] = [
|
||||
(
|
||||
"cache_affinity_total",
|
||||
"Total cache affinities stored",
|
||||
@@ -596,7 +596,7 @@ class AdminCacheMetricsAdapter(AdminApiAdapter):
|
||||
class AdminGetUserAffinityAdapter(AdminApiAdapter):
|
||||
user_identifier: str
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
db = context.db
|
||||
try:
|
||||
user_id = resolve_user_identifier(db, self.user_identifier)
|
||||
@@ -673,11 +673,11 @@ class AdminGetUserAffinityAdapter(AdminApiAdapter):
|
||||
|
||||
@dataclass
|
||||
class AdminListAffinitiesAdapter(AdminApiAdapter):
|
||||
keyword: Optional[str]
|
||||
keyword: str | None
|
||||
limit: int
|
||||
offset: int
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
db = context.db
|
||||
redis_client = get_redis_client_sync()
|
||||
if not redis_client:
|
||||
@@ -686,7 +686,7 @@ class AdminListAffinitiesAdapter(AdminApiAdapter):
|
||||
affinity_mgr = await get_affinity_manager(redis_client)
|
||||
matched_user_id = None
|
||||
matched_api_key_id = None
|
||||
raw_affinities: List[Dict[str, Any]] = []
|
||||
raw_affinities: list[dict[str, Any]] = []
|
||||
|
||||
if self.keyword:
|
||||
# 首先检查是否是 API Key ID(affinity_key)
|
||||
@@ -724,14 +724,14 @@ class AdminListAffinitiesAdapter(AdminApiAdapter):
|
||||
}
|
||||
|
||||
# 批量查询用户 API Key 信息
|
||||
user_api_key_map: Dict[str, ApiKey] = {}
|
||||
user_api_key_map: dict[str, ApiKey] = {}
|
||||
if affinity_keys:
|
||||
user_api_keys = db.query(ApiKey).filter(ApiKey.id.in_(list(affinity_keys))).all()
|
||||
user_api_key_map = {str(k.id): k for k in user_api_keys}
|
||||
|
||||
# 收集所有 user_id
|
||||
user_ids = {str(k.user_id) for k in user_api_key_map.values()}
|
||||
user_map: Dict[str, User] = {}
|
||||
user_map: dict[str, User] = {}
|
||||
if user_ids:
|
||||
users = db.query(User).filter(User.id.in_(list(user_ids))).all()
|
||||
user_map = {str(user.id): user for user in users}
|
||||
@@ -771,7 +771,7 @@ class AdminListAffinitiesAdapter(AdminApiAdapter):
|
||||
global_model_ids = {
|
||||
item.get("model_name") for item in raw_affinities if item.get("model_name")
|
||||
}
|
||||
global_model_map: Dict[str, GlobalModel] = {}
|
||||
global_model_map: dict[str, GlobalModel] = {}
|
||||
if global_model_ids:
|
||||
# model_name 可能是 UUID 格式的 global_model_id,也可能是原始模型名称
|
||||
global_models = db.query(GlobalModel).filter(
|
||||
@@ -885,7 +885,7 @@ class AdminListAffinitiesAdapter(AdminApiAdapter):
|
||||
class AdminClearUserCacheAdapter(AdminApiAdapter):
|
||||
user_identifier: str
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
db = context.db
|
||||
try:
|
||||
redis_client = get_redis_client_sync()
|
||||
@@ -995,7 +995,7 @@ class AdminClearSingleAffinityAdapter(AdminApiAdapter):
|
||||
model_id: str
|
||||
api_format: str
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
db = context.db
|
||||
try:
|
||||
redis_client = get_redis_client_sync()
|
||||
@@ -1048,7 +1048,7 @@ class AdminClearSingleAffinityAdapter(AdminApiAdapter):
|
||||
|
||||
|
||||
class AdminClearAllCacheAdapter(AdminApiAdapter):
|
||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
try:
|
||||
redis_client = get_redis_client_sync()
|
||||
affinity_mgr = await get_affinity_manager(redis_client)
|
||||
@@ -1068,7 +1068,7 @@ class AdminClearAllCacheAdapter(AdminApiAdapter):
|
||||
class AdminClearProviderCacheAdapter(AdminApiAdapter):
|
||||
provider_id: str
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
try:
|
||||
redis_client = get_redis_client_sync()
|
||||
affinity_mgr = await get_affinity_manager(redis_client)
|
||||
@@ -1091,7 +1091,7 @@ class AdminClearProviderCacheAdapter(AdminApiAdapter):
|
||||
|
||||
|
||||
class AdminCacheConfigAdapter(AdminApiAdapter):
|
||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
from src.services.cache.affinity_manager import CacheAffinityManager
|
||||
from src.config.constants import ConcurrencyDefaults
|
||||
from src.services.rate_limit.adaptive_reservation import get_adaptive_reservation_manager
|
||||
@@ -1260,7 +1260,7 @@ async def clear_provider_model_mapping_cache(
|
||||
|
||||
|
||||
class AdminModelMappingCacheStatsAdapter(AdminApiAdapter):
|
||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
import json
|
||||
|
||||
from src.clients.redis_client import get_redis_client
|
||||
@@ -1510,7 +1510,7 @@ class AdminModelMappingCacheStatsAdapter(AdminApiAdapter):
|
||||
|
||||
|
||||
class AdminClearAllModelMappingCacheAdapter(AdminApiAdapter):
|
||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
from src.clients.redis_client import get_redis_client
|
||||
|
||||
try:
|
||||
@@ -1552,7 +1552,7 @@ class AdminClearAllModelMappingCacheAdapter(AdminApiAdapter):
|
||||
class AdminClearModelMappingCacheByNameAdapter(AdminApiAdapter):
|
||||
model_name: str
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
from src.clients.redis_client import get_redis_client
|
||||
|
||||
try:
|
||||
@@ -1599,7 +1599,7 @@ class AdminClearProviderModelMappingCacheAdapter(AdminApiAdapter):
|
||||
provider_id: str
|
||||
global_model_id: str
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
from src.clients.redis_client import get_redis_client
|
||||
|
||||
try:
|
||||
|
||||
@@ -4,7 +4,6 @@
|
||||
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from typing import List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
@@ -28,29 +27,29 @@ class CandidateResponse(BaseModel):
|
||||
request_id: str
|
||||
candidate_index: int
|
||||
retry_index: int = 0 # 重试序号(从0开始)
|
||||
provider_id: Optional[str] = None
|
||||
provider_name: Optional[str] = None
|
||||
provider_website: Optional[str] = None # Provider 官网
|
||||
endpoint_id: Optional[str] = None
|
||||
endpoint_name: Optional[str] = None # 端点显示名称(api_format)
|
||||
key_id: Optional[str] = None
|
||||
key_name: Optional[str] = None # 密钥名称
|
||||
key_preview: Optional[str] = None # 密钥脱敏预览(如 sk-***abc)
|
||||
key_capabilities: Optional[dict] = None # Key 支持的能力
|
||||
required_capabilities: Optional[dict] = None # 请求实际需要的能力标签
|
||||
provider_id: str | None = None
|
||||
provider_name: str | None = None
|
||||
provider_website: str | None = None # Provider 官网
|
||||
endpoint_id: str | None = None
|
||||
endpoint_name: str | None = None # 端点显示名称(api_format)
|
||||
key_id: str | None = None
|
||||
key_name: str | None = None # 密钥名称
|
||||
key_preview: str | None = None # 密钥脱敏预览(如 sk-***abc)
|
||||
key_capabilities: dict | None = None # Key 支持的能力
|
||||
required_capabilities: dict | None = None # 请求实际需要的能力标签
|
||||
status: str # 'pending', 'success', 'failed', 'skipped'
|
||||
skip_reason: Optional[str] = None
|
||||
skip_reason: str | None = None
|
||||
is_cached: bool = False
|
||||
# 执行结果字段
|
||||
status_code: Optional[int] = None
|
||||
error_type: Optional[str] = None
|
||||
error_message: Optional[str] = None
|
||||
latency_ms: Optional[int] = None
|
||||
concurrent_requests: Optional[int] = None
|
||||
extra_data: Optional[dict] = None
|
||||
status_code: int | None = None
|
||||
error_type: str | None = None
|
||||
error_message: str | None = None
|
||||
latency_ms: int | None = None
|
||||
concurrent_requests: int | None = None
|
||||
extra_data: dict | None = None
|
||||
created_at: datetime
|
||||
started_at: Optional[datetime] = None
|
||||
finished_at: Optional[datetime] = None
|
||||
started_at: datetime | None = None
|
||||
finished_at: datetime | None = None
|
||||
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
@@ -62,7 +61,7 @@ class RequestTraceResponse(BaseModel):
|
||||
total_candidates: int
|
||||
final_status: str # 'success', 'failed', 'cancelled', 'streaming', 'pending'
|
||||
total_latency_ms: int
|
||||
candidates: List[CandidateResponse]
|
||||
candidates: list[CandidateResponse]
|
||||
|
||||
|
||||
@router.get("/{request_id}", response_model=RequestTraceResponse)
|
||||
@@ -253,7 +252,7 @@ class AdminGetRequestTraceAdapter(AdminApiAdapter):
|
||||
key_preview_map[k.id] = "***"
|
||||
|
||||
# 构建 candidate 响应列表
|
||||
candidate_responses: List[CandidateResponse] = []
|
||||
candidate_responses: list[CandidateResponse] = []
|
||||
for candidate in candidates:
|
||||
provider_name = (
|
||||
provider_map.get(candidate.provider_id) if candidate.provider_id else None
|
||||
|
||||
@@ -9,7 +9,7 @@ Provider 操作 API 路由
|
||||
"""
|
||||
|
||||
from dataclasses import asdict, is_dataclass
|
||||
from typing import Any, Dict, List, Optional
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
from pydantic import BaseModel, Field
|
||||
@@ -18,9 +18,7 @@ from sqlalchemy.orm import Session
|
||||
from src.database import get_db
|
||||
from src.models.database import Provider, User
|
||||
from src.services.provider_ops import (
|
||||
ActionStatus,
|
||||
ConnectorAuthType,
|
||||
ConnectorStatus,
|
||||
ProviderActionType,
|
||||
ProviderOpsConfig,
|
||||
ProviderOpsService,
|
||||
@@ -40,46 +38,46 @@ class ArchitectureInfo(BaseModel):
|
||||
architecture_id: str
|
||||
display_name: str
|
||||
description: str
|
||||
supported_auth_types: List[Dict[str, str]]
|
||||
supported_actions: List[Dict[str, Any]]
|
||||
default_connector: Optional[str]
|
||||
supported_auth_types: list[dict[str, str]]
|
||||
supported_actions: list[dict[str, Any]]
|
||||
default_connector: str | None
|
||||
|
||||
|
||||
class ConnectorConfigRequest(BaseModel):
|
||||
"""连接器配置请求"""
|
||||
|
||||
auth_type: str = Field(..., description="认证类型")
|
||||
config: Dict[str, Any] = Field(default_factory=dict, description="连接器配置")
|
||||
credentials: Dict[str, Any] = Field(default_factory=dict, description="凭据信息")
|
||||
config: dict[str, Any] = Field(default_factory=dict, description="连接器配置")
|
||||
credentials: dict[str, Any] = Field(default_factory=dict, description="凭据信息")
|
||||
|
||||
|
||||
class ActionConfigRequest(BaseModel):
|
||||
"""操作配置请求"""
|
||||
|
||||
enabled: bool = Field(True, description="是否启用")
|
||||
config: Dict[str, Any] = Field(default_factory=dict, description="操作配置")
|
||||
config: dict[str, Any] = Field(default_factory=dict, description="操作配置")
|
||||
|
||||
|
||||
class SaveConfigRequest(BaseModel):
|
||||
"""保存配置请求"""
|
||||
|
||||
architecture_id: str = Field("generic_api", description="架构 ID")
|
||||
base_url: Optional[str] = Field(None, description="API 基础地址")
|
||||
base_url: str | None = Field(None, description="API 基础地址")
|
||||
connector: ConnectorConfigRequest
|
||||
actions: Dict[str, ActionConfigRequest] = Field(default_factory=dict)
|
||||
schedule: Dict[str, str] = Field(default_factory=dict, description="定时任务配置")
|
||||
actions: dict[str, ActionConfigRequest] = Field(default_factory=dict)
|
||||
schedule: dict[str, str] = Field(default_factory=dict, description="定时任务配置")
|
||||
|
||||
|
||||
class ConnectRequest(BaseModel):
|
||||
"""连接请求"""
|
||||
|
||||
credentials: Optional[Dict[str, Any]] = Field(None, description="凭据(可选,使用已保存的)")
|
||||
credentials: dict[str, Any] | None = Field(None, description="凭据(可选,使用已保存的)")
|
||||
|
||||
|
||||
class ExecuteActionRequest(BaseModel):
|
||||
"""执行操作请求"""
|
||||
|
||||
config: Optional[Dict[str, Any]] = Field(None, description="操作配置(覆盖默认)")
|
||||
config: dict[str, Any] | None = Field(None, description="操作配置(覆盖默认)")
|
||||
|
||||
|
||||
class ConnectionStatusResponse(BaseModel):
|
||||
@@ -87,9 +85,9 @@ class ConnectionStatusResponse(BaseModel):
|
||||
|
||||
status: str
|
||||
auth_type: str
|
||||
connected_at: Optional[str]
|
||||
expires_at: Optional[str]
|
||||
last_error: Optional[str]
|
||||
connected_at: str | None
|
||||
expires_at: str | None
|
||||
last_error: str | None
|
||||
|
||||
|
||||
class ActionResultResponse(BaseModel):
|
||||
@@ -97,10 +95,10 @@ class ActionResultResponse(BaseModel):
|
||||
|
||||
status: str
|
||||
action_type: str
|
||||
data: Optional[Any]
|
||||
message: Optional[str]
|
||||
data: Any | None
|
||||
message: str | None
|
||||
executed_at: str
|
||||
response_time_ms: Optional[int]
|
||||
response_time_ms: int | None
|
||||
cache_ttl_seconds: int
|
||||
|
||||
|
||||
@@ -109,9 +107,9 @@ class ProviderOpsStatusResponse(BaseModel):
|
||||
|
||||
provider_id: str
|
||||
is_configured: bool
|
||||
architecture_id: Optional[str]
|
||||
architecture_id: str | None
|
||||
connection_status: ConnectionStatusResponse
|
||||
enabled_actions: List[str]
|
||||
enabled_actions: list[str]
|
||||
|
||||
|
||||
class ProviderOpsConfigResponse(BaseModel):
|
||||
@@ -119,17 +117,17 @@ class ProviderOpsConfigResponse(BaseModel):
|
||||
|
||||
provider_id: str
|
||||
is_configured: bool
|
||||
architecture_id: Optional[str] = None
|
||||
base_url: Optional[str] = None
|
||||
connector: Optional[Dict[str, Any]] = None # 脱敏后的连接器配置
|
||||
architecture_id: str | None = None
|
||||
base_url: str | None = None
|
||||
connector: dict[str, Any] | None = None # 脱敏后的连接器配置
|
||||
|
||||
|
||||
class VerifyAuthResponse(BaseModel):
|
||||
"""验证认证响应"""
|
||||
|
||||
success: bool
|
||||
message: Optional[str] = None
|
||||
data: Optional[Dict[str, Any]] = None
|
||||
message: str | None = None
|
||||
data: dict[str, Any] | None = None
|
||||
|
||||
|
||||
# ==================== Helper Functions ====================
|
||||
@@ -147,7 +145,7 @@ def _serialize_data(data: Any) -> Any:
|
||||
# ==================== Routes ====================
|
||||
|
||||
|
||||
@router.get("/architectures", response_model=List[ArchitectureInfo])
|
||||
@router.get("/architectures", response_model=list[ArchitectureInfo])
|
||||
async def list_architectures(_: User = Depends(require_admin)):
|
||||
"""获取所有可用的架构"""
|
||||
registry = get_registry()
|
||||
@@ -490,7 +488,7 @@ async def checkin(
|
||||
|
||||
@router.post("/batch/balance")
|
||||
async def batch_query_balance(
|
||||
provider_ids: Optional[List[str]] = None,
|
||||
provider_ids: list[str] | None = None,
|
||||
db: Session = Depends(get_db),
|
||||
_: User = Depends(require_admin),
|
||||
):
|
||||
|
||||
@@ -4,7 +4,6 @@ Provider Query API 端点
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from typing import Optional
|
||||
|
||||
import httpx
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
@@ -40,7 +39,7 @@ class ModelsQueryRequest(BaseModel):
|
||||
"""模型列表查询请求"""
|
||||
|
||||
provider_id: str
|
||||
api_key_id: Optional[str] = None
|
||||
api_key_id: str | None = None
|
||||
force_refresh: bool = False # 强制刷新,跳过缓存
|
||||
|
||||
|
||||
@@ -49,11 +48,11 @@ class TestModelRequest(BaseModel):
|
||||
|
||||
provider_id: str
|
||||
model_name: str
|
||||
api_key_id: Optional[str] = None
|
||||
endpoint_id: Optional[str] = None # 指定使用的端点ID
|
||||
api_key_id: str | None = None
|
||||
endpoint_id: str | None = None # 指定使用的端点ID
|
||||
stream: bool = False
|
||||
message: Optional[str] = "你好"
|
||||
api_format: Optional[str] = None # 指定使用的API格式,如果不指定则使用端点的默认格式
|
||||
message: str | None = "你好"
|
||||
api_format: str | None = None # 指定使用的API格式,如果不指定则使用端点的默认格式
|
||||
|
||||
|
||||
# ============ API Endpoints ============
|
||||
|
||||
@@ -3,7 +3,6 @@
|
||||
"""
|
||||
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request
|
||||
from fastapi.responses import JSONResponse
|
||||
@@ -25,11 +24,11 @@ pipeline = ApiRequestPipeline()
|
||||
|
||||
class ProviderBillingUpdate(BaseModel):
|
||||
billing_type: ProviderBillingType
|
||||
monthly_quota_usd: Optional[float] = None
|
||||
monthly_quota_usd: float | None = None
|
||||
quota_reset_day: int = Field(default=30, ge=1, le=365) # 重置周期(天数)
|
||||
quota_last_reset_at: Optional[str] = None # 当前周期开始时间
|
||||
quota_expires_at: Optional[str] = None
|
||||
rpm_limit: Optional[int] = Field(default=None, ge=0)
|
||||
quota_last_reset_at: str | None = None # 当前周期开始时间
|
||||
quota_expires_at: str | None = None
|
||||
rpm_limit: int | None = Field(default=None, ge=0)
|
||||
provider_priority: int = Field(default=100, ge=0, le=200)
|
||||
|
||||
|
||||
@@ -163,13 +162,12 @@ class AdminProviderBillingAdapter(AdminApiAdapter):
|
||||
provider.quota_reset_day = config.quota_reset_day
|
||||
provider.provider_priority = config.provider_priority
|
||||
|
||||
from dateutil import parser
|
||||
from sqlalchemy import func
|
||||
|
||||
from src.models.database import Usage
|
||||
|
||||
if config.quota_last_reset_at:
|
||||
new_reset_at = parser.parse(config.quota_last_reset_at)
|
||||
new_reset_at = datetime.fromisoformat(config.quota_last_reset_at)
|
||||
# 确保有时区信息,如果没有则假设为 UTC
|
||||
if new_reset_at.tzinfo is None:
|
||||
new_reset_at = new_reset_at.replace(tzinfo=timezone.utc)
|
||||
@@ -188,7 +186,7 @@ class AdminProviderBillingAdapter(AdminApiAdapter):
|
||||
logger.info(f"Synced usage for provider {provider.name}: ${period_usage:.4f} since {new_reset_at}")
|
||||
|
||||
if config.quota_expires_at:
|
||||
expires_at = parser.parse(config.quota_expires_at)
|
||||
expires_at = datetime.fromisoformat(config.quota_expires_at)
|
||||
# 确保有时区信息,如果没有则假设为 UTC
|
||||
if expires_at.tzinfo is None:
|
||||
expires_at = expires_at.replace(tzinfo=timezone.utc)
|
||||
|
||||
@@ -3,7 +3,7 @@ Provider 模型管理 API
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, List, Optional
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, Request
|
||||
from sqlalchemy.orm import Session, joinedload
|
||||
@@ -40,15 +40,15 @@ router = APIRouter(tags=["Model Management"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
|
||||
|
||||
@router.get("/{provider_id}/models", response_model=List[ModelResponse])
|
||||
@router.get("/{provider_id}/models", response_model=list[ModelResponse])
|
||||
async def list_provider_models(
|
||||
provider_id: str,
|
||||
request: Request,
|
||||
is_active: Optional[bool] = None,
|
||||
is_active: bool | None = None,
|
||||
skip: int = 0,
|
||||
limit: int = 100,
|
||||
db: Session = Depends(get_db),
|
||||
) -> List[ModelResponse]:
|
||||
) -> list[ModelResponse]:
|
||||
"""
|
||||
获取提供商的所有模型
|
||||
|
||||
@@ -222,13 +222,13 @@ async def delete_provider_model(
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.post("/{provider_id}/models/batch", response_model=List[ModelResponse])
|
||||
@router.post("/{provider_id}/models/batch", response_model=list[ModelResponse])
|
||||
async def batch_create_provider_models(
|
||||
provider_id: str,
|
||||
models_data: List[ModelCreate],
|
||||
models_data: list[ModelCreate],
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
) -> List[ModelResponse]:
|
||||
) -> list[ModelResponse]:
|
||||
"""
|
||||
批量创建模型
|
||||
|
||||
@@ -375,7 +375,7 @@ async def import_models_from_upstream(
|
||||
@dataclass
|
||||
class AdminListProviderModelsAdapter(AdminApiAdapter):
|
||||
provider_id: str
|
||||
is_active: Optional[bool]
|
||||
is_active: bool | None
|
||||
skip: int
|
||||
limit: int
|
||||
|
||||
@@ -482,7 +482,7 @@ class AdminDeleteProviderModelAdapter(AdminApiAdapter):
|
||||
@dataclass
|
||||
class AdminBatchCreateModelsAdapter(AdminApiAdapter):
|
||||
provider_id: str
|
||||
models_data: List[ModelCreate]
|
||||
models_data: list[ModelCreate]
|
||||
|
||||
async def handle(self, context): # type: ignore[override]
|
||||
db = context.db
|
||||
@@ -525,7 +525,7 @@ class AdminGetProviderAvailableSourceModelsAdapter(AdminApiAdapter):
|
||||
)
|
||||
|
||||
# 2. 构建以 GlobalModel 为主键的字典
|
||||
global_models_dict: Dict[str, Dict[str, Any]] = {}
|
||||
global_models_dict: dict[str, dict[str, Any]] = {}
|
||||
|
||||
for model in models:
|
||||
global_model = model.global_model
|
||||
|
||||
@@ -2,7 +2,6 @@
|
||||
|
||||
import asyncio
|
||||
from datetime import datetime, timezone
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, Request
|
||||
from pydantic import BaseModel, ConfigDict, Field, ValidationError
|
||||
@@ -48,7 +47,7 @@ class MappingMatchingGlobalModel(BaseModel):
|
||||
global_model_name: str
|
||||
display_name: str
|
||||
is_active: bool
|
||||
matched_models: List[MappingMatchedModel] = Field(
|
||||
matched_models: list[MappingMatchedModel] = Field(
|
||||
default_factory=list, description="匹配到的模型列表"
|
||||
)
|
||||
|
||||
@@ -62,8 +61,8 @@ class MappingMatchingKey(BaseModel):
|
||||
key_name: str
|
||||
masked_key: str
|
||||
is_active: bool
|
||||
allowed_models: List[str] = Field(default_factory=list, description="Key 的模型白名单")
|
||||
matching_global_models: List[MappingMatchingGlobalModel] = Field(
|
||||
allowed_models: list[str] = Field(default_factory=list, description="Key 的模型白名单")
|
||||
matching_global_models: list[MappingMatchingGlobalModel] = Field(
|
||||
default_factory=list, description="匹配到的 GlobalModel 列表"
|
||||
)
|
||||
|
||||
@@ -75,7 +74,7 @@ class ProviderMappingPreviewResponse(BaseModel):
|
||||
|
||||
provider_id: str
|
||||
provider_name: str
|
||||
keys: List[MappingMatchingKey] = Field(
|
||||
keys: list[MappingMatchingKey] = Field(
|
||||
default_factory=list, description="有白名单配置且匹配到映射的 Key 列表"
|
||||
)
|
||||
total_keys: int = Field(0, description="有匹配结果的 Key 数量")
|
||||
@@ -95,7 +94,7 @@ async def list_providers(
|
||||
request: Request,
|
||||
skip: int = Query(0, ge=0),
|
||||
limit: int = Query(100, ge=1, le=500),
|
||||
is_active: Optional[bool] = None,
|
||||
is_active: bool | None = None,
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
"""
|
||||
@@ -209,7 +208,7 @@ async def delete_provider(provider_id: str, request: Request, db: Session = Depe
|
||||
|
||||
|
||||
class AdminListProvidersAdapter(AdminApiAdapter):
|
||||
def __init__(self, skip: int, limit: int, is_active: Optional[bool]):
|
||||
def __init__(self, skip: int, limit: int, is_active: bool | None):
|
||||
self.skip = skip
|
||||
self.limit = limit
|
||||
self.is_active = is_active
|
||||
@@ -473,7 +472,7 @@ async def get_provider_mapping_preview(
|
||||
pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode),
|
||||
timeout=MAPPING_PREVIEW_TIMEOUT_SECONDS,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
except TimeoutError:
|
||||
logger.warning(f"映射预览超时: provider_id={provider_id}")
|
||||
raise InvalidRequestException("映射预览超时,请简化配置或稍后重试")
|
||||
|
||||
@@ -565,7 +564,7 @@ class AdminGetProviderMappingPreviewAdapter(AdminApiAdapter):
|
||||
truncated_models = total_models_with_mappings - MAPPING_PREVIEW_MAX_MODELS
|
||||
|
||||
# 构建有映射配置的 GlobalModel 映射
|
||||
models_with_mappings: Dict[str, tuple] = {} # id -> (model_info, mappings)
|
||||
models_with_mappings: dict[str, tuple] = {} # id -> (model_info, mappings)
|
||||
for gm in global_models:
|
||||
config = gm.config or {}
|
||||
mappings = config.get("model_mappings", [])
|
||||
@@ -585,7 +584,7 @@ class AdminGetProviderMappingPreviewAdapter(AdminApiAdapter):
|
||||
truncated_models=0,
|
||||
)
|
||||
|
||||
key_infos: List[MappingMatchingKey] = []
|
||||
key_infos: list[MappingMatchingKey] = []
|
||||
total_matches = 0
|
||||
|
||||
# 创建 CryptoService 实例
|
||||
@@ -611,10 +610,10 @@ class AdminGetProviderMappingPreviewAdapter(AdminApiAdapter):
|
||||
pass
|
||||
|
||||
# 查找匹配的 GlobalModel
|
||||
matching_global_models: List[MappingMatchingGlobalModel] = []
|
||||
matching_global_models: list[MappingMatchingGlobalModel] = []
|
||||
|
||||
for gm_id, (gm, mappings) in models_with_mappings.items():
|
||||
matched_models: List[MappingMatchedModel] = []
|
||||
matched_models: list[MappingMatchedModel] = []
|
||||
|
||||
for allowed_model in allowed_models_list:
|
||||
for mapping_pattern in mappings:
|
||||
|
||||
@@ -4,7 +4,6 @@ Provider 摘要与健康监控 API
|
||||
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Dict, List
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, Request
|
||||
from sqlalchemy import case, func
|
||||
@@ -35,11 +34,11 @@ router = APIRouter(tags=["Provider Summary"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
|
||||
|
||||
@router.get("/summary", response_model=List[ProviderWithEndpointsSummary])
|
||||
@router.get("/summary", response_model=list[ProviderWithEndpointsSummary])
|
||||
async def get_providers_summary(
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
) -> List[ProviderWithEndpointsSummary]:
|
||||
) -> list[ProviderWithEndpointsSummary]:
|
||||
"""
|
||||
获取所有提供商摘要信息
|
||||
|
||||
@@ -381,8 +380,8 @@ class AdminProviderHealthMonitorAdapter(AdminApiAdapter):
|
||||
)
|
||||
attempts = attempts_query.limit(limit_rows).all()
|
||||
|
||||
buffered_attempts: Dict[str, List[RequestCandidate]] = {eid: [] for eid in endpoint_ids}
|
||||
counters: Dict[str, int] = {eid: 0 for eid in endpoint_ids}
|
||||
buffered_attempts: dict[str, list[RequestCandidate]] = {eid: [] for eid in endpoint_ids}
|
||||
counters: dict[str, int] = {eid: 0 for eid in endpoint_ids}
|
||||
|
||||
for attempt in attempts:
|
||||
if not attempt.endpoint_id or attempt.endpoint_id not in buffered_attempts:
|
||||
@@ -392,10 +391,10 @@ class AdminProviderHealthMonitorAdapter(AdminApiAdapter):
|
||||
buffered_attempts[attempt.endpoint_id].append(attempt)
|
||||
counters[attempt.endpoint_id] += 1
|
||||
|
||||
endpoint_monitors: List[EndpointHealthMonitor] = []
|
||||
endpoint_monitors: list[EndpointHealthMonitor] = []
|
||||
for endpoint in endpoints:
|
||||
attempt_list = list(reversed(buffered_attempts.get(endpoint.id, [])))
|
||||
events: List[EndpointHealthEvent] = []
|
||||
events: list[EndpointHealthEvent] = []
|
||||
for attempt in attempt_list:
|
||||
event_timestamp = attempt.finished_at or attempt.started_at or attempt.created_at
|
||||
events.append(
|
||||
|
||||
@@ -4,8 +4,6 @@ IP 安全管理接口
|
||||
提供 IP 黑白名单管理和速率限制统计
|
||||
"""
|
||||
|
||||
from typing import List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from pydantic import BaseModel, Field, ValidationError
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -14,7 +12,6 @@ from src.api.base.adapter import ApiMode
|
||||
from src.api.base.authenticated_adapter import AuthenticatedApiAdapter
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.core.exceptions import InvalidRequestException, translate_pydantic_error
|
||||
from src.core.logger import logger
|
||||
from src.database import get_db
|
||||
from src.services.rate_limit.ip_limiter import IPRateLimiter
|
||||
|
||||
@@ -30,7 +27,7 @@ class AddIPToBlacklistRequest(BaseModel):
|
||||
|
||||
ip_address: str = Field(..., description="IP 地址")
|
||||
reason: str = Field(..., min_length=1, max_length=200, description="加入黑名单的原因")
|
||||
ttl: Optional[int] = Field(None, gt=0, description="过期时间(秒),None 表示永久")
|
||||
ttl: int | None = Field(None, gt=0, description="过期时间(秒),None 表示永久")
|
||||
|
||||
|
||||
class RemoveIPFromBlacklistRequest(BaseModel):
|
||||
|
||||
@@ -1,10 +1,8 @@
|
||||
"""系统设置API端点。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import copy
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from pydantic import ValidationError
|
||||
@@ -649,9 +647,8 @@ class AdminSystemStatsAdapter(AdminApiAdapter):
|
||||
class AdminTriggerCleanupAdapter(AdminApiAdapter):
|
||||
async def handle(self, context): # type: ignore[override]
|
||||
"""手动触发清理任务"""
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from sqlalchemy import func
|
||||
|
||||
from src.services.system.maintenance_scheduler import get_maintenance_scheduler
|
||||
|
||||
|
||||
@@ -3,7 +3,6 @@
|
||||
from collections import defaultdict
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from sqlalchemy import func
|
||||
@@ -34,8 +33,8 @@ pipeline = ApiRequestPipeline()
|
||||
async def get_usage_aggregation(
|
||||
request: Request,
|
||||
group_by: str = Query(..., description="Aggregation dimension: model, user, provider, or api_format"),
|
||||
start_date: Optional[datetime] = None,
|
||||
end_date: Optional[datetime] = None,
|
||||
start_date: datetime | None = None,
|
||||
end_date: datetime | None = None,
|
||||
limit: int = Query(20, ge=1, le=100),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
@@ -75,8 +74,8 @@ async def get_usage_aggregation(
|
||||
@router.get("/stats")
|
||||
async def get_usage_stats(
|
||||
request: Request,
|
||||
start_date: Optional[datetime] = None,
|
||||
end_date: Optional[datetime] = None,
|
||||
start_date: datetime | None = None,
|
||||
end_date: datetime | None = None,
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
"""
|
||||
@@ -122,14 +121,14 @@ async def get_activity_heatmap(
|
||||
@router.get("/records")
|
||||
async def get_usage_records(
|
||||
request: Request,
|
||||
start_date: Optional[datetime] = None,
|
||||
end_date: Optional[datetime] = None,
|
||||
search: Optional[str] = None, # 通用搜索:用户名、密钥名、模型名、提供商名
|
||||
user_id: Optional[str] = None,
|
||||
username: Optional[str] = None,
|
||||
model: Optional[str] = None,
|
||||
provider: Optional[str] = None,
|
||||
status: Optional[str] = None, # stream, standard, error
|
||||
start_date: datetime | None = None,
|
||||
end_date: datetime | None = None,
|
||||
search: str | None = None, # 通用搜索:用户名、密钥名、模型名、提供商名
|
||||
user_id: str | None = None,
|
||||
username: str | None = None,
|
||||
model: str | None = None,
|
||||
provider: str | None = None,
|
||||
status: str | None = None, # stream, standard, error
|
||||
limit: int = Query(100, ge=1, le=500),
|
||||
offset: int = Query(0, ge=0),
|
||||
db: Session = Depends(get_db),
|
||||
@@ -179,7 +178,7 @@ async def get_usage_records(
|
||||
@router.get("/active")
|
||||
async def get_active_requests(
|
||||
request: Request,
|
||||
ids: Optional[str] = Query(None, description="逗号分隔的请求 ID 列表,用于查询特定请求的状态"),
|
||||
ids: str | None = Query(None, description="逗号分隔的请求 ID 列表,用于查询特定请求的状态"),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
"""
|
||||
@@ -259,7 +258,7 @@ async def get_usage_detail(
|
||||
|
||||
|
||||
class AdminUsageStatsAdapter(AdminApiAdapter):
|
||||
def __init__(self, start_date: Optional[datetime], end_date: Optional[datetime]):
|
||||
def __init__(self, start_date: datetime | None, end_date: datetime | None):
|
||||
self.start_date = start_date
|
||||
self.end_date = end_date
|
||||
|
||||
@@ -339,7 +338,7 @@ class AdminActivityHeatmapAdapter(AdminApiAdapter):
|
||||
|
||||
|
||||
class AdminUsageByModelAdapter(AdminApiAdapter):
|
||||
def __init__(self, start_date: Optional[datetime], end_date: Optional[datetime], limit: int):
|
||||
def __init__(self, start_date: datetime | None, end_date: datetime | None, limit: int):
|
||||
self.start_date = start_date
|
||||
self.end_date = end_date
|
||||
self.limit = limit
|
||||
@@ -386,7 +385,7 @@ class AdminUsageByModelAdapter(AdminApiAdapter):
|
||||
|
||||
|
||||
class AdminUsageByUserAdapter(AdminApiAdapter):
|
||||
def __init__(self, start_date: Optional[datetime], end_date: Optional[datetime], limit: int):
|
||||
def __init__(self, start_date: datetime | None, end_date: datetime | None, limit: int):
|
||||
self.start_date = start_date
|
||||
self.end_date = end_date
|
||||
self.limit = limit
|
||||
@@ -436,7 +435,7 @@ class AdminUsageByUserAdapter(AdminApiAdapter):
|
||||
|
||||
|
||||
class AdminUsageByProviderAdapter(AdminApiAdapter):
|
||||
def __init__(self, start_date: Optional[datetime], end_date: Optional[datetime], limit: int):
|
||||
def __init__(self, start_date: datetime | None, end_date: datetime | None, limit: int):
|
||||
self.start_date = start_date
|
||||
self.end_date = end_date
|
||||
self.limit = limit
|
||||
@@ -446,7 +445,7 @@ class AdminUsageByProviderAdapter(AdminApiAdapter):
|
||||
|
||||
# 从 request_candidates 表统计每个 Provider 的尝试次数和成功率
|
||||
# 这样可以正确统计 Fallback 场景(一个请求可能尝试多个 Provider)
|
||||
from sqlalchemy import case, Integer
|
||||
from sqlalchemy import case
|
||||
|
||||
attempt_query = db.query(
|
||||
RequestCandidate.provider_id,
|
||||
@@ -550,7 +549,7 @@ class AdminUsageByProviderAdapter(AdminApiAdapter):
|
||||
|
||||
|
||||
class AdminUsageByApiFormatAdapter(AdminApiAdapter):
|
||||
def __init__(self, start_date: Optional[datetime], end_date: Optional[datetime], limit: int):
|
||||
def __init__(self, start_date: datetime | None, end_date: datetime | None, limit: int):
|
||||
self.start_date = start_date
|
||||
self.end_date = end_date
|
||||
self.limit = limit
|
||||
@@ -608,14 +607,14 @@ class AdminUsageByApiFormatAdapter(AdminApiAdapter):
|
||||
class AdminUsageRecordsAdapter(AdminApiAdapter):
|
||||
def __init__(
|
||||
self,
|
||||
start_date: Optional[datetime],
|
||||
end_date: Optional[datetime],
|
||||
search: Optional[str],
|
||||
user_id: Optional[str],
|
||||
username: Optional[str],
|
||||
model: Optional[str],
|
||||
provider: Optional[str],
|
||||
status: Optional[str],
|
||||
start_date: datetime | None,
|
||||
end_date: datetime | None,
|
||||
search: str | None,
|
||||
user_id: str | None,
|
||||
username: str | None,
|
||||
model: str | None,
|
||||
provider: str | None,
|
||||
status: str | None,
|
||||
limit: int,
|
||||
offset: int,
|
||||
):
|
||||
@@ -744,7 +743,7 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
|
||||
|
||||
for req_id, candidates in request_candidates.items():
|
||||
# 提取所有不同的 candidate_index
|
||||
unique_candidates = set(c[0] for c in candidates)
|
||||
unique_candidates = {c[0] for c in candidates}
|
||||
# 如果有多个不同的 candidate_index,说明发生了 Fallback(Provider 切换)
|
||||
fallback_map[req_id] = len(unique_candidates) > 1
|
||||
|
||||
@@ -877,7 +876,7 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
|
||||
class AdminActiveRequestsAdapter(AdminApiAdapter):
|
||||
"""轻量级活跃请求状态查询适配器"""
|
||||
|
||||
def __init__(self, ids: Optional[str]):
|
||||
def __init__(self, ids: str | None):
|
||||
self.ids = ids
|
||||
|
||||
async def handle(self, context): # type: ignore[override]
|
||||
@@ -1033,8 +1032,8 @@ class AdminUsageDetailAdapter(AdminApiAdapter):
|
||||
@router.get("/cache-affinity/ttl-analysis")
|
||||
async def analyze_cache_affinity_ttl(
|
||||
request: Request,
|
||||
user_id: Optional[str] = Query(None, description="指定用户 ID"),
|
||||
api_key_id: Optional[str] = Query(None, description="指定 API Key ID"),
|
||||
user_id: str | None = Query(None, description="指定用户 ID"),
|
||||
api_key_id: str | None = Query(None, description="指定 API Key ID"),
|
||||
hours: int = Query(168, ge=1, le=720, description="分析最近多少小时的数据"),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
@@ -1057,8 +1056,8 @@ async def analyze_cache_affinity_ttl(
|
||||
@router.get("/cache-affinity/hit-analysis")
|
||||
async def analyze_cache_hit(
|
||||
request: Request,
|
||||
user_id: Optional[str] = Query(None, description="指定用户 ID"),
|
||||
api_key_id: Optional[str] = Query(None, description="指定 API Key ID"),
|
||||
user_id: str | None = Query(None, description="指定用户 ID"),
|
||||
api_key_id: str | None = Query(None, description="指定 API Key ID"),
|
||||
hours: int = Query(168, ge=1, le=720, description="分析最近多少小时的数据"),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
@@ -1080,8 +1079,8 @@ class CacheAffinityTTLAnalysisAdapter(AdminApiAdapter):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
user_id: Optional[str],
|
||||
api_key_id: Optional[str],
|
||||
user_id: str | None,
|
||||
api_key_id: str | None,
|
||||
hours: int,
|
||||
):
|
||||
self.user_id = user_id
|
||||
@@ -1114,8 +1113,8 @@ class CacheHitAnalysisAdapter(AdminApiAdapter):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
user_id: Optional[str],
|
||||
api_key_id: Optional[str],
|
||||
user_id: str | None,
|
||||
api_key_id: str | None,
|
||||
hours: int,
|
||||
):
|
||||
self.user_id = user_id
|
||||
@@ -1147,7 +1146,7 @@ async def get_interval_timeline(
|
||||
request: Request,
|
||||
hours: int = Query(24, ge=1, le=720, description="分析最近多少小时的数据"),
|
||||
limit: int = Query(10000, ge=100, le=50000, description="最大返回数据点数量"),
|
||||
user_id: Optional[str] = Query(None, description="指定用户 ID"),
|
||||
user_id: str | None = Query(None, description="指定用户 ID"),
|
||||
include_user_info: bool = Query(False, description="是否包含用户信息(用于管理员多用户视图)"),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
@@ -1177,7 +1176,7 @@ class IntervalTimelineAdapter(AdminApiAdapter):
|
||||
self,
|
||||
hours: int,
|
||||
limit: int,
|
||||
user_id: Optional[str] = None,
|
||||
user_id: str | None = None,
|
||||
include_user_info: bool = False,
|
||||
):
|
||||
self.hours = hours
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
"""用户管理 API 端点。"""
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request, status
|
||||
from pydantic import ValidationError
|
||||
@@ -48,8 +47,8 @@ async def list_users(
|
||||
request: Request,
|
||||
skip: int = Query(0, ge=0, description="跳过记录数"),
|
||||
limit: int = Query(100, ge=1, le=1000, description="返回记录数"),
|
||||
role: Optional[str] = Query(None, description="按角色筛选(user/admin)"),
|
||||
is_active: Optional[bool] = Query(None, description="按状态筛选"),
|
||||
role: str | None = Query(None, description="按角色筛选(user/admin)"),
|
||||
is_active: bool | None = Query(None, description="按状态筛选"),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
"""
|
||||
@@ -136,7 +135,7 @@ async def reset_user_quota(user_id: str, request: Request, db: Session = Depends
|
||||
async def get_user_api_keys(
|
||||
user_id: str,
|
||||
request: Request,
|
||||
is_active: Optional[bool] = Query(None, description="按状态筛选"),
|
||||
is_active: bool | None = Query(None, description="按状态筛选"),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
"""
|
||||
@@ -274,7 +273,7 @@ class AdminCreateUserAdapter(AdminApiAdapter):
|
||||
|
||||
|
||||
class AdminListUsersAdapter(AdminApiAdapter):
|
||||
def __init__(self, skip: int, limit: int, role: Optional[str], is_active: Optional[bool]):
|
||||
def __init__(self, skip: int, limit: int, role: str | None, is_active: bool | None):
|
||||
self.skip = skip
|
||||
self.limit = limit
|
||||
self.role = role
|
||||
@@ -467,7 +466,7 @@ class AdminResetUserQuotaAdapter(AdminApiAdapter):
|
||||
class AdminGetUserKeysAdapter(AdminApiAdapter):
|
||||
"""获取用户的API Keys"""
|
||||
|
||||
def __init__(self, user_id: str, is_active: Optional[bool]):
|
||||
def __init__(self, user_id: str, is_active: bool | None):
|
||||
self.user_id = user_id
|
||||
self.is_active = is_active
|
||||
|
||||
|
||||
@@ -1,9 +1,8 @@
|
||||
"""公告系统 API 端点。"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request, status
|
||||
from fastapi import APIRouter, Depends, Query, Request
|
||||
from pydantic import ValidationError
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
@@ -12,7 +11,6 @@ from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.authenticated_adapter import AuthenticatedApiAdapter
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.core.exceptions import InvalidRequestException, translate_pydantic_error
|
||||
from src.core.logger import logger
|
||||
from src.database import get_db
|
||||
from src.models.api import CreateAnnouncementRequest, UpdateAnnouncementRequest
|
||||
from src.models.database import User
|
||||
@@ -251,7 +249,7 @@ class AnnouncementOptionalAuthAdapter(ApiAdapter):
|
||||
context.extra["optional_user"] = await self._resolve_optional_user(context)
|
||||
return None
|
||||
|
||||
async def _resolve_optional_user(self, context) -> Optional[User]:
|
||||
async def _resolve_optional_user(self, context) -> User | None:
|
||||
if context.user:
|
||||
return context.user
|
||||
|
||||
@@ -285,7 +283,7 @@ class AnnouncementOptionalAuthAdapter(ApiAdapter):
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def get_optional_user(self, context) -> Optional[User]:
|
||||
def get_optional_user(self, context) -> User | None:
|
||||
return context.extra.get("optional_user")
|
||||
|
||||
|
||||
|
||||
@@ -2,10 +2,9 @@
|
||||
认证相关API端点
|
||||
"""
|
||||
|
||||
from typing import Optional, Tuple
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
|
||||
from fastapi.security import HTTPBearer
|
||||
from pydantic import ValidationError
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
@@ -42,7 +41,7 @@ from src.services.email import EmailSenderService, EmailVerificationService
|
||||
from src.utils.request_utils import get_client_ip, get_user_agent
|
||||
|
||||
|
||||
def validate_email_suffix(db: Session, email: str) -> Tuple[bool, Optional[str]]:
|
||||
def validate_email_suffix(db: Session, email: str) -> tuple[bool, str | None]:
|
||||
"""
|
||||
验证邮箱后缀是否允许注册
|
||||
|
||||
|
||||
@@ -1,8 +1,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from enum import Enum
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import Any
|
||||
|
||||
from fastapi import Request, Response
|
||||
|
||||
@@ -23,7 +21,7 @@ class ApiAdapter(ABC):
|
||||
|
||||
name: str = "base"
|
||||
mode: ApiMode = ApiMode.STANDARD
|
||||
api_format: Optional[str] = None # 对应 Provider API 格式提示
|
||||
api_format: str | None = None # 对应 Provider API 格式提示
|
||||
audit_log_enabled: bool = True
|
||||
audit_success_event = None
|
||||
audit_failure_event = None
|
||||
@@ -36,7 +34,7 @@ class ApiAdapter(ABC):
|
||||
"""可选的授权钩子,默认允许通过。"""
|
||||
return None
|
||||
|
||||
def extract_api_key(self, request: Request) -> Optional[str]:
|
||||
def extract_api_key(self, request: Request) -> str | None:
|
||||
"""
|
||||
从请求中提取客户端 API 密钥。
|
||||
|
||||
@@ -55,17 +53,17 @@ class ApiAdapter(ABC):
|
||||
context: ApiRequestContext,
|
||||
*,
|
||||
success: bool,
|
||||
status_code: Optional[int],
|
||||
error: Optional[str] = None,
|
||||
) -> Dict[str, Any]:
|
||||
status_code: int | None,
|
||||
error: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""允许适配器在审计日志中追加自定义字段。"""
|
||||
return {}
|
||||
|
||||
def detect_capability_requirements(
|
||||
self,
|
||||
headers: Dict[str, str],
|
||||
request_body: Optional[Dict[str, Any]] = None,
|
||||
) -> Dict[str, bool]:
|
||||
headers: dict[str, str],
|
||||
request_body: dict[str, Any] | None = None,
|
||||
) -> dict[str, bool]:
|
||||
"""
|
||||
检测请求中隐含的能力需求(子类可覆盖)
|
||||
|
||||
|
||||
@@ -1,5 +1,3 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
from src.models.database import UserRole
|
||||
|
||||
@@ -1,10 +1,9 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import time
|
||||
import uuid
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import Any
|
||||
|
||||
from fastapi import HTTPException, Request
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -21,34 +20,34 @@ class ApiRequestContext:
|
||||
|
||||
request: Request
|
||||
db: Session
|
||||
user: Optional[User]
|
||||
api_key: Optional[ApiKey]
|
||||
user: User | None
|
||||
api_key: ApiKey | None
|
||||
request_id: str
|
||||
start_time: float
|
||||
client_ip: str
|
||||
user_agent: str
|
||||
original_headers: Dict[str, str]
|
||||
query_params: Dict[str, str]
|
||||
original_headers: dict[str, str]
|
||||
query_params: dict[str, str]
|
||||
raw_body: bytes | None = None
|
||||
json_body: Optional[Dict[str, Any]] = None
|
||||
quota_remaining: Optional[float] = None
|
||||
json_body: dict[str, Any] | None = None
|
||||
quota_remaining: float | None = None
|
||||
mode: str = "standard" # standard / proxy
|
||||
api_format_hint: Optional[str] = None
|
||||
api_format_hint: str | None = None
|
||||
|
||||
# URL 路径参数(如 Gemini API 的 /v1beta/models/{model}:generateContent)
|
||||
path_params: Dict[str, Any] = field(default_factory=dict)
|
||||
path_params: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
# Management Token(用于管理 API 认证)
|
||||
management_token: Optional[ManagementToken] = None
|
||||
management_token: ManagementToken | None = None
|
||||
|
||||
# 供适配器扩展的状态存储
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
audit_metadata: Dict[str, Any] = field(default_factory=dict)
|
||||
extra: dict[str, Any] = field(default_factory=dict)
|
||||
audit_metadata: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
# 高频轮询端点日志抑制标志
|
||||
quiet_logging: bool = False
|
||||
|
||||
def ensure_json_body(self) -> Dict[str, Any]:
|
||||
def ensure_json_body(self) -> dict[str, Any]:
|
||||
"""确保请求体已解析为JSON并返回。"""
|
||||
if self.json_body is not None:
|
||||
return self.json_body
|
||||
@@ -70,7 +69,7 @@ class ApiRequestContext:
|
||||
if value is not None:
|
||||
self.audit_metadata[key] = value
|
||||
|
||||
def extend_audit_metadata(self, data: Dict[str, Any]) -> None:
|
||||
def extend_audit_metadata(self, data: dict[str, Any]) -> None:
|
||||
"""批量附加审计字段。"""
|
||||
for key, value in data.items():
|
||||
if value is not None:
|
||||
@@ -81,13 +80,13 @@ class ApiRequestContext:
|
||||
cls,
|
||||
request: Request,
|
||||
db: Session,
|
||||
user: Optional[User],
|
||||
api_key: Optional[ApiKey],
|
||||
raw_body: Optional[bytes] = None,
|
||||
user: User | None,
|
||||
api_key: ApiKey | None,
|
||||
raw_body: bytes | None = None,
|
||||
mode: str = "standard",
|
||||
api_format_hint: Optional[str] = None,
|
||||
path_params: Optional[Dict[str, Any]] = None,
|
||||
) -> "ApiRequestContext":
|
||||
api_format_hint: str | None = None,
|
||||
path_params: dict[str, Any] | None = None,
|
||||
) -> ApiRequestContext:
|
||||
"""创建上下文实例并提前读取必要的元数据。"""
|
||||
request_id = getattr(request.state, "request_id", None) or str(uuid.uuid4())[:8]
|
||||
setattr(request.state, "request_id", request_id)
|
||||
|
||||
@@ -10,8 +10,9 @@
|
||||
4. Key 的 allowed_models 允许该模型(null = 允许所有)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
from dataclasses import asdict, dataclass
|
||||
from typing import Any, Optional
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
@@ -27,7 +28,7 @@ _CACHE_KEY_PREFIX = "models:list"
|
||||
_CACHE_TTL = CacheTTL.MODEL # 300 秒
|
||||
|
||||
|
||||
def _get_cache_key(api_formats: list[str], client_format: Optional[str] = None) -> str:
|
||||
def _get_cache_key(api_formats: list[str], client_format: str | None = None) -> str:
|
||||
"""生成缓存 key"""
|
||||
formats_str = ",".join(sorted(api_formats))
|
||||
format_key = (client_format or "any").lower()
|
||||
@@ -35,8 +36,8 @@ def _get_cache_key(api_formats: list[str], client_format: Optional[str] = None)
|
||||
|
||||
|
||||
async def _get_cached_models(
|
||||
api_formats: list[str], client_format: Optional[str] = None
|
||||
) -> Optional[list["ModelInfo"]]:
|
||||
api_formats: list[str], client_format: str | None = None
|
||||
) -> list[ModelInfo] | None:
|
||||
"""从缓存获取模型列表"""
|
||||
cache_key = _get_cache_key(api_formats, client_format)
|
||||
try:
|
||||
@@ -51,8 +52,8 @@ async def _get_cached_models(
|
||||
|
||||
async def _set_cached_models(
|
||||
api_formats: list[str],
|
||||
models: list["ModelInfo"],
|
||||
client_format: Optional[str] = None,
|
||||
models: list[ModelInfo],
|
||||
client_format: str | None = None,
|
||||
) -> None:
|
||||
"""将模型列表写入缓存"""
|
||||
cache_key = _get_cache_key(api_formats, client_format)
|
||||
@@ -87,8 +88,8 @@ class ModelInfo:
|
||||
|
||||
id: str # 模型 ID (GlobalModel.name 或 provider_model_name)
|
||||
display_name: str
|
||||
description: Optional[str]
|
||||
created_at: Optional[str] # ISO 格式
|
||||
description: str | None
|
||||
created_at: str | None # ISO 格式
|
||||
created_timestamp: int # Unix 时间戳
|
||||
provider_name: str
|
||||
provider_id: str = "" # Provider ID,用于权限过滤
|
||||
@@ -100,27 +101,27 @@ class ModelInfo:
|
||||
image_generation: bool = False
|
||||
structured_output: bool = False
|
||||
# 规格参数
|
||||
context_limit: Optional[int] = None
|
||||
output_limit: Optional[int] = None
|
||||
context_limit: int | None = None
|
||||
output_limit: int | None = None
|
||||
# 元信息
|
||||
family: Optional[str] = None
|
||||
knowledge_cutoff: Optional[str] = None
|
||||
input_modalities: Optional[list[str]] = None
|
||||
output_modalities: Optional[list[str]] = None
|
||||
family: str | None = None
|
||||
knowledge_cutoff: str | None = None
|
||||
input_modalities: list[str] | None = None
|
||||
output_modalities: list[str] | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class AccessRestrictions:
|
||||
"""API Key 或 User 的访问限制"""
|
||||
|
||||
allowed_providers: Optional[list[str]] = None # 允许的 Provider ID 列表
|
||||
allowed_models: Optional[list[str]] = None # 允许的模型名称列表
|
||||
allowed_api_formats: Optional[list[str]] = None # 允许的 API 格式列表
|
||||
allowed_providers: list[str] | None = None # 允许的 Provider ID 列表
|
||||
allowed_models: list[str] | None = None # 允许的模型名称列表
|
||||
allowed_api_formats: list[str] | None = None # 允许的 API 格式列表
|
||||
|
||||
@classmethod
|
||||
def from_api_key_and_user(
|
||||
cls, api_key: Optional[ApiKey], user: Optional[User]
|
||||
) -> "AccessRestrictions":
|
||||
cls, api_key: ApiKey | None, user: User | None
|
||||
) -> AccessRestrictions:
|
||||
"""
|
||||
从 API Key 和 User 合并访问限制
|
||||
|
||||
@@ -130,9 +131,9 @@ class AccessRestrictions:
|
||||
- 如果 API Key 无限制但 User 有限制,使用 User 的限制
|
||||
- 两者都无限制则返回空限制
|
||||
"""
|
||||
allowed_providers: Optional[list[str]] = None
|
||||
allowed_models: Optional[list[str]] = None
|
||||
allowed_api_formats: Optional[list[str]] = None
|
||||
allowed_providers: list[str] | None = None
|
||||
allowed_models: list[str] | None = None
|
||||
allowed_api_formats: list[str] | None = None
|
||||
|
||||
# 优先使用 API Key 的限制
|
||||
if api_key:
|
||||
@@ -197,8 +198,8 @@ class AccessRestrictions:
|
||||
|
||||
|
||||
def _normalize_api_formats(
|
||||
api_formats: Optional[list[str]],
|
||||
provider_to_formats: Optional[dict[str, set[str]]] = None,
|
||||
api_formats: list[str] | None,
|
||||
provider_to_formats: dict[str, set[str]] | None = None,
|
||||
) -> list[str]:
|
||||
"""规范化 API 格式列表(大写),必要时从 provider_to_formats 兜底"""
|
||||
if api_formats:
|
||||
@@ -212,7 +213,7 @@ def _normalize_api_formats(
|
||||
|
||||
|
||||
def _get_provider_model_names_for_formats(
|
||||
model: Model, usable_formats: Optional[set[str]] = None
|
||||
model: Model, usable_formats: set[str] | None = None
|
||||
) -> set[str]:
|
||||
"""
|
||||
获取模型在指定格式下支持的 Provider 模型名称集合
|
||||
@@ -305,7 +306,7 @@ def get_compatible_provider_formats(
|
||||
def get_available_provider_ids(
|
||||
db: Session,
|
||||
api_formats: list[str],
|
||||
provider_to_formats: Optional[dict[str, set[str]]] = None,
|
||||
provider_to_formats: dict[str, set[str]] | None = None,
|
||||
) -> set[str]:
|
||||
"""
|
||||
返回有可用端点的 Provider IDs
|
||||
@@ -334,7 +335,7 @@ def get_available_provider_ids(
|
||||
def _get_available_model_ids_for_format(
|
||||
db: Session,
|
||||
api_formats: list[str],
|
||||
provider_to_formats: Optional[dict[str, set[str]]] = None,
|
||||
provider_to_formats: dict[str, set[str]] | None = None,
|
||||
) -> set[str]:
|
||||
"""
|
||||
获取指定格式下真正可用的模型 ID 集合
|
||||
@@ -410,7 +411,7 @@ def _get_available_model_ids_for_format(
|
||||
return available_model_ids
|
||||
|
||||
|
||||
def _extract_model_info(model: Any) -> Optional[ModelInfo]:
|
||||
def _extract_model_info(model: Any) -> ModelInfo | None:
|
||||
"""
|
||||
从 Model 对象提取 ModelInfo
|
||||
|
||||
@@ -424,7 +425,7 @@ def _extract_model_info(model: Any) -> Optional[ModelInfo]:
|
||||
|
||||
model_id: str = global_model.name
|
||||
display_name: str = global_model.display_name
|
||||
created_at: Optional[str] = (
|
||||
created_at: str | None = (
|
||||
model.created_at.strftime("%Y-%m-%dT%H:%M:%SZ") if model.created_at else None
|
||||
)
|
||||
created_timestamp: int = int(model.created_at.timestamp()) if model.created_at else 0
|
||||
@@ -433,7 +434,7 @@ def _extract_model_info(model: Any) -> Optional[ModelInfo]:
|
||||
|
||||
# 从 GlobalModel.config 提取配置信息
|
||||
config: dict = global_model.config or {}
|
||||
description: Optional[str] = config.get("description")
|
||||
description: str | None = config.get("description")
|
||||
|
||||
return ModelInfo(
|
||||
id=model_id,
|
||||
@@ -464,10 +465,10 @@ def _extract_model_info(model: Any) -> Optional[ModelInfo]:
|
||||
async def list_available_models(
|
||||
db: Session,
|
||||
available_provider_ids: set[str],
|
||||
api_formats: Optional[list[str]] = None,
|
||||
restrictions: Optional[AccessRestrictions] = None,
|
||||
provider_to_formats: Optional[dict[str, set[str]]] = None,
|
||||
client_format: Optional[str] = None,
|
||||
api_formats: list[str] | None = None,
|
||||
restrictions: AccessRestrictions | None = None,
|
||||
provider_to_formats: dict[str, set[str]] | None = None,
|
||||
client_format: str | None = None,
|
||||
) -> list[ModelInfo]:
|
||||
"""
|
||||
获取可用模型列表(已去重,带缓存)
|
||||
@@ -503,7 +504,7 @@ async def list_available_models(
|
||||
return cached
|
||||
|
||||
# 如果提供了 api_formats,获取真正可用的模型 ID
|
||||
available_model_ids: Optional[set[str]] = None
|
||||
available_model_ids: set[str] | None = None
|
||||
if normalized_formats:
|
||||
available_model_ids = _get_available_model_ids_for_format(
|
||||
db, normalized_formats, provider_to_formats
|
||||
@@ -551,10 +552,10 @@ def find_model_by_id(
|
||||
db: Session,
|
||||
model_id: str,
|
||||
available_provider_ids: set[str],
|
||||
api_formats: Optional[list[str]] = None,
|
||||
restrictions: Optional[AccessRestrictions] = None,
|
||||
provider_to_formats: Optional[dict[str, set[str]]] = None,
|
||||
) -> Optional[ModelInfo]:
|
||||
api_formats: list[str] | None = None,
|
||||
restrictions: AccessRestrictions | None = None,
|
||||
provider_to_formats: dict[str, set[str]] | None = None,
|
||||
) -> ModelInfo | None:
|
||||
"""
|
||||
按 ID 查找模型(仅支持 GlobalModel.name)
|
||||
|
||||
@@ -575,7 +576,7 @@ def find_model_by_id(
|
||||
normalized_formats = _normalize_api_formats(api_formats, provider_to_formats)
|
||||
|
||||
# 如果提供了 api_formats,获取真正可用的模型 ID
|
||||
available_model_ids: Optional[set[str]] = None
|
||||
available_model_ids: set[str] | None = None
|
||||
if normalized_formats:
|
||||
available_model_ids = _get_available_model_ids_for_format(
|
||||
db, normalized_formats, provider_to_formats
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import asdict, dataclass
|
||||
from typing import Any, List, Sequence, Tuple, TypeVar
|
||||
from typing import Any, TypeVar
|
||||
from collections.abc import Sequence
|
||||
|
||||
from sqlalchemy.orm import Query
|
||||
|
||||
@@ -19,7 +18,7 @@ class PaginationMeta:
|
||||
return asdict(self)
|
||||
|
||||
|
||||
def paginate_query(query: Query, limit: int, offset: int) -> Tuple[int, List[T]]:
|
||||
def paginate_query(query: Query, limit: int, offset: int) -> tuple[int, list[T]]:
|
||||
"""
|
||||
对 SQLAlchemy 查询应用 limit/offset,并返回总数与结果列表。
|
||||
"""
|
||||
@@ -30,7 +29,7 @@ def paginate_query(query: Query, limit: int, offset: int) -> Tuple[int, List[T]]
|
||||
|
||||
def paginate_sequence(
|
||||
items: Sequence[T], limit: int, offset: int
|
||||
) -> Tuple[List[T], PaginationMeta]:
|
||||
) -> tuple[list[T], PaginationMeta]:
|
||||
"""
|
||||
对内存序列应用分页,返回切片和元数据。
|
||||
"""
|
||||
@@ -40,7 +39,7 @@ def paginate_sequence(
|
||||
return sliced, meta
|
||||
|
||||
|
||||
def build_pagination_payload(items: List[dict], meta: PaginationMeta, **extra: Any) -> dict:
|
||||
def build_pagination_payload(items: list[dict], meta: PaginationMeta, **extra: Any) -> dict:
|
||||
"""
|
||||
构建标准分页响应 payload。
|
||||
"""
|
||||
|
||||
@@ -2,7 +2,7 @@ from __future__ import annotations
|
||||
|
||||
import time
|
||||
from enum import Enum
|
||||
from typing import TYPE_CHECKING, Any, Optional, Tuple
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from fastapi import HTTPException, Request
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -52,8 +52,8 @@ class ApiRequestPipeline:
|
||||
db: Session,
|
||||
*,
|
||||
mode: ApiMode = ApiMode.STANDARD,
|
||||
api_format_hint: Optional[str] = None,
|
||||
path_params: Optional[dict[str, Any]] = None,
|
||||
api_format_hint: str | None = None,
|
||||
path_params: dict[str, Any] | None = None,
|
||||
):
|
||||
# 高频轮询端点抑制 debug 日志
|
||||
is_quiet = http_request.url.path in QUIET_POLLING_PATHS
|
||||
@@ -95,7 +95,7 @@ class ApiRequestPipeline:
|
||||
)
|
||||
if not is_quiet:
|
||||
logger.debug("[Pipeline] Raw body读取完成 | size=%d bytes", len(raw_body) if raw_body is not None else 0)
|
||||
except asyncio.TimeoutError:
|
||||
except TimeoutError:
|
||||
timeout_sec = int(config.request_body_timeout)
|
||||
logger.error(f"读取请求体超时({timeout_sec}s),可能客户端未发送完整请求体")
|
||||
raise HTTPException(
|
||||
@@ -166,7 +166,7 @@ class ApiRequestPipeline:
|
||||
|
||||
def _authenticate_client(
|
||||
self, request: Request, db: Session, adapter: ApiAdapter, *, quiet: bool = False
|
||||
) -> Tuple[User, ApiKey]:
|
||||
) -> tuple[User, ApiKey]:
|
||||
if not quiet:
|
||||
logger.debug("[Pipeline._authenticate_client] 开始")
|
||||
# 使用 adapter 的 extract_api_key 方法,支持不同 API 格式的认证头
|
||||
@@ -215,7 +215,7 @@ class ApiRequestPipeline:
|
||||
|
||||
async def _authenticate_admin(
|
||||
self, request: Request, db: Session
|
||||
) -> Tuple[User, Optional["ManagementToken"]]:
|
||||
) -> tuple[User, ManagementToken | None]:
|
||||
"""管理员认证,支持 JWT 和 Management Token 两种方式"""
|
||||
from src.models.database import ManagementToken
|
||||
from src.utils.request_utils import get_client_ip
|
||||
@@ -278,7 +278,7 @@ class ApiRequestPipeline:
|
||||
|
||||
async def _authenticate_user(
|
||||
self, request: Request, db: Session
|
||||
) -> Tuple[User, Optional["ManagementToken"]]:
|
||||
) -> tuple[User, ManagementToken | None]:
|
||||
"""用户认证,支持 JWT 和 Management Token 两种方式"""
|
||||
from src.models.database import ManagementToken
|
||||
from src.utils.request_utils import get_client_ip
|
||||
@@ -329,7 +329,7 @@ class ApiRequestPipeline:
|
||||
|
||||
async def _authenticate_management(
|
||||
self, request: Request, db: Session
|
||||
) -> Tuple[User, "ManagementToken"]:
|
||||
) -> tuple[User, ManagementToken]:
|
||||
"""Management Token 认证"""
|
||||
from src.models.database import ManagementToken
|
||||
from src.utils.request_utils import get_client_ip
|
||||
@@ -362,7 +362,7 @@ class ApiRequestPipeline:
|
||||
|
||||
return user, management_token
|
||||
|
||||
def _calculate_quota_remaining(self, user: Optional[User]) -> Optional[float]:
|
||||
def _calculate_quota_remaining(self, user: User | None) -> float | None:
|
||||
if not user:
|
||||
return None
|
||||
if user.quota_usd is None or user.quota_usd < 0:
|
||||
@@ -375,8 +375,8 @@ class ApiRequestPipeline:
|
||||
adapter: ApiAdapter,
|
||||
*,
|
||||
success: bool,
|
||||
status_code: Optional[int] = None,
|
||||
error: Optional[str] = None,
|
||||
status_code: int | None = None,
|
||||
error: str | None = None,
|
||||
) -> None:
|
||||
"""记录审计事件
|
||||
|
||||
@@ -432,8 +432,8 @@ class ApiRequestPipeline:
|
||||
adapter: ApiAdapter,
|
||||
*,
|
||||
success: bool,
|
||||
status_code: Optional[int],
|
||||
error: Optional[str],
|
||||
status_code: int | None,
|
||||
error: str | None,
|
||||
) -> dict:
|
||||
duration_ms = max((time.time() - context.start_time) * 1000, 0.0)
|
||||
request = context.request
|
||||
|
||||
@@ -2,7 +2,6 @@
|
||||
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import List
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from sqlalchemy import and_, func
|
||||
@@ -952,7 +951,7 @@ class DashboardDailyStatsAdapter(DashboardAdapter):
|
||||
# 构建完整日期序列(使用业务时区日期)
|
||||
current_date = start_date_local.date()
|
||||
end_date_date = end_date_local.date()
|
||||
formatted: List[dict] = []
|
||||
formatted: list[dict] = []
|
||||
while current_date <= end_date_date:
|
||||
date_str = current_date.isoformat()
|
||||
stat = stats_map.get(date_str)
|
||||
|
||||
@@ -27,21 +27,20 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
import asyncio
|
||||
import time
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
Awaitable,
|
||||
Callable,
|
||||
Coroutine,
|
||||
Dict,
|
||||
Optional,
|
||||
Protocol,
|
||||
TypeVar,
|
||||
runtime_checkable,
|
||||
)
|
||||
|
||||
from collections.abc import Callable
|
||||
from collections.abc import Awaitable, Coroutine
|
||||
|
||||
from fastapi import Request
|
||||
from fastapi.responses import JSONResponse, StreamingResponse
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -57,6 +56,9 @@ from src.services.usage.service import UsageService
|
||||
if TYPE_CHECKING:
|
||||
from src.api.handlers.base.stream_context import StreamContext
|
||||
|
||||
# Adapter 检测器类型:接受 headers 和可选的 request_body,返回能力需求字典
|
||||
type AdapterDetectorType = Callable[[dict[str, str], dict[str, Any] | None], dict[str, bool]]
|
||||
|
||||
|
||||
class MessageTelemetry:
|
||||
"""
|
||||
@@ -105,29 +107,29 @@ class MessageTelemetry:
|
||||
output_tokens: int,
|
||||
response_time_ms: int,
|
||||
status_code: int,
|
||||
request_body: Dict[str, Any],
|
||||
request_headers: Dict[str, Any],
|
||||
request_body: dict[str, Any],
|
||||
request_headers: dict[str, Any],
|
||||
response_body: Any,
|
||||
response_headers: Dict[str, Any],
|
||||
client_response_headers: Optional[Dict[str, Any]] = None,
|
||||
response_headers: dict[str, Any],
|
||||
client_response_headers: dict[str, Any] | None = None,
|
||||
cache_creation_tokens: int = 0,
|
||||
cache_read_tokens: int = 0,
|
||||
is_stream: bool = False,
|
||||
provider_request_headers: Optional[Dict[str, Any]] = None,
|
||||
provider_request_headers: dict[str, Any] | None = None,
|
||||
# 时间指标
|
||||
first_byte_time_ms: Optional[int] = None, # 首字时间/TTFB
|
||||
first_byte_time_ms: int | None = None, # 首字时间/TTFB
|
||||
# Provider 侧追踪信息(用于记录真实成本)
|
||||
provider_id: Optional[str] = None,
|
||||
provider_endpoint_id: Optional[str] = None,
|
||||
provider_api_key_id: Optional[str] = None,
|
||||
api_format: Optional[str] = None,
|
||||
provider_id: str | None = None,
|
||||
provider_endpoint_id: str | None = None,
|
||||
provider_api_key_id: str | None = None,
|
||||
api_format: str | None = None,
|
||||
# 格式转换追踪
|
||||
endpoint_api_format: Optional[str] = None, # 端点原生 API 格式
|
||||
endpoint_api_format: str | None = None, # 端点原生 API 格式
|
||||
has_format_conversion: bool = False, # 是否发生了格式转换
|
||||
# 模型映射信息
|
||||
target_model: Optional[str] = None,
|
||||
target_model: str | None = None,
|
||||
# Provider 响应元数据(如 Gemini 的 modelVersion)
|
||||
response_metadata: Optional[Dict[str, Any]] = None,
|
||||
response_metadata: dict[str, Any] | None = None,
|
||||
) -> float:
|
||||
total_cost = await self.calculate_cost(
|
||||
provider,
|
||||
@@ -199,24 +201,24 @@ class MessageTelemetry:
|
||||
response_time_ms: int,
|
||||
status_code: int,
|
||||
error_message: str,
|
||||
request_body: Dict[str, Any],
|
||||
request_headers: Dict[str, Any],
|
||||
request_body: dict[str, Any],
|
||||
request_headers: dict[str, Any],
|
||||
is_stream: bool,
|
||||
api_format: Optional[str] = None,
|
||||
provider_request_headers: Optional[Dict[str, Any]] = None,
|
||||
api_format: str | None = None,
|
||||
provider_request_headers: dict[str, Any] | None = None,
|
||||
# 预估 token 信息(来自 message_start 事件,用于中断请求的成本估算)
|
||||
input_tokens: int = 0,
|
||||
output_tokens: int = 0,
|
||||
cache_creation_tokens: int = 0,
|
||||
cache_read_tokens: int = 0,
|
||||
response_body: Optional[Dict[str, Any]] = None,
|
||||
response_headers: Optional[Dict[str, Any]] = None,
|
||||
client_response_headers: Optional[Dict[str, Any]] = None,
|
||||
response_body: dict[str, Any] | None = None,
|
||||
response_headers: dict[str, Any] | None = None,
|
||||
client_response_headers: dict[str, Any] | None = None,
|
||||
# 格式转换追踪
|
||||
endpoint_api_format: Optional[str] = None,
|
||||
endpoint_api_format: str | None = None,
|
||||
has_format_conversion: bool = False,
|
||||
# 模型映射信息
|
||||
target_model: Optional[str] = None,
|
||||
target_model: str | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
记录失败请求
|
||||
@@ -273,24 +275,24 @@ class MessageTelemetry:
|
||||
provider: str,
|
||||
model: str,
|
||||
response_time_ms: int,
|
||||
first_byte_time_ms: Optional[int],
|
||||
first_byte_time_ms: int | None,
|
||||
status_code: int,
|
||||
request_body: Dict[str, Any],
|
||||
request_headers: Dict[str, Any],
|
||||
request_body: dict[str, Any],
|
||||
request_headers: dict[str, Any],
|
||||
is_stream: bool,
|
||||
api_format: Optional[str] = None,
|
||||
provider_request_headers: Optional[Dict[str, Any]] = None,
|
||||
api_format: str | None = None,
|
||||
provider_request_headers: dict[str, Any] | None = None,
|
||||
input_tokens: int = 0,
|
||||
output_tokens: int = 0,
|
||||
cache_creation_tokens: int = 0,
|
||||
cache_read_tokens: int = 0,
|
||||
response_body: Optional[Dict[str, Any]] = None,
|
||||
response_headers: Optional[Dict[str, Any]] = None,
|
||||
client_response_headers: Optional[Dict[str, Any]] = None,
|
||||
response_body: dict[str, Any] | None = None,
|
||||
response_headers: dict[str, Any] | None = None,
|
||||
client_response_headers: dict[str, Any] | None = None,
|
||||
# 格式转换追踪
|
||||
endpoint_api_format: Optional[str] = None,
|
||||
endpoint_api_format: str | None = None,
|
||||
has_format_conversion: bool = False,
|
||||
target_model: Optional[str] = None,
|
||||
target_model: str | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
记录客户端取消的请求
|
||||
@@ -341,9 +343,9 @@ class MessageHandlerProtocol(Protocol):
|
||||
self,
|
||||
request: Any,
|
||||
http_request: Request,
|
||||
original_headers: Dict[str, str],
|
||||
original_request_body: Dict[str, Any],
|
||||
query_params: Optional[Dict[str, str]] = None,
|
||||
original_headers: dict[str, str],
|
||||
original_request_body: dict[str, Any],
|
||||
query_params: dict[str, str] | None = None,
|
||||
) -> StreamingResponse:
|
||||
"""处理流式请求"""
|
||||
...
|
||||
@@ -352,9 +354,9 @@ class MessageHandlerProtocol(Protocol):
|
||||
self,
|
||||
request: Any,
|
||||
http_request: Request,
|
||||
original_headers: Dict[str, str],
|
||||
original_request_body: Dict[str, Any],
|
||||
query_params: Optional[Dict[str, str]] = None,
|
||||
original_headers: dict[str, str],
|
||||
original_request_body: dict[str, Any],
|
||||
query_params: dict[str, str] | None = None,
|
||||
) -> JSONResponse:
|
||||
"""处理非流式请求"""
|
||||
...
|
||||
@@ -371,9 +373,6 @@ class BaseMessageHandler:
|
||||
推荐使用 MessageHandlerProtocol 中定义的签名。
|
||||
"""
|
||||
|
||||
# Adapter 检测器类型
|
||||
AdapterDetectorType = Callable[[Dict[str, str], Optional[Dict[str, Any]]], Dict[str, bool]]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
@@ -384,8 +383,8 @@ class BaseMessageHandler:
|
||||
client_ip: str,
|
||||
user_agent: str,
|
||||
start_time: float,
|
||||
allowed_api_formats: Optional[list[str]] = None,
|
||||
adapter_detector: Optional[AdapterDetectorType] = None,
|
||||
allowed_api_formats: list[str] | None = None,
|
||||
adapter_detector: AdapterDetectorType | None = None,
|
||||
) -> None:
|
||||
self.db = db
|
||||
self.user = user
|
||||
@@ -408,9 +407,9 @@ class BaseMessageHandler:
|
||||
def _resolve_capability_requirements(
|
||||
self,
|
||||
model_name: str,
|
||||
request_headers: Optional[Dict[str, str]] = None,
|
||||
request_body: Optional[Dict[str, Any]] = None,
|
||||
) -> Dict[str, bool]:
|
||||
request_headers: dict[str, str] | None = None,
|
||||
request_body: dict[str, Any] | None = None,
|
||||
) -> dict[str, bool]:
|
||||
"""
|
||||
解析请求的能力需求
|
||||
|
||||
@@ -442,12 +441,12 @@ class BaseMessageHandler:
|
||||
async def _resolve_preferred_key_ids(
|
||||
self,
|
||||
model_name: str,
|
||||
request_body: Optional[Dict[str, Any]] = None,
|
||||
) -> Optional[list[str]]:
|
||||
request_body: dict[str, Any] | None = None,
|
||||
) -> list[str] | None:
|
||||
"""可选的 Key 优先级解析钩子(默认不启用)。"""
|
||||
return None
|
||||
|
||||
def get_api_format(self, provider_type: Optional[str] = None) -> APIFormat:
|
||||
def get_api_format(self, provider_type: str | None = None) -> APIFormat:
|
||||
"""根据 provider_type 解析 API 格式,未知类型默认 OPENAI"""
|
||||
if provider_type:
|
||||
result = resolve_api_format(provider_type, default=APIFormat.OPENAI)
|
||||
@@ -456,17 +455,17 @@ class BaseMessageHandler:
|
||||
|
||||
def build_provider_payload(
|
||||
self,
|
||||
original_body: Dict[str, Any],
|
||||
original_body: dict[str, Any],
|
||||
*,
|
||||
mapped_model: Optional[str] = None,
|
||||
) -> Dict[str, Any]:
|
||||
mapped_model: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""构建发送给 Provider 的请求体,替换 model 名称"""
|
||||
payload = dict(original_body)
|
||||
if mapped_model:
|
||||
payload["model"] = mapped_model
|
||||
return payload
|
||||
|
||||
def _update_usage_to_streaming(self, request_id: Optional[str] = None) -> None:
|
||||
def _update_usage_to_streaming(self, request_id: str | None = None) -> None:
|
||||
"""更新 Usage 状态为 streaming(流式传输开始时调用)
|
||||
|
||||
使用 asyncio 后台任务执行数据库更新,避免阻塞流式传输
|
||||
@@ -500,7 +499,7 @@ class BaseMessageHandler:
|
||||
# 创建后台任务,不阻塞当前流
|
||||
asyncio.create_task(_do_update())
|
||||
|
||||
def _update_usage_to_streaming_with_ctx(self, ctx: "StreamContext") -> None:
|
||||
def _update_usage_to_streaming_with_ctx(self, ctx: StreamContext) -> None:
|
||||
"""更新 Usage 状态为 streaming,同时更新 provider 相关信息
|
||||
|
||||
使用 asyncio 后台任务执行数据库更新,避免阻塞流式传输
|
||||
|
||||
@@ -19,7 +19,7 @@ Chat Adapter 通用基类
|
||||
import time
|
||||
import traceback
|
||||
from abc import abstractmethod
|
||||
from typing import Any, Dict, Optional, Tuple, Type
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from fastapi import HTTPException, Request
|
||||
@@ -65,7 +65,7 @@ class ChatAdapterBase(ApiAdapter):
|
||||
|
||||
# 子类必须覆盖
|
||||
FORMAT_ID: str = "UNKNOWN"
|
||||
HANDLER_CLASS: Type[ChatHandlerBase]
|
||||
HANDLER_CLASS: type[ChatHandlerBase]
|
||||
|
||||
# 适配器配置
|
||||
name: str = "chat.base"
|
||||
@@ -90,7 +90,7 @@ class ChatAdapterBase(ApiAdapter):
|
||||
return base_url
|
||||
|
||||
@classmethod
|
||||
def build_base_headers(cls, api_key: str) -> Dict[str, str]:
|
||||
def build_base_headers(cls, api_key: str) -> dict[str, str]:
|
||||
"""构建基础请求头,使用统一的 headers.py 实现"""
|
||||
return build_adapter_base_headers(cls._get_api_format(), api_key)
|
||||
|
||||
@@ -101,13 +101,13 @@ class ChatAdapterBase(ApiAdapter):
|
||||
|
||||
@classmethod
|
||||
def build_headers_with_extra(
|
||||
cls, api_key: str, extra_headers: Optional[Dict[str, str]] = None
|
||||
) -> Dict[str, str]:
|
||||
cls, api_key: str, extra_headers: dict[str, str] | None = None
|
||||
) -> dict[str, str]:
|
||||
"""构建完整请求头(包含 extra_headers),使用统一的 headers.py 实现"""
|
||||
return build_adapter_headers(cls._get_api_format(), api_key, extra_headers)
|
||||
|
||||
@classmethod
|
||||
def build_request_body(cls, request_data: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
|
||||
def build_request_body(cls, request_data: dict[str, Any] | None = None) -> dict[str, Any]:
|
||||
"""构建测试请求体,使用转换器注册表自动处理格式转换
|
||||
|
||||
Args:
|
||||
@@ -120,11 +120,11 @@ class ChatAdapterBase(ApiAdapter):
|
||||
|
||||
return build_test_request_body(cls.FORMAT_ID, request_data)
|
||||
|
||||
def extract_api_key(self, request: Request) -> Optional[str]:
|
||||
def extract_api_key(self, request: Request) -> str | None:
|
||||
"""从请求中提取 API 密钥,使用统一的 headers.py 实现"""
|
||||
return extract_client_api_key(dict(request.headers), self._get_api_format())
|
||||
|
||||
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
|
||||
def __init__(self, allowed_api_formats: list[str] | None = None):
|
||||
self.allowed_api_formats = allowed_api_formats or [self.FORMAT_ID]
|
||||
|
||||
async def handle(self, context: ApiRequestContext):
|
||||
@@ -282,8 +282,8 @@ class ChatAdapterBase(ApiAdapter):
|
||||
)
|
||||
|
||||
def _merge_path_params(
|
||||
self, original_request_body: Dict[str, Any], path_params: Dict[str, Any]
|
||||
) -> Dict[str, Any]:
|
||||
self, original_request_body: dict[str, Any], path_params: dict[str, Any]
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
合并 URL 路径参数到请求体 - 子类可覆盖
|
||||
|
||||
@@ -316,7 +316,7 @@ class ChatAdapterBase(ApiAdapter):
|
||||
"""
|
||||
pass
|
||||
|
||||
def _extract_message_count(self, payload: Dict[str, Any], request_obj) -> int:
|
||||
def _extract_message_count(self, payload: dict[str, Any], request_obj) -> int:
|
||||
"""
|
||||
提取消息数量 - 子类可覆盖
|
||||
|
||||
@@ -327,7 +327,7 @@ class ChatAdapterBase(ApiAdapter):
|
||||
messages = request_obj.messages
|
||||
return len(messages) if isinstance(messages, list) else 0
|
||||
|
||||
def _build_audit_metadata(self, payload: Dict[str, Any], request_obj) -> Dict[str, Any]:
|
||||
def _build_audit_metadata(self, payload: dict[str, Any], request_obj) -> dict[str, Any]:
|
||||
"""
|
||||
构建审计日志元数据 - 子类可覆盖
|
||||
"""
|
||||
@@ -355,8 +355,8 @@ class ChatAdapterBase(ApiAdapter):
|
||||
model: str,
|
||||
stream: bool,
|
||||
start_time: float,
|
||||
original_headers: Dict[str, str],
|
||||
original_request_body: Dict[str, Any],
|
||||
original_headers: dict[str, str],
|
||||
original_request_body: dict[str, Any],
|
||||
client_ip: str,
|
||||
request_id: str,
|
||||
) -> JSONResponse:
|
||||
@@ -426,8 +426,8 @@ class ChatAdapterBase(ApiAdapter):
|
||||
model: str,
|
||||
stream: bool,
|
||||
start_time: float,
|
||||
original_headers: Dict[str, str],
|
||||
original_request_body: Dict[str, Any],
|
||||
original_headers: dict[str, str],
|
||||
original_request_body: dict[str, Any],
|
||||
client_ip: str,
|
||||
request_id: str,
|
||||
) -> JSONResponse:
|
||||
@@ -527,12 +527,12 @@ class ChatAdapterBase(ApiAdapter):
|
||||
cache_read_input_tokens: int,
|
||||
input_price_per_1m: float,
|
||||
output_price_per_1m: float,
|
||||
cache_creation_price_per_1m: Optional[float],
|
||||
cache_read_price_per_1m: Optional[float],
|
||||
price_per_request: Optional[float],
|
||||
tiered_pricing: Optional[dict] = None,
|
||||
cache_ttl_minutes: Optional[int] = None,
|
||||
) -> Dict[str, Any]:
|
||||
cache_creation_price_per_1m: float | None,
|
||||
cache_read_price_per_1m: float | None,
|
||||
price_per_request: float | None,
|
||||
tiered_pricing: dict | None = None,
|
||||
cache_ttl_minutes: int | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
计算请求成本
|
||||
|
||||
@@ -597,8 +597,8 @@ class ChatAdapterBase(ApiAdapter):
|
||||
client: httpx.AsyncClient,
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
) -> Tuple[list, Optional[str]]:
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
) -> tuple[list, str | None]:
|
||||
"""
|
||||
查询上游 API 支持的模型列表
|
||||
|
||||
@@ -626,16 +626,16 @@ class ChatAdapterBase(ApiAdapter):
|
||||
client: httpx.AsyncClient,
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
request_data: Dict[str, Any],
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
request_data: dict[str, Any],
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
# 用量计算参数(现在强制记录)
|
||||
db: Optional[Any] = None,
|
||||
user: Optional[Any] = None,
|
||||
provider_name: Optional[str] = None,
|
||||
provider_id: Optional[str] = None,
|
||||
api_key_id: Optional[str] = None,
|
||||
model_name: Optional[str] = None,
|
||||
) -> Dict[str, Any]:
|
||||
db: Any | None = None,
|
||||
user: Any | None = None,
|
||||
provider_name: str | None = None,
|
||||
provider_id: str | None = None,
|
||||
api_key_id: str | None = None,
|
||||
model_name: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
测试模型连接性(非流式)
|
||||
|
||||
@@ -682,11 +682,11 @@ class ChatAdapterBase(ApiAdapter):
|
||||
# Adapter 注册表 - 用于根据 API format 获取 Adapter 实例
|
||||
# =========================================================================
|
||||
|
||||
_ADAPTER_REGISTRY: Dict[str, Type["ChatAdapterBase"]] = {}
|
||||
_ADAPTER_REGISTRY: dict[str, type[ChatAdapterBase]] = {}
|
||||
_ADAPTERS_LOADED = False
|
||||
|
||||
|
||||
def register_adapter(adapter_class: Type["ChatAdapterBase"]) -> Type["ChatAdapterBase"]:
|
||||
def register_adapter(adapter_class: type[ChatAdapterBase]) -> type[ChatAdapterBase]:
|
||||
"""
|
||||
注册 Adapter 类到注册表
|
||||
|
||||
@@ -731,7 +731,7 @@ def _ensure_adapters_loaded():
|
||||
_ADAPTERS_LOADED = True
|
||||
|
||||
|
||||
def get_adapter_class(api_format: str) -> Optional[Type["ChatAdapterBase"]]:
|
||||
def get_adapter_class(api_format: str) -> type[ChatAdapterBase] | None:
|
||||
"""
|
||||
根据 API format 获取 Adapter 类
|
||||
|
||||
@@ -745,7 +745,7 @@ def get_adapter_class(api_format: str) -> Optional[Type["ChatAdapterBase"]]:
|
||||
return _ADAPTER_REGISTRY.get(api_format.upper()) if api_format else None
|
||||
|
||||
|
||||
def get_adapter_instance(api_format: str) -> Optional["ChatAdapterBase"]:
|
||||
def get_adapter_instance(api_format: str) -> ChatAdapterBase | None:
|
||||
"""
|
||||
根据 API format 获取 Adapter 实例
|
||||
|
||||
|
||||
@@ -22,7 +22,10 @@ Chat Handler Base - Chat API 格式的通用基类
|
||||
import asyncio
|
||||
import json
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Any, AsyncGenerator, Awaitable, Callable, Dict, Optional, Union
|
||||
from typing import Any
|
||||
|
||||
from collections.abc import Callable
|
||||
from collections.abc import AsyncGenerator, Awaitable
|
||||
|
||||
import httpx
|
||||
from fastapi import BackgroundTasks, Request
|
||||
@@ -75,10 +78,10 @@ def _get_error_status_code(e: Exception, default: int = 400) -> int:
|
||||
|
||||
|
||||
def _convert_error_response_best_effort(
|
||||
error_response: Dict[str, Any],
|
||||
error_response: dict[str, Any],
|
||||
source_format: str,
|
||||
target_format: str,
|
||||
) -> Dict[str, Any]:
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
将上游错误响应 best-effort 转换为客户端格式。
|
||||
|
||||
@@ -97,7 +100,7 @@ def _convert_error_response_best_effort(
|
||||
def _build_client_error_response_best_effort(
|
||||
message: str,
|
||||
target_format: str,
|
||||
) -> Dict[str, Any]:
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
当无法解析上游错误 body 时,构造一个目标格式的错误响应(best-effort)。
|
||||
"""
|
||||
@@ -117,11 +120,11 @@ def _build_client_error_response_best_effort(
|
||||
|
||||
|
||||
def _build_error_json_payload(
|
||||
e: Union[ThinkingSignatureException, UpstreamClientException],
|
||||
e: ThinkingSignatureException | UpstreamClientException,
|
||||
client_format: str,
|
||||
provider_format: str,
|
||||
needs_conversion: bool = True,
|
||||
) -> Dict[str, Any]:
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
构建错误 JSON 响应 payload(公共逻辑)。
|
||||
|
||||
@@ -185,10 +188,10 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
client_ip: str,
|
||||
user_agent: str,
|
||||
start_time: float,
|
||||
allowed_api_formats: Optional[list] = None,
|
||||
adapter_detector: Optional[
|
||||
Callable[[Dict[str, str], Optional[Dict[str, Any]]], Dict[str, bool]]
|
||||
] = None,
|
||||
allowed_api_formats: list | None = None,
|
||||
adapter_detector: None | (
|
||||
Callable[[dict[str, str], dict[str, Any] | None], dict[str, bool]]
|
||||
) = None,
|
||||
):
|
||||
allowed = allowed_api_formats or [self.FORMAT_ID]
|
||||
super().__init__(
|
||||
@@ -202,7 +205,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
allowed_api_formats=allowed,
|
||||
adapter_detector=adapter_detector,
|
||||
)
|
||||
self._parser: Optional[ResponseParser] = None
|
||||
self._parser: ResponseParser | None = None
|
||||
self._request_builder = PassthroughRequestBuilder()
|
||||
|
||||
@property
|
||||
@@ -228,7 +231,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def _extract_usage(self, response: Dict) -> Dict[str, int]:
|
||||
def _extract_usage(self, response: dict) -> dict[str, int]:
|
||||
"""
|
||||
从响应中提取 token 使用情况
|
||||
|
||||
@@ -241,7 +244,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
"""
|
||||
pass
|
||||
|
||||
def _normalize_response(self, response: Dict) -> Dict:
|
||||
def _normalize_response(self, response: dict) -> dict:
|
||||
"""
|
||||
规范化响应(可选覆盖)
|
||||
|
||||
@@ -257,8 +260,8 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
|
||||
def extract_model_from_request(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
path_params: Optional[Dict[str, Any]] = None, # noqa: ARG002 - 子类使用
|
||||
request_body: dict[str, Any],
|
||||
path_params: dict[str, Any] | None = None, # noqa: ARG002 - 子类使用
|
||||
) -> str:
|
||||
"""
|
||||
从请求中提取模型名 - 子类可覆盖
|
||||
@@ -282,9 +285,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
|
||||
def apply_mapped_model(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
request_body: dict[str, Any],
|
||||
mapped_model: str, # noqa: ARG002 - 子类使用
|
||||
) -> Dict[str, Any]:
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
将映射后的模型名应用到请求体
|
||||
|
||||
@@ -303,9 +306,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
|
||||
def get_model_for_url(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
mapped_model: Optional[str],
|
||||
) -> Optional[str]:
|
||||
request_body: dict[str, Any],
|
||||
mapped_model: str | None,
|
||||
) -> str | None:
|
||||
"""
|
||||
获取用于 URL 路径的模型名
|
||||
|
||||
@@ -323,8 +326,8 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
|
||||
def prepare_provider_request_body(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
) -> Dict[str, Any]:
|
||||
request_body: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
准备发送给 Provider 的请求体 - 子类可覆盖
|
||||
|
||||
@@ -341,9 +344,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
|
||||
def _set_model_after_conversion(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
request_body: dict[str, Any],
|
||||
provider_api_format: str,
|
||||
mapped_model: Optional[str],
|
||||
mapped_model: str | None,
|
||||
fallback_model: str,
|
||||
) -> None:
|
||||
"""
|
||||
@@ -372,7 +375,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
|
||||
def _set_stream_after_conversion(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
request_body: dict[str, Any],
|
||||
client_api_format: str,
|
||||
provider_api_format: str,
|
||||
is_stream: bool,
|
||||
@@ -414,8 +417,8 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
self,
|
||||
source_model: str,
|
||||
provider_id: str,
|
||||
api_format: Optional[str] = None,
|
||||
) -> Optional[str]:
|
||||
api_format: str | None = None,
|
||||
) -> str | None:
|
||||
"""
|
||||
获取模型映射后的实际模型名
|
||||
|
||||
@@ -452,10 +455,10 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
self,
|
||||
request: Any,
|
||||
http_request: Request,
|
||||
original_headers: Dict[str, Any],
|
||||
original_request_body: Dict[str, Any],
|
||||
query_params: Optional[Dict[str, str]] = None,
|
||||
) -> Union[StreamingResponse, JSONResponse]:
|
||||
original_headers: dict[str, Any],
|
||||
original_request_body: dict[str, Any],
|
||||
query_params: dict[str, str] | None = None,
|
||||
) -> StreamingResponse | JSONResponse:
|
||||
"""处理流式响应"""
|
||||
logger.debug(f"开始流式响应处理 ({self.FORMAT_ID})")
|
||||
|
||||
@@ -466,7 +469,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
|
||||
# 可变请求体容器:允许 orchestrator 在遇到 Thinking 签名错误时整流请求体后重试
|
||||
# 结构: {"body": 实际请求体, "_rectified": 是否已整流, "_rectified_this_turn": 本轮是否整流}
|
||||
request_body_ref: Dict[str, Any] = {"body": original_request_body}
|
||||
request_body_ref: dict[str, Any] = {"body": original_request_body}
|
||||
|
||||
# 创建类型安全的流式上下文
|
||||
ctx = StreamContext(model=model, api_format=api_format)
|
||||
@@ -492,7 +495,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
endpoint: ProviderEndpoint,
|
||||
key: ProviderAPIKey,
|
||||
candidate: ProviderCandidate,
|
||||
) -> AsyncGenerator[bytes, None]:
|
||||
) -> AsyncGenerator[bytes]:
|
||||
return await self._execute_stream_request(
|
||||
ctx,
|
||||
stream_processor,
|
||||
@@ -615,12 +618,12 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
provider: Provider,
|
||||
endpoint: ProviderEndpoint,
|
||||
key: ProviderAPIKey,
|
||||
original_request_body: Dict[str, Any],
|
||||
original_headers: Dict[str, str],
|
||||
query_params: Optional[Dict[str, str]] = None,
|
||||
candidate: Optional[ProviderCandidate] = None,
|
||||
is_disconnected: Optional[Callable[[], Awaitable[bool]]] = None,
|
||||
) -> AsyncGenerator[bytes, None]:
|
||||
original_request_body: dict[str, Any],
|
||||
original_headers: dict[str, str],
|
||||
query_params: dict[str, str] | None = None,
|
||||
candidate: ProviderCandidate | None = None,
|
||||
is_disconnected: Callable[[], Awaitable[bool]] | None = None,
|
||||
) -> AsyncGenerator[bytes]:
|
||||
"""执行流式请求并返回流生成器"""
|
||||
# 重置上下文状态(重试时清除之前的数据)
|
||||
ctx.reset_for_retry()
|
||||
@@ -799,7 +802,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
ctx.error_message = "client_disconnected_during_prefetch"
|
||||
raise
|
||||
|
||||
except asyncio.TimeoutError:
|
||||
except TimeoutError:
|
||||
# 整体请求超时(建立连接 + 获取首字节)
|
||||
# 清理可能已建立的连接上下文
|
||||
if response_ctx is not None:
|
||||
@@ -856,8 +859,8 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
self,
|
||||
ctx: StreamContext,
|
||||
error: Exception,
|
||||
original_headers: Dict[str, str],
|
||||
original_request_body: Dict[str, Any],
|
||||
original_headers: dict[str, str],
|
||||
original_request_body: dict[str, Any],
|
||||
) -> None:
|
||||
"""记录流式请求失败"""
|
||||
response_time_ms = self.elapsed_ms()
|
||||
@@ -904,9 +907,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
self,
|
||||
request: Any,
|
||||
http_request: Request,
|
||||
original_headers: Dict[str, Any],
|
||||
original_request_body: Dict[str, Any],
|
||||
query_params: Optional[Dict[str, str]] = None,
|
||||
original_headers: dict[str, Any],
|
||||
original_request_body: dict[str, Any],
|
||||
query_params: dict[str, str] | None = None,
|
||||
) -> JSONResponse:
|
||||
"""处理非流式响应"""
|
||||
logger.debug(f"开始非流式响应处理 ({self.FORMAT_ID})")
|
||||
@@ -918,29 +921,29 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
|
||||
# 可变请求体容器:允许 orchestrator 在遇到 Thinking 签名错误时整流请求体后重试
|
||||
# 结构: {"body": 实际请求体, "_rectified": 是否已整流, "_rectified_this_turn": 本轮是否整流}
|
||||
request_body_ref: Dict[str, Any] = {"body": original_request_body}
|
||||
request_body_ref: dict[str, Any] = {"body": original_request_body}
|
||||
|
||||
# 用于跟踪的变量
|
||||
provider_name: Optional[str] = None
|
||||
response_json: Optional[Dict[str, Any]] = None
|
||||
provider_name: str | None = None
|
||||
response_json: dict[str, Any] | None = None
|
||||
status_code = 200
|
||||
response_headers: Dict[str, str] = {}
|
||||
provider_request_headers: Dict[str, str] = {}
|
||||
provider_request_body: Optional[Dict[str, Any]] = None
|
||||
provider_api_format_for_error: Optional[str] = None
|
||||
client_api_format_for_error: Optional[str] = None
|
||||
response_headers: dict[str, str] = {}
|
||||
provider_request_headers: dict[str, str] = {}
|
||||
provider_request_body: dict[str, Any] | None = None
|
||||
provider_api_format_for_error: str | None = None
|
||||
client_api_format_for_error: str | None = None
|
||||
needs_conversion_for_error: bool = False
|
||||
provider_id: Optional[str] = None # Provider ID(用于失败记录)
|
||||
endpoint_id: Optional[str] = None # Endpoint ID(用于失败记录)
|
||||
key_id: Optional[str] = None # Key ID(用于失败记录)
|
||||
mapped_model_result: Optional[str] = None # 映射后的目标模型名(用于 Usage 记录)
|
||||
provider_id: str | None = None # Provider ID(用于失败记录)
|
||||
endpoint_id: str | None = None # Endpoint ID(用于失败记录)
|
||||
key_id: str | None = None # Key ID(用于失败记录)
|
||||
mapped_model_result: str | None = None # 映射后的目标模型名(用于 Usage 记录)
|
||||
|
||||
async def sync_request_func(
|
||||
provider: Provider,
|
||||
endpoint: ProviderEndpoint,
|
||||
key: ProviderAPIKey,
|
||||
candidate: ProviderCandidate,
|
||||
) -> Dict[str, Any]:
|
||||
) -> dict[str, Any]:
|
||||
nonlocal provider_name, response_json, status_code, response_headers
|
||||
nonlocal provider_request_headers, provider_request_body, mapped_model_result
|
||||
nonlocal provider_api_format_for_error, client_api_format_for_error, needs_conversion_for_error
|
||||
@@ -1293,7 +1296,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
actual_request_body = provider_request_body or original_request_body
|
||||
|
||||
# 尝试从异常中提取响应头
|
||||
error_response_headers: Dict[str, str] = {}
|
||||
error_response_headers: dict[str, str] = {}
|
||||
if isinstance(e, ProviderRateLimitException) and e.response_headers:
|
||||
error_response_headers = e.response_headers
|
||||
elif isinstance(e, httpx.HTTPStatusError) and hasattr(e, "response"):
|
||||
|
||||
@@ -17,7 +17,7 @@ CLI Adapter 通用基类
|
||||
|
||||
import time
|
||||
import traceback
|
||||
from typing import Any, Dict, Optional, Tuple, Type
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from fastapi import HTTPException, Request
|
||||
@@ -63,7 +63,7 @@ class CliAdapterBase(ApiAdapter):
|
||||
|
||||
# 子类必须覆盖
|
||||
FORMAT_ID: str = "UNKNOWN"
|
||||
HANDLER_CLASS: Type[CliMessageHandlerBase]
|
||||
HANDLER_CLASS: type[CliMessageHandlerBase]
|
||||
|
||||
# 适配器配置
|
||||
name: str = "cli.base"
|
||||
@@ -72,7 +72,7 @@ class CliAdapterBase(ApiAdapter):
|
||||
# 计费模板配置(子类可覆盖,如 "claude", "openai", "gemini")
|
||||
BILLING_TEMPLATE: str = "claude"
|
||||
|
||||
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
|
||||
def __init__(self, allowed_api_formats: list[str] | None = None):
|
||||
self.allowed_api_formats = allowed_api_formats or [self.FORMAT_ID]
|
||||
|
||||
# =========================================================================
|
||||
@@ -87,7 +87,7 @@ class CliAdapterBase(ApiAdapter):
|
||||
except KeyError:
|
||||
return APIFormat.OPENAI
|
||||
|
||||
def extract_api_key(self, request: Request) -> Optional[str]:
|
||||
def extract_api_key(self, request: Request) -> str | None:
|
||||
"""
|
||||
从请求中提取 API 密钥
|
||||
|
||||
@@ -96,7 +96,7 @@ class CliAdapterBase(ApiAdapter):
|
||||
return extract_client_api_key(dict(request.headers), self._get_api_format())
|
||||
|
||||
@classmethod
|
||||
def build_base_headers(cls, api_key: str) -> Dict[str, str]:
|
||||
def build_base_headers(cls, api_key: str) -> dict[str, str]:
|
||||
"""
|
||||
构建 CLI API 认证头
|
||||
|
||||
@@ -106,8 +106,8 @@ class CliAdapterBase(ApiAdapter):
|
||||
|
||||
@classmethod
|
||||
def build_headers_with_extra(
|
||||
cls, api_key: str, extra_headers: Optional[Dict[str, str]] = None
|
||||
) -> Dict[str, str]:
|
||||
cls, api_key: str, extra_headers: dict[str, str] | None = None
|
||||
) -> dict[str, str]:
|
||||
"""
|
||||
构建带额外头部的完整请求头
|
||||
|
||||
@@ -260,8 +260,8 @@ class CliAdapterBase(ApiAdapter):
|
||||
)
|
||||
|
||||
def _merge_path_params(
|
||||
self, original_request_body: Dict[str, Any], path_params: Dict[str, Any]
|
||||
) -> Dict[str, Any]:
|
||||
self, original_request_body: dict[str, Any], path_params: dict[str, Any]
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
合并 URL 路径参数到请求体 - 子类可覆盖
|
||||
|
||||
@@ -280,7 +280,7 @@ class CliAdapterBase(ApiAdapter):
|
||||
merged[key] = value
|
||||
return merged
|
||||
|
||||
def _extract_message_count(self, payload: Dict[str, Any]) -> int:
|
||||
def _extract_message_count(self, payload: dict[str, Any]) -> int:
|
||||
"""
|
||||
提取消息数量 - 子类可覆盖
|
||||
|
||||
@@ -297,9 +297,9 @@ class CliAdapterBase(ApiAdapter):
|
||||
|
||||
def _build_audit_metadata(
|
||||
self,
|
||||
payload: Dict[str, Any],
|
||||
path_params: Optional[Dict[str, Any]] = None,
|
||||
) -> Dict[str, Any]:
|
||||
payload: dict[str, Any],
|
||||
path_params: dict[str, Any] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
构建审计日志元数据 - 子类可覆盖
|
||||
|
||||
@@ -338,8 +338,8 @@ class CliAdapterBase(ApiAdapter):
|
||||
model: str,
|
||||
stream: bool,
|
||||
start_time: float,
|
||||
original_headers: Dict[str, str],
|
||||
original_request_body: Dict[str, Any],
|
||||
original_headers: dict[str, str],
|
||||
original_request_body: dict[str, Any],
|
||||
client_ip: str,
|
||||
request_id: str,
|
||||
) -> JSONResponse:
|
||||
@@ -409,8 +409,8 @@ class CliAdapterBase(ApiAdapter):
|
||||
model: str,
|
||||
stream: bool,
|
||||
start_time: float,
|
||||
original_headers: Dict[str, str],
|
||||
original_request_body: Dict[str, Any],
|
||||
original_headers: dict[str, str],
|
||||
original_request_body: dict[str, Any],
|
||||
client_ip: str,
|
||||
request_id: str,
|
||||
) -> JSONResponse:
|
||||
@@ -507,12 +507,12 @@ class CliAdapterBase(ApiAdapter):
|
||||
cache_read_input_tokens: int,
|
||||
input_price_per_1m: float,
|
||||
output_price_per_1m: float,
|
||||
cache_creation_price_per_1m: Optional[float],
|
||||
cache_read_price_per_1m: Optional[float],
|
||||
price_per_request: Optional[float],
|
||||
tiered_pricing: Optional[dict] = None,
|
||||
cache_ttl_minutes: Optional[int] = None,
|
||||
) -> Dict[str, Any]:
|
||||
cache_creation_price_per_1m: float | None,
|
||||
cache_read_price_per_1m: float | None,
|
||||
price_per_request: float | None,
|
||||
tiered_pricing: dict | None = None,
|
||||
cache_ttl_minutes: int | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
计算请求成本
|
||||
|
||||
@@ -567,8 +567,8 @@ class CliAdapterBase(ApiAdapter):
|
||||
client: httpx.AsyncClient,
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
) -> Tuple[list, Optional[str]]:
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
) -> tuple[list, str | None]:
|
||||
"""
|
||||
查询上游 API 支持的模型列表
|
||||
|
||||
@@ -596,16 +596,16 @@ class CliAdapterBase(ApiAdapter):
|
||||
client: httpx.AsyncClient,
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
request_data: Dict[str, Any],
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
request_data: dict[str, Any],
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
# 用量计算参数
|
||||
db: Optional[Any] = None,
|
||||
user: Optional[Any] = None,
|
||||
provider_name: Optional[str] = None,
|
||||
provider_id: Optional[str] = None,
|
||||
api_key_id: Optional[str] = None,
|
||||
model_name: Optional[str] = None,
|
||||
) -> Dict[str, Any]:
|
||||
db: Any | None = None,
|
||||
user: Any | None = None,
|
||||
provider_name: str | None = None,
|
||||
provider_id: str | None = None,
|
||||
api_key_id: str | None = None,
|
||||
model_name: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
测试模型连接性(非流式)
|
||||
|
||||
@@ -669,7 +669,7 @@ class CliAdapterBase(ApiAdapter):
|
||||
# =========================================================================
|
||||
|
||||
@classmethod
|
||||
def build_endpoint_url(cls, base_url: str, request_data: Dict[str, Any], model_name: Optional[str] = None) -> str:
|
||||
def build_endpoint_url(cls, base_url: str, request_data: dict[str, Any], model_name: str | None = None) -> str:
|
||||
"""
|
||||
构建CLI API端点URL - 子类应覆盖
|
||||
|
||||
@@ -684,7 +684,7 @@ class CliAdapterBase(ApiAdapter):
|
||||
raise NotImplementedError(f"{cls.FORMAT_ID} adapter must implement build_endpoint_url")
|
||||
|
||||
@classmethod
|
||||
def build_request_body(cls, request_data: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
|
||||
def build_request_body(cls, request_data: dict[str, Any] | None = None) -> dict[str, Any]:
|
||||
"""构建测试请求体,使用转换器注册表自动处理格式转换
|
||||
|
||||
Args:
|
||||
@@ -698,7 +698,7 @@ class CliAdapterBase(ApiAdapter):
|
||||
return build_test_request_body(cls.FORMAT_ID, request_data)
|
||||
|
||||
@classmethod
|
||||
def get_cli_user_agent(cls) -> Optional[str]:
|
||||
def get_cli_user_agent(cls) -> str | None:
|
||||
"""
|
||||
获取CLI User-Agent - 子类可覆盖
|
||||
|
||||
@@ -708,7 +708,7 @@ class CliAdapterBase(ApiAdapter):
|
||||
return None
|
||||
|
||||
@classmethod
|
||||
def get_cli_extra_headers(cls) -> Dict[str, str]:
|
||||
def get_cli_extra_headers(cls) -> dict[str, str]:
|
||||
"""
|
||||
获取CLI额外请求头 - 子类可覆盖
|
||||
|
||||
@@ -718,7 +718,7 @@ class CliAdapterBase(ApiAdapter):
|
||||
Returns:
|
||||
额外请求头字典
|
||||
"""
|
||||
headers: Dict[str, str] = {}
|
||||
headers: dict[str, str] = {}
|
||||
cli_user_agent = cls.get_cli_user_agent()
|
||||
if cli_user_agent:
|
||||
headers["User-Agent"] = cli_user_agent
|
||||
@@ -728,11 +728,11 @@ class CliAdapterBase(ApiAdapter):
|
||||
# CLI Adapter 注册表 - 用于根据 API format 获取 CLI Adapter 实例
|
||||
# =========================================================================
|
||||
|
||||
_CLI_ADAPTER_REGISTRY: Dict[str, Type["CliAdapterBase"]] = {}
|
||||
_CLI_ADAPTER_REGISTRY: dict[str, type[CliAdapterBase]] = {}
|
||||
_CLI_ADAPTERS_LOADED = False
|
||||
|
||||
|
||||
def register_cli_adapter(adapter_class: Type["CliAdapterBase"]) -> Type["CliAdapterBase"]:
|
||||
def register_cli_adapter(adapter_class: type[CliAdapterBase]) -> type[CliAdapterBase]:
|
||||
"""
|
||||
注册 CLI Adapter 类到注册表
|
||||
|
||||
@@ -771,13 +771,13 @@ def _ensure_cli_adapters_loaded():
|
||||
_CLI_ADAPTERS_LOADED = True
|
||||
|
||||
|
||||
def get_cli_adapter_class(api_format: str) -> Optional[Type["CliAdapterBase"]]:
|
||||
def get_cli_adapter_class(api_format: str) -> type[CliAdapterBase] | None:
|
||||
"""根据 API format 获取 CLI Adapter 类"""
|
||||
_ensure_cli_adapters_loaded()
|
||||
return _CLI_ADAPTER_REGISTRY.get(api_format.upper()) if api_format else None
|
||||
|
||||
|
||||
def get_cli_adapter_instance(api_format: str) -> Optional["CliAdapterBase"]:
|
||||
def get_cli_adapter_instance(api_format: str) -> CliAdapterBase | None:
|
||||
"""根据 API format 获取 CLI Adapter 实例"""
|
||||
adapter_class = get_cli_adapter_class(api_format)
|
||||
if adapter_class:
|
||||
|
||||
@@ -10,6 +10,8 @@ CLI Message Handler 通用基类
|
||||
3. 简化新格式接入 - 只需实现 ResponseParser 和少量钩子方法
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import codecs
|
||||
import json
|
||||
@@ -17,14 +19,11 @@ import time
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
AsyncGenerator,
|
||||
Callable,
|
||||
Dict,
|
||||
List,
|
||||
Optional,
|
||||
Tuple,
|
||||
)
|
||||
|
||||
from collections.abc import Callable
|
||||
from collections.abc import AsyncGenerator
|
||||
|
||||
import httpx
|
||||
from fastapi import BackgroundTasks, Request
|
||||
from fastapi.responses import JSONResponse, StreamingResponse
|
||||
@@ -45,7 +44,6 @@ from src.api.handlers.base.request_builder import PassthroughRequestBuilder, get
|
||||
# 直接从具体模块导入,避免循环依赖
|
||||
from src.api.handlers.base.response_parser import (
|
||||
ResponseParser,
|
||||
StreamStats,
|
||||
)
|
||||
from src.api.handlers.base.stream_context import StreamContext
|
||||
from src.api.handlers.base.utils import (
|
||||
@@ -86,7 +84,7 @@ from src.utils.timeout import read_first_chunk_with_ttfb_timeout
|
||||
# ==============================================================================
|
||||
|
||||
|
||||
def _parse_sse_data_line(line: str) -> Tuple[Optional[Any], str]:
|
||||
def _parse_sse_data_line(line: str) -> tuple[Any | None, str]:
|
||||
"""
|
||||
解析标准 SSE data 行
|
||||
|
||||
@@ -108,7 +106,7 @@ def _parse_sse_data_line(line: str) -> Tuple[Optional[Any], str]:
|
||||
return None, "invalid"
|
||||
|
||||
|
||||
def _parse_sse_event_data_line(line: str) -> Tuple[Optional[Any], str]:
|
||||
def _parse_sse_event_data_line(line: str) -> tuple[Any | None, str]:
|
||||
"""
|
||||
解析 event + data 同行格式(如 "event: xxx data: {...}")
|
||||
|
||||
@@ -126,7 +124,7 @@ def _parse_sse_event_data_line(line: str) -> Tuple[Optional[Any], str]:
|
||||
return None, "invalid"
|
||||
|
||||
|
||||
def _parse_gemini_json_array_line(line: str) -> Tuple[Optional[Any], str]:
|
||||
def _parse_gemini_json_array_line(line: str) -> tuple[Any | None, str]:
|
||||
"""
|
||||
解析 Gemini JSON-array 格式的裸 JSON 行
|
||||
|
||||
@@ -151,9 +149,9 @@ def _parse_gemini_json_array_line(line: str) -> Tuple[Optional[Any], str]:
|
||||
|
||||
|
||||
def _format_converted_events_to_sse(
|
||||
converted_events: List[Dict[str, Any]],
|
||||
converted_events: list[dict[str, Any]],
|
||||
client_format: str,
|
||||
) -> List[str]:
|
||||
) -> list[str]:
|
||||
"""
|
||||
将转换后的事件格式化为 SSE 行
|
||||
|
||||
@@ -164,7 +162,7 @@ def _format_converted_events_to_sse(
|
||||
Returns:
|
||||
SSE 行列表(每个元素是完整的 SSE 事件,包含尾部空行)
|
||||
"""
|
||||
result: List[str] = []
|
||||
result: list[str] = []
|
||||
needs_event_line = client_format.upper() in ("CLAUDE", "CLAUDE_CLI")
|
||||
|
||||
for evt in converted_events:
|
||||
@@ -213,10 +211,10 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
client_ip: str,
|
||||
user_agent: str,
|
||||
start_time: float,
|
||||
allowed_api_formats: Optional[list] = None,
|
||||
adapter_detector: Optional[
|
||||
Callable[[Dict[str, str], Optional[Dict[str, Any]]], Dict[str, bool]]
|
||||
] = None,
|
||||
allowed_api_formats: list | None = None,
|
||||
adapter_detector: None | (
|
||||
Callable[[dict[str, str], dict[str, Any] | None], dict[str, bool]]
|
||||
) = None,
|
||||
):
|
||||
allowed = allowed_api_formats or [self.FORMAT_ID]
|
||||
super().__init__(
|
||||
@@ -230,7 +228,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
allowed_api_formats=allowed,
|
||||
adapter_detector=adapter_detector,
|
||||
)
|
||||
self._parser: Optional[ResponseParser] = None
|
||||
self._parser: ResponseParser | None = None
|
||||
self._request_builder = PassthroughRequestBuilder()
|
||||
|
||||
@property
|
||||
@@ -253,7 +251,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
self,
|
||||
source_model: str,
|
||||
provider_id: str,
|
||||
) -> Optional[str]:
|
||||
) -> str | None:
|
||||
"""
|
||||
获取模型映射后的实际模型名
|
||||
|
||||
@@ -296,8 +294,8 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
|
||||
def extract_model_from_request(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
path_params: Optional[Dict[str, Any]] = None, # noqa: ARG002 - 子类使用
|
||||
request_body: dict[str, Any],
|
||||
path_params: dict[str, Any] | None = None, # noqa: ARG002 - 子类使用
|
||||
) -> str:
|
||||
"""
|
||||
从请求中提取模型名 - 子类可覆盖
|
||||
@@ -321,9 +319,9 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
|
||||
def apply_mapped_model(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
request_body: dict[str, Any],
|
||||
mapped_model: str, # noqa: ARG002 - 子类使用
|
||||
) -> Dict[str, Any]:
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
将映射后的模型名应用到请求体
|
||||
|
||||
@@ -342,8 +340,8 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
|
||||
def prepare_provider_request_body(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
) -> Dict[str, Any]:
|
||||
request_body: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
准备发送给 Provider 的请求体 - 子类可覆盖
|
||||
|
||||
@@ -359,7 +357,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
return request_body
|
||||
|
||||
@staticmethod
|
||||
def _get_format_metadata(format_id: str) -> Optional["ApiFormatDefinition"]:
|
||||
def _get_format_metadata(format_id: str) -> ApiFormatDefinition | None:
|
||||
"""获取格式元数据(解析失败返回 None)"""
|
||||
from src.core.api_format import APIFormat
|
||||
from src.core.api_format.metadata import API_FORMAT_DEFINITIONS
|
||||
@@ -372,10 +370,10 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
|
||||
def _finalize_converted_request(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
request_body: dict[str, Any],
|
||||
client_api_format: str,
|
||||
provider_api_format: str,
|
||||
mapped_model: Optional[str],
|
||||
mapped_model: str | None,
|
||||
fallback_model: str,
|
||||
is_stream: bool,
|
||||
) -> None:
|
||||
@@ -418,13 +416,13 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
|
||||
def _convert_request_for_cross_format(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
request_body: dict[str, Any],
|
||||
client_api_format: str,
|
||||
provider_api_format: str,
|
||||
mapped_model: Optional[str],
|
||||
mapped_model: str | None,
|
||||
fallback_model: str,
|
||||
is_stream: bool,
|
||||
) -> Tuple[Dict[str, Any], str]:
|
||||
) -> tuple[dict[str, Any], str]:
|
||||
"""
|
||||
跨格式请求转换的公共逻辑
|
||||
|
||||
@@ -465,9 +463,9 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
|
||||
def get_model_for_url(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
mapped_model: Optional[str],
|
||||
) -> Optional[str]:
|
||||
request_body: dict[str, Any],
|
||||
mapped_model: str | None,
|
||||
) -> str | None:
|
||||
"""
|
||||
获取用于 URL 路径的模型名
|
||||
|
||||
@@ -485,8 +483,8 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
|
||||
def _extract_response_metadata(
|
||||
self,
|
||||
response: Dict[str, Any],
|
||||
) -> Dict[str, Any]:
|
||||
response: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
从响应中提取 Provider 特有的元数据 - 子类可覆盖
|
||||
|
||||
@@ -503,11 +501,11 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
|
||||
async def process_stream(
|
||||
self,
|
||||
original_request_body: Dict[str, Any],
|
||||
original_headers: Dict[str, str],
|
||||
query_params: Optional[Dict[str, str]] = None,
|
||||
path_params: Optional[Dict[str, Any]] = None,
|
||||
http_request: Optional[Request] = None,
|
||||
original_request_body: dict[str, Any],
|
||||
original_headers: dict[str, str],
|
||||
query_params: dict[str, str] | None = None,
|
||||
path_params: dict[str, Any] | None = None,
|
||||
http_request: Request | None = None,
|
||||
) -> StreamingResponse:
|
||||
"""
|
||||
处理流式请求
|
||||
@@ -529,7 +527,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
|
||||
# 可变请求体容器:允许 orchestrator 在遇到 Thinking 签名错误时整流请求体后重试
|
||||
# 结构: {"body": 实际请求体, "_rectified": 是否已整流, "_rectified_this_turn": 本轮是否整流}
|
||||
request_body_ref: Dict[str, Any] = {"body": original_request_body}
|
||||
request_body_ref: dict[str, Any] = {"body": original_request_body}
|
||||
|
||||
# 使用子类实现的方法提取 model(不同 API 格式的 model 位置不同)
|
||||
# 注意:使用 original_request_body,因为整流只修改 messages,不影响 model 字段
|
||||
@@ -550,7 +548,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
endpoint: ProviderEndpoint,
|
||||
key: ProviderAPIKey,
|
||||
candidate: ProviderCandidate,
|
||||
) -> AsyncGenerator[bytes, None]:
|
||||
) -> AsyncGenerator[bytes]:
|
||||
return await self._execute_stream_request(
|
||||
ctx,
|
||||
provider,
|
||||
@@ -653,12 +651,12 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
provider: Provider,
|
||||
endpoint: ProviderEndpoint,
|
||||
key: ProviderAPIKey,
|
||||
original_request_body: Dict[str, Any],
|
||||
original_headers: Dict[str, str],
|
||||
query_params: Optional[Dict[str, str]] = None,
|
||||
candidate: Optional[ProviderCandidate] = None,
|
||||
http_request: Optional[Request] = None,
|
||||
) -> AsyncGenerator[bytes, None]:
|
||||
original_request_body: dict[str, Any],
|
||||
original_headers: dict[str, str],
|
||||
query_params: dict[str, str] | None = None,
|
||||
candidate: ProviderCandidate | None = None,
|
||||
http_request: Request | None = None,
|
||||
) -> AsyncGenerator[bytes]:
|
||||
"""执行流式请求并返回流生成器"""
|
||||
# 重置上下文状态(重试时清除之前的数据,避免累积)
|
||||
ctx.parsed_chunks = []
|
||||
@@ -824,7 +822,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
else:
|
||||
await asyncio.wait_for(_connect_and_prefetch(), timeout=request_timeout)
|
||||
|
||||
except asyncio.TimeoutError:
|
||||
except TimeoutError:
|
||||
# 整体请求超时(建立连接 + 获取首字节)
|
||||
# 清理可能已建立的连接上下文
|
||||
if response_ctx is not None:
|
||||
@@ -898,7 +896,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
stream_response: httpx.Response,
|
||||
response_ctx: Any,
|
||||
http_client: httpx.AsyncClient,
|
||||
) -> AsyncGenerator[bytes, None]:
|
||||
) -> AsyncGenerator[bytes]:
|
||||
"""创建响应流生成器(使用字节流)"""
|
||||
try:
|
||||
sse_parser = SSEEventParser()
|
||||
@@ -961,9 +959,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
},
|
||||
}
|
||||
self._mark_first_output(ctx, output_state)
|
||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode(
|
||||
"utf-8"
|
||||
)
|
||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
|
||||
return # 结束生成器
|
||||
|
||||
# 格式转换或直接透传
|
||||
@@ -1015,7 +1011,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
"message": ctx.error_message,
|
||||
},
|
||||
}
|
||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode("utf-8")
|
||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
|
||||
else:
|
||||
logger.debug("流式数据转发完成")
|
||||
# 为 OpenAI 客户端补齐 [DONE] 标记(非 CLI 格式)
|
||||
@@ -1040,7 +1036,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
"message": ctx.error_message,
|
||||
},
|
||||
}
|
||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode("utf-8")
|
||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
|
||||
except httpx.RemoteProtocolError:
|
||||
if ctx.data_count > 0:
|
||||
error_event = {
|
||||
@@ -1050,7 +1046,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
"message": "上游连接意外关闭,部分响应已成功传输",
|
||||
},
|
||||
}
|
||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode("utf-8")
|
||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
|
||||
else:
|
||||
raise
|
||||
finally:
|
||||
@@ -1241,7 +1237,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
except (EmbeddedErrorException, ProviderTimeoutException, ProviderNotAvailableException):
|
||||
# 重新抛出可重试的 Provider 异常,触发故障转移
|
||||
raise
|
||||
except (OSError, IOError) as e:
|
||||
except OSError as e:
|
||||
# 网络 I/O 异常:记录警告,可能需要重试
|
||||
logger.warning(f" [{self.request_id}] 预读流时发生网络异常: {type(e).__name__}: {e}")
|
||||
except Exception as e:
|
||||
@@ -1261,7 +1257,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
response_ctx: Any,
|
||||
http_client: httpx.AsyncClient,
|
||||
prefetched_chunks: list,
|
||||
) -> AsyncGenerator[bytes, None]:
|
||||
) -> AsyncGenerator[bytes]:
|
||||
"""创建响应流生成器(带预读数据,使用字节流)"""
|
||||
try:
|
||||
sse_parser = SSEEventParser()
|
||||
@@ -1382,9 +1378,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
},
|
||||
}
|
||||
self._mark_first_output(ctx, output_state)
|
||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode(
|
||||
"utf-8"
|
||||
)
|
||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
|
||||
return
|
||||
|
||||
# 格式转换或直接透传
|
||||
@@ -1439,7 +1433,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
"message": ctx.error_message,
|
||||
},
|
||||
}
|
||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode("utf-8")
|
||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
|
||||
else:
|
||||
logger.debug("流式数据转发完成")
|
||||
# 为 OpenAI 客户端补齐 [DONE] 标记(非 CLI 格式)
|
||||
@@ -1463,7 +1457,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
"message": ctx.error_message,
|
||||
},
|
||||
}
|
||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode("utf-8")
|
||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
|
||||
except httpx.RemoteProtocolError:
|
||||
if ctx.data_count > 0:
|
||||
error_event = {
|
||||
@@ -1473,7 +1467,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
"message": "上游连接意外关闭,部分响应已成功传输",
|
||||
},
|
||||
}
|
||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode("utf-8")
|
||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
|
||||
else:
|
||||
raise
|
||||
finally:
|
||||
@@ -1489,7 +1483,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
def _handle_sse_event(
|
||||
self,
|
||||
ctx: StreamContext,
|
||||
event_name: Optional[str],
|
||||
event_name: str | None,
|
||||
data_str: str,
|
||||
record_chunk: bool = False,
|
||||
) -> None:
|
||||
@@ -1538,7 +1532,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
self,
|
||||
ctx: StreamContext,
|
||||
event_type: str,
|
||||
data: Dict[str, Any],
|
||||
data: dict[str, Any],
|
||||
) -> None:
|
||||
"""
|
||||
处理解析后的事件数据 - 子类应覆盖此方法
|
||||
@@ -1612,7 +1606,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
def _record_converted_chunks(
|
||||
self,
|
||||
ctx: StreamContext,
|
||||
converted_events: List[Dict[str, Any]],
|
||||
converted_events: list[dict[str, Any]],
|
||||
) -> None:
|
||||
"""
|
||||
记录转换后的 chunk 数据到 parsed_chunks,并更新统计信息
|
||||
@@ -1656,7 +1650,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
def _extract_usage_from_converted_event(
|
||||
self,
|
||||
ctx: StreamContext,
|
||||
evt: Dict[str, Any],
|
||||
evt: dict[str, Any],
|
||||
event_type: str,
|
||||
) -> None:
|
||||
"""
|
||||
@@ -1672,7 +1666,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
evt: 转换后的事件
|
||||
event_type: 事件类型
|
||||
"""
|
||||
usage: Optional[Dict[str, Any]] = None
|
||||
usage: dict[str, Any] | None = None
|
||||
|
||||
# Claude 格式: message_delta 或 message_start
|
||||
if event_type == "message_delta":
|
||||
@@ -1737,9 +1731,9 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
async def _create_monitored_stream(
|
||||
self,
|
||||
ctx: StreamContext,
|
||||
stream_generator: AsyncGenerator[bytes, None],
|
||||
http_request: Optional[Request] = None,
|
||||
) -> AsyncGenerator[bytes, None]:
|
||||
stream_generator: AsyncGenerator[bytes],
|
||||
http_request: Request | None = None,
|
||||
) -> AsyncGenerator[bytes]:
|
||||
"""
|
||||
创建带监控的流生成器
|
||||
|
||||
@@ -1833,8 +1827,8 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
async def _record_stream_stats(
|
||||
self,
|
||||
ctx: StreamContext,
|
||||
original_headers: Dict[str, str],
|
||||
original_request_body: Dict[str, Any],
|
||||
original_headers: dict[str, str],
|
||||
original_request_body: dict[str, Any],
|
||||
) -> None:
|
||||
"""在流完成后记录统计信息"""
|
||||
try:
|
||||
@@ -1996,7 +1990,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
from src.services.request.candidate import RequestCandidateService
|
||||
|
||||
# 计算候选自身的 TTFB
|
||||
candidate_first_byte_time_ms: Optional[int] = None
|
||||
candidate_first_byte_time_ms: int | None = None
|
||||
if ctx.first_byte_time_ms is not None:
|
||||
candidate_first_byte_time_ms = (
|
||||
RequestCandidateService.calculate_candidate_ttfb(
|
||||
@@ -2061,8 +2055,8 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
self,
|
||||
ctx: StreamContext,
|
||||
error: Exception,
|
||||
original_headers: Dict[str, str],
|
||||
original_request_body: Dict[str, Any],
|
||||
original_headers: dict[str, str],
|
||||
original_request_body: dict[str, Any],
|
||||
) -> None:
|
||||
"""记录流式请求失败"""
|
||||
# 使用 self.start_time 作为时间基准,与首字时间保持一致
|
||||
@@ -2111,10 +2105,10 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
|
||||
async def process_sync(
|
||||
self,
|
||||
original_request_body: Dict[str, Any],
|
||||
original_headers: Dict[str, str],
|
||||
query_params: Optional[Dict[str, str]] = None,
|
||||
path_params: Optional[Dict[str, Any]] = None,
|
||||
original_request_body: dict[str, Any],
|
||||
original_headers: dict[str, str],
|
||||
query_params: dict[str, str] | None = None,
|
||||
path_params: dict[str, Any] | None = None,
|
||||
) -> JSONResponse:
|
||||
"""
|
||||
处理非流式请求
|
||||
@@ -2142,19 +2136,19 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
endpoint_id = None # Endpoint ID(用于失败记录)
|
||||
key_id = None # Key ID(用于失败记录)
|
||||
mapped_model_result = None # 映射后的目标模型名(用于 Usage 记录)
|
||||
response_metadata_result: Dict[str, Any] = {} # Provider 响应元数据
|
||||
response_metadata_result: dict[str, Any] = {} # Provider 响应元数据
|
||||
needs_conversion = False # 是否需要格式转换(由 candidate 决定)
|
||||
|
||||
# 可变请求体容器:允许 orchestrator 在遇到 Thinking 签名错误时整流请求体后重试
|
||||
# 结构: {"body": 实际请求体, "_rectified": 是否已整流, "_rectified_this_turn": 本轮是否整流}
|
||||
request_body_ref: Dict[str, Any] = {"body": original_request_body}
|
||||
request_body_ref: dict[str, Any] = {"body": original_request_body}
|
||||
|
||||
async def sync_request_func(
|
||||
provider: Provider,
|
||||
endpoint: ProviderEndpoint,
|
||||
key: ProviderAPIKey,
|
||||
candidate: ProviderCandidate,
|
||||
) -> Dict[str, Any]:
|
||||
) -> dict[str, Any]:
|
||||
nonlocal provider_name, response_json, status_code, response_headers, provider_api_format, provider_request_headers, provider_request_body, mapped_model_result, response_metadata_result, needs_conversion
|
||||
provider_name = str(provider.name)
|
||||
provider_api_format = str(endpoint.api_format) if endpoint.api_format else ""
|
||||
@@ -2470,7 +2464,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
actual_request_body = provider_request_body or original_request_body
|
||||
|
||||
# 尝试从异常中提取响应头
|
||||
error_response_headers: Dict[str, str] = {}
|
||||
error_response_headers: dict[str, str] = {}
|
||||
if isinstance(e, ProviderRateLimitException) and e.response_headers:
|
||||
error_response_headers = e.response_headers
|
||||
elif isinstance(e, httpx.HTTPStatusError) and hasattr(e, "response"):
|
||||
@@ -2581,7 +2575,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
)
|
||||
return True
|
||||
|
||||
def _mark_first_output(self, ctx: StreamContext, state: Dict[str, bool]) -> None:
|
||||
def _mark_first_output(self, ctx: StreamContext, state: dict[str, bool]) -> None:
|
||||
"""
|
||||
标记首次输出:记录 TTFB 并更新 streaming 状态
|
||||
|
||||
@@ -2605,7 +2599,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
ctx: StreamContext,
|
||||
line: str,
|
||||
events: list, # noqa: ARG002 - 预留给上下文感知转换
|
||||
) -> Tuple[List[str], List[Dict[str, Any]]]:
|
||||
) -> tuple[list[str], list[dict[str, Any]]]:
|
||||
"""
|
||||
将 SSE 行从 Provider 格式转换为客户端格式
|
||||
|
||||
@@ -2690,7 +2684,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
|
||||
def _parse_sse_line_to_json(
|
||||
self, line: str, provider_format: str
|
||||
) -> Tuple[Optional[Any], str]:
|
||||
) -> tuple[Any | None, str]:
|
||||
"""
|
||||
解析 SSE 行为 JSON 对象
|
||||
|
||||
|
||||
@@ -8,7 +8,6 @@ StreamSmoother 使用这些提取器来处理不同格式的 SSE 事件。
|
||||
import copy
|
||||
import json
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Optional
|
||||
|
||||
|
||||
class ContentExtractor(ABC):
|
||||
@@ -20,7 +19,7 @@ class ContentExtractor(ABC):
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def extract_content(self, data: dict) -> Optional[str]:
|
||||
def extract_content(self, data: dict) -> str | None:
|
||||
"""
|
||||
从 SSE 数据中提取可拆分的文本内容
|
||||
|
||||
@@ -64,7 +63,7 @@ class OpenAIContentExtractor(ContentExtractor):
|
||||
- 只在 delta 仅包含 role/content 时允许拆分,避免破坏 tool_calls 等结构
|
||||
"""
|
||||
|
||||
def extract_content(self, data: dict) -> Optional[str]:
|
||||
def extract_content(self, data: dict) -> str | None:
|
||||
if not isinstance(data, dict):
|
||||
return None
|
||||
|
||||
@@ -115,7 +114,7 @@ class OpenAIContentExtractor(ContentExtractor):
|
||||
new_choices.append(new_choice)
|
||||
new_data["choices"] = new_choices
|
||||
|
||||
return f"data: {json.dumps(new_data, ensure_ascii=False)}\n\n".encode("utf-8")
|
||||
return f"data: {json.dumps(new_data, ensure_ascii=False)}\n\n".encode()
|
||||
|
||||
|
||||
class ClaudeContentExtractor(ContentExtractor):
|
||||
@@ -127,7 +126,7 @@ class ClaudeContentExtractor(ContentExtractor):
|
||||
- 数据结构: delta.type=text_delta, delta.text
|
||||
"""
|
||||
|
||||
def extract_content(self, data: dict) -> Optional[str]:
|
||||
def extract_content(self, data: dict) -> str | None:
|
||||
if not isinstance(data, dict):
|
||||
return None
|
||||
|
||||
@@ -165,9 +164,7 @@ class ClaudeContentExtractor(ContentExtractor):
|
||||
|
||||
# Claude 格式需要 event: 前缀
|
||||
event_name = event_type or "content_block_delta"
|
||||
return f"event: {event_name}\ndata: {json.dumps(new_data, ensure_ascii=False)}\n\n".encode(
|
||||
"utf-8"
|
||||
)
|
||||
return f"event: {event_name}\ndata: {json.dumps(new_data, ensure_ascii=False)}\n\n".encode()
|
||||
|
||||
|
||||
class GeminiContentExtractor(ContentExtractor):
|
||||
@@ -179,7 +176,7 @@ class GeminiContentExtractor(ContentExtractor):
|
||||
- 只有纯文本块才拆分
|
||||
"""
|
||||
|
||||
def extract_content(self, data: dict) -> Optional[str]:
|
||||
def extract_content(self, data: dict) -> str | None:
|
||||
if not isinstance(data, dict):
|
||||
return None
|
||||
|
||||
@@ -226,7 +223,7 @@ class GeminiContentExtractor(ContentExtractor):
|
||||
if "parts" in content and content["parts"]:
|
||||
content["parts"][0]["text"] = new_content
|
||||
|
||||
return f"data: {json.dumps(new_data, ensure_ascii=False)}\n\n".encode("utf-8")
|
||||
return f"data: {json.dumps(new_data, ensure_ascii=False)}\n\n".encode()
|
||||
|
||||
|
||||
# 提取器注册表
|
||||
@@ -237,7 +234,7 @@ _EXTRACTORS: dict[str, type[ContentExtractor]] = {
|
||||
}
|
||||
|
||||
|
||||
def get_extractor(format_name: str) -> Optional[ContentExtractor]:
|
||||
def get_extractor(format_name: str) -> ContentExtractor | None:
|
||||
"""
|
||||
根据格式名获取对应的内容提取器实例
|
||||
|
||||
|
||||
@@ -14,15 +14,14 @@
|
||||
- EndpointCheckOrchestrator: 协调整个流程
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, AsyncIterator, Dict, Iterable, Optional, Union, List
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Any
|
||||
from collections.abc import Iterable
|
||||
import time
|
||||
import uuid
|
||||
import json
|
||||
from functools import lru_cache
|
||||
import asyncio
|
||||
from collections import defaultdict
|
||||
|
||||
import httpx
|
||||
|
||||
@@ -31,7 +30,7 @@ from src.core.api_format import CORE_REDACT_HEADERS, merge_headers_with_protecti
|
||||
from src.utils.ssl_utils import get_ssl_context
|
||||
|
||||
|
||||
def _redact_headers(headers: Dict[str, str]) -> Dict[str, str]:
|
||||
def _redact_headers(headers: dict[str, str]) -> dict[str, str]:
|
||||
return redact_headers_for_log(headers, CORE_REDACT_HEADERS)
|
||||
|
||||
|
||||
@@ -46,10 +45,10 @@ def _truncate_repr(value: Any, limit: int = 1200) -> str:
|
||||
|
||||
|
||||
def build_safe_headers(
|
||||
base_headers: Dict[str, str],
|
||||
extra_headers: Optional[Dict[str, str]],
|
||||
base_headers: dict[str, str],
|
||||
extra_headers: dict[str, str] | None,
|
||||
protected_keys: Iterable[str],
|
||||
) -> Dict[str, str]:
|
||||
) -> dict[str, str]:
|
||||
"""
|
||||
合并 extra_headers,但防止覆盖 protected_keys(大小写不敏感)。
|
||||
"""
|
||||
@@ -60,16 +59,16 @@ async def run_endpoint_check(
|
||||
*,
|
||||
client: httpx.AsyncClient, # 保持兼容性,但内部不使用
|
||||
url: str,
|
||||
headers: Dict[str, str],
|
||||
json_body: Dict[str, Any],
|
||||
headers: dict[str, str],
|
||||
json_body: dict[str, Any],
|
||||
api_format: str,
|
||||
provider_name: Optional[str] = None,
|
||||
model_name: Optional[str] = None,
|
||||
api_key_id: Optional[str] = None,
|
||||
provider_id: Optional[str] = None,
|
||||
db: Optional[Any] = None, # Session对象,需要时才导入
|
||||
user: Optional[Any] = None, # User对象
|
||||
) -> Dict[str, Any]:
|
||||
provider_name: str | None = None,
|
||||
model_name: str | None = None,
|
||||
api_key_id: str | None = None,
|
||||
provider_id: str | None = None,
|
||||
db: Any | None = None, # Session对象,需要时才导入
|
||||
user: Any | None = None, # User对象
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
执行端点检查(重构版本,使用新的架构):
|
||||
- 使用新的架构类来分离关注点
|
||||
@@ -123,21 +122,21 @@ async def _calculate_and_record_usage(
|
||||
provider_id: str,
|
||||
api_key_id: str,
|
||||
model_name: str,
|
||||
request_data: Dict[str, Any],
|
||||
response_data: Optional[Dict[str, Any]],
|
||||
request_data: dict[str, Any],
|
||||
response_data: dict[str, Any] | None,
|
||||
request_id: str,
|
||||
response_time_ms: int,
|
||||
request_headers: Dict[str, str],
|
||||
response_headers: Optional[Dict[str, str]] = None,
|
||||
request_headers: dict[str, str],
|
||||
response_headers: dict[str, str] | None = None,
|
||||
status_code: int = 0,
|
||||
error_message: Optional[str] = None,
|
||||
error_message: str | None = None,
|
||||
# 新增:支持直接传递token数据
|
||||
input_tokens: Optional[int] = None,
|
||||
output_tokens: Optional[int] = None,
|
||||
cache_creation_input_tokens: Optional[int] = None,
|
||||
cache_read_input_tokens: Optional[int] = None,
|
||||
api_format: Optional[str] = None,
|
||||
) -> Dict[str, Any]:
|
||||
input_tokens: int | None = None,
|
||||
output_tokens: int | None = None,
|
||||
cache_creation_input_tokens: int | None = None,
|
||||
cache_read_input_tokens: int | None = None,
|
||||
api_format: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
计算并记录用量数据(遗留函数)
|
||||
|
||||
@@ -149,7 +148,7 @@ async def _calculate_and_record_usage(
|
||||
"""
|
||||
from src.services.usage.service import UsageService
|
||||
from src.services.request.candidate import RequestCandidateService
|
||||
from src.models.database import ApiKey, ProviderAPIKey, ProviderEndpoint
|
||||
from src.models.database import ApiKey, ProviderAPIKey
|
||||
|
||||
# 获取Provider API Key对象(不是用户API Key)
|
||||
provider_api_key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == api_key_id).first()
|
||||
@@ -360,7 +359,7 @@ async def _calculate_and_record_usage(
|
||||
}
|
||||
|
||||
|
||||
def _extract_tokens_from_response(api_identifier: str, response_data: Optional[Dict[str, Any]]) -> tuple[int, int, int, int]:
|
||||
def _extract_tokens_from_response(api_identifier: str, response_data: dict[str, Any] | None) -> tuple[int, int, int, int]:
|
||||
"""
|
||||
从响应中提取Token计数信息
|
||||
|
||||
@@ -446,7 +445,7 @@ def _extract_tokens_from_response(api_identifier: str, response_data: Optional[D
|
||||
|
||||
|
||||
|
||||
def _fallback_token_counting(request_data: Dict[str, Any], response_data: Optional[Dict[str, Any]]) -> tuple[int, int, int, int]:
|
||||
def _fallback_token_counting(request_data: dict[str, Any], response_data: dict[str, Any] | None) -> tuple[int, int, int, int]:
|
||||
"""
|
||||
回退的Token计数方法(简单估算)
|
||||
|
||||
@@ -508,16 +507,16 @@ def _fallback_token_counting(request_data: Dict[str, Any], response_data: Option
|
||||
class EndpointCheckRequest:
|
||||
"""端点检查请求数据类"""
|
||||
url: str
|
||||
headers: Dict[str, str]
|
||||
json_body: Dict[str, Any]
|
||||
headers: dict[str, str]
|
||||
json_body: dict[str, Any]
|
||||
api_format: str
|
||||
provider_name: Optional[str] = None
|
||||
model_name: Optional[str] = None
|
||||
api_key_id: Optional[str] = None
|
||||
provider_id: Optional[str] = None
|
||||
db: Optional[Any] = None
|
||||
user: Optional[Any] = None
|
||||
request_id: Optional[str] = None
|
||||
provider_name: str | None = None
|
||||
model_name: str | None = None
|
||||
api_key_id: str | None = None
|
||||
provider_id: str | None = None
|
||||
db: Any | None = None
|
||||
user: Any | None = None
|
||||
request_id: str | None = None
|
||||
timeout: float = 30.0
|
||||
|
||||
|
||||
@@ -525,12 +524,12 @@ class EndpointCheckRequest:
|
||||
class EndpointCheckResult:
|
||||
"""端点检查结果数据类"""
|
||||
status_code: int
|
||||
headers: Dict[str, str]
|
||||
headers: dict[str, str]
|
||||
response_time_ms: int
|
||||
request_id: str
|
||||
response_data: Optional[Dict[str, Any]] = None
|
||||
error_message: Optional[str] = None
|
||||
usage_data: Optional[Dict[str, Any]] = None
|
||||
response_data: dict[str, Any] | None = None
|
||||
error_message: str | None = None
|
||||
usage_data: dict[str, Any] | None = None
|
||||
|
||||
|
||||
class HttpRequestExecutor:
|
||||
@@ -613,7 +612,7 @@ class UsageCalculator:
|
||||
return _extract_tokens_from_response(api_identifier, result.response_data)
|
||||
|
||||
@staticmethod
|
||||
def _fallback_token_counting(request_data: Dict[str, Any], response_data: Optional[Dict[str, Any]]) -> tuple[int, int, int, int]:
|
||||
def _fallback_token_counting(request_data: dict[str, Any], response_data: dict[str, Any] | None) -> tuple[int, int, int, int]:
|
||||
"""回退的Token计数方法(简单估算)"""
|
||||
# 估算输入Token
|
||||
messages = request_data.get("messages", request_data.get("contents", []))
|
||||
@@ -665,12 +664,12 @@ class AsyncBatchUsageRecorder:
|
||||
def __init__(self, batch_size: int = 10, flush_interval: float = 2.0):
|
||||
self.batch_size = batch_size
|
||||
self.flush_interval = flush_interval
|
||||
self.pending_records: List[Dict[str, Any]] = []
|
||||
self._flush_task: Optional[asyncio.Task] = None
|
||||
self.pending_records: list[dict[str, Any]] = []
|
||||
self._flush_task: asyncio.Task | None = None
|
||||
self._lock = asyncio.Lock()
|
||||
self._running = True
|
||||
|
||||
async def add_record(self, usage_data: Dict[str, Any]) -> None:
|
||||
async def add_record(self, usage_data: dict[str, Any]) -> None:
|
||||
"""添加用量记录到批处理队列"""
|
||||
async with self._lock:
|
||||
self.pending_records.append(usage_data)
|
||||
@@ -740,7 +739,7 @@ class AsyncBatchUsageRecorder:
|
||||
|
||||
|
||||
# 全局批处理器实例(单例)
|
||||
_global_batch_recorder: Optional[AsyncBatchUsageRecorder] = None
|
||||
_global_batch_recorder: AsyncBatchUsageRecorder | None = None
|
||||
|
||||
def get_batch_recorder() -> AsyncBatchUsageRecorder:
|
||||
"""获取全局批处理器实例"""
|
||||
@@ -756,7 +755,7 @@ def get_batch_recorder() -> AsyncBatchUsageRecorder:
|
||||
|
||||
class EndpointCheckError(Exception):
|
||||
"""端点检查错误基类"""
|
||||
def __init__(self, message: str, error_type: str, status_code: int = 500, details: Optional[Dict[str, Any]] = None):
|
||||
def __init__(self, message: str, error_type: str, status_code: int = 500, details: dict[str, Any] | None = None):
|
||||
super().__init__(message)
|
||||
self.message = message
|
||||
self.error_type = error_type
|
||||
@@ -765,22 +764,22 @@ class EndpointCheckError(Exception):
|
||||
|
||||
class NetworkError(EndpointCheckError):
|
||||
"""网络请求错误"""
|
||||
def __init__(self, message: str, details: Optional[Dict[str, Any]] = None):
|
||||
def __init__(self, message: str, details: dict[str, Any] | None = None):
|
||||
super().__init__(message, "network_error", 0, details)
|
||||
|
||||
class AuthenticationError(EndpointCheckError):
|
||||
"""认证错误"""
|
||||
def __init__(self, message: str, details: Optional[Dict[str, Any]] = None):
|
||||
def __init__(self, message: str, details: dict[str, Any] | None = None):
|
||||
super().__init__(message, "authentication_error", 401, details)
|
||||
|
||||
class RateLimitError(EndpointCheckError):
|
||||
"""速率限制错误"""
|
||||
def __init__(self, message: str, details: Optional[Dict[str, Any]] = None):
|
||||
def __init__(self, message: str, details: dict[str, Any] | None = None):
|
||||
super().__init__(message, "rate_limit_error", 429, details)
|
||||
|
||||
class UpstreamError(EndpointCheckError):
|
||||
"""上游服务错误"""
|
||||
def __init__(self, message: str, status_code: int, details: Optional[Dict[str, Any]] = None):
|
||||
def __init__(self, message: str, status_code: int, details: dict[str, Any] | None = None):
|
||||
super().__init__(message, "upstream_error", status_code, details)
|
||||
|
||||
|
||||
@@ -982,7 +981,7 @@ class EndpointCheckConfig:
|
||||
retry_on_timeouts: bool = True
|
||||
|
||||
@classmethod
|
||||
def from_env(cls) -> 'EndpointCheckConfig':
|
||||
def from_env(cls) -> EndpointCheckConfig:
|
||||
"""从环境变量创建配置"""
|
||||
import os
|
||||
|
||||
@@ -1004,7 +1003,7 @@ class EndpointCheckConfig:
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, config_dict: Dict[str, Any]) -> 'EndpointCheckConfig':
|
||||
def from_dict(cls, config_dict: dict[str, Any]) -> EndpointCheckConfig:
|
||||
"""从字典创建配置"""
|
||||
return cls(**{k: v for k, v in config_dict.items() if hasattr(cls, k)})
|
||||
|
||||
@@ -1012,7 +1011,7 @@ class EndpointCheckConfig:
|
||||
class ConfigurableEndpointChecker:
|
||||
"""可配置的端点检查器"""
|
||||
|
||||
def __init__(self, config: Optional[EndpointCheckConfig] = None):
|
||||
def __init__(self, config: EndpointCheckConfig | None = None):
|
||||
self.config = config or EndpointCheckConfig()
|
||||
self.executor = HttpRequestExecutor(timeout=self.config.timeout)
|
||||
self.usage_calculator = UsageCalculator()
|
||||
@@ -1171,9 +1170,9 @@ class ConfigurableEndpointChecker:
|
||||
|
||||
|
||||
# 全局配置检查器实例
|
||||
_global_configured_checker: Optional[ConfigurableEndpointChecker] = None
|
||||
_global_configured_checker: ConfigurableEndpointChecker | None = None
|
||||
|
||||
def get_configured_checker(config: Optional[EndpointCheckConfig] = None) -> ConfigurableEndpointChecker:
|
||||
def get_configured_checker(config: EndpointCheckConfig | None = None) -> ConfigurableEndpointChecker:
|
||||
"""获取全局配置检查器实例"""
|
||||
global _global_configured_checker
|
||||
if _global_configured_checker is None or config is not None:
|
||||
@@ -1186,8 +1185,8 @@ def get_configured_checker(config: Optional[EndpointCheckConfig] = None) -> Conf
|
||||
class EndpointCheckOrchestrator:
|
||||
"""端点检查协调器 - 协调整个流程"""
|
||||
|
||||
def __init__(self, executor: Optional[HttpRequestExecutor] = None,
|
||||
usage_calculator: Optional[UsageCalculator] = None):
|
||||
def __init__(self, executor: HttpRequestExecutor | None = None,
|
||||
usage_calculator: UsageCalculator | None = None):
|
||||
self.executor = executor or HttpRequestExecutor()
|
||||
self.usage_calculator = usage_calculator or UsageCalculator()
|
||||
|
||||
|
||||
@@ -6,7 +6,7 @@
|
||||
"""
|
||||
|
||||
import re
|
||||
from typing import Any, Dict, Optional, Tuple, Type
|
||||
from typing import Any
|
||||
|
||||
from src.api.handlers.base.response_parser import (
|
||||
ParsedChunk,
|
||||
@@ -20,7 +20,7 @@ from src.api.handlers.base.utils import extract_cache_creation_tokens
|
||||
from src.core.api_format import is_cli_format
|
||||
|
||||
|
||||
def _check_nested_error(response: Dict[str, Any]) -> Tuple[bool, Optional[Dict[str, Any]]]:
|
||||
def _check_nested_error(response: dict[str, Any]) -> tuple[bool, dict[str, Any] | None]:
|
||||
"""
|
||||
检查响应中是否存在嵌套错误(某些代理服务返回 HTTP 200 但在响应体中包含错误)
|
||||
|
||||
@@ -62,7 +62,7 @@ def _check_nested_error(response: Dict[str, Any]) -> Tuple[bool, Optional[Dict[s
|
||||
return False, None
|
||||
|
||||
|
||||
def _extract_embedded_status_code(error_info: Optional[Dict[str, Any]]) -> Optional[int]:
|
||||
def _extract_embedded_status_code(error_info: dict[str, Any] | None) -> int | None:
|
||||
"""
|
||||
从错误信息中提取嵌套的状态码
|
||||
|
||||
@@ -137,7 +137,7 @@ class OpenAIResponseParser(ResponseParser):
|
||||
self.name = "OPENAI"
|
||||
self.api_format = "OPENAI"
|
||||
|
||||
def parse_sse_line(self, line: str, stats: StreamStats) -> Optional[ParsedChunk]:
|
||||
def parse_sse_line(self, line: str, stats: StreamStats) -> ParsedChunk | None:
|
||||
if not line or not line.strip():
|
||||
return None
|
||||
|
||||
@@ -186,7 +186,7 @@ class OpenAIResponseParser(ResponseParser):
|
||||
|
||||
return chunk
|
||||
|
||||
def parse_response(self, response: Dict[str, Any], status_code: int) -> ParsedResponse:
|
||||
def parse_response(self, response: dict[str, Any], status_code: int) -> ParsedResponse:
|
||||
result = ParsedResponse(
|
||||
raw_response=response,
|
||||
status_code=status_code,
|
||||
@@ -217,7 +217,7 @@ class OpenAIResponseParser(ResponseParser):
|
||||
|
||||
return result
|
||||
|
||||
def extract_usage_from_response(self, response: Dict[str, Any]) -> Dict[str, int]:
|
||||
def extract_usage_from_response(self, response: dict[str, Any]) -> dict[str, int]:
|
||||
usage = response.get("usage") or {}
|
||||
return {
|
||||
"input_tokens": usage.get("prompt_tokens", 0),
|
||||
@@ -226,7 +226,7 @@ class OpenAIResponseParser(ResponseParser):
|
||||
"cache_read_tokens": 0,
|
||||
}
|
||||
|
||||
def extract_text_content(self, response: Dict[str, Any]) -> str:
|
||||
def extract_text_content(self, response: dict[str, Any]) -> str:
|
||||
choices = response.get("choices", [])
|
||||
if choices:
|
||||
message = choices[0].get("message", {})
|
||||
@@ -235,7 +235,7 @@ class OpenAIResponseParser(ResponseParser):
|
||||
return content
|
||||
return ""
|
||||
|
||||
def is_error_response(self, response: Dict[str, Any]) -> bool:
|
||||
def is_error_response(self, response: dict[str, Any]) -> bool:
|
||||
is_error, _ = _check_nested_error(response)
|
||||
return is_error
|
||||
|
||||
@@ -259,7 +259,7 @@ class ClaudeResponseParser(ResponseParser):
|
||||
self.name = "CLAUDE"
|
||||
self.api_format = "CLAUDE"
|
||||
|
||||
def parse_sse_line(self, line: str, stats: StreamStats) -> Optional[ParsedChunk]:
|
||||
def parse_sse_line(self, line: str, stats: StreamStats) -> ParsedChunk | None:
|
||||
if not line or not line.strip():
|
||||
return None
|
||||
|
||||
@@ -324,7 +324,7 @@ class ClaudeResponseParser(ResponseParser):
|
||||
|
||||
return chunk
|
||||
|
||||
def parse_response(self, response: Dict[str, Any], status_code: int) -> ParsedResponse:
|
||||
def parse_response(self, response: dict[str, Any], status_code: int) -> ParsedResponse:
|
||||
result = ParsedResponse(
|
||||
raw_response=response,
|
||||
status_code=status_code,
|
||||
@@ -358,7 +358,7 @@ class ClaudeResponseParser(ResponseParser):
|
||||
|
||||
return result
|
||||
|
||||
def extract_usage_from_response(self, response: Dict[str, Any]) -> Dict[str, int]:
|
||||
def extract_usage_from_response(self, response: dict[str, Any]) -> dict[str, int]:
|
||||
# 对于 message_start 事件,usage 在 message.usage 路径下
|
||||
# 对于其他响应,usage 在顶层
|
||||
usage = response.get("usage") or {}
|
||||
@@ -372,7 +372,7 @@ class ClaudeResponseParser(ResponseParser):
|
||||
"cache_read_tokens": usage.get("cache_read_input_tokens", 0),
|
||||
}
|
||||
|
||||
def extract_text_content(self, response: Dict[str, Any]) -> str:
|
||||
def extract_text_content(self, response: dict[str, Any]) -> str:
|
||||
content = response.get("content", [])
|
||||
if isinstance(content, list):
|
||||
text_parts = []
|
||||
@@ -382,7 +382,7 @@ class ClaudeResponseParser(ResponseParser):
|
||||
return "".join(text_parts)
|
||||
return ""
|
||||
|
||||
def is_error_response(self, response: Dict[str, Any]) -> bool:
|
||||
def is_error_response(self, response: dict[str, Any]) -> bool:
|
||||
is_error, _ = _check_nested_error(response)
|
||||
return is_error
|
||||
|
||||
@@ -406,7 +406,7 @@ class GeminiResponseParser(ResponseParser):
|
||||
self.name = "GEMINI"
|
||||
self.api_format = "GEMINI"
|
||||
|
||||
def parse_sse_line(self, line: str, stats: StreamStats) -> Optional[ParsedChunk]:
|
||||
def parse_sse_line(self, line: str, stats: StreamStats) -> ParsedChunk | None:
|
||||
"""
|
||||
解析 Gemini SSE 行
|
||||
|
||||
@@ -473,7 +473,7 @@ class GeminiResponseParser(ResponseParser):
|
||||
|
||||
return chunk
|
||||
|
||||
def parse_response(self, response: Dict[str, Any], status_code: int) -> ParsedResponse:
|
||||
def parse_response(self, response: dict[str, Any], status_code: int) -> ParsedResponse:
|
||||
result = ParsedResponse(
|
||||
raw_response=response,
|
||||
status_code=status_code,
|
||||
@@ -509,7 +509,7 @@ class GeminiResponseParser(ResponseParser):
|
||||
|
||||
return result
|
||||
|
||||
def extract_usage_from_response(self, response: Dict[str, Any]) -> Dict[str, int]:
|
||||
def extract_usage_from_response(self, response: dict[str, Any]) -> dict[str, int]:
|
||||
"""
|
||||
从 Gemini 响应中提取 token 使用量
|
||||
|
||||
@@ -531,7 +531,7 @@ class GeminiResponseParser(ResponseParser):
|
||||
"cache_read_tokens": usage.get("cached_tokens", 0),
|
||||
}
|
||||
|
||||
def extract_text_content(self, response: Dict[str, Any]) -> str:
|
||||
def extract_text_content(self, response: dict[str, Any]) -> str:
|
||||
candidates = response.get("candidates", [])
|
||||
if candidates:
|
||||
content = candidates[0].get("content", {})
|
||||
@@ -543,7 +543,7 @@ class GeminiResponseParser(ResponseParser):
|
||||
return "".join(text_parts)
|
||||
return ""
|
||||
|
||||
def is_error_response(self, response: Dict[str, Any]) -> bool:
|
||||
def is_error_response(self, response: dict[str, Any]) -> bool:
|
||||
"""
|
||||
判断响应是否为错误响应
|
||||
|
||||
@@ -562,7 +562,7 @@ class GeminiCliResponseParser(GeminiResponseParser):
|
||||
|
||||
|
||||
# 解析器注册表
|
||||
_PARSERS: Dict[str, Type[ResponseParser]] = {
|
||||
_PARSERS: dict[str, type[ResponseParser]] = {
|
||||
"CLAUDE": ClaudeResponseParser,
|
||||
"CLAUDE_CLI": ClaudeCliResponseParser,
|
||||
"OPENAI": OpenAIResponseParser,
|
||||
|
||||
@@ -11,12 +11,11 @@
|
||||
payload, headers = builder.build(original_body, original_headers, endpoint, key)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Any, Dict, FrozenSet, Optional, Tuple
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from src.core.crypto import crypto_service
|
||||
from src.core.api_format import HeaderBuilder, UPSTREAM_DROP_HEADERS
|
||||
@@ -37,9 +36,9 @@ class ProviderAuthInfo:
|
||||
auth_header: str
|
||||
auth_value: str
|
||||
# 解密后的认证配置(用于 URL 构建等场景,避免重复解密)
|
||||
decrypted_auth_config: Optional[Dict[str, Any]] = None
|
||||
decrypted_auth_config: dict[str, Any] | None = None
|
||||
|
||||
def as_tuple(self) -> Tuple[str, str]:
|
||||
def as_tuple(self) -> tuple[str, str]:
|
||||
"""返回 (auth_header, auth_value) 元组"""
|
||||
return (self.auth_header, self.auth_value)
|
||||
|
||||
@@ -48,7 +47,7 @@ class ProviderAuthInfo:
|
||||
# ==============================================================================
|
||||
|
||||
# 兼容别名:历史代码使用 SENSITIVE_HEADERS 命名
|
||||
SENSITIVE_HEADERS: FrozenSet[str] = UPSTREAM_DROP_HEADERS
|
||||
SENSITIVE_HEADERS: frozenset[str] = UPSTREAM_DROP_HEADERS
|
||||
|
||||
|
||||
# ==============================================================================
|
||||
@@ -57,14 +56,14 @@ SENSITIVE_HEADERS: FrozenSet[str] = UPSTREAM_DROP_HEADERS
|
||||
|
||||
# 标准测试请求体(OpenAI 格式)
|
||||
# 用于 check_endpoint 等测试场景,使用简单安全的消息内容避免触发安全过滤
|
||||
DEFAULT_TEST_REQUEST: Dict[str, Any] = {
|
||||
DEFAULT_TEST_REQUEST: dict[str, Any] = {
|
||||
"messages": [{"role": "user", "content": "Hi"}],
|
||||
"max_tokens": 5,
|
||||
"temperature": 0,
|
||||
}
|
||||
|
||||
|
||||
def get_test_request_data(request_data: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
|
||||
def get_test_request_data(request_data: dict[str, Any] | None = None) -> dict[str, Any]:
|
||||
"""获取测试请求数据
|
||||
|
||||
如果传入 request_data,则合并到默认测试请求中;
|
||||
@@ -85,8 +84,8 @@ def get_test_request_data(request_data: Optional[Dict[str, Any]] = None) -> Dict
|
||||
|
||||
def build_test_request_body(
|
||||
format_id: str,
|
||||
request_data: Optional[Dict[str, Any]] = None,
|
||||
) -> Dict[str, Any]:
|
||||
request_data: dict[str, Any] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""构建测试请求体,自动处理格式转换
|
||||
|
||||
使用格式转换注册表将 OpenAI 格式的测试请求转换为目标格式。
|
||||
@@ -127,39 +126,39 @@ class RequestBuilder(ABC):
|
||||
@abstractmethod
|
||||
def build_payload(
|
||||
self,
|
||||
original_body: Dict[str, Any],
|
||||
original_body: dict[str, Any],
|
||||
*,
|
||||
mapped_model: Optional[str] = None,
|
||||
mapped_model: str | None = None,
|
||||
is_stream: bool = False,
|
||||
) -> Dict[str, Any]:
|
||||
) -> dict[str, Any]:
|
||||
"""构建请求体"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def build_headers(
|
||||
self,
|
||||
original_headers: Dict[str, str],
|
||||
original_headers: dict[str, str],
|
||||
endpoint: Any,
|
||||
key: Any,
|
||||
*,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
pre_computed_auth: Optional[Tuple[str, str]] = None,
|
||||
) -> Dict[str, str]:
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
pre_computed_auth: tuple[str, str] | None = None,
|
||||
) -> dict[str, str]:
|
||||
"""构建请求头"""
|
||||
pass
|
||||
|
||||
def build(
|
||||
self,
|
||||
original_body: Dict[str, Any],
|
||||
original_headers: Dict[str, str],
|
||||
original_body: dict[str, Any],
|
||||
original_headers: dict[str, str],
|
||||
endpoint: Any,
|
||||
key: Any,
|
||||
*,
|
||||
mapped_model: Optional[str] = None,
|
||||
mapped_model: str | None = None,
|
||||
is_stream: bool = False,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
pre_computed_auth: Optional[Tuple[str, str]] = None,
|
||||
) -> Tuple[Dict[str, Any], Dict[str, str]]:
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
pre_computed_auth: tuple[str, str] | None = None,
|
||||
) -> tuple[dict[str, Any], dict[str, str]]:
|
||||
"""
|
||||
构建完整的请求(请求体 + 请求头)
|
||||
|
||||
@@ -202,11 +201,11 @@ class PassthroughRequestBuilder(RequestBuilder):
|
||||
|
||||
def build_payload(
|
||||
self,
|
||||
original_body: Dict[str, Any],
|
||||
original_body: dict[str, Any],
|
||||
*,
|
||||
mapped_model: Optional[str] = None, # noqa: ARG002 - 由 apply_mapped_model 处理
|
||||
mapped_model: str | None = None, # noqa: ARG002 - 由 apply_mapped_model 处理
|
||||
is_stream: bool = False, # noqa: ARG002 - 保留原始值,不自动添加
|
||||
) -> Dict[str, Any]:
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
透传请求体 - 原样复制,不做任何修改
|
||||
|
||||
@@ -218,13 +217,13 @@ class PassthroughRequestBuilder(RequestBuilder):
|
||||
|
||||
def build_headers(
|
||||
self,
|
||||
original_headers: Dict[str, str],
|
||||
original_headers: dict[str, str],
|
||||
endpoint: Any,
|
||||
key: Any,
|
||||
*,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
pre_computed_auth: Optional[Tuple[str, str]] = None,
|
||||
) -> Dict[str, str]:
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
pre_computed_auth: tuple[str, str] | None = None,
|
||||
) -> dict[str, str]:
|
||||
"""
|
||||
透传请求头 - 清理敏感头部(黑名单),透传其他所有头部
|
||||
|
||||
@@ -289,11 +288,11 @@ class PassthroughRequestBuilder(RequestBuilder):
|
||||
|
||||
|
||||
def build_passthrough_request(
|
||||
original_body: Dict[str, Any],
|
||||
original_headers: Dict[str, str],
|
||||
original_body: dict[str, Any],
|
||||
original_headers: dict[str, str],
|
||||
endpoint: Any,
|
||||
key: Any,
|
||||
) -> Tuple[Dict[str, Any], Dict[str, str]]:
|
||||
) -> tuple[dict[str, Any], dict[str, str]]:
|
||||
"""
|
||||
构建透传模式的请求
|
||||
|
||||
|
||||
@@ -4,7 +4,7 @@
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import Any
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -13,14 +13,14 @@ class ParsedChunk:
|
||||
|
||||
# 原始数据
|
||||
raw_line: str
|
||||
event_type: Optional[str] = None
|
||||
data: Optional[Dict[str, Any]] = None
|
||||
event_type: str | None = None
|
||||
data: dict[str, Any] | None = None
|
||||
|
||||
# 提取的内容
|
||||
text_delta: str = ""
|
||||
is_done: bool = False
|
||||
is_error: bool = False
|
||||
error_message: Optional[str] = None
|
||||
error_message: str | None = None
|
||||
|
||||
# 使用量信息(通常在最后一个 chunk 中)
|
||||
input_tokens: int = 0
|
||||
@@ -29,7 +29,7 @@ class ParsedChunk:
|
||||
cache_read_tokens: int = 0
|
||||
|
||||
# 响应 ID
|
||||
response_id: Optional[str] = None
|
||||
response_id: str | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -48,21 +48,21 @@ class StreamStats:
|
||||
|
||||
# 内容
|
||||
collected_text: str = ""
|
||||
response_id: Optional[str] = None
|
||||
response_id: str | None = None
|
||||
|
||||
# 状态
|
||||
has_completion: bool = False
|
||||
status_code: int = 200
|
||||
error_message: Optional[str] = None
|
||||
error_message: str | None = None
|
||||
|
||||
# Provider 信息
|
||||
provider_name: Optional[str] = None
|
||||
endpoint_id: Optional[str] = None
|
||||
key_id: Optional[str] = None
|
||||
provider_name: str | None = None
|
||||
endpoint_id: str | None = None
|
||||
key_id: str | None = None
|
||||
|
||||
# 响应头和完整响应
|
||||
response_headers: Dict[str, str] = field(default_factory=dict)
|
||||
final_response: Optional[Dict[str, Any]] = None
|
||||
response_headers: dict[str, str] = field(default_factory=dict)
|
||||
final_response: dict[str, Any] | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -70,12 +70,12 @@ class ParsedResponse:
|
||||
"""解析后的非流式响应"""
|
||||
|
||||
# 原始响应
|
||||
raw_response: Dict[str, Any]
|
||||
raw_response: dict[str, Any]
|
||||
status_code: int
|
||||
|
||||
# 提取的内容
|
||||
text_content: str = ""
|
||||
response_id: Optional[str] = None
|
||||
response_id: str | None = None
|
||||
|
||||
# 使用量
|
||||
input_tokens: int = 0
|
||||
@@ -85,10 +85,10 @@ class ParsedResponse:
|
||||
|
||||
# 错误信息
|
||||
is_error: bool = False
|
||||
error_type: Optional[str] = None
|
||||
error_message: Optional[str] = None
|
||||
error_type: str | None = None
|
||||
error_message: str | None = None
|
||||
# 从响应体解析出的嵌套状态码(当 HTTP 200 但响应体含错误时使用)
|
||||
embedded_status_code: Optional[int] = None
|
||||
embedded_status_code: int | None = None
|
||||
|
||||
|
||||
class ResponseParser(ABC):
|
||||
@@ -106,7 +106,7 @@ class ResponseParser(ABC):
|
||||
api_format: str = "UNKNOWN"
|
||||
|
||||
@abstractmethod
|
||||
def parse_sse_line(self, line: str, stats: StreamStats) -> Optional[ParsedChunk]:
|
||||
def parse_sse_line(self, line: str, stats: StreamStats) -> ParsedChunk | None:
|
||||
"""
|
||||
解析单行 SSE 数据
|
||||
|
||||
@@ -120,7 +120,7 @@ class ResponseParser(ABC):
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def parse_response(self, response: Dict[str, Any], status_code: int) -> ParsedResponse:
|
||||
def parse_response(self, response: dict[str, Any], status_code: int) -> ParsedResponse:
|
||||
"""
|
||||
解析非流式响应
|
||||
|
||||
@@ -134,7 +134,7 @@ class ResponseParser(ABC):
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def extract_usage_from_response(self, response: Dict[str, Any]) -> Dict[str, int]:
|
||||
def extract_usage_from_response(self, response: dict[str, Any]) -> dict[str, int]:
|
||||
"""
|
||||
从响应中提取 token 使用量
|
||||
|
||||
@@ -147,7 +147,7 @@ class ResponseParser(ABC):
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def extract_text_content(self, response: Dict[str, Any]) -> str:
|
||||
def extract_text_content(self, response: dict[str, Any]) -> str:
|
||||
"""
|
||||
从响应中提取文本内容
|
||||
|
||||
@@ -159,7 +159,7 @@ class ResponseParser(ABC):
|
||||
"""
|
||||
pass
|
||||
|
||||
def is_error_response(self, response: Dict[str, Any]) -> bool:
|
||||
def is_error_response(self, response: dict[str, Any]) -> bool:
|
||||
"""
|
||||
判断响应是否为错误响应
|
||||
|
||||
|
||||
@@ -8,9 +8,11 @@
|
||||
- 请求/响应数据
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.core.api_format.conversion.stream_state import StreamState
|
||||
@@ -35,16 +37,16 @@ class StreamContext:
|
||||
api_key_id: int = 0
|
||||
|
||||
# Provider 信息(在请求执行时填充)
|
||||
provider_name: Optional[str] = None
|
||||
provider_id: Optional[str] = None
|
||||
endpoint_id: Optional[str] = None
|
||||
key_id: Optional[str] = None
|
||||
attempt_id: Optional[str] = None
|
||||
provider_name: str | None = None
|
||||
provider_id: str | None = None
|
||||
endpoint_id: str | None = None
|
||||
key_id: str | None = None
|
||||
attempt_id: str | None = None
|
||||
attempt_synced: bool = False
|
||||
provider_api_format: Optional[str] = None # Provider 的响应格式
|
||||
provider_api_format: str | None = None # Provider 的响应格式
|
||||
|
||||
# 模型映射
|
||||
mapped_model: Optional[str] = None
|
||||
mapped_model: str | None = None
|
||||
|
||||
# Token 统计
|
||||
input_tokens: int = 0
|
||||
@@ -53,33 +55,33 @@ class StreamContext:
|
||||
cache_creation_tokens: int = 0
|
||||
|
||||
# 响应内容
|
||||
_collected_text_parts: List[str] = field(default_factory=list, repr=False)
|
||||
response_id: Optional[str] = None
|
||||
final_usage: Optional[Dict[str, Any]] = None
|
||||
final_response: Optional[Dict[str, Any]] = None
|
||||
_collected_text_parts: list[str] = field(default_factory=list, repr=False)
|
||||
response_id: str | None = None
|
||||
final_usage: dict[str, Any] | None = None
|
||||
final_response: dict[str, Any] | None = None
|
||||
|
||||
# 时间指标
|
||||
first_byte_time_ms: Optional[int] = None # 首字时间 (TTFB - Time To First Byte)
|
||||
first_byte_time_ms: int | None = None # 首字时间 (TTFB - Time To First Byte)
|
||||
start_time: float = field(default_factory=time.time)
|
||||
|
||||
# 响应状态
|
||||
status_code: int = 200
|
||||
error_message: Optional[str] = None # 客户端友好的错误消息
|
||||
upstream_response: Optional[str] = None # 原始 Provider 响应(用于请求链路追踪)
|
||||
error_message: str | None = None # 客户端友好的错误消息
|
||||
upstream_response: str | None = None # 原始 Provider 响应(用于请求链路追踪)
|
||||
has_completion: bool = False
|
||||
|
||||
# 请求/响应数据
|
||||
response_headers: Dict[str, str] = field(default_factory=dict) # 提供商响应头
|
||||
client_response_headers: Dict[str, str] = field(default_factory=dict) # 返回给客户端的响应头
|
||||
provider_request_headers: Dict[str, str] = field(default_factory=dict)
|
||||
provider_request_body: Optional[Dict[str, Any]] = None
|
||||
response_headers: dict[str, str] = field(default_factory=dict) # 提供商响应头
|
||||
client_response_headers: dict[str, str] = field(default_factory=dict) # 返回给客户端的响应头
|
||||
provider_request_headers: dict[str, str] = field(default_factory=dict)
|
||||
provider_request_body: dict[str, Any] | None = None
|
||||
|
||||
# 格式转换信息(CLI handler 需要)
|
||||
client_api_format: str = ""
|
||||
needs_conversion: bool = False # 是否需要跨格式转换(由 handler 层设置)
|
||||
|
||||
# Provider 响应元数据(CLI handler 需要)
|
||||
response_metadata: Dict[str, Any] = field(default_factory=dict)
|
||||
response_metadata: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
# 整流标记(Thinking Rectifier)
|
||||
rectified: bool = False # 请求是否经过整流(移除 thinking 块后重试)
|
||||
@@ -87,10 +89,10 @@ class StreamContext:
|
||||
# 流式处理统计
|
||||
data_count: int = 0
|
||||
chunk_count: int = 0
|
||||
parsed_chunks: List[Dict[str, Any]] = field(default_factory=list)
|
||||
parsed_chunks: list[dict[str, Any]] = field(default_factory=list)
|
||||
|
||||
# 流式格式转换状态(跨 chunk 追踪)
|
||||
stream_conversion_state: Optional["StreamState"] = None
|
||||
stream_conversion_state: StreamState | None = None
|
||||
|
||||
def reset_for_retry(self) -> None:
|
||||
"""
|
||||
@@ -138,7 +140,7 @@ class StreamContext:
|
||||
provider_id: str,
|
||||
endpoint_id: str,
|
||||
key_id: str,
|
||||
provider_api_format: Optional[str] = None,
|
||||
provider_api_format: str | None = None,
|
||||
) -> None:
|
||||
"""更新 Provider 信息"""
|
||||
self.provider_name = provider_name
|
||||
@@ -149,10 +151,10 @@ class StreamContext:
|
||||
|
||||
def update_usage(
|
||||
self,
|
||||
input_tokens: Optional[int] = None,
|
||||
output_tokens: Optional[int] = None,
|
||||
cached_tokens: Optional[int] = None,
|
||||
cache_creation_tokens: Optional[int] = None,
|
||||
input_tokens: int | None = None,
|
||||
output_tokens: int | None = None,
|
||||
cached_tokens: int | None = None,
|
||||
cache_creation_tokens: int | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
更新 Token 使用统计
|
||||
@@ -194,7 +196,7 @@ class StreamContext:
|
||||
self,
|
||||
status_code: int,
|
||||
error_message: str,
|
||||
upstream_response: Optional[str] = None,
|
||||
upstream_response: str | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
标记请求失败
|
||||
@@ -230,7 +232,7 @@ class StreamContext:
|
||||
"""检查是否因客户端断开连接而结束"""
|
||||
return self.status_code == 499
|
||||
|
||||
def build_response_body(self, response_time_ms: int) -> Dict[str, Any]:
|
||||
def build_response_body(self, response_time_ms: int) -> dict[str, Any]:
|
||||
"""
|
||||
构建响应体元数据
|
||||
|
||||
|
||||
@@ -13,7 +13,10 @@ import asyncio
|
||||
import codecs
|
||||
import json
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, AsyncGenerator, Callable, Optional
|
||||
from typing import Any
|
||||
|
||||
from collections.abc import Callable
|
||||
from collections.abc import AsyncGenerator
|
||||
|
||||
import httpx
|
||||
|
||||
@@ -65,10 +68,10 @@ class StreamProcessor:
|
||||
self,
|
||||
request_id: str,
|
||||
default_parser: ResponseParser,
|
||||
on_streaming_start: Optional[Callable[[], None]] = None,
|
||||
on_streaming_start: Callable[[], None] | None = None,
|
||||
*,
|
||||
collect_text: bool = False,
|
||||
smoothing_config: Optional[StreamSmoothingConfig] = None,
|
||||
smoothing_config: StreamSmoothingConfig | None = None,
|
||||
):
|
||||
"""
|
||||
初始化流处理器
|
||||
@@ -105,7 +108,7 @@ class StreamProcessor:
|
||||
def handle_sse_event(
|
||||
self,
|
||||
ctx: StreamContext,
|
||||
event_name: Optional[str],
|
||||
event_name: str | None,
|
||||
data_str: str,
|
||||
*,
|
||||
skip_record: bool = False,
|
||||
@@ -363,7 +366,7 @@ class StreamProcessor:
|
||||
):
|
||||
# 重新抛出可重试的 Provider 异常,触发故障转移
|
||||
raise
|
||||
except (OSError, IOError) as e:
|
||||
except OSError as e:
|
||||
# 网络 I/O 异常:记录警告,可能需要重试
|
||||
logger.warning(f" [{self.request_id}] 预读流时发生网络异常: {type(e).__name__}: {e}")
|
||||
except Exception as e:
|
||||
@@ -382,10 +385,10 @@ class StreamProcessor:
|
||||
byte_iterator: Any,
|
||||
response_ctx: Any,
|
||||
http_client: httpx.AsyncClient,
|
||||
prefetched_chunks: Optional[list] = None,
|
||||
prefetched_chunks: list | None = None,
|
||||
*,
|
||||
start_time: Optional[float] = None,
|
||||
) -> AsyncGenerator[bytes, None]:
|
||||
start_time: float | None = None,
|
||||
) -> AsyncGenerator[bytes]:
|
||||
"""
|
||||
创建响应流生成器
|
||||
|
||||
@@ -547,9 +550,7 @@ class StreamProcessor:
|
||||
logger.warning(f"[{self.request_id}] 流式格式转换失败: {conv_err}")
|
||||
payload = _build_stream_error_payload("响应格式转换失败,请稍后重试")
|
||||
error_bytes = (
|
||||
f"data: {json.dumps(payload, ensure_ascii=False)}\n\n".encode(
|
||||
"utf-8"
|
||||
)
|
||||
f"data: {json.dumps(payload, ensure_ascii=False)}\n\n".encode()
|
||||
)
|
||||
done_bytes = (
|
||||
b"data: [DONE]\n\n" if client_format.startswith("OPENAI") else b""
|
||||
@@ -570,7 +571,7 @@ class StreamProcessor:
|
||||
# 统一使用 SSE 格式输出(Gemini streamGenerateContent 也使用 SSE)
|
||||
# 参考: https://ai.google.dev/api/generate-content
|
||||
out.append(
|
||||
f"data: {json.dumps(evt, ensure_ascii=False)}\n\n".encode("utf-8")
|
||||
f"data: {json.dumps(evt, ensure_ascii=False)}\n\n".encode()
|
||||
)
|
||||
return out
|
||||
|
||||
@@ -769,9 +770,9 @@ class StreamProcessor:
|
||||
async def create_monitored_stream(
|
||||
self,
|
||||
ctx: StreamContext,
|
||||
stream_generator: AsyncGenerator[bytes, None],
|
||||
stream_generator: AsyncGenerator[bytes],
|
||||
is_disconnected: Callable[[], Any],
|
||||
) -> AsyncGenerator[bytes, None]:
|
||||
) -> AsyncGenerator[bytes]:
|
||||
"""
|
||||
创建带监控的流生成器
|
||||
|
||||
@@ -833,8 +834,8 @@ class StreamProcessor:
|
||||
|
||||
async def create_smoothed_stream(
|
||||
self,
|
||||
stream_generator: AsyncGenerator[bytes, None],
|
||||
) -> AsyncGenerator[bytes, None]:
|
||||
stream_generator: AsyncGenerator[bytes],
|
||||
) -> AsyncGenerator[bytes]:
|
||||
"""
|
||||
创建平滑输出的流生成器
|
||||
|
||||
@@ -933,7 +934,7 @@ class StreamProcessor:
|
||||
if buffer:
|
||||
yield buffer
|
||||
|
||||
def _get_extractor(self, format_name: str) -> Optional[ContentExtractor]:
|
||||
def _get_extractor(self, format_name: str) -> ContentExtractor | None:
|
||||
"""获取或创建格式对应的提取器(带缓存)"""
|
||||
if format_name not in self._extractors:
|
||||
extractor = get_extractor(format_name)
|
||||
@@ -943,7 +944,7 @@ class StreamProcessor:
|
||||
|
||||
def _detect_format_and_extract(
|
||||
self, data: dict
|
||||
) -> tuple[Optional[str], Optional[ContentExtractor]]:
|
||||
) -> tuple[str | None, ContentExtractor | None]:
|
||||
"""
|
||||
检测数据格式并提取内容
|
||||
|
||||
@@ -998,10 +999,10 @@ class StreamProcessor:
|
||||
|
||||
|
||||
async def create_smoothed_stream(
|
||||
stream_generator: AsyncGenerator[bytes, None],
|
||||
stream_generator: AsyncGenerator[bytes],
|
||||
chunk_size: int = 20,
|
||||
delay_ms: int = 8,
|
||||
) -> AsyncGenerator[bytes, None]:
|
||||
) -> AsyncGenerator[bytes]:
|
||||
"""
|
||||
独立的平滑流生成函数
|
||||
|
||||
@@ -1032,7 +1033,7 @@ class _LightweightSmoother:
|
||||
self.delay_ms = delay_ms
|
||||
self._extractors: dict[str, ContentExtractor] = {}
|
||||
|
||||
def _get_extractor(self, format_name: str) -> Optional[ContentExtractor]:
|
||||
def _get_extractor(self, format_name: str) -> ContentExtractor | None:
|
||||
if format_name not in self._extractors:
|
||||
extractor = get_extractor(format_name)
|
||||
if extractor:
|
||||
@@ -1041,7 +1042,7 @@ class _LightweightSmoother:
|
||||
|
||||
def _detect_format_and_extract(
|
||||
self, data: dict
|
||||
) -> tuple[Optional[str], Optional[ContentExtractor]]:
|
||||
) -> tuple[str | None, ContentExtractor | None]:
|
||||
for format_name in get_extractor_formats():
|
||||
extractor = self._get_extractor(format_name)
|
||||
if extractor:
|
||||
@@ -1060,8 +1061,8 @@ class _LightweightSmoother:
|
||||
return [content[i : i + self.chunk_size] for i in range(0, text_length, self.chunk_size)]
|
||||
|
||||
async def smooth(
|
||||
self, stream_generator: AsyncGenerator[bytes, None]
|
||||
) -> AsyncGenerator[bytes, None]:
|
||||
self, stream_generator: AsyncGenerator[bytes]
|
||||
) -> AsyncGenerator[bytes]:
|
||||
buffer = b""
|
||||
is_first_content = True
|
||||
|
||||
|
||||
@@ -9,7 +9,7 @@
|
||||
|
||||
import asyncio
|
||||
import time
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
@@ -58,8 +58,8 @@ class StreamTelemetryRecorder:
|
||||
async def record_stream_stats(
|
||||
self,
|
||||
ctx: StreamContext,
|
||||
original_headers: Dict[str, str],
|
||||
original_request_body: Dict[str, Any],
|
||||
original_headers: dict[str, str],
|
||||
original_request_body: dict[str, Any],
|
||||
start_time: float,
|
||||
) -> None:
|
||||
"""
|
||||
@@ -144,9 +144,9 @@ class StreamTelemetryRecorder:
|
||||
self,
|
||||
writer: TelemetryWriter,
|
||||
ctx: StreamContext,
|
||||
original_headers: Dict[str, str],
|
||||
actual_request_body: Dict[str, Any],
|
||||
response_body: Optional[Dict[str, Any]],
|
||||
original_headers: dict[str, str],
|
||||
actual_request_body: dict[str, Any],
|
||||
response_body: dict[str, Any] | None,
|
||||
response_time_ms: int,
|
||||
) -> None:
|
||||
"""记录成功的请求"""
|
||||
@@ -193,9 +193,9 @@ class StreamTelemetryRecorder:
|
||||
self,
|
||||
writer: TelemetryWriter,
|
||||
ctx: StreamContext,
|
||||
original_headers: Dict[str, str],
|
||||
actual_request_body: Dict[str, Any],
|
||||
response_body: Optional[Dict[str, Any]],
|
||||
original_headers: dict[str, str],
|
||||
actual_request_body: dict[str, Any],
|
||||
response_body: dict[str, Any] | None,
|
||||
response_time_ms: int,
|
||||
) -> None:
|
||||
"""记录失败的请求"""
|
||||
@@ -236,9 +236,9 @@ class StreamTelemetryRecorder:
|
||||
self,
|
||||
writer: TelemetryWriter,
|
||||
ctx: StreamContext,
|
||||
original_headers: Dict[str, str],
|
||||
actual_request_body: Dict[str, Any],
|
||||
response_body: Optional[Dict[str, Any]],
|
||||
original_headers: dict[str, str],
|
||||
actual_request_body: dict[str, Any],
|
||||
response_body: dict[str, Any] | None,
|
||||
response_time_ms: int,
|
||||
) -> None:
|
||||
"""记录客户端取消的请求"""
|
||||
@@ -285,7 +285,7 @@ class StreamTelemetryRecorder:
|
||||
|
||||
from src.services.request.candidate import RequestCandidateService
|
||||
|
||||
extra_data: Dict[str, Any] = {
|
||||
extra_data: dict[str, Any] = {
|
||||
"stream_completed": ctx.is_success(),
|
||||
"data_count": ctx.data_count,
|
||||
}
|
||||
@@ -358,7 +358,7 @@ class StreamTelemetryRecorder:
|
||||
status: str,
|
||||
response_time_ms: int,
|
||||
status_code: int = 200,
|
||||
error_message: Optional[str] = None,
|
||||
error_message: str | None = None,
|
||||
) -> None:
|
||||
"""直接更新 Usage 表的状态字段"""
|
||||
try:
|
||||
@@ -378,7 +378,7 @@ class StreamTelemetryRecorder:
|
||||
|
||||
async def _get_telemetry_writer(
|
||||
self, bg_db: Session, ctx: StreamContext, response_time_ms: int
|
||||
) -> Optional[TelemetryWriter]:
|
||||
) -> TelemetryWriter | None:
|
||||
if config.usage_queue_enabled and self.user_id and self.api_key_id:
|
||||
return QueueTelemetryWriter(
|
||||
request_id=self.request_id,
|
||||
@@ -400,9 +400,9 @@ class StreamTelemetryRecorder:
|
||||
self,
|
||||
writer: TelemetryWriter,
|
||||
ctx: StreamContext,
|
||||
original_headers: Dict[str, str],
|
||||
actual_request_body: Dict[str, Any],
|
||||
response_body: Optional[Dict[str, Any]],
|
||||
original_headers: dict[str, str],
|
||||
actual_request_body: dict[str, Any],
|
||||
response_body: dict[str, Any] | None,
|
||||
response_time_ms: int,
|
||||
) -> None:
|
||||
"""根据上下文状态分发到对应的记录方法"""
|
||||
@@ -430,7 +430,7 @@ class StreamTelemetryRecorder:
|
||||
return "cancelled"
|
||||
return "failed"
|
||||
|
||||
def _build_db_writer(self, bg_db: Session) -> Optional[DbTelemetryWriter]:
|
||||
def _build_db_writer(self, bg_db: Session) -> DbTelemetryWriter | None:
|
||||
user = bg_db.query(User).filter(User.id == self.user_id).first()
|
||||
api_key_obj = bg_db.query(ApiKey).filter(ApiKey.id == self.api_key_id).first()
|
||||
|
||||
|
||||
@@ -4,8 +4,9 @@ Handler 基础工具函数
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
import json
|
||||
from typing import TYPE_CHECKING, Any, Dict, Optional
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from src.core.exceptions import EmbeddedErrorException, ProviderNotAvailableException
|
||||
from src.core.api_format import filter_response_headers
|
||||
@@ -15,7 +16,7 @@ if TYPE_CHECKING:
|
||||
from src.core.api_format.conversion.registry import FormatConversionRegistry
|
||||
|
||||
|
||||
def get_format_converter_registry() -> "FormatConversionRegistry":
|
||||
def get_format_converter_registry() -> FormatConversionRegistry:
|
||||
"""
|
||||
获取格式转换注册表(线程安全)
|
||||
|
||||
@@ -31,7 +32,7 @@ def get_format_converter_registry() -> "FormatConversionRegistry":
|
||||
return format_conversion_registry
|
||||
|
||||
|
||||
def extract_cache_creation_tokens(usage: Dict[str, Any]) -> int:
|
||||
def extract_cache_creation_tokens(usage: dict[str, Any]) -> int:
|
||||
"""
|
||||
提取缓存创建 tokens(兼容三种格式)
|
||||
|
||||
@@ -99,7 +100,7 @@ def extract_cache_creation_tokens(usage: Dict[str, Any]) -> int:
|
||||
return old_format
|
||||
|
||||
|
||||
def build_sse_headers(extra_headers: Optional[Dict[str, str]] = None) -> Dict[str, str]:
|
||||
def build_sse_headers(extra_headers: dict[str, str] | None = None) -> dict[str, str]:
|
||||
"""
|
||||
构建 SSE(text/event-stream)推荐响应头,用于减少代理缓冲带来的卡顿/成段输出。
|
||||
|
||||
@@ -107,7 +108,7 @@ def build_sse_headers(extra_headers: Optional[Dict[str, str]] = None) -> Dict[st
|
||||
- Cache-Control: no-transform 可避免部分代理对流做压缩/改写导致缓冲
|
||||
- X-Accel-Buffering: no 可显式提示 Nginx 关闭缓冲(即使全局已关闭也无害)
|
||||
"""
|
||||
headers: Dict[str, str] = {
|
||||
headers: dict[str, str] = {
|
||||
"Cache-Control": "no-cache, no-transform",
|
||||
"X-Accel-Buffering": "no",
|
||||
}
|
||||
@@ -116,7 +117,7 @@ def build_sse_headers(extra_headers: Optional[Dict[str, str]] = None) -> Dict[st
|
||||
return headers
|
||||
|
||||
|
||||
def filter_proxy_response_headers(headers: Optional[Dict[str, str]]) -> Dict[str, str]:
|
||||
def filter_proxy_response_headers(headers: dict[str, str] | None) -> dict[str, str]:
|
||||
"""
|
||||
过滤上游响应头中不应透传给客户端的字段。
|
||||
|
||||
@@ -148,8 +149,8 @@ def check_prefetched_response_error(
|
||||
parser: Any,
|
||||
request_id: str,
|
||||
provider_name: str,
|
||||
endpoint_id: Optional[str],
|
||||
base_url: Optional[str],
|
||||
endpoint_id: str | None,
|
||||
base_url: str | None,
|
||||
) -> None:
|
||||
"""
|
||||
检查预读的响应是否为非 SSE 格式的错误响应(HTML 或纯 JSON 错误)
|
||||
|
||||
@@ -4,7 +4,7 @@ Claude Chat Adapter - 基于 ChatAdapterBase 的 Claude Chat API 适配器
|
||||
处理 /v1/messages 端点的 Claude Chat 格式请求。
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, Optional, Tuple, Type
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from fastapi import HTTPException, Request
|
||||
@@ -25,9 +25,9 @@ class ClaudeCapabilityDetector:
|
||||
|
||||
@staticmethod
|
||||
def detect_from_headers(
|
||||
headers: Dict[str, str],
|
||||
request_body: Optional[Dict[str, Any]] = None,
|
||||
) -> Dict[str, bool]:
|
||||
headers: dict[str, str],
|
||||
request_body: dict[str, Any] | None = None,
|
||||
) -> dict[str, bool]:
|
||||
"""
|
||||
从 Claude 请求头检测能力需求
|
||||
|
||||
@@ -38,7 +38,7 @@ class ClaudeCapabilityDetector:
|
||||
headers: 请求头字典
|
||||
request_body: 请求体(Claude 不使用,保留用于接口统一)
|
||||
"""
|
||||
requirements: Dict[str, bool] = {}
|
||||
requirements: dict[str, bool] = {}
|
||||
|
||||
# 使用统一的大小写不敏感获取
|
||||
beta_header = get_header_value(headers, "anthropic-beta")
|
||||
@@ -61,21 +61,21 @@ class ClaudeChatAdapter(ChatAdapterBase):
|
||||
name = "claude.chat"
|
||||
|
||||
@property
|
||||
def HANDLER_CLASS(self) -> Type[ChatHandlerBase]:
|
||||
def HANDLER_CLASS(self) -> type[ChatHandlerBase]:
|
||||
"""延迟导入 Handler 类避免循环依赖"""
|
||||
from src.api.handlers.claude.handler import ClaudeChatHandler
|
||||
|
||||
return ClaudeChatHandler
|
||||
|
||||
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
|
||||
def __init__(self, allowed_api_formats: list[str] | None = None):
|
||||
super().__init__(allowed_api_formats or ["CLAUDE"])
|
||||
logger.info(f"[{self.name}] 初始化Chat模式适配器 | API格式: {self.allowed_api_formats}")
|
||||
|
||||
def detect_capability_requirements(
|
||||
self,
|
||||
headers: Dict[str, str],
|
||||
request_body: Optional[Dict[str, Any]] = None,
|
||||
) -> Dict[str, bool]:
|
||||
headers: dict[str, str],
|
||||
request_body: dict[str, Any] | None = None,
|
||||
) -> dict[str, bool]:
|
||||
"""检测 Claude 请求中隐含的能力需求"""
|
||||
return ClaudeCapabilityDetector.detect_from_headers(headers)
|
||||
|
||||
@@ -124,7 +124,7 @@ class ClaudeChatAdapter(ChatAdapterBase):
|
||||
)
|
||||
return request
|
||||
|
||||
def _build_audit_metadata(self, _payload: Dict[str, Any], request_obj) -> Dict[str, Any]:
|
||||
def _build_audit_metadata(self, _payload: dict[str, Any], request_obj) -> dict[str, Any]:
|
||||
"""构建 Claude Chat 特定的审计元数据"""
|
||||
role_counts: dict[str, int] = {}
|
||||
for message in request_obj.messages:
|
||||
@@ -153,8 +153,8 @@ class ClaudeChatAdapter(ChatAdapterBase):
|
||||
client: httpx.AsyncClient,
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
) -> Tuple[list, Optional[str]]:
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
) -> tuple[list, str | None]:
|
||||
"""查询 Claude API 支持的模型列表"""
|
||||
headers = cls.build_headers_with_extra(api_key, extra_headers)
|
||||
|
||||
@@ -201,7 +201,7 @@ class ClaudeChatAdapter(ChatAdapterBase):
|
||||
# build_request_body 使用基类实现,通过 format_conversion_registry 自动转换 OPENAI -> CLAUDE
|
||||
|
||||
|
||||
def build_claude_adapter(x_app_header: Optional[str]):
|
||||
def build_claude_adapter(x_app_header: str | None):
|
||||
"""根据 x-app 头部构造 Chat 或 Claude Code 适配器。"""
|
||||
if x_app_header and x_app_header.lower() == "cli":
|
||||
from src.api.handlers.claude_cli.adapter import ClaudeCliAdapter
|
||||
@@ -216,7 +216,7 @@ class ClaudeTokenCountAdapter(ApiAdapter):
|
||||
name = "claude.token_count"
|
||||
mode = ApiMode.STANDARD
|
||||
|
||||
def extract_api_key(self, request: Request) -> Optional[str]:
|
||||
def extract_api_key(self, request: Request) -> str | None:
|
||||
"""从请求中提取 API 密钥 (x-api-key 或 Authorization: Bearer)"""
|
||||
# 优先检查 x-api-key
|
||||
api_key = request.headers.get("x-api-key")
|
||||
|
||||
@@ -5,7 +5,7 @@ Claude Chat Handler - 基于通用 Chat Handler 基类的简化实现
|
||||
代码量从原来的 ~1470 行减少到 ~120 行。
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import Any
|
||||
|
||||
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
|
||||
from src.api.handlers.base.utils import extract_cache_creation_tokens
|
||||
@@ -25,8 +25,8 @@ class ClaudeChatHandler(ChatHandlerBase):
|
||||
|
||||
def extract_model_from_request(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
path_params: Optional[Dict[str, Any]] = None, # noqa: ARG002
|
||||
request_body: dict[str, Any],
|
||||
path_params: dict[str, Any] | None = None, # noqa: ARG002
|
||||
) -> str:
|
||||
"""
|
||||
从请求中提取模型名 - Claude 格式实现
|
||||
@@ -45,9 +45,9 @@ class ClaudeChatHandler(ChatHandlerBase):
|
||||
|
||||
def apply_mapped_model(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
request_body: dict[str, Any],
|
||||
mapped_model: str,
|
||||
) -> Dict[str, Any]:
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
将映射后的模型名应用到请求体
|
||||
|
||||
@@ -90,7 +90,7 @@ class ClaudeChatHandler(ChatHandlerBase):
|
||||
|
||||
return request
|
||||
|
||||
def _extract_usage(self, response: Dict) -> Dict[str, int]:
|
||||
def _extract_usage(self, response: dict) -> dict[str, int]:
|
||||
"""
|
||||
从 Claude 响应中提取 token 使用情况
|
||||
|
||||
@@ -108,7 +108,7 @@ class ClaudeChatHandler(ChatHandlerBase):
|
||||
"cache_read_input_tokens": usage.get("cache_read_input_tokens", 0),
|
||||
}
|
||||
|
||||
def _normalize_response(self, response: Dict[str, Any]) -> Dict[str, Any]:
|
||||
def _normalize_response(self, response: dict[str, Any]) -> dict[str, Any]:
|
||||
"""
|
||||
规范化 Claude 响应
|
||||
|
||||
|
||||
@@ -4,10 +4,9 @@ Claude SSE 流解析器
|
||||
解析 Claude Messages API 的 Server-Sent Events 流。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any, Dict, List, Optional
|
||||
from typing import Any
|
||||
|
||||
from src.api.handlers.base.utils import extract_cache_creation_tokens
|
||||
|
||||
@@ -43,7 +42,7 @@ class ClaudeStreamParser:
|
||||
DELTA_TEXT = "text_delta"
|
||||
DELTA_INPUT_JSON = "input_json_delta"
|
||||
|
||||
def parse_chunk(self, chunk: bytes | str) -> List[Dict[str, Any]]:
|
||||
def parse_chunk(self, chunk: bytes | str) -> list[dict[str, Any]]:
|
||||
"""
|
||||
解析 SSE 数据块
|
||||
|
||||
@@ -58,10 +57,10 @@ class ClaudeStreamParser:
|
||||
else:
|
||||
text = chunk
|
||||
|
||||
events: List[Dict[str, Any]] = []
|
||||
events: list[dict[str, Any]] = []
|
||||
lines = text.strip().split("\n")
|
||||
|
||||
current_event_type: Optional[str] = None
|
||||
current_event_type: str | None = None
|
||||
|
||||
for line in lines:
|
||||
line = line.strip()
|
||||
@@ -96,7 +95,7 @@ class ClaudeStreamParser:
|
||||
|
||||
return events
|
||||
|
||||
def parse_line(self, line: str) -> Optional[Dict[str, Any]]:
|
||||
def parse_line(self, line: str) -> dict[str, Any] | None:
|
||||
"""
|
||||
解析单行 SSE 数据
|
||||
|
||||
@@ -117,7 +116,7 @@ class ClaudeStreamParser:
|
||||
except json.JSONDecodeError:
|
||||
return None
|
||||
|
||||
def is_done_event(self, event: Dict[str, Any]) -> bool:
|
||||
def is_done_event(self, event: dict[str, Any]) -> bool:
|
||||
"""
|
||||
判断是否为结束事件
|
||||
|
||||
@@ -130,7 +129,7 @@ class ClaudeStreamParser:
|
||||
event_type = event.get("type")
|
||||
return event_type in (self.EVENT_MESSAGE_STOP, "__done__")
|
||||
|
||||
def is_error_event(self, event: Dict[str, Any]) -> bool:
|
||||
def is_error_event(self, event: dict[str, Any]) -> bool:
|
||||
"""
|
||||
判断是否为错误事件
|
||||
|
||||
@@ -142,7 +141,7 @@ class ClaudeStreamParser:
|
||||
"""
|
||||
return event.get("type") == self.EVENT_ERROR
|
||||
|
||||
def get_event_type(self, event: Dict[str, Any]) -> Optional[str]:
|
||||
def get_event_type(self, event: dict[str, Any]) -> str | None:
|
||||
"""
|
||||
获取事件类型
|
||||
|
||||
@@ -155,7 +154,7 @@ class ClaudeStreamParser:
|
||||
event_type = event.get("type")
|
||||
return str(event_type) if event_type is not None else None
|
||||
|
||||
def extract_text_delta(self, event: Dict[str, Any]) -> Optional[str]:
|
||||
def extract_text_delta(self, event: dict[str, Any]) -> str | None:
|
||||
"""
|
||||
从 content_block_delta 事件中提取文本增量
|
||||
|
||||
@@ -175,7 +174,7 @@ class ClaudeStreamParser:
|
||||
|
||||
return None
|
||||
|
||||
def extract_usage(self, event: Dict[str, Any]) -> Optional[Dict[str, int]]:
|
||||
def extract_usage(self, event: dict[str, Any]) -> dict[str, int] | None:
|
||||
"""
|
||||
从事件中提取 token 使用量
|
||||
|
||||
@@ -212,7 +211,7 @@ class ClaudeStreamParser:
|
||||
|
||||
return None
|
||||
|
||||
def extract_message_id(self, event: Dict[str, Any]) -> Optional[str]:
|
||||
def extract_message_id(self, event: dict[str, Any]) -> str | None:
|
||||
"""
|
||||
从 message_start 事件中提取消息 ID
|
||||
|
||||
@@ -229,7 +228,7 @@ class ClaudeStreamParser:
|
||||
msg_id = message.get("id")
|
||||
return str(msg_id) if msg_id is not None else None
|
||||
|
||||
def extract_stop_reason(self, event: Dict[str, Any]) -> Optional[str]:
|
||||
def extract_stop_reason(self, event: dict[str, Any]) -> str | None:
|
||||
"""
|
||||
从 message_delta 事件中提取停止原因
|
||||
|
||||
|
||||
@@ -4,7 +4,7 @@ Claude CLI Adapter - 基于通用 CLI Adapter 基类的简化实现
|
||||
继承 CliAdapterBase,只需配置 FORMAT_ID 和 HANDLER_CLASS。
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, Optional, Tuple, Type
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
@@ -27,20 +27,20 @@ class ClaudeCliAdapter(CliAdapterBase):
|
||||
name = "claude.cli"
|
||||
|
||||
@property
|
||||
def HANDLER_CLASS(self) -> Type[CliMessageHandlerBase]:
|
||||
def HANDLER_CLASS(self) -> type[CliMessageHandlerBase]:
|
||||
"""延迟导入 Handler 类避免循环依赖"""
|
||||
from src.api.handlers.claude_cli.handler import ClaudeCliMessageHandler
|
||||
|
||||
return ClaudeCliMessageHandler
|
||||
|
||||
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
|
||||
def __init__(self, allowed_api_formats: list[str] | None = None):
|
||||
super().__init__(allowed_api_formats or ["CLAUDE_CLI"])
|
||||
|
||||
def detect_capability_requirements(
|
||||
self,
|
||||
headers: Dict[str, str],
|
||||
request_body: Optional[Dict[str, Any]] = None,
|
||||
) -> Dict[str, bool]:
|
||||
headers: dict[str, str],
|
||||
request_body: dict[str, Any] | None = None,
|
||||
) -> dict[str, bool]:
|
||||
"""检测 Claude CLI 请求中隐含的能力需求"""
|
||||
return ClaudeCapabilityDetector.detect_from_headers(headers)
|
||||
|
||||
@@ -61,16 +61,16 @@ class ClaudeCliAdapter(CliAdapterBase):
|
||||
"""
|
||||
return input_tokens + cache_creation_input_tokens + cache_read_input_tokens
|
||||
|
||||
def _extract_message_count(self, payload: Dict[str, Any]) -> int:
|
||||
def _extract_message_count(self, payload: dict[str, Any]) -> int:
|
||||
"""Claude CLI 使用 messages 字段"""
|
||||
messages = payload.get("messages", [])
|
||||
return len(messages) if isinstance(messages, list) else 0
|
||||
|
||||
def _build_audit_metadata(
|
||||
self,
|
||||
payload: Dict[str, Any],
|
||||
path_params: Optional[Dict[str, Any]] = None, # noqa: ARG002
|
||||
) -> Dict[str, Any]:
|
||||
payload: dict[str, Any],
|
||||
path_params: dict[str, Any] | None = None, # noqa: ARG002
|
||||
) -> dict[str, Any]:
|
||||
"""Claude CLI 特定的审计元数据"""
|
||||
model = payload.get("model", "unknown")
|
||||
stream = payload.get("stream", False)
|
||||
@@ -104,8 +104,8 @@ class ClaudeCliAdapter(CliAdapterBase):
|
||||
client: httpx.AsyncClient,
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
) -> Tuple[list, Optional[str]]:
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
) -> tuple[list, str | None]:
|
||||
"""查询 Claude API 支持的模型列表(带 CLI User-Agent)"""
|
||||
# 复用 ClaudeChatAdapter 的实现,添加 CLI User-Agent
|
||||
cli_headers = {"User-Agent": config.internal_user_agent_claude_cli}
|
||||
@@ -120,7 +120,7 @@ class ClaudeCliAdapter(CliAdapterBase):
|
||||
return models, error
|
||||
|
||||
@classmethod
|
||||
def build_endpoint_url(cls, base_url: str, request_data: Dict[str, Any], model_name: Optional[str] = None) -> str:
|
||||
def build_endpoint_url(cls, base_url: str, request_data: dict[str, Any], model_name: str | None = None) -> str:
|
||||
"""构建Claude CLI API端点URL"""
|
||||
base_url = base_url.rstrip("/")
|
||||
if base_url.endswith("/v1"):
|
||||
@@ -131,12 +131,12 @@ class ClaudeCliAdapter(CliAdapterBase):
|
||||
# build_request_body 使用基类实现,通过 format_conversion_registry 自动转换 OPENAI -> CLAUDE_CLI
|
||||
|
||||
@classmethod
|
||||
def get_cli_user_agent(cls) -> Optional[str]:
|
||||
def get_cli_user_agent(cls) -> str | None:
|
||||
"""获取Claude CLI User-Agent"""
|
||||
return config.internal_user_agent_claude_cli
|
||||
|
||||
@classmethod
|
||||
def get_cli_extra_headers(cls) -> Dict[str, str]:
|
||||
def get_cli_extra_headers(cls) -> dict[str, str]:
|
||||
"""获取Claude CLI额外请求头,包含 x-app: cli 标识"""
|
||||
headers = super().get_cli_extra_headers()
|
||||
headers["x-app"] = "cli" # 标识 CLI 模式,让上游使用正确的认证方式
|
||||
|
||||
@@ -4,7 +4,7 @@ Claude CLI Message Handler - 基于通用 CLI Handler 基类的简化实现
|
||||
继承 CliMessageHandlerBase,只需覆盖格式特定的配置和事件处理逻辑。
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import Any
|
||||
|
||||
from src.api.handlers.base.cli_handler_base import (
|
||||
CliMessageHandlerBase,
|
||||
@@ -33,8 +33,8 @@ class ClaudeCliMessageHandler(CliMessageHandlerBase):
|
||||
|
||||
def extract_model_from_request(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
path_params: Optional[Dict[str, Any]] = None, # noqa: ARG002
|
||||
request_body: dict[str, Any],
|
||||
path_params: dict[str, Any] | None = None, # noqa: ARG002
|
||||
) -> str:
|
||||
"""
|
||||
从请求中提取模型名 - Claude 格式实现
|
||||
@@ -53,9 +53,9 @@ class ClaudeCliMessageHandler(CliMessageHandlerBase):
|
||||
|
||||
def apply_mapped_model(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
request_body: dict[str, Any],
|
||||
mapped_model: str,
|
||||
) -> Dict[str, Any]:
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Claude API 的 model 在请求体顶级
|
||||
|
||||
@@ -74,7 +74,7 @@ class ClaudeCliMessageHandler(CliMessageHandlerBase):
|
||||
self,
|
||||
ctx: StreamContext,
|
||||
event_type: str,
|
||||
data: Dict[str, Any],
|
||||
data: dict[str, Any],
|
||||
) -> None:
|
||||
"""
|
||||
处理 Claude CLI 格式的 SSE 事件
|
||||
@@ -142,8 +142,8 @@ class ClaudeCliMessageHandler(CliMessageHandlerBase):
|
||||
|
||||
def _extract_response_metadata(
|
||||
self,
|
||||
response: Dict[str, Any],
|
||||
) -> Dict[str, Any]:
|
||||
response: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
从 Claude 响应中提取元数据
|
||||
|
||||
@@ -155,7 +155,7 @@ class ClaudeCliMessageHandler(CliMessageHandlerBase):
|
||||
Returns:
|
||||
提取的元数据字典
|
||||
"""
|
||||
metadata: Dict[str, Any] = {}
|
||||
metadata: dict[str, Any] = {}
|
||||
|
||||
# 提取模型名称(实际使用的模型)
|
||||
if "model" in response:
|
||||
|
||||
@@ -4,7 +4,7 @@ Gemini Chat Adapter
|
||||
处理 Gemini API 格式的请求适配
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, Optional, Tuple, Type
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from fastapi import HTTPException, Request
|
||||
@@ -33,17 +33,17 @@ class GeminiChatAdapter(ChatAdapterBase):
|
||||
name = "gemini.chat"
|
||||
|
||||
@property
|
||||
def HANDLER_CLASS(self) -> Type[ChatHandlerBase]:
|
||||
def HANDLER_CLASS(self) -> type[ChatHandlerBase]:
|
||||
"""延迟导入 Handler 类避免循环依赖"""
|
||||
from src.api.handlers.gemini.handler import GeminiChatHandler
|
||||
|
||||
return GeminiChatHandler
|
||||
|
||||
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
|
||||
def __init__(self, allowed_api_formats: list[str] | None = None):
|
||||
super().__init__(allowed_api_formats or ["GEMINI"])
|
||||
logger.info(f"[{self.name}] 初始化 Gemini Chat 适配器 | API格式: {self.allowed_api_formats}")
|
||||
|
||||
def extract_api_key(self, request: Request) -> Optional[str]:
|
||||
def extract_api_key(self, request: Request) -> str | None:
|
||||
"""
|
||||
从请求中提取 API 密钥 - Gemini 支持 header 和 query 两种方式
|
||||
|
||||
@@ -68,8 +68,8 @@ class GeminiChatAdapter(ChatAdapterBase):
|
||||
return {}
|
||||
|
||||
def _merge_path_params(
|
||||
self, original_request_body: Dict[str, Any], path_params: Dict[str, Any] # noqa: ARG002
|
||||
) -> Dict[str, Any]:
|
||||
self, original_request_body: dict[str, Any], path_params: dict[str, Any] # noqa: ARG002
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
合并 URL 路径参数到请求体 - Gemini 特化版本
|
||||
|
||||
@@ -122,14 +122,14 @@ class GeminiChatAdapter(ChatAdapterBase):
|
||||
request.stream = is_stream
|
||||
return request
|
||||
|
||||
def _extract_message_count(self, payload: Dict[str, Any], request_obj) -> int:
|
||||
def _extract_message_count(self, payload: dict[str, Any], request_obj) -> int:
|
||||
"""提取消息数量"""
|
||||
contents = payload.get("contents", [])
|
||||
if hasattr(request_obj, "contents"):
|
||||
contents = request_obj.contents
|
||||
return len(contents) if isinstance(contents, list) else 0
|
||||
|
||||
def _build_audit_metadata(self, payload: Dict[str, Any], request_obj) -> Dict[str, Any]:
|
||||
def _build_audit_metadata(self, payload: dict[str, Any], request_obj) -> dict[str, Any]:
|
||||
"""构建 Gemini Chat 特定的审计元数据"""
|
||||
role_counts: dict[str, int] = {}
|
||||
|
||||
@@ -182,8 +182,8 @@ class GeminiChatAdapter(ChatAdapterBase):
|
||||
client: httpx.AsyncClient,
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
) -> Tuple[list, Optional[str]]:
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
) -> tuple[list, str | None]:
|
||||
"""查询 Gemini API 支持的模型列表"""
|
||||
# Gemini 使用 URL 参数传递 key,不需要 headers 中的认证
|
||||
base_url_clean = base_url.rstrip("/")
|
||||
@@ -192,7 +192,7 @@ class GeminiChatAdapter(ChatAdapterBase):
|
||||
else:
|
||||
models_url = f"{base_url_clean}/v1beta/models?key={api_key}"
|
||||
|
||||
headers: Dict[str, str] = {}
|
||||
headers: dict[str, str] = {}
|
||||
if extra_headers:
|
||||
headers.update(extra_headers)
|
||||
|
||||
@@ -242,16 +242,16 @@ class GeminiChatAdapter(ChatAdapterBase):
|
||||
client: httpx.AsyncClient,
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
request_data: Dict[str, Any],
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
request_data: dict[str, Any],
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
# 用量计算参数
|
||||
db: Optional[Any] = None,
|
||||
user: Optional[Any] = None,
|
||||
provider_name: Optional[str] = None,
|
||||
provider_id: Optional[str] = None,
|
||||
api_key_id: Optional[str] = None,
|
||||
model_name: Optional[str] = None,
|
||||
) -> Dict[str, Any]:
|
||||
db: Any | None = None,
|
||||
user: Any | None = None,
|
||||
provider_name: str | None = None,
|
||||
provider_id: str | None = None,
|
||||
api_key_id: str | None = None,
|
||||
model_name: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""测试 Gemini API 模型连接性(非流式)"""
|
||||
# Gemini需要从request_data或model_name参数获取model名称
|
||||
effective_model_name = model_name or request_data.get("model", "")
|
||||
|
||||
@@ -4,7 +4,7 @@ Gemini Chat Handler
|
||||
处理 Gemini API 格式的请求
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import Any
|
||||
|
||||
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
|
||||
|
||||
@@ -76,8 +76,8 @@ class GeminiChatHandler(ChatHandlerBase):
|
||||
|
||||
def extract_model_from_request(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
path_params: Optional[Dict[str, Any]] = None,
|
||||
request_body: dict[str, Any],
|
||||
path_params: dict[str, Any] | None = None,
|
||||
) -> str:
|
||||
"""
|
||||
从请求中提取模型名 - Gemini Chat 格式实现
|
||||
@@ -126,7 +126,7 @@ class GeminiChatHandler(ChatHandlerBase):
|
||||
|
||||
return request
|
||||
|
||||
def _extract_usage(self, response: Dict) -> Dict[str, int]:
|
||||
def _extract_usage(self, response: dict) -> dict[str, int]:
|
||||
"""
|
||||
从 Gemini 响应中提取 token 使用情况
|
||||
|
||||
@@ -151,7 +151,7 @@ class GeminiChatHandler(ChatHandlerBase):
|
||||
"cache_read_input_tokens": usage.get("cached_tokens", 0),
|
||||
}
|
||||
|
||||
def _normalize_response(self, response: Dict) -> Dict:
|
||||
def _normalize_response(self, response: dict) -> dict:
|
||||
"""
|
||||
规范化 Gemini 响应
|
||||
|
||||
|
||||
@@ -15,7 +15,7 @@ Gemini streamGenerateContent 常见两种返回(与上游/代理实现有关
|
||||
"""
|
||||
|
||||
import json
|
||||
from typing import Any, Dict, List, Optional, Union
|
||||
from typing import Any
|
||||
|
||||
|
||||
class GeminiStreamParser:
|
||||
@@ -43,7 +43,7 @@ class GeminiStreamParser:
|
||||
self._in_array = False
|
||||
self._brace_depth = 0
|
||||
|
||||
def parse_chunk(self, chunk: Union[bytes, str]) -> List[Dict[str, Any]]:
|
||||
def parse_chunk(self, chunk: bytes | str) -> list[dict[str, Any]]:
|
||||
"""
|
||||
解析流式数据块
|
||||
|
||||
@@ -58,7 +58,7 @@ class GeminiStreamParser:
|
||||
else:
|
||||
text = chunk
|
||||
|
||||
events: List[Dict[str, Any]] = []
|
||||
events: list[dict[str, Any]] = []
|
||||
|
||||
for char in text:
|
||||
if char == "[" and not self._in_array:
|
||||
@@ -97,7 +97,7 @@ class GeminiStreamParser:
|
||||
|
||||
return events
|
||||
|
||||
def parse_line(self, line: str) -> Optional[Dict[str, Any]]:
|
||||
def parse_line(self, line: str) -> dict[str, Any] | None:
|
||||
"""
|
||||
解析单行 JSON 数据
|
||||
|
||||
@@ -118,7 +118,7 @@ class GeminiStreamParser:
|
||||
except json.JSONDecodeError:
|
||||
return None
|
||||
|
||||
def is_done_event(self, event: Dict[str, Any]) -> bool:
|
||||
def is_done_event(self, event: dict[str, Any]) -> bool:
|
||||
"""
|
||||
判断是否为结束事件
|
||||
|
||||
@@ -143,7 +143,7 @@ class GeminiStreamParser:
|
||||
|
||||
return False
|
||||
|
||||
def is_error_event(self, event: Dict[str, Any]) -> bool:
|
||||
def is_error_event(self, event: dict[str, Any]) -> bool:
|
||||
"""
|
||||
判断是否为错误事件
|
||||
|
||||
@@ -171,7 +171,7 @@ class GeminiStreamParser:
|
||||
|
||||
return False
|
||||
|
||||
def extract_error_info(self, event: Dict[str, Any]) -> Optional[Dict[str, Any]]:
|
||||
def extract_error_info(self, event: dict[str, Any]) -> dict[str, Any] | None:
|
||||
"""
|
||||
从事件中提取错误信息
|
||||
|
||||
@@ -208,7 +208,7 @@ class GeminiStreamParser:
|
||||
|
||||
return None
|
||||
|
||||
def get_finish_reason(self, event: Dict[str, Any]) -> Optional[str]:
|
||||
def get_finish_reason(self, event: dict[str, Any]) -> str | None:
|
||||
"""
|
||||
获取结束原因
|
||||
|
||||
@@ -224,7 +224,7 @@ class GeminiStreamParser:
|
||||
return str(reason) if reason is not None else None
|
||||
return None
|
||||
|
||||
def extract_text_delta(self, event: Dict[str, Any]) -> Optional[str]:
|
||||
def extract_text_delta(self, event: dict[str, Any]) -> str | None:
|
||||
"""
|
||||
从响应中提取文本内容
|
||||
|
||||
@@ -248,7 +248,7 @@ class GeminiStreamParser:
|
||||
|
||||
return "".join(text_parts) if text_parts else None
|
||||
|
||||
def extract_usage(self, event: Dict[str, Any]) -> Optional[Dict[str, int]]:
|
||||
def extract_usage(self, event: dict[str, Any]) -> dict[str, int] | None:
|
||||
"""
|
||||
从事件中提取 token 使用量
|
||||
|
||||
@@ -280,7 +280,7 @@ class GeminiStreamParser:
|
||||
"cached_tokens": usage_metadata.get("cachedContentTokenCount", 0),
|
||||
}
|
||||
|
||||
def extract_model_version(self, event: Dict[str, Any]) -> Optional[str]:
|
||||
def extract_model_version(self, event: dict[str, Any]) -> str | None:
|
||||
"""
|
||||
从响应中提取模型版本
|
||||
|
||||
@@ -293,7 +293,7 @@ class GeminiStreamParser:
|
||||
version = event.get("modelVersion")
|
||||
return str(version) if version is not None else None
|
||||
|
||||
def extract_safety_ratings(self, event: Dict[str, Any]) -> Optional[List[Dict[str, Any]]]:
|
||||
def extract_safety_ratings(self, event: dict[str, Any]) -> list[dict[str, Any]] | None:
|
||||
"""
|
||||
从响应中提取安全评级
|
||||
|
||||
|
||||
@@ -4,7 +4,7 @@ Gemini CLI Adapter - 基于通用 CLI Adapter 基类的实现
|
||||
继承 CliAdapterBase,处理 Gemini CLI 格式的请求。
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, Optional, Tuple, Type
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from fastapi import Request
|
||||
@@ -29,16 +29,16 @@ class GeminiCliAdapter(CliAdapterBase):
|
||||
name = "gemini.cli"
|
||||
|
||||
@property
|
||||
def HANDLER_CLASS(self) -> Type[CliMessageHandlerBase]:
|
||||
def HANDLER_CLASS(self) -> type[CliMessageHandlerBase]:
|
||||
"""延迟导入 Handler 类避免循环依赖"""
|
||||
from src.api.handlers.gemini_cli.handler import GeminiCliMessageHandler
|
||||
|
||||
return GeminiCliMessageHandler
|
||||
|
||||
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
|
||||
def __init__(self, allowed_api_formats: list[str] | None = None):
|
||||
super().__init__(allowed_api_formats or ["GEMINI_CLI"])
|
||||
|
||||
def extract_api_key(self, request: Request) -> Optional[str]:
|
||||
def extract_api_key(self, request: Request) -> str | None:
|
||||
"""
|
||||
从请求中提取 API 密钥 - Gemini CLI 支持 header 和 query 两种方式
|
||||
|
||||
@@ -53,8 +53,8 @@ class GeminiCliAdapter(CliAdapterBase):
|
||||
)
|
||||
|
||||
def _merge_path_params(
|
||||
self, original_request_body: Dict[str, Any], path_params: Dict[str, Any] # noqa: ARG002
|
||||
) -> Dict[str, Any]:
|
||||
self, original_request_body: dict[str, Any], path_params: dict[str, Any] # noqa: ARG002
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
合并 URL 路径参数到请求体 - Gemini CLI 特化版本
|
||||
|
||||
@@ -74,23 +74,23 @@ class GeminiCliAdapter(CliAdapterBase):
|
||||
# Gemini: 不合并任何 path_params 到请求体
|
||||
return original_request_body.copy()
|
||||
|
||||
def _extract_message_count(self, payload: Dict[str, Any]) -> int:
|
||||
def _extract_message_count(self, payload: dict[str, Any]) -> int:
|
||||
"""Gemini CLI 使用 contents 字段"""
|
||||
contents = payload.get("contents", [])
|
||||
return len(contents) if isinstance(contents, list) else 0
|
||||
|
||||
def _build_audit_metadata(
|
||||
self,
|
||||
payload: Dict[str, Any],
|
||||
path_params: Optional[Dict[str, Any]] = None,
|
||||
) -> Dict[str, Any]:
|
||||
payload: dict[str, Any],
|
||||
path_params: dict[str, Any] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""Gemini CLI 特定的审计元数据"""
|
||||
# 从 path_params 获取 model(Gemini 请求体不含 model)
|
||||
model = path_params.get("model", "unknown") if path_params else "unknown"
|
||||
contents = payload.get("contents", [])
|
||||
generation_config = payload.get("generation_config", {}) or {}
|
||||
|
||||
role_counts: Dict[str, int] = {}
|
||||
role_counts: dict[str, int] = {}
|
||||
for content in contents:
|
||||
role = content.get("role", "unknown") if isinstance(content, dict) else "unknown"
|
||||
role_counts[role] = role_counts.get(role, 0) + 1
|
||||
@@ -120,8 +120,8 @@ class GeminiCliAdapter(CliAdapterBase):
|
||||
client: httpx.AsyncClient,
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
) -> Tuple[list, Optional[str]]:
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
) -> tuple[list, str | None]:
|
||||
"""查询 Gemini API 支持的模型列表(带 CLI User-Agent)"""
|
||||
# 复用 GeminiChatAdapter 的实现,添加 CLI User-Agent
|
||||
cli_headers = {"User-Agent": config.internal_user_agent_gemini_cli}
|
||||
@@ -136,7 +136,7 @@ class GeminiCliAdapter(CliAdapterBase):
|
||||
return models, error
|
||||
|
||||
@classmethod
|
||||
def build_endpoint_url(cls, base_url: str, request_data: Dict[str, Any], model_name: Optional[str] = None) -> str:
|
||||
def build_endpoint_url(cls, base_url: str, request_data: dict[str, Any], model_name: str | None = None) -> str:
|
||||
"""构建Gemini CLI API端点URL"""
|
||||
effective_model_name = model_name or request_data.get("model", "")
|
||||
if not effective_model_name:
|
||||
@@ -152,12 +152,12 @@ class GeminiCliAdapter(CliAdapterBase):
|
||||
# build_request_body 使用基类实现,通过 format_conversion_registry 自动转换 OPENAI -> GEMINI_CLI
|
||||
|
||||
@classmethod
|
||||
def get_cli_user_agent(cls) -> Optional[str]:
|
||||
def get_cli_user_agent(cls) -> str | None:
|
||||
"""获取Gemini CLI User-Agent"""
|
||||
return config.internal_user_agent_gemini_cli
|
||||
|
||||
@classmethod
|
||||
def get_cli_extra_headers(cls) -> Dict[str, str]:
|
||||
def get_cli_extra_headers(cls) -> dict[str, str]:
|
||||
"""获取Gemini CLI额外请求头,包含 x-app: cli 标识"""
|
||||
headers = super().get_cli_extra_headers()
|
||||
headers["x-app"] = "cli" # 标识 CLI 模式,让上游使用正确的 adapter
|
||||
|
||||
@@ -4,7 +4,7 @@ Gemini CLI Message Handler - 基于通用 CLI Handler 基类的实现
|
||||
继承 CliMessageHandlerBase,处理 Gemini CLI API 格式的请求。
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import Any
|
||||
|
||||
from src.api.handlers.base.cli_handler_base import (
|
||||
CliMessageHandlerBase,
|
||||
@@ -34,8 +34,8 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
|
||||
|
||||
def extract_model_from_request(
|
||||
self,
|
||||
request_body: Dict[str, Any], # noqa: ARG002 - 基类签名要求
|
||||
path_params: Optional[Dict[str, Any]] = None,
|
||||
request_body: dict[str, Any], # noqa: ARG002 - 基类签名要求
|
||||
path_params: dict[str, Any] | None = None,
|
||||
) -> str:
|
||||
"""
|
||||
从请求中提取模型名 - Gemini 格式实现
|
||||
@@ -57,8 +57,8 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
|
||||
|
||||
def prepare_provider_request_body(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
) -> Dict[str, Any]:
|
||||
request_body: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
准备发送给 Gemini API 的请求体 - 移除 model 字段
|
||||
|
||||
@@ -77,9 +77,9 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
|
||||
|
||||
def get_model_for_url(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
mapped_model: Optional[str],
|
||||
) -> Optional[str]:
|
||||
request_body: dict[str, Any],
|
||||
mapped_model: str | None,
|
||||
) -> str | None:
|
||||
"""
|
||||
Gemini 需要将 model 放入 URL 路径中
|
||||
|
||||
@@ -93,7 +93,7 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
|
||||
# 优先使用映射后的模型名,否则使用请求体中的
|
||||
return mapped_model or request_body.get("model")
|
||||
|
||||
def _extract_usage_from_event(self, event: Dict[str, Any]) -> Dict[str, int]:
|
||||
def _extract_usage_from_event(self, event: dict[str, Any]) -> dict[str, int]:
|
||||
"""
|
||||
从 Gemini 事件中提取 token 使用情况
|
||||
|
||||
@@ -126,7 +126,7 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
|
||||
self,
|
||||
ctx: StreamContext,
|
||||
_event_type: str,
|
||||
data: Dict[str, Any],
|
||||
data: dict[str, Any],
|
||||
) -> None:
|
||||
"""
|
||||
处理 Gemini CLI 格式的流式事件
|
||||
@@ -190,8 +190,8 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
|
||||
|
||||
def _extract_response_metadata(
|
||||
self,
|
||||
response: Dict[str, Any],
|
||||
) -> Dict[str, Any]:
|
||||
response: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
从 Gemini 响应中提取元数据
|
||||
|
||||
@@ -203,7 +203,7 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
|
||||
Returns:
|
||||
包含 model_version 的元数据字典
|
||||
"""
|
||||
metadata: Dict[str, Any] = {}
|
||||
metadata: dict[str, Any] = {}
|
||||
model_version = response.get("modelVersion")
|
||||
if model_version:
|
||||
metadata["model_version"] = model_version
|
||||
|
||||
@@ -4,7 +4,7 @@ OpenAI Chat Adapter - 基于 ChatAdapterBase 的 OpenAI Chat API 适配器
|
||||
处理 /v1/chat/completions 端点的 OpenAI Chat 格式请求。
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, Optional, Tuple, Type
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from fastapi.responses import JSONResponse
|
||||
@@ -28,13 +28,13 @@ class OpenAIChatAdapter(ChatAdapterBase):
|
||||
name = "openai.chat"
|
||||
|
||||
@property
|
||||
def HANDLER_CLASS(self) -> Type[ChatHandlerBase]:
|
||||
def HANDLER_CLASS(self) -> type[ChatHandlerBase]:
|
||||
"""延迟导入 Handler 类避免循环依赖"""
|
||||
from src.api.handlers.openai.handler import OpenAIChatHandler
|
||||
|
||||
return OpenAIChatHandler
|
||||
|
||||
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
|
||||
def __init__(self, allowed_api_formats: list[str] | None = None):
|
||||
super().__init__(allowed_api_formats or ["OPENAI"])
|
||||
|
||||
def _validate_request_body(self, original_request_body: dict, path_params: dict = None):
|
||||
@@ -66,7 +66,7 @@ class OpenAIChatAdapter(ChatAdapterBase):
|
||||
max_tokens=original_request_body.get("max_tokens"),
|
||||
)
|
||||
|
||||
def _build_audit_metadata(self, payload: Dict[str, Any], request_obj) -> Dict[str, Any]:
|
||||
def _build_audit_metadata(self, payload: dict[str, Any], request_obj) -> dict[str, Any]:
|
||||
"""构建 OpenAI Chat 特定的审计元数据"""
|
||||
role_counts = {}
|
||||
for message in request_obj.messages:
|
||||
@@ -105,8 +105,8 @@ class OpenAIChatAdapter(ChatAdapterBase):
|
||||
client: httpx.AsyncClient,
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
) -> Tuple[list, Optional[str]]:
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
) -> tuple[list, str | None]:
|
||||
"""查询 OpenAI 兼容 API 支持的模型列表"""
|
||||
headers = cls.build_headers_with_extra(api_key, extra_headers)
|
||||
|
||||
|
||||
@@ -5,7 +5,7 @@ OpenAI Chat Handler - 基于通用 Chat Handler 基类的简化实现
|
||||
代码量从原来的 ~1315 行减少到 ~100 行。
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import Any
|
||||
|
||||
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
|
||||
|
||||
@@ -24,8 +24,8 @@ class OpenAIChatHandler(ChatHandlerBase):
|
||||
|
||||
def extract_model_from_request(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
path_params: Optional[Dict[str, Any]] = None, # noqa: ARG002
|
||||
request_body: dict[str, Any],
|
||||
path_params: dict[str, Any] | None = None, # noqa: ARG002
|
||||
) -> str:
|
||||
"""
|
||||
从请求中提取模型名 - OpenAI 格式实现
|
||||
@@ -44,9 +44,9 @@ class OpenAIChatHandler(ChatHandlerBase):
|
||||
|
||||
def apply_mapped_model(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
request_body: dict[str, Any],
|
||||
mapped_model: str,
|
||||
) -> Dict[str, Any]:
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
将映射后的模型名应用到请求体
|
||||
|
||||
@@ -89,7 +89,7 @@ class OpenAIChatHandler(ChatHandlerBase):
|
||||
|
||||
return request
|
||||
|
||||
def _extract_usage(self, response: Dict) -> Dict[str, int]:
|
||||
def _extract_usage(self, response: dict) -> dict[str, int]:
|
||||
"""
|
||||
从 OpenAI 响应中提取 token 使用情况
|
||||
|
||||
@@ -106,7 +106,7 @@ class OpenAIChatHandler(ChatHandlerBase):
|
||||
"cache_read_input_tokens": 0,
|
||||
}
|
||||
|
||||
def _normalize_response(self, response: Dict) -> Dict:
|
||||
def _normalize_response(self, response: dict) -> dict:
|
||||
"""
|
||||
规范化 OpenAI 响应
|
||||
|
||||
|
||||
@@ -4,10 +4,9 @@ OpenAI SSE 流解析器
|
||||
解析 OpenAI Chat Completions API 的 Server-Sent Events 流。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any, Dict, List, Optional
|
||||
from typing import Any
|
||||
|
||||
|
||||
class OpenAIStreamParser:
|
||||
@@ -23,7 +22,7 @@ class OpenAIStreamParser:
|
||||
- 流结束时发送 data: [DONE]
|
||||
"""
|
||||
|
||||
def parse_chunk(self, chunk: bytes | str) -> List[Dict[str, Any]]:
|
||||
def parse_chunk(self, chunk: bytes | str) -> list[dict[str, Any]]:
|
||||
"""
|
||||
解析 SSE 数据块
|
||||
|
||||
@@ -38,7 +37,7 @@ class OpenAIStreamParser:
|
||||
else:
|
||||
text = chunk
|
||||
|
||||
chunks: List[Dict[str, Any]] = []
|
||||
chunks: list[dict[str, Any]] = []
|
||||
lines = text.strip().split("\n")
|
||||
|
||||
for line in lines:
|
||||
@@ -64,7 +63,7 @@ class OpenAIStreamParser:
|
||||
|
||||
return chunks
|
||||
|
||||
def parse_line(self, line: str) -> Optional[Dict[str, Any]]:
|
||||
def parse_line(self, line: str) -> dict[str, Any] | None:
|
||||
"""
|
||||
解析单行 SSE 数据
|
||||
|
||||
@@ -85,7 +84,7 @@ class OpenAIStreamParser:
|
||||
except json.JSONDecodeError:
|
||||
return None
|
||||
|
||||
def is_done_chunk(self, chunk: Dict[str, Any]) -> bool:
|
||||
def is_done_chunk(self, chunk: dict[str, Any]) -> bool:
|
||||
"""
|
||||
判断是否为结束 chunk
|
||||
|
||||
@@ -107,7 +106,7 @@ class OpenAIStreamParser:
|
||||
|
||||
return False
|
||||
|
||||
def get_finish_reason(self, chunk: Dict[str, Any]) -> Optional[str]:
|
||||
def get_finish_reason(self, chunk: dict[str, Any]) -> str | None:
|
||||
"""
|
||||
获取结束原因
|
||||
|
||||
@@ -123,7 +122,7 @@ class OpenAIStreamParser:
|
||||
return str(reason) if reason is not None else None
|
||||
return None
|
||||
|
||||
def extract_text_delta(self, chunk: Dict[str, Any]) -> Optional[str]:
|
||||
def extract_text_delta(self, chunk: dict[str, Any]) -> str | None:
|
||||
"""
|
||||
从 chunk 中提取文本增量
|
||||
|
||||
@@ -145,7 +144,7 @@ class OpenAIStreamParser:
|
||||
|
||||
return None
|
||||
|
||||
def extract_tool_calls_delta(self, chunk: Dict[str, Any]) -> Optional[List[Dict[str, Any]]]:
|
||||
def extract_tool_calls_delta(self, chunk: dict[str, Any]) -> list[dict[str, Any]] | None:
|
||||
"""
|
||||
从 chunk 中提取工具调用增量
|
||||
|
||||
@@ -165,7 +164,7 @@ class OpenAIStreamParser:
|
||||
return tool_calls
|
||||
return None
|
||||
|
||||
def extract_role(self, chunk: Dict[str, Any]) -> Optional[str]:
|
||||
def extract_role(self, chunk: dict[str, Any]) -> str | None:
|
||||
"""
|
||||
从 chunk 中提取角色
|
||||
|
||||
|
||||
@@ -4,7 +4,7 @@ OpenAI CLI Adapter - 基于通用 CLI Adapter 基类的简化实现
|
||||
继承 CliAdapterBase,只需配置 FORMAT_ID 和 HANDLER_CLASS。
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, Optional, Tuple, Type
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
@@ -27,13 +27,13 @@ class OpenAICliAdapter(CliAdapterBase):
|
||||
name = "openai.cli"
|
||||
|
||||
@property
|
||||
def HANDLER_CLASS(self) -> Type[CliMessageHandlerBase]:
|
||||
def HANDLER_CLASS(self) -> type[CliMessageHandlerBase]:
|
||||
"""延迟导入 Handler 类避免循环依赖"""
|
||||
from src.api.handlers.openai_cli.handler import OpenAICliMessageHandler
|
||||
|
||||
return OpenAICliMessageHandler
|
||||
|
||||
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
|
||||
def __init__(self, allowed_api_formats: list[str] | None = None):
|
||||
super().__init__(allowed_api_formats or ["OPENAI_CLI"])
|
||||
|
||||
# =========================================================================
|
||||
@@ -46,8 +46,8 @@ class OpenAICliAdapter(CliAdapterBase):
|
||||
client: httpx.AsyncClient,
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
) -> Tuple[list, Optional[str]]:
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
) -> tuple[list, str | None]:
|
||||
"""查询 OpenAI 兼容 API 支持的模型列表(带 CLI User-Agent)"""
|
||||
# 复用 OpenAIChatAdapter 的实现,添加 CLI User-Agent
|
||||
cli_headers = {"User-Agent": config.internal_user_agent_openai_cli}
|
||||
@@ -62,7 +62,7 @@ class OpenAICliAdapter(CliAdapterBase):
|
||||
return models, error
|
||||
|
||||
@classmethod
|
||||
def build_endpoint_url(cls, base_url: str, request_data: Dict[str, Any], model_name: Optional[str] = None) -> str:
|
||||
def build_endpoint_url(cls, base_url: str, request_data: dict[str, Any], model_name: str | None = None) -> str:
|
||||
"""构建OpenAI CLI API端点URL"""
|
||||
base_url = base_url.rstrip("/")
|
||||
if base_url.endswith("/v1"):
|
||||
@@ -74,7 +74,7 @@ class OpenAICliAdapter(CliAdapterBase):
|
||||
# OPENAI -> OPENAI_CLI 无转换器,会直接透传原始请求
|
||||
|
||||
@classmethod
|
||||
def get_cli_user_agent(cls) -> Optional[str]:
|
||||
def get_cli_user_agent(cls) -> str | None:
|
||||
"""获取OpenAI CLI User-Agent"""
|
||||
return config.internal_user_agent_openai_cli
|
||||
|
||||
|
||||
@@ -5,7 +5,7 @@ OpenAI CLI Message Handler - 基于通用 CLI Handler 基类的简化实现
|
||||
代码量从原来的 900+ 行减少到 ~100 行。
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import Any
|
||||
|
||||
from src.api.handlers.base.cli_handler_base import (
|
||||
CliMessageHandlerBase,
|
||||
@@ -32,8 +32,8 @@ class OpenAICliMessageHandler(CliMessageHandlerBase):
|
||||
|
||||
def extract_model_from_request(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
path_params: Optional[Dict[str, Any]] = None, # noqa: ARG002
|
||||
request_body: dict[str, Any],
|
||||
path_params: dict[str, Any] | None = None, # noqa: ARG002
|
||||
) -> str:
|
||||
"""
|
||||
从请求中提取模型名 - OpenAI 格式实现
|
||||
@@ -52,9 +52,9 @@ class OpenAICliMessageHandler(CliMessageHandlerBase):
|
||||
|
||||
def apply_mapped_model(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
request_body: dict[str, Any],
|
||||
mapped_model: str,
|
||||
) -> Dict[str, Any]:
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
OpenAI CLI (Responses API) 的 model 在请求体顶级字段。
|
||||
|
||||
@@ -73,7 +73,7 @@ class OpenAICliMessageHandler(CliMessageHandlerBase):
|
||||
self,
|
||||
ctx: StreamContext,
|
||||
event_type: str,
|
||||
data: Dict[str, Any],
|
||||
data: dict[str, Any],
|
||||
) -> None:
|
||||
"""
|
||||
处理 OpenAI CLI 格式的 SSE 事件
|
||||
@@ -144,8 +144,8 @@ class OpenAICliMessageHandler(CliMessageHandlerBase):
|
||||
|
||||
def _extract_response_metadata(
|
||||
self,
|
||||
response: Dict[str, Any],
|
||||
) -> Dict[str, Any]:
|
||||
response: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
从 OpenAI 响应中提取元数据
|
||||
|
||||
@@ -157,7 +157,7 @@ class OpenAICliMessageHandler(CliMessageHandlerBase):
|
||||
Returns:
|
||||
提取的元数据字典
|
||||
"""
|
||||
metadata: Dict[str, Any] = {}
|
||||
metadata: dict[str, Any] = {}
|
||||
|
||||
# 提取模型名称(实际使用的模型)
|
||||
if "model" in response:
|
||||
|
||||
@@ -2,7 +2,6 @@
|
||||
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -22,7 +21,7 @@ pipeline = ApiRequestPipeline()
|
||||
@router.get("/my-audit-logs")
|
||||
async def get_my_audit_logs(
|
||||
request: Request,
|
||||
event_type: Optional[str] = Query(None, description="事件类型筛选"),
|
||||
event_type: str | None = Query(None, description="事件类型筛选"),
|
||||
days: int = Query(30, description="查询天数"),
|
||||
limit: int = Query(50, description="返回数量限制"),
|
||||
offset: int = Query(0, ge=0, description="偏移量"),
|
||||
@@ -86,7 +85,7 @@ class AuthenticatedApiAdapter(ApiAdapter):
|
||||
|
||||
@dataclass
|
||||
class UserAuditLogsAdapter(AuthenticatedApiAdapter):
|
||||
event_type: Optional[str]
|
||||
event_type: str | None
|
||||
days: int
|
||||
limit: int
|
||||
offset: int
|
||||
|
||||
@@ -1,8 +1,7 @@
|
||||
"""OAuth 管理端点(管理员)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Dict, List, Optional
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, Request
|
||||
from pydantic import BaseModel, Field, ValidationError
|
||||
@@ -27,24 +26,24 @@ class SupportedOAuthType(BaseModel):
|
||||
default_authorization_url: str
|
||||
default_token_url: str
|
||||
default_userinfo_url: str
|
||||
default_scopes: List[str]
|
||||
default_scopes: list[str]
|
||||
|
||||
|
||||
class OAuthProviderUpsertRequest(BaseModel):
|
||||
display_name: str = Field(..., min_length=1, max_length=100)
|
||||
client_id: str = Field(..., min_length=1, max_length=255)
|
||||
client_secret: Optional[str] = Field(None, max_length=2048)
|
||||
client_secret: str | None = Field(None, max_length=2048)
|
||||
|
||||
authorization_url_override: Optional[str] = Field(None, max_length=500)
|
||||
token_url_override: Optional[str] = Field(None, max_length=500)
|
||||
userinfo_url_override: Optional[str] = Field(None, max_length=500)
|
||||
scopes: Optional[List[str]] = None
|
||||
authorization_url_override: str | None = Field(None, max_length=500)
|
||||
token_url_override: str | None = Field(None, max_length=500)
|
||||
userinfo_url_override: str | None = Field(None, max_length=500)
|
||||
scopes: list[str] | None = None
|
||||
|
||||
redirect_uri: str = Field(..., min_length=1, max_length=500)
|
||||
frontend_callback_url: str = Field(..., min_length=1, max_length=500)
|
||||
|
||||
attribute_mapping: Optional[Dict[str, Any]] = None
|
||||
extra_config: Optional[Dict[str, Any]] = None
|
||||
attribute_mapping: dict[str, Any] | None = None
|
||||
extra_config: dict[str, Any] | None = None
|
||||
|
||||
is_enabled: bool = False
|
||||
force: bool = False
|
||||
@@ -55,14 +54,14 @@ class OAuthProviderAdminResponse(BaseModel):
|
||||
display_name: str
|
||||
client_id: str
|
||||
has_secret: bool
|
||||
authorization_url_override: Optional[str] = None
|
||||
token_url_override: Optional[str] = None
|
||||
userinfo_url_override: Optional[str] = None
|
||||
scopes: Optional[List[str]] = None
|
||||
authorization_url_override: str | None = None
|
||||
token_url_override: str | None = None
|
||||
userinfo_url_override: str | None = None
|
||||
scopes: list[str] | None = None
|
||||
redirect_uri: str
|
||||
frontend_callback_url: str
|
||||
attribute_mapping: Optional[Dict[str, Any]] = None
|
||||
extra_config: Optional[Dict[str, Any]] = None
|
||||
attribute_mapping: dict[str, Any] | None = None
|
||||
extra_config: dict[str, Any] | None = None
|
||||
is_enabled: bool
|
||||
|
||||
|
||||
@@ -77,19 +76,19 @@ class OAuthProviderTestRequest(BaseModel):
|
||||
"""测试请求,使用表单数据而非数据库配置"""
|
||||
|
||||
client_id: str = Field(..., min_length=1)
|
||||
client_secret: Optional[str] = None
|
||||
authorization_url_override: Optional[str] = None
|
||||
token_url_override: Optional[str] = None
|
||||
client_secret: str | None = None
|
||||
authorization_url_override: str | None = None
|
||||
token_url_override: str | None = None
|
||||
redirect_uri: str = Field(..., min_length=1)
|
||||
|
||||
|
||||
@router.get("/supported-types", response_model=List[SupportedOAuthType])
|
||||
@router.get("/supported-types", response_model=list[SupportedOAuthType])
|
||||
async def get_supported_types(request: Request, db: Session = Depends(get_db)) -> Any:
|
||||
adapter = GetSupportedTypesAdapter()
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.get("/providers", response_model=List[OAuthProviderAdminResponse])
|
||||
@router.get("/providers", response_model=list[OAuthProviderAdminResponse])
|
||||
async def list_provider_configs(request: Request, db: Session = Depends(get_db)) -> Any:
|
||||
adapter = ListOAuthProviderConfigsAdapter()
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
@@ -1,8 +1,7 @@
|
||||
"""OAuth 公开端点(无需登录)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Optional
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, status
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -38,10 +37,10 @@ async def oauth_authorize(provider_type: str, db: Session = Depends(get_db)) ->
|
||||
async def oauth_callback(
|
||||
provider_type: str,
|
||||
db: Session = Depends(get_db),
|
||||
code: Optional[str] = Query(None),
|
||||
state: Optional[str] = Query(None),
|
||||
error: Optional[str] = Query(None),
|
||||
error_description: Optional[str] = Query(None),
|
||||
code: str | None = Query(None),
|
||||
state: str | None = Query(None),
|
||||
error: str | None = Query(None),
|
||||
error_description: str | None = Query(None),
|
||||
) -> RedirectResponse:
|
||||
"""
|
||||
OAuth 回调端点。
|
||||
|
||||
@@ -1,8 +1,7 @@
|
||||
"""OAuth 用户端点(需登录)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Optional, cast
|
||||
from typing import Any, cast
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -53,7 +52,7 @@ async def bind_oauth_provider(
|
||||
provider_type: str,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
bind_token: Optional[str] = None,
|
||||
bind_token: str | None = None,
|
||||
) -> RedirectResponse:
|
||||
"""发起 OAuth 绑定流程,支持通过 bind_token 参数进行安全认证"""
|
||||
adapter = BindOAuthProviderAdapter(provider_type=provider_type, bind_token=bind_token)
|
||||
@@ -115,10 +114,10 @@ class BindOAuthProviderAdapter(AuthenticatedApiAdapter):
|
||||
2. bind_token 参数 (浏览器跳转场景)
|
||||
"""
|
||||
|
||||
def __init__(self, provider_type: str, bind_token: Optional[str] = None):
|
||||
def __init__(self, provider_type: str, bind_token: str | None = None):
|
||||
self.provider_type = provider_type
|
||||
self.bind_token = bind_token
|
||||
self._user_from_bind_token: Optional[User] = None
|
||||
self._user_from_bind_token: User | None = None
|
||||
|
||||
@property
|
||||
def mode(self) -> ApiMode: # type: ignore[override]
|
||||
@@ -136,7 +135,7 @@ class BindOAuthProviderAdapter(AuthenticatedApiAdapter):
|
||||
raise HTTPException(status_code=401, detail="未登录")
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> RedirectResponse: # type: ignore[override]
|
||||
user: Optional[User] = context.user
|
||||
user: User | None = context.user
|
||||
|
||||
# 如果使用 bind_token,验证并获取用户
|
||||
if self.bind_token:
|
||||
|
||||
@@ -6,10 +6,9 @@
|
||||
from collections import defaultdict
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, Request
|
||||
from sqlalchemy import and_, func, or_
|
||||
from sqlalchemy import and_, or_
|
||||
from sqlalchemy.orm import Session, joinedload
|
||||
|
||||
from src.api.base.adapter import ApiAdapter, ApiMode
|
||||
@@ -41,10 +40,10 @@ router = APIRouter(prefix="/api/public", tags=["System Catalog"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
|
||||
|
||||
@router.get("/providers", response_model=List[PublicProviderResponse])
|
||||
@router.get("/providers", response_model=list[PublicProviderResponse])
|
||||
async def get_public_providers(
|
||||
request: Request,
|
||||
is_active: Optional[bool] = Query(None, description="过滤活跃状态"),
|
||||
is_active: bool | None = Query(None, description="过滤活跃状态"),
|
||||
skip: int = Query(0, description="跳过记录数"),
|
||||
limit: int = Query(100, description="返回记录数限制"),
|
||||
db: Session = Depends(get_db),
|
||||
@@ -77,11 +76,11 @@ async def get_public_providers(
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=ApiMode.PUBLIC)
|
||||
|
||||
|
||||
@router.get("/models", response_model=List[PublicModelResponse])
|
||||
@router.get("/models", response_model=list[PublicModelResponse])
|
||||
async def get_public_models(
|
||||
request: Request,
|
||||
provider_id: Optional[str] = Query(None, description="提供商ID过滤"),
|
||||
is_active: Optional[bool] = Query(None, description="过滤活跃状态"),
|
||||
provider_id: str | None = Query(None, description="提供商ID过滤"),
|
||||
is_active: bool | None = Query(None, description="过滤活跃状态"),
|
||||
skip: int = Query(0, description="跳过记录数"),
|
||||
limit: int = Query(100, description="返回记录数限制"),
|
||||
db: Session = Depends(get_db),
|
||||
@@ -145,7 +144,7 @@ async def get_public_stats(request: Request, db: Session = Depends(get_db)):
|
||||
async def search_models(
|
||||
request: Request,
|
||||
q: str = Query(..., description="搜索关键词"),
|
||||
provider_id: Optional[int] = Query(None, description="提供商ID过滤"),
|
||||
provider_id: int | None = Query(None, description="提供商ID过滤"),
|
||||
limit: int = Query(20, description="返回记录数限制"),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
@@ -234,8 +233,8 @@ async def get_public_global_models(
|
||||
request: Request,
|
||||
skip: int = Query(0, ge=0, description="跳过记录数"),
|
||||
limit: int = Query(100, ge=1, le=1000, description="返回记录数限制"),
|
||||
is_active: Optional[bool] = Query(None, description="过滤活跃状态"),
|
||||
search: Optional[str] = Query(None, description="搜索关键词"),
|
||||
is_active: bool | None = Query(None, description="过滤活跃状态"),
|
||||
search: str | None = Query(None, description="搜索关键词"),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
"""
|
||||
@@ -283,7 +282,7 @@ class PublicApiAdapter(ApiAdapter):
|
||||
|
||||
@dataclass
|
||||
class PublicProvidersAdapter(PublicApiAdapter):
|
||||
is_active: Optional[bool]
|
||||
is_active: bool | None
|
||||
skip: int
|
||||
limit: int
|
||||
|
||||
@@ -338,8 +337,8 @@ class PublicProvidersAdapter(PublicApiAdapter):
|
||||
|
||||
@dataclass
|
||||
class PublicModelsAdapter(PublicApiAdapter):
|
||||
provider_id: Optional[str]
|
||||
is_active: Optional[bool]
|
||||
provider_id: str | None
|
||||
is_active: bool | None
|
||||
skip: int
|
||||
limit: int
|
||||
|
||||
@@ -426,7 +425,7 @@ class PublicStatsAdapter(PublicApiAdapter):
|
||||
@dataclass
|
||||
class PublicSearchModelsAdapter(PublicApiAdapter):
|
||||
query: str
|
||||
provider_id: Optional[int]
|
||||
provider_id: int | None
|
||||
limit: int
|
||||
|
||||
async def handle(self, context): # type: ignore[override]
|
||||
@@ -508,7 +507,7 @@ class PublicApiFormatHealthMonitorAdapter(PublicApiAdapter):
|
||||
.all()
|
||||
)
|
||||
|
||||
all_formats: List[str] = []
|
||||
all_formats: list[str] = []
|
||||
for (api_format_enum,) in active_formats:
|
||||
api_format = (
|
||||
api_format_enum.value if hasattr(api_format_enum, "value") else str(api_format_enum)
|
||||
@@ -525,7 +524,7 @@ class PublicApiFormatHealthMonitorAdapter(PublicApiAdapter):
|
||||
)
|
||||
.all()
|
||||
)
|
||||
endpoint_map: Dict[str, List[str]] = defaultdict(list)
|
||||
endpoint_map: dict[str, list[str]] = defaultdict(list)
|
||||
for api_format_enum, endpoint_id in endpoint_rows:
|
||||
api_format = (
|
||||
api_format_enum.value if hasattr(api_format_enum, "value") else str(api_format_enum)
|
||||
@@ -551,7 +550,7 @@ class PublicApiFormatHealthMonitorAdapter(PublicApiAdapter):
|
||||
.all()
|
||||
)
|
||||
|
||||
grouped_candidates: Dict[str, List[RequestCandidate]] = {}
|
||||
grouped_candidates: dict[str, list[RequestCandidate]] = {}
|
||||
|
||||
for candidate, api_format_enum in rows:
|
||||
api_format = (
|
||||
@@ -564,7 +563,7 @@ class PublicApiFormatHealthMonitorAdapter(PublicApiAdapter):
|
||||
grouped_candidates[api_format].append(candidate)
|
||||
|
||||
# 3. 为所有活跃格式生成监控数据
|
||||
monitors: List[PublicApiFormatHealthMonitor] = []
|
||||
monitors: list[PublicApiFormatHealthMonitor] = []
|
||||
for api_format in all_formats:
|
||||
candidates = grouped_candidates.get(api_format, [])
|
||||
|
||||
@@ -579,7 +578,7 @@ class PublicApiFormatHealthMonitorAdapter(PublicApiAdapter):
|
||||
success_rate = success_count / actual_completed if actual_completed > 0 else 1.0
|
||||
|
||||
# 转换为公开版事件列表(不含敏感信息如 provider_id, key_id)
|
||||
events: List[PublicHealthEvent] = []
|
||||
events: list[PublicHealthEvent] = []
|
||||
for c in candidates:
|
||||
event_time = c.finished_at or c.started_at or c.created_at
|
||||
events.append(
|
||||
@@ -649,8 +648,8 @@ class PublicGlobalModelsAdapter(PublicApiAdapter):
|
||||
|
||||
skip: int
|
||||
limit: int
|
||||
is_active: Optional[bool]
|
||||
search: Optional[str]
|
||||
is_active: bool | None
|
||||
search: str | None
|
||||
|
||||
async def handle(self, context): # type: ignore[override]
|
||||
db = context.db
|
||||
|
||||
@@ -7,8 +7,6 @@
|
||||
- Authorization: Bearer (bearer) -> OpenAI 格式
|
||||
"""
|
||||
|
||||
from typing import Optional, Tuple, Union
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, Request
|
||||
from fastapi.responses import JSONResponse
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -35,7 +33,6 @@ from src.core.logger import logger
|
||||
from src.database import get_db
|
||||
from src.models.database import ApiKey, User
|
||||
from src.services.auth.service import AuthService
|
||||
from src.services.system.config import SystemConfigService
|
||||
|
||||
router = APIRouter(tags=["System Catalog"])
|
||||
|
||||
@@ -54,7 +51,7 @@ _ALL_CHAT_FORMATS = [
|
||||
|
||||
def _extract_api_key_from_request(
|
||||
request: Request, definition: ApiFormatDefinition
|
||||
) -> Optional[str]:
|
||||
) -> str | None:
|
||||
"""根据格式定义从请求中提取 API Key"""
|
||||
auth_header = definition.auth_header.lower()
|
||||
auth_type = definition.auth_type
|
||||
@@ -76,7 +73,7 @@ def _extract_api_key_from_request(
|
||||
return header_value
|
||||
|
||||
|
||||
def _detect_api_format_and_key(request: Request) -> Tuple[str, Optional[str]]:
|
||||
def _detect_api_format_and_key(request: Request) -> tuple[str, str | None]:
|
||||
"""
|
||||
根据请求头检测 API 格式并提取 API Key
|
||||
|
||||
@@ -163,7 +160,7 @@ def _build_empty_list_response(api_format: str) -> dict:
|
||||
|
||||
def _filter_formats_by_restrictions(
|
||||
formats: list[str], restrictions: AccessRestrictions, api_format: str
|
||||
) -> Tuple[list[str], Optional[dict]]:
|
||||
) -> tuple[list[str], dict | None]:
|
||||
"""
|
||||
根据访问限制过滤 API 格式
|
||||
|
||||
@@ -182,7 +179,7 @@ def _filter_formats_by_restrictions(
|
||||
return filtered, None
|
||||
|
||||
|
||||
def _authenticate(db: Session, api_key: Optional[str]) -> Tuple[Optional[User], Optional[ApiKey]]:
|
||||
def _authenticate(db: Session, api_key: str | None) -> tuple[User | None, ApiKey | None]:
|
||||
"""
|
||||
认证 API Key
|
||||
|
||||
@@ -248,8 +245,8 @@ def _build_auth_error_response(api_format: str) -> JSONResponse:
|
||||
|
||||
def _build_claude_list_response(
|
||||
models: list[ModelInfo],
|
||||
before_id: Optional[str],
|
||||
after_id: Optional[str],
|
||||
before_id: str | None,
|
||||
after_id: str | None,
|
||||
limit: int,
|
||||
) -> dict:
|
||||
"""构建 Claude 格式的列表响应"""
|
||||
@@ -309,7 +306,7 @@ def _build_openai_list_response(models: list[ModelInfo]) -> dict:
|
||||
def _build_gemini_list_response(
|
||||
models: list[ModelInfo],
|
||||
page_size: int,
|
||||
page_token: Optional[str],
|
||||
page_token: str | None,
|
||||
) -> dict:
|
||||
"""构建 Gemini 格式的列表响应"""
|
||||
# 处理分页
|
||||
@@ -435,14 +432,14 @@ def _build_404_response(model_id: str, api_format: str) -> JSONResponse:
|
||||
async def list_models(
|
||||
request: Request,
|
||||
# Claude 分页参数
|
||||
before_id: Optional[str] = Query(None, description="返回此 ID 之前的结果 (Claude)"),
|
||||
after_id: Optional[str] = Query(None, description="返回此 ID 之后的结果 (Claude)"),
|
||||
before_id: str | None = Query(None, description="返回此 ID 之前的结果 (Claude)"),
|
||||
after_id: str | None = Query(None, description="返回此 ID 之后的结果 (Claude)"),
|
||||
limit: int = Query(20, ge=1, le=1000, description="返回数量限制 (Claude)"),
|
||||
# Gemini 分页参数
|
||||
page_size: int = Query(50, alias="pageSize", ge=1, le=1000, description="每页数量 (Gemini)"),
|
||||
page_token: Optional[str] = Query(None, alias="pageToken", description="分页 token (Gemini)"),
|
||||
page_token: str | None = Query(None, alias="pageToken", description="分页 token (Gemini)"),
|
||||
db: Session = Depends(get_db),
|
||||
) -> Union[dict, JSONResponse]:
|
||||
) -> dict | JSONResponse:
|
||||
"""
|
||||
列出可用模型(统一端点)
|
||||
|
||||
@@ -556,7 +553,7 @@ async def retrieve_model(
|
||||
model_id: str,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
) -> Union[dict, JSONResponse]:
|
||||
) -> dict | JSONResponse:
|
||||
"""
|
||||
获取单个模型详情(统一端点)
|
||||
|
||||
@@ -658,9 +655,9 @@ async def retrieve_model(
|
||||
async def list_models_gemini(
|
||||
request: Request,
|
||||
page_size: int = Query(50, alias="pageSize", ge=1, le=1000),
|
||||
page_token: Optional[str] = Query(None, alias="pageToken"),
|
||||
page_token: str | None = Query(None, alias="pageToken"),
|
||||
db: Session = Depends(get_db),
|
||||
) -> Union[dict, JSONResponse]:
|
||||
) -> dict | JSONResponse:
|
||||
"""
|
||||
列出可用模型(Gemini v1beta 专用端点)
|
||||
|
||||
@@ -741,7 +738,7 @@ async def get_model_gemini(
|
||||
request: Request,
|
||||
model_name: str,
|
||||
db: Session = Depends(get_db),
|
||||
) -> Union[dict, JSONResponse]:
|
||||
) -> dict | JSONResponse:
|
||||
"""
|
||||
获取单个模型详情(Gemini v1beta 专用端点)
|
||||
|
||||
|
||||
@@ -1,12 +1,11 @@
|
||||
"""公开模块状态 API(供登录页等使用)"""
|
||||
|
||||
from typing import List
|
||||
|
||||
from fastapi import APIRouter, Depends
|
||||
from pydantic import BaseModel
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.core.modules import ModuleCategory, get_module_registry
|
||||
from src.core.modules import get_module_registry
|
||||
from src.database import get_db
|
||||
|
||||
router = APIRouter(prefix="/api/modules", tags=["Modules"])
|
||||
@@ -20,7 +19,7 @@ class AuthModuleInfo(BaseModel):
|
||||
active: bool
|
||||
|
||||
|
||||
@router.get("/auth-status", response_model=List[AuthModuleInfo])
|
||||
@router.get("/auth-status", response_model=list[AuthModuleInfo])
|
||||
async def get_auth_modules_status(db: Session = Depends(get_db)):
|
||||
"""
|
||||
获取认证模块状态(公开接口)
|
||||
|
||||
@@ -5,7 +5,7 @@ System Catalog / 健康检查相关端点
|
||||
"""
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
@@ -28,7 +28,7 @@ router = APIRouter(tags=["System Catalog"])
|
||||
# ============== 辅助函数 ==============
|
||||
|
||||
|
||||
def _as_bool(value: Optional[str], default: bool) -> bool:
|
||||
def _as_bool(value: str | None, default: bool) -> bool:
|
||||
"""将字符串转换为布尔值"""
|
||||
if value is None:
|
||||
return default
|
||||
@@ -39,9 +39,9 @@ def _serialize_provider(
|
||||
provider: Provider,
|
||||
include_models: bool,
|
||||
include_endpoints: bool,
|
||||
) -> Dict[str, Any]:
|
||||
) -> dict[str, Any]:
|
||||
"""序列化 Provider 对象"""
|
||||
provider_data: Dict[str, Any] = {
|
||||
provider_data: dict[str, Any] = {
|
||||
"id": provider.id,
|
||||
"name": provider.name,
|
||||
"is_active": provider.is_active,
|
||||
@@ -81,7 +81,7 @@ def _serialize_provider(
|
||||
return provider_data
|
||||
|
||||
|
||||
def _select_provider(db: Session, provider_name: Optional[str]) -> Optional[Provider]:
|
||||
def _select_provider(db: Session, provider_name: str | None) -> Provider | None:
|
||||
"""选择 Provider(按 provider_priority 优先级选择)"""
|
||||
query = db.query(Provider).filter(Provider.is_active == True)
|
||||
if provider_name:
|
||||
@@ -104,7 +104,7 @@ async def service_health(db: Session = Depends(get_db)):
|
||||
)
|
||||
active_models = db.query(func.count(Model.id)).filter(Model.is_active == True).scalar() or 0
|
||||
|
||||
redis_info: Dict[str, Any] = {"status": "unknown"}
|
||||
redis_info: dict[str, Any] = {"status": "unknown"}
|
||||
try:
|
||||
redis = await get_redis_client()
|
||||
if redis:
|
||||
@@ -245,9 +245,9 @@ async def provider_detail(
|
||||
async def test_connection(
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
provider: Optional[str] = Query(None),
|
||||
provider: str | None = Query(None),
|
||||
model: str = Query("claude-3-haiku-20240307"),
|
||||
api_format: Optional[str] = Query(None),
|
||||
api_format: str | None = Query(None),
|
||||
):
|
||||
"""测试 Provider 连接"""
|
||||
selected_provider = _select_provider(db, provider)
|
||||
|
||||
@@ -2,7 +2,6 @@
|
||||
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from fastapi.responses import JSONResponse
|
||||
@@ -55,13 +54,13 @@ class CreateManagementTokenRequest(BaseModel):
|
||||
"""创建 Management Token 请求"""
|
||||
|
||||
name: str = Field(..., min_length=1, max_length=100, description="Token 名称")
|
||||
description: Optional[str] = Field(None, max_length=500, description="描述")
|
||||
allowed_ips: Optional[list[str]] = Field(None, description="IP 白名单")
|
||||
expires_at: Optional[datetime] = Field(None, description="过期时间")
|
||||
description: str | None = Field(None, max_length=500, description="描述")
|
||||
allowed_ips: list[str] | None = Field(None, description="IP 白名单")
|
||||
expires_at: datetime | None = Field(None, description="过期时间")
|
||||
|
||||
@field_validator("allowed_ips")
|
||||
@classmethod
|
||||
def validate_allowed_ips(cls, v: Optional[list[str]]) -> Optional[list[str]]:
|
||||
def validate_allowed_ips(cls, v: list[str] | None) -> list[str] | None:
|
||||
return validate_ip_list(v)
|
||||
|
||||
@field_validator("expires_at", mode="before")
|
||||
@@ -81,10 +80,10 @@ class UpdateManagementTokenRequest(BaseModel):
|
||||
|
||||
model_config = {"extra": "allow"} # 允许额外字段以便检测哪些字段被显式提供
|
||||
|
||||
name: Optional[str] = Field(None, min_length=1, max_length=100)
|
||||
description: Optional[str] = Field(None, max_length=500)
|
||||
allowed_ips: Optional[list[str]] = None
|
||||
expires_at: Optional[datetime] = None
|
||||
name: str | None = Field(None, min_length=1, max_length=100)
|
||||
description: str | None = Field(None, max_length=500)
|
||||
allowed_ips: list[str] | None = None
|
||||
expires_at: datetime | None = None
|
||||
|
||||
# 用于追踪哪些字段被显式提供(包括显式设为 null 的情况)
|
||||
_provided_fields: set[str] = set()
|
||||
@@ -101,7 +100,7 @@ class UpdateManagementTokenRequest(BaseModel):
|
||||
|
||||
@field_validator("allowed_ips")
|
||||
@classmethod
|
||||
def validate_allowed_ips(cls, v: Optional[list[str]]) -> Optional[list[str]]:
|
||||
def validate_allowed_ips(cls, v: list[str] | None) -> list[str] | None:
|
||||
# 如果是 None,表示要清空,直接返回
|
||||
if v is None:
|
||||
return None
|
||||
@@ -122,7 +121,7 @@ class UpdateManagementTokenRequest(BaseModel):
|
||||
@router.get("")
|
||||
async def list_my_management_tokens(
|
||||
request: Request,
|
||||
is_active: Optional[bool] = Query(None, description="筛选激活状态"),
|
||||
is_active: bool | None = Query(None, description="筛选激活状态"),
|
||||
skip: int = Query(0, ge=0),
|
||||
limit: int = Query(50, ge=1, le=100),
|
||||
db: Session = Depends(get_db),
|
||||
@@ -347,7 +346,7 @@ class ListMyManagementTokensAdapter(ManagementTokenApiAdapter):
|
||||
"""列出用户的 Management Tokens"""
|
||||
|
||||
name: str = "list_my_management_tokens"
|
||||
is_active: Optional[bool] = None
|
||||
is_active: bool | None = None
|
||||
skip: int = 0
|
||||
limit: int = 50
|
||||
|
||||
|
||||
@@ -2,14 +2,12 @@
|
||||
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from pydantic import ValidationError
|
||||
from sqlalchemy import and_, func
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.base.adapter import ApiAdapter, ApiMode
|
||||
from src.api.base.authenticated_adapter import AuthenticatedApiAdapter
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.core.crypto import crypto_service
|
||||
@@ -170,9 +168,9 @@ async def toggle_my_api_key(key_id: str, request: Request, db: Session = Depends
|
||||
@router.get("/usage")
|
||||
async def get_my_usage(
|
||||
request: Request,
|
||||
start_date: Optional[datetime] = Query(None, description="开始时间(ISO 格式)"),
|
||||
end_date: Optional[datetime] = Query(None, description="结束时间(ISO 格式)"),
|
||||
search: Optional[str] = Query(None, description="搜索关键词(密钥名、模型名)"),
|
||||
start_date: datetime | None = Query(None, description="开始时间(ISO 格式)"),
|
||||
end_date: datetime | None = Query(None, description="结束时间(ISO 格式)"),
|
||||
search: str | None = Query(None, description="搜索关键词(密钥名、模型名)"),
|
||||
limit: int = Query(100, ge=1, le=200, description="每页记录数,默认100,最大200"),
|
||||
offset: int = Query(0, ge=0, le=2000, description="偏移量,用于分页,最大2000"),
|
||||
db: Session = Depends(get_db),
|
||||
@@ -200,7 +198,7 @@ async def get_my_usage(
|
||||
@router.get("/usage/active")
|
||||
async def get_my_active_requests(
|
||||
request: Request,
|
||||
ids: Optional[str] = Query(None, description="请求 ID 列表,逗号分隔"),
|
||||
ids: str | None = Query(None, description="请求 ID 列表,逗号分隔"),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
"""
|
||||
@@ -268,7 +266,7 @@ async def list_available_models(
|
||||
request: Request,
|
||||
skip: int = Query(0, ge=0, description="跳过记录数"),
|
||||
limit: int = Query(100, ge=1, le=1000, description="返回记录数限制"),
|
||||
search: Optional[str] = Query(None, description="搜索关键词"),
|
||||
search: str | None = Query(None, description="搜索关键词"),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
"""
|
||||
@@ -721,9 +719,9 @@ class ToggleMyApiKeyAdapter(AuthenticatedApiAdapter):
|
||||
class GetUsageAdapter(AuthenticatedApiAdapter):
|
||||
"""获取用户使用统计的适配器"""
|
||||
|
||||
start_date: Optional[datetime]
|
||||
end_date: Optional[datetime]
|
||||
search: Optional[str] = None
|
||||
start_date: datetime | None
|
||||
end_date: datetime | None
|
||||
search: str | None = None
|
||||
limit: int = 100
|
||||
offset: int = 0
|
||||
|
||||
@@ -983,7 +981,7 @@ class GetUsageAdapter(AuthenticatedApiAdapter):
|
||||
class GetActiveRequestsAdapter(AuthenticatedApiAdapter):
|
||||
"""轻量级活跃请求状态查询适配器(用于用户端轮询)"""
|
||||
|
||||
ids: Optional[str] = None
|
||||
ids: str | None = None
|
||||
|
||||
async def handle(self, context): # type: ignore[override]
|
||||
from src.services.usage import UsageService
|
||||
@@ -1045,7 +1043,7 @@ class ListAvailableModelsAdapter(AuthenticatedApiAdapter):
|
||||
|
||||
skip: int
|
||||
limit: int
|
||||
search: Optional[str]
|
||||
search: str | None
|
||||
|
||||
async def handle(self, context): # type: ignore[override]
|
||||
from sqlalchemy import or_
|
||||
@@ -1220,7 +1218,6 @@ class ListAvailableProvidersAdapter(AuthenticatedApiAdapter):
|
||||
async def handle(self, context): # type: ignore[override]
|
||||
from sqlalchemy.orm import selectinload
|
||||
|
||||
from src.models.database import ProviderEndpoint
|
||||
|
||||
db = context.db
|
||||
|
||||
|
||||
@@ -8,11 +8,13 @@
|
||||
3. 连接池复用:Keep-alive 连接减少 TCP 握手开销
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import hashlib
|
||||
import time
|
||||
from contextlib import asynccontextmanager
|
||||
from typing import Any, Dict, Optional, Tuple
|
||||
from typing import Any
|
||||
from urllib.parse import quote, urlparse
|
||||
|
||||
import httpx
|
||||
@@ -26,7 +28,7 @@ _proxy_clients_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]}"
|
||||
|
||||
|
||||
def build_proxy_url(proxy_config: Dict[str, Any]) -> Optional[str]:
|
||||
def build_proxy_url(proxy_config: dict[str, Any]) -> str | None:
|
||||
"""
|
||||
根据代理配置构建完整的代理 URL
|
||||
|
||||
@@ -103,11 +105,11 @@ class HTTPClientPool:
|
||||
3. LRU 淘汰:代理客户端超过上限时淘汰最久未使用的
|
||||
"""
|
||||
|
||||
_instance: Optional["HTTPClientPool"] = None
|
||||
_default_client: Optional[httpx.AsyncClient] = None
|
||||
_clients: Dict[str, httpx.AsyncClient] = {}
|
||||
_instance: HTTPClientPool | None = None
|
||||
_default_client: httpx.AsyncClient | None = None
|
||||
_clients: dict[str, httpx.AsyncClient] = {}
|
||||
# 代理客户端缓存:{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
|
||||
|
||||
@@ -242,7 +244,7 @@ class HTTPClientPool:
|
||||
@classmethod
|
||||
async def get_proxy_client(
|
||||
cls,
|
||||
proxy_config: Optional[Dict[str, Any]] = None,
|
||||
proxy_config: dict[str, Any] | None = None,
|
||||
) -> httpx.AsyncClient:
|
||||
"""
|
||||
获取代理客户端(带缓存复用)
|
||||
@@ -280,7 +282,7 @@ class HTTPClientPool:
|
||||
await cls._evict_lru_proxy_client()
|
||||
|
||||
# 创建新客户端(使用默认超时,请求时可覆盖)
|
||||
client_config: Dict[str, Any] = {
|
||||
client_config: dict[str, Any] = {
|
||||
"http2": False,
|
||||
"verify": get_ssl_context(),
|
||||
"follow_redirects": True,
|
||||
@@ -370,8 +372,8 @@ class HTTPClientPool:
|
||||
@classmethod
|
||||
def create_client_with_proxy(
|
||||
cls,
|
||||
proxy_config: Optional[Dict[str, Any]] = None,
|
||||
timeout: Optional[httpx.Timeout] = None,
|
||||
proxy_config: dict[str, Any] | None = None,
|
||||
timeout: httpx.Timeout | None = None,
|
||||
**kwargs: Any,
|
||||
) -> httpx.AsyncClient:
|
||||
"""
|
||||
@@ -387,7 +389,7 @@ class HTTPClientPool:
|
||||
Returns:
|
||||
配置好的 httpx.AsyncClient 实例(调用者需要负责关闭)
|
||||
"""
|
||||
client_config: Dict[str, Any] = {
|
||||
client_config: dict[str, Any] = {
|
||||
"http2": False,
|
||||
"verify": get_ssl_context(),
|
||||
"follow_redirects": True,
|
||||
@@ -413,7 +415,7 @@ class HTTPClientPool:
|
||||
return httpx.AsyncClient(**client_config)
|
||||
|
||||
@classmethod
|
||||
def get_pool_stats(cls) -> Dict[str, Any]:
|
||||
def get_pool_stats(cls) -> dict[str, Any]:
|
||||
"""获取连接池统计信息"""
|
||||
return {
|
||||
"default_client_active": cls._default_client is not None,
|
||||
|
||||
@@ -9,10 +9,11 @@
|
||||
- 调用方可以根据状态决定降级策略
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import time
|
||||
from enum import Enum
|
||||
from typing import Optional
|
||||
|
||||
import redis.asyncio as aioredis
|
||||
from src.core.logger import logger
|
||||
@@ -35,8 +36,8 @@ class RedisClientManager:
|
||||
提供 Redis 连接管理、熔断器保护和状态监控。
|
||||
"""
|
||||
|
||||
_instance: Optional["RedisClientManager"] = None
|
||||
_redis: Optional[aioredis.Redis] = None
|
||||
_instance: RedisClientManager | None = None
|
||||
_redis: aioredis.Redis | None = None
|
||||
|
||||
def __new__(cls):
|
||||
"""单例模式"""
|
||||
@@ -50,11 +51,11 @@ class RedisClientManager:
|
||||
return
|
||||
|
||||
self._initialized = True
|
||||
self._circuit_open_until: Optional[float] = None
|
||||
self._circuit_open_until: float | None = None
|
||||
self._consecutive_failures: int = 0
|
||||
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._last_error: Optional[str] = None # 记录最后一次错误
|
||||
self._last_error: str | None = None # 记录最后一次错误
|
||||
|
||||
def get_state(self) -> RedisState:
|
||||
"""
|
||||
@@ -100,7 +101,7 @@ class RedisClientManager:
|
||||
self._consecutive_failures = 0
|
||||
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连接
|
||||
|
||||
@@ -236,7 +237,7 @@ class RedisClientManager:
|
||||
self._redis = None
|
||||
logger.info("全局Redis客户端已关闭")
|
||||
|
||||
def get_client(self) -> Optional[aioredis.Redis]:
|
||||
def get_client(self) -> aioredis.Redis | None:
|
||||
"""
|
||||
获取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客户端
|
||||
|
||||
@@ -277,7 +278,7 @@ async def get_redis_client(require_redis: bool = False) -> Optional[aioredis.Red
|
||||
return _redis_manager.get_client()
|
||||
|
||||
|
||||
def get_redis_client_sync() -> Optional[aioredis.Redis]:
|
||||
def get_redis_client_sync() -> aioredis.Redis | None:
|
||||
"""
|
||||
同步获取Redis客户端(不会初始化)
|
||||
|
||||
|
||||
@@ -12,8 +12,9 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
import logging
|
||||
from typing import TYPE_CHECKING, Optional, Tuple
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.core.api_format.conversion.registry import FormatConversionRegistry
|
||||
@@ -26,11 +27,11 @@ logger = logging.getLogger(__name__)
|
||||
def is_format_compatible(
|
||||
client_format: str,
|
||||
endpoint_api_format: str,
|
||||
endpoint_format_acceptance_config: Optional[dict],
|
||||
endpoint_format_acceptance_config: dict | None,
|
||||
is_stream: bool,
|
||||
global_conversion_enabled: bool,
|
||||
registry: Optional["FormatConversionRegistry"] = None,
|
||||
) -> Tuple[bool, bool, Optional[str]]:
|
||||
registry: FormatConversionRegistry | None = None,
|
||||
) -> tuple[bool, bool, str | None]:
|
||||
"""
|
||||
检查端点是否兼容客户端格式
|
||||
|
||||
|
||||
@@ -4,7 +4,6 @@
|
||||
用于严格模式转换失败时抛出,让编排器可以尝试下一个候选。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
class FormatConversionError(Exception):
|
||||
|
||||
@@ -9,13 +9,11 @@
|
||||
这些应复用 `src/core/api_format/metadata.py`(API_FORMAT_DEFINITIONS)作为单一事实来源。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Dict, Set
|
||||
|
||||
|
||||
# 角色映射(仅作为辅助;system/tool 的具体落点以 Normalizer 规则为准)
|
||||
ROLE_MAPPINGS: Dict[str, Dict[str, str]] = {
|
||||
ROLE_MAPPINGS: dict[str, dict[str, str]] = {
|
||||
"OPENAI": {
|
||||
"user": "user",
|
||||
"assistant": "assistant",
|
||||
@@ -29,7 +27,7 @@ ROLE_MAPPINGS: Dict[str, Dict[str, str]] = {
|
||||
|
||||
|
||||
# 停止原因映射(internal -> provider),未知值使用 UNKNOWN 并写入 extra/raw
|
||||
STOP_REASON_MAPPINGS: Dict[str, Dict[str, str]] = {
|
||||
STOP_REASON_MAPPINGS: dict[str, dict[str, str]] = {
|
||||
"CLAUDE": {
|
||||
"end_turn": "end_turn",
|
||||
"max_tokens": "max_tokens",
|
||||
@@ -60,7 +58,7 @@ STOP_REASON_MAPPINGS: Dict[str, Dict[str, str]] = {
|
||||
|
||||
|
||||
# 使用量字段映射(provider usage field -> internal UsageInfo field)
|
||||
USAGE_FIELD_MAPPINGS: Dict[str, Dict[str, str]] = {
|
||||
USAGE_FIELD_MAPPINGS: dict[str, dict[str, str]] = {
|
||||
"CLAUDE": {
|
||||
"input_tokens": "input_tokens",
|
||||
"output_tokens": "output_tokens",
|
||||
@@ -82,7 +80,7 @@ USAGE_FIELD_MAPPINGS: Dict[str, Dict[str, str]] = {
|
||||
|
||||
|
||||
# 错误类型映射(provider -> internal ErrorType.value)
|
||||
ERROR_TYPE_MAPPINGS: Dict[str, Dict[str, str]] = {
|
||||
ERROR_TYPE_MAPPINGS: dict[str, dict[str, str]] = {
|
||||
"CLAUDE": {
|
||||
"invalid_request_error": "invalid_request",
|
||||
"authentication_error": "authentication",
|
||||
@@ -116,7 +114,7 @@ ERROR_TYPE_MAPPINGS: Dict[str, Dict[str, str]] = {
|
||||
|
||||
|
||||
# 可重试的错误类型(internal ErrorType.value)
|
||||
RETRYABLE_ERROR_TYPES: Set[str] = {
|
||||
RETRYABLE_ERROR_TYPES: set[str] = {
|
||||
"rate_limit",
|
||||
"overloaded",
|
||||
"server_error",
|
||||
|
||||
@@ -10,11 +10,10 @@
|
||||
- 兼容优先:UnknownBlock 在内部保留,但默认在输出阶段丢弃(可观测、可随时调整策略)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from enum import Enum
|
||||
from typing import Any, Dict, FrozenSet, List, Optional, Union
|
||||
from typing import Any
|
||||
|
||||
|
||||
class Role(str, Enum):
|
||||
@@ -65,7 +64,7 @@ class TextBlock:
|
||||
|
||||
type: ContentType = field(default=ContentType.TEXT, init=False)
|
||||
text: str = ""
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
extra: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -74,11 +73,11 @@ class ImageBlock:
|
||||
|
||||
type: ContentType = field(default=ContentType.IMAGE, init=False)
|
||||
# base64 编码的图片数据(二选一)
|
||||
data: Optional[str] = None
|
||||
media_type: Optional[str] = None
|
||||
data: str | None = None
|
||||
media_type: str | None = None
|
||||
# 或者 URL 引用
|
||||
url: Optional[str] = None
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
url: str | None = None
|
||||
extra: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -88,8 +87,8 @@ class ToolUseBlock:
|
||||
type: ContentType = field(default=ContentType.TOOL_USE, init=False)
|
||||
tool_id: str = ""
|
||||
tool_name: str = ""
|
||||
tool_input: Dict[str, Any] = field(default_factory=dict)
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
tool_input: dict[str, Any] = field(default_factory=dict)
|
||||
extra: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -100,9 +99,9 @@ class ToolResultBlock:
|
||||
tool_use_id: str = "" # 对应的 ToolUseBlock.tool_id
|
||||
# 工具输出可能是纯文本,也可能是结构化 JSON(Gemini functionResponse 等)
|
||||
output: Any = None
|
||||
content_text: Optional[str] = None
|
||||
content_text: str | None = None
|
||||
is_error: bool = False
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
extra: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -111,11 +110,11 @@ class UnknownBlock:
|
||||
|
||||
type: ContentType = field(default=ContentType.UNKNOWN, init=False)
|
||||
raw_type: str = "" # 原始的类型字符串(各格式不一致)
|
||||
payload: Dict[str, Any] = field(default_factory=dict) # 原始结构(尽量保持)
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
payload: 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
|
||||
@@ -123,8 +122,8 @@ class InternalMessage:
|
||||
"""统一的消息表示"""
|
||||
|
||||
role: Role
|
||||
content: List[ContentBlock] # 统一使用列表,纯文本用单个 TextBlock
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
content: list[ContentBlock] # 统一使用列表,纯文本用单个 TextBlock
|
||||
extra: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -132,9 +131,9 @@ class ToolDefinition:
|
||||
"""统一的工具定义"""
|
||||
|
||||
name: str
|
||||
description: Optional[str] = None
|
||||
parameters: Optional[Dict[str, Any]] = None # JSON Schema
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
description: str | None = None
|
||||
parameters: dict[str, Any] | None = None # JSON Schema
|
||||
extra: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
class ToolChoiceType(str, Enum):
|
||||
@@ -149,8 +148,8 @@ class ToolChoice:
|
||||
"""统一的工具选择"""
|
||||
|
||||
type: ToolChoiceType
|
||||
tool_name: Optional[str] = None
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
tool_name: str | None = None
|
||||
extra: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -159,7 +158,7 @@ class InstructionSegment:
|
||||
|
||||
role: Role # 仅允许 Role.SYSTEM / Role.DEVELOPER
|
||||
text: str = ""
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
extra: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -167,25 +166,25 @@ class InternalRequest:
|
||||
"""统一的请求表示"""
|
||||
|
||||
model: str
|
||||
messages: List[InternalMessage]
|
||||
messages: list[InternalMessage]
|
||||
|
||||
# 指令层:保留 system/developer 结构与顺序
|
||||
instructions: List[InstructionSegment] = field(default_factory=list)
|
||||
instructions: list[InstructionSegment] = field(default_factory=list)
|
||||
|
||||
# 兼容字段:instructions 的 join 文本(无 role 标签),用于 Claude/Gemini 这类仅接受字符串 system 的格式
|
||||
system: Optional[str] = None
|
||||
system: str | None = None
|
||||
|
||||
max_tokens: Optional[int] = None
|
||||
temperature: Optional[float] = None
|
||||
top_p: Optional[float] = None
|
||||
top_k: Optional[int] = None
|
||||
stop_sequences: Optional[List[str]] = None
|
||||
max_tokens: int | None = None
|
||||
temperature: float | None = None
|
||||
top_p: float | None = None
|
||||
top_k: int | None = None
|
||||
stop_sequences: list[str] | None = None
|
||||
stream: bool = False
|
||||
tools: Optional[List[ToolDefinition]] = None
|
||||
tool_choice: Optional[ToolChoice] = None # auto/none/required 或指定 tool_name
|
||||
extra: Dict[str, Any] = field(default_factory=dict) # 未识别字段透传
|
||||
tools: list[ToolDefinition] | None = None
|
||||
tool_choice: ToolChoice | None = None # auto/none/required 或指定 tool_name
|
||||
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 {
|
||||
"model": self.model,
|
||||
@@ -208,7 +207,7 @@ class UsageInfo:
|
||||
total_tokens: int = 0
|
||||
cache_read_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
|
||||
@@ -217,12 +216,12 @@ class InternalResponse:
|
||||
|
||||
id: str
|
||||
model: str
|
||||
content: List[ContentBlock]
|
||||
stop_reason: Optional[StopReason] = None
|
||||
usage: Optional[UsageInfo] = None
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
content: list[ContentBlock]
|
||||
stop_reason: StopReason | None = None
|
||||
usage: UsageInfo | None = None
|
||||
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
|
||||
if self.usage:
|
||||
@@ -246,12 +245,12 @@ class InternalError:
|
||||
|
||||
type: ErrorType
|
||||
message: str
|
||||
code: Optional[str] = None # 原始错误码
|
||||
param: Optional[str] = None # 导致错误的参数
|
||||
code: str | None = None # 原始错误码
|
||||
param: str | None = None # 导致错误的参数
|
||||
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 {
|
||||
"type": self.type.value,
|
||||
@@ -269,7 +268,7 @@ class FormatCapabilities:
|
||||
supports_error_conversion: bool = True
|
||||
supports_tools: bool = True
|
||||
supports_images: bool = False
|
||||
supported_features: FrozenSet[str] = field(default_factory=frozenset)
|
||||
supported_features: frozenset[str] = field(default_factory=frozenset)
|
||||
|
||||
|
||||
__all__ = [
|
||||
|
||||
@@ -5,10 +5,9 @@
|
||||
再从 internal 输出到目标格式。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Any, Dict, List, Optional
|
||||
from typing import Any
|
||||
|
||||
from .internal import FormatCapabilities, InternalError, InternalRequest, InternalResponse
|
||||
from .stream_events import InternalStreamEvent
|
||||
@@ -24,19 +23,19 @@ class FormatNormalizer(ABC):
|
||||
# ============ 请求转换 ============
|
||||
|
||||
@abstractmethod
|
||||
def request_to_internal(self, request: Dict[str, Any]) -> InternalRequest:
|
||||
def request_to_internal(self, request: dict[str, Any]) -> InternalRequest:
|
||||
"""将格式特定请求转换为内部表示"""
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
def request_from_internal(self, internal: InternalRequest) -> Dict[str, Any]:
|
||||
def request_from_internal(self, internal: InternalRequest) -> dict[str, Any]:
|
||||
"""将内部表示转换为格式特定请求"""
|
||||
raise NotImplementedError
|
||||
|
||||
# ============ 响应转换 ============
|
||||
|
||||
@abstractmethod
|
||||
def response_to_internal(self, response: Dict[str, Any]) -> InternalResponse:
|
||||
def response_to_internal(self, response: dict[str, Any]) -> InternalResponse:
|
||||
"""将格式特定响应转换为内部表示"""
|
||||
raise NotImplementedError
|
||||
|
||||
@@ -45,8 +44,8 @@ class FormatNormalizer(ABC):
|
||||
self,
|
||||
internal: InternalResponse,
|
||||
*,
|
||||
requested_model: Optional[str] = None,
|
||||
) -> Dict[str, Any]:
|
||||
requested_model: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""将内部表示转换为格式特定响应
|
||||
|
||||
Args:
|
||||
@@ -61,9 +60,9 @@ class FormatNormalizer(ABC):
|
||||
|
||||
def stream_chunk_to_internal(
|
||||
self,
|
||||
chunk: Dict[str, Any],
|
||||
chunk: dict[str, Any],
|
||||
state: StreamState,
|
||||
) -> List[InternalStreamEvent]:
|
||||
) -> list[InternalStreamEvent]:
|
||||
"""将格式特定流式块转换为内部事件"""
|
||||
raise NotImplementedError
|
||||
|
||||
@@ -71,21 +70,21 @@ class FormatNormalizer(ABC):
|
||||
self,
|
||||
event: InternalStreamEvent,
|
||||
state: StreamState,
|
||||
) -> List[Dict[str, Any]]:
|
||||
) -> list[dict[str, Any]]:
|
||||
"""将内部事件转换为格式特定流式块"""
|
||||
raise NotImplementedError
|
||||
|
||||
# ============ 错误转换(可选) ============
|
||||
|
||||
def is_error_response(self, response: Dict[str, Any]) -> bool:
|
||||
def is_error_response(self, response: dict[str, Any]) -> bool:
|
||||
"""基于 body 的兜底判断(不可靠),子类可覆盖"""
|
||||
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
|
||||
|
||||
def error_from_internal(self, internal: InternalError) -> Dict[str, Any]:
|
||||
def error_from_internal(self, internal: InternalError) -> dict[str, Any]:
|
||||
"""将内部错误表示转换为格式特定错误"""
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
@@ -6,7 +6,6 @@ Normalizers
|
||||
本目录在 Phase 1 仅创建结构;具体实现将在 Phase 2+ 补齐。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
__all__: list[str] = []
|
||||
|
||||
|
||||
@@ -7,10 +7,9 @@ Claude Messages API Normalizer
|
||||
- 可选:Claude error <-> InternalError
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
from typing import Any
|
||||
|
||||
from src.core.api_format.conversion.field_mappings import (
|
||||
ERROR_TYPE_MAPPINGS,
|
||||
@@ -63,7 +62,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
||||
supports_images=True,
|
||||
)
|
||||
|
||||
_CLAUDE_STOP_TO_INTERNAL: Dict[str, StopReason] = {
|
||||
_CLAUDE_STOP_TO_INTERNAL: dict[str, StopReason] = {
|
||||
"end_turn": StopReason.END_TURN,
|
||||
"max_tokens": StopReason.MAX_TOKENS,
|
||||
"stop_sequence": StopReason.STOP_SEQUENCE,
|
||||
@@ -73,7 +72,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
||||
"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.AUTHENTICATION: "authentication_error",
|
||||
ErrorType.PERMISSION_DENIED: "permission_error",
|
||||
@@ -90,11 +89,11 @@ class ClaudeNormalizer(FormatNormalizer):
|
||||
# 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 "")
|
||||
dropped: Dict[str, int] = {}
|
||||
dropped: dict[str, int] = {}
|
||||
|
||||
instructions: List[InstructionSegment] = []
|
||||
instructions: list[InstructionSegment] = []
|
||||
|
||||
# 顶层 system 先进入 instructions(保持确定性优先级)
|
||||
sys_value = request.get("system")
|
||||
@@ -103,7 +102,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
||||
if sys_text:
|
||||
instructions.append(InstructionSegment(role=Role.SYSTEM, text=sys_text))
|
||||
|
||||
messages: List[InternalMessage] = []
|
||||
messages: list[InternalMessage] = []
|
||||
for msg in request.get("messages") or []:
|
||||
if not isinstance(msg, dict):
|
||||
dropped["claude_message_non_dict"] = dropped.get("claude_message_non_dict", 0) + 1
|
||||
@@ -156,15 +155,15 @@ class ClaudeNormalizer(FormatNormalizer):
|
||||
|
||||
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)
|
||||
|
||||
# Claude Messages API: messages[] 仅允许 user/assistant,且需要交替;这里做最小修复
|
||||
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,
|
||||
"messages": out_messages,
|
||||
"max_tokens": internal.max_tokens if internal.max_tokens is not None else 4096,
|
||||
@@ -210,20 +209,20 @@ class ClaudeNormalizer(FormatNormalizer):
|
||||
# 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 "")
|
||||
model = str(response.get("model") or "")
|
||||
|
||||
blocks, dropped = self._claude_content_to_blocks(response.get("content"))
|
||||
|
||||
raw_stop = response.get("stop_reason")
|
||||
stop_reason: Optional[StopReason] = None
|
||||
stop_reason: StopReason | None = None
|
||||
if raw_stop is not None:
|
||||
stop_reason = self._CLAUDE_STOP_TO_INTERNAL.get(str(raw_stop), StopReason.UNKNOWN)
|
||||
|
||||
usage_info = self._claude_usage_to_internal(response.get("usage"))
|
||||
|
||||
extra: Dict[str, Any] = {}
|
||||
extra: dict[str, Any] = {}
|
||||
if raw_stop is not None:
|
||||
extra.setdefault("raw", {})["stop_reason"] = raw_stop
|
||||
|
||||
@@ -245,13 +244,13 @@ class ClaudeNormalizer(FormatNormalizer):
|
||||
self,
|
||||
internal: InternalResponse,
|
||||
*,
|
||||
requested_model: Optional[str] = None,
|
||||
) -> Dict[str, Any]:
|
||||
requested_model: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
cid = internal.id or "unknown"
|
||||
if not cid.startswith("msg_"):
|
||||
cid = f"msg_{cid}"
|
||||
|
||||
content: List[Dict[str, Any]] = []
|
||||
content: list[dict[str, Any]] = []
|
||||
for b in internal.content:
|
||||
if isinstance(b, TextBlock):
|
||||
if b.text:
|
||||
@@ -288,7 +287,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
||||
if internal.stop_reason is not None:
|
||||
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:
|
||||
usage = {
|
||||
"input_tokens": int(internal.usage.input_tokens),
|
||||
@@ -319,11 +318,11 @@ class ClaudeNormalizer(FormatNormalizer):
|
||||
|
||||
def stream_chunk_to_internal(
|
||||
self,
|
||||
chunk: Dict[str, Any],
|
||||
chunk: dict[str, Any],
|
||||
state: StreamState,
|
||||
) -> List[InternalStreamEvent]:
|
||||
) -> list[InternalStreamEvent]:
|
||||
ss = state.substate(self.FORMAT_ID)
|
||||
events: List[InternalStreamEvent] = []
|
||||
events: list[InternalStreamEvent] = []
|
||||
|
||||
event_type = chunk.get("type")
|
||||
if event_type is None:
|
||||
@@ -335,7 +334,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
||||
|
||||
if event_type == "message_start":
|
||||
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 "")
|
||||
# 保留初始化时设置的 model(客户端请求的模型),仅在空时用上游值
|
||||
model = state.model or str(message.get("model") or "")
|
||||
@@ -350,7 +349,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
||||
if event_type == "content_block_start":
|
||||
index = int(chunk.get("index") or 0)
|
||||
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")
|
||||
|
||||
if btype == "text":
|
||||
@@ -385,7 +384,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
||||
if event_type == "content_block_delta":
|
||||
index = int(chunk.get("index") or 0)
|
||||
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")
|
||||
|
||||
if dtype == "text_delta":
|
||||
@@ -415,7 +414,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
||||
|
||||
if event_type == "message_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")
|
||||
if raw_stop is not None:
|
||||
ss["stop_reason"] = str(raw_stop)
|
||||
@@ -426,7 +425,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
||||
|
||||
if event_type == "message_stop":
|
||||
raw_stop = ss.get("stop_reason")
|
||||
stop_reason: Optional[StopReason] = None
|
||||
stop_reason: StopReason | None = None
|
||||
if raw_stop is not None:
|
||||
stop_reason = self._CLAUDE_STOP_TO_INTERNAL.get(str(raw_stop), StopReason.UNKNOWN)
|
||||
usage_info = self._claude_usage_to_internal(ss.get("usage"))
|
||||
@@ -444,9 +443,9 @@ class ClaudeNormalizer(FormatNormalizer):
|
||||
self,
|
||||
event: InternalStreamEvent,
|
||||
state: StreamState,
|
||||
) -> List[Dict[str, Any]]:
|
||||
) -> list[dict[str, Any]]:
|
||||
ss = state.substate(self.FORMAT_ID)
|
||||
out: List[Dict[str, Any]] = []
|
||||
out: list[dict[str, Any]] = []
|
||||
|
||||
if isinstance(event, MessageStartEvent):
|
||||
state.message_id = event.message_id or state.message_id
|
||||
@@ -454,7 +453,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
||||
if not state.model:
|
||||
state.model = event.model or ""
|
||||
ss.setdefault("block_index_to_tool_id", {})
|
||||
message_obj: Dict[str, Any] = {
|
||||
message_obj: dict[str, Any] = {
|
||||
"id": state.message_id or "msg_stream",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
@@ -531,7 +530,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
||||
if event.stop_reason is not None:
|
||||
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",
|
||||
"delta": {"stop_reason": stop_reason},
|
||||
}
|
||||
@@ -553,15 +552,15 @@ class ClaudeNormalizer(FormatNormalizer):
|
||||
# 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):
|
||||
return False
|
||||
if response.get("type") == "error":
|
||||
return True
|
||||
return "error" in response
|
||||
|
||||
def error_to_internal(self, error_response: Dict[str, Any]) -> InternalError:
|
||||
err: Dict[str, Any] = {}
|
||||
def error_to_internal(self, error_response: dict[str, Any]) -> InternalError:
|
||||
err: dict[str, Any] = {}
|
||||
if isinstance(error_response, dict):
|
||||
err_raw = error_response.get("error")
|
||||
err = err_raw if isinstance(err_raw, dict) else {}
|
||||
@@ -580,9 +579,9 @@ class ClaudeNormalizer(FormatNormalizer):
|
||||
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")
|
||||
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:
|
||||
payload["param"] = internal.param
|
||||
if internal.code is not None:
|
||||
@@ -593,8 +592,8 @@ class ClaudeNormalizer(FormatNormalizer):
|
||||
# Helpers
|
||||
# =========================
|
||||
|
||||
def _claude_message_to_internal(self, msg: Dict[str, Any]) -> Tuple[Optional[InternalMessage], Dict[str, int]]:
|
||||
dropped: Dict[str, int] = {}
|
||||
def _claude_message_to_internal(self, msg: dict[str, Any]) -> tuple[InternalMessage | None, dict[str, int]]:
|
||||
dropped: dict[str, int] = {}
|
||||
role_raw = str(msg.get("role") or "unknown")
|
||||
|
||||
if role_raw == "user":
|
||||
@@ -616,8 +615,8 @@ class ClaudeNormalizer(FormatNormalizer):
|
||||
dropped,
|
||||
)
|
||||
|
||||
def _claude_content_to_blocks(self, content: Any) -> Tuple[List[ContentBlock], Dict[str, int]]:
|
||||
dropped: Dict[str, int] = {}
|
||||
def _claude_content_to_blocks(self, content: Any) -> tuple[list[ContentBlock], dict[str, int]]:
|
||||
dropped: dict[str, int] = {}
|
||||
if content is None:
|
||||
return [], dropped
|
||||
if isinstance(content, str):
|
||||
@@ -626,7 +625,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
||||
dropped["claude_content_non_list"] = dropped.get("claude_content_non_list", 0) + 1
|
||||
return [], dropped
|
||||
|
||||
blocks: List[ContentBlock] = []
|
||||
blocks: list[ContentBlock] = []
|
||||
for block in content:
|
||||
if not isinstance(block, dict):
|
||||
dropped["claude_block_non_dict"] = dropped.get("claude_block_non_dict", 0) + 1
|
||||
@@ -641,7 +640,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
||||
|
||||
if btype == "image":
|
||||
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")
|
||||
if stype == "base64":
|
||||
data = src.get("data")
|
||||
@@ -687,7 +686,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
||||
tool_use_id: str,
|
||||
raw_content: Any,
|
||||
is_error: bool,
|
||||
raw_block: Dict[str, Any],
|
||||
raw_block: dict[str, Any],
|
||||
) -> ToolResultBlock:
|
||||
if raw_content is None:
|
||||
return ToolResultBlock(
|
||||
@@ -723,7 +722,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
||||
)
|
||||
|
||||
if isinstance(raw_content, list):
|
||||
text_parts: List[str] = []
|
||||
text_parts: list[str] = []
|
||||
for part in raw_content:
|
||||
if isinstance(part, dict) and part.get("type") == "text":
|
||||
text = part.get("text")
|
||||
@@ -747,15 +746,15 @@ class ClaudeNormalizer(FormatNormalizer):
|
||||
extra={"claude": raw_block},
|
||||
)
|
||||
|
||||
def _collapse_claude_system(self, system_value: Any) -> Tuple[Optional[str], Dict[str, int]]:
|
||||
dropped: Dict[str, int] = {}
|
||||
def _collapse_claude_system(self, system_value: Any) -> tuple[str | None, dict[str, int]]:
|
||||
dropped: dict[str, int] = {}
|
||||
if system_value is None:
|
||||
return None, dropped
|
||||
if isinstance(system_value, str):
|
||||
return (system_value or None), dropped
|
||||
|
||||
if isinstance(system_value, list):
|
||||
texts: List[str] = []
|
||||
texts: list[str] = []
|
||||
for item in system_value:
|
||||
if not isinstance(item, dict):
|
||||
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
|
||||
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]
|
||||
joined = "\n\n".join(parts)
|
||||
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):
|
||||
return None
|
||||
|
||||
out: List[ToolDefinition] = []
|
||||
out: list[ToolDefinition] = []
|
||||
for tool in tools:
|
||||
if not isinstance(tool, dict):
|
||||
continue
|
||||
@@ -799,7 +798,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
||||
)
|
||||
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:
|
||||
return None
|
||||
if not isinstance(tool_choice, dict):
|
||||
@@ -818,7 +817,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
||||
|
||||
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:
|
||||
return {"type": "none"}
|
||||
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": "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"
|
||||
|
||||
blocks: List[Dict[str, Any]] = []
|
||||
text_parts: List[str] = []
|
||||
blocks: list[dict[str, Any]] = []
|
||||
text_parts: list[str] = []
|
||||
|
||||
for b in msg.content:
|
||||
if isinstance(b, UnknownBlock):
|
||||
@@ -902,8 +901,8 @@ class ClaudeNormalizer(FormatNormalizer):
|
||||
|
||||
return {"role": role, "content": blocks}
|
||||
|
||||
def _coerce_claude_message_sequence(self, messages: List[InternalMessage]) -> List[InternalMessage]:
|
||||
normalized: List[InternalMessage] = []
|
||||
def _coerce_claude_message_sequence(self, messages: list[InternalMessage]) -> list[InternalMessage]:
|
||||
normalized: list[InternalMessage] = []
|
||||
for m in messages:
|
||||
role = m.role
|
||||
if role not in (Role.USER, Role.ASSISTANT):
|
||||
@@ -916,7 +915,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
||||
if normalized[0].role != Role.USER:
|
||||
normalized = [InternalMessage(role=Role.USER, content=[])] + normalized
|
||||
|
||||
merged: List[InternalMessage] = []
|
||||
merged: list[InternalMessage] = []
|
||||
for m in normalized:
|
||||
if merged and merged[-1].role == m.role:
|
||||
merged[-1].content.extend(m.content)
|
||||
@@ -925,12 +924,12 @@ class ClaudeNormalizer(FormatNormalizer):
|
||||
|
||||
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):
|
||||
return None
|
||||
|
||||
mapping = USAGE_FIELD_MAPPINGS.get("CLAUDE", {})
|
||||
fields: Dict[str, int] = {}
|
||||
fields: dict[str, int] = {}
|
||||
extra = self._extract_extra(usage, set(mapping.keys()))
|
||||
|
||||
for provider_key, internal_key in mapping.items():
|
||||
@@ -952,8 +951,8 @@ class ClaudeNormalizer(FormatNormalizer):
|
||||
extra={"claude": extra} if extra else {},
|
||||
)
|
||||
|
||||
def _usage_to_claude(self, usage: UsageInfo) -> Dict[str, Any]:
|
||||
result: Dict[str, Any] = {
|
||||
def _usage_to_claude(self, usage: UsageInfo) -> dict[str, Any]:
|
||||
result: dict[str, Any] = {
|
||||
"input_tokens": int(usage.input_tokens),
|
||||
"output_tokens": int(usage.output_tokens),
|
||||
}
|
||||
@@ -969,7 +968,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
||||
except ValueError:
|
||||
return ErrorType.UNKNOWN
|
||||
|
||||
def _optional_int(self, value: Any) -> Optional[int]:
|
||||
def _optional_int(self, value: Any) -> int | None:
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
@@ -977,7 +976,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
|
||||
def _optional_float(self, value: Any) -> Optional[float]:
|
||||
def _optional_float(self, value: Any) -> float | None:
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
@@ -985,7 +984,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
||||
except (TypeError, ValueError):
|
||||
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:
|
||||
return None
|
||||
if isinstance(value, str):
|
||||
@@ -994,10 +993,10 @@ class ClaudeNormalizer(FormatNormalizer):
|
||||
return [str(x) for x in value if x is not 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}
|
||||
|
||||
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():
|
||||
target[k] = target.get(k, 0) + int(v)
|
||||
|
||||
|
||||
@@ -7,7 +7,6 @@ CLAUDE_CLI 的请求/响应 body 与 CLAUDE 一致(Anthropic Messages API)
|
||||
如需 CLI 特殊处理,可覆盖 request_from_internal / request_to_internal 等方法。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from src.core.api_format.conversion.normalizers.claude import ClaudeNormalizer
|
||||
|
||||
|
||||
@@ -11,11 +11,9 @@ Gemini (GenerateContent / streamGenerateContent) Normalizer
|
||||
- 响应/流式通常为 camelCase(candidates/finishReason/usageMetadata/modelVersion)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import time
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
from typing import Any
|
||||
|
||||
from src.core.api_format.conversion.field_mappings import (
|
||||
ERROR_TYPE_MAPPINGS,
|
||||
@@ -68,7 +66,7 @@ class GeminiNormalizer(FormatNormalizer):
|
||||
supports_images=True,
|
||||
)
|
||||
|
||||
_FINISH_REASON_TO_STOP: Dict[str, StopReason] = {
|
||||
_FINISH_REASON_TO_STOP: dict[str, StopReason] = {
|
||||
"STOP": StopReason.END_TURN,
|
||||
"MAX_TOKENS": StopReason.MAX_TOKENS,
|
||||
"SAFETY": StopReason.CONTENT_FILTERED,
|
||||
@@ -77,7 +75,7 @@ class GeminiNormalizer(FormatNormalizer):
|
||||
"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.AUTHENTICATION: "UNAUTHENTICATED",
|
||||
ErrorType.PERMISSION_DENIED: "PERMISSION_DENIED",
|
||||
@@ -94,11 +92,11 @@ class GeminiNormalizer(FormatNormalizer):
|
||||
# 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 "")
|
||||
dropped: Dict[str, int] = {}
|
||||
dropped: dict[str, int] = {}
|
||||
|
||||
instructions: List[InstructionSegment] = []
|
||||
instructions: list[InstructionSegment] = []
|
||||
system_text, sys_dropped = self._collapse_system_instruction(
|
||||
request.get("system_instruction")
|
||||
if "system_instruction" in request
|
||||
@@ -108,7 +106,7 @@ class GeminiNormalizer(FormatNormalizer):
|
||||
if system_text:
|
||||
instructions.append(InstructionSegment(role=Role.SYSTEM, text=system_text))
|
||||
|
||||
messages: List[InternalMessage] = []
|
||||
messages: list[InternalMessage] = []
|
||||
contents = request.get("contents") or []
|
||||
if isinstance(contents, list):
|
||||
for content in contents:
|
||||
@@ -152,7 +150,7 @@ class GeminiNormalizer(FormatNormalizer):
|
||||
)
|
||||
|
||||
# 构建 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 等)
|
||||
# 这些字段在 _get_generation_config 中已提取,需要单独存储以便转换时使用
|
||||
@@ -160,7 +158,7 @@ class GeminiNormalizer(FormatNormalizer):
|
||||
response_modalities = generation_config.get("response_modalities")
|
||||
thinking_config = generation_config.get("thinking_config")
|
||||
if response_modalities or thinking_config:
|
||||
google_extra: Dict[str, Any] = {}
|
||||
google_extra: dict[str, Any] = {}
|
||||
if response_modalities:
|
||||
google_extra["response_modalities"] = response_modalities
|
||||
if thinking_config:
|
||||
@@ -188,7 +186,7 @@ class GeminiNormalizer(FormatNormalizer):
|
||||
|
||||
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)
|
||||
|
||||
# tools/tool_choice
|
||||
@@ -212,7 +210,7 @@ class GeminiNormalizer(FormatNormalizer):
|
||||
if 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:
|
||||
generation_config["max_output_tokens"] = internal.max_tokens
|
||||
if internal.temperature is not None:
|
||||
@@ -231,7 +229,7 @@ class GeminiNormalizer(FormatNormalizer):
|
||||
thinking_config = google_extra.get("thinking_config")
|
||||
if isinstance(thinking_config, dict):
|
||||
# snake_case -> camelCase 转换
|
||||
gemini_thinking: Dict[str, Any] = {}
|
||||
gemini_thinking: dict[str, Any] = {}
|
||||
if "thinking_budget" in thinking_config:
|
||||
gemini_thinking["thinkingBudget"] = thinking_config["thinking_budget"]
|
||||
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:
|
||||
generation_config["thinkingConfig"] = orig_gc["thinking_config"]
|
||||
|
||||
contents: List[Dict[str, Any]] = []
|
||||
contents: list[dict[str, Any]] = []
|
||||
for msg in internal.messages:
|
||||
contents.append(self._internal_message_to_content(msg))
|
||||
|
||||
result: Dict[str, Any] = {
|
||||
result: dict[str, Any] = {
|
||||
"contents": contents,
|
||||
}
|
||||
|
||||
@@ -295,7 +293,7 @@ class GeminiNormalizer(FormatNormalizer):
|
||||
# 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 "")
|
||||
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"))
|
||||
|
||||
extra: Dict[str, Any] = {}
|
||||
extra: dict[str, Any] = {}
|
||||
if finish_reason is not None:
|
||||
extra.setdefault("raw", {})["finishReason"] = finish_reason
|
||||
|
||||
@@ -337,9 +335,9 @@ class GeminiNormalizer(FormatNormalizer):
|
||||
self,
|
||||
internal: InternalResponse,
|
||||
*,
|
||||
requested_model: Optional[str] = None,
|
||||
) -> Dict[str, Any]:
|
||||
parts: List[Dict[str, Any]] = []
|
||||
requested_model: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
parts: list[dict[str, Any]] = []
|
||||
for b in internal.content:
|
||||
if isinstance(b, TextBlock):
|
||||
if b.text:
|
||||
@@ -374,7 +372,7 @@ class GeminiNormalizer(FormatNormalizer):
|
||||
if internal.stop_reason is not None:
|
||||
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:
|
||||
usage_metadata = {
|
||||
"promptTokenCount": int(internal.usage.input_tokens),
|
||||
@@ -384,7 +382,7 @@ class GeminiNormalizer(FormatNormalizer):
|
||||
if 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"},
|
||||
"index": 0,
|
||||
}
|
||||
@@ -394,7 +392,7 @@ class GeminiNormalizer(FormatNormalizer):
|
||||
# 优先使用用户请求的原始模型名,回退到上游返回的模型名
|
||||
model_name = requested_model if requested_model else (internal.model or "gemini")
|
||||
|
||||
out: Dict[str, Any] = {
|
||||
out: dict[str, Any] = {
|
||||
"candidates": [candidate],
|
||||
"modelVersion": model_name,
|
||||
}
|
||||
@@ -412,9 +410,9 @@ class GeminiNormalizer(FormatNormalizer):
|
||||
# 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)
|
||||
events: List[InternalStreamEvent] = []
|
||||
events: list[InternalStreamEvent] = []
|
||||
|
||||
if not ss.get("message_started"):
|
||||
# 保留初始化时设置的 model(客户端请求的模型),仅在空时用上游值
|
||||
@@ -544,11 +542,11 @@ class GeminiNormalizer(FormatNormalizer):
|
||||
self,
|
||||
event: InternalStreamEvent,
|
||||
state: StreamState,
|
||||
) -> List[Dict[str, Any]]:
|
||||
) -> list[dict[str, Any]]:
|
||||
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 {
|
||||
"candidates": [
|
||||
{
|
||||
@@ -622,7 +620,7 @@ class GeminiNormalizer(FormatNormalizer):
|
||||
|
||||
name = str(entry.get("name") or "")
|
||||
raw_json = str(entry.get("json") or "")
|
||||
args: Dict[str, Any] = {}
|
||||
args: dict[str, Any] = {}
|
||||
if raw_json:
|
||||
try:
|
||||
parsed = json.loads(raw_json)
|
||||
@@ -639,7 +637,7 @@ class GeminiNormalizer(FormatNormalizer):
|
||||
if event.stop_reason is not None:
|
||||
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:
|
||||
chunk["candidates"][0]["finishReason"] = finish_reason
|
||||
|
||||
@@ -665,10 +663,10 @@ class GeminiNormalizer(FormatNormalizer):
|
||||
# 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
|
||||
|
||||
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 = err if isinstance(err, dict) else {}
|
||||
|
||||
@@ -691,9 +689,9 @@ class GeminiNormalizer(FormatNormalizer):
|
||||
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")
|
||||
payload: Dict[str, Any] = {
|
||||
payload: dict[str, Any] = {
|
||||
"code": 400 if internal.type == ErrorType.INVALID_REQUEST else 500,
|
||||
"message": internal.message,
|
||||
"status": status,
|
||||
@@ -704,8 +702,8 @@ class GeminiNormalizer(FormatNormalizer):
|
||||
# Helpers
|
||||
# =========================
|
||||
|
||||
def _content_to_internal_message(self, content: Dict[str, Any]) -> Tuple[Optional[InternalMessage], Dict[str, int]]:
|
||||
dropped: Dict[str, int] = {}
|
||||
def _content_to_internal_message(self, content: dict[str, Any]) -> tuple[InternalMessage | None, dict[str, int]]:
|
||||
dropped: dict[str, int] = {}
|
||||
|
||||
role_raw = str(content.get("role") or "user")
|
||||
if role_raw == "model":
|
||||
@@ -727,15 +725,15 @@ class GeminiNormalizer(FormatNormalizer):
|
||||
dropped,
|
||||
)
|
||||
|
||||
def _parts_to_blocks(self, parts: Any) -> Tuple[List[ContentBlock], Dict[str, int]]:
|
||||
dropped: Dict[str, int] = {}
|
||||
def _parts_to_blocks(self, parts: Any) -> tuple[list[ContentBlock], dict[str, int]]:
|
||||
dropped: dict[str, int] = {}
|
||||
if parts is None:
|
||||
return [], dropped
|
||||
if not isinstance(parts, list):
|
||||
dropped["gemini_parts_non_list"] = dropped.get("gemini_parts_non_list", 0) + 1
|
||||
return [], dropped
|
||||
|
||||
blocks: List[ContentBlock] = []
|
||||
blocks: list[ContentBlock] = []
|
||||
for part in parts:
|
||||
if not isinstance(part, dict):
|
||||
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 "")
|
||||
response = func_resp.get("response")
|
||||
output: Any = None
|
||||
content_text: Optional[str] = None
|
||||
content_text: str | None = None
|
||||
|
||||
# 兼容历史:response 常见结构为 {"result": ...}
|
||||
if isinstance(response, dict) and "result" in response:
|
||||
@@ -815,10 +813,10 @@ class GeminiNormalizer(FormatNormalizer):
|
||||
|
||||
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"
|
||||
|
||||
parts: List[Dict[str, Any]] = []
|
||||
parts: list[dict[str, Any]] = []
|
||||
for b in msg.content:
|
||||
if isinstance(b, UnknownBlock):
|
||||
continue
|
||||
@@ -861,8 +859,8 @@ class GeminiNormalizer(FormatNormalizer):
|
||||
|
||||
return {"role": role, "parts": parts}
|
||||
|
||||
def _collapse_system_instruction(self, system_instruction: Any) -> Tuple[Optional[str], Dict[str, int]]:
|
||||
dropped: Dict[str, int] = {}
|
||||
def _collapse_system_instruction(self, system_instruction: Any) -> tuple[str | None, dict[str, int]]:
|
||||
dropped: dict[str, int] = {}
|
||||
if system_instruction is None:
|
||||
return None, dropped
|
||||
|
||||
@@ -870,7 +868,7 @@ class GeminiNormalizer(FormatNormalizer):
|
||||
if isinstance(system_instruction, dict):
|
||||
parts = system_instruction.get("parts")
|
||||
if isinstance(parts, list):
|
||||
texts: List[str] = []
|
||||
texts: list[str] = []
|
||||
for part in parts:
|
||||
if isinstance(part, dict) and "text" in part and 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
|
||||
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
|
||||
gc = request.get("generation_config") if "generation_config" in request else request.get("generationConfig")
|
||||
if not isinstance(gc, dict):
|
||||
@@ -893,7 +891,7 @@ class GeminiNormalizer(FormatNormalizer):
|
||||
return gc.get(k)
|
||||
return None
|
||||
|
||||
normalized: Dict[str, Any] = {}
|
||||
normalized: dict[str, Any] = {}
|
||||
normalized["max_output_tokens"] = pick("max_output_tokens", "maxOutputTokens")
|
||||
normalized["temperature"] = pick("temperature")
|
||||
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}
|
||||
|
||||
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):
|
||||
return None
|
||||
|
||||
out: List[ToolDefinition] = []
|
||||
out: list[ToolDefinition] = []
|
||||
for tool in tools:
|
||||
if not isinstance(tool, dict):
|
||||
continue
|
||||
@@ -945,7 +943,7 @@ class GeminiNormalizer(FormatNormalizer):
|
||||
|
||||
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:
|
||||
return None
|
||||
if not isinstance(tool_config, dict):
|
||||
@@ -972,9 +970,9 @@ class GeminiNormalizer(FormatNormalizer):
|
||||
|
||||
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"
|
||||
cfg: Dict[str, Any] = {}
|
||||
cfg: dict[str, Any] = {}
|
||||
|
||||
if tool_choice.type == ToolChoiceType.NONE:
|
||||
mode = "NONE"
|
||||
@@ -987,12 +985,12 @@ class GeminiNormalizer(FormatNormalizer):
|
||||
cfg["mode"] = mode
|
||||
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):
|
||||
return None
|
||||
|
||||
mapping = USAGE_FIELD_MAPPINGS.get("GEMINI", {})
|
||||
fields: Dict[str, int] = {}
|
||||
fields: dict[str, int] = {}
|
||||
extra = self._extract_extra(usage_metadata, set(mapping.keys()))
|
||||
|
||||
# promptTokenCount/candidatesTokenCount/totalTokenCount/cachedContentTokenCount
|
||||
@@ -1023,7 +1021,7 @@ class GeminiNormalizer(FormatNormalizer):
|
||||
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]
|
||||
joined = "\n\n".join(parts)
|
||||
return joined or None
|
||||
@@ -1034,7 +1032,7 @@ class GeminiNormalizer(FormatNormalizer):
|
||||
except ValueError:
|
||||
return ErrorType.UNKNOWN
|
||||
|
||||
def _optional_int(self, value: Any) -> Optional[int]:
|
||||
def _optional_int(self, value: Any) -> int | None:
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
@@ -1042,7 +1040,7 @@ class GeminiNormalizer(FormatNormalizer):
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
|
||||
def _optional_float(self, value: Any) -> Optional[float]:
|
||||
def _optional_float(self, value: Any) -> float | None:
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
@@ -1050,7 +1048,7 @@ class GeminiNormalizer(FormatNormalizer):
|
||||
except (TypeError, ValueError):
|
||||
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:
|
||||
return None
|
||||
if isinstance(value, str):
|
||||
@@ -1059,10 +1057,10 @@ class GeminiNormalizer(FormatNormalizer):
|
||||
return [str(x) for x in value if x is not 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}
|
||||
|
||||
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():
|
||||
target[k] = target.get(k, 0) + int(v)
|
||||
|
||||
|
||||
@@ -7,7 +7,6 @@ GEMINI_CLI 的请求/响应 body 与 GEMINI 一致(Google Gemini API),差
|
||||
如需 CLI 特殊处理,可覆盖 request_from_internal / request_to_internal 等方法。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from src.core.api_format.conversion.normalizers.gemini import GeminiNormalizer
|
||||
|
||||
|
||||
@@ -7,11 +7,10 @@ OpenAI Chat Completions Normalizer
|
||||
- 可选:OpenAI error <-> InternalError
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import time
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
from typing import Any
|
||||
|
||||
from src.core.logger import logger
|
||||
from src.core.api_format.conversion.field_mappings import (
|
||||
@@ -49,7 +48,6 @@ from src.core.api_format.conversion.stream_events import (
|
||||
InternalStreamEvent,
|
||||
MessageStartEvent,
|
||||
MessageStopEvent,
|
||||
StreamEventType,
|
||||
ToolCallDeltaEvent,
|
||||
)
|
||||
from src.core.api_format.conversion.stream_state import StreamState
|
||||
@@ -65,7 +63,7 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
)
|
||||
|
||||
# finish_reason -> StopReason
|
||||
_FINISH_REASON_TO_STOP: Dict[str, StopReason] = {
|
||||
_FINISH_REASON_TO_STOP: dict[str, StopReason] = {
|
||||
"stop": StopReason.END_TURN,
|
||||
"length": StopReason.MAX_TOKENS,
|
||||
"tool_calls": StopReason.TOOL_USE,
|
||||
@@ -74,7 +72,7 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
}
|
||||
|
||||
# StopReason -> finish_reason
|
||||
_STOP_TO_FINISH_REASON: Dict[StopReason, str] = {
|
||||
_STOP_TO_FINISH_REASON: dict[StopReason, str] = {
|
||||
StopReason.END_TURN: "stop",
|
||||
StopReason.MAX_TOKENS: "length",
|
||||
StopReason.STOP_SEQUENCE: "stop",
|
||||
@@ -84,7 +82,7 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
}
|
||||
|
||||
# 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.AUTHENTICATION: "invalid_api_key",
|
||||
ErrorType.PERMISSION_DENIED: "invalid_request_error",
|
||||
@@ -101,13 +99,13 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
# 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 "")
|
||||
|
||||
dropped: Dict[str, int] = {}
|
||||
dropped: dict[str, int] = {}
|
||||
|
||||
instructions: List[InstructionSegment] = []
|
||||
messages: List[InternalMessage] = []
|
||||
instructions: list[InstructionSegment] = []
|
||||
messages: list[InternalMessage] = []
|
||||
|
||||
for msg in request.get("messages") or []:
|
||||
if not isinstance(msg, dict):
|
||||
@@ -146,7 +144,7 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
)
|
||||
|
||||
# 构建 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 = request.get("extra_body")
|
||||
@@ -175,8 +173,8 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
|
||||
return internal
|
||||
|
||||
def request_from_internal(self, internal: InternalRequest) -> Dict[str, Any]:
|
||||
out_messages: List[Dict[str, Any]] = []
|
||||
def request_from_internal(self, internal: InternalRequest) -> dict[str, Any]:
|
||||
out_messages: list[dict[str, Any]] = []
|
||||
|
||||
if internal.instructions:
|
||||
for seg in internal.instructions:
|
||||
@@ -189,7 +187,7 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
for msg in internal.messages:
|
||||
out_messages.extend(self._internal_message_to_openai_messages(msg))
|
||||
|
||||
result: Dict[str, Any] = {
|
||||
result: dict[str, Any] = {
|
||||
"model": internal.model,
|
||||
"messages": out_messages,
|
||||
}
|
||||
@@ -233,11 +231,11 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
# 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 "")
|
||||
model = str(response.get("model") or "")
|
||||
|
||||
extra: Dict[str, Any] = {}
|
||||
extra: dict[str, Any] = {}
|
||||
|
||||
choices = response.get("choices") or []
|
||||
if isinstance(choices, list) and len(choices) > 1:
|
||||
@@ -285,12 +283,12 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
self,
|
||||
internal: InternalResponse,
|
||||
*,
|
||||
requested_model: Optional[str] = None,
|
||||
) -> Dict[str, Any]:
|
||||
requested_model: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
# OpenAI Chat Completions response envelope
|
||||
# 优先使用用户请求的原始模型名,回退到上游返回的模型名
|
||||
model_name = requested_model if requested_model else internal.model
|
||||
out: Dict[str, Any] = {
|
||||
out: dict[str, Any] = {
|
||||
"id": internal.id or "chatcmpl-unknown",
|
||||
"object": "chat.completion",
|
||||
"created": int(time.time()),
|
||||
@@ -298,7 +296,7 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
"choices": [],
|
||||
}
|
||||
|
||||
message: Dict[str, Any] = {"role": "assistant"}
|
||||
message: dict[str, Any] = {"role": "assistant"}
|
||||
|
||||
content_blocks, tool_blocks = self._split_blocks(internal.content)
|
||||
content_value = self._blocks_to_openai_content(content_blocks)
|
||||
@@ -336,9 +334,9 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
# 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)
|
||||
events: List[InternalStreamEvent] = []
|
||||
events: list[InternalStreamEvent] = []
|
||||
|
||||
# OpenAI streaming error(通常是单个 {"error": {...}})
|
||||
if isinstance(chunk, dict) and "error" in chunk:
|
||||
@@ -437,11 +435,11 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
self,
|
||||
event: InternalStreamEvent,
|
||||
state: StreamState,
|
||||
) -> List[Dict[str, Any]]:
|
||||
) -> list[dict[str, Any]]:
|
||||
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 {
|
||||
"id": state.message_id or "chatcmpl-stream",
|
||||
"object": "chat.completion.chunk",
|
||||
@@ -570,10 +568,10 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
# 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
|
||||
|
||||
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 = err if isinstance(err, dict) else {}
|
||||
|
||||
@@ -592,9 +590,9 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
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")
|
||||
payload: Dict[str, Any] = {
|
||||
payload: dict[str, Any] = {
|
||||
"message": internal.message,
|
||||
"type": type_str,
|
||||
}
|
||||
@@ -608,8 +606,8 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
# Helpers
|
||||
# =========================
|
||||
|
||||
def _openai_message_to_internal(self, msg: Dict[str, Any]) -> Tuple[Optional[InternalMessage], Dict[str, int]]:
|
||||
dropped: Dict[str, int] = {}
|
||||
def _openai_message_to_internal(self, msg: dict[str, Any]) -> tuple[InternalMessage | None, dict[str, int]]:
|
||||
dropped: dict[str, int] = {}
|
||||
|
||||
role_raw = str(msg.get("role") or "unknown")
|
||||
role = self._role_from_openai(role_raw)
|
||||
@@ -656,8 +654,8 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
dropped,
|
||||
)
|
||||
|
||||
def _openai_content_to_blocks(self, content: Any) -> Tuple[List[ContentBlock], Dict[str, int]]:
|
||||
dropped: Dict[str, int] = {}
|
||||
def _openai_content_to_blocks(self, content: Any) -> tuple[list[ContentBlock], dict[str, int]]:
|
||||
dropped: dict[str, int] = {}
|
||||
|
||||
if content is None:
|
||||
return [], dropped
|
||||
@@ -667,7 +665,7 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
dropped["openai_content_non_list"] = dropped.get("openai_content_non_list", 0) + 1
|
||||
return [], dropped
|
||||
|
||||
blocks: List[ContentBlock] = []
|
||||
blocks: list[ContentBlock] = []
|
||||
for part in content:
|
||||
if not isinstance(part, dict):
|
||||
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
|
||||
|
||||
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)
|
||||
text_parts = [b.text for b in blocks if isinstance(b, TextBlock) and b.text]
|
||||
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]
|
||||
joined = "\n\n".join(parts)
|
||||
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):
|
||||
return None
|
||||
|
||||
out: List[ToolDefinition] = []
|
||||
out: list[ToolDefinition] = []
|
||||
for tool in tools:
|
||||
if not isinstance(tool, dict):
|
||||
continue
|
||||
@@ -720,7 +718,7 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
continue
|
||||
|
||||
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 "")
|
||||
if not name:
|
||||
continue
|
||||
@@ -740,7 +738,7 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
|
||||
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:
|
||||
return None
|
||||
|
||||
@@ -756,13 +754,13 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
|
||||
if isinstance(tool_choice, dict) and tool_choice.get("type") == "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 "")
|
||||
return ToolChoice(type=ToolChoiceType.TOOL, tool_name=name, 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:
|
||||
return "none"
|
||||
if tool_choice.type == ToolChoiceType.AUTO:
|
||||
@@ -773,8 +771,8 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
return {"type": "function", "function": {"name": tool_choice.tool_name or ""}}
|
||||
return "auto"
|
||||
|
||||
def _openai_tool_call_to_block(self, tool_call: Any) -> Tuple[Optional[ToolUseBlock], Dict[str, int]]:
|
||||
dropped: Dict[str, int] = {}
|
||||
def _openai_tool_call_to_block(self, tool_call: Any) -> tuple[ToolUseBlock | None, dict[str, int]]:
|
||||
dropped: dict[str, int] = {}
|
||||
if not isinstance(tool_call, dict):
|
||||
dropped["openai_tool_call_non_dict"] = dropped.get("openai_tool_call_non_dict", 0) + 1
|
||||
return None, dropped
|
||||
@@ -785,12 +783,12 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
return None, dropped
|
||||
|
||||
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 "")
|
||||
args_str = str(fn.get("arguments") or "")
|
||||
tool_id = str(tool_call.get("id") or "")
|
||||
|
||||
tool_input: Dict[str, Any]
|
||||
tool_input: dict[str, Any]
|
||||
if args_str:
|
||||
try:
|
||||
parsed = json.loads(args_str)
|
||||
@@ -810,15 +808,15 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
dropped,
|
||||
)
|
||||
|
||||
def _legacy_function_call_to_block(self, func_call: Dict[str, Any]) -> Tuple[Optional[ToolUseBlock], Dict[str, int]]:
|
||||
dropped: 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] = {}
|
||||
name = str(func_call.get("name") or "")
|
||||
args_str = str(func_call.get("arguments") or "")
|
||||
if not name:
|
||||
dropped["openai_function_call_missing_name"] = dropped.get("openai_function_call_missing_name", 0) + 1
|
||||
return None, dropped
|
||||
|
||||
tool_input: Dict[str, Any]
|
||||
tool_input: dict[str, Any]
|
||||
if args_str:
|
||||
try:
|
||||
parsed = json.loads(args_str)
|
||||
@@ -840,10 +838,10 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
|
||||
def _openai_tool_result_message_to_block(
|
||||
self,
|
||||
msg: Dict[str, Any],
|
||||
msg: dict[str, Any],
|
||||
tool_call_id: str,
|
||||
) -> Tuple[Optional[ToolResultBlock], Dict[str, int]]:
|
||||
dropped: Dict[str, int] = {}
|
||||
) -> tuple[ToolResultBlock | None, dict[str, int]]:
|
||||
dropped: dict[str, int] = {}
|
||||
content = msg.get("content")
|
||||
if content is None:
|
||||
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,
|
||||
)
|
||||
|
||||
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):
|
||||
return None
|
||||
|
||||
@@ -908,9 +906,9 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
extra={"openai": extra} if extra else {},
|
||||
)
|
||||
|
||||
def _blocks_to_openai_content(self, blocks: List[ContentBlock]) -> Optional[Union[str, List[Dict[str, Any]]]]:
|
||||
parts: List[Dict[str, Any]] = []
|
||||
text_parts: List[str] = []
|
||||
def _blocks_to_openai_content(self, blocks: list[ContentBlock]) -> str | list[dict[str, Any]] | None:
|
||||
parts: list[dict[str, Any]] = []
|
||||
text_parts: list[str] = []
|
||||
|
||||
for b in blocks:
|
||||
if isinstance(b, TextBlock):
|
||||
@@ -948,9 +946,9 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
# OpenAI content 可以是空字符串;但作为响应 message.content 通常允许为 ""/None。
|
||||
return ""
|
||||
|
||||
def _split_blocks(self, blocks: List[ContentBlock]) -> Tuple[List[ContentBlock], List[ToolUseBlock]]:
|
||||
content_blocks: List[ContentBlock] = []
|
||||
tool_blocks: List[ToolUseBlock] = []
|
||||
def _split_blocks(self, blocks: list[ContentBlock]) -> tuple[list[ContentBlock], list[ToolUseBlock]]:
|
||||
content_blocks: list[ContentBlock] = []
|
||||
tool_blocks: list[ToolUseBlock] = []
|
||||
for b in blocks:
|
||||
if isinstance(b, ToolUseBlock):
|
||||
tool_blocks.append(b)
|
||||
@@ -963,7 +961,7 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
content_blocks.append(b)
|
||||
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:
|
||||
return self._user_message_to_openai(msg)
|
||||
if msg.role == Role.ASSISTANT:
|
||||
@@ -974,9 +972,9 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
return [{"role": "tool", "content": content_value or ""}]
|
||||
return [{"role": "user", "content": ""}]
|
||||
|
||||
def _user_message_to_openai(self, msg: InternalMessage) -> List[Dict[str, Any]]:
|
||||
out: List[Dict[str, Any]] = []
|
||||
pending: List[ContentBlock] = []
|
||||
def _user_message_to_openai(self, msg: InternalMessage) -> list[dict[str, Any]]:
|
||||
out: list[dict[str, Any]] = []
|
||||
pending: list[ContentBlock] = []
|
||||
|
||||
def flush_user() -> None:
|
||||
nonlocal pending
|
||||
@@ -1008,9 +1006,9 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
|
||||
return out
|
||||
|
||||
def _assistant_message_to_openai(self, msg: InternalMessage) -> Dict[str, Any]:
|
||||
content_blocks: List[ContentBlock] = []
|
||||
tool_blocks: List[ToolUseBlock] = []
|
||||
def _assistant_message_to_openai(self, msg: InternalMessage) -> dict[str, Any]:
|
||||
content_blocks: list[ContentBlock] = []
|
||||
tool_blocks: list[ToolUseBlock] = []
|
||||
|
||||
for b in msg.content:
|
||||
if isinstance(b, ToolUseBlock):
|
||||
@@ -1022,7 +1020,7 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
continue
|
||||
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)
|
||||
out["content"] = content_value if content_value is not None else ""
|
||||
|
||||
@@ -1031,7 +1029,7 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
|
||||
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
|
||||
if block.content_text is not None:
|
||||
content = block.content_text
|
||||
@@ -1048,7 +1046,7 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
"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 {
|
||||
"index": index,
|
||||
"id": block.tool_id or f"call_{index}",
|
||||
@@ -1085,7 +1083,7 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
except ValueError:
|
||||
return ErrorType.UNKNOWN
|
||||
|
||||
def _optional_int(self, value: Any) -> Optional[int]:
|
||||
def _optional_int(self, value: Any) -> int | None:
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
@@ -1093,7 +1091,7 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
|
||||
def _optional_float(self, value: Any) -> Optional[float]:
|
||||
def _optional_float(self, value: Any) -> float | None:
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
@@ -1101,7 +1099,7 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
except (TypeError, ValueError):
|
||||
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:
|
||||
return None
|
||||
if isinstance(value, str):
|
||||
@@ -1110,14 +1108,14 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
return [str(x) for x in value if x is not 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}
|
||||
|
||||
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():
|
||||
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")
|
||||
if not isinstance(mapping, dict):
|
||||
mapping = {}
|
||||
@@ -1131,7 +1129,7 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
ss["next_block_index"] = next_idx + 1
|
||||
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")
|
||||
if not isinstance(mapping, dict):
|
||||
mapping = {}
|
||||
|
||||
@@ -10,11 +10,10 @@ OpenAI CLI / Responses Normalizer (OPENAI_CLI)
|
||||
- 未识别的字段会进入 extra/raw,未知内容块保留在 internal,但默认输出阶段会丢弃。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import time
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
from typing import Any
|
||||
|
||||
from src.core.api_format.conversion.field_mappings import (
|
||||
ERROR_TYPE_MAPPINGS,
|
||||
@@ -65,7 +64,7 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
supports_images=True,
|
||||
)
|
||||
|
||||
_ERROR_TYPE_TO_OPENAI: Dict[ErrorType, str] = {
|
||||
_ERROR_TYPE_TO_OPENAI: dict[ErrorType, str] = {
|
||||
ErrorType.INVALID_REQUEST: "invalid_request_error",
|
||||
ErrorType.AUTHENTICATION: "invalid_api_key",
|
||||
ErrorType.PERMISSION_DENIED: "invalid_request_error",
|
||||
@@ -82,12 +81,12 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
# 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 "")
|
||||
|
||||
instructions_text = request.get("instructions")
|
||||
instructions: List[InstructionSegment] = []
|
||||
system_text: Optional[str] = None
|
||||
instructions: list[InstructionSegment] = []
|
||||
system_text: str | None = None
|
||||
if isinstance(instructions_text, str) and instructions_text.strip():
|
||||
system_text = instructions_text
|
||||
instructions.append(InstructionSegment(role=Role.SYSTEM, text=instructions_text))
|
||||
@@ -118,8 +117,8 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
|
||||
return internal
|
||||
|
||||
def request_from_internal(self, internal: InternalRequest) -> Dict[str, Any]:
|
||||
result: Dict[str, Any] = {
|
||||
def request_from_internal(self, internal: InternalRequest) -> dict[str, Any]:
|
||||
result: dict[str, Any] = {
|
||||
"model": internal.model,
|
||||
"input": self._internal_messages_to_input(internal.messages),
|
||||
}
|
||||
@@ -164,7 +163,7 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
# 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)
|
||||
|
||||
rid = str(payload.get("id") or "")
|
||||
@@ -191,8 +190,8 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
self,
|
||||
internal: InternalResponse,
|
||||
*,
|
||||
requested_model: Optional[str] = None,
|
||||
) -> Dict[str, Any]:
|
||||
requested_model: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
text = self._collapse_internal_text(internal.content)
|
||||
|
||||
output_message = {
|
||||
@@ -203,7 +202,7 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
}
|
||||
|
||||
usage = internal.usage or UsageInfo()
|
||||
usage_obj: Dict[str, Any] = {
|
||||
usage_obj: dict[str, Any] = {
|
||||
"input_tokens": usage.input_tokens,
|
||||
"output_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(
|
||||
self,
|
||||
chunk: Dict[str, Any],
|
||||
chunk: dict[str, Any],
|
||||
state: StreamState,
|
||||
) -> List[InternalStreamEvent]:
|
||||
) -> list[InternalStreamEvent]:
|
||||
ss = state.substate(self.FORMAT_ID)
|
||||
events: List[InternalStreamEvent] = []
|
||||
events: list[InternalStreamEvent] = []
|
||||
|
||||
# 统一错误结构(最佳努力)
|
||||
if isinstance(chunk, dict) and "error" in chunk:
|
||||
@@ -392,11 +391,11 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
self,
|
||||
event: InternalStreamEvent,
|
||||
state: StreamState,
|
||||
) -> List[Dict[str, Any]]:
|
||||
) -> list[dict[str, Any]]:
|
||||
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 字段;这里强制保证
|
||||
return payload
|
||||
|
||||
@@ -463,10 +462,10 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
# 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
|
||||
|
||||
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 = err if isinstance(err, dict) else {}
|
||||
|
||||
@@ -484,9 +483,9 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
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")
|
||||
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:
|
||||
payload["code"] = internal.code
|
||||
if internal.param is not None:
|
||||
@@ -497,7 +496,7 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
# 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):
|
||||
return {}
|
||||
resp_inner = response.get("response")
|
||||
@@ -506,8 +505,8 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
return resp_inner
|
||||
return response
|
||||
|
||||
def _extract_output_text_blocks(self, payload: Dict[str, Any]) -> Tuple[List[ContentBlock], Dict[str, Any]]:
|
||||
text_parts: List[str] = []
|
||||
def _extract_output_text_blocks(self, payload: dict[str, Any]) -> tuple[list[ContentBlock], dict[str, Any]]:
|
||||
text_parts: list[str] = []
|
||||
|
||||
output = payload.get("output")
|
||||
if isinstance(output, list):
|
||||
@@ -532,12 +531,12 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
if not text_parts and isinstance(payload.get("output_text"), str):
|
||||
text_parts.append(payload.get("output_text") or "")
|
||||
|
||||
blocks: List[ContentBlock] = []
|
||||
blocks: list[ContentBlock] = []
|
||||
text = "".join(text_parts)
|
||||
if 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
|
||||
|
||||
def _usage_to_internal(self, usage: Any) -> UsageInfo:
|
||||
@@ -553,14 +552,14 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
extra={"openai_cli": {"usage": usage}},
|
||||
)
|
||||
|
||||
def _collapse_internal_text(self, blocks: List[ContentBlock]) -> str:
|
||||
parts: List[str] = []
|
||||
def _collapse_internal_text(self, blocks: list[ContentBlock]) -> str:
|
||||
parts: list[str] = []
|
||||
for block in blocks:
|
||||
if isinstance(block, TextBlock) and block.text:
|
||||
parts.append(block.text)
|
||||
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:
|
||||
return []
|
||||
|
||||
@@ -575,7 +574,7 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
if not isinstance(input_data, list):
|
||||
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:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
@@ -624,7 +623,7 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
|
||||
# reasoning -> assistant 消息,提取 summary 作为文本
|
||||
if item_type == "reasoning":
|
||||
summary_parts: List[str] = []
|
||||
summary_parts: list[str] = []
|
||||
summary = item.get("summary")
|
||||
if isinstance(summary, list):
|
||||
for s in summary:
|
||||
@@ -638,7 +637,7 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
summary_parts.append(summary)
|
||||
|
||||
# 如果有 summary 文本,创建一个 UnknownBlock 保留原始结构
|
||||
reasoning_blocks: List[ContentBlock] = []
|
||||
reasoning_blocks: list[ContentBlock] = []
|
||||
if summary_parts:
|
||||
# 保留 reasoning 的 summary 作为 UnknownBlock,便于输出时决策
|
||||
reasoning_blocks.append(UnknownBlock(
|
||||
@@ -662,7 +661,7 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
|
||||
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:
|
||||
return []
|
||||
if isinstance(content, str):
|
||||
@@ -673,7 +672,7 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
if not isinstance(content, list):
|
||||
return [UnknownBlock(raw_type="content", payload={"content": content})]
|
||||
|
||||
blocks: List[ContentBlock] = []
|
||||
blocks: list[ContentBlock] = []
|
||||
for part in content:
|
||||
if isinstance(part, str):
|
||||
if part:
|
||||
@@ -690,8 +689,8 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
blocks.append(UnknownBlock(raw_type=ptype or "unknown", payload=part))
|
||||
return blocks
|
||||
|
||||
def _internal_messages_to_input(self, messages: List[InternalMessage]) -> List[Dict[str, Any]]:
|
||||
out: List[Dict[str, Any]] = []
|
||||
def _internal_messages_to_input(self, messages: list[InternalMessage]) -> list[dict[str, Any]]:
|
||||
out: list[dict[str, Any]] = []
|
||||
for msg in messages:
|
||||
# ToolUseBlock -> function_call
|
||||
for block in msg.content:
|
||||
@@ -729,7 +728,7 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
|
||||
# 普通 message(TextBlock)
|
||||
role = self._role_to_openai(msg.role)
|
||||
content_items: List[Dict[str, Any]] = []
|
||||
content_items: list[dict[str, Any]] = []
|
||||
has_text = False
|
||||
|
||||
for block in msg.content:
|
||||
@@ -748,10 +747,10 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
|
||||
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):
|
||||
return None
|
||||
out: List[ToolDefinition] = []
|
||||
out: list[ToolDefinition] = []
|
||||
for tool in tools:
|
||||
if not isinstance(tool, dict):
|
||||
continue
|
||||
@@ -783,7 +782,7 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
)
|
||||
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:
|
||||
return None
|
||||
if isinstance(tool_choice, str):
|
||||
@@ -802,7 +801,7 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
|
||||
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:
|
||||
return "none"
|
||||
if tool_choice.type == ToolChoiceType.AUTO:
|
||||
@@ -840,7 +839,7 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
return "tool"
|
||||
return "user"
|
||||
|
||||
def _optional_int(self, value: Any) -> Optional[int]:
|
||||
def _optional_int(self, value: Any) -> int | None:
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
@@ -848,7 +847,7 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
|
||||
def _optional_float(self, value: Any) -> Optional[float]:
|
||||
def _optional_float(self, value: Any) -> float | None:
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
@@ -856,13 +855,13 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
except (TypeError, ValueError):
|
||||
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:
|
||||
return None
|
||||
if isinstance(value, str):
|
||||
return [value]
|
||||
if isinstance(value, list):
|
||||
out: List[str] = []
|
||||
out: list[str] = []
|
||||
for item in value:
|
||||
if item is None:
|
||||
continue
|
||||
@@ -870,14 +869,14 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
return out
|
||||
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):
|
||||
return {}
|
||||
return {k: v for k, v in payload.items() if k not in keep_keys}
|
||||
|
||||
def _join_instructions(self, internal: InternalRequest) -> str:
|
||||
if internal.instructions:
|
||||
parts: List[str] = []
|
||||
parts: list[str] = []
|
||||
for seg in internal.instructions:
|
||||
if seg.text:
|
||||
parts.append(seg.text)
|
||||
|
||||
@@ -9,12 +9,12 @@ source -> internal -> target
|
||||
- 转换失败将抛出 `FormatConversionError`(不再静默回退)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
import time
|
||||
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.metrics import format_conversion_duration_seconds, format_conversion_total
|
||||
@@ -29,7 +29,7 @@ def _track_conversion_metrics(
|
||||
direction: str,
|
||||
source: str,
|
||||
target: str,
|
||||
) -> Generator[None, None, None]:
|
||||
) -> Generator[None]:
|
||||
start = time.perf_counter()
|
||||
try:
|
||||
yield
|
||||
@@ -47,13 +47,13 @@ class FormatConversionRegistry:
|
||||
"""基于 Normalizer 的格式转换注册表"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._normalizers: Dict[str, FormatNormalizer] = {}
|
||||
self._normalizers: dict[str, FormatNormalizer] = {}
|
||||
|
||||
def register(self, normalizer: FormatNormalizer) -> None:
|
||||
self._normalizers[str(normalizer.FORMAT_ID).upper()] = normalizer
|
||||
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())
|
||||
|
||||
def _require_normalizer(self, format_id: str) -> FormatNormalizer:
|
||||
@@ -66,10 +66,10 @@ class FormatConversionRegistry:
|
||||
|
||||
def convert_request(
|
||||
self,
|
||||
request: Dict[str, Any],
|
||||
request: dict[str, Any],
|
||||
source_format: str,
|
||||
target_format: str,
|
||||
) -> Dict[str, Any]:
|
||||
) -> dict[str, Any]:
|
||||
if str(source_format).upper() == str(target_format).upper():
|
||||
return request
|
||||
|
||||
@@ -85,12 +85,12 @@ class FormatConversionRegistry:
|
||||
|
||||
def convert_response(
|
||||
self,
|
||||
response: Dict[str, Any],
|
||||
response: dict[str, Any],
|
||||
source_format: str,
|
||||
target_format: str,
|
||||
*,
|
||||
requested_model: Optional[str] = None,
|
||||
) -> Dict[str, Any]:
|
||||
requested_model: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""转换响应格式
|
||||
|
||||
Args:
|
||||
@@ -124,10 +124,10 @@ class FormatConversionRegistry:
|
||||
|
||||
def convert_error_response(
|
||||
self,
|
||||
error_response: Dict[str, Any],
|
||||
error_response: dict[str, Any],
|
||||
source_format: str,
|
||||
target_format: str,
|
||||
) -> Dict[str, Any]:
|
||||
) -> dict[str, Any]:
|
||||
if str(source_format).upper() == str(target_format).upper():
|
||||
return error_response
|
||||
|
||||
@@ -152,11 +152,11 @@ class FormatConversionRegistry:
|
||||
|
||||
def convert_stream_chunk(
|
||||
self,
|
||||
chunk: Dict[str, Any],
|
||||
chunk: dict[str, Any],
|
||||
source_format: str,
|
||||
target_format: str,
|
||||
state: Optional[StreamState] = None,
|
||||
) -> List[Dict[str, Any]]:
|
||||
state: StreamState | None = None,
|
||||
) -> list[dict[str, Any]]:
|
||||
if str(source_format).upper() == str(target_format).upper():
|
||||
return [chunk]
|
||||
|
||||
@@ -182,7 +182,7 @@ class FormatConversionRegistry:
|
||||
with _track_conversion_metrics("stream", str(source_format).upper(), str(target_format).upper()):
|
||||
try:
|
||||
events = src.stream_chunk_to_internal(chunk, state)
|
||||
out: List[Dict[str, Any]] = []
|
||||
out: list[dict[str, Any]] = []
|
||||
for event in events:
|
||||
out.extend(tgt.stream_event_from_internal(event, state))
|
||||
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 True
|
||||
|
||||
def list_normalizers(self) -> List[str]:
|
||||
def list_normalizers(self) -> list[str]:
|
||||
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()
|
||||
if src not in self._normalizers:
|
||||
return []
|
||||
|
||||
@@ -4,11 +4,10 @@
|
||||
用于把 OpenAI/Claude/Gemini 的流式协议映射为统一事件序列,再由目标格式 Normalizer 输出。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from enum import Enum
|
||||
from typing import Any, Dict, Optional, Union
|
||||
from typing import Any
|
||||
|
||||
from .internal import ContentType, InternalError, StopReason, UsageInfo
|
||||
|
||||
@@ -34,8 +33,8 @@ class MessageStartEvent:
|
||||
type: StreamEventType = field(default=StreamEventType.MESSAGE_START, init=False)
|
||||
message_id: str = ""
|
||||
model: str = ""
|
||||
usage: Optional[UsageInfo] = None # Claude 流式响应的 message_start 可能包含 usage
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
usage: UsageInfo | None = None # Claude 流式响应的 message_start 可能包含 usage
|
||||
extra: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -46,9 +45,9 @@ class ContentBlockStartEvent:
|
||||
block_index: int = 0
|
||||
block_type: ContentType = ContentType.TEXT
|
||||
# 工具调用时使用(TOOL_USE block)
|
||||
tool_id: Optional[str] = None
|
||||
tool_name: Optional[str] = None
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
tool_id: str | None = None
|
||||
tool_name: str | None = None
|
||||
extra: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -58,7 +57,7 @@ class ContentDeltaEvent:
|
||||
type: StreamEventType = field(default=StreamEventType.CONTENT_DELTA, init=False)
|
||||
block_index: int = 0
|
||||
text_delta: str = ""
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
extra: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -69,7 +68,7 @@ class ToolCallDeltaEvent:
|
||||
block_index: int = 0
|
||||
tool_id: str = ""
|
||||
input_delta: str = "" # JSON 字符串片段
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
extra: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -78,7 +77,7 @@ class ContentBlockStopEvent:
|
||||
|
||||
type: StreamEventType = field(default=StreamEventType.CONTENT_BLOCK_STOP, init=False)
|
||||
block_index: int = 0
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
extra: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -86,9 +85,9 @@ class MessageStopEvent:
|
||||
"""消息结束事件"""
|
||||
|
||||
type: StreamEventType = field(default=StreamEventType.MESSAGE_STOP, init=False)
|
||||
stop_reason: Optional[StopReason] = None
|
||||
usage: Optional[UsageInfo] = None
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
stop_reason: StopReason | None = None
|
||||
usage: UsageInfo | None = None
|
||||
extra: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -97,7 +96,7 @@ class UsageEvent:
|
||||
|
||||
type: StreamEventType = field(default=StreamEventType.USAGE, init=False)
|
||||
usage: UsageInfo = field(default_factory=UsageInfo)
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
extra: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -106,7 +105,7 @@ class ErrorEvent:
|
||||
|
||||
type: StreamEventType = field(default=StreamEventType.ERROR, init=False)
|
||||
error: InternalError
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
extra: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -115,21 +114,21 @@ class UnknownStreamEvent:
|
||||
|
||||
type: StreamEventType = field(default=StreamEventType.UNKNOWN, init=False)
|
||||
raw_type: str = ""
|
||||
payload: Dict[str, Any] = field(default_factory=dict)
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
payload: dict[str, Any] = field(default_factory=dict)
|
||||
extra: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
InternalStreamEvent = Union[
|
||||
MessageStartEvent,
|
||||
ContentBlockStartEvent,
|
||||
ContentDeltaEvent,
|
||||
ToolCallDeltaEvent,
|
||||
ContentBlockStopEvent,
|
||||
MessageStopEvent,
|
||||
UsageEvent,
|
||||
ErrorEvent,
|
||||
UnknownStreamEvent,
|
||||
]
|
||||
InternalStreamEvent = (
|
||||
MessageStartEvent
|
||||
| ContentBlockStartEvent
|
||||
| ContentDeltaEvent
|
||||
| ToolCallDeltaEvent
|
||||
| ContentBlockStopEvent
|
||||
| MessageStopEvent
|
||||
| UsageEvent
|
||||
| ErrorEvent
|
||||
| UnknownStreamEvent
|
||||
)
|
||||
|
||||
|
||||
__all__ = [
|
||||
|
||||
@@ -5,10 +5,9 @@
|
||||
每个 Normalizer 通过 `substate(format_id)` 获取自己的隔离状态字典。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Dict
|
||||
from typing import Any
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -26,12 +25,12 @@ class StreamState:
|
||||
message_id: str = ""
|
||||
|
||||
# Registry/调用层的通用扩展信息(与具体格式无关)
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
extra: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
# 各 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()
|
||||
return self.by_format.setdefault(key, {})
|
||||
|
||||
@@ -6,7 +6,8 @@ API 格式检测
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Dict, Optional, Tuple
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
if TYPE_CHECKING:
|
||||
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(
|
||||
headers: Dict[str, str],
|
||||
query_params: Optional[Dict[str, str]],
|
||||
headers: dict[str, str],
|
||||
query_params: dict[str, str] | None,
|
||||
definition: ApiFormatDefinition,
|
||||
) -> Tuple[Optional[str], str]:
|
||||
) -> tuple[str | None, str]:
|
||||
"""
|
||||
根据格式定义从请求中提取 API Key
|
||||
|
||||
@@ -64,9 +65,9 @@ def _extract_api_key_by_definition(
|
||||
|
||||
|
||||
def detect_format_from_request(
|
||||
headers: Dict[str, str],
|
||||
query_params: Optional[Dict[str, str]] = None,
|
||||
) -> Tuple[APIFormat, Optional[str], str]:
|
||||
headers: dict[str, str],
|
||||
query_params: dict[str, str] | None = None,
|
||||
) -> tuple[APIFormat, str | None, str]:
|
||||
"""
|
||||
从请求头检测 API 格式和 API Key
|
||||
|
||||
@@ -107,8 +108,8 @@ def detect_format_from_request(
|
||||
|
||||
|
||||
def detect_format_and_key_from_starlette(
|
||||
request: "Request",
|
||||
) -> Tuple[str, Optional[str], str]:
|
||||
request: Request,
|
||||
) -> tuple[str, str | None, str]:
|
||||
"""
|
||||
从 Starlette Request 对象检测 API 格式和 API Key
|
||||
|
||||
@@ -135,7 +136,7 @@ def detect_format_and_key_from_starlette(
|
||||
|
||||
def detect_format_from_response(
|
||||
response_data: dict,
|
||||
) -> Optional[APIFormat]:
|
||||
) -> APIFormat | None:
|
||||
"""
|
||||
从响应内容检测 API 格式
|
||||
|
||||
|
||||
@@ -12,7 +12,8 @@
|
||||
|
||||
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.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 的认证
|
||||
"authorization",
|
||||
@@ -41,7 +42,7 @@ UPSTREAM_DROP_HEADERS: FrozenSet[str] = frozenset(
|
||||
|
||||
# 最小必脱敏集合(编译时常量,用于快速路径)
|
||||
# 完整脱敏应使用 SystemConfigService.get_sensitive_headers()
|
||||
CORE_REDACT_HEADERS: FrozenSet[str] = frozenset(
|
||||
CORE_REDACT_HEADERS: frozenset[str] = frozenset(
|
||||
{
|
||||
"authorization",
|
||||
"x-api-key",
|
||||
@@ -50,7 +51,7 @@ CORE_REDACT_HEADERS: FrozenSet[str] = frozenset(
|
||||
)
|
||||
|
||||
# Hop-by-hop 头部 (RFC 7230)
|
||||
HOP_BY_HOP_HEADERS: FrozenSet[str] = frozenset(
|
||||
HOP_BY_HOP_HEADERS: frozenset[str] = frozenset(
|
||||
{
|
||||
"connection",
|
||||
"keep-alive",
|
||||
@@ -64,7 +65,7 @@ HOP_BY_HOP_HEADERS: FrozenSet[str] = frozenset(
|
||||
)
|
||||
|
||||
# 响应时需要过滤的头部(body-dependent + hop-by-hop)
|
||||
RESPONSE_DROP_HEADERS: FrozenSet[str] = (
|
||||
RESPONSE_DROP_HEADERS: frozenset[str] = (
|
||||
frozenset(
|
||||
{
|
||||
"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 统一为小写
|
||||
|
||||
@@ -92,7 +93,7 @@ def normalize_headers(headers: Dict[str, str]) -> Dict[str, str]:
|
||||
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
|
||||
|
||||
@@ -147,10 +148,10 @@ def extract_client_api_key(headers: Dict[str, str], api_format: APIFormat) -> Op
|
||||
|
||||
|
||||
def extract_client_api_key_with_query(
|
||||
headers: Dict[str, str],
|
||||
query_params: Optional[Dict[str, str]],
|
||||
headers: dict[str, str],
|
||||
query_params: dict[str, str] | None,
|
||||
api_format: APIFormat,
|
||||
) -> Optional[str]:
|
||||
) -> str | None:
|
||||
"""
|
||||
从客户端请求头或 URL 参数提取 API Key
|
||||
|
||||
@@ -184,10 +185,10 @@ def extract_client_api_key_with_query(
|
||||
|
||||
|
||||
def detect_capabilities(
|
||||
headers: Dict[str, str],
|
||||
headers: dict[str, str],
|
||||
api_format: APIFormat,
|
||||
request_body: Optional[Dict[str, Any]] = None, # noqa: ARG001 - 预留给部分格式使用
|
||||
) -> Dict[str, bool]:
|
||||
request_body: dict[str, Any] | None = None, # noqa: ARG001 - 预留给部分格式使用
|
||||
) -> dict[str, bool]:
|
||||
"""
|
||||
从请求头检测能力需求
|
||||
|
||||
@@ -203,7 +204,7 @@ def detect_capabilities(
|
||||
能力需求字典,如 {"context_1m": True}
|
||||
"""
|
||||
|
||||
requirements: Dict[str, bool] = {}
|
||||
requirements: dict[str, bool] = {}
|
||||
|
||||
if api_format in (APIFormat.CLAUDE, APIFormat.CLAUDE_CLI):
|
||||
beta_header = get_header_value(headers, "anthropic-beta")
|
||||
@@ -228,20 +229,20 @@ class HeaderBuilder:
|
||||
|
||||
def __init__(self) -> None:
|
||||
# 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)
|
||||
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():
|
||||
self.add(k, v)
|
||||
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 不被覆盖
|
||||
|
||||
@@ -253,13 +254,13 @@ class HeaderBuilder:
|
||||
self.add(k, v)
|
||||
return self
|
||||
|
||||
def remove(self, keys: FrozenSet[str]) -> "HeaderBuilder":
|
||||
def remove(self, keys: frozenset[str]) -> HeaderBuilder:
|
||||
"""移除指定的头部"""
|
||||
for k in keys:
|
||||
self._headers.pop(k.lower(), None)
|
||||
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(
|
||||
self,
|
||||
rules: list[Dict[str, Any]],
|
||||
protected_keys: Optional[AbstractSet[str]] = None,
|
||||
) -> "HeaderBuilder":
|
||||
rules: list[dict[str, Any]],
|
||||
protected_keys: AbstractSet[str] | None = None,
|
||||
) -> HeaderBuilder:
|
||||
"""
|
||||
应用请求头规则
|
||||
|
||||
@@ -314,20 +315,20 @@ class HeaderBuilder:
|
||||
|
||||
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()}
|
||||
|
||||
|
||||
def build_upstream_headers(
|
||||
original_headers: Dict[str, str],
|
||||
original_headers: dict[str, str],
|
||||
api_format: APIFormat,
|
||||
provider_api_key: str,
|
||||
*,
|
||||
endpoint_headers: Optional[Dict[str, str]] = None,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
drop_headers: Optional[FrozenSet[str]] = None,
|
||||
) -> Dict[str, str]:
|
||||
endpoint_headers: dict[str, str] | None = None,
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
drop_headers: frozenset[str] | None = None,
|
||||
) -> dict[str, str]:
|
||||
"""
|
||||
构建发送给上游 Provider 的请求头
|
||||
|
||||
@@ -386,10 +387,10 @@ def build_upstream_headers(
|
||||
|
||||
|
||||
def merge_headers_with_protection(
|
||||
base_headers: Dict[str, str],
|
||||
extra_headers: Optional[Dict[str, str]],
|
||||
protected_keys: FrozenSet[str] | Set[str],
|
||||
) -> Dict[str, str]:
|
||||
base_headers: dict[str, str],
|
||||
extra_headers: dict[str, str] | None,
|
||||
protected_keys: frozenset[str] | set[str],
|
||||
) -> dict[str, str]:
|
||||
"""
|
||||
合并头部但保护指定的 key 不被覆盖
|
||||
|
||||
@@ -418,9 +419,9 @@ def merge_headers_with_protection(
|
||||
|
||||
|
||||
def filter_response_headers(
|
||||
headers: Optional[Dict[str, str]],
|
||||
drop_headers: Optional[FrozenSet[str]] = None,
|
||||
) -> Dict[str, str]:
|
||||
headers: dict[str, str] | None,
|
||||
drop_headers: frozenset[str] | None = None,
|
||||
) -> dict[str, str]:
|
||||
"""
|
||||
过滤上游响应头中不应透传给客户端的字段
|
||||
|
||||
@@ -446,9 +447,9 @@ def filter_response_headers(
|
||||
|
||||
|
||||
def redact_headers_for_log(
|
||||
headers: Dict[str, str],
|
||||
redact_keys: Optional[FrozenSet[str]] = None,
|
||||
) -> Dict[str, str]:
|
||||
headers: dict[str, str],
|
||||
redact_keys: frozenset[str] | None = None,
|
||||
) -> dict[str, str]:
|
||||
"""
|
||||
将敏感头部值替换为 *** 用于日志记录
|
||||
|
||||
@@ -487,7 +488,7 @@ def build_adapter_base_headers(
|
||||
api_key: str,
|
||||
*,
|
||||
include_extra: bool = True,
|
||||
) -> Dict[str, str]:
|
||||
) -> dict[str, str]:
|
||||
"""
|
||||
根据 API 格式构建基础请求头
|
||||
|
||||
@@ -504,7 +505,7 @@ def build_adapter_base_headers(
|
||||
auth_header, auth_type = get_auth_config(api_format)
|
||||
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,
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
@@ -520,8 +521,8 @@ def build_adapter_base_headers(
|
||||
def build_adapter_headers(
|
||||
api_format: APIFormat,
|
||||
api_key: str,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
) -> Dict[str, str]:
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
) -> dict[str, str]:
|
||||
"""
|
||||
构建完整的 Adapter 请求头
|
||||
|
||||
@@ -565,8 +566,8 @@ def get_adapter_protected_keys(api_format: APIFormat) -> tuple[str, ...]:
|
||||
|
||||
|
||||
def extract_set_headers_from_rules(
|
||||
header_rules: Optional[list[Dict[str, Any]]],
|
||||
) -> Optional[Dict[str, str]]:
|
||||
header_rules: list[dict[str, Any]] | None,
|
||||
) -> dict[str, str] | None:
|
||||
"""
|
||||
从 header_rules 中提取 set 操作生成的头部字典
|
||||
|
||||
@@ -582,7 +583,7 @@ def extract_set_headers_from_rules(
|
||||
if not header_rules:
|
||||
return None
|
||||
|
||||
headers: Dict[str, str] = {}
|
||||
headers: dict[str, str] = {}
|
||||
for rule in header_rules:
|
||||
if rule.get("action") == "set":
|
||||
key = rule.get("key", "")
|
||||
@@ -593,7 +594,7 @@ def extract_set_headers_from_rules(
|
||||
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 提取额外请求头
|
||||
|
||||
|
||||
@@ -13,13 +13,12 @@ API 格式元数据定义
|
||||
definition = get_api_format_definition(APIFormat.CLAUDE)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from dataclasses import dataclass, field
|
||||
from functools import lru_cache
|
||||
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
|
||||
|
||||
@@ -64,7 +63,7 @@ class ApiFormatDefinition:
|
||||
yield normalized
|
||||
|
||||
|
||||
_DEFINITIONS: Dict[APIFormat, ApiFormatDefinition] = {
|
||||
_DEFINITIONS: dict[APIFormat, ApiFormatDefinition] = {
|
||||
APIFormat.CLAUDE: ApiFormatDefinition(
|
||||
api_format=APIFormat.CLAUDE,
|
||||
aliases=("claude", "anthropic", "claude_compatible"),
|
||||
@@ -151,12 +150,12 @@ def get_api_format_definition(api_format: APIFormat) -> ApiFormatDefinition:
|
||||
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())
|
||||
|
||||
|
||||
def build_alias_lookup() -> Dict[str, APIFormat]:
|
||||
def build_alias_lookup() -> dict[str, APIFormat]:
|
||||
"""
|
||||
构建 alias -> APIFormat 的查找表。
|
||||
每次调用都会返回新的 dict,避免可变全局引发并发问题。
|
||||
@@ -237,7 +236,7 @@ def get_protected_keys(api_format: APIFormat) -> frozenset[str]:
|
||||
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()
|
||||
|
||||
|
||||
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)
|
||||
def _alias_lookup_cache() -> Dict[str, APIFormat]:
|
||||
def _alias_lookup_cache() -> dict[str, APIFormat]:
|
||||
"""缓存 alias -> APIFormat 查找表,减少重复构建。"""
|
||||
return build_alias_lookup()
|
||||
|
||||
|
||||
def resolve_api_format_alias(value: str) -> Optional[APIFormat]:
|
||||
def resolve_api_format_alias(value: str) -> APIFormat | None:
|
||||
"""根据别名查找 APIFormat,找不到时返回 None。"""
|
||||
if not value:
|
||||
return None
|
||||
@@ -310,9 +309,9 @@ def resolve_api_format_alias(value: str) -> Optional[APIFormat]:
|
||||
|
||||
|
||||
def resolve_api_format(
|
||||
value: Union[str, APIFormat, None],
|
||||
default: Optional[APIFormat] = None,
|
||||
) -> Optional[APIFormat]:
|
||||
value: str | APIFormat | None,
|
||||
default: APIFormat | None = None,
|
||||
) -> APIFormat | None:
|
||||
"""
|
||||
将任意字符串/枚举值解析为 APIFormat。
|
||||
|
||||
|
||||
@@ -6,13 +6,13 @@ API 格式工具函数
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Optional, Union
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
if TYPE_CHECKING:
|
||||
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 透传格式
|
||||
|
||||
@@ -40,7 +40,7 @@ def is_cli_format(format_id: Union[str, "APIFormat", None]) -> bool:
|
||||
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 后缀)
|
||||
|
||||
@@ -66,7 +66,7 @@ def get_base_format(format_id: Union[str, "APIFormat", None]) -> Optional[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(
|
||||
format1: Union[str, "APIFormat", None],
|
||||
format2: Union[str, "APIFormat", None],
|
||||
format1: str | APIFormat | None,
|
||||
format2: str | APIFormat | None,
|
||||
) -> bool:
|
||||
"""
|
||||
判断两个格式是否相同
|
||||
@@ -95,7 +95,7 @@ def is_same_format(
|
||||
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:
|
||||
"""
|
||||
判断是否为可转换格式
|
||||
|
||||
|
||||
@@ -8,7 +8,6 @@
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from typing import Set
|
||||
|
||||
from src.core.logger import logger
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -23,7 +22,7 @@ class BatchCommitter:
|
||||
interval_seconds: 批量提交间隔(秒)
|
||||
"""
|
||||
self.interval_seconds = interval_seconds
|
||||
self._pending_sessions: Set[Session] = set()
|
||||
self._pending_sessions: set[Session] = set()
|
||||
self._lock = asyncio.Lock()
|
||||
self._task = None
|
||||
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user