mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 09:20:22 +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/
|
COPY frontend/ ./frontend/
|
||||||
RUN cd frontend && npm run build
|
RUN cd frontend && npm run build
|
||||||
# ==================== 运行时镜像 ====================
|
# ==================== 运行时镜像 ====================
|
||||||
FROM python:3.12-slim
|
FROM python:3.14-slim
|
||||||
WORKDIR /app
|
WORKDIR /app
|
||||||
# 运行时依赖(无 gcc/nodejs/npm)
|
# 运行时依赖(无 gcc/nodejs/npm)
|
||||||
RUN apt-get update && apt-get install -y \
|
RUN apt-get update && apt-get install -y \
|
||||||
@@ -17,7 +17,7 @@ RUN apt-get update && apt-get install -y \
|
|||||||
curl \
|
curl \
|
||||||
&& rm -rf /var/lib/apt/lists/*
|
&& rm -rf /var/lib/apt/lists/*
|
||||||
# 从 base 镜像复制 Python 包
|
# 从 base 镜像复制 Python 包
|
||||||
COPY --from=builder /usr/local/lib/python3.12/site-packages /usr/local/lib/python3.12/site-packages
|
COPY --from=builder /usr/local/lib/python3.14/site-packages /usr/local/lib/python3.14/site-packages
|
||||||
# 只复制需要的 Python 可执行文件
|
# 只复制需要的 Python 可执行文件
|
||||||
COPY --from=builder /usr/local/bin/gunicorn /usr/local/bin/
|
COPY --from=builder /usr/local/bin/gunicorn /usr/local/bin/
|
||||||
COPY --from=builder /usr/local/bin/uvicorn /usr/local/bin/
|
COPY --from=builder /usr/local/bin/uvicorn /usr/local/bin/
|
||||||
|
|||||||
@@ -10,7 +10,7 @@ COPY frontend/ ./frontend/
|
|||||||
RUN cd frontend && npm run build
|
RUN cd frontend && npm run build
|
||||||
|
|
||||||
# ==================== 运行时镜像 ====================
|
# ==================== 运行时镜像 ====================
|
||||||
FROM python:3.12-slim
|
FROM python:3.14-slim
|
||||||
|
|
||||||
WORKDIR /app
|
WORKDIR /app
|
||||||
|
|
||||||
@@ -24,7 +24,7 @@ RUN sed -i 's/deb.debian.org/mirrors.tuna.tsinghua.edu.cn/g' /etc/apt/sources.li
|
|||||||
&& rm -rf /var/lib/apt/lists/*
|
&& rm -rf /var/lib/apt/lists/*
|
||||||
|
|
||||||
# 从 base 镜像复制 Python 包
|
# 从 base 镜像复制 Python 包
|
||||||
COPY --from=builder /usr/local/lib/python3.12/site-packages /usr/local/lib/python3.12/site-packages
|
COPY --from=builder /usr/local/lib/python3.14/site-packages /usr/local/lib/python3.14/site-packages
|
||||||
|
|
||||||
# 只复制需要的 Python 可执行文件
|
# 只复制需要的 Python 可执行文件
|
||||||
COPY --from=builder /usr/local/bin/gunicorn /usr/local/bin/
|
COPY --from=builder /usr/local/bin/gunicorn /usr/local/bin/
|
||||||
|
|||||||
@@ -2,7 +2,7 @@
|
|||||||
# 用于 GitHub Actions CI 构建(不使用国内镜像源)
|
# 用于 GitHub Actions CI 构建(不使用国内镜像源)
|
||||||
# 构建命令: docker build -f Dockerfile.base -t aether-base:latest .
|
# 构建命令: docker build -f Dockerfile.base -t aether-base:latest .
|
||||||
# 只在 pyproject.toml 或 frontend/package*.json 变化时需要重建
|
# 只在 pyproject.toml 或 frontend/package*.json 变化时需要重建
|
||||||
FROM python:3.12-slim
|
FROM python:3.14-slim
|
||||||
|
|
||||||
WORKDIR /app
|
WORKDIR /app
|
||||||
|
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
# 构建镜像:编译环境 + 预编译的依赖(国内镜像源版本)
|
# 构建镜像:编译环境 + 预编译的依赖(国内镜像源版本)
|
||||||
# 构建命令: docker build -f Dockerfile.base.local -t aether-base:latest .
|
# 构建命令: docker build -f Dockerfile.base.local -t aether-base:latest .
|
||||||
# 只在 pyproject.toml 或 frontend/package*.json 变化时需要重建
|
# 只在 pyproject.toml 或 frontend/package*.json 变化时需要重建
|
||||||
FROM python:3.12-slim
|
FROM python:3.14-slim
|
||||||
|
|
||||||
WORKDIR /app
|
WORKDIR /app
|
||||||
|
|
||||||
|
|||||||
@@ -15,13 +15,11 @@ classifiers = [
|
|||||||
"Intended Audience :: Developers",
|
"Intended Audience :: Developers",
|
||||||
"License :: Other/Proprietary License",
|
"License :: Other/Proprietary License",
|
||||||
"Programming Language :: Python :: 3",
|
"Programming Language :: Python :: 3",
|
||||||
"Programming Language :: Python :: 3.8",
|
|
||||||
"Programming Language :: Python :: 3.9",
|
|
||||||
"Programming Language :: Python :: 3.10",
|
|
||||||
"Programming Language :: Python :: 3.11",
|
|
||||||
"Programming Language :: Python :: 3.12",
|
"Programming Language :: Python :: 3.12",
|
||||||
|
"Programming Language :: Python :: 3.13",
|
||||||
|
"Programming Language :: Python :: 3.14",
|
||||||
]
|
]
|
||||||
requires-python = ">=3.9"
|
requires-python = ">=3.12"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"fastapi[standard]>=0.115.11",
|
"fastapi[standard]>=0.115.11",
|
||||||
"uvicorn>=0.34.0",
|
"uvicorn>=0.34.0",
|
||||||
@@ -83,7 +81,7 @@ dev-dependencies = [
|
|||||||
|
|
||||||
[tool.black]
|
[tool.black]
|
||||||
line-length = 100
|
line-length = 100
|
||||||
target-version = ['py38']
|
target-version = ['py312']
|
||||||
|
|
||||||
[tool.isort]
|
[tool.isort]
|
||||||
profile = "black"
|
profile = "black"
|
||||||
@@ -99,7 +97,7 @@ source = "vcs"
|
|||||||
version-file = "src/_version.py"
|
version-file = "src/_version.py"
|
||||||
|
|
||||||
[tool.mypy]
|
[tool.mypy]
|
||||||
python_version = "3.9"
|
python_version = "3.12"
|
||||||
warn_return_any = true
|
warn_return_any = true
|
||||||
warn_unused_configs = true
|
warn_unused_configs = true
|
||||||
disallow_untyped_defs = true
|
disallow_untyped_defs = true
|
||||||
|
|||||||
@@ -10,7 +10,6 @@
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import List, Optional
|
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||||
from pydantic import BaseModel, Field, ValidationError
|
from pydantic import BaseModel, Field, ValidationError
|
||||||
@@ -34,7 +33,7 @@ class EnableAdaptiveRequest(BaseModel):
|
|||||||
"""启用自适应模式请求"""
|
"""启用自适应模式请求"""
|
||||||
|
|
||||||
enabled: bool = Field(..., description="是否启用自适应模式(true=自适应,false=固定限制)")
|
enabled: bool = Field(..., description="是否启用自适应模式(true=自适应,false=固定限制)")
|
||||||
fixed_limit: Optional[int] = Field(
|
fixed_limit: int | None = Field(
|
||||||
None, ge=1, le=100, description="固定 RPM 限制(仅当 enabled=false 时生效,1-100)"
|
None, ge=1, le=100, description="固定 RPM 限制(仅当 enabled=false 时生效,1-100)"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -43,30 +42,30 @@ class AdaptiveStatsResponse(BaseModel):
|
|||||||
"""自适应统计响应"""
|
"""自适应统计响应"""
|
||||||
|
|
||||||
adaptive_mode: bool = Field(..., description="是否为自适应模式(rpm_limit=NULL)")
|
adaptive_mode: bool = Field(..., description="是否为自适应模式(rpm_limit=NULL)")
|
||||||
rpm_limit: Optional[int] = Field(None, description="用户配置的固定限制(NULL=自适应)")
|
rpm_limit: int | None = Field(None, description="用户配置的固定限制(NULL=自适应)")
|
||||||
effective_limit: Optional[int] = Field(
|
effective_limit: int | None = Field(
|
||||||
None, description="当前有效限制(自适应使用学习值,固定使用配置值)"
|
None, description="当前有效限制(自适应使用学习值,固定使用配置值)"
|
||||||
)
|
)
|
||||||
learned_limit: Optional[int] = Field(None, description="学习到的 RPM 限制")
|
learned_limit: int | None = Field(None, description="学习到的 RPM 限制")
|
||||||
concurrent_429_count: int
|
concurrent_429_count: int
|
||||||
rpm_429_count: int
|
rpm_429_count: int
|
||||||
last_429_at: Optional[str]
|
last_429_at: str | None
|
||||||
last_429_type: Optional[str]
|
last_429_type: str | None
|
||||||
adjustment_count: int
|
adjustment_count: int
|
||||||
recent_adjustments: List[dict]
|
recent_adjustments: list[dict]
|
||||||
|
|
||||||
|
|
||||||
class KeyListItem(BaseModel):
|
class KeyListItem(BaseModel):
|
||||||
"""Key 列表项"""
|
"""Key 列表项"""
|
||||||
|
|
||||||
id: str
|
id: str
|
||||||
name: Optional[str]
|
name: str | None
|
||||||
provider_id: str
|
provider_id: str
|
||||||
api_formats: List[str] = Field(default_factory=list)
|
api_formats: list[str] = Field(default_factory=list)
|
||||||
is_adaptive: bool = Field(..., description="是否为自适应模式(rpm_limit=NULL)")
|
is_adaptive: bool = Field(..., description="是否为自适应模式(rpm_limit=NULL)")
|
||||||
rpm_limit: Optional[int] = Field(None, description="固定 RPM 限制(NULL=自适应)")
|
rpm_limit: int | None = Field(None, description="固定 RPM 限制(NULL=自适应)")
|
||||||
effective_limit: Optional[int] = Field(None, description="当前有效限制")
|
effective_limit: int | None = Field(None, description="当前有效限制")
|
||||||
learned_rpm_limit: Optional[int] = Field(None, description="学习到的 RPM 限制")
|
learned_rpm_limit: int | None = Field(None, description="学习到的 RPM 限制")
|
||||||
concurrent_429_count: int
|
concurrent_429_count: int
|
||||||
rpm_429_count: int
|
rpm_429_count: int
|
||||||
|
|
||||||
@@ -76,12 +75,12 @@ class KeyListItem(BaseModel):
|
|||||||
|
|
||||||
@router.get(
|
@router.get(
|
||||||
"/keys",
|
"/keys",
|
||||||
response_model=List[KeyListItem],
|
response_model=list[KeyListItem],
|
||||||
summary="获取所有启用自适应模式的Key",
|
summary="获取所有启用自适应模式的Key",
|
||||||
)
|
)
|
||||||
async def list_adaptive_keys(
|
async def list_adaptive_keys(
|
||||||
request: Request,
|
request: Request,
|
||||||
provider_id: Optional[str] = Query(None, description="按 Provider 过滤"),
|
provider_id: str | None = Query(None, description="按 Provider 过滤"),
|
||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
@@ -207,7 +206,7 @@ async def get_adaptive_summary(
|
|||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class ListAdaptiveKeysAdapter(AdminApiAdapter):
|
class ListAdaptiveKeysAdapter(AdminApiAdapter):
|
||||||
provider_id: Optional[str] = None
|
provider_id: str | None = None
|
||||||
|
|
||||||
async def handle(self, context): # type: ignore[override]
|
async def handle(self, context): # type: ignore[override]
|
||||||
# 自适应模式:rpm_limit = NULL
|
# 自适应模式:rpm_limit = NULL
|
||||||
|
|||||||
@@ -5,7 +5,6 @@
|
|||||||
|
|
||||||
import os
|
import os
|
||||||
from datetime import datetime, timedelta, timezone
|
from datetime import datetime, timedelta, timezone
|
||||||
from typing import Optional
|
|
||||||
from zoneinfo import ZoneInfo
|
from zoneinfo import ZoneInfo
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||||
@@ -25,7 +24,7 @@ from src.services.user.apikey import ApiKeyService
|
|||||||
APP_TIMEZONE = ZoneInfo(os.getenv("APP_TIMEZONE", "Asia/Shanghai"))
|
APP_TIMEZONE = ZoneInfo(os.getenv("APP_TIMEZONE", "Asia/Shanghai"))
|
||||||
|
|
||||||
|
|
||||||
def parse_expiry_date(date_str: Optional[str]) -> Optional[datetime]:
|
def parse_expiry_date(date_str: str | None) -> datetime | None:
|
||||||
"""解析过期日期字符串为 datetime 对象。
|
"""解析过期日期字符串为 datetime 对象。
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
@@ -70,7 +69,7 @@ async def list_standalone_api_keys(
|
|||||||
request: Request,
|
request: Request,
|
||||||
skip: int = Query(0, ge=0),
|
skip: int = Query(0, ge=0),
|
||||||
limit: int = Query(100, ge=1, le=500),
|
limit: int = Query(100, ge=1, le=500),
|
||||||
is_active: Optional[bool] = None,
|
is_active: bool | None = None,
|
||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
@@ -330,7 +329,7 @@ class AdminListStandaloneKeysAdapter(AdminApiAdapter):
|
|||||||
self,
|
self,
|
||||||
skip: int,
|
skip: int,
|
||||||
limit: int,
|
limit: int,
|
||||||
is_active: Optional[bool],
|
is_active: bool | None,
|
||||||
):
|
):
|
||||||
self.skip = skip
|
self.skip = skip
|
||||||
self.limit = limit
|
self.limit = limit
|
||||||
|
|||||||
@@ -5,7 +5,6 @@ Endpoint 健康监控 API
|
|||||||
from collections import defaultdict
|
from collections import defaultdict
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from datetime import datetime, timedelta, timezone
|
from datetime import datetime, timedelta, timezone
|
||||||
from typing import Dict, List, Optional
|
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, Query, Request
|
from fastapi import APIRouter, Depends, Query, Request
|
||||||
from sqlalchemy import func
|
from sqlalchemy import func
|
||||||
@@ -128,7 +127,7 @@ async def get_api_format_health_monitor(
|
|||||||
async def get_key_health(
|
async def get_key_health(
|
||||||
key_id: str,
|
key_id: str,
|
||||||
request: Request,
|
request: Request,
|
||||||
api_format: Optional[str] = Query(None, description="API 格式(可选,如 CLAUDE、OPENAI)"),
|
api_format: str | None = Query(None, description="API 格式(可选,如 CLAUDE、OPENAI)"),
|
||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
) -> HealthStatusResponse:
|
) -> HealthStatusResponse:
|
||||||
"""
|
"""
|
||||||
@@ -161,7 +160,7 @@ async def get_key_health(
|
|||||||
async def recover_key_health(
|
async def recover_key_health(
|
||||||
key_id: str,
|
key_id: str,
|
||||||
request: Request,
|
request: Request,
|
||||||
api_format: Optional[str] = Query(None, description="API 格式(可选,不指定则恢复所有格式)"),
|
api_format: str | None = Query(None, description="API 格式(可选,不指定则恢复所有格式)"),
|
||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
) -> dict:
|
) -> dict:
|
||||||
"""
|
"""
|
||||||
@@ -278,7 +277,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# 构建所有格式的 provider_count 映射
|
# 构建所有格式的 provider_count 映射
|
||||||
all_formats: Dict[str, int] = {}
|
all_formats: dict[str, int] = {}
|
||||||
for api_format_enum, provider_count in active_formats:
|
for api_format_enum, provider_count in active_formats:
|
||||||
api_format = (
|
api_format = (
|
||||||
api_format_enum.value if hasattr(api_format_enum, "value") else str(api_format_enum)
|
api_format_enum.value if hasattr(api_format_enum, "value") else str(api_format_enum)
|
||||||
@@ -295,7 +294,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
|||||||
)
|
)
|
||||||
.all()
|
.all()
|
||||||
)
|
)
|
||||||
endpoint_map: Dict[str, List[str]] = defaultdict(list)
|
endpoint_map: dict[str, list[str]] = defaultdict(list)
|
||||||
active_provider_formats: set[tuple[str, str]] = set()
|
active_provider_formats: set[tuple[str, str]] = set()
|
||||||
for api_format_enum, endpoint_id, provider_id in endpoint_rows:
|
for api_format_enum, endpoint_id, provider_id in endpoint_rows:
|
||||||
api_format = (
|
api_format = (
|
||||||
@@ -305,7 +304,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
|||||||
active_provider_formats.add((str(provider_id), api_format))
|
active_provider_formats.add((str(provider_id), api_format))
|
||||||
|
|
||||||
# 1.2 统计每个 API 格式可用的活跃 Key 数量(Key 属于 Provider,通过 api_formats 关联格式)
|
# 1.2 统计每个 API 格式可用的活跃 Key 数量(Key 属于 Provider,通过 api_formats 关联格式)
|
||||||
key_counts: Dict[str, int] = {}
|
key_counts: dict[str, int] = {}
|
||||||
if active_provider_formats:
|
if active_provider_formats:
|
||||||
active_provider_keys = (
|
active_provider_keys = (
|
||||||
db.query(ProviderAPIKey.provider_id, ProviderAPIKey.api_formats)
|
db.query(ProviderAPIKey.provider_id, ProviderAPIKey.api_formats)
|
||||||
@@ -342,7 +341,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# 构建每个格式的状态统计
|
# 构建每个格式的状态统计
|
||||||
status_counts: Dict[str, Dict[str, int]] = {}
|
status_counts: dict[str, dict[str, int]] = {}
|
||||||
for api_format_enum, status, count in status_counts_query:
|
for api_format_enum, status, count in status_counts_query:
|
||||||
api_format = (
|
api_format = (
|
||||||
api_format_enum.value if hasattr(api_format_enum, "value") else str(api_format_enum)
|
api_format_enum.value if hasattr(api_format_enum, "value") else str(api_format_enum)
|
||||||
@@ -370,7 +369,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
|||||||
.all()
|
.all()
|
||||||
)
|
)
|
||||||
|
|
||||||
grouped_attempts: Dict[str, List[RequestCandidate]] = {}
|
grouped_attempts: dict[str, list[RequestCandidate]] = {}
|
||||||
|
|
||||||
for attempt, api_format_enum, provider_id in rows:
|
for attempt, api_format_enum, provider_id in rows:
|
||||||
api_format = (
|
api_format = (
|
||||||
@@ -384,7 +383,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
|||||||
grouped_attempts[api_format].append(attempt)
|
grouped_attempts[api_format].append(attempt)
|
||||||
|
|
||||||
# 4. 为所有活跃格式生成监控数据(包括没有请求记录的)
|
# 4. 为所有活跃格式生成监控数据(包括没有请求记录的)
|
||||||
monitors: List[ApiFormatHealthMonitor] = []
|
monitors: list[ApiFormatHealthMonitor] = []
|
||||||
for api_format in all_formats:
|
for api_format in all_formats:
|
||||||
attempts = grouped_attempts.get(api_format, [])
|
attempts = grouped_attempts.get(api_format, [])
|
||||||
# 获取窗口内的真实统计数据
|
# 获取窗口内的真实统计数据
|
||||||
@@ -399,7 +398,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
|||||||
|
|
||||||
# 时间线按时间正序
|
# 时间线按时间正序
|
||||||
attempts_sorted = list(reversed(attempts))
|
attempts_sorted = list(reversed(attempts))
|
||||||
events: List[EndpointHealthEvent] = []
|
events: list[EndpointHealthEvent] = []
|
||||||
for attempt in attempts_sorted:
|
for attempt in attempts_sorted:
|
||||||
event_timestamp = attempt.finished_at or attempt.started_at or attempt.created_at
|
event_timestamp = attempt.finished_at or attempt.started_at or attempt.created_at
|
||||||
events.append(
|
events.append(
|
||||||
@@ -462,7 +461,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
|||||||
@dataclass
|
@dataclass
|
||||||
class AdminKeyHealthAdapter(AdminApiAdapter):
|
class AdminKeyHealthAdapter(AdminApiAdapter):
|
||||||
key_id: str
|
key_id: str
|
||||||
api_format: Optional[str] = None
|
api_format: str | None = None
|
||||||
|
|
||||||
async def handle(self, context): # type: ignore[override]
|
async def handle(self, context): # type: ignore[override]
|
||||||
health_data = health_monitor.get_key_health(context.db, self.key_id, self.api_format)
|
health_data = health_monitor.get_key_health(context.db, self.key_id, self.api_format)
|
||||||
@@ -500,7 +499,7 @@ class AdminKeyHealthAdapter(AdminApiAdapter):
|
|||||||
@dataclass
|
@dataclass
|
||||||
class AdminRecoverKeyHealthAdapter(AdminApiAdapter):
|
class AdminRecoverKeyHealthAdapter(AdminApiAdapter):
|
||||||
key_id: str
|
key_id: str
|
||||||
api_format: Optional[str] = None
|
api_format: str | None = None
|
||||||
|
|
||||||
async def handle(self, context): # type: ignore[override]
|
async def handle(self, context): # type: ignore[override]
|
||||||
db = context.db
|
db = context.db
|
||||||
|
|||||||
@@ -6,7 +6,6 @@ import json
|
|||||||
import uuid
|
import uuid
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from datetime import datetime, timezone
|
from datetime import datetime, timezone
|
||||||
from typing import Dict, List
|
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, Query, Request
|
from fastapi import APIRouter, Depends, Query, Request
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
@@ -146,14 +145,14 @@ async def delete_endpoint_key(
|
|||||||
# ========== Provider Keys API ==========
|
# ========== Provider Keys API ==========
|
||||||
|
|
||||||
|
|
||||||
@router.get("/providers/{provider_id}/keys", response_model=List[EndpointAPIKeyResponse])
|
@router.get("/providers/{provider_id}/keys", response_model=list[EndpointAPIKeyResponse])
|
||||||
async def list_provider_keys(
|
async def list_provider_keys(
|
||||||
provider_id: str,
|
provider_id: str,
|
||||||
request: Request,
|
request: Request,
|
||||||
skip: int = Query(0, ge=0, description="跳过的记录数"),
|
skip: int = Query(0, ge=0, description="跳过的记录数"),
|
||||||
limit: int = Query(100, ge=1, le=1000, description="返回的最大记录数"),
|
limit: int = Query(100, ge=1, le=1000, description="返回的最大记录数"),
|
||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
) -> List[EndpointAPIKeyResponse]:
|
) -> list[EndpointAPIKeyResponse]:
|
||||||
"""
|
"""
|
||||||
获取 Provider 的所有 Keys
|
获取 Provider 的所有 Keys
|
||||||
|
|
||||||
@@ -503,12 +502,12 @@ class AdminGetKeysGroupedByFormatAdapter(AdminApiAdapter):
|
|||||||
)
|
)
|
||||||
.all()
|
.all()
|
||||||
)
|
)
|
||||||
endpoint_base_url_map: Dict[tuple[str, str], str] = {}
|
endpoint_base_url_map: dict[tuple[str, str], str] = {}
|
||||||
for provider_id, api_format, base_url in endpoints:
|
for provider_id, api_format, base_url in endpoints:
|
||||||
fmt = api_format.value if hasattr(api_format, "value") else str(api_format)
|
fmt = api_format.value if hasattr(api_format, "value") else str(api_format)
|
||||||
endpoint_base_url_map[(str(provider_id), fmt)] = base_url
|
endpoint_base_url_map[(str(provider_id), fmt)] = base_url
|
||||||
|
|
||||||
grouped: Dict[str, List[dict]] = {}
|
grouped: dict[str, list[dict]] = {}
|
||||||
for key, provider in keys:
|
for key, provider in keys:
|
||||||
api_formats = key.api_formats or []
|
api_formats = key.api_formats or []
|
||||||
|
|
||||||
|
|||||||
@@ -5,10 +5,9 @@ ProviderEndpoint CRUD 管理 API
|
|||||||
import uuid
|
import uuid
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from datetime import datetime, timezone
|
from datetime import datetime, timezone
|
||||||
from typing import List, Optional
|
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, Query, Request
|
from fastapi import APIRouter, Depends, Query, Request
|
||||||
from sqlalchemy import and_, func
|
from sqlalchemy import and_
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
from sqlalchemy.orm.attributes import flag_modified
|
from sqlalchemy.orm.attributes import flag_modified
|
||||||
|
|
||||||
@@ -29,7 +28,7 @@ router = APIRouter(tags=["Endpoint Management"])
|
|||||||
pipeline = ApiRequestPipeline()
|
pipeline = ApiRequestPipeline()
|
||||||
|
|
||||||
|
|
||||||
def mask_proxy_password(proxy_config: Optional[dict]) -> Optional[dict]:
|
def mask_proxy_password(proxy_config: dict | None) -> dict | None:
|
||||||
"""对代理配置中的密码进行脱敏处理"""
|
"""对代理配置中的密码进行脱敏处理"""
|
||||||
if not proxy_config:
|
if not proxy_config:
|
||||||
return None
|
return None
|
||||||
@@ -39,14 +38,14 @@ def mask_proxy_password(proxy_config: Optional[dict]) -> Optional[dict]:
|
|||||||
return masked
|
return masked
|
||||||
|
|
||||||
|
|
||||||
@router.get("/providers/{provider_id}/endpoints", response_model=List[ProviderEndpointResponse])
|
@router.get("/providers/{provider_id}/endpoints", response_model=list[ProviderEndpointResponse])
|
||||||
async def list_provider_endpoints(
|
async def list_provider_endpoints(
|
||||||
provider_id: str,
|
provider_id: str,
|
||||||
request: Request,
|
request: Request,
|
||||||
skip: int = Query(0, ge=0, description="跳过的记录数"),
|
skip: int = Query(0, ge=0, description="跳过的记录数"),
|
||||||
limit: int = Query(100, ge=1, le=1000, description="返回的最大记录数"),
|
limit: int = Query(100, ge=1, le=1000, description="返回的最大记录数"),
|
||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
) -> List[ProviderEndpointResponse]:
|
) -> list[ProviderEndpointResponse]:
|
||||||
"""
|
"""
|
||||||
获取指定 Provider 的所有 Endpoints
|
获取指定 Provider 的所有 Endpoints
|
||||||
|
|
||||||
@@ -245,7 +244,7 @@ class AdminListProviderEndpointsAdapter(AdminApiAdapter):
|
|||||||
if is_active:
|
if is_active:
|
||||||
active_keys_map[fmt] = active_keys_map.get(fmt, 0) + 1
|
active_keys_map[fmt] = active_keys_map.get(fmt, 0) + 1
|
||||||
|
|
||||||
result: List[ProviderEndpointResponse] = []
|
result: list[ProviderEndpointResponse] = []
|
||||||
for endpoint in endpoints:
|
for endpoint in endpoints:
|
||||||
endpoint_format = (
|
endpoint_format = (
|
||||||
endpoint.api_format
|
endpoint.api_format
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
"""LDAP配置管理API端点。"""
|
"""LDAP配置管理API端点。"""
|
||||||
|
|
||||||
import re
|
import re
|
||||||
from typing import Any, Dict, Optional
|
from typing import Any
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, Request
|
from fastapi import APIRouter, Depends, Request
|
||||||
from pydantic import BaseModel, Field, ValidationError, field_validator
|
from pydantic import BaseModel, Field, ValidationError, field_validator
|
||||||
@@ -30,9 +30,9 @@ BCRYPT_HASH_PATTERN = re.compile(r"^\$2[aby]\$\d{2}\$.{53}$")
|
|||||||
class LDAPConfigResponse(BaseModel):
|
class LDAPConfigResponse(BaseModel):
|
||||||
"""LDAP配置响应(不返回密码)"""
|
"""LDAP配置响应(不返回密码)"""
|
||||||
|
|
||||||
server_url: Optional[str] = None
|
server_url: str | None = None
|
||||||
bind_dn: Optional[str] = None
|
bind_dn: str | None = None
|
||||||
base_dn: Optional[str] = None
|
base_dn: str | None = None
|
||||||
has_bind_password: bool = False
|
has_bind_password: bool = False
|
||||||
user_search_filter: str
|
user_search_filter: str
|
||||||
username_attr: str
|
username_attr: str
|
||||||
@@ -50,7 +50,7 @@ class LDAPConfigUpdate(BaseModel):
|
|||||||
server_url: str = Field(..., min_length=1, max_length=255)
|
server_url: str = Field(..., min_length=1, max_length=255)
|
||||||
bind_dn: str = Field(..., min_length=1, max_length=255)
|
bind_dn: str = Field(..., min_length=1, max_length=255)
|
||||||
# 允许空字符串表示"清除密码";非空时自动 strip 并校验不能为空
|
# 允许空字符串表示"清除密码";非空时自动 strip 并校验不能为空
|
||||||
bind_password: Optional[str] = Field(None, max_length=1024)
|
bind_password: str | None = Field(None, max_length=1024)
|
||||||
base_dn: str = Field(..., min_length=1, max_length=255)
|
base_dn: str = Field(..., min_length=1, max_length=255)
|
||||||
user_search_filter: str = Field(default="(uid={username})", max_length=500)
|
user_search_filter: str = Field(default="(uid={username})", max_length=500)
|
||||||
username_attr: str = Field(default="uid", max_length=50)
|
username_attr: str = Field(default="uid", max_length=50)
|
||||||
@@ -63,7 +63,7 @@ class LDAPConfigUpdate(BaseModel):
|
|||||||
|
|
||||||
@field_validator("bind_password")
|
@field_validator("bind_password")
|
||||||
@classmethod
|
@classmethod
|
||||||
def validate_bind_password(cls, v: Optional[str]) -> Optional[str]:
|
def validate_bind_password(cls, v: str | None) -> str | None:
|
||||||
if v is None or v == "":
|
if v is None or v == "":
|
||||||
return v
|
return v
|
||||||
v = v.strip()
|
v = v.strip()
|
||||||
@@ -114,22 +114,22 @@ class LDAPTestResponse(BaseModel):
|
|||||||
class LDAPConfigTest(BaseModel):
|
class LDAPConfigTest(BaseModel):
|
||||||
"""LDAP配置测试请求(全部可选,用于临时覆盖)"""
|
"""LDAP配置测试请求(全部可选,用于临时覆盖)"""
|
||||||
|
|
||||||
server_url: Optional[str] = Field(None, min_length=1, max_length=255)
|
server_url: str | None = Field(None, min_length=1, max_length=255)
|
||||||
bind_dn: Optional[str] = Field(None, min_length=1, max_length=255)
|
bind_dn: str | None = Field(None, min_length=1, max_length=255)
|
||||||
bind_password: Optional[str] = Field(None, min_length=1)
|
bind_password: str | None = Field(None, min_length=1)
|
||||||
base_dn: Optional[str] = Field(None, min_length=1, max_length=255)
|
base_dn: str | None = Field(None, min_length=1, max_length=255)
|
||||||
user_search_filter: Optional[str] = Field(None, max_length=500)
|
user_search_filter: str | None = Field(None, max_length=500)
|
||||||
username_attr: Optional[str] = Field(None, max_length=50)
|
username_attr: str | None = Field(None, max_length=50)
|
||||||
email_attr: Optional[str] = Field(None, max_length=50)
|
email_attr: str | None = Field(None, max_length=50)
|
||||||
display_name_attr: Optional[str] = Field(None, max_length=50)
|
display_name_attr: str | None = Field(None, max_length=50)
|
||||||
is_enabled: Optional[bool] = None
|
is_enabled: bool | None = None
|
||||||
is_exclusive: Optional[bool] = None
|
is_exclusive: bool | None = None
|
||||||
use_starttls: Optional[bool] = None
|
use_starttls: bool | None = None
|
||||||
connect_timeout: Optional[int] = Field(None, ge=1, le=60)
|
connect_timeout: int | None = Field(None, ge=1, le=60)
|
||||||
|
|
||||||
@field_validator("user_search_filter")
|
@field_validator("user_search_filter")
|
||||||
@classmethod
|
@classmethod
|
||||||
def validate_search_filter(cls, v: Optional[str]) -> Optional[str]:
|
def validate_search_filter(cls, v: str | None) -> str | None:
|
||||||
if v is None:
|
if v is None:
|
||||||
return v
|
return v
|
||||||
if "{username}" not in v:
|
if "{username}" not in v:
|
||||||
@@ -263,7 +263,7 @@ async def test_ldap_connection(request: Request, db: Session = Depends(get_db))
|
|||||||
|
|
||||||
|
|
||||||
class AdminGetLDAPConfigAdapter(AdminApiAdapter):
|
class AdminGetLDAPConfigAdapter(AdminApiAdapter):
|
||||||
async def handle(self, context) -> Dict[str, Any]: # type: ignore[override]
|
async def handle(self, context) -> dict[str, Any]: # type: ignore[override]
|
||||||
db = context.db
|
db = context.db
|
||||||
config = db.query(LDAPConfig).first()
|
config = db.query(LDAPConfig).first()
|
||||||
|
|
||||||
@@ -300,7 +300,7 @@ class AdminGetLDAPConfigAdapter(AdminApiAdapter):
|
|||||||
|
|
||||||
|
|
||||||
class AdminUpdateLDAPConfigAdapter(AdminApiAdapter):
|
class AdminUpdateLDAPConfigAdapter(AdminApiAdapter):
|
||||||
async def handle(self, context) -> Dict[str, str]: # type: ignore[override]
|
async def handle(self, context) -> dict[str, str]: # type: ignore[override]
|
||||||
db = context.db
|
db = context.db
|
||||||
payload = context.ensure_json_body()
|
payload = context.ensure_json_body()
|
||||||
|
|
||||||
@@ -421,7 +421,7 @@ class AdminUpdateLDAPConfigAdapter(AdminApiAdapter):
|
|||||||
|
|
||||||
|
|
||||||
class AdminTestLDAPConnectionAdapter(AdminApiAdapter):
|
class AdminTestLDAPConnectionAdapter(AdminApiAdapter):
|
||||||
async def handle(self, context) -> Dict[str, Any]: # type: ignore[override]
|
async def handle(self, context) -> dict[str, Any]: # type: ignore[override]
|
||||||
from src.services.auth.ldap import LDAPService
|
from src.services.auth.ldap import LDAPService
|
||||||
|
|
||||||
db = context.db
|
db = context.db
|
||||||
@@ -442,7 +442,7 @@ class AdminTestLDAPConnectionAdapter(AdminApiAdapter):
|
|||||||
raise InvalidRequestException(translate_pydantic_error(errors[0]))
|
raise InvalidRequestException(translate_pydantic_error(errors[0]))
|
||||||
raise InvalidRequestException("请求数据验证失败")
|
raise InvalidRequestException("请求数据验证失败")
|
||||||
|
|
||||||
config_data: Dict[str, Any] = {}
|
config_data: dict[str, Any] = {}
|
||||||
|
|
||||||
if saved_config:
|
if saved_config:
|
||||||
config_data = {
|
config_data = {
|
||||||
|
|||||||
@@ -1,7 +1,6 @@
|
|||||||
"""管理员 Management Token 管理端点"""
|
"""管理员 Management Token 管理端点"""
|
||||||
|
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import Optional
|
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||||
from fastapi.responses import JSONResponse
|
from fastapi.responses import JSONResponse
|
||||||
@@ -46,8 +45,8 @@ class AdminManagementTokenApiAdapter(AdminApiAdapter):
|
|||||||
@router.get("")
|
@router.get("")
|
||||||
async def list_all_management_tokens(
|
async def list_all_management_tokens(
|
||||||
request: Request,
|
request: Request,
|
||||||
user_id: Optional[str] = Query(None, description="筛选用户 ID"),
|
user_id: str | None = Query(None, description="筛选用户 ID"),
|
||||||
is_active: Optional[bool] = Query(None, description="筛选激活状态"),
|
is_active: bool | None = Query(None, description="筛选激活状态"),
|
||||||
skip: int = Query(0, ge=0),
|
skip: int = Query(0, ge=0),
|
||||||
limit: int = Query(50, ge=1, le=100),
|
limit: int = Query(50, ge=1, le=100),
|
||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
@@ -174,8 +173,8 @@ class AdminListManagementTokensAdapter(AdminManagementTokenApiAdapter):
|
|||||||
"""列出所有 Management Tokens"""
|
"""列出所有 Management Tokens"""
|
||||||
|
|
||||||
name: str = "admin_list_management_tokens"
|
name: str = "admin_list_management_tokens"
|
||||||
user_id: Optional[str] = None
|
user_id: str | None = None
|
||||||
is_active: Optional[bool] = None
|
is_active: bool | None = None
|
||||||
skip: int = 0
|
skip: int = 0
|
||||||
limit: int = 50
|
limit: int = 50
|
||||||
|
|
||||||
@@ -197,7 +196,7 @@ class AdminListManagementTokensAdapter(AdminManagementTokenApiAdapter):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# 预加载用户信息
|
# 预加载用户信息
|
||||||
user_ids = list(set(t.user_id for t in tokens))
|
user_ids = list({t.user_id for t in tokens})
|
||||||
users = {u.id: u for u in context.db.query(User).filter(User.id.in_(user_ids)).all()}
|
users = {u.id: u for u in context.db.query(User).filter(User.id.in_(user_ids)).all()}
|
||||||
for token in tokens:
|
for token in tokens:
|
||||||
token.user = users.get(token.user_id)
|
token.user = users.get(token.user_id)
|
||||||
|
|||||||
@@ -5,7 +5,6 @@
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import Dict, List
|
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, Request
|
from fastapi import APIRouter, Depends, Request
|
||||||
from sqlalchemy.orm import Session, joinedload
|
from sqlalchemy.orm import Session, joinedload
|
||||||
@@ -64,12 +63,12 @@ class AdminGetModelCatalogAdapter(AdminApiAdapter):
|
|||||||
db: Session = context.db
|
db: Session = context.db
|
||||||
|
|
||||||
# 1. 获取所有活跃的 GlobalModel
|
# 1. 获取所有活跃的 GlobalModel
|
||||||
global_models: List[GlobalModel] = (
|
global_models: list[GlobalModel] = (
|
||||||
db.query(GlobalModel).filter(GlobalModel.is_active == True).all()
|
db.query(GlobalModel).filter(GlobalModel.is_active == True).all()
|
||||||
)
|
)
|
||||||
|
|
||||||
# 2. 获取所有活跃的 Model 实现(包含 global_model 以便计算有效价格)
|
# 2. 获取所有活跃的 Model 实现(包含 global_model 以便计算有效价格)
|
||||||
models: List[Model] = (
|
models: list[Model] = (
|
||||||
db.query(Model)
|
db.query(Model)
|
||||||
.options(joinedload(Model.provider), joinedload(Model.global_model))
|
.options(joinedload(Model.provider), joinedload(Model.global_model))
|
||||||
.filter(Model.is_active == True)
|
.filter(Model.is_active == True)
|
||||||
@@ -77,17 +76,17 @@ class AdminGetModelCatalogAdapter(AdminApiAdapter):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# 按 GlobalModel ID 组织关联提供商
|
# 按 GlobalModel ID 组织关联提供商
|
||||||
models_by_global_model: Dict[str, List[Model]] = {}
|
models_by_global_model: dict[str, list[Model]] = {}
|
||||||
for model in models:
|
for model in models:
|
||||||
if model.global_model_id:
|
if model.global_model_id:
|
||||||
models_by_global_model.setdefault(model.global_model_id, []).append(model)
|
models_by_global_model.setdefault(model.global_model_id, []).append(model)
|
||||||
|
|
||||||
# 3. 为每个 GlobalModel 构建 catalog item
|
# 3. 为每个 GlobalModel 构建 catalog item
|
||||||
catalog_items: List[ModelCatalogItem] = []
|
catalog_items: list[ModelCatalogItem] = []
|
||||||
|
|
||||||
for gm in global_models:
|
for gm in global_models:
|
||||||
gm_id = gm.id
|
gm_id = gm.id
|
||||||
provider_entries: List[ModelCatalogProviderDetail] = []
|
provider_entries: list[ModelCatalogProviderDetail] = []
|
||||||
# 从 config JSON 读取能力标志
|
# 从 config JSON 读取能力标志
|
||||||
gm_config = gm.config or {}
|
gm_config = gm.config or {}
|
||||||
capability_flags = {
|
capability_flags = {
|
||||||
|
|||||||
@@ -3,8 +3,7 @@ models.dev 外部模型数据代理
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import json
|
import json
|
||||||
|
from typing import Any
|
||||||
from typing import Any, Optional
|
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
from fastapi import APIRouter, Depends, HTTPException
|
from fastapi import APIRouter, Depends, HTTPException
|
||||||
@@ -43,7 +42,7 @@ OFFICIAL_PROVIDERS = {
|
|||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
async def _get_cached_data() -> Optional[dict[str, Any]]:
|
async def _get_cached_data() -> dict[str, Any] | None:
|
||||||
"""从 Redis 获取缓存数据"""
|
"""从 Redis 获取缓存数据"""
|
||||||
redis = await get_redis_client()
|
redis = await get_redis_client()
|
||||||
if redis is None:
|
if redis is None:
|
||||||
|
|||||||
@@ -5,7 +5,6 @@ GlobalModel Admin API
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import Optional
|
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, Query, Request
|
from fastapi import APIRouter, Depends, Query, Request
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
@@ -37,8 +36,8 @@ async def list_global_models(
|
|||||||
request: Request,
|
request: Request,
|
||||||
skip: int = Query(0, ge=0),
|
skip: int = Query(0, ge=0),
|
||||||
limit: int = Query(100, ge=1, le=1000),
|
limit: int = Query(100, ge=1, le=1000),
|
||||||
is_active: Optional[bool] = Query(None),
|
is_active: bool | None = Query(None),
|
||||||
search: Optional[str] = Query(None),
|
search: str | None = Query(None),
|
||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
) -> GlobalModelListResponse:
|
) -> GlobalModelListResponse:
|
||||||
"""
|
"""
|
||||||
@@ -254,8 +253,8 @@ class AdminListGlobalModelsAdapter(AdminApiAdapter):
|
|||||||
|
|
||||||
skip: int
|
skip: int
|
||||||
limit: int
|
limit: int
|
||||||
is_active: Optional[bool]
|
is_active: bool | None
|
||||||
search: Optional[str]
|
search: str | None
|
||||||
|
|
||||||
async def handle(self, context): # type: ignore[override]
|
async def handle(self, context): # type: ignore[override]
|
||||||
from sqlalchemy import func
|
from sqlalchemy import func
|
||||||
|
|||||||
@@ -9,7 +9,6 @@ GlobalModel 请求链路预览 API
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import Dict, List, Optional
|
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, Request
|
from fastapi import APIRouter, Depends, Request
|
||||||
from pydantic import BaseModel, ConfigDict, Field
|
from pydantic import BaseModel, ConfigDict, Field
|
||||||
@@ -47,22 +46,22 @@ class RoutingKeyInfo(BaseModel):
|
|||||||
name: str
|
name: str
|
||||||
masked_key: str = Field("", description="脱敏的 API Key")
|
masked_key: str = Field("", description="脱敏的 API Key")
|
||||||
internal_priority: int = Field(..., description="Key 内部优先级")
|
internal_priority: int = Field(..., description="Key 内部优先级")
|
||||||
global_priority_by_format: Optional[Dict[str, int]] = Field(None, description="按 API 格式的全局优先级")
|
global_priority_by_format: dict[str, int] | None = Field(None, description="按 API 格式的全局优先级")
|
||||||
rpm_limit: Optional[int] = Field(None, description="RPM 限制,null 表示自适应")
|
rpm_limit: int | None = Field(None, description="RPM 限制,null 表示自适应")
|
||||||
is_adaptive: bool = Field(False, description="是否为自适应 RPM 模式")
|
is_adaptive: bool = Field(False, description="是否为自适应 RPM 模式")
|
||||||
effective_rpm: Optional[int] = Field(None, description="有效 RPM 限制")
|
effective_rpm: int | None = Field(None, description="有效 RPM 限制")
|
||||||
cache_ttl_minutes: int = Field(0, description="缓存 TTL(分钟)")
|
cache_ttl_minutes: int = Field(0, description="缓存 TTL(分钟)")
|
||||||
health_score: float = Field(1.0, description="健康度分数(0-1 小数格式)")
|
health_score: float = Field(1.0, description="健康度分数(0-1 小数格式)")
|
||||||
is_active: bool
|
is_active: bool
|
||||||
api_formats: List[str] = Field(default_factory=list, description="支持的 API 格式")
|
api_formats: list[str] = Field(default_factory=list, description="支持的 API 格式")
|
||||||
# 模型白名单
|
# 模型白名单
|
||||||
allowed_models: Optional[List[str]] = Field(None, description="允许的模型列表,null 表示不限制")
|
allowed_models: list[str] | None = Field(None, description="允许的模型列表,null 表示不限制")
|
||||||
# 熔断状态
|
# 熔断状态
|
||||||
circuit_breaker_open: bool = Field(False, description="熔断器是否打开")
|
circuit_breaker_open: bool = Field(False, description="熔断器是否打开")
|
||||||
circuit_breaker_formats: List[str] = Field(
|
circuit_breaker_formats: list[str] = Field(
|
||||||
default_factory=list, description="熔断的 API 格式列表"
|
default_factory=list, description="熔断的 API 格式列表"
|
||||||
)
|
)
|
||||||
next_probe_at: Optional[str] = Field(None, description="下次探测时间(ISO格式)")
|
next_probe_at: str | None = Field(None, description="下次探测时间(ISO格式)")
|
||||||
|
|
||||||
model_config = ConfigDict(from_attributes=True)
|
model_config = ConfigDict(from_attributes=True)
|
||||||
|
|
||||||
@@ -73,9 +72,9 @@ class RoutingEndpointInfo(BaseModel):
|
|||||||
id: str
|
id: str
|
||||||
api_format: str
|
api_format: str
|
||||||
base_url: str
|
base_url: str
|
||||||
custom_path: Optional[str] = None
|
custom_path: str | None = None
|
||||||
is_active: bool
|
is_active: bool
|
||||||
keys: List[RoutingKeyInfo] = Field(default_factory=list)
|
keys: list[RoutingKeyInfo] = Field(default_factory=list)
|
||||||
total_keys: int = 0
|
total_keys: int = 0
|
||||||
active_keys: int = 0
|
active_keys: int = 0
|
||||||
|
|
||||||
@@ -87,7 +86,7 @@ class RoutingModelMapping(BaseModel):
|
|||||||
|
|
||||||
name: str = Field(..., description="映射名称")
|
name: str = Field(..., description="映射名称")
|
||||||
priority: int = Field(..., description="优先级(数字越小优先级越高)")
|
priority: int = Field(..., description="优先级(数字越小优先级越高)")
|
||||||
api_formats: Optional[List[str]] = Field(None, description="作用域(适用的 API 格式)")
|
api_formats: list[str] | None = Field(None, description="作用域(适用的 API 格式)")
|
||||||
|
|
||||||
|
|
||||||
class RoutingProviderInfo(BaseModel):
|
class RoutingProviderInfo(BaseModel):
|
||||||
@@ -97,18 +96,18 @@ class RoutingProviderInfo(BaseModel):
|
|||||||
name: str
|
name: str
|
||||||
model_id: str = Field(..., description="Model ID(GlobalModel 与 Provider 的关联记录 ID)")
|
model_id: str = Field(..., description="Model ID(GlobalModel 与 Provider 的关联记录 ID)")
|
||||||
provider_priority: int = Field(..., description="提供商优先级(数字越小优先级越高)")
|
provider_priority: int = Field(..., description="提供商优先级(数字越小优先级越高)")
|
||||||
billing_type: Optional[str] = Field(None, description="计费类型")
|
billing_type: str | None = Field(None, description="计费类型")
|
||||||
monthly_quota_usd: Optional[float] = Field(None, description="月额度(美元)")
|
monthly_quota_usd: float | None = Field(None, description="月额度(美元)")
|
||||||
monthly_used_usd: Optional[float] = Field(None, description="已用额度(美元)")
|
monthly_used_usd: float | None = Field(None, description="已用额度(美元)")
|
||||||
is_active: bool
|
is_active: bool
|
||||||
# 模型映射信息
|
# 模型映射信息
|
||||||
provider_model_name: str = Field(..., description="提供商侧的模型名称")
|
provider_model_name: str = Field(..., description="提供商侧的模型名称")
|
||||||
model_mappings: List[RoutingModelMapping] = Field(
|
model_mappings: list[RoutingModelMapping] = Field(
|
||||||
default_factory=list, description="模型名称映射列表"
|
default_factory=list, description="模型名称映射列表"
|
||||||
)
|
)
|
||||||
model_is_active: bool = Field(True, description="Model 是否活跃")
|
model_is_active: bool = Field(True, description="Model 是否活跃")
|
||||||
# Endpoint 和 Key 信息
|
# Endpoint 和 Key 信息
|
||||||
endpoints: List[RoutingEndpointInfo] = Field(default_factory=list)
|
endpoints: list[RoutingEndpointInfo] = Field(default_factory=list)
|
||||||
total_endpoints: int = 0
|
total_endpoints: int = 0
|
||||||
active_endpoints: int = 0
|
active_endpoints: int = 0
|
||||||
|
|
||||||
@@ -123,7 +122,7 @@ class GlobalKeyWhitelistItem(BaseModel):
|
|||||||
masked_key: str = Field(..., description="脱敏的 API Key")
|
masked_key: str = Field(..., description="脱敏的 API Key")
|
||||||
provider_id: str = Field(..., description="Provider ID")
|
provider_id: str = Field(..., description="Provider ID")
|
||||||
provider_name: str = Field(..., description="Provider 名称")
|
provider_name: str = Field(..., description="Provider 名称")
|
||||||
allowed_models: List[str] = Field(default_factory=list, description="Key 白名单模型列表")
|
allowed_models: list[str] = Field(default_factory=list, description="Key 白名单模型列表")
|
||||||
|
|
||||||
model_config = ConfigDict(from_attributes=True)
|
model_config = ConfigDict(from_attributes=True)
|
||||||
|
|
||||||
@@ -136,11 +135,11 @@ class ModelRoutingPreviewResponse(BaseModel):
|
|||||||
display_name: str
|
display_name: str
|
||||||
is_active: bool
|
is_active: bool
|
||||||
# GlobalModel 的模型映射(用于前端匹配 Key 白名单)
|
# GlobalModel 的模型映射(用于前端匹配 Key 白名单)
|
||||||
global_model_mappings: List[str] = Field(
|
global_model_mappings: list[str] = Field(
|
||||||
default_factory=list, description="GlobalModel 的模型映射规则(正则模式)"
|
default_factory=list, description="GlobalModel 的模型映射规则(正则模式)"
|
||||||
)
|
)
|
||||||
# 链路信息
|
# 链路信息
|
||||||
providers: List[RoutingProviderInfo] = Field(
|
providers: list[RoutingProviderInfo] = Field(
|
||||||
default_factory=list, description="按优先级排序的提供商列表"
|
default_factory=list, description="按优先级排序的提供商列表"
|
||||||
)
|
)
|
||||||
total_providers: int = 0
|
total_providers: int = 0
|
||||||
@@ -149,7 +148,7 @@ class ModelRoutingPreviewResponse(BaseModel):
|
|||||||
scheduling_mode: str = Field("cache_affinity", description="调度模式")
|
scheduling_mode: str = Field("cache_affinity", description="调度模式")
|
||||||
priority_mode: str = Field("provider", description="优先级模式")
|
priority_mode: str = Field("provider", description="优先级模式")
|
||||||
# 全局 Key 白名单数据(供前端实时匹配,包含所有 Provider 的 Key)
|
# 全局 Key 白名单数据(供前端实时匹配,包含所有 Provider 的 Key)
|
||||||
all_keys_whitelist: List[GlobalKeyWhitelistItem] = Field(
|
all_keys_whitelist: list[GlobalKeyWhitelistItem] = Field(
|
||||||
default_factory=list, description="所有 Provider 的 Key 白名单数据"
|
default_factory=list, description="所有 Provider 的 Key 白名单数据"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -227,7 +226,7 @@ class AdminGetModelRoutingPreviewAdapter(AdminApiAdapter):
|
|||||||
provider_ids = [m.provider_id for m in models if m.provider_id]
|
provider_ids = [m.provider_id for m in models if m.provider_id]
|
||||||
|
|
||||||
# 批量获取 Provider 的 Endpoints
|
# 批量获取 Provider 的 Endpoints
|
||||||
endpoints_by_provider: Dict[str, List[ProviderEndpoint]] = {}
|
endpoints_by_provider: dict[str, list[ProviderEndpoint]] = {}
|
||||||
if provider_ids:
|
if provider_ids:
|
||||||
endpoints = (
|
endpoints = (
|
||||||
db.query(ProviderEndpoint)
|
db.query(ProviderEndpoint)
|
||||||
@@ -240,7 +239,7 @@ class AdminGetModelRoutingPreviewAdapter(AdminApiAdapter):
|
|||||||
endpoints_by_provider[ep.provider_id].append(ep)
|
endpoints_by_provider[ep.provider_id].append(ep)
|
||||||
|
|
||||||
# 批量获取 Provider 的 Keys
|
# 批量获取 Provider 的 Keys
|
||||||
keys_by_provider: Dict[str, List[ProviderAPIKey]] = {}
|
keys_by_provider: dict[str, list[ProviderAPIKey]] = {}
|
||||||
if provider_ids:
|
if provider_ids:
|
||||||
keys = (
|
keys = (
|
||||||
db.query(ProviderAPIKey).filter(ProviderAPIKey.provider_id.in_(provider_ids)).all()
|
db.query(ProviderAPIKey).filter(ProviderAPIKey.provider_id.in_(provider_ids)).all()
|
||||||
@@ -251,14 +250,14 @@ class AdminGetModelRoutingPreviewAdapter(AdminApiAdapter):
|
|||||||
keys_by_provider[key.provider_id].append(key)
|
keys_by_provider[key.provider_id].append(key)
|
||||||
|
|
||||||
# 提取 GlobalModel 的 model_mappings(用于 Key 白名单匹配)
|
# 提取 GlobalModel 的 model_mappings(用于 Key 白名单匹配)
|
||||||
global_model_mappings: List[str] = []
|
global_model_mappings: list[str] = []
|
||||||
if global_model.config and isinstance(global_model.config, dict):
|
if global_model.config and isinstance(global_model.config, dict):
|
||||||
mappings = global_model.config.get("model_mappings")
|
mappings = global_model.config.get("model_mappings")
|
||||||
if isinstance(mappings, list):
|
if isinstance(mappings, list):
|
||||||
global_model_mappings = [m for m in mappings if isinstance(m, str)]
|
global_model_mappings = [m for m in mappings if isinstance(m, str)]
|
||||||
|
|
||||||
# 构建 Provider 路由信息
|
# 构建 Provider 路由信息
|
||||||
provider_infos: List[RoutingProviderInfo] = []
|
provider_infos: list[RoutingProviderInfo] = []
|
||||||
for model in models:
|
for model in models:
|
||||||
provider = model.provider
|
provider = model.provider
|
||||||
if not provider:
|
if not provider:
|
||||||
@@ -281,7 +280,7 @@ class AdminGetModelRoutingPreviewAdapter(AdminApiAdapter):
|
|||||||
provider_keys = keys_by_provider.get(provider.id, [])
|
provider_keys = keys_by_provider.get(provider.id, [])
|
||||||
|
|
||||||
# 按 api_format 组织 Keys
|
# 按 api_format 组织 Keys
|
||||||
keys_by_endpoint: Dict[str, List[ProviderAPIKey]] = {}
|
keys_by_endpoint: dict[str, list[ProviderAPIKey]] = {}
|
||||||
for key in provider_keys:
|
for key in provider_keys:
|
||||||
# 每个 Key 可能支持多个 api_formats
|
# 每个 Key 可能支持多个 api_formats
|
||||||
for fmt in key.api_formats or []:
|
for fmt in key.api_formats or []:
|
||||||
@@ -355,8 +354,8 @@ class AdminGetModelRoutingPreviewAdapter(AdminApiAdapter):
|
|||||||
|
|
||||||
# 检查熔断状态
|
# 检查熔断状态
|
||||||
circuit_breaker_open = False
|
circuit_breaker_open = False
|
||||||
circuit_breaker_formats: List[str] = []
|
circuit_breaker_formats: list[str] = []
|
||||||
next_probe_at: Optional[str] = None
|
next_probe_at: str | None = None
|
||||||
if key.circuit_breaker_by_format:
|
if key.circuit_breaker_by_format:
|
||||||
for fmt, cb_state in key.circuit_breaker_by_format.items():
|
for fmt, cb_state in key.circuit_breaker_by_format.items():
|
||||||
if isinstance(cb_state, dict) and cb_state.get("open"):
|
if isinstance(cb_state, dict) and cb_state.get("open"):
|
||||||
@@ -462,7 +461,7 @@ class AdminGetModelRoutingPreviewAdapter(AdminApiAdapter):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# 获取所有活跃 Provider 的 Key 白名单数据(供前端实时匹配)
|
# 获取所有活跃 Provider 的 Key 白名单数据(供前端实时匹配)
|
||||||
all_keys_whitelist: List[GlobalKeyWhitelistItem] = []
|
all_keys_whitelist: list[GlobalKeyWhitelistItem] = []
|
||||||
crypto = CryptoService()
|
crypto = CryptoService()
|
||||||
|
|
||||||
# 获取所有活跃的 Key(带白名单),使用 selectinload 避免 N+1 查询
|
# 获取所有活跃的 Key(带白名单),使用 selectinload 避免 N+1 查询
|
||||||
|
|||||||
@@ -1,7 +1,8 @@
|
|||||||
"""模块管理 API 端点"""
|
"""模块管理 API 端点"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import Any, Dict, Optional
|
from typing import Any
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, Request
|
from fastapi import APIRouter, Depends, Request
|
||||||
from pydantic import BaseModel
|
from pydantic import BaseModel
|
||||||
@@ -28,18 +29,18 @@ class ModuleStatusResponse(BaseModel):
|
|||||||
enabled: bool
|
enabled: bool
|
||||||
active: bool
|
active: bool
|
||||||
config_validated: bool
|
config_validated: bool
|
||||||
config_error: Optional[str]
|
config_error: str | None
|
||||||
display_name: str
|
display_name: str
|
||||||
description: str
|
description: str
|
||||||
category: str
|
category: str
|
||||||
admin_route: Optional[str]
|
admin_route: str | None
|
||||||
admin_menu_icon: Optional[str]
|
admin_menu_icon: str | None
|
||||||
admin_menu_group: Optional[str]
|
admin_menu_group: str | None
|
||||||
admin_menu_order: int
|
admin_menu_order: int
|
||||||
health: str
|
health: str
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_status(cls, status: ModuleStatus) -> "ModuleStatusResponse":
|
def from_status(cls, status: ModuleStatus) -> ModuleStatusResponse:
|
||||||
return cls(
|
return cls(
|
||||||
name=status.name,
|
name=status.name,
|
||||||
available=status.available,
|
available=status.available,
|
||||||
@@ -130,7 +131,7 @@ async def set_module_enabled(
|
|||||||
class AdminGetAllModulesStatusAdapter(AdminApiAdapter):
|
class AdminGetAllModulesStatusAdapter(AdminApiAdapter):
|
||||||
"""获取所有模块状态"""
|
"""获取所有模块状态"""
|
||||||
|
|
||||||
async def handle(self, context) -> Dict[str, Any]:
|
async def handle(self, context) -> dict[str, Any]:
|
||||||
registry = get_module_registry()
|
registry = get_module_registry()
|
||||||
all_status = await registry.get_all_status_async(context.db)
|
all_status = await registry.get_all_status_async(context.db)
|
||||||
|
|
||||||
@@ -146,7 +147,7 @@ class AdminGetModuleStatusAdapter(AdminApiAdapter):
|
|||||||
|
|
||||||
module_name: str
|
module_name: str
|
||||||
|
|
||||||
async def handle(self, context) -> Dict[str, Any]:
|
async def handle(self, context) -> dict[str, Any]:
|
||||||
registry = get_module_registry()
|
registry = get_module_registry()
|
||||||
status = await registry.get_module_status_async(self.module_name, context.db)
|
status = await registry.get_module_status_async(self.module_name, context.db)
|
||||||
|
|
||||||
@@ -162,7 +163,7 @@ class AdminSetModuleEnabledAdapter(AdminApiAdapter):
|
|||||||
|
|
||||||
module_name: str
|
module_name: str
|
||||||
|
|
||||||
async def handle(self, context) -> Dict[str, Any]:
|
async def handle(self, context) -> dict[str, Any]:
|
||||||
registry = get_module_registry()
|
registry = get_module_registry()
|
||||||
|
|
||||||
# 检查模块是否存在
|
# 检查模块是否存在
|
||||||
|
|||||||
@@ -2,7 +2,6 @@
|
|||||||
|
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from datetime import datetime, timedelta, timezone
|
from datetime import datetime, timedelta, timezone
|
||||||
from typing import List, Optional
|
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||||
from sqlalchemy import func
|
from sqlalchemy import func
|
||||||
@@ -33,8 +32,8 @@ pipeline = ApiRequestPipeline()
|
|||||||
@router.get("/audit-logs")
|
@router.get("/audit-logs")
|
||||||
async def get_audit_logs(
|
async def get_audit_logs(
|
||||||
request: Request,
|
request: Request,
|
||||||
username: Optional[str] = Query(None, description="用户名筛选 (模糊匹配)"),
|
username: str | None = Query(None, description="用户名筛选 (模糊匹配)"),
|
||||||
event_type: Optional[str] = Query(None, description="事件类型筛选"),
|
event_type: str | None = Query(None, description="事件类型筛选"),
|
||||||
days: int = Query(7, description="查询天数"),
|
days: int = Query(7, description="查询天数"),
|
||||||
limit: int = Query(100, description="返回数量限制"),
|
limit: int = Query(100, description="返回数量限制"),
|
||||||
offset: int = Query(0, description="偏移量"),
|
offset: int = Query(0, description="偏移量"),
|
||||||
@@ -212,8 +211,8 @@ async def get_circuit_history(
|
|||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class AdminGetAuditLogsAdapter(AdminApiAdapter):
|
class AdminGetAuditLogsAdapter(AdminApiAdapter):
|
||||||
username: Optional[str]
|
username: str | None
|
||||||
event_type: Optional[str]
|
event_type: str | None
|
||||||
days: int
|
days: int
|
||||||
limit: int
|
limit: int
|
||||||
offset: int
|
offset: int
|
||||||
@@ -497,8 +496,8 @@ class AdminCircuitHistoryAdapter(AdminApiAdapter):
|
|||||||
return {"items": history, "count": len(history)}
|
return {"items": history, "count": len(history)}
|
||||||
|
|
||||||
|
|
||||||
def _get_health_recommendations(error_stats: dict, health_score: int) -> List[str]:
|
def _get_health_recommendations(error_stats: dict, health_score: int) -> list[str]:
|
||||||
recommendations: List[str] = []
|
recommendations: list[str] = []
|
||||||
if health_score < 50:
|
if health_score < 50:
|
||||||
recommendations.append("系统健康状况严重,请立即检查错误日志")
|
recommendations.append("系统健康状况严重,请立即检查错误日志")
|
||||||
if error_stats.get("total_errors", 0) > 100:
|
if error_stats.get("total_errors", 0) > 100:
|
||||||
|
|||||||
@@ -5,7 +5,7 @@
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import Any, Dict, List, Optional, Tuple
|
from typing import Any
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||||
from fastapi.responses import PlainTextResponse
|
from fastapi.responses import PlainTextResponse
|
||||||
@@ -13,7 +13,7 @@ from sqlalchemy.orm import Session
|
|||||||
|
|
||||||
from src.api.base.admin_adapter import AdminApiAdapter
|
from src.api.base.admin_adapter import AdminApiAdapter
|
||||||
from src.api.base.context import ApiRequestContext
|
from src.api.base.context import ApiRequestContext
|
||||||
from src.api.base.pagination import PaginationMeta, build_pagination_payload, paginate_sequence
|
from src.api.base.pagination import build_pagination_payload, paginate_sequence
|
||||||
from src.api.base.pipeline import ApiRequestPipeline
|
from src.api.base.pipeline import ApiRequestPipeline
|
||||||
from src.clients.redis_client import get_redis_client_sync
|
from src.clients.redis_client import get_redis_client_sync
|
||||||
from src.core.crypto import crypto_service
|
from src.core.crypto import crypto_service
|
||||||
@@ -28,7 +28,7 @@ router = APIRouter(prefix="/api/admin/monitoring/cache", tags=["Admin - Monitori
|
|||||||
pipeline = ApiRequestPipeline()
|
pipeline = ApiRequestPipeline()
|
||||||
|
|
||||||
|
|
||||||
def mask_api_key(api_key: Optional[str], prefix_len: int = 8, suffix_len: int = 4) -> Optional[str]:
|
def mask_api_key(api_key: str | None, prefix_len: int = 8, suffix_len: int = 4) -> str | None:
|
||||||
"""
|
"""
|
||||||
脱敏 API Key,显示前缀 + 星号 + 后缀
|
脱敏 API Key,显示前缀 + 星号 + 后缀
|
||||||
例如: sk-jhiId-xxxxxxxxxxxAABB -> sk-jhiId-********AABB
|
例如: sk-jhiId-xxxxxxxxxxxAABB -> sk-jhiId-********AABB
|
||||||
@@ -47,7 +47,7 @@ def mask_api_key(api_key: Optional[str], prefix_len: int = 8, suffix_len: int =
|
|||||||
return f"{api_key[:prefix_len]}********{api_key[-suffix_len:]}"
|
return f"{api_key[:prefix_len]}********{api_key[-suffix_len:]}"
|
||||||
|
|
||||||
|
|
||||||
def decrypt_and_mask(encrypted_key: Optional[str], prefix_len: int = 8) -> Optional[str]:
|
def decrypt_and_mask(encrypted_key: str | None, prefix_len: int = 8) -> str | None:
|
||||||
"""
|
"""
|
||||||
解密 API Key 后脱敏显示
|
解密 API Key 后脱敏显示
|
||||||
|
|
||||||
@@ -65,7 +65,7 @@ def decrypt_and_mask(encrypted_key: Optional[str], prefix_len: int = 8) -> Optio
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
def resolve_user_identifier(db: Session, identifier: str) -> Optional[str]:
|
def resolve_user_identifier(db: Session, identifier: str) -> str | None:
|
||||||
"""
|
"""
|
||||||
将用户标识符(username/email/user_id/api_key_id)解析为 user_id
|
将用户标识符(username/email/user_id/api_key_id)解析为 user_id
|
||||||
|
|
||||||
@@ -181,7 +181,7 @@ async def get_user_affinity(
|
|||||||
@router.get("/affinities")
|
@router.get("/affinities")
|
||||||
async def list_affinities(
|
async def list_affinities(
|
||||||
request: Request,
|
request: Request,
|
||||||
keyword: Optional[str] = None,
|
keyword: str | None = None,
|
||||||
limit: int = Query(100, ge=1, le=1000, description="返回数量限制"),
|
limit: int = Query(100, ge=1, le=1000, description="返回数量限制"),
|
||||||
offset: int = Query(0, ge=0, description="偏移量"),
|
offset: int = Query(0, ge=0, description="偏移量"),
|
||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
@@ -421,7 +421,7 @@ async def get_cache_metrics(
|
|||||||
|
|
||||||
|
|
||||||
class AdminCacheStatsAdapter(AdminApiAdapter):
|
class AdminCacheStatsAdapter(AdminApiAdapter):
|
||||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||||
try:
|
try:
|
||||||
redis_client = get_redis_client_sync()
|
redis_client = get_redis_client_sync()
|
||||||
# 读取系统配置,确保监控接口与编排器使用一致的模式
|
# 读取系统配置,确保监控接口与编排器使用一致的模式
|
||||||
@@ -487,14 +487,14 @@ class AdminCacheMetricsAdapter(AdminApiAdapter):
|
|||||||
logger.exception(f"导出缓存指标失败: {exc}")
|
logger.exception(f"导出缓存指标失败: {exc}")
|
||||||
raise HTTPException(status_code=500, detail=f"导出缓存指标失败: {exc}")
|
raise HTTPException(status_code=500, detail=f"导出缓存指标失败: {exc}")
|
||||||
|
|
||||||
def _format_prometheus(self, stats: Dict[str, Any]) -> str:
|
def _format_prometheus(self, stats: dict[str, Any]) -> str:
|
||||||
"""
|
"""
|
||||||
将 scheduler/affinity 指标转换为 Prometheus 文本格式。
|
将 scheduler/affinity 指标转换为 Prometheus 文本格式。
|
||||||
"""
|
"""
|
||||||
scheduler_metrics = stats.get("scheduler_metrics", {})
|
scheduler_metrics = stats.get("scheduler_metrics", {})
|
||||||
affinity_stats = stats.get("affinity_stats", {})
|
affinity_stats = stats.get("affinity_stats", {})
|
||||||
|
|
||||||
metric_map: List[Tuple[str, str, float]] = [
|
metric_map: list[tuple[str, str, float]] = [
|
||||||
(
|
(
|
||||||
"cache_scheduler_total_batches",
|
"cache_scheduler_total_batches",
|
||||||
"Total batches pulled from provider list",
|
"Total batches pulled from provider list",
|
||||||
@@ -542,7 +542,7 @@ class AdminCacheMetricsAdapter(AdminApiAdapter):
|
|||||||
),
|
),
|
||||||
]
|
]
|
||||||
|
|
||||||
affinity_map: List[Tuple[str, str, float]] = [
|
affinity_map: list[tuple[str, str, float]] = [
|
||||||
(
|
(
|
||||||
"cache_affinity_total",
|
"cache_affinity_total",
|
||||||
"Total cache affinities stored",
|
"Total cache affinities stored",
|
||||||
@@ -596,7 +596,7 @@ class AdminCacheMetricsAdapter(AdminApiAdapter):
|
|||||||
class AdminGetUserAffinityAdapter(AdminApiAdapter):
|
class AdminGetUserAffinityAdapter(AdminApiAdapter):
|
||||||
user_identifier: str
|
user_identifier: str
|
||||||
|
|
||||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||||
db = context.db
|
db = context.db
|
||||||
try:
|
try:
|
||||||
user_id = resolve_user_identifier(db, self.user_identifier)
|
user_id = resolve_user_identifier(db, self.user_identifier)
|
||||||
@@ -673,11 +673,11 @@ class AdminGetUserAffinityAdapter(AdminApiAdapter):
|
|||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class AdminListAffinitiesAdapter(AdminApiAdapter):
|
class AdminListAffinitiesAdapter(AdminApiAdapter):
|
||||||
keyword: Optional[str]
|
keyword: str | None
|
||||||
limit: int
|
limit: int
|
||||||
offset: int
|
offset: int
|
||||||
|
|
||||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||||
db = context.db
|
db = context.db
|
||||||
redis_client = get_redis_client_sync()
|
redis_client = get_redis_client_sync()
|
||||||
if not redis_client:
|
if not redis_client:
|
||||||
@@ -686,7 +686,7 @@ class AdminListAffinitiesAdapter(AdminApiAdapter):
|
|||||||
affinity_mgr = await get_affinity_manager(redis_client)
|
affinity_mgr = await get_affinity_manager(redis_client)
|
||||||
matched_user_id = None
|
matched_user_id = None
|
||||||
matched_api_key_id = None
|
matched_api_key_id = None
|
||||||
raw_affinities: List[Dict[str, Any]] = []
|
raw_affinities: list[dict[str, Any]] = []
|
||||||
|
|
||||||
if self.keyword:
|
if self.keyword:
|
||||||
# 首先检查是否是 API Key ID(affinity_key)
|
# 首先检查是否是 API Key ID(affinity_key)
|
||||||
@@ -724,14 +724,14 @@ class AdminListAffinitiesAdapter(AdminApiAdapter):
|
|||||||
}
|
}
|
||||||
|
|
||||||
# 批量查询用户 API Key 信息
|
# 批量查询用户 API Key 信息
|
||||||
user_api_key_map: Dict[str, ApiKey] = {}
|
user_api_key_map: dict[str, ApiKey] = {}
|
||||||
if affinity_keys:
|
if affinity_keys:
|
||||||
user_api_keys = db.query(ApiKey).filter(ApiKey.id.in_(list(affinity_keys))).all()
|
user_api_keys = db.query(ApiKey).filter(ApiKey.id.in_(list(affinity_keys))).all()
|
||||||
user_api_key_map = {str(k.id): k for k in user_api_keys}
|
user_api_key_map = {str(k.id): k for k in user_api_keys}
|
||||||
|
|
||||||
# 收集所有 user_id
|
# 收集所有 user_id
|
||||||
user_ids = {str(k.user_id) for k in user_api_key_map.values()}
|
user_ids = {str(k.user_id) for k in user_api_key_map.values()}
|
||||||
user_map: Dict[str, User] = {}
|
user_map: dict[str, User] = {}
|
||||||
if user_ids:
|
if user_ids:
|
||||||
users = db.query(User).filter(User.id.in_(list(user_ids))).all()
|
users = db.query(User).filter(User.id.in_(list(user_ids))).all()
|
||||||
user_map = {str(user.id): user for user in users}
|
user_map = {str(user.id): user for user in users}
|
||||||
@@ -771,7 +771,7 @@ class AdminListAffinitiesAdapter(AdminApiAdapter):
|
|||||||
global_model_ids = {
|
global_model_ids = {
|
||||||
item.get("model_name") for item in raw_affinities if item.get("model_name")
|
item.get("model_name") for item in raw_affinities if item.get("model_name")
|
||||||
}
|
}
|
||||||
global_model_map: Dict[str, GlobalModel] = {}
|
global_model_map: dict[str, GlobalModel] = {}
|
||||||
if global_model_ids:
|
if global_model_ids:
|
||||||
# model_name 可能是 UUID 格式的 global_model_id,也可能是原始模型名称
|
# model_name 可能是 UUID 格式的 global_model_id,也可能是原始模型名称
|
||||||
global_models = db.query(GlobalModel).filter(
|
global_models = db.query(GlobalModel).filter(
|
||||||
@@ -885,7 +885,7 @@ class AdminListAffinitiesAdapter(AdminApiAdapter):
|
|||||||
class AdminClearUserCacheAdapter(AdminApiAdapter):
|
class AdminClearUserCacheAdapter(AdminApiAdapter):
|
||||||
user_identifier: str
|
user_identifier: str
|
||||||
|
|
||||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||||
db = context.db
|
db = context.db
|
||||||
try:
|
try:
|
||||||
redis_client = get_redis_client_sync()
|
redis_client = get_redis_client_sync()
|
||||||
@@ -995,7 +995,7 @@ class AdminClearSingleAffinityAdapter(AdminApiAdapter):
|
|||||||
model_id: str
|
model_id: str
|
||||||
api_format: str
|
api_format: str
|
||||||
|
|
||||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||||
db = context.db
|
db = context.db
|
||||||
try:
|
try:
|
||||||
redis_client = get_redis_client_sync()
|
redis_client = get_redis_client_sync()
|
||||||
@@ -1048,7 +1048,7 @@ class AdminClearSingleAffinityAdapter(AdminApiAdapter):
|
|||||||
|
|
||||||
|
|
||||||
class AdminClearAllCacheAdapter(AdminApiAdapter):
|
class AdminClearAllCacheAdapter(AdminApiAdapter):
|
||||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||||
try:
|
try:
|
||||||
redis_client = get_redis_client_sync()
|
redis_client = get_redis_client_sync()
|
||||||
affinity_mgr = await get_affinity_manager(redis_client)
|
affinity_mgr = await get_affinity_manager(redis_client)
|
||||||
@@ -1068,7 +1068,7 @@ class AdminClearAllCacheAdapter(AdminApiAdapter):
|
|||||||
class AdminClearProviderCacheAdapter(AdminApiAdapter):
|
class AdminClearProviderCacheAdapter(AdminApiAdapter):
|
||||||
provider_id: str
|
provider_id: str
|
||||||
|
|
||||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||||
try:
|
try:
|
||||||
redis_client = get_redis_client_sync()
|
redis_client = get_redis_client_sync()
|
||||||
affinity_mgr = await get_affinity_manager(redis_client)
|
affinity_mgr = await get_affinity_manager(redis_client)
|
||||||
@@ -1091,7 +1091,7 @@ class AdminClearProviderCacheAdapter(AdminApiAdapter):
|
|||||||
|
|
||||||
|
|
||||||
class AdminCacheConfigAdapter(AdminApiAdapter):
|
class AdminCacheConfigAdapter(AdminApiAdapter):
|
||||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||||
from src.services.cache.affinity_manager import CacheAffinityManager
|
from src.services.cache.affinity_manager import CacheAffinityManager
|
||||||
from src.config.constants import ConcurrencyDefaults
|
from src.config.constants import ConcurrencyDefaults
|
||||||
from src.services.rate_limit.adaptive_reservation import get_adaptive_reservation_manager
|
from src.services.rate_limit.adaptive_reservation import get_adaptive_reservation_manager
|
||||||
@@ -1260,7 +1260,7 @@ async def clear_provider_model_mapping_cache(
|
|||||||
|
|
||||||
|
|
||||||
class AdminModelMappingCacheStatsAdapter(AdminApiAdapter):
|
class AdminModelMappingCacheStatsAdapter(AdminApiAdapter):
|
||||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||||
import json
|
import json
|
||||||
|
|
||||||
from src.clients.redis_client import get_redis_client
|
from src.clients.redis_client import get_redis_client
|
||||||
@@ -1510,7 +1510,7 @@ class AdminModelMappingCacheStatsAdapter(AdminApiAdapter):
|
|||||||
|
|
||||||
|
|
||||||
class AdminClearAllModelMappingCacheAdapter(AdminApiAdapter):
|
class AdminClearAllModelMappingCacheAdapter(AdminApiAdapter):
|
||||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||||
from src.clients.redis_client import get_redis_client
|
from src.clients.redis_client import get_redis_client
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -1552,7 +1552,7 @@ class AdminClearAllModelMappingCacheAdapter(AdminApiAdapter):
|
|||||||
class AdminClearModelMappingCacheByNameAdapter(AdminApiAdapter):
|
class AdminClearModelMappingCacheByNameAdapter(AdminApiAdapter):
|
||||||
model_name: str
|
model_name: str
|
||||||
|
|
||||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||||
from src.clients.redis_client import get_redis_client
|
from src.clients.redis_client import get_redis_client
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -1599,7 +1599,7 @@ class AdminClearProviderModelMappingCacheAdapter(AdminApiAdapter):
|
|||||||
provider_id: str
|
provider_id: str
|
||||||
global_model_id: str
|
global_model_id: str
|
||||||
|
|
||||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||||
from src.clients.redis_client import get_redis_client
|
from src.clients.redis_client import get_redis_client
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
|||||||
@@ -4,7 +4,6 @@
|
|||||||
|
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
from typing import List, Optional
|
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||||
from pydantic import BaseModel, ConfigDict
|
from pydantic import BaseModel, ConfigDict
|
||||||
@@ -28,29 +27,29 @@ class CandidateResponse(BaseModel):
|
|||||||
request_id: str
|
request_id: str
|
||||||
candidate_index: int
|
candidate_index: int
|
||||||
retry_index: int = 0 # 重试序号(从0开始)
|
retry_index: int = 0 # 重试序号(从0开始)
|
||||||
provider_id: Optional[str] = None
|
provider_id: str | None = None
|
||||||
provider_name: Optional[str] = None
|
provider_name: str | None = None
|
||||||
provider_website: Optional[str] = None # Provider 官网
|
provider_website: str | None = None # Provider 官网
|
||||||
endpoint_id: Optional[str] = None
|
endpoint_id: str | None = None
|
||||||
endpoint_name: Optional[str] = None # 端点显示名称(api_format)
|
endpoint_name: str | None = None # 端点显示名称(api_format)
|
||||||
key_id: Optional[str] = None
|
key_id: str | None = None
|
||||||
key_name: Optional[str] = None # 密钥名称
|
key_name: str | None = None # 密钥名称
|
||||||
key_preview: Optional[str] = None # 密钥脱敏预览(如 sk-***abc)
|
key_preview: str | None = None # 密钥脱敏预览(如 sk-***abc)
|
||||||
key_capabilities: Optional[dict] = None # Key 支持的能力
|
key_capabilities: dict | None = None # Key 支持的能力
|
||||||
required_capabilities: Optional[dict] = None # 请求实际需要的能力标签
|
required_capabilities: dict | None = None # 请求实际需要的能力标签
|
||||||
status: str # 'pending', 'success', 'failed', 'skipped'
|
status: str # 'pending', 'success', 'failed', 'skipped'
|
||||||
skip_reason: Optional[str] = None
|
skip_reason: str | None = None
|
||||||
is_cached: bool = False
|
is_cached: bool = False
|
||||||
# 执行结果字段
|
# 执行结果字段
|
||||||
status_code: Optional[int] = None
|
status_code: int | None = None
|
||||||
error_type: Optional[str] = None
|
error_type: str | None = None
|
||||||
error_message: Optional[str] = None
|
error_message: str | None = None
|
||||||
latency_ms: Optional[int] = None
|
latency_ms: int | None = None
|
||||||
concurrent_requests: Optional[int] = None
|
concurrent_requests: int | None = None
|
||||||
extra_data: Optional[dict] = None
|
extra_data: dict | None = None
|
||||||
created_at: datetime
|
created_at: datetime
|
||||||
started_at: Optional[datetime] = None
|
started_at: datetime | None = None
|
||||||
finished_at: Optional[datetime] = None
|
finished_at: datetime | None = None
|
||||||
|
|
||||||
model_config = ConfigDict(from_attributes=True)
|
model_config = ConfigDict(from_attributes=True)
|
||||||
|
|
||||||
@@ -62,7 +61,7 @@ class RequestTraceResponse(BaseModel):
|
|||||||
total_candidates: int
|
total_candidates: int
|
||||||
final_status: str # 'success', 'failed', 'cancelled', 'streaming', 'pending'
|
final_status: str # 'success', 'failed', 'cancelled', 'streaming', 'pending'
|
||||||
total_latency_ms: int
|
total_latency_ms: int
|
||||||
candidates: List[CandidateResponse]
|
candidates: list[CandidateResponse]
|
||||||
|
|
||||||
|
|
||||||
@router.get("/{request_id}", response_model=RequestTraceResponse)
|
@router.get("/{request_id}", response_model=RequestTraceResponse)
|
||||||
@@ -253,7 +252,7 @@ class AdminGetRequestTraceAdapter(AdminApiAdapter):
|
|||||||
key_preview_map[k.id] = "***"
|
key_preview_map[k.id] = "***"
|
||||||
|
|
||||||
# 构建 candidate 响应列表
|
# 构建 candidate 响应列表
|
||||||
candidate_responses: List[CandidateResponse] = []
|
candidate_responses: list[CandidateResponse] = []
|
||||||
for candidate in candidates:
|
for candidate in candidates:
|
||||||
provider_name = (
|
provider_name = (
|
||||||
provider_map.get(candidate.provider_id) if candidate.provider_id else None
|
provider_map.get(candidate.provider_id) if candidate.provider_id else None
|
||||||
|
|||||||
@@ -9,7 +9,7 @@ Provider 操作 API 路由
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
from dataclasses import asdict, is_dataclass
|
from dataclasses import asdict, is_dataclass
|
||||||
from typing import Any, Dict, List, Optional
|
from typing import Any
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, HTTPException
|
from fastapi import APIRouter, Depends, HTTPException
|
||||||
from pydantic import BaseModel, Field
|
from pydantic import BaseModel, Field
|
||||||
@@ -18,9 +18,7 @@ from sqlalchemy.orm import Session
|
|||||||
from src.database import get_db
|
from src.database import get_db
|
||||||
from src.models.database import Provider, User
|
from src.models.database import Provider, User
|
||||||
from src.services.provider_ops import (
|
from src.services.provider_ops import (
|
||||||
ActionStatus,
|
|
||||||
ConnectorAuthType,
|
ConnectorAuthType,
|
||||||
ConnectorStatus,
|
|
||||||
ProviderActionType,
|
ProviderActionType,
|
||||||
ProviderOpsConfig,
|
ProviderOpsConfig,
|
||||||
ProviderOpsService,
|
ProviderOpsService,
|
||||||
@@ -40,46 +38,46 @@ class ArchitectureInfo(BaseModel):
|
|||||||
architecture_id: str
|
architecture_id: str
|
||||||
display_name: str
|
display_name: str
|
||||||
description: str
|
description: str
|
||||||
supported_auth_types: List[Dict[str, str]]
|
supported_auth_types: list[dict[str, str]]
|
||||||
supported_actions: List[Dict[str, Any]]
|
supported_actions: list[dict[str, Any]]
|
||||||
default_connector: Optional[str]
|
default_connector: str | None
|
||||||
|
|
||||||
|
|
||||||
class ConnectorConfigRequest(BaseModel):
|
class ConnectorConfigRequest(BaseModel):
|
||||||
"""连接器配置请求"""
|
"""连接器配置请求"""
|
||||||
|
|
||||||
auth_type: str = Field(..., description="认证类型")
|
auth_type: str = Field(..., description="认证类型")
|
||||||
config: Dict[str, Any] = Field(default_factory=dict, description="连接器配置")
|
config: dict[str, Any] = Field(default_factory=dict, description="连接器配置")
|
||||||
credentials: Dict[str, Any] = Field(default_factory=dict, description="凭据信息")
|
credentials: dict[str, Any] = Field(default_factory=dict, description="凭据信息")
|
||||||
|
|
||||||
|
|
||||||
class ActionConfigRequest(BaseModel):
|
class ActionConfigRequest(BaseModel):
|
||||||
"""操作配置请求"""
|
"""操作配置请求"""
|
||||||
|
|
||||||
enabled: bool = Field(True, description="是否启用")
|
enabled: bool = Field(True, description="是否启用")
|
||||||
config: Dict[str, Any] = Field(default_factory=dict, description="操作配置")
|
config: dict[str, Any] = Field(default_factory=dict, description="操作配置")
|
||||||
|
|
||||||
|
|
||||||
class SaveConfigRequest(BaseModel):
|
class SaveConfigRequest(BaseModel):
|
||||||
"""保存配置请求"""
|
"""保存配置请求"""
|
||||||
|
|
||||||
architecture_id: str = Field("generic_api", description="架构 ID")
|
architecture_id: str = Field("generic_api", description="架构 ID")
|
||||||
base_url: Optional[str] = Field(None, description="API 基础地址")
|
base_url: str | None = Field(None, description="API 基础地址")
|
||||||
connector: ConnectorConfigRequest
|
connector: ConnectorConfigRequest
|
||||||
actions: Dict[str, ActionConfigRequest] = Field(default_factory=dict)
|
actions: dict[str, ActionConfigRequest] = Field(default_factory=dict)
|
||||||
schedule: Dict[str, str] = Field(default_factory=dict, description="定时任务配置")
|
schedule: dict[str, str] = Field(default_factory=dict, description="定时任务配置")
|
||||||
|
|
||||||
|
|
||||||
class ConnectRequest(BaseModel):
|
class ConnectRequest(BaseModel):
|
||||||
"""连接请求"""
|
"""连接请求"""
|
||||||
|
|
||||||
credentials: Optional[Dict[str, Any]] = Field(None, description="凭据(可选,使用已保存的)")
|
credentials: dict[str, Any] | None = Field(None, description="凭据(可选,使用已保存的)")
|
||||||
|
|
||||||
|
|
||||||
class ExecuteActionRequest(BaseModel):
|
class ExecuteActionRequest(BaseModel):
|
||||||
"""执行操作请求"""
|
"""执行操作请求"""
|
||||||
|
|
||||||
config: Optional[Dict[str, Any]] = Field(None, description="操作配置(覆盖默认)")
|
config: dict[str, Any] | None = Field(None, description="操作配置(覆盖默认)")
|
||||||
|
|
||||||
|
|
||||||
class ConnectionStatusResponse(BaseModel):
|
class ConnectionStatusResponse(BaseModel):
|
||||||
@@ -87,9 +85,9 @@ class ConnectionStatusResponse(BaseModel):
|
|||||||
|
|
||||||
status: str
|
status: str
|
||||||
auth_type: str
|
auth_type: str
|
||||||
connected_at: Optional[str]
|
connected_at: str | None
|
||||||
expires_at: Optional[str]
|
expires_at: str | None
|
||||||
last_error: Optional[str]
|
last_error: str | None
|
||||||
|
|
||||||
|
|
||||||
class ActionResultResponse(BaseModel):
|
class ActionResultResponse(BaseModel):
|
||||||
@@ -97,10 +95,10 @@ class ActionResultResponse(BaseModel):
|
|||||||
|
|
||||||
status: str
|
status: str
|
||||||
action_type: str
|
action_type: str
|
||||||
data: Optional[Any]
|
data: Any | None
|
||||||
message: Optional[str]
|
message: str | None
|
||||||
executed_at: str
|
executed_at: str
|
||||||
response_time_ms: Optional[int]
|
response_time_ms: int | None
|
||||||
cache_ttl_seconds: int
|
cache_ttl_seconds: int
|
||||||
|
|
||||||
|
|
||||||
@@ -109,9 +107,9 @@ class ProviderOpsStatusResponse(BaseModel):
|
|||||||
|
|
||||||
provider_id: str
|
provider_id: str
|
||||||
is_configured: bool
|
is_configured: bool
|
||||||
architecture_id: Optional[str]
|
architecture_id: str | None
|
||||||
connection_status: ConnectionStatusResponse
|
connection_status: ConnectionStatusResponse
|
||||||
enabled_actions: List[str]
|
enabled_actions: list[str]
|
||||||
|
|
||||||
|
|
||||||
class ProviderOpsConfigResponse(BaseModel):
|
class ProviderOpsConfigResponse(BaseModel):
|
||||||
@@ -119,17 +117,17 @@ class ProviderOpsConfigResponse(BaseModel):
|
|||||||
|
|
||||||
provider_id: str
|
provider_id: str
|
||||||
is_configured: bool
|
is_configured: bool
|
||||||
architecture_id: Optional[str] = None
|
architecture_id: str | None = None
|
||||||
base_url: Optional[str] = None
|
base_url: str | None = None
|
||||||
connector: Optional[Dict[str, Any]] = None # 脱敏后的连接器配置
|
connector: dict[str, Any] | None = None # 脱敏后的连接器配置
|
||||||
|
|
||||||
|
|
||||||
class VerifyAuthResponse(BaseModel):
|
class VerifyAuthResponse(BaseModel):
|
||||||
"""验证认证响应"""
|
"""验证认证响应"""
|
||||||
|
|
||||||
success: bool
|
success: bool
|
||||||
message: Optional[str] = None
|
message: str | None = None
|
||||||
data: Optional[Dict[str, Any]] = None
|
data: dict[str, Any] | None = None
|
||||||
|
|
||||||
|
|
||||||
# ==================== Helper Functions ====================
|
# ==================== Helper Functions ====================
|
||||||
@@ -147,7 +145,7 @@ def _serialize_data(data: Any) -> Any:
|
|||||||
# ==================== Routes ====================
|
# ==================== Routes ====================
|
||||||
|
|
||||||
|
|
||||||
@router.get("/architectures", response_model=List[ArchitectureInfo])
|
@router.get("/architectures", response_model=list[ArchitectureInfo])
|
||||||
async def list_architectures(_: User = Depends(require_admin)):
|
async def list_architectures(_: User = Depends(require_admin)):
|
||||||
"""获取所有可用的架构"""
|
"""获取所有可用的架构"""
|
||||||
registry = get_registry()
|
registry = get_registry()
|
||||||
@@ -490,7 +488,7 @@ async def checkin(
|
|||||||
|
|
||||||
@router.post("/batch/balance")
|
@router.post("/batch/balance")
|
||||||
async def batch_query_balance(
|
async def batch_query_balance(
|
||||||
provider_ids: Optional[List[str]] = None,
|
provider_ids: list[str] | None = None,
|
||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
_: User = Depends(require_admin),
|
_: User = Depends(require_admin),
|
||||||
):
|
):
|
||||||
|
|||||||
@@ -4,7 +4,6 @@ Provider Query API 端点
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
from typing import Optional
|
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
from fastapi import APIRouter, Depends, HTTPException
|
from fastapi import APIRouter, Depends, HTTPException
|
||||||
@@ -40,7 +39,7 @@ class ModelsQueryRequest(BaseModel):
|
|||||||
"""模型列表查询请求"""
|
"""模型列表查询请求"""
|
||||||
|
|
||||||
provider_id: str
|
provider_id: str
|
||||||
api_key_id: Optional[str] = None
|
api_key_id: str | None = None
|
||||||
force_refresh: bool = False # 强制刷新,跳过缓存
|
force_refresh: bool = False # 强制刷新,跳过缓存
|
||||||
|
|
||||||
|
|
||||||
@@ -49,11 +48,11 @@ class TestModelRequest(BaseModel):
|
|||||||
|
|
||||||
provider_id: str
|
provider_id: str
|
||||||
model_name: str
|
model_name: str
|
||||||
api_key_id: Optional[str] = None
|
api_key_id: str | None = None
|
||||||
endpoint_id: Optional[str] = None # 指定使用的端点ID
|
endpoint_id: str | None = None # 指定使用的端点ID
|
||||||
stream: bool = False
|
stream: bool = False
|
||||||
message: Optional[str] = "你好"
|
message: str | None = "你好"
|
||||||
api_format: Optional[str] = None # 指定使用的API格式,如果不指定则使用端点的默认格式
|
api_format: str | None = None # 指定使用的API格式,如果不指定则使用端点的默认格式
|
||||||
|
|
||||||
|
|
||||||
# ============ API Endpoints ============
|
# ============ API Endpoints ============
|
||||||
|
|||||||
@@ -3,7 +3,6 @@
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
from datetime import datetime, timedelta, timezone
|
from datetime import datetime, timedelta, timezone
|
||||||
from typing import Optional
|
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, HTTPException, Request
|
from fastapi import APIRouter, Depends, HTTPException, Request
|
||||||
from fastapi.responses import JSONResponse
|
from fastapi.responses import JSONResponse
|
||||||
@@ -25,11 +24,11 @@ pipeline = ApiRequestPipeline()
|
|||||||
|
|
||||||
class ProviderBillingUpdate(BaseModel):
|
class ProviderBillingUpdate(BaseModel):
|
||||||
billing_type: ProviderBillingType
|
billing_type: ProviderBillingType
|
||||||
monthly_quota_usd: Optional[float] = None
|
monthly_quota_usd: float | None = None
|
||||||
quota_reset_day: int = Field(default=30, ge=1, le=365) # 重置周期(天数)
|
quota_reset_day: int = Field(default=30, ge=1, le=365) # 重置周期(天数)
|
||||||
quota_last_reset_at: Optional[str] = None # 当前周期开始时间
|
quota_last_reset_at: str | None = None # 当前周期开始时间
|
||||||
quota_expires_at: Optional[str] = None
|
quota_expires_at: str | None = None
|
||||||
rpm_limit: Optional[int] = Field(default=None, ge=0)
|
rpm_limit: int | None = Field(default=None, ge=0)
|
||||||
provider_priority: int = Field(default=100, ge=0, le=200)
|
provider_priority: int = Field(default=100, ge=0, le=200)
|
||||||
|
|
||||||
|
|
||||||
@@ -163,13 +162,12 @@ class AdminProviderBillingAdapter(AdminApiAdapter):
|
|||||||
provider.quota_reset_day = config.quota_reset_day
|
provider.quota_reset_day = config.quota_reset_day
|
||||||
provider.provider_priority = config.provider_priority
|
provider.provider_priority = config.provider_priority
|
||||||
|
|
||||||
from dateutil import parser
|
|
||||||
from sqlalchemy import func
|
from sqlalchemy import func
|
||||||
|
|
||||||
from src.models.database import Usage
|
from src.models.database import Usage
|
||||||
|
|
||||||
if config.quota_last_reset_at:
|
if config.quota_last_reset_at:
|
||||||
new_reset_at = parser.parse(config.quota_last_reset_at)
|
new_reset_at = datetime.fromisoformat(config.quota_last_reset_at)
|
||||||
# 确保有时区信息,如果没有则假设为 UTC
|
# 确保有时区信息,如果没有则假设为 UTC
|
||||||
if new_reset_at.tzinfo is None:
|
if new_reset_at.tzinfo is None:
|
||||||
new_reset_at = new_reset_at.replace(tzinfo=timezone.utc)
|
new_reset_at = new_reset_at.replace(tzinfo=timezone.utc)
|
||||||
@@ -188,7 +186,7 @@ class AdminProviderBillingAdapter(AdminApiAdapter):
|
|||||||
logger.info(f"Synced usage for provider {provider.name}: ${period_usage:.4f} since {new_reset_at}")
|
logger.info(f"Synced usage for provider {provider.name}: ${period_usage:.4f} since {new_reset_at}")
|
||||||
|
|
||||||
if config.quota_expires_at:
|
if config.quota_expires_at:
|
||||||
expires_at = parser.parse(config.quota_expires_at)
|
expires_at = datetime.fromisoformat(config.quota_expires_at)
|
||||||
# 确保有时区信息,如果没有则假设为 UTC
|
# 确保有时区信息,如果没有则假设为 UTC
|
||||||
if expires_at.tzinfo is None:
|
if expires_at.tzinfo is None:
|
||||||
expires_at = expires_at.replace(tzinfo=timezone.utc)
|
expires_at = expires_at.replace(tzinfo=timezone.utc)
|
||||||
|
|||||||
@@ -3,7 +3,7 @@ Provider 模型管理 API
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import Any, Dict, List, Optional
|
from typing import Any
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, Request
|
from fastapi import APIRouter, Depends, Request
|
||||||
from sqlalchemy.orm import Session, joinedload
|
from sqlalchemy.orm import Session, joinedload
|
||||||
@@ -40,15 +40,15 @@ router = APIRouter(tags=["Model Management"])
|
|||||||
pipeline = ApiRequestPipeline()
|
pipeline = ApiRequestPipeline()
|
||||||
|
|
||||||
|
|
||||||
@router.get("/{provider_id}/models", response_model=List[ModelResponse])
|
@router.get("/{provider_id}/models", response_model=list[ModelResponse])
|
||||||
async def list_provider_models(
|
async def list_provider_models(
|
||||||
provider_id: str,
|
provider_id: str,
|
||||||
request: Request,
|
request: Request,
|
||||||
is_active: Optional[bool] = None,
|
is_active: bool | None = None,
|
||||||
skip: int = 0,
|
skip: int = 0,
|
||||||
limit: int = 100,
|
limit: int = 100,
|
||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
) -> List[ModelResponse]:
|
) -> list[ModelResponse]:
|
||||||
"""
|
"""
|
||||||
获取提供商的所有模型
|
获取提供商的所有模型
|
||||||
|
|
||||||
@@ -222,13 +222,13 @@ async def delete_provider_model(
|
|||||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||||
|
|
||||||
|
|
||||||
@router.post("/{provider_id}/models/batch", response_model=List[ModelResponse])
|
@router.post("/{provider_id}/models/batch", response_model=list[ModelResponse])
|
||||||
async def batch_create_provider_models(
|
async def batch_create_provider_models(
|
||||||
provider_id: str,
|
provider_id: str,
|
||||||
models_data: List[ModelCreate],
|
models_data: list[ModelCreate],
|
||||||
request: Request,
|
request: Request,
|
||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
) -> List[ModelResponse]:
|
) -> list[ModelResponse]:
|
||||||
"""
|
"""
|
||||||
批量创建模型
|
批量创建模型
|
||||||
|
|
||||||
@@ -375,7 +375,7 @@ async def import_models_from_upstream(
|
|||||||
@dataclass
|
@dataclass
|
||||||
class AdminListProviderModelsAdapter(AdminApiAdapter):
|
class AdminListProviderModelsAdapter(AdminApiAdapter):
|
||||||
provider_id: str
|
provider_id: str
|
||||||
is_active: Optional[bool]
|
is_active: bool | None
|
||||||
skip: int
|
skip: int
|
||||||
limit: int
|
limit: int
|
||||||
|
|
||||||
@@ -482,7 +482,7 @@ class AdminDeleteProviderModelAdapter(AdminApiAdapter):
|
|||||||
@dataclass
|
@dataclass
|
||||||
class AdminBatchCreateModelsAdapter(AdminApiAdapter):
|
class AdminBatchCreateModelsAdapter(AdminApiAdapter):
|
||||||
provider_id: str
|
provider_id: str
|
||||||
models_data: List[ModelCreate]
|
models_data: list[ModelCreate]
|
||||||
|
|
||||||
async def handle(self, context): # type: ignore[override]
|
async def handle(self, context): # type: ignore[override]
|
||||||
db = context.db
|
db = context.db
|
||||||
@@ -525,7 +525,7 @@ class AdminGetProviderAvailableSourceModelsAdapter(AdminApiAdapter):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# 2. 构建以 GlobalModel 为主键的字典
|
# 2. 构建以 GlobalModel 为主键的字典
|
||||||
global_models_dict: Dict[str, Dict[str, Any]] = {}
|
global_models_dict: dict[str, dict[str, Any]] = {}
|
||||||
|
|
||||||
for model in models:
|
for model in models:
|
||||||
global_model = model.global_model
|
global_model = model.global_model
|
||||||
|
|||||||
@@ -2,7 +2,6 @@
|
|||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
from datetime import datetime, timezone
|
from datetime import datetime, timezone
|
||||||
from typing import Dict, List, Optional
|
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, Query, Request
|
from fastapi import APIRouter, Depends, Query, Request
|
||||||
from pydantic import BaseModel, ConfigDict, Field, ValidationError
|
from pydantic import BaseModel, ConfigDict, Field, ValidationError
|
||||||
@@ -48,7 +47,7 @@ class MappingMatchingGlobalModel(BaseModel):
|
|||||||
global_model_name: str
|
global_model_name: str
|
||||||
display_name: str
|
display_name: str
|
||||||
is_active: bool
|
is_active: bool
|
||||||
matched_models: List[MappingMatchedModel] = Field(
|
matched_models: list[MappingMatchedModel] = Field(
|
||||||
default_factory=list, description="匹配到的模型列表"
|
default_factory=list, description="匹配到的模型列表"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -62,8 +61,8 @@ class MappingMatchingKey(BaseModel):
|
|||||||
key_name: str
|
key_name: str
|
||||||
masked_key: str
|
masked_key: str
|
||||||
is_active: bool
|
is_active: bool
|
||||||
allowed_models: List[str] = Field(default_factory=list, description="Key 的模型白名单")
|
allowed_models: list[str] = Field(default_factory=list, description="Key 的模型白名单")
|
||||||
matching_global_models: List[MappingMatchingGlobalModel] = Field(
|
matching_global_models: list[MappingMatchingGlobalModel] = Field(
|
||||||
default_factory=list, description="匹配到的 GlobalModel 列表"
|
default_factory=list, description="匹配到的 GlobalModel 列表"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -75,7 +74,7 @@ class ProviderMappingPreviewResponse(BaseModel):
|
|||||||
|
|
||||||
provider_id: str
|
provider_id: str
|
||||||
provider_name: str
|
provider_name: str
|
||||||
keys: List[MappingMatchingKey] = Field(
|
keys: list[MappingMatchingKey] = Field(
|
||||||
default_factory=list, description="有白名单配置且匹配到映射的 Key 列表"
|
default_factory=list, description="有白名单配置且匹配到映射的 Key 列表"
|
||||||
)
|
)
|
||||||
total_keys: int = Field(0, description="有匹配结果的 Key 数量")
|
total_keys: int = Field(0, description="有匹配结果的 Key 数量")
|
||||||
@@ -95,7 +94,7 @@ async def list_providers(
|
|||||||
request: Request,
|
request: Request,
|
||||||
skip: int = Query(0, ge=0),
|
skip: int = Query(0, ge=0),
|
||||||
limit: int = Query(100, ge=1, le=500),
|
limit: int = Query(100, ge=1, le=500),
|
||||||
is_active: Optional[bool] = None,
|
is_active: bool | None = None,
|
||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
@@ -209,7 +208,7 @@ async def delete_provider(provider_id: str, request: Request, db: Session = Depe
|
|||||||
|
|
||||||
|
|
||||||
class AdminListProvidersAdapter(AdminApiAdapter):
|
class AdminListProvidersAdapter(AdminApiAdapter):
|
||||||
def __init__(self, skip: int, limit: int, is_active: Optional[bool]):
|
def __init__(self, skip: int, limit: int, is_active: bool | None):
|
||||||
self.skip = skip
|
self.skip = skip
|
||||||
self.limit = limit
|
self.limit = limit
|
||||||
self.is_active = is_active
|
self.is_active = is_active
|
||||||
@@ -473,7 +472,7 @@ async def get_provider_mapping_preview(
|
|||||||
pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode),
|
pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode),
|
||||||
timeout=MAPPING_PREVIEW_TIMEOUT_SECONDS,
|
timeout=MAPPING_PREVIEW_TIMEOUT_SECONDS,
|
||||||
)
|
)
|
||||||
except asyncio.TimeoutError:
|
except TimeoutError:
|
||||||
logger.warning(f"映射预览超时: provider_id={provider_id}")
|
logger.warning(f"映射预览超时: provider_id={provider_id}")
|
||||||
raise InvalidRequestException("映射预览超时,请简化配置或稍后重试")
|
raise InvalidRequestException("映射预览超时,请简化配置或稍后重试")
|
||||||
|
|
||||||
@@ -565,7 +564,7 @@ class AdminGetProviderMappingPreviewAdapter(AdminApiAdapter):
|
|||||||
truncated_models = total_models_with_mappings - MAPPING_PREVIEW_MAX_MODELS
|
truncated_models = total_models_with_mappings - MAPPING_PREVIEW_MAX_MODELS
|
||||||
|
|
||||||
# 构建有映射配置的 GlobalModel 映射
|
# 构建有映射配置的 GlobalModel 映射
|
||||||
models_with_mappings: Dict[str, tuple] = {} # id -> (model_info, mappings)
|
models_with_mappings: dict[str, tuple] = {} # id -> (model_info, mappings)
|
||||||
for gm in global_models:
|
for gm in global_models:
|
||||||
config = gm.config or {}
|
config = gm.config or {}
|
||||||
mappings = config.get("model_mappings", [])
|
mappings = config.get("model_mappings", [])
|
||||||
@@ -585,7 +584,7 @@ class AdminGetProviderMappingPreviewAdapter(AdminApiAdapter):
|
|||||||
truncated_models=0,
|
truncated_models=0,
|
||||||
)
|
)
|
||||||
|
|
||||||
key_infos: List[MappingMatchingKey] = []
|
key_infos: list[MappingMatchingKey] = []
|
||||||
total_matches = 0
|
total_matches = 0
|
||||||
|
|
||||||
# 创建 CryptoService 实例
|
# 创建 CryptoService 实例
|
||||||
@@ -611,10 +610,10 @@ class AdminGetProviderMappingPreviewAdapter(AdminApiAdapter):
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
# 查找匹配的 GlobalModel
|
# 查找匹配的 GlobalModel
|
||||||
matching_global_models: List[MappingMatchingGlobalModel] = []
|
matching_global_models: list[MappingMatchingGlobalModel] = []
|
||||||
|
|
||||||
for gm_id, (gm, mappings) in models_with_mappings.items():
|
for gm_id, (gm, mappings) in models_with_mappings.items():
|
||||||
matched_models: List[MappingMatchedModel] = []
|
matched_models: list[MappingMatchedModel] = []
|
||||||
|
|
||||||
for allowed_model in allowed_models_list:
|
for allowed_model in allowed_models_list:
|
||||||
for mapping_pattern in mappings:
|
for mapping_pattern in mappings:
|
||||||
|
|||||||
@@ -4,7 +4,6 @@ Provider 摘要与健康监控 API
|
|||||||
|
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from datetime import datetime, timedelta, timezone
|
from datetime import datetime, timedelta, timezone
|
||||||
from typing import Dict, List
|
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, Query, Request
|
from fastapi import APIRouter, Depends, Query, Request
|
||||||
from sqlalchemy import case, func
|
from sqlalchemy import case, func
|
||||||
@@ -35,11 +34,11 @@ router = APIRouter(tags=["Provider Summary"])
|
|||||||
pipeline = ApiRequestPipeline()
|
pipeline = ApiRequestPipeline()
|
||||||
|
|
||||||
|
|
||||||
@router.get("/summary", response_model=List[ProviderWithEndpointsSummary])
|
@router.get("/summary", response_model=list[ProviderWithEndpointsSummary])
|
||||||
async def get_providers_summary(
|
async def get_providers_summary(
|
||||||
request: Request,
|
request: Request,
|
||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
) -> List[ProviderWithEndpointsSummary]:
|
) -> list[ProviderWithEndpointsSummary]:
|
||||||
"""
|
"""
|
||||||
获取所有提供商摘要信息
|
获取所有提供商摘要信息
|
||||||
|
|
||||||
@@ -381,8 +380,8 @@ class AdminProviderHealthMonitorAdapter(AdminApiAdapter):
|
|||||||
)
|
)
|
||||||
attempts = attempts_query.limit(limit_rows).all()
|
attempts = attempts_query.limit(limit_rows).all()
|
||||||
|
|
||||||
buffered_attempts: Dict[str, List[RequestCandidate]] = {eid: [] for eid in endpoint_ids}
|
buffered_attempts: dict[str, list[RequestCandidate]] = {eid: [] for eid in endpoint_ids}
|
||||||
counters: Dict[str, int] = {eid: 0 for eid in endpoint_ids}
|
counters: dict[str, int] = {eid: 0 for eid in endpoint_ids}
|
||||||
|
|
||||||
for attempt in attempts:
|
for attempt in attempts:
|
||||||
if not attempt.endpoint_id or attempt.endpoint_id not in buffered_attempts:
|
if not attempt.endpoint_id or attempt.endpoint_id not in buffered_attempts:
|
||||||
@@ -392,10 +391,10 @@ class AdminProviderHealthMonitorAdapter(AdminApiAdapter):
|
|||||||
buffered_attempts[attempt.endpoint_id].append(attempt)
|
buffered_attempts[attempt.endpoint_id].append(attempt)
|
||||||
counters[attempt.endpoint_id] += 1
|
counters[attempt.endpoint_id] += 1
|
||||||
|
|
||||||
endpoint_monitors: List[EndpointHealthMonitor] = []
|
endpoint_monitors: list[EndpointHealthMonitor] = []
|
||||||
for endpoint in endpoints:
|
for endpoint in endpoints:
|
||||||
attempt_list = list(reversed(buffered_attempts.get(endpoint.id, [])))
|
attempt_list = list(reversed(buffered_attempts.get(endpoint.id, [])))
|
||||||
events: List[EndpointHealthEvent] = []
|
events: list[EndpointHealthEvent] = []
|
||||||
for attempt in attempt_list:
|
for attempt in attempt_list:
|
||||||
event_timestamp = attempt.finished_at or attempt.started_at or attempt.created_at
|
event_timestamp = attempt.finished_at or attempt.started_at or attempt.created_at
|
||||||
events.append(
|
events.append(
|
||||||
|
|||||||
@@ -4,8 +4,6 @@ IP 安全管理接口
|
|||||||
提供 IP 黑白名单管理和速率限制统计
|
提供 IP 黑白名单管理和速率限制统计
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from typing import List, Optional
|
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||||
from pydantic import BaseModel, Field, ValidationError
|
from pydantic import BaseModel, Field, ValidationError
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
@@ -14,7 +12,6 @@ from src.api.base.adapter import ApiMode
|
|||||||
from src.api.base.authenticated_adapter import AuthenticatedApiAdapter
|
from src.api.base.authenticated_adapter import AuthenticatedApiAdapter
|
||||||
from src.api.base.pipeline import ApiRequestPipeline
|
from src.api.base.pipeline import ApiRequestPipeline
|
||||||
from src.core.exceptions import InvalidRequestException, translate_pydantic_error
|
from src.core.exceptions import InvalidRequestException, translate_pydantic_error
|
||||||
from src.core.logger import logger
|
|
||||||
from src.database import get_db
|
from src.database import get_db
|
||||||
from src.services.rate_limit.ip_limiter import IPRateLimiter
|
from src.services.rate_limit.ip_limiter import IPRateLimiter
|
||||||
|
|
||||||
@@ -30,7 +27,7 @@ class AddIPToBlacklistRequest(BaseModel):
|
|||||||
|
|
||||||
ip_address: str = Field(..., description="IP 地址")
|
ip_address: str = Field(..., description="IP 地址")
|
||||||
reason: str = Field(..., min_length=1, max_length=200, description="加入黑名单的原因")
|
reason: str = Field(..., min_length=1, max_length=200, description="加入黑名单的原因")
|
||||||
ttl: Optional[int] = Field(None, gt=0, description="过期时间(秒),None 表示永久")
|
ttl: int | None = Field(None, gt=0, description="过期时间(秒),None 表示永久")
|
||||||
|
|
||||||
|
|
||||||
class RemoveIPFromBlacklistRequest(BaseModel):
|
class RemoveIPFromBlacklistRequest(BaseModel):
|
||||||
|
|||||||
@@ -1,10 +1,8 @@
|
|||||||
"""系统设置API端点。"""
|
"""系统设置API端点。"""
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import copy
|
import copy
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import Optional
|
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||||
from pydantic import ValidationError
|
from pydantic import ValidationError
|
||||||
@@ -649,9 +647,8 @@ class AdminSystemStatsAdapter(AdminApiAdapter):
|
|||||||
class AdminTriggerCleanupAdapter(AdminApiAdapter):
|
class AdminTriggerCleanupAdapter(AdminApiAdapter):
|
||||||
async def handle(self, context): # type: ignore[override]
|
async def handle(self, context): # type: ignore[override]
|
||||||
"""手动触发清理任务"""
|
"""手动触发清理任务"""
|
||||||
from datetime import datetime, timedelta, timezone
|
from datetime import datetime, timezone
|
||||||
|
|
||||||
from sqlalchemy import func
|
|
||||||
|
|
||||||
from src.services.system.maintenance_scheduler import get_maintenance_scheduler
|
from src.services.system.maintenance_scheduler import get_maintenance_scheduler
|
||||||
|
|
||||||
|
|||||||
@@ -3,7 +3,6 @@
|
|||||||
from collections import defaultdict
|
from collections import defaultdict
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
from typing import Optional
|
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||||
from sqlalchemy import func
|
from sqlalchemy import func
|
||||||
@@ -34,8 +33,8 @@ pipeline = ApiRequestPipeline()
|
|||||||
async def get_usage_aggregation(
|
async def get_usage_aggregation(
|
||||||
request: Request,
|
request: Request,
|
||||||
group_by: str = Query(..., description="Aggregation dimension: model, user, provider, or api_format"),
|
group_by: str = Query(..., description="Aggregation dimension: model, user, provider, or api_format"),
|
||||||
start_date: Optional[datetime] = None,
|
start_date: datetime | None = None,
|
||||||
end_date: Optional[datetime] = None,
|
end_date: datetime | None = None,
|
||||||
limit: int = Query(20, ge=1, le=100),
|
limit: int = Query(20, ge=1, le=100),
|
||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
):
|
):
|
||||||
@@ -75,8 +74,8 @@ async def get_usage_aggregation(
|
|||||||
@router.get("/stats")
|
@router.get("/stats")
|
||||||
async def get_usage_stats(
|
async def get_usage_stats(
|
||||||
request: Request,
|
request: Request,
|
||||||
start_date: Optional[datetime] = None,
|
start_date: datetime | None = None,
|
||||||
end_date: Optional[datetime] = None,
|
end_date: datetime | None = None,
|
||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
@@ -122,14 +121,14 @@ async def get_activity_heatmap(
|
|||||||
@router.get("/records")
|
@router.get("/records")
|
||||||
async def get_usage_records(
|
async def get_usage_records(
|
||||||
request: Request,
|
request: Request,
|
||||||
start_date: Optional[datetime] = None,
|
start_date: datetime | None = None,
|
||||||
end_date: Optional[datetime] = None,
|
end_date: datetime | None = None,
|
||||||
search: Optional[str] = None, # 通用搜索:用户名、密钥名、模型名、提供商名
|
search: str | None = None, # 通用搜索:用户名、密钥名、模型名、提供商名
|
||||||
user_id: Optional[str] = None,
|
user_id: str | None = None,
|
||||||
username: Optional[str] = None,
|
username: str | None = None,
|
||||||
model: Optional[str] = None,
|
model: str | None = None,
|
||||||
provider: Optional[str] = None,
|
provider: str | None = None,
|
||||||
status: Optional[str] = None, # stream, standard, error
|
status: str | None = None, # stream, standard, error
|
||||||
limit: int = Query(100, ge=1, le=500),
|
limit: int = Query(100, ge=1, le=500),
|
||||||
offset: int = Query(0, ge=0),
|
offset: int = Query(0, ge=0),
|
||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
@@ -179,7 +178,7 @@ async def get_usage_records(
|
|||||||
@router.get("/active")
|
@router.get("/active")
|
||||||
async def get_active_requests(
|
async def get_active_requests(
|
||||||
request: Request,
|
request: Request,
|
||||||
ids: Optional[str] = Query(None, description="逗号分隔的请求 ID 列表,用于查询特定请求的状态"),
|
ids: str | None = Query(None, description="逗号分隔的请求 ID 列表,用于查询特定请求的状态"),
|
||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
@@ -259,7 +258,7 @@ async def get_usage_detail(
|
|||||||
|
|
||||||
|
|
||||||
class AdminUsageStatsAdapter(AdminApiAdapter):
|
class AdminUsageStatsAdapter(AdminApiAdapter):
|
||||||
def __init__(self, start_date: Optional[datetime], end_date: Optional[datetime]):
|
def __init__(self, start_date: datetime | None, end_date: datetime | None):
|
||||||
self.start_date = start_date
|
self.start_date = start_date
|
||||||
self.end_date = end_date
|
self.end_date = end_date
|
||||||
|
|
||||||
@@ -339,7 +338,7 @@ class AdminActivityHeatmapAdapter(AdminApiAdapter):
|
|||||||
|
|
||||||
|
|
||||||
class AdminUsageByModelAdapter(AdminApiAdapter):
|
class AdminUsageByModelAdapter(AdminApiAdapter):
|
||||||
def __init__(self, start_date: Optional[datetime], end_date: Optional[datetime], limit: int):
|
def __init__(self, start_date: datetime | None, end_date: datetime | None, limit: int):
|
||||||
self.start_date = start_date
|
self.start_date = start_date
|
||||||
self.end_date = end_date
|
self.end_date = end_date
|
||||||
self.limit = limit
|
self.limit = limit
|
||||||
@@ -386,7 +385,7 @@ class AdminUsageByModelAdapter(AdminApiAdapter):
|
|||||||
|
|
||||||
|
|
||||||
class AdminUsageByUserAdapter(AdminApiAdapter):
|
class AdminUsageByUserAdapter(AdminApiAdapter):
|
||||||
def __init__(self, start_date: Optional[datetime], end_date: Optional[datetime], limit: int):
|
def __init__(self, start_date: datetime | None, end_date: datetime | None, limit: int):
|
||||||
self.start_date = start_date
|
self.start_date = start_date
|
||||||
self.end_date = end_date
|
self.end_date = end_date
|
||||||
self.limit = limit
|
self.limit = limit
|
||||||
@@ -436,7 +435,7 @@ class AdminUsageByUserAdapter(AdminApiAdapter):
|
|||||||
|
|
||||||
|
|
||||||
class AdminUsageByProviderAdapter(AdminApiAdapter):
|
class AdminUsageByProviderAdapter(AdminApiAdapter):
|
||||||
def __init__(self, start_date: Optional[datetime], end_date: Optional[datetime], limit: int):
|
def __init__(self, start_date: datetime | None, end_date: datetime | None, limit: int):
|
||||||
self.start_date = start_date
|
self.start_date = start_date
|
||||||
self.end_date = end_date
|
self.end_date = end_date
|
||||||
self.limit = limit
|
self.limit = limit
|
||||||
@@ -446,7 +445,7 @@ class AdminUsageByProviderAdapter(AdminApiAdapter):
|
|||||||
|
|
||||||
# 从 request_candidates 表统计每个 Provider 的尝试次数和成功率
|
# 从 request_candidates 表统计每个 Provider 的尝试次数和成功率
|
||||||
# 这样可以正确统计 Fallback 场景(一个请求可能尝试多个 Provider)
|
# 这样可以正确统计 Fallback 场景(一个请求可能尝试多个 Provider)
|
||||||
from sqlalchemy import case, Integer
|
from sqlalchemy import case
|
||||||
|
|
||||||
attempt_query = db.query(
|
attempt_query = db.query(
|
||||||
RequestCandidate.provider_id,
|
RequestCandidate.provider_id,
|
||||||
@@ -550,7 +549,7 @@ class AdminUsageByProviderAdapter(AdminApiAdapter):
|
|||||||
|
|
||||||
|
|
||||||
class AdminUsageByApiFormatAdapter(AdminApiAdapter):
|
class AdminUsageByApiFormatAdapter(AdminApiAdapter):
|
||||||
def __init__(self, start_date: Optional[datetime], end_date: Optional[datetime], limit: int):
|
def __init__(self, start_date: datetime | None, end_date: datetime | None, limit: int):
|
||||||
self.start_date = start_date
|
self.start_date = start_date
|
||||||
self.end_date = end_date
|
self.end_date = end_date
|
||||||
self.limit = limit
|
self.limit = limit
|
||||||
@@ -608,14 +607,14 @@ class AdminUsageByApiFormatAdapter(AdminApiAdapter):
|
|||||||
class AdminUsageRecordsAdapter(AdminApiAdapter):
|
class AdminUsageRecordsAdapter(AdminApiAdapter):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
start_date: Optional[datetime],
|
start_date: datetime | None,
|
||||||
end_date: Optional[datetime],
|
end_date: datetime | None,
|
||||||
search: Optional[str],
|
search: str | None,
|
||||||
user_id: Optional[str],
|
user_id: str | None,
|
||||||
username: Optional[str],
|
username: str | None,
|
||||||
model: Optional[str],
|
model: str | None,
|
||||||
provider: Optional[str],
|
provider: str | None,
|
||||||
status: Optional[str],
|
status: str | None,
|
||||||
limit: int,
|
limit: int,
|
||||||
offset: int,
|
offset: int,
|
||||||
):
|
):
|
||||||
@@ -744,7 +743,7 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
|
|||||||
|
|
||||||
for req_id, candidates in request_candidates.items():
|
for req_id, candidates in request_candidates.items():
|
||||||
# 提取所有不同的 candidate_index
|
# 提取所有不同的 candidate_index
|
||||||
unique_candidates = set(c[0] for c in candidates)
|
unique_candidates = {c[0] for c in candidates}
|
||||||
# 如果有多个不同的 candidate_index,说明发生了 Fallback(Provider 切换)
|
# 如果有多个不同的 candidate_index,说明发生了 Fallback(Provider 切换)
|
||||||
fallback_map[req_id] = len(unique_candidates) > 1
|
fallback_map[req_id] = len(unique_candidates) > 1
|
||||||
|
|
||||||
@@ -877,7 +876,7 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
|
|||||||
class AdminActiveRequestsAdapter(AdminApiAdapter):
|
class AdminActiveRequestsAdapter(AdminApiAdapter):
|
||||||
"""轻量级活跃请求状态查询适配器"""
|
"""轻量级活跃请求状态查询适配器"""
|
||||||
|
|
||||||
def __init__(self, ids: Optional[str]):
|
def __init__(self, ids: str | None):
|
||||||
self.ids = ids
|
self.ids = ids
|
||||||
|
|
||||||
async def handle(self, context): # type: ignore[override]
|
async def handle(self, context): # type: ignore[override]
|
||||||
@@ -1033,8 +1032,8 @@ class AdminUsageDetailAdapter(AdminApiAdapter):
|
|||||||
@router.get("/cache-affinity/ttl-analysis")
|
@router.get("/cache-affinity/ttl-analysis")
|
||||||
async def analyze_cache_affinity_ttl(
|
async def analyze_cache_affinity_ttl(
|
||||||
request: Request,
|
request: Request,
|
||||||
user_id: Optional[str] = Query(None, description="指定用户 ID"),
|
user_id: str | None = Query(None, description="指定用户 ID"),
|
||||||
api_key_id: Optional[str] = Query(None, description="指定 API Key ID"),
|
api_key_id: str | None = Query(None, description="指定 API Key ID"),
|
||||||
hours: int = Query(168, ge=1, le=720, description="分析最近多少小时的数据"),
|
hours: int = Query(168, ge=1, le=720, description="分析最近多少小时的数据"),
|
||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
):
|
):
|
||||||
@@ -1057,8 +1056,8 @@ async def analyze_cache_affinity_ttl(
|
|||||||
@router.get("/cache-affinity/hit-analysis")
|
@router.get("/cache-affinity/hit-analysis")
|
||||||
async def analyze_cache_hit(
|
async def analyze_cache_hit(
|
||||||
request: Request,
|
request: Request,
|
||||||
user_id: Optional[str] = Query(None, description="指定用户 ID"),
|
user_id: str | None = Query(None, description="指定用户 ID"),
|
||||||
api_key_id: Optional[str] = Query(None, description="指定 API Key ID"),
|
api_key_id: str | None = Query(None, description="指定 API Key ID"),
|
||||||
hours: int = Query(168, ge=1, le=720, description="分析最近多少小时的数据"),
|
hours: int = Query(168, ge=1, le=720, description="分析最近多少小时的数据"),
|
||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
):
|
):
|
||||||
@@ -1080,8 +1079,8 @@ class CacheAffinityTTLAnalysisAdapter(AdminApiAdapter):
|
|||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
user_id: Optional[str],
|
user_id: str | None,
|
||||||
api_key_id: Optional[str],
|
api_key_id: str | None,
|
||||||
hours: int,
|
hours: int,
|
||||||
):
|
):
|
||||||
self.user_id = user_id
|
self.user_id = user_id
|
||||||
@@ -1114,8 +1113,8 @@ class CacheHitAnalysisAdapter(AdminApiAdapter):
|
|||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
user_id: Optional[str],
|
user_id: str | None,
|
||||||
api_key_id: Optional[str],
|
api_key_id: str | None,
|
||||||
hours: int,
|
hours: int,
|
||||||
):
|
):
|
||||||
self.user_id = user_id
|
self.user_id = user_id
|
||||||
@@ -1147,7 +1146,7 @@ async def get_interval_timeline(
|
|||||||
request: Request,
|
request: Request,
|
||||||
hours: int = Query(24, ge=1, le=720, description="分析最近多少小时的数据"),
|
hours: int = Query(24, ge=1, le=720, description="分析最近多少小时的数据"),
|
||||||
limit: int = Query(10000, ge=100, le=50000, description="最大返回数据点数量"),
|
limit: int = Query(10000, ge=100, le=50000, description="最大返回数据点数量"),
|
||||||
user_id: Optional[str] = Query(None, description="指定用户 ID"),
|
user_id: str | None = Query(None, description="指定用户 ID"),
|
||||||
include_user_info: bool = Query(False, description="是否包含用户信息(用于管理员多用户视图)"),
|
include_user_info: bool = Query(False, description="是否包含用户信息(用于管理员多用户视图)"),
|
||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
):
|
):
|
||||||
@@ -1177,7 +1176,7 @@ class IntervalTimelineAdapter(AdminApiAdapter):
|
|||||||
self,
|
self,
|
||||||
hours: int,
|
hours: int,
|
||||||
limit: int,
|
limit: int,
|
||||||
user_id: Optional[str] = None,
|
user_id: str | None = None,
|
||||||
include_user_info: bool = False,
|
include_user_info: bool = False,
|
||||||
):
|
):
|
||||||
self.hours = hours
|
self.hours = hours
|
||||||
|
|||||||
@@ -1,7 +1,6 @@
|
|||||||
"""用户管理 API 端点。"""
|
"""用户管理 API 端点。"""
|
||||||
|
|
||||||
from datetime import datetime, timezone
|
from datetime import datetime, timezone
|
||||||
from typing import Optional
|
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request, status
|
from fastapi import APIRouter, Depends, HTTPException, Query, Request, status
|
||||||
from pydantic import ValidationError
|
from pydantic import ValidationError
|
||||||
@@ -48,8 +47,8 @@ async def list_users(
|
|||||||
request: Request,
|
request: Request,
|
||||||
skip: int = Query(0, ge=0, description="跳过记录数"),
|
skip: int = Query(0, ge=0, description="跳过记录数"),
|
||||||
limit: int = Query(100, ge=1, le=1000, description="返回记录数"),
|
limit: int = Query(100, ge=1, le=1000, description="返回记录数"),
|
||||||
role: Optional[str] = Query(None, description="按角色筛选(user/admin)"),
|
role: str | None = Query(None, description="按角色筛选(user/admin)"),
|
||||||
is_active: Optional[bool] = Query(None, description="按状态筛选"),
|
is_active: bool | None = Query(None, description="按状态筛选"),
|
||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
@@ -136,7 +135,7 @@ async def reset_user_quota(user_id: str, request: Request, db: Session = Depends
|
|||||||
async def get_user_api_keys(
|
async def get_user_api_keys(
|
||||||
user_id: str,
|
user_id: str,
|
||||||
request: Request,
|
request: Request,
|
||||||
is_active: Optional[bool] = Query(None, description="按状态筛选"),
|
is_active: bool | None = Query(None, description="按状态筛选"),
|
||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
@@ -274,7 +273,7 @@ class AdminCreateUserAdapter(AdminApiAdapter):
|
|||||||
|
|
||||||
|
|
||||||
class AdminListUsersAdapter(AdminApiAdapter):
|
class AdminListUsersAdapter(AdminApiAdapter):
|
||||||
def __init__(self, skip: int, limit: int, role: Optional[str], is_active: Optional[bool]):
|
def __init__(self, skip: int, limit: int, role: str | None, is_active: bool | None):
|
||||||
self.skip = skip
|
self.skip = skip
|
||||||
self.limit = limit
|
self.limit = limit
|
||||||
self.role = role
|
self.role = role
|
||||||
@@ -467,7 +466,7 @@ class AdminResetUserQuotaAdapter(AdminApiAdapter):
|
|||||||
class AdminGetUserKeysAdapter(AdminApiAdapter):
|
class AdminGetUserKeysAdapter(AdminApiAdapter):
|
||||||
"""获取用户的API Keys"""
|
"""获取用户的API Keys"""
|
||||||
|
|
||||||
def __init__(self, user_id: str, is_active: Optional[bool]):
|
def __init__(self, user_id: str, is_active: bool | None):
|
||||||
self.user_id = user_id
|
self.user_id = user_id
|
||||||
self.is_active = is_active
|
self.is_active = is_active
|
||||||
|
|
||||||
|
|||||||
@@ -1,9 +1,8 @@
|
|||||||
"""公告系统 API 端点。"""
|
"""公告系统 API 端点。"""
|
||||||
|
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import Optional
|
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request, status
|
from fastapi import APIRouter, Depends, Query, Request
|
||||||
from pydantic import ValidationError
|
from pydantic import ValidationError
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
@@ -12,7 +11,6 @@ from src.api.base.admin_adapter import AdminApiAdapter
|
|||||||
from src.api.base.authenticated_adapter import AuthenticatedApiAdapter
|
from src.api.base.authenticated_adapter import AuthenticatedApiAdapter
|
||||||
from src.api.base.pipeline import ApiRequestPipeline
|
from src.api.base.pipeline import ApiRequestPipeline
|
||||||
from src.core.exceptions import InvalidRequestException, translate_pydantic_error
|
from src.core.exceptions import InvalidRequestException, translate_pydantic_error
|
||||||
from src.core.logger import logger
|
|
||||||
from src.database import get_db
|
from src.database import get_db
|
||||||
from src.models.api import CreateAnnouncementRequest, UpdateAnnouncementRequest
|
from src.models.api import CreateAnnouncementRequest, UpdateAnnouncementRequest
|
||||||
from src.models.database import User
|
from src.models.database import User
|
||||||
@@ -251,7 +249,7 @@ class AnnouncementOptionalAuthAdapter(ApiAdapter):
|
|||||||
context.extra["optional_user"] = await self._resolve_optional_user(context)
|
context.extra["optional_user"] = await self._resolve_optional_user(context)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
async def _resolve_optional_user(self, context) -> Optional[User]:
|
async def _resolve_optional_user(self, context) -> User | None:
|
||||||
if context.user:
|
if context.user:
|
||||||
return context.user
|
return context.user
|
||||||
|
|
||||||
@@ -285,7 +283,7 @@ class AnnouncementOptionalAuthAdapter(ApiAdapter):
|
|||||||
except Exception:
|
except Exception:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def get_optional_user(self, context) -> Optional[User]:
|
def get_optional_user(self, context) -> User | None:
|
||||||
return context.extra.get("optional_user")
|
return context.extra.get("optional_user")
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -2,10 +2,9 @@
|
|||||||
认证相关API端点
|
认证相关API端点
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from typing import Optional, Tuple
|
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||||
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
|
from fastapi.security import HTTPBearer
|
||||||
from pydantic import ValidationError
|
from pydantic import ValidationError
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
@@ -42,7 +41,7 @@ from src.services.email import EmailSenderService, EmailVerificationService
|
|||||||
from src.utils.request_utils import get_client_ip, get_user_agent
|
from src.utils.request_utils import get_client_ip, get_user_agent
|
||||||
|
|
||||||
|
|
||||||
def validate_email_suffix(db: Session, email: str) -> Tuple[bool, Optional[str]]:
|
def validate_email_suffix(db: Session, email: str) -> tuple[bool, str | None]:
|
||||||
"""
|
"""
|
||||||
验证邮箱后缀是否允许注册
|
验证邮箱后缀是否允许注册
|
||||||
|
|
||||||
|
|||||||
@@ -1,8 +1,6 @@
|
|||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
from typing import Any, Dict, Optional
|
from typing import Any
|
||||||
|
|
||||||
from fastapi import Request, Response
|
from fastapi import Request, Response
|
||||||
|
|
||||||
@@ -23,7 +21,7 @@ class ApiAdapter(ABC):
|
|||||||
|
|
||||||
name: str = "base"
|
name: str = "base"
|
||||||
mode: ApiMode = ApiMode.STANDARD
|
mode: ApiMode = ApiMode.STANDARD
|
||||||
api_format: Optional[str] = None # 对应 Provider API 格式提示
|
api_format: str | None = None # 对应 Provider API 格式提示
|
||||||
audit_log_enabled: bool = True
|
audit_log_enabled: bool = True
|
||||||
audit_success_event = None
|
audit_success_event = None
|
||||||
audit_failure_event = None
|
audit_failure_event = None
|
||||||
@@ -36,7 +34,7 @@ class ApiAdapter(ABC):
|
|||||||
"""可选的授权钩子,默认允许通过。"""
|
"""可选的授权钩子,默认允许通过。"""
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def extract_api_key(self, request: Request) -> Optional[str]:
|
def extract_api_key(self, request: Request) -> str | None:
|
||||||
"""
|
"""
|
||||||
从请求中提取客户端 API 密钥。
|
从请求中提取客户端 API 密钥。
|
||||||
|
|
||||||
@@ -55,17 +53,17 @@ class ApiAdapter(ABC):
|
|||||||
context: ApiRequestContext,
|
context: ApiRequestContext,
|
||||||
*,
|
*,
|
||||||
success: bool,
|
success: bool,
|
||||||
status_code: Optional[int],
|
status_code: int | None,
|
||||||
error: Optional[str] = None,
|
error: str | None = None,
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""允许适配器在审计日志中追加自定义字段。"""
|
"""允许适配器在审计日志中追加自定义字段。"""
|
||||||
return {}
|
return {}
|
||||||
|
|
||||||
def detect_capability_requirements(
|
def detect_capability_requirements(
|
||||||
self,
|
self,
|
||||||
headers: Dict[str, str],
|
headers: dict[str, str],
|
||||||
request_body: Optional[Dict[str, Any]] = None,
|
request_body: dict[str, Any] | None = None,
|
||||||
) -> Dict[str, bool]:
|
) -> dict[str, bool]:
|
||||||
"""
|
"""
|
||||||
检测请求中隐含的能力需求(子类可覆盖)
|
检测请求中隐含的能力需求(子类可覆盖)
|
||||||
|
|
||||||
|
|||||||
@@ -1,5 +1,3 @@
|
|||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from fastapi import HTTPException
|
from fastapi import HTTPException
|
||||||
|
|
||||||
from src.models.database import UserRole
|
from src.models.database import UserRole
|
||||||
|
|||||||
@@ -1,10 +1,9 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import json
|
import json
|
||||||
import time
|
import time
|
||||||
import uuid
|
import uuid
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from typing import Any, Dict, Optional
|
from typing import Any
|
||||||
|
|
||||||
from fastapi import HTTPException, Request
|
from fastapi import HTTPException, Request
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
@@ -21,34 +20,34 @@ class ApiRequestContext:
|
|||||||
|
|
||||||
request: Request
|
request: Request
|
||||||
db: Session
|
db: Session
|
||||||
user: Optional[User]
|
user: User | None
|
||||||
api_key: Optional[ApiKey]
|
api_key: ApiKey | None
|
||||||
request_id: str
|
request_id: str
|
||||||
start_time: float
|
start_time: float
|
||||||
client_ip: str
|
client_ip: str
|
||||||
user_agent: str
|
user_agent: str
|
||||||
original_headers: Dict[str, str]
|
original_headers: dict[str, str]
|
||||||
query_params: Dict[str, str]
|
query_params: dict[str, str]
|
||||||
raw_body: bytes | None = None
|
raw_body: bytes | None = None
|
||||||
json_body: Optional[Dict[str, Any]] = None
|
json_body: dict[str, Any] | None = None
|
||||||
quota_remaining: Optional[float] = None
|
quota_remaining: float | None = None
|
||||||
mode: str = "standard" # standard / proxy
|
mode: str = "standard" # standard / proxy
|
||||||
api_format_hint: Optional[str] = None
|
api_format_hint: str | None = None
|
||||||
|
|
||||||
# URL 路径参数(如 Gemini API 的 /v1beta/models/{model}:generateContent)
|
# URL 路径参数(如 Gemini API 的 /v1beta/models/{model}:generateContent)
|
||||||
path_params: Dict[str, Any] = field(default_factory=dict)
|
path_params: dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
# Management Token(用于管理 API 认证)
|
# Management Token(用于管理 API 认证)
|
||||||
management_token: Optional[ManagementToken] = None
|
management_token: ManagementToken | None = None
|
||||||
|
|
||||||
# 供适配器扩展的状态存储
|
# 供适配器扩展的状态存储
|
||||||
extra: Dict[str, Any] = field(default_factory=dict)
|
extra: dict[str, Any] = field(default_factory=dict)
|
||||||
audit_metadata: Dict[str, Any] = field(default_factory=dict)
|
audit_metadata: dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
# 高频轮询端点日志抑制标志
|
# 高频轮询端点日志抑制标志
|
||||||
quiet_logging: bool = False
|
quiet_logging: bool = False
|
||||||
|
|
||||||
def ensure_json_body(self) -> Dict[str, Any]:
|
def ensure_json_body(self) -> dict[str, Any]:
|
||||||
"""确保请求体已解析为JSON并返回。"""
|
"""确保请求体已解析为JSON并返回。"""
|
||||||
if self.json_body is not None:
|
if self.json_body is not None:
|
||||||
return self.json_body
|
return self.json_body
|
||||||
@@ -70,7 +69,7 @@ class ApiRequestContext:
|
|||||||
if value is not None:
|
if value is not None:
|
||||||
self.audit_metadata[key] = value
|
self.audit_metadata[key] = value
|
||||||
|
|
||||||
def extend_audit_metadata(self, data: Dict[str, Any]) -> None:
|
def extend_audit_metadata(self, data: dict[str, Any]) -> None:
|
||||||
"""批量附加审计字段。"""
|
"""批量附加审计字段。"""
|
||||||
for key, value in data.items():
|
for key, value in data.items():
|
||||||
if value is not None:
|
if value is not None:
|
||||||
@@ -81,13 +80,13 @@ class ApiRequestContext:
|
|||||||
cls,
|
cls,
|
||||||
request: Request,
|
request: Request,
|
||||||
db: Session,
|
db: Session,
|
||||||
user: Optional[User],
|
user: User | None,
|
||||||
api_key: Optional[ApiKey],
|
api_key: ApiKey | None,
|
||||||
raw_body: Optional[bytes] = None,
|
raw_body: bytes | None = None,
|
||||||
mode: str = "standard",
|
mode: str = "standard",
|
||||||
api_format_hint: Optional[str] = None,
|
api_format_hint: str | None = None,
|
||||||
path_params: Optional[Dict[str, Any]] = None,
|
path_params: dict[str, Any] | None = None,
|
||||||
) -> "ApiRequestContext":
|
) -> ApiRequestContext:
|
||||||
"""创建上下文实例并提前读取必要的元数据。"""
|
"""创建上下文实例并提前读取必要的元数据。"""
|
||||||
request_id = getattr(request.state, "request_id", None) or str(uuid.uuid4())[:8]
|
request_id = getattr(request.state, "request_id", None) or str(uuid.uuid4())[:8]
|
||||||
setattr(request.state, "request_id", request_id)
|
setattr(request.state, "request_id", request_id)
|
||||||
|
|||||||
@@ -10,8 +10,9 @@
|
|||||||
4. Key 的 allowed_models 允许该模型(null = 允许所有)
|
4. Key 的 allowed_models 允许该模型(null = 允许所有)
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
from dataclasses import asdict, dataclass
|
from dataclasses import asdict, dataclass
|
||||||
from typing import Any, Optional
|
from typing import Any
|
||||||
|
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
@@ -27,7 +28,7 @@ _CACHE_KEY_PREFIX = "models:list"
|
|||||||
_CACHE_TTL = CacheTTL.MODEL # 300 秒
|
_CACHE_TTL = CacheTTL.MODEL # 300 秒
|
||||||
|
|
||||||
|
|
||||||
def _get_cache_key(api_formats: list[str], client_format: Optional[str] = None) -> str:
|
def _get_cache_key(api_formats: list[str], client_format: str | None = None) -> str:
|
||||||
"""生成缓存 key"""
|
"""生成缓存 key"""
|
||||||
formats_str = ",".join(sorted(api_formats))
|
formats_str = ",".join(sorted(api_formats))
|
||||||
format_key = (client_format or "any").lower()
|
format_key = (client_format or "any").lower()
|
||||||
@@ -35,8 +36,8 @@ def _get_cache_key(api_formats: list[str], client_format: Optional[str] = None)
|
|||||||
|
|
||||||
|
|
||||||
async def _get_cached_models(
|
async def _get_cached_models(
|
||||||
api_formats: list[str], client_format: Optional[str] = None
|
api_formats: list[str], client_format: str | None = None
|
||||||
) -> Optional[list["ModelInfo"]]:
|
) -> list[ModelInfo] | None:
|
||||||
"""从缓存获取模型列表"""
|
"""从缓存获取模型列表"""
|
||||||
cache_key = _get_cache_key(api_formats, client_format)
|
cache_key = _get_cache_key(api_formats, client_format)
|
||||||
try:
|
try:
|
||||||
@@ -51,8 +52,8 @@ async def _get_cached_models(
|
|||||||
|
|
||||||
async def _set_cached_models(
|
async def _set_cached_models(
|
||||||
api_formats: list[str],
|
api_formats: list[str],
|
||||||
models: list["ModelInfo"],
|
models: list[ModelInfo],
|
||||||
client_format: Optional[str] = None,
|
client_format: str | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""将模型列表写入缓存"""
|
"""将模型列表写入缓存"""
|
||||||
cache_key = _get_cache_key(api_formats, client_format)
|
cache_key = _get_cache_key(api_formats, client_format)
|
||||||
@@ -87,8 +88,8 @@ class ModelInfo:
|
|||||||
|
|
||||||
id: str # 模型 ID (GlobalModel.name 或 provider_model_name)
|
id: str # 模型 ID (GlobalModel.name 或 provider_model_name)
|
||||||
display_name: str
|
display_name: str
|
||||||
description: Optional[str]
|
description: str | None
|
||||||
created_at: Optional[str] # ISO 格式
|
created_at: str | None # ISO 格式
|
||||||
created_timestamp: int # Unix 时间戳
|
created_timestamp: int # Unix 时间戳
|
||||||
provider_name: str
|
provider_name: str
|
||||||
provider_id: str = "" # Provider ID,用于权限过滤
|
provider_id: str = "" # Provider ID,用于权限过滤
|
||||||
@@ -100,27 +101,27 @@ class ModelInfo:
|
|||||||
image_generation: bool = False
|
image_generation: bool = False
|
||||||
structured_output: bool = False
|
structured_output: bool = False
|
||||||
# 规格参数
|
# 规格参数
|
||||||
context_limit: Optional[int] = None
|
context_limit: int | None = None
|
||||||
output_limit: Optional[int] = None
|
output_limit: int | None = None
|
||||||
# 元信息
|
# 元信息
|
||||||
family: Optional[str] = None
|
family: str | None = None
|
||||||
knowledge_cutoff: Optional[str] = None
|
knowledge_cutoff: str | None = None
|
||||||
input_modalities: Optional[list[str]] = None
|
input_modalities: list[str] | None = None
|
||||||
output_modalities: Optional[list[str]] = None
|
output_modalities: list[str] | None = None
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class AccessRestrictions:
|
class AccessRestrictions:
|
||||||
"""API Key 或 User 的访问限制"""
|
"""API Key 或 User 的访问限制"""
|
||||||
|
|
||||||
allowed_providers: Optional[list[str]] = None # 允许的 Provider ID 列表
|
allowed_providers: list[str] | None = None # 允许的 Provider ID 列表
|
||||||
allowed_models: Optional[list[str]] = None # 允许的模型名称列表
|
allowed_models: list[str] | None = None # 允许的模型名称列表
|
||||||
allowed_api_formats: Optional[list[str]] = None # 允许的 API 格式列表
|
allowed_api_formats: list[str] | None = None # 允许的 API 格式列表
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_api_key_and_user(
|
def from_api_key_and_user(
|
||||||
cls, api_key: Optional[ApiKey], user: Optional[User]
|
cls, api_key: ApiKey | None, user: User | None
|
||||||
) -> "AccessRestrictions":
|
) -> AccessRestrictions:
|
||||||
"""
|
"""
|
||||||
从 API Key 和 User 合并访问限制
|
从 API Key 和 User 合并访问限制
|
||||||
|
|
||||||
@@ -130,9 +131,9 @@ class AccessRestrictions:
|
|||||||
- 如果 API Key 无限制但 User 有限制,使用 User 的限制
|
- 如果 API Key 无限制但 User 有限制,使用 User 的限制
|
||||||
- 两者都无限制则返回空限制
|
- 两者都无限制则返回空限制
|
||||||
"""
|
"""
|
||||||
allowed_providers: Optional[list[str]] = None
|
allowed_providers: list[str] | None = None
|
||||||
allowed_models: Optional[list[str]] = None
|
allowed_models: list[str] | None = None
|
||||||
allowed_api_formats: Optional[list[str]] = None
|
allowed_api_formats: list[str] | None = None
|
||||||
|
|
||||||
# 优先使用 API Key 的限制
|
# 优先使用 API Key 的限制
|
||||||
if api_key:
|
if api_key:
|
||||||
@@ -197,8 +198,8 @@ class AccessRestrictions:
|
|||||||
|
|
||||||
|
|
||||||
def _normalize_api_formats(
|
def _normalize_api_formats(
|
||||||
api_formats: Optional[list[str]],
|
api_formats: list[str] | None,
|
||||||
provider_to_formats: Optional[dict[str, set[str]]] = None,
|
provider_to_formats: dict[str, set[str]] | None = None,
|
||||||
) -> list[str]:
|
) -> list[str]:
|
||||||
"""规范化 API 格式列表(大写),必要时从 provider_to_formats 兜底"""
|
"""规范化 API 格式列表(大写),必要时从 provider_to_formats 兜底"""
|
||||||
if api_formats:
|
if api_formats:
|
||||||
@@ -212,7 +213,7 @@ def _normalize_api_formats(
|
|||||||
|
|
||||||
|
|
||||||
def _get_provider_model_names_for_formats(
|
def _get_provider_model_names_for_formats(
|
||||||
model: Model, usable_formats: Optional[set[str]] = None
|
model: Model, usable_formats: set[str] | None = None
|
||||||
) -> set[str]:
|
) -> set[str]:
|
||||||
"""
|
"""
|
||||||
获取模型在指定格式下支持的 Provider 模型名称集合
|
获取模型在指定格式下支持的 Provider 模型名称集合
|
||||||
@@ -305,7 +306,7 @@ def get_compatible_provider_formats(
|
|||||||
def get_available_provider_ids(
|
def get_available_provider_ids(
|
||||||
db: Session,
|
db: Session,
|
||||||
api_formats: list[str],
|
api_formats: list[str],
|
||||||
provider_to_formats: Optional[dict[str, set[str]]] = None,
|
provider_to_formats: dict[str, set[str]] | None = None,
|
||||||
) -> set[str]:
|
) -> set[str]:
|
||||||
"""
|
"""
|
||||||
返回有可用端点的 Provider IDs
|
返回有可用端点的 Provider IDs
|
||||||
@@ -334,7 +335,7 @@ def get_available_provider_ids(
|
|||||||
def _get_available_model_ids_for_format(
|
def _get_available_model_ids_for_format(
|
||||||
db: Session,
|
db: Session,
|
||||||
api_formats: list[str],
|
api_formats: list[str],
|
||||||
provider_to_formats: Optional[dict[str, set[str]]] = None,
|
provider_to_formats: dict[str, set[str]] | None = None,
|
||||||
) -> set[str]:
|
) -> set[str]:
|
||||||
"""
|
"""
|
||||||
获取指定格式下真正可用的模型 ID 集合
|
获取指定格式下真正可用的模型 ID 集合
|
||||||
@@ -410,7 +411,7 @@ def _get_available_model_ids_for_format(
|
|||||||
return available_model_ids
|
return available_model_ids
|
||||||
|
|
||||||
|
|
||||||
def _extract_model_info(model: Any) -> Optional[ModelInfo]:
|
def _extract_model_info(model: Any) -> ModelInfo | None:
|
||||||
"""
|
"""
|
||||||
从 Model 对象提取 ModelInfo
|
从 Model 对象提取 ModelInfo
|
||||||
|
|
||||||
@@ -424,7 +425,7 @@ def _extract_model_info(model: Any) -> Optional[ModelInfo]:
|
|||||||
|
|
||||||
model_id: str = global_model.name
|
model_id: str = global_model.name
|
||||||
display_name: str = global_model.display_name
|
display_name: str = global_model.display_name
|
||||||
created_at: Optional[str] = (
|
created_at: str | None = (
|
||||||
model.created_at.strftime("%Y-%m-%dT%H:%M:%SZ") if model.created_at else None
|
model.created_at.strftime("%Y-%m-%dT%H:%M:%SZ") if model.created_at else None
|
||||||
)
|
)
|
||||||
created_timestamp: int = int(model.created_at.timestamp()) if model.created_at else 0
|
created_timestamp: int = int(model.created_at.timestamp()) if model.created_at else 0
|
||||||
@@ -433,7 +434,7 @@ def _extract_model_info(model: Any) -> Optional[ModelInfo]:
|
|||||||
|
|
||||||
# 从 GlobalModel.config 提取配置信息
|
# 从 GlobalModel.config 提取配置信息
|
||||||
config: dict = global_model.config or {}
|
config: dict = global_model.config or {}
|
||||||
description: Optional[str] = config.get("description")
|
description: str | None = config.get("description")
|
||||||
|
|
||||||
return ModelInfo(
|
return ModelInfo(
|
||||||
id=model_id,
|
id=model_id,
|
||||||
@@ -464,10 +465,10 @@ def _extract_model_info(model: Any) -> Optional[ModelInfo]:
|
|||||||
async def list_available_models(
|
async def list_available_models(
|
||||||
db: Session,
|
db: Session,
|
||||||
available_provider_ids: set[str],
|
available_provider_ids: set[str],
|
||||||
api_formats: Optional[list[str]] = None,
|
api_formats: list[str] | None = None,
|
||||||
restrictions: Optional[AccessRestrictions] = None,
|
restrictions: AccessRestrictions | None = None,
|
||||||
provider_to_formats: Optional[dict[str, set[str]]] = None,
|
provider_to_formats: dict[str, set[str]] | None = None,
|
||||||
client_format: Optional[str] = None,
|
client_format: str | None = None,
|
||||||
) -> list[ModelInfo]:
|
) -> list[ModelInfo]:
|
||||||
"""
|
"""
|
||||||
获取可用模型列表(已去重,带缓存)
|
获取可用模型列表(已去重,带缓存)
|
||||||
@@ -503,7 +504,7 @@ async def list_available_models(
|
|||||||
return cached
|
return cached
|
||||||
|
|
||||||
# 如果提供了 api_formats,获取真正可用的模型 ID
|
# 如果提供了 api_formats,获取真正可用的模型 ID
|
||||||
available_model_ids: Optional[set[str]] = None
|
available_model_ids: set[str] | None = None
|
||||||
if normalized_formats:
|
if normalized_formats:
|
||||||
available_model_ids = _get_available_model_ids_for_format(
|
available_model_ids = _get_available_model_ids_for_format(
|
||||||
db, normalized_formats, provider_to_formats
|
db, normalized_formats, provider_to_formats
|
||||||
@@ -551,10 +552,10 @@ def find_model_by_id(
|
|||||||
db: Session,
|
db: Session,
|
||||||
model_id: str,
|
model_id: str,
|
||||||
available_provider_ids: set[str],
|
available_provider_ids: set[str],
|
||||||
api_formats: Optional[list[str]] = None,
|
api_formats: list[str] | None = None,
|
||||||
restrictions: Optional[AccessRestrictions] = None,
|
restrictions: AccessRestrictions | None = None,
|
||||||
provider_to_formats: Optional[dict[str, set[str]]] = None,
|
provider_to_formats: dict[str, set[str]] | None = None,
|
||||||
) -> Optional[ModelInfo]:
|
) -> ModelInfo | None:
|
||||||
"""
|
"""
|
||||||
按 ID 查找模型(仅支持 GlobalModel.name)
|
按 ID 查找模型(仅支持 GlobalModel.name)
|
||||||
|
|
||||||
@@ -575,7 +576,7 @@ def find_model_by_id(
|
|||||||
normalized_formats = _normalize_api_formats(api_formats, provider_to_formats)
|
normalized_formats = _normalize_api_formats(api_formats, provider_to_formats)
|
||||||
|
|
||||||
# 如果提供了 api_formats,获取真正可用的模型 ID
|
# 如果提供了 api_formats,获取真正可用的模型 ID
|
||||||
available_model_ids: Optional[set[str]] = None
|
available_model_ids: set[str] | None = None
|
||||||
if normalized_formats:
|
if normalized_formats:
|
||||||
available_model_ids = _get_available_model_ids_for_format(
|
available_model_ids = _get_available_model_ids_for_format(
|
||||||
db, normalized_formats, provider_to_formats
|
db, normalized_formats, provider_to_formats
|
||||||
|
|||||||
@@ -1,7 +1,6 @@
|
|||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from dataclasses import asdict, dataclass
|
from dataclasses import asdict, dataclass
|
||||||
from typing import Any, List, Sequence, Tuple, TypeVar
|
from typing import Any, TypeVar
|
||||||
|
from collections.abc import Sequence
|
||||||
|
|
||||||
from sqlalchemy.orm import Query
|
from sqlalchemy.orm import Query
|
||||||
|
|
||||||
@@ -19,7 +18,7 @@ class PaginationMeta:
|
|||||||
return asdict(self)
|
return asdict(self)
|
||||||
|
|
||||||
|
|
||||||
def paginate_query(query: Query, limit: int, offset: int) -> Tuple[int, List[T]]:
|
def paginate_query(query: Query, limit: int, offset: int) -> tuple[int, list[T]]:
|
||||||
"""
|
"""
|
||||||
对 SQLAlchemy 查询应用 limit/offset,并返回总数与结果列表。
|
对 SQLAlchemy 查询应用 limit/offset,并返回总数与结果列表。
|
||||||
"""
|
"""
|
||||||
@@ -30,7 +29,7 @@ def paginate_query(query: Query, limit: int, offset: int) -> Tuple[int, List[T]]
|
|||||||
|
|
||||||
def paginate_sequence(
|
def paginate_sequence(
|
||||||
items: Sequence[T], limit: int, offset: int
|
items: Sequence[T], limit: int, offset: int
|
||||||
) -> Tuple[List[T], PaginationMeta]:
|
) -> tuple[list[T], PaginationMeta]:
|
||||||
"""
|
"""
|
||||||
对内存序列应用分页,返回切片和元数据。
|
对内存序列应用分页,返回切片和元数据。
|
||||||
"""
|
"""
|
||||||
@@ -40,7 +39,7 @@ def paginate_sequence(
|
|||||||
return sliced, meta
|
return sliced, meta
|
||||||
|
|
||||||
|
|
||||||
def build_pagination_payload(items: List[dict], meta: PaginationMeta, **extra: Any) -> dict:
|
def build_pagination_payload(items: list[dict], meta: PaginationMeta, **extra: Any) -> dict:
|
||||||
"""
|
"""
|
||||||
构建标准分页响应 payload。
|
构建标准分页响应 payload。
|
||||||
"""
|
"""
|
||||||
|
|||||||
@@ -2,7 +2,7 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import time
|
import time
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
from typing import TYPE_CHECKING, Any, Optional, Tuple
|
from typing import TYPE_CHECKING, Any
|
||||||
|
|
||||||
from fastapi import HTTPException, Request
|
from fastapi import HTTPException, Request
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
@@ -52,8 +52,8 @@ class ApiRequestPipeline:
|
|||||||
db: Session,
|
db: Session,
|
||||||
*,
|
*,
|
||||||
mode: ApiMode = ApiMode.STANDARD,
|
mode: ApiMode = ApiMode.STANDARD,
|
||||||
api_format_hint: Optional[str] = None,
|
api_format_hint: str | None = None,
|
||||||
path_params: Optional[dict[str, Any]] = None,
|
path_params: dict[str, Any] | None = None,
|
||||||
):
|
):
|
||||||
# 高频轮询端点抑制 debug 日志
|
# 高频轮询端点抑制 debug 日志
|
||||||
is_quiet = http_request.url.path in QUIET_POLLING_PATHS
|
is_quiet = http_request.url.path in QUIET_POLLING_PATHS
|
||||||
@@ -95,7 +95,7 @@ class ApiRequestPipeline:
|
|||||||
)
|
)
|
||||||
if not is_quiet:
|
if not is_quiet:
|
||||||
logger.debug("[Pipeline] Raw body读取完成 | size=%d bytes", len(raw_body) if raw_body is not None else 0)
|
logger.debug("[Pipeline] Raw body读取完成 | size=%d bytes", len(raw_body) if raw_body is not None else 0)
|
||||||
except asyncio.TimeoutError:
|
except TimeoutError:
|
||||||
timeout_sec = int(config.request_body_timeout)
|
timeout_sec = int(config.request_body_timeout)
|
||||||
logger.error(f"读取请求体超时({timeout_sec}s),可能客户端未发送完整请求体")
|
logger.error(f"读取请求体超时({timeout_sec}s),可能客户端未发送完整请求体")
|
||||||
raise HTTPException(
|
raise HTTPException(
|
||||||
@@ -166,7 +166,7 @@ class ApiRequestPipeline:
|
|||||||
|
|
||||||
def _authenticate_client(
|
def _authenticate_client(
|
||||||
self, request: Request, db: Session, adapter: ApiAdapter, *, quiet: bool = False
|
self, request: Request, db: Session, adapter: ApiAdapter, *, quiet: bool = False
|
||||||
) -> Tuple[User, ApiKey]:
|
) -> tuple[User, ApiKey]:
|
||||||
if not quiet:
|
if not quiet:
|
||||||
logger.debug("[Pipeline._authenticate_client] 开始")
|
logger.debug("[Pipeline._authenticate_client] 开始")
|
||||||
# 使用 adapter 的 extract_api_key 方法,支持不同 API 格式的认证头
|
# 使用 adapter 的 extract_api_key 方法,支持不同 API 格式的认证头
|
||||||
@@ -215,7 +215,7 @@ class ApiRequestPipeline:
|
|||||||
|
|
||||||
async def _authenticate_admin(
|
async def _authenticate_admin(
|
||||||
self, request: Request, db: Session
|
self, request: Request, db: Session
|
||||||
) -> Tuple[User, Optional["ManagementToken"]]:
|
) -> tuple[User, ManagementToken | None]:
|
||||||
"""管理员认证,支持 JWT 和 Management Token 两种方式"""
|
"""管理员认证,支持 JWT 和 Management Token 两种方式"""
|
||||||
from src.models.database import ManagementToken
|
from src.models.database import ManagementToken
|
||||||
from src.utils.request_utils import get_client_ip
|
from src.utils.request_utils import get_client_ip
|
||||||
@@ -278,7 +278,7 @@ class ApiRequestPipeline:
|
|||||||
|
|
||||||
async def _authenticate_user(
|
async def _authenticate_user(
|
||||||
self, request: Request, db: Session
|
self, request: Request, db: Session
|
||||||
) -> Tuple[User, Optional["ManagementToken"]]:
|
) -> tuple[User, ManagementToken | None]:
|
||||||
"""用户认证,支持 JWT 和 Management Token 两种方式"""
|
"""用户认证,支持 JWT 和 Management Token 两种方式"""
|
||||||
from src.models.database import ManagementToken
|
from src.models.database import ManagementToken
|
||||||
from src.utils.request_utils import get_client_ip
|
from src.utils.request_utils import get_client_ip
|
||||||
@@ -329,7 +329,7 @@ class ApiRequestPipeline:
|
|||||||
|
|
||||||
async def _authenticate_management(
|
async def _authenticate_management(
|
||||||
self, request: Request, db: Session
|
self, request: Request, db: Session
|
||||||
) -> Tuple[User, "ManagementToken"]:
|
) -> tuple[User, ManagementToken]:
|
||||||
"""Management Token 认证"""
|
"""Management Token 认证"""
|
||||||
from src.models.database import ManagementToken
|
from src.models.database import ManagementToken
|
||||||
from src.utils.request_utils import get_client_ip
|
from src.utils.request_utils import get_client_ip
|
||||||
@@ -362,7 +362,7 @@ class ApiRequestPipeline:
|
|||||||
|
|
||||||
return user, management_token
|
return user, management_token
|
||||||
|
|
||||||
def _calculate_quota_remaining(self, user: Optional[User]) -> Optional[float]:
|
def _calculate_quota_remaining(self, user: User | None) -> float | None:
|
||||||
if not user:
|
if not user:
|
||||||
return None
|
return None
|
||||||
if user.quota_usd is None or user.quota_usd < 0:
|
if user.quota_usd is None or user.quota_usd < 0:
|
||||||
@@ -375,8 +375,8 @@ class ApiRequestPipeline:
|
|||||||
adapter: ApiAdapter,
|
adapter: ApiAdapter,
|
||||||
*,
|
*,
|
||||||
success: bool,
|
success: bool,
|
||||||
status_code: Optional[int] = None,
|
status_code: int | None = None,
|
||||||
error: Optional[str] = None,
|
error: str | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""记录审计事件
|
"""记录审计事件
|
||||||
|
|
||||||
@@ -432,8 +432,8 @@ class ApiRequestPipeline:
|
|||||||
adapter: ApiAdapter,
|
adapter: ApiAdapter,
|
||||||
*,
|
*,
|
||||||
success: bool,
|
success: bool,
|
||||||
status_code: Optional[int],
|
status_code: int | None,
|
||||||
error: Optional[str],
|
error: str | None,
|
||||||
) -> dict:
|
) -> dict:
|
||||||
duration_ms = max((time.time() - context.start_time) * 1000, 0.0)
|
duration_ms = max((time.time() - context.start_time) * 1000, 0.0)
|
||||||
request = context.request
|
request = context.request
|
||||||
|
|||||||
@@ -2,7 +2,6 @@
|
|||||||
|
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from datetime import datetime, timedelta, timezone
|
from datetime import datetime, timedelta, timezone
|
||||||
from typing import List
|
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||||
from sqlalchemy import and_, func
|
from sqlalchemy import and_, func
|
||||||
@@ -952,7 +951,7 @@ class DashboardDailyStatsAdapter(DashboardAdapter):
|
|||||||
# 构建完整日期序列(使用业务时区日期)
|
# 构建完整日期序列(使用业务时区日期)
|
||||||
current_date = start_date_local.date()
|
current_date = start_date_local.date()
|
||||||
end_date_date = end_date_local.date()
|
end_date_date = end_date_local.date()
|
||||||
formatted: List[dict] = []
|
formatted: list[dict] = []
|
||||||
while current_date <= end_date_date:
|
while current_date <= end_date_date:
|
||||||
date_str = current_date.isoformat()
|
date_str = current_date.isoformat()
|
||||||
stat = stats_map.get(date_str)
|
stat = stats_map.get(date_str)
|
||||||
|
|||||||
@@ -27,21 +27,20 @@
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import time
|
import time
|
||||||
from typing import (
|
from typing import (
|
||||||
TYPE_CHECKING,
|
TYPE_CHECKING,
|
||||||
Any,
|
Any,
|
||||||
Awaitable,
|
|
||||||
Callable,
|
|
||||||
Coroutine,
|
|
||||||
Dict,
|
|
||||||
Optional,
|
|
||||||
Protocol,
|
Protocol,
|
||||||
TypeVar,
|
TypeVar,
|
||||||
runtime_checkable,
|
runtime_checkable,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
from collections.abc import Callable
|
||||||
|
from collections.abc import Awaitable, Coroutine
|
||||||
|
|
||||||
from fastapi import Request
|
from fastapi import Request
|
||||||
from fastapi.responses import JSONResponse, StreamingResponse
|
from fastapi.responses import JSONResponse, StreamingResponse
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
@@ -57,6 +56,9 @@ from src.services.usage.service import UsageService
|
|||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from src.api.handlers.base.stream_context import StreamContext
|
from src.api.handlers.base.stream_context import StreamContext
|
||||||
|
|
||||||
|
# Adapter 检测器类型:接受 headers 和可选的 request_body,返回能力需求字典
|
||||||
|
type AdapterDetectorType = Callable[[dict[str, str], dict[str, Any] | None], dict[str, bool]]
|
||||||
|
|
||||||
|
|
||||||
class MessageTelemetry:
|
class MessageTelemetry:
|
||||||
"""
|
"""
|
||||||
@@ -105,29 +107,29 @@ class MessageTelemetry:
|
|||||||
output_tokens: int,
|
output_tokens: int,
|
||||||
response_time_ms: int,
|
response_time_ms: int,
|
||||||
status_code: int,
|
status_code: int,
|
||||||
request_body: Dict[str, Any],
|
request_body: dict[str, Any],
|
||||||
request_headers: Dict[str, Any],
|
request_headers: dict[str, Any],
|
||||||
response_body: Any,
|
response_body: Any,
|
||||||
response_headers: Dict[str, Any],
|
response_headers: dict[str, Any],
|
||||||
client_response_headers: Optional[Dict[str, Any]] = None,
|
client_response_headers: dict[str, Any] | None = None,
|
||||||
cache_creation_tokens: int = 0,
|
cache_creation_tokens: int = 0,
|
||||||
cache_read_tokens: int = 0,
|
cache_read_tokens: int = 0,
|
||||||
is_stream: bool = False,
|
is_stream: bool = False,
|
||||||
provider_request_headers: Optional[Dict[str, Any]] = None,
|
provider_request_headers: dict[str, Any] | None = None,
|
||||||
# 时间指标
|
# 时间指标
|
||||||
first_byte_time_ms: Optional[int] = None, # 首字时间/TTFB
|
first_byte_time_ms: int | None = None, # 首字时间/TTFB
|
||||||
# Provider 侧追踪信息(用于记录真实成本)
|
# Provider 侧追踪信息(用于记录真实成本)
|
||||||
provider_id: Optional[str] = None,
|
provider_id: str | None = None,
|
||||||
provider_endpoint_id: Optional[str] = None,
|
provider_endpoint_id: str | None = None,
|
||||||
provider_api_key_id: Optional[str] = None,
|
provider_api_key_id: str | None = None,
|
||||||
api_format: Optional[str] = None,
|
api_format: str | None = None,
|
||||||
# 格式转换追踪
|
# 格式转换追踪
|
||||||
endpoint_api_format: Optional[str] = None, # 端点原生 API 格式
|
endpoint_api_format: str | None = None, # 端点原生 API 格式
|
||||||
has_format_conversion: bool = False, # 是否发生了格式转换
|
has_format_conversion: bool = False, # 是否发生了格式转换
|
||||||
# 模型映射信息
|
# 模型映射信息
|
||||||
target_model: Optional[str] = None,
|
target_model: str | None = None,
|
||||||
# Provider 响应元数据(如 Gemini 的 modelVersion)
|
# Provider 响应元数据(如 Gemini 的 modelVersion)
|
||||||
response_metadata: Optional[Dict[str, Any]] = None,
|
response_metadata: dict[str, Any] | None = None,
|
||||||
) -> float:
|
) -> float:
|
||||||
total_cost = await self.calculate_cost(
|
total_cost = await self.calculate_cost(
|
||||||
provider,
|
provider,
|
||||||
@@ -199,24 +201,24 @@ class MessageTelemetry:
|
|||||||
response_time_ms: int,
|
response_time_ms: int,
|
||||||
status_code: int,
|
status_code: int,
|
||||||
error_message: str,
|
error_message: str,
|
||||||
request_body: Dict[str, Any],
|
request_body: dict[str, Any],
|
||||||
request_headers: Dict[str, Any],
|
request_headers: dict[str, Any],
|
||||||
is_stream: bool,
|
is_stream: bool,
|
||||||
api_format: Optional[str] = None,
|
api_format: str | None = None,
|
||||||
provider_request_headers: Optional[Dict[str, Any]] = None,
|
provider_request_headers: dict[str, Any] | None = None,
|
||||||
# 预估 token 信息(来自 message_start 事件,用于中断请求的成本估算)
|
# 预估 token 信息(来自 message_start 事件,用于中断请求的成本估算)
|
||||||
input_tokens: int = 0,
|
input_tokens: int = 0,
|
||||||
output_tokens: int = 0,
|
output_tokens: int = 0,
|
||||||
cache_creation_tokens: int = 0,
|
cache_creation_tokens: int = 0,
|
||||||
cache_read_tokens: int = 0,
|
cache_read_tokens: int = 0,
|
||||||
response_body: Optional[Dict[str, Any]] = None,
|
response_body: dict[str, Any] | None = None,
|
||||||
response_headers: Optional[Dict[str, Any]] = None,
|
response_headers: dict[str, Any] | None = None,
|
||||||
client_response_headers: Optional[Dict[str, Any]] = None,
|
client_response_headers: dict[str, Any] | None = None,
|
||||||
# 格式转换追踪
|
# 格式转换追踪
|
||||||
endpoint_api_format: Optional[str] = None,
|
endpoint_api_format: str | None = None,
|
||||||
has_format_conversion: bool = False,
|
has_format_conversion: bool = False,
|
||||||
# 模型映射信息
|
# 模型映射信息
|
||||||
target_model: Optional[str] = None,
|
target_model: str | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""
|
"""
|
||||||
记录失败请求
|
记录失败请求
|
||||||
@@ -273,24 +275,24 @@ class MessageTelemetry:
|
|||||||
provider: str,
|
provider: str,
|
||||||
model: str,
|
model: str,
|
||||||
response_time_ms: int,
|
response_time_ms: int,
|
||||||
first_byte_time_ms: Optional[int],
|
first_byte_time_ms: int | None,
|
||||||
status_code: int,
|
status_code: int,
|
||||||
request_body: Dict[str, Any],
|
request_body: dict[str, Any],
|
||||||
request_headers: Dict[str, Any],
|
request_headers: dict[str, Any],
|
||||||
is_stream: bool,
|
is_stream: bool,
|
||||||
api_format: Optional[str] = None,
|
api_format: str | None = None,
|
||||||
provider_request_headers: Optional[Dict[str, Any]] = None,
|
provider_request_headers: dict[str, Any] | None = None,
|
||||||
input_tokens: int = 0,
|
input_tokens: int = 0,
|
||||||
output_tokens: int = 0,
|
output_tokens: int = 0,
|
||||||
cache_creation_tokens: int = 0,
|
cache_creation_tokens: int = 0,
|
||||||
cache_read_tokens: int = 0,
|
cache_read_tokens: int = 0,
|
||||||
response_body: Optional[Dict[str, Any]] = None,
|
response_body: dict[str, Any] | None = None,
|
||||||
response_headers: Optional[Dict[str, Any]] = None,
|
response_headers: dict[str, Any] | None = None,
|
||||||
client_response_headers: Optional[Dict[str, Any]] = None,
|
client_response_headers: dict[str, Any] | None = None,
|
||||||
# 格式转换追踪
|
# 格式转换追踪
|
||||||
endpoint_api_format: Optional[str] = None,
|
endpoint_api_format: str | None = None,
|
||||||
has_format_conversion: bool = False,
|
has_format_conversion: bool = False,
|
||||||
target_model: Optional[str] = None,
|
target_model: str | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""
|
"""
|
||||||
记录客户端取消的请求
|
记录客户端取消的请求
|
||||||
@@ -341,9 +343,9 @@ class MessageHandlerProtocol(Protocol):
|
|||||||
self,
|
self,
|
||||||
request: Any,
|
request: Any,
|
||||||
http_request: Request,
|
http_request: Request,
|
||||||
original_headers: Dict[str, str],
|
original_headers: dict[str, str],
|
||||||
original_request_body: Dict[str, Any],
|
original_request_body: dict[str, Any],
|
||||||
query_params: Optional[Dict[str, str]] = None,
|
query_params: dict[str, str] | None = None,
|
||||||
) -> StreamingResponse:
|
) -> StreamingResponse:
|
||||||
"""处理流式请求"""
|
"""处理流式请求"""
|
||||||
...
|
...
|
||||||
@@ -352,9 +354,9 @@ class MessageHandlerProtocol(Protocol):
|
|||||||
self,
|
self,
|
||||||
request: Any,
|
request: Any,
|
||||||
http_request: Request,
|
http_request: Request,
|
||||||
original_headers: Dict[str, str],
|
original_headers: dict[str, str],
|
||||||
original_request_body: Dict[str, Any],
|
original_request_body: dict[str, Any],
|
||||||
query_params: Optional[Dict[str, str]] = None,
|
query_params: dict[str, str] | None = None,
|
||||||
) -> JSONResponse:
|
) -> JSONResponse:
|
||||||
"""处理非流式请求"""
|
"""处理非流式请求"""
|
||||||
...
|
...
|
||||||
@@ -371,9 +373,6 @@ class BaseMessageHandler:
|
|||||||
推荐使用 MessageHandlerProtocol 中定义的签名。
|
推荐使用 MessageHandlerProtocol 中定义的签名。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
# Adapter 检测器类型
|
|
||||||
AdapterDetectorType = Callable[[Dict[str, str], Optional[Dict[str, Any]]], Dict[str, bool]]
|
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
@@ -384,8 +383,8 @@ class BaseMessageHandler:
|
|||||||
client_ip: str,
|
client_ip: str,
|
||||||
user_agent: str,
|
user_agent: str,
|
||||||
start_time: float,
|
start_time: float,
|
||||||
allowed_api_formats: Optional[list[str]] = None,
|
allowed_api_formats: list[str] | None = None,
|
||||||
adapter_detector: Optional[AdapterDetectorType] = None,
|
adapter_detector: AdapterDetectorType | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
self.db = db
|
self.db = db
|
||||||
self.user = user
|
self.user = user
|
||||||
@@ -408,9 +407,9 @@ class BaseMessageHandler:
|
|||||||
def _resolve_capability_requirements(
|
def _resolve_capability_requirements(
|
||||||
self,
|
self,
|
||||||
model_name: str,
|
model_name: str,
|
||||||
request_headers: Optional[Dict[str, str]] = None,
|
request_headers: dict[str, str] | None = None,
|
||||||
request_body: Optional[Dict[str, Any]] = None,
|
request_body: dict[str, Any] | None = None,
|
||||||
) -> Dict[str, bool]:
|
) -> dict[str, bool]:
|
||||||
"""
|
"""
|
||||||
解析请求的能力需求
|
解析请求的能力需求
|
||||||
|
|
||||||
@@ -442,12 +441,12 @@ class BaseMessageHandler:
|
|||||||
async def _resolve_preferred_key_ids(
|
async def _resolve_preferred_key_ids(
|
||||||
self,
|
self,
|
||||||
model_name: str,
|
model_name: str,
|
||||||
request_body: Optional[Dict[str, Any]] = None,
|
request_body: dict[str, Any] | None = None,
|
||||||
) -> Optional[list[str]]:
|
) -> list[str] | None:
|
||||||
"""可选的 Key 优先级解析钩子(默认不启用)。"""
|
"""可选的 Key 优先级解析钩子(默认不启用)。"""
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def get_api_format(self, provider_type: Optional[str] = None) -> APIFormat:
|
def get_api_format(self, provider_type: str | None = None) -> APIFormat:
|
||||||
"""根据 provider_type 解析 API 格式,未知类型默认 OPENAI"""
|
"""根据 provider_type 解析 API 格式,未知类型默认 OPENAI"""
|
||||||
if provider_type:
|
if provider_type:
|
||||||
result = resolve_api_format(provider_type, default=APIFormat.OPENAI)
|
result = resolve_api_format(provider_type, default=APIFormat.OPENAI)
|
||||||
@@ -456,17 +455,17 @@ class BaseMessageHandler:
|
|||||||
|
|
||||||
def build_provider_payload(
|
def build_provider_payload(
|
||||||
self,
|
self,
|
||||||
original_body: Dict[str, Any],
|
original_body: dict[str, Any],
|
||||||
*,
|
*,
|
||||||
mapped_model: Optional[str] = None,
|
mapped_model: str | None = None,
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""构建发送给 Provider 的请求体,替换 model 名称"""
|
"""构建发送给 Provider 的请求体,替换 model 名称"""
|
||||||
payload = dict(original_body)
|
payload = dict(original_body)
|
||||||
if mapped_model:
|
if mapped_model:
|
||||||
payload["model"] = mapped_model
|
payload["model"] = mapped_model
|
||||||
return payload
|
return payload
|
||||||
|
|
||||||
def _update_usage_to_streaming(self, request_id: Optional[str] = None) -> None:
|
def _update_usage_to_streaming(self, request_id: str | None = None) -> None:
|
||||||
"""更新 Usage 状态为 streaming(流式传输开始时调用)
|
"""更新 Usage 状态为 streaming(流式传输开始时调用)
|
||||||
|
|
||||||
使用 asyncio 后台任务执行数据库更新,避免阻塞流式传输
|
使用 asyncio 后台任务执行数据库更新,避免阻塞流式传输
|
||||||
@@ -500,7 +499,7 @@ class BaseMessageHandler:
|
|||||||
# 创建后台任务,不阻塞当前流
|
# 创建后台任务,不阻塞当前流
|
||||||
asyncio.create_task(_do_update())
|
asyncio.create_task(_do_update())
|
||||||
|
|
||||||
def _update_usage_to_streaming_with_ctx(self, ctx: "StreamContext") -> None:
|
def _update_usage_to_streaming_with_ctx(self, ctx: StreamContext) -> None:
|
||||||
"""更新 Usage 状态为 streaming,同时更新 provider 相关信息
|
"""更新 Usage 状态为 streaming,同时更新 provider 相关信息
|
||||||
|
|
||||||
使用 asyncio 后台任务执行数据库更新,避免阻塞流式传输
|
使用 asyncio 后台任务执行数据库更新,避免阻塞流式传输
|
||||||
|
|||||||
@@ -19,7 +19,7 @@ Chat Adapter 通用基类
|
|||||||
import time
|
import time
|
||||||
import traceback
|
import traceback
|
||||||
from abc import abstractmethod
|
from abc import abstractmethod
|
||||||
from typing import Any, Dict, Optional, Tuple, Type
|
from typing import Any
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
from fastapi import HTTPException, Request
|
from fastapi import HTTPException, Request
|
||||||
@@ -65,7 +65,7 @@ class ChatAdapterBase(ApiAdapter):
|
|||||||
|
|
||||||
# 子类必须覆盖
|
# 子类必须覆盖
|
||||||
FORMAT_ID: str = "UNKNOWN"
|
FORMAT_ID: str = "UNKNOWN"
|
||||||
HANDLER_CLASS: Type[ChatHandlerBase]
|
HANDLER_CLASS: type[ChatHandlerBase]
|
||||||
|
|
||||||
# 适配器配置
|
# 适配器配置
|
||||||
name: str = "chat.base"
|
name: str = "chat.base"
|
||||||
@@ -90,7 +90,7 @@ class ChatAdapterBase(ApiAdapter):
|
|||||||
return base_url
|
return base_url
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def build_base_headers(cls, api_key: str) -> Dict[str, str]:
|
def build_base_headers(cls, api_key: str) -> dict[str, str]:
|
||||||
"""构建基础请求头,使用统一的 headers.py 实现"""
|
"""构建基础请求头,使用统一的 headers.py 实现"""
|
||||||
return build_adapter_base_headers(cls._get_api_format(), api_key)
|
return build_adapter_base_headers(cls._get_api_format(), api_key)
|
||||||
|
|
||||||
@@ -101,13 +101,13 @@ class ChatAdapterBase(ApiAdapter):
|
|||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def build_headers_with_extra(
|
def build_headers_with_extra(
|
||||||
cls, api_key: str, extra_headers: Optional[Dict[str, str]] = None
|
cls, api_key: str, extra_headers: dict[str, str] | None = None
|
||||||
) -> Dict[str, str]:
|
) -> dict[str, str]:
|
||||||
"""构建完整请求头(包含 extra_headers),使用统一的 headers.py 实现"""
|
"""构建完整请求头(包含 extra_headers),使用统一的 headers.py 实现"""
|
||||||
return build_adapter_headers(cls._get_api_format(), api_key, extra_headers)
|
return build_adapter_headers(cls._get_api_format(), api_key, extra_headers)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def build_request_body(cls, request_data: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
|
def build_request_body(cls, request_data: dict[str, Any] | None = None) -> dict[str, Any]:
|
||||||
"""构建测试请求体,使用转换器注册表自动处理格式转换
|
"""构建测试请求体,使用转换器注册表自动处理格式转换
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
@@ -120,11 +120,11 @@ class ChatAdapterBase(ApiAdapter):
|
|||||||
|
|
||||||
return build_test_request_body(cls.FORMAT_ID, request_data)
|
return build_test_request_body(cls.FORMAT_ID, request_data)
|
||||||
|
|
||||||
def extract_api_key(self, request: Request) -> Optional[str]:
|
def extract_api_key(self, request: Request) -> str | None:
|
||||||
"""从请求中提取 API 密钥,使用统一的 headers.py 实现"""
|
"""从请求中提取 API 密钥,使用统一的 headers.py 实现"""
|
||||||
return extract_client_api_key(dict(request.headers), self._get_api_format())
|
return extract_client_api_key(dict(request.headers), self._get_api_format())
|
||||||
|
|
||||||
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
|
def __init__(self, allowed_api_formats: list[str] | None = None):
|
||||||
self.allowed_api_formats = allowed_api_formats or [self.FORMAT_ID]
|
self.allowed_api_formats = allowed_api_formats or [self.FORMAT_ID]
|
||||||
|
|
||||||
async def handle(self, context: ApiRequestContext):
|
async def handle(self, context: ApiRequestContext):
|
||||||
@@ -282,8 +282,8 @@ class ChatAdapterBase(ApiAdapter):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def _merge_path_params(
|
def _merge_path_params(
|
||||||
self, original_request_body: Dict[str, Any], path_params: Dict[str, Any]
|
self, original_request_body: dict[str, Any], path_params: dict[str, Any]
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
合并 URL 路径参数到请求体 - 子类可覆盖
|
合并 URL 路径参数到请求体 - 子类可覆盖
|
||||||
|
|
||||||
@@ -316,7 +316,7 @@ class ChatAdapterBase(ApiAdapter):
|
|||||||
"""
|
"""
|
||||||
pass
|
pass
|
||||||
|
|
||||||
def _extract_message_count(self, payload: Dict[str, Any], request_obj) -> int:
|
def _extract_message_count(self, payload: dict[str, Any], request_obj) -> int:
|
||||||
"""
|
"""
|
||||||
提取消息数量 - 子类可覆盖
|
提取消息数量 - 子类可覆盖
|
||||||
|
|
||||||
@@ -327,7 +327,7 @@ class ChatAdapterBase(ApiAdapter):
|
|||||||
messages = request_obj.messages
|
messages = request_obj.messages
|
||||||
return len(messages) if isinstance(messages, list) else 0
|
return len(messages) if isinstance(messages, list) else 0
|
||||||
|
|
||||||
def _build_audit_metadata(self, payload: Dict[str, Any], request_obj) -> Dict[str, Any]:
|
def _build_audit_metadata(self, payload: dict[str, Any], request_obj) -> dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
构建审计日志元数据 - 子类可覆盖
|
构建审计日志元数据 - 子类可覆盖
|
||||||
"""
|
"""
|
||||||
@@ -355,8 +355,8 @@ class ChatAdapterBase(ApiAdapter):
|
|||||||
model: str,
|
model: str,
|
||||||
stream: bool,
|
stream: bool,
|
||||||
start_time: float,
|
start_time: float,
|
||||||
original_headers: Dict[str, str],
|
original_headers: dict[str, str],
|
||||||
original_request_body: Dict[str, Any],
|
original_request_body: dict[str, Any],
|
||||||
client_ip: str,
|
client_ip: str,
|
||||||
request_id: str,
|
request_id: str,
|
||||||
) -> JSONResponse:
|
) -> JSONResponse:
|
||||||
@@ -426,8 +426,8 @@ class ChatAdapterBase(ApiAdapter):
|
|||||||
model: str,
|
model: str,
|
||||||
stream: bool,
|
stream: bool,
|
||||||
start_time: float,
|
start_time: float,
|
||||||
original_headers: Dict[str, str],
|
original_headers: dict[str, str],
|
||||||
original_request_body: Dict[str, Any],
|
original_request_body: dict[str, Any],
|
||||||
client_ip: str,
|
client_ip: str,
|
||||||
request_id: str,
|
request_id: str,
|
||||||
) -> JSONResponse:
|
) -> JSONResponse:
|
||||||
@@ -527,12 +527,12 @@ class ChatAdapterBase(ApiAdapter):
|
|||||||
cache_read_input_tokens: int,
|
cache_read_input_tokens: int,
|
||||||
input_price_per_1m: float,
|
input_price_per_1m: float,
|
||||||
output_price_per_1m: float,
|
output_price_per_1m: float,
|
||||||
cache_creation_price_per_1m: Optional[float],
|
cache_creation_price_per_1m: float | None,
|
||||||
cache_read_price_per_1m: Optional[float],
|
cache_read_price_per_1m: float | None,
|
||||||
price_per_request: Optional[float],
|
price_per_request: float | None,
|
||||||
tiered_pricing: Optional[dict] = None,
|
tiered_pricing: dict | None = None,
|
||||||
cache_ttl_minutes: Optional[int] = None,
|
cache_ttl_minutes: int | None = None,
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
计算请求成本
|
计算请求成本
|
||||||
|
|
||||||
@@ -597,8 +597,8 @@ class ChatAdapterBase(ApiAdapter):
|
|||||||
client: httpx.AsyncClient,
|
client: httpx.AsyncClient,
|
||||||
base_url: str,
|
base_url: str,
|
||||||
api_key: str,
|
api_key: str,
|
||||||
extra_headers: Optional[Dict[str, str]] = None,
|
extra_headers: dict[str, str] | None = None,
|
||||||
) -> Tuple[list, Optional[str]]:
|
) -> tuple[list, str | None]:
|
||||||
"""
|
"""
|
||||||
查询上游 API 支持的模型列表
|
查询上游 API 支持的模型列表
|
||||||
|
|
||||||
@@ -626,16 +626,16 @@ class ChatAdapterBase(ApiAdapter):
|
|||||||
client: httpx.AsyncClient,
|
client: httpx.AsyncClient,
|
||||||
base_url: str,
|
base_url: str,
|
||||||
api_key: str,
|
api_key: str,
|
||||||
request_data: Dict[str, Any],
|
request_data: dict[str, Any],
|
||||||
extra_headers: Optional[Dict[str, str]] = None,
|
extra_headers: dict[str, str] | None = None,
|
||||||
# 用量计算参数(现在强制记录)
|
# 用量计算参数(现在强制记录)
|
||||||
db: Optional[Any] = None,
|
db: Any | None = None,
|
||||||
user: Optional[Any] = None,
|
user: Any | None = None,
|
||||||
provider_name: Optional[str] = None,
|
provider_name: str | None = None,
|
||||||
provider_id: Optional[str] = None,
|
provider_id: str | None = None,
|
||||||
api_key_id: Optional[str] = None,
|
api_key_id: str | None = None,
|
||||||
model_name: Optional[str] = None,
|
model_name: str | None = None,
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
测试模型连接性(非流式)
|
测试模型连接性(非流式)
|
||||||
|
|
||||||
@@ -682,11 +682,11 @@ class ChatAdapterBase(ApiAdapter):
|
|||||||
# Adapter 注册表 - 用于根据 API format 获取 Adapter 实例
|
# Adapter 注册表 - 用于根据 API format 获取 Adapter 实例
|
||||||
# =========================================================================
|
# =========================================================================
|
||||||
|
|
||||||
_ADAPTER_REGISTRY: Dict[str, Type["ChatAdapterBase"]] = {}
|
_ADAPTER_REGISTRY: dict[str, type[ChatAdapterBase]] = {}
|
||||||
_ADAPTERS_LOADED = False
|
_ADAPTERS_LOADED = False
|
||||||
|
|
||||||
|
|
||||||
def register_adapter(adapter_class: Type["ChatAdapterBase"]) -> Type["ChatAdapterBase"]:
|
def register_adapter(adapter_class: type[ChatAdapterBase]) -> type[ChatAdapterBase]:
|
||||||
"""
|
"""
|
||||||
注册 Adapter 类到注册表
|
注册 Adapter 类到注册表
|
||||||
|
|
||||||
@@ -731,7 +731,7 @@ def _ensure_adapters_loaded():
|
|||||||
_ADAPTERS_LOADED = True
|
_ADAPTERS_LOADED = True
|
||||||
|
|
||||||
|
|
||||||
def get_adapter_class(api_format: str) -> Optional[Type["ChatAdapterBase"]]:
|
def get_adapter_class(api_format: str) -> type[ChatAdapterBase] | None:
|
||||||
"""
|
"""
|
||||||
根据 API format 获取 Adapter 类
|
根据 API format 获取 Adapter 类
|
||||||
|
|
||||||
@@ -745,7 +745,7 @@ def get_adapter_class(api_format: str) -> Optional[Type["ChatAdapterBase"]]:
|
|||||||
return _ADAPTER_REGISTRY.get(api_format.upper()) if api_format else None
|
return _ADAPTER_REGISTRY.get(api_format.upper()) if api_format else None
|
||||||
|
|
||||||
|
|
||||||
def get_adapter_instance(api_format: str) -> Optional["ChatAdapterBase"]:
|
def get_adapter_instance(api_format: str) -> ChatAdapterBase | None:
|
||||||
"""
|
"""
|
||||||
根据 API format 获取 Adapter 实例
|
根据 API format 获取 Adapter 实例
|
||||||
|
|
||||||
|
|||||||
@@ -22,7 +22,10 @@ Chat Handler Base - Chat API 格式的通用基类
|
|||||||
import asyncio
|
import asyncio
|
||||||
import json
|
import json
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from typing import Any, AsyncGenerator, Awaitable, Callable, Dict, Optional, Union
|
from typing import Any
|
||||||
|
|
||||||
|
from collections.abc import Callable
|
||||||
|
from collections.abc import AsyncGenerator, Awaitable
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
from fastapi import BackgroundTasks, Request
|
from fastapi import BackgroundTasks, Request
|
||||||
@@ -75,10 +78,10 @@ def _get_error_status_code(e: Exception, default: int = 400) -> int:
|
|||||||
|
|
||||||
|
|
||||||
def _convert_error_response_best_effort(
|
def _convert_error_response_best_effort(
|
||||||
error_response: Dict[str, Any],
|
error_response: dict[str, Any],
|
||||||
source_format: str,
|
source_format: str,
|
||||||
target_format: str,
|
target_format: str,
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
将上游错误响应 best-effort 转换为客户端格式。
|
将上游错误响应 best-effort 转换为客户端格式。
|
||||||
|
|
||||||
@@ -97,7 +100,7 @@ def _convert_error_response_best_effort(
|
|||||||
def _build_client_error_response_best_effort(
|
def _build_client_error_response_best_effort(
|
||||||
message: str,
|
message: str,
|
||||||
target_format: str,
|
target_format: str,
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
当无法解析上游错误 body 时,构造一个目标格式的错误响应(best-effort)。
|
当无法解析上游错误 body 时,构造一个目标格式的错误响应(best-effort)。
|
||||||
"""
|
"""
|
||||||
@@ -117,11 +120,11 @@ def _build_client_error_response_best_effort(
|
|||||||
|
|
||||||
|
|
||||||
def _build_error_json_payload(
|
def _build_error_json_payload(
|
||||||
e: Union[ThinkingSignatureException, UpstreamClientException],
|
e: ThinkingSignatureException | UpstreamClientException,
|
||||||
client_format: str,
|
client_format: str,
|
||||||
provider_format: str,
|
provider_format: str,
|
||||||
needs_conversion: bool = True,
|
needs_conversion: bool = True,
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
构建错误 JSON 响应 payload(公共逻辑)。
|
构建错误 JSON 响应 payload(公共逻辑)。
|
||||||
|
|
||||||
@@ -185,10 +188,10 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
client_ip: str,
|
client_ip: str,
|
||||||
user_agent: str,
|
user_agent: str,
|
||||||
start_time: float,
|
start_time: float,
|
||||||
allowed_api_formats: Optional[list] = None,
|
allowed_api_formats: list | None = None,
|
||||||
adapter_detector: Optional[
|
adapter_detector: None | (
|
||||||
Callable[[Dict[str, str], Optional[Dict[str, Any]]], Dict[str, bool]]
|
Callable[[dict[str, str], dict[str, Any] | None], dict[str, bool]]
|
||||||
] = None,
|
) = None,
|
||||||
):
|
):
|
||||||
allowed = allowed_api_formats or [self.FORMAT_ID]
|
allowed = allowed_api_formats or [self.FORMAT_ID]
|
||||||
super().__init__(
|
super().__init__(
|
||||||
@@ -202,7 +205,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
allowed_api_formats=allowed,
|
allowed_api_formats=allowed,
|
||||||
adapter_detector=adapter_detector,
|
adapter_detector=adapter_detector,
|
||||||
)
|
)
|
||||||
self._parser: Optional[ResponseParser] = None
|
self._parser: ResponseParser | None = None
|
||||||
self._request_builder = PassthroughRequestBuilder()
|
self._request_builder = PassthroughRequestBuilder()
|
||||||
|
|
||||||
@property
|
@property
|
||||||
@@ -228,7 +231,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def _extract_usage(self, response: Dict) -> Dict[str, int]:
|
def _extract_usage(self, response: dict) -> dict[str, int]:
|
||||||
"""
|
"""
|
||||||
从响应中提取 token 使用情况
|
从响应中提取 token 使用情况
|
||||||
|
|
||||||
@@ -241,7 +244,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
"""
|
"""
|
||||||
pass
|
pass
|
||||||
|
|
||||||
def _normalize_response(self, response: Dict) -> Dict:
|
def _normalize_response(self, response: dict) -> dict:
|
||||||
"""
|
"""
|
||||||
规范化响应(可选覆盖)
|
规范化响应(可选覆盖)
|
||||||
|
|
||||||
@@ -257,8 +260,8 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
|
|
||||||
def extract_model_from_request(
|
def extract_model_from_request(
|
||||||
self,
|
self,
|
||||||
request_body: Dict[str, Any],
|
request_body: dict[str, Any],
|
||||||
path_params: Optional[Dict[str, Any]] = None, # noqa: ARG002 - 子类使用
|
path_params: dict[str, Any] | None = None, # noqa: ARG002 - 子类使用
|
||||||
) -> str:
|
) -> str:
|
||||||
"""
|
"""
|
||||||
从请求中提取模型名 - 子类可覆盖
|
从请求中提取模型名 - 子类可覆盖
|
||||||
@@ -282,9 +285,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
|
|
||||||
def apply_mapped_model(
|
def apply_mapped_model(
|
||||||
self,
|
self,
|
||||||
request_body: Dict[str, Any],
|
request_body: dict[str, Any],
|
||||||
mapped_model: str, # noqa: ARG002 - 子类使用
|
mapped_model: str, # noqa: ARG002 - 子类使用
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
将映射后的模型名应用到请求体
|
将映射后的模型名应用到请求体
|
||||||
|
|
||||||
@@ -303,9 +306,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
|
|
||||||
def get_model_for_url(
|
def get_model_for_url(
|
||||||
self,
|
self,
|
||||||
request_body: Dict[str, Any],
|
request_body: dict[str, Any],
|
||||||
mapped_model: Optional[str],
|
mapped_model: str | None,
|
||||||
) -> Optional[str]:
|
) -> str | None:
|
||||||
"""
|
"""
|
||||||
获取用于 URL 路径的模型名
|
获取用于 URL 路径的模型名
|
||||||
|
|
||||||
@@ -323,8 +326,8 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
|
|
||||||
def prepare_provider_request_body(
|
def prepare_provider_request_body(
|
||||||
self,
|
self,
|
||||||
request_body: Dict[str, Any],
|
request_body: dict[str, Any],
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
准备发送给 Provider 的请求体 - 子类可覆盖
|
准备发送给 Provider 的请求体 - 子类可覆盖
|
||||||
|
|
||||||
@@ -341,9 +344,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
|
|
||||||
def _set_model_after_conversion(
|
def _set_model_after_conversion(
|
||||||
self,
|
self,
|
||||||
request_body: Dict[str, Any],
|
request_body: dict[str, Any],
|
||||||
provider_api_format: str,
|
provider_api_format: str,
|
||||||
mapped_model: Optional[str],
|
mapped_model: str | None,
|
||||||
fallback_model: str,
|
fallback_model: str,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""
|
"""
|
||||||
@@ -372,7 +375,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
|
|
||||||
def _set_stream_after_conversion(
|
def _set_stream_after_conversion(
|
||||||
self,
|
self,
|
||||||
request_body: Dict[str, Any],
|
request_body: dict[str, Any],
|
||||||
client_api_format: str,
|
client_api_format: str,
|
||||||
provider_api_format: str,
|
provider_api_format: str,
|
||||||
is_stream: bool,
|
is_stream: bool,
|
||||||
@@ -414,8 +417,8 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
self,
|
self,
|
||||||
source_model: str,
|
source_model: str,
|
||||||
provider_id: str,
|
provider_id: str,
|
||||||
api_format: Optional[str] = None,
|
api_format: str | None = None,
|
||||||
) -> Optional[str]:
|
) -> str | None:
|
||||||
"""
|
"""
|
||||||
获取模型映射后的实际模型名
|
获取模型映射后的实际模型名
|
||||||
|
|
||||||
@@ -452,10 +455,10 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
self,
|
self,
|
||||||
request: Any,
|
request: Any,
|
||||||
http_request: Request,
|
http_request: Request,
|
||||||
original_headers: Dict[str, Any],
|
original_headers: dict[str, Any],
|
||||||
original_request_body: Dict[str, Any],
|
original_request_body: dict[str, Any],
|
||||||
query_params: Optional[Dict[str, str]] = None,
|
query_params: dict[str, str] | None = None,
|
||||||
) -> Union[StreamingResponse, JSONResponse]:
|
) -> StreamingResponse | JSONResponse:
|
||||||
"""处理流式响应"""
|
"""处理流式响应"""
|
||||||
logger.debug(f"开始流式响应处理 ({self.FORMAT_ID})")
|
logger.debug(f"开始流式响应处理 ({self.FORMAT_ID})")
|
||||||
|
|
||||||
@@ -466,7 +469,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
|
|
||||||
# 可变请求体容器:允许 orchestrator 在遇到 Thinking 签名错误时整流请求体后重试
|
# 可变请求体容器:允许 orchestrator 在遇到 Thinking 签名错误时整流请求体后重试
|
||||||
# 结构: {"body": 实际请求体, "_rectified": 是否已整流, "_rectified_this_turn": 本轮是否整流}
|
# 结构: {"body": 实际请求体, "_rectified": 是否已整流, "_rectified_this_turn": 本轮是否整流}
|
||||||
request_body_ref: Dict[str, Any] = {"body": original_request_body}
|
request_body_ref: dict[str, Any] = {"body": original_request_body}
|
||||||
|
|
||||||
# 创建类型安全的流式上下文
|
# 创建类型安全的流式上下文
|
||||||
ctx = StreamContext(model=model, api_format=api_format)
|
ctx = StreamContext(model=model, api_format=api_format)
|
||||||
@@ -492,7 +495,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
endpoint: ProviderEndpoint,
|
endpoint: ProviderEndpoint,
|
||||||
key: ProviderAPIKey,
|
key: ProviderAPIKey,
|
||||||
candidate: ProviderCandidate,
|
candidate: ProviderCandidate,
|
||||||
) -> AsyncGenerator[bytes, None]:
|
) -> AsyncGenerator[bytes]:
|
||||||
return await self._execute_stream_request(
|
return await self._execute_stream_request(
|
||||||
ctx,
|
ctx,
|
||||||
stream_processor,
|
stream_processor,
|
||||||
@@ -615,12 +618,12 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
provider: Provider,
|
provider: Provider,
|
||||||
endpoint: ProviderEndpoint,
|
endpoint: ProviderEndpoint,
|
||||||
key: ProviderAPIKey,
|
key: ProviderAPIKey,
|
||||||
original_request_body: Dict[str, Any],
|
original_request_body: dict[str, Any],
|
||||||
original_headers: Dict[str, str],
|
original_headers: dict[str, str],
|
||||||
query_params: Optional[Dict[str, str]] = None,
|
query_params: dict[str, str] | None = None,
|
||||||
candidate: Optional[ProviderCandidate] = None,
|
candidate: ProviderCandidate | None = None,
|
||||||
is_disconnected: Optional[Callable[[], Awaitable[bool]]] = None,
|
is_disconnected: Callable[[], Awaitable[bool]] | None = None,
|
||||||
) -> AsyncGenerator[bytes, None]:
|
) -> AsyncGenerator[bytes]:
|
||||||
"""执行流式请求并返回流生成器"""
|
"""执行流式请求并返回流生成器"""
|
||||||
# 重置上下文状态(重试时清除之前的数据)
|
# 重置上下文状态(重试时清除之前的数据)
|
||||||
ctx.reset_for_retry()
|
ctx.reset_for_retry()
|
||||||
@@ -799,7 +802,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
ctx.error_message = "client_disconnected_during_prefetch"
|
ctx.error_message = "client_disconnected_during_prefetch"
|
||||||
raise
|
raise
|
||||||
|
|
||||||
except asyncio.TimeoutError:
|
except TimeoutError:
|
||||||
# 整体请求超时(建立连接 + 获取首字节)
|
# 整体请求超时(建立连接 + 获取首字节)
|
||||||
# 清理可能已建立的连接上下文
|
# 清理可能已建立的连接上下文
|
||||||
if response_ctx is not None:
|
if response_ctx is not None:
|
||||||
@@ -856,8 +859,8 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
self,
|
self,
|
||||||
ctx: StreamContext,
|
ctx: StreamContext,
|
||||||
error: Exception,
|
error: Exception,
|
||||||
original_headers: Dict[str, str],
|
original_headers: dict[str, str],
|
||||||
original_request_body: Dict[str, Any],
|
original_request_body: dict[str, Any],
|
||||||
) -> None:
|
) -> None:
|
||||||
"""记录流式请求失败"""
|
"""记录流式请求失败"""
|
||||||
response_time_ms = self.elapsed_ms()
|
response_time_ms = self.elapsed_ms()
|
||||||
@@ -904,9 +907,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
self,
|
self,
|
||||||
request: Any,
|
request: Any,
|
||||||
http_request: Request,
|
http_request: Request,
|
||||||
original_headers: Dict[str, Any],
|
original_headers: dict[str, Any],
|
||||||
original_request_body: Dict[str, Any],
|
original_request_body: dict[str, Any],
|
||||||
query_params: Optional[Dict[str, str]] = None,
|
query_params: dict[str, str] | None = None,
|
||||||
) -> JSONResponse:
|
) -> JSONResponse:
|
||||||
"""处理非流式响应"""
|
"""处理非流式响应"""
|
||||||
logger.debug(f"开始非流式响应处理 ({self.FORMAT_ID})")
|
logger.debug(f"开始非流式响应处理 ({self.FORMAT_ID})")
|
||||||
@@ -918,29 +921,29 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
|
|
||||||
# 可变请求体容器:允许 orchestrator 在遇到 Thinking 签名错误时整流请求体后重试
|
# 可变请求体容器:允许 orchestrator 在遇到 Thinking 签名错误时整流请求体后重试
|
||||||
# 结构: {"body": 实际请求体, "_rectified": 是否已整流, "_rectified_this_turn": 本轮是否整流}
|
# 结构: {"body": 实际请求体, "_rectified": 是否已整流, "_rectified_this_turn": 本轮是否整流}
|
||||||
request_body_ref: Dict[str, Any] = {"body": original_request_body}
|
request_body_ref: dict[str, Any] = {"body": original_request_body}
|
||||||
|
|
||||||
# 用于跟踪的变量
|
# 用于跟踪的变量
|
||||||
provider_name: Optional[str] = None
|
provider_name: str | None = None
|
||||||
response_json: Optional[Dict[str, Any]] = None
|
response_json: dict[str, Any] | None = None
|
||||||
status_code = 200
|
status_code = 200
|
||||||
response_headers: Dict[str, str] = {}
|
response_headers: dict[str, str] = {}
|
||||||
provider_request_headers: Dict[str, str] = {}
|
provider_request_headers: dict[str, str] = {}
|
||||||
provider_request_body: Optional[Dict[str, Any]] = None
|
provider_request_body: dict[str, Any] | None = None
|
||||||
provider_api_format_for_error: Optional[str] = None
|
provider_api_format_for_error: str | None = None
|
||||||
client_api_format_for_error: Optional[str] = None
|
client_api_format_for_error: str | None = None
|
||||||
needs_conversion_for_error: bool = False
|
needs_conversion_for_error: bool = False
|
||||||
provider_id: Optional[str] = None # Provider ID(用于失败记录)
|
provider_id: str | None = None # Provider ID(用于失败记录)
|
||||||
endpoint_id: Optional[str] = None # Endpoint ID(用于失败记录)
|
endpoint_id: str | None = None # Endpoint ID(用于失败记录)
|
||||||
key_id: Optional[str] = None # Key ID(用于失败记录)
|
key_id: str | None = None # Key ID(用于失败记录)
|
||||||
mapped_model_result: Optional[str] = None # 映射后的目标模型名(用于 Usage 记录)
|
mapped_model_result: str | None = None # 映射后的目标模型名(用于 Usage 记录)
|
||||||
|
|
||||||
async def sync_request_func(
|
async def sync_request_func(
|
||||||
provider: Provider,
|
provider: Provider,
|
||||||
endpoint: ProviderEndpoint,
|
endpoint: ProviderEndpoint,
|
||||||
key: ProviderAPIKey,
|
key: ProviderAPIKey,
|
||||||
candidate: ProviderCandidate,
|
candidate: ProviderCandidate,
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
nonlocal provider_name, response_json, status_code, response_headers
|
nonlocal provider_name, response_json, status_code, response_headers
|
||||||
nonlocal provider_request_headers, provider_request_body, mapped_model_result
|
nonlocal provider_request_headers, provider_request_body, mapped_model_result
|
||||||
nonlocal provider_api_format_for_error, client_api_format_for_error, needs_conversion_for_error
|
nonlocal provider_api_format_for_error, client_api_format_for_error, needs_conversion_for_error
|
||||||
@@ -1293,7 +1296,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
|||||||
actual_request_body = provider_request_body or original_request_body
|
actual_request_body = provider_request_body or original_request_body
|
||||||
|
|
||||||
# 尝试从异常中提取响应头
|
# 尝试从异常中提取响应头
|
||||||
error_response_headers: Dict[str, str] = {}
|
error_response_headers: dict[str, str] = {}
|
||||||
if isinstance(e, ProviderRateLimitException) and e.response_headers:
|
if isinstance(e, ProviderRateLimitException) and e.response_headers:
|
||||||
error_response_headers = e.response_headers
|
error_response_headers = e.response_headers
|
||||||
elif isinstance(e, httpx.HTTPStatusError) and hasattr(e, "response"):
|
elif isinstance(e, httpx.HTTPStatusError) and hasattr(e, "response"):
|
||||||
|
|||||||
@@ -17,7 +17,7 @@ CLI Adapter 通用基类
|
|||||||
|
|
||||||
import time
|
import time
|
||||||
import traceback
|
import traceback
|
||||||
from typing import Any, Dict, Optional, Tuple, Type
|
from typing import Any
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
from fastapi import HTTPException, Request
|
from fastapi import HTTPException, Request
|
||||||
@@ -63,7 +63,7 @@ class CliAdapterBase(ApiAdapter):
|
|||||||
|
|
||||||
# 子类必须覆盖
|
# 子类必须覆盖
|
||||||
FORMAT_ID: str = "UNKNOWN"
|
FORMAT_ID: str = "UNKNOWN"
|
||||||
HANDLER_CLASS: Type[CliMessageHandlerBase]
|
HANDLER_CLASS: type[CliMessageHandlerBase]
|
||||||
|
|
||||||
# 适配器配置
|
# 适配器配置
|
||||||
name: str = "cli.base"
|
name: str = "cli.base"
|
||||||
@@ -72,7 +72,7 @@ class CliAdapterBase(ApiAdapter):
|
|||||||
# 计费模板配置(子类可覆盖,如 "claude", "openai", "gemini")
|
# 计费模板配置(子类可覆盖,如 "claude", "openai", "gemini")
|
||||||
BILLING_TEMPLATE: str = "claude"
|
BILLING_TEMPLATE: str = "claude"
|
||||||
|
|
||||||
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
|
def __init__(self, allowed_api_formats: list[str] | None = None):
|
||||||
self.allowed_api_formats = allowed_api_formats or [self.FORMAT_ID]
|
self.allowed_api_formats = allowed_api_formats or [self.FORMAT_ID]
|
||||||
|
|
||||||
# =========================================================================
|
# =========================================================================
|
||||||
@@ -87,7 +87,7 @@ class CliAdapterBase(ApiAdapter):
|
|||||||
except KeyError:
|
except KeyError:
|
||||||
return APIFormat.OPENAI
|
return APIFormat.OPENAI
|
||||||
|
|
||||||
def extract_api_key(self, request: Request) -> Optional[str]:
|
def extract_api_key(self, request: Request) -> str | None:
|
||||||
"""
|
"""
|
||||||
从请求中提取 API 密钥
|
从请求中提取 API 密钥
|
||||||
|
|
||||||
@@ -96,7 +96,7 @@ class CliAdapterBase(ApiAdapter):
|
|||||||
return extract_client_api_key(dict(request.headers), self._get_api_format())
|
return extract_client_api_key(dict(request.headers), self._get_api_format())
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def build_base_headers(cls, api_key: str) -> Dict[str, str]:
|
def build_base_headers(cls, api_key: str) -> dict[str, str]:
|
||||||
"""
|
"""
|
||||||
构建 CLI API 认证头
|
构建 CLI API 认证头
|
||||||
|
|
||||||
@@ -106,8 +106,8 @@ class CliAdapterBase(ApiAdapter):
|
|||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def build_headers_with_extra(
|
def build_headers_with_extra(
|
||||||
cls, api_key: str, extra_headers: Optional[Dict[str, str]] = None
|
cls, api_key: str, extra_headers: dict[str, str] | None = None
|
||||||
) -> Dict[str, str]:
|
) -> dict[str, str]:
|
||||||
"""
|
"""
|
||||||
构建带额外头部的完整请求头
|
构建带额外头部的完整请求头
|
||||||
|
|
||||||
@@ -260,8 +260,8 @@ class CliAdapterBase(ApiAdapter):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def _merge_path_params(
|
def _merge_path_params(
|
||||||
self, original_request_body: Dict[str, Any], path_params: Dict[str, Any]
|
self, original_request_body: dict[str, Any], path_params: dict[str, Any]
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
合并 URL 路径参数到请求体 - 子类可覆盖
|
合并 URL 路径参数到请求体 - 子类可覆盖
|
||||||
|
|
||||||
@@ -280,7 +280,7 @@ class CliAdapterBase(ApiAdapter):
|
|||||||
merged[key] = value
|
merged[key] = value
|
||||||
return merged
|
return merged
|
||||||
|
|
||||||
def _extract_message_count(self, payload: Dict[str, Any]) -> int:
|
def _extract_message_count(self, payload: dict[str, Any]) -> int:
|
||||||
"""
|
"""
|
||||||
提取消息数量 - 子类可覆盖
|
提取消息数量 - 子类可覆盖
|
||||||
|
|
||||||
@@ -297,9 +297,9 @@ class CliAdapterBase(ApiAdapter):
|
|||||||
|
|
||||||
def _build_audit_metadata(
|
def _build_audit_metadata(
|
||||||
self,
|
self,
|
||||||
payload: Dict[str, Any],
|
payload: dict[str, Any],
|
||||||
path_params: Optional[Dict[str, Any]] = None,
|
path_params: dict[str, Any] | None = None,
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
构建审计日志元数据 - 子类可覆盖
|
构建审计日志元数据 - 子类可覆盖
|
||||||
|
|
||||||
@@ -338,8 +338,8 @@ class CliAdapterBase(ApiAdapter):
|
|||||||
model: str,
|
model: str,
|
||||||
stream: bool,
|
stream: bool,
|
||||||
start_time: float,
|
start_time: float,
|
||||||
original_headers: Dict[str, str],
|
original_headers: dict[str, str],
|
||||||
original_request_body: Dict[str, Any],
|
original_request_body: dict[str, Any],
|
||||||
client_ip: str,
|
client_ip: str,
|
||||||
request_id: str,
|
request_id: str,
|
||||||
) -> JSONResponse:
|
) -> JSONResponse:
|
||||||
@@ -409,8 +409,8 @@ class CliAdapterBase(ApiAdapter):
|
|||||||
model: str,
|
model: str,
|
||||||
stream: bool,
|
stream: bool,
|
||||||
start_time: float,
|
start_time: float,
|
||||||
original_headers: Dict[str, str],
|
original_headers: dict[str, str],
|
||||||
original_request_body: Dict[str, Any],
|
original_request_body: dict[str, Any],
|
||||||
client_ip: str,
|
client_ip: str,
|
||||||
request_id: str,
|
request_id: str,
|
||||||
) -> JSONResponse:
|
) -> JSONResponse:
|
||||||
@@ -507,12 +507,12 @@ class CliAdapterBase(ApiAdapter):
|
|||||||
cache_read_input_tokens: int,
|
cache_read_input_tokens: int,
|
||||||
input_price_per_1m: float,
|
input_price_per_1m: float,
|
||||||
output_price_per_1m: float,
|
output_price_per_1m: float,
|
||||||
cache_creation_price_per_1m: Optional[float],
|
cache_creation_price_per_1m: float | None,
|
||||||
cache_read_price_per_1m: Optional[float],
|
cache_read_price_per_1m: float | None,
|
||||||
price_per_request: Optional[float],
|
price_per_request: float | None,
|
||||||
tiered_pricing: Optional[dict] = None,
|
tiered_pricing: dict | None = None,
|
||||||
cache_ttl_minutes: Optional[int] = None,
|
cache_ttl_minutes: int | None = None,
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
计算请求成本
|
计算请求成本
|
||||||
|
|
||||||
@@ -567,8 +567,8 @@ class CliAdapterBase(ApiAdapter):
|
|||||||
client: httpx.AsyncClient,
|
client: httpx.AsyncClient,
|
||||||
base_url: str,
|
base_url: str,
|
||||||
api_key: str,
|
api_key: str,
|
||||||
extra_headers: Optional[Dict[str, str]] = None,
|
extra_headers: dict[str, str] | None = None,
|
||||||
) -> Tuple[list, Optional[str]]:
|
) -> tuple[list, str | None]:
|
||||||
"""
|
"""
|
||||||
查询上游 API 支持的模型列表
|
查询上游 API 支持的模型列表
|
||||||
|
|
||||||
@@ -596,16 +596,16 @@ class CliAdapterBase(ApiAdapter):
|
|||||||
client: httpx.AsyncClient,
|
client: httpx.AsyncClient,
|
||||||
base_url: str,
|
base_url: str,
|
||||||
api_key: str,
|
api_key: str,
|
||||||
request_data: Dict[str, Any],
|
request_data: dict[str, Any],
|
||||||
extra_headers: Optional[Dict[str, str]] = None,
|
extra_headers: dict[str, str] | None = None,
|
||||||
# 用量计算参数
|
# 用量计算参数
|
||||||
db: Optional[Any] = None,
|
db: Any | None = None,
|
||||||
user: Optional[Any] = None,
|
user: Any | None = None,
|
||||||
provider_name: Optional[str] = None,
|
provider_name: str | None = None,
|
||||||
provider_id: Optional[str] = None,
|
provider_id: str | None = None,
|
||||||
api_key_id: Optional[str] = None,
|
api_key_id: str | None = None,
|
||||||
model_name: Optional[str] = None,
|
model_name: str | None = None,
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
测试模型连接性(非流式)
|
测试模型连接性(非流式)
|
||||||
|
|
||||||
@@ -669,7 +669,7 @@ class CliAdapterBase(ApiAdapter):
|
|||||||
# =========================================================================
|
# =========================================================================
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def build_endpoint_url(cls, base_url: str, request_data: Dict[str, Any], model_name: Optional[str] = None) -> str:
|
def build_endpoint_url(cls, base_url: str, request_data: dict[str, Any], model_name: str | None = None) -> str:
|
||||||
"""
|
"""
|
||||||
构建CLI API端点URL - 子类应覆盖
|
构建CLI API端点URL - 子类应覆盖
|
||||||
|
|
||||||
@@ -684,7 +684,7 @@ class CliAdapterBase(ApiAdapter):
|
|||||||
raise NotImplementedError(f"{cls.FORMAT_ID} adapter must implement build_endpoint_url")
|
raise NotImplementedError(f"{cls.FORMAT_ID} adapter must implement build_endpoint_url")
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def build_request_body(cls, request_data: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
|
def build_request_body(cls, request_data: dict[str, Any] | None = None) -> dict[str, Any]:
|
||||||
"""构建测试请求体,使用转换器注册表自动处理格式转换
|
"""构建测试请求体,使用转换器注册表自动处理格式转换
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
@@ -698,7 +698,7 @@ class CliAdapterBase(ApiAdapter):
|
|||||||
return build_test_request_body(cls.FORMAT_ID, request_data)
|
return build_test_request_body(cls.FORMAT_ID, request_data)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def get_cli_user_agent(cls) -> Optional[str]:
|
def get_cli_user_agent(cls) -> str | None:
|
||||||
"""
|
"""
|
||||||
获取CLI User-Agent - 子类可覆盖
|
获取CLI User-Agent - 子类可覆盖
|
||||||
|
|
||||||
@@ -708,7 +708,7 @@ class CliAdapterBase(ApiAdapter):
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def get_cli_extra_headers(cls) -> Dict[str, str]:
|
def get_cli_extra_headers(cls) -> dict[str, str]:
|
||||||
"""
|
"""
|
||||||
获取CLI额外请求头 - 子类可覆盖
|
获取CLI额外请求头 - 子类可覆盖
|
||||||
|
|
||||||
@@ -718,7 +718,7 @@ class CliAdapterBase(ApiAdapter):
|
|||||||
Returns:
|
Returns:
|
||||||
额外请求头字典
|
额外请求头字典
|
||||||
"""
|
"""
|
||||||
headers: Dict[str, str] = {}
|
headers: dict[str, str] = {}
|
||||||
cli_user_agent = cls.get_cli_user_agent()
|
cli_user_agent = cls.get_cli_user_agent()
|
||||||
if cli_user_agent:
|
if cli_user_agent:
|
||||||
headers["User-Agent"] = cli_user_agent
|
headers["User-Agent"] = cli_user_agent
|
||||||
@@ -728,11 +728,11 @@ class CliAdapterBase(ApiAdapter):
|
|||||||
# CLI Adapter 注册表 - 用于根据 API format 获取 CLI Adapter 实例
|
# CLI Adapter 注册表 - 用于根据 API format 获取 CLI Adapter 实例
|
||||||
# =========================================================================
|
# =========================================================================
|
||||||
|
|
||||||
_CLI_ADAPTER_REGISTRY: Dict[str, Type["CliAdapterBase"]] = {}
|
_CLI_ADAPTER_REGISTRY: dict[str, type[CliAdapterBase]] = {}
|
||||||
_CLI_ADAPTERS_LOADED = False
|
_CLI_ADAPTERS_LOADED = False
|
||||||
|
|
||||||
|
|
||||||
def register_cli_adapter(adapter_class: Type["CliAdapterBase"]) -> Type["CliAdapterBase"]:
|
def register_cli_adapter(adapter_class: type[CliAdapterBase]) -> type[CliAdapterBase]:
|
||||||
"""
|
"""
|
||||||
注册 CLI Adapter 类到注册表
|
注册 CLI Adapter 类到注册表
|
||||||
|
|
||||||
@@ -771,13 +771,13 @@ def _ensure_cli_adapters_loaded():
|
|||||||
_CLI_ADAPTERS_LOADED = True
|
_CLI_ADAPTERS_LOADED = True
|
||||||
|
|
||||||
|
|
||||||
def get_cli_adapter_class(api_format: str) -> Optional[Type["CliAdapterBase"]]:
|
def get_cli_adapter_class(api_format: str) -> type[CliAdapterBase] | None:
|
||||||
"""根据 API format 获取 CLI Adapter 类"""
|
"""根据 API format 获取 CLI Adapter 类"""
|
||||||
_ensure_cli_adapters_loaded()
|
_ensure_cli_adapters_loaded()
|
||||||
return _CLI_ADAPTER_REGISTRY.get(api_format.upper()) if api_format else None
|
return _CLI_ADAPTER_REGISTRY.get(api_format.upper()) if api_format else None
|
||||||
|
|
||||||
|
|
||||||
def get_cli_adapter_instance(api_format: str) -> Optional["CliAdapterBase"]:
|
def get_cli_adapter_instance(api_format: str) -> CliAdapterBase | None:
|
||||||
"""根据 API format 获取 CLI Adapter 实例"""
|
"""根据 API format 获取 CLI Adapter 实例"""
|
||||||
adapter_class = get_cli_adapter_class(api_format)
|
adapter_class = get_cli_adapter_class(api_format)
|
||||||
if adapter_class:
|
if adapter_class:
|
||||||
|
|||||||
@@ -10,6 +10,8 @@ CLI Message Handler 通用基类
|
|||||||
3. 简化新格式接入 - 只需实现 ResponseParser 和少量钩子方法
|
3. 简化新格式接入 - 只需实现 ResponseParser 和少量钩子方法
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import codecs
|
import codecs
|
||||||
import json
|
import json
|
||||||
@@ -17,14 +19,11 @@ import time
|
|||||||
from typing import (
|
from typing import (
|
||||||
TYPE_CHECKING,
|
TYPE_CHECKING,
|
||||||
Any,
|
Any,
|
||||||
AsyncGenerator,
|
|
||||||
Callable,
|
|
||||||
Dict,
|
|
||||||
List,
|
|
||||||
Optional,
|
|
||||||
Tuple,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
from collections.abc import Callable
|
||||||
|
from collections.abc import AsyncGenerator
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
from fastapi import BackgroundTasks, Request
|
from fastapi import BackgroundTasks, Request
|
||||||
from fastapi.responses import JSONResponse, StreamingResponse
|
from fastapi.responses import JSONResponse, StreamingResponse
|
||||||
@@ -45,7 +44,6 @@ from src.api.handlers.base.request_builder import PassthroughRequestBuilder, get
|
|||||||
# 直接从具体模块导入,避免循环依赖
|
# 直接从具体模块导入,避免循环依赖
|
||||||
from src.api.handlers.base.response_parser import (
|
from src.api.handlers.base.response_parser import (
|
||||||
ResponseParser,
|
ResponseParser,
|
||||||
StreamStats,
|
|
||||||
)
|
)
|
||||||
from src.api.handlers.base.stream_context import StreamContext
|
from src.api.handlers.base.stream_context import StreamContext
|
||||||
from src.api.handlers.base.utils import (
|
from src.api.handlers.base.utils import (
|
||||||
@@ -86,7 +84,7 @@ from src.utils.timeout import read_first_chunk_with_ttfb_timeout
|
|||||||
# ==============================================================================
|
# ==============================================================================
|
||||||
|
|
||||||
|
|
||||||
def _parse_sse_data_line(line: str) -> Tuple[Optional[Any], str]:
|
def _parse_sse_data_line(line: str) -> tuple[Any | None, str]:
|
||||||
"""
|
"""
|
||||||
解析标准 SSE data 行
|
解析标准 SSE data 行
|
||||||
|
|
||||||
@@ -108,7 +106,7 @@ def _parse_sse_data_line(line: str) -> Tuple[Optional[Any], str]:
|
|||||||
return None, "invalid"
|
return None, "invalid"
|
||||||
|
|
||||||
|
|
||||||
def _parse_sse_event_data_line(line: str) -> Tuple[Optional[Any], str]:
|
def _parse_sse_event_data_line(line: str) -> tuple[Any | None, str]:
|
||||||
"""
|
"""
|
||||||
解析 event + data 同行格式(如 "event: xxx data: {...}")
|
解析 event + data 同行格式(如 "event: xxx data: {...}")
|
||||||
|
|
||||||
@@ -126,7 +124,7 @@ def _parse_sse_event_data_line(line: str) -> Tuple[Optional[Any], str]:
|
|||||||
return None, "invalid"
|
return None, "invalid"
|
||||||
|
|
||||||
|
|
||||||
def _parse_gemini_json_array_line(line: str) -> Tuple[Optional[Any], str]:
|
def _parse_gemini_json_array_line(line: str) -> tuple[Any | None, str]:
|
||||||
"""
|
"""
|
||||||
解析 Gemini JSON-array 格式的裸 JSON 行
|
解析 Gemini JSON-array 格式的裸 JSON 行
|
||||||
|
|
||||||
@@ -151,9 +149,9 @@ def _parse_gemini_json_array_line(line: str) -> Tuple[Optional[Any], str]:
|
|||||||
|
|
||||||
|
|
||||||
def _format_converted_events_to_sse(
|
def _format_converted_events_to_sse(
|
||||||
converted_events: List[Dict[str, Any]],
|
converted_events: list[dict[str, Any]],
|
||||||
client_format: str,
|
client_format: str,
|
||||||
) -> List[str]:
|
) -> list[str]:
|
||||||
"""
|
"""
|
||||||
将转换后的事件格式化为 SSE 行
|
将转换后的事件格式化为 SSE 行
|
||||||
|
|
||||||
@@ -164,7 +162,7 @@ def _format_converted_events_to_sse(
|
|||||||
Returns:
|
Returns:
|
||||||
SSE 行列表(每个元素是完整的 SSE 事件,包含尾部空行)
|
SSE 行列表(每个元素是完整的 SSE 事件,包含尾部空行)
|
||||||
"""
|
"""
|
||||||
result: List[str] = []
|
result: list[str] = []
|
||||||
needs_event_line = client_format.upper() in ("CLAUDE", "CLAUDE_CLI")
|
needs_event_line = client_format.upper() in ("CLAUDE", "CLAUDE_CLI")
|
||||||
|
|
||||||
for evt in converted_events:
|
for evt in converted_events:
|
||||||
@@ -213,10 +211,10 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
client_ip: str,
|
client_ip: str,
|
||||||
user_agent: str,
|
user_agent: str,
|
||||||
start_time: float,
|
start_time: float,
|
||||||
allowed_api_formats: Optional[list] = None,
|
allowed_api_formats: list | None = None,
|
||||||
adapter_detector: Optional[
|
adapter_detector: None | (
|
||||||
Callable[[Dict[str, str], Optional[Dict[str, Any]]], Dict[str, bool]]
|
Callable[[dict[str, str], dict[str, Any] | None], dict[str, bool]]
|
||||||
] = None,
|
) = None,
|
||||||
):
|
):
|
||||||
allowed = allowed_api_formats or [self.FORMAT_ID]
|
allowed = allowed_api_formats or [self.FORMAT_ID]
|
||||||
super().__init__(
|
super().__init__(
|
||||||
@@ -230,7 +228,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
allowed_api_formats=allowed,
|
allowed_api_formats=allowed,
|
||||||
adapter_detector=adapter_detector,
|
adapter_detector=adapter_detector,
|
||||||
)
|
)
|
||||||
self._parser: Optional[ResponseParser] = None
|
self._parser: ResponseParser | None = None
|
||||||
self._request_builder = PassthroughRequestBuilder()
|
self._request_builder = PassthroughRequestBuilder()
|
||||||
|
|
||||||
@property
|
@property
|
||||||
@@ -253,7 +251,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
self,
|
self,
|
||||||
source_model: str,
|
source_model: str,
|
||||||
provider_id: str,
|
provider_id: str,
|
||||||
) -> Optional[str]:
|
) -> str | None:
|
||||||
"""
|
"""
|
||||||
获取模型映射后的实际模型名
|
获取模型映射后的实际模型名
|
||||||
|
|
||||||
@@ -296,8 +294,8 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
|
|
||||||
def extract_model_from_request(
|
def extract_model_from_request(
|
||||||
self,
|
self,
|
||||||
request_body: Dict[str, Any],
|
request_body: dict[str, Any],
|
||||||
path_params: Optional[Dict[str, Any]] = None, # noqa: ARG002 - 子类使用
|
path_params: dict[str, Any] | None = None, # noqa: ARG002 - 子类使用
|
||||||
) -> str:
|
) -> str:
|
||||||
"""
|
"""
|
||||||
从请求中提取模型名 - 子类可覆盖
|
从请求中提取模型名 - 子类可覆盖
|
||||||
@@ -321,9 +319,9 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
|
|
||||||
def apply_mapped_model(
|
def apply_mapped_model(
|
||||||
self,
|
self,
|
||||||
request_body: Dict[str, Any],
|
request_body: dict[str, Any],
|
||||||
mapped_model: str, # noqa: ARG002 - 子类使用
|
mapped_model: str, # noqa: ARG002 - 子类使用
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
将映射后的模型名应用到请求体
|
将映射后的模型名应用到请求体
|
||||||
|
|
||||||
@@ -342,8 +340,8 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
|
|
||||||
def prepare_provider_request_body(
|
def prepare_provider_request_body(
|
||||||
self,
|
self,
|
||||||
request_body: Dict[str, Any],
|
request_body: dict[str, Any],
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
准备发送给 Provider 的请求体 - 子类可覆盖
|
准备发送给 Provider 的请求体 - 子类可覆盖
|
||||||
|
|
||||||
@@ -359,7 +357,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
return request_body
|
return request_body
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _get_format_metadata(format_id: str) -> Optional["ApiFormatDefinition"]:
|
def _get_format_metadata(format_id: str) -> ApiFormatDefinition | None:
|
||||||
"""获取格式元数据(解析失败返回 None)"""
|
"""获取格式元数据(解析失败返回 None)"""
|
||||||
from src.core.api_format import APIFormat
|
from src.core.api_format import APIFormat
|
||||||
from src.core.api_format.metadata import API_FORMAT_DEFINITIONS
|
from src.core.api_format.metadata import API_FORMAT_DEFINITIONS
|
||||||
@@ -372,10 +370,10 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
|
|
||||||
def _finalize_converted_request(
|
def _finalize_converted_request(
|
||||||
self,
|
self,
|
||||||
request_body: Dict[str, Any],
|
request_body: dict[str, Any],
|
||||||
client_api_format: str,
|
client_api_format: str,
|
||||||
provider_api_format: str,
|
provider_api_format: str,
|
||||||
mapped_model: Optional[str],
|
mapped_model: str | None,
|
||||||
fallback_model: str,
|
fallback_model: str,
|
||||||
is_stream: bool,
|
is_stream: bool,
|
||||||
) -> None:
|
) -> None:
|
||||||
@@ -418,13 +416,13 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
|
|
||||||
def _convert_request_for_cross_format(
|
def _convert_request_for_cross_format(
|
||||||
self,
|
self,
|
||||||
request_body: Dict[str, Any],
|
request_body: dict[str, Any],
|
||||||
client_api_format: str,
|
client_api_format: str,
|
||||||
provider_api_format: str,
|
provider_api_format: str,
|
||||||
mapped_model: Optional[str],
|
mapped_model: str | None,
|
||||||
fallback_model: str,
|
fallback_model: str,
|
||||||
is_stream: bool,
|
is_stream: bool,
|
||||||
) -> Tuple[Dict[str, Any], str]:
|
) -> tuple[dict[str, Any], str]:
|
||||||
"""
|
"""
|
||||||
跨格式请求转换的公共逻辑
|
跨格式请求转换的公共逻辑
|
||||||
|
|
||||||
@@ -465,9 +463,9 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
|
|
||||||
def get_model_for_url(
|
def get_model_for_url(
|
||||||
self,
|
self,
|
||||||
request_body: Dict[str, Any],
|
request_body: dict[str, Any],
|
||||||
mapped_model: Optional[str],
|
mapped_model: str | None,
|
||||||
) -> Optional[str]:
|
) -> str | None:
|
||||||
"""
|
"""
|
||||||
获取用于 URL 路径的模型名
|
获取用于 URL 路径的模型名
|
||||||
|
|
||||||
@@ -485,8 +483,8 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
|
|
||||||
def _extract_response_metadata(
|
def _extract_response_metadata(
|
||||||
self,
|
self,
|
||||||
response: Dict[str, Any],
|
response: dict[str, Any],
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
从响应中提取 Provider 特有的元数据 - 子类可覆盖
|
从响应中提取 Provider 特有的元数据 - 子类可覆盖
|
||||||
|
|
||||||
@@ -503,11 +501,11 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
|
|
||||||
async def process_stream(
|
async def process_stream(
|
||||||
self,
|
self,
|
||||||
original_request_body: Dict[str, Any],
|
original_request_body: dict[str, Any],
|
||||||
original_headers: Dict[str, str],
|
original_headers: dict[str, str],
|
||||||
query_params: Optional[Dict[str, str]] = None,
|
query_params: dict[str, str] | None = None,
|
||||||
path_params: Optional[Dict[str, Any]] = None,
|
path_params: dict[str, Any] | None = None,
|
||||||
http_request: Optional[Request] = None,
|
http_request: Request | None = None,
|
||||||
) -> StreamingResponse:
|
) -> StreamingResponse:
|
||||||
"""
|
"""
|
||||||
处理流式请求
|
处理流式请求
|
||||||
@@ -529,7 +527,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
|
|
||||||
# 可变请求体容器:允许 orchestrator 在遇到 Thinking 签名错误时整流请求体后重试
|
# 可变请求体容器:允许 orchestrator 在遇到 Thinking 签名错误时整流请求体后重试
|
||||||
# 结构: {"body": 实际请求体, "_rectified": 是否已整流, "_rectified_this_turn": 本轮是否整流}
|
# 结构: {"body": 实际请求体, "_rectified": 是否已整流, "_rectified_this_turn": 本轮是否整流}
|
||||||
request_body_ref: Dict[str, Any] = {"body": original_request_body}
|
request_body_ref: dict[str, Any] = {"body": original_request_body}
|
||||||
|
|
||||||
# 使用子类实现的方法提取 model(不同 API 格式的 model 位置不同)
|
# 使用子类实现的方法提取 model(不同 API 格式的 model 位置不同)
|
||||||
# 注意:使用 original_request_body,因为整流只修改 messages,不影响 model 字段
|
# 注意:使用 original_request_body,因为整流只修改 messages,不影响 model 字段
|
||||||
@@ -550,7 +548,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
endpoint: ProviderEndpoint,
|
endpoint: ProviderEndpoint,
|
||||||
key: ProviderAPIKey,
|
key: ProviderAPIKey,
|
||||||
candidate: ProviderCandidate,
|
candidate: ProviderCandidate,
|
||||||
) -> AsyncGenerator[bytes, None]:
|
) -> AsyncGenerator[bytes]:
|
||||||
return await self._execute_stream_request(
|
return await self._execute_stream_request(
|
||||||
ctx,
|
ctx,
|
||||||
provider,
|
provider,
|
||||||
@@ -653,12 +651,12 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
provider: Provider,
|
provider: Provider,
|
||||||
endpoint: ProviderEndpoint,
|
endpoint: ProviderEndpoint,
|
||||||
key: ProviderAPIKey,
|
key: ProviderAPIKey,
|
||||||
original_request_body: Dict[str, Any],
|
original_request_body: dict[str, Any],
|
||||||
original_headers: Dict[str, str],
|
original_headers: dict[str, str],
|
||||||
query_params: Optional[Dict[str, str]] = None,
|
query_params: dict[str, str] | None = None,
|
||||||
candidate: Optional[ProviderCandidate] = None,
|
candidate: ProviderCandidate | None = None,
|
||||||
http_request: Optional[Request] = None,
|
http_request: Request | None = None,
|
||||||
) -> AsyncGenerator[bytes, None]:
|
) -> AsyncGenerator[bytes]:
|
||||||
"""执行流式请求并返回流生成器"""
|
"""执行流式请求并返回流生成器"""
|
||||||
# 重置上下文状态(重试时清除之前的数据,避免累积)
|
# 重置上下文状态(重试时清除之前的数据,避免累积)
|
||||||
ctx.parsed_chunks = []
|
ctx.parsed_chunks = []
|
||||||
@@ -824,7 +822,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
else:
|
else:
|
||||||
await asyncio.wait_for(_connect_and_prefetch(), timeout=request_timeout)
|
await asyncio.wait_for(_connect_and_prefetch(), timeout=request_timeout)
|
||||||
|
|
||||||
except asyncio.TimeoutError:
|
except TimeoutError:
|
||||||
# 整体请求超时(建立连接 + 获取首字节)
|
# 整体请求超时(建立连接 + 获取首字节)
|
||||||
# 清理可能已建立的连接上下文
|
# 清理可能已建立的连接上下文
|
||||||
if response_ctx is not None:
|
if response_ctx is not None:
|
||||||
@@ -898,7 +896,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
stream_response: httpx.Response,
|
stream_response: httpx.Response,
|
||||||
response_ctx: Any,
|
response_ctx: Any,
|
||||||
http_client: httpx.AsyncClient,
|
http_client: httpx.AsyncClient,
|
||||||
) -> AsyncGenerator[bytes, None]:
|
) -> AsyncGenerator[bytes]:
|
||||||
"""创建响应流生成器(使用字节流)"""
|
"""创建响应流生成器(使用字节流)"""
|
||||||
try:
|
try:
|
||||||
sse_parser = SSEEventParser()
|
sse_parser = SSEEventParser()
|
||||||
@@ -961,9 +959,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
},
|
},
|
||||||
}
|
}
|
||||||
self._mark_first_output(ctx, output_state)
|
self._mark_first_output(ctx, output_state)
|
||||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode(
|
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
|
||||||
"utf-8"
|
|
||||||
)
|
|
||||||
return # 结束生成器
|
return # 结束生成器
|
||||||
|
|
||||||
# 格式转换或直接透传
|
# 格式转换或直接透传
|
||||||
@@ -1015,7 +1011,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
"message": ctx.error_message,
|
"message": ctx.error_message,
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode("utf-8")
|
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
|
||||||
else:
|
else:
|
||||||
logger.debug("流式数据转发完成")
|
logger.debug("流式数据转发完成")
|
||||||
# 为 OpenAI 客户端补齐 [DONE] 标记(非 CLI 格式)
|
# 为 OpenAI 客户端补齐 [DONE] 标记(非 CLI 格式)
|
||||||
@@ -1040,7 +1036,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
"message": ctx.error_message,
|
"message": ctx.error_message,
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode("utf-8")
|
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
|
||||||
except httpx.RemoteProtocolError:
|
except httpx.RemoteProtocolError:
|
||||||
if ctx.data_count > 0:
|
if ctx.data_count > 0:
|
||||||
error_event = {
|
error_event = {
|
||||||
@@ -1050,7 +1046,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
"message": "上游连接意外关闭,部分响应已成功传输",
|
"message": "上游连接意外关闭,部分响应已成功传输",
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode("utf-8")
|
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
|
||||||
else:
|
else:
|
||||||
raise
|
raise
|
||||||
finally:
|
finally:
|
||||||
@@ -1241,7 +1237,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
except (EmbeddedErrorException, ProviderTimeoutException, ProviderNotAvailableException):
|
except (EmbeddedErrorException, ProviderTimeoutException, ProviderNotAvailableException):
|
||||||
# 重新抛出可重试的 Provider 异常,触发故障转移
|
# 重新抛出可重试的 Provider 异常,触发故障转移
|
||||||
raise
|
raise
|
||||||
except (OSError, IOError) as e:
|
except OSError as e:
|
||||||
# 网络 I/O 异常:记录警告,可能需要重试
|
# 网络 I/O 异常:记录警告,可能需要重试
|
||||||
logger.warning(f" [{self.request_id}] 预读流时发生网络异常: {type(e).__name__}: {e}")
|
logger.warning(f" [{self.request_id}] 预读流时发生网络异常: {type(e).__name__}: {e}")
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -1261,7 +1257,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
response_ctx: Any,
|
response_ctx: Any,
|
||||||
http_client: httpx.AsyncClient,
|
http_client: httpx.AsyncClient,
|
||||||
prefetched_chunks: list,
|
prefetched_chunks: list,
|
||||||
) -> AsyncGenerator[bytes, None]:
|
) -> AsyncGenerator[bytes]:
|
||||||
"""创建响应流生成器(带预读数据,使用字节流)"""
|
"""创建响应流生成器(带预读数据,使用字节流)"""
|
||||||
try:
|
try:
|
||||||
sse_parser = SSEEventParser()
|
sse_parser = SSEEventParser()
|
||||||
@@ -1382,9 +1378,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
},
|
},
|
||||||
}
|
}
|
||||||
self._mark_first_output(ctx, output_state)
|
self._mark_first_output(ctx, output_state)
|
||||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode(
|
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
|
||||||
"utf-8"
|
|
||||||
)
|
|
||||||
return
|
return
|
||||||
|
|
||||||
# 格式转换或直接透传
|
# 格式转换或直接透传
|
||||||
@@ -1439,7 +1433,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
"message": ctx.error_message,
|
"message": ctx.error_message,
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode("utf-8")
|
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
|
||||||
else:
|
else:
|
||||||
logger.debug("流式数据转发完成")
|
logger.debug("流式数据转发完成")
|
||||||
# 为 OpenAI 客户端补齐 [DONE] 标记(非 CLI 格式)
|
# 为 OpenAI 客户端补齐 [DONE] 标记(非 CLI 格式)
|
||||||
@@ -1463,7 +1457,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
"message": ctx.error_message,
|
"message": ctx.error_message,
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode("utf-8")
|
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
|
||||||
except httpx.RemoteProtocolError:
|
except httpx.RemoteProtocolError:
|
||||||
if ctx.data_count > 0:
|
if ctx.data_count > 0:
|
||||||
error_event = {
|
error_event = {
|
||||||
@@ -1473,7 +1467,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
"message": "上游连接意外关闭,部分响应已成功传输",
|
"message": "上游连接意外关闭,部分响应已成功传输",
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode("utf-8")
|
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
|
||||||
else:
|
else:
|
||||||
raise
|
raise
|
||||||
finally:
|
finally:
|
||||||
@@ -1489,7 +1483,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
def _handle_sse_event(
|
def _handle_sse_event(
|
||||||
self,
|
self,
|
||||||
ctx: StreamContext,
|
ctx: StreamContext,
|
||||||
event_name: Optional[str],
|
event_name: str | None,
|
||||||
data_str: str,
|
data_str: str,
|
||||||
record_chunk: bool = False,
|
record_chunk: bool = False,
|
||||||
) -> None:
|
) -> None:
|
||||||
@@ -1538,7 +1532,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
self,
|
self,
|
||||||
ctx: StreamContext,
|
ctx: StreamContext,
|
||||||
event_type: str,
|
event_type: str,
|
||||||
data: Dict[str, Any],
|
data: dict[str, Any],
|
||||||
) -> None:
|
) -> None:
|
||||||
"""
|
"""
|
||||||
处理解析后的事件数据 - 子类应覆盖此方法
|
处理解析后的事件数据 - 子类应覆盖此方法
|
||||||
@@ -1612,7 +1606,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
def _record_converted_chunks(
|
def _record_converted_chunks(
|
||||||
self,
|
self,
|
||||||
ctx: StreamContext,
|
ctx: StreamContext,
|
||||||
converted_events: List[Dict[str, Any]],
|
converted_events: list[dict[str, Any]],
|
||||||
) -> None:
|
) -> None:
|
||||||
"""
|
"""
|
||||||
记录转换后的 chunk 数据到 parsed_chunks,并更新统计信息
|
记录转换后的 chunk 数据到 parsed_chunks,并更新统计信息
|
||||||
@@ -1656,7 +1650,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
def _extract_usage_from_converted_event(
|
def _extract_usage_from_converted_event(
|
||||||
self,
|
self,
|
||||||
ctx: StreamContext,
|
ctx: StreamContext,
|
||||||
evt: Dict[str, Any],
|
evt: dict[str, Any],
|
||||||
event_type: str,
|
event_type: str,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""
|
"""
|
||||||
@@ -1672,7 +1666,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
evt: 转换后的事件
|
evt: 转换后的事件
|
||||||
event_type: 事件类型
|
event_type: 事件类型
|
||||||
"""
|
"""
|
||||||
usage: Optional[Dict[str, Any]] = None
|
usage: dict[str, Any] | None = None
|
||||||
|
|
||||||
# Claude 格式: message_delta 或 message_start
|
# Claude 格式: message_delta 或 message_start
|
||||||
if event_type == "message_delta":
|
if event_type == "message_delta":
|
||||||
@@ -1737,9 +1731,9 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
async def _create_monitored_stream(
|
async def _create_monitored_stream(
|
||||||
self,
|
self,
|
||||||
ctx: StreamContext,
|
ctx: StreamContext,
|
||||||
stream_generator: AsyncGenerator[bytes, None],
|
stream_generator: AsyncGenerator[bytes],
|
||||||
http_request: Optional[Request] = None,
|
http_request: Request | None = None,
|
||||||
) -> AsyncGenerator[bytes, None]:
|
) -> AsyncGenerator[bytes]:
|
||||||
"""
|
"""
|
||||||
创建带监控的流生成器
|
创建带监控的流生成器
|
||||||
|
|
||||||
@@ -1833,8 +1827,8 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
async def _record_stream_stats(
|
async def _record_stream_stats(
|
||||||
self,
|
self,
|
||||||
ctx: StreamContext,
|
ctx: StreamContext,
|
||||||
original_headers: Dict[str, str],
|
original_headers: dict[str, str],
|
||||||
original_request_body: Dict[str, Any],
|
original_request_body: dict[str, Any],
|
||||||
) -> None:
|
) -> None:
|
||||||
"""在流完成后记录统计信息"""
|
"""在流完成后记录统计信息"""
|
||||||
try:
|
try:
|
||||||
@@ -1996,7 +1990,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
from src.services.request.candidate import RequestCandidateService
|
from src.services.request.candidate import RequestCandidateService
|
||||||
|
|
||||||
# 计算候选自身的 TTFB
|
# 计算候选自身的 TTFB
|
||||||
candidate_first_byte_time_ms: Optional[int] = None
|
candidate_first_byte_time_ms: int | None = None
|
||||||
if ctx.first_byte_time_ms is not None:
|
if ctx.first_byte_time_ms is not None:
|
||||||
candidate_first_byte_time_ms = (
|
candidate_first_byte_time_ms = (
|
||||||
RequestCandidateService.calculate_candidate_ttfb(
|
RequestCandidateService.calculate_candidate_ttfb(
|
||||||
@@ -2061,8 +2055,8 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
self,
|
self,
|
||||||
ctx: StreamContext,
|
ctx: StreamContext,
|
||||||
error: Exception,
|
error: Exception,
|
||||||
original_headers: Dict[str, str],
|
original_headers: dict[str, str],
|
||||||
original_request_body: Dict[str, Any],
|
original_request_body: dict[str, Any],
|
||||||
) -> None:
|
) -> None:
|
||||||
"""记录流式请求失败"""
|
"""记录流式请求失败"""
|
||||||
# 使用 self.start_time 作为时间基准,与首字时间保持一致
|
# 使用 self.start_time 作为时间基准,与首字时间保持一致
|
||||||
@@ -2111,10 +2105,10 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
|
|
||||||
async def process_sync(
|
async def process_sync(
|
||||||
self,
|
self,
|
||||||
original_request_body: Dict[str, Any],
|
original_request_body: dict[str, Any],
|
||||||
original_headers: Dict[str, str],
|
original_headers: dict[str, str],
|
||||||
query_params: Optional[Dict[str, str]] = None,
|
query_params: dict[str, str] | None = None,
|
||||||
path_params: Optional[Dict[str, Any]] = None,
|
path_params: dict[str, Any] | None = None,
|
||||||
) -> JSONResponse:
|
) -> JSONResponse:
|
||||||
"""
|
"""
|
||||||
处理非流式请求
|
处理非流式请求
|
||||||
@@ -2142,19 +2136,19 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
endpoint_id = None # Endpoint ID(用于失败记录)
|
endpoint_id = None # Endpoint ID(用于失败记录)
|
||||||
key_id = None # Key ID(用于失败记录)
|
key_id = None # Key ID(用于失败记录)
|
||||||
mapped_model_result = None # 映射后的目标模型名(用于 Usage 记录)
|
mapped_model_result = None # 映射后的目标模型名(用于 Usage 记录)
|
||||||
response_metadata_result: Dict[str, Any] = {} # Provider 响应元数据
|
response_metadata_result: dict[str, Any] = {} # Provider 响应元数据
|
||||||
needs_conversion = False # 是否需要格式转换(由 candidate 决定)
|
needs_conversion = False # 是否需要格式转换(由 candidate 决定)
|
||||||
|
|
||||||
# 可变请求体容器:允许 orchestrator 在遇到 Thinking 签名错误时整流请求体后重试
|
# 可变请求体容器:允许 orchestrator 在遇到 Thinking 签名错误时整流请求体后重试
|
||||||
# 结构: {"body": 实际请求体, "_rectified": 是否已整流, "_rectified_this_turn": 本轮是否整流}
|
# 结构: {"body": 实际请求体, "_rectified": 是否已整流, "_rectified_this_turn": 本轮是否整流}
|
||||||
request_body_ref: Dict[str, Any] = {"body": original_request_body}
|
request_body_ref: dict[str, Any] = {"body": original_request_body}
|
||||||
|
|
||||||
async def sync_request_func(
|
async def sync_request_func(
|
||||||
provider: Provider,
|
provider: Provider,
|
||||||
endpoint: ProviderEndpoint,
|
endpoint: ProviderEndpoint,
|
||||||
key: ProviderAPIKey,
|
key: ProviderAPIKey,
|
||||||
candidate: ProviderCandidate,
|
candidate: ProviderCandidate,
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
nonlocal provider_name, response_json, status_code, response_headers, provider_api_format, provider_request_headers, provider_request_body, mapped_model_result, response_metadata_result, needs_conversion
|
nonlocal provider_name, response_json, status_code, response_headers, provider_api_format, provider_request_headers, provider_request_body, mapped_model_result, response_metadata_result, needs_conversion
|
||||||
provider_name = str(provider.name)
|
provider_name = str(provider.name)
|
||||||
provider_api_format = str(endpoint.api_format) if endpoint.api_format else ""
|
provider_api_format = str(endpoint.api_format) if endpoint.api_format else ""
|
||||||
@@ -2470,7 +2464,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
actual_request_body = provider_request_body or original_request_body
|
actual_request_body = provider_request_body or original_request_body
|
||||||
|
|
||||||
# 尝试从异常中提取响应头
|
# 尝试从异常中提取响应头
|
||||||
error_response_headers: Dict[str, str] = {}
|
error_response_headers: dict[str, str] = {}
|
||||||
if isinstance(e, ProviderRateLimitException) and e.response_headers:
|
if isinstance(e, ProviderRateLimitException) and e.response_headers:
|
||||||
error_response_headers = e.response_headers
|
error_response_headers = e.response_headers
|
||||||
elif isinstance(e, httpx.HTTPStatusError) and hasattr(e, "response"):
|
elif isinstance(e, httpx.HTTPStatusError) and hasattr(e, "response"):
|
||||||
@@ -2581,7 +2575,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
)
|
)
|
||||||
return True
|
return True
|
||||||
|
|
||||||
def _mark_first_output(self, ctx: StreamContext, state: Dict[str, bool]) -> None:
|
def _mark_first_output(self, ctx: StreamContext, state: dict[str, bool]) -> None:
|
||||||
"""
|
"""
|
||||||
标记首次输出:记录 TTFB 并更新 streaming 状态
|
标记首次输出:记录 TTFB 并更新 streaming 状态
|
||||||
|
|
||||||
@@ -2605,7 +2599,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
ctx: StreamContext,
|
ctx: StreamContext,
|
||||||
line: str,
|
line: str,
|
||||||
events: list, # noqa: ARG002 - 预留给上下文感知转换
|
events: list, # noqa: ARG002 - 预留给上下文感知转换
|
||||||
) -> Tuple[List[str], List[Dict[str, Any]]]:
|
) -> tuple[list[str], list[dict[str, Any]]]:
|
||||||
"""
|
"""
|
||||||
将 SSE 行从 Provider 格式转换为客户端格式
|
将 SSE 行从 Provider 格式转换为客户端格式
|
||||||
|
|
||||||
@@ -2690,7 +2684,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
|
|
||||||
def _parse_sse_line_to_json(
|
def _parse_sse_line_to_json(
|
||||||
self, line: str, provider_format: str
|
self, line: str, provider_format: str
|
||||||
) -> Tuple[Optional[Any], str]:
|
) -> tuple[Any | None, str]:
|
||||||
"""
|
"""
|
||||||
解析 SSE 行为 JSON 对象
|
解析 SSE 行为 JSON 对象
|
||||||
|
|
||||||
|
|||||||
@@ -8,7 +8,6 @@ StreamSmoother 使用这些提取器来处理不同格式的 SSE 事件。
|
|||||||
import copy
|
import copy
|
||||||
import json
|
import json
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from typing import Optional
|
|
||||||
|
|
||||||
|
|
||||||
class ContentExtractor(ABC):
|
class ContentExtractor(ABC):
|
||||||
@@ -20,7 +19,7 @@ class ContentExtractor(ABC):
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def extract_content(self, data: dict) -> Optional[str]:
|
def extract_content(self, data: dict) -> str | None:
|
||||||
"""
|
"""
|
||||||
从 SSE 数据中提取可拆分的文本内容
|
从 SSE 数据中提取可拆分的文本内容
|
||||||
|
|
||||||
@@ -64,7 +63,7 @@ class OpenAIContentExtractor(ContentExtractor):
|
|||||||
- 只在 delta 仅包含 role/content 时允许拆分,避免破坏 tool_calls 等结构
|
- 只在 delta 仅包含 role/content 时允许拆分,避免破坏 tool_calls 等结构
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def extract_content(self, data: dict) -> Optional[str]:
|
def extract_content(self, data: dict) -> str | None:
|
||||||
if not isinstance(data, dict):
|
if not isinstance(data, dict):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@@ -115,7 +114,7 @@ class OpenAIContentExtractor(ContentExtractor):
|
|||||||
new_choices.append(new_choice)
|
new_choices.append(new_choice)
|
||||||
new_data["choices"] = new_choices
|
new_data["choices"] = new_choices
|
||||||
|
|
||||||
return f"data: {json.dumps(new_data, ensure_ascii=False)}\n\n".encode("utf-8")
|
return f"data: {json.dumps(new_data, ensure_ascii=False)}\n\n".encode()
|
||||||
|
|
||||||
|
|
||||||
class ClaudeContentExtractor(ContentExtractor):
|
class ClaudeContentExtractor(ContentExtractor):
|
||||||
@@ -127,7 +126,7 @@ class ClaudeContentExtractor(ContentExtractor):
|
|||||||
- 数据结构: delta.type=text_delta, delta.text
|
- 数据结构: delta.type=text_delta, delta.text
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def extract_content(self, data: dict) -> Optional[str]:
|
def extract_content(self, data: dict) -> str | None:
|
||||||
if not isinstance(data, dict):
|
if not isinstance(data, dict):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@@ -165,9 +164,7 @@ class ClaudeContentExtractor(ContentExtractor):
|
|||||||
|
|
||||||
# Claude 格式需要 event: 前缀
|
# Claude 格式需要 event: 前缀
|
||||||
event_name = event_type or "content_block_delta"
|
event_name = event_type or "content_block_delta"
|
||||||
return f"event: {event_name}\ndata: {json.dumps(new_data, ensure_ascii=False)}\n\n".encode(
|
return f"event: {event_name}\ndata: {json.dumps(new_data, ensure_ascii=False)}\n\n".encode()
|
||||||
"utf-8"
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
class GeminiContentExtractor(ContentExtractor):
|
class GeminiContentExtractor(ContentExtractor):
|
||||||
@@ -179,7 +176,7 @@ class GeminiContentExtractor(ContentExtractor):
|
|||||||
- 只有纯文本块才拆分
|
- 只有纯文本块才拆分
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def extract_content(self, data: dict) -> Optional[str]:
|
def extract_content(self, data: dict) -> str | None:
|
||||||
if not isinstance(data, dict):
|
if not isinstance(data, dict):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@@ -226,7 +223,7 @@ class GeminiContentExtractor(ContentExtractor):
|
|||||||
if "parts" in content and content["parts"]:
|
if "parts" in content and content["parts"]:
|
||||||
content["parts"][0]["text"] = new_content
|
content["parts"][0]["text"] = new_content
|
||||||
|
|
||||||
return f"data: {json.dumps(new_data, ensure_ascii=False)}\n\n".encode("utf-8")
|
return f"data: {json.dumps(new_data, ensure_ascii=False)}\n\n".encode()
|
||||||
|
|
||||||
|
|
||||||
# 提取器注册表
|
# 提取器注册表
|
||||||
@@ -237,7 +234,7 @@ _EXTRACTORS: dict[str, type[ContentExtractor]] = {
|
|||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
def get_extractor(format_name: str) -> Optional[ContentExtractor]:
|
def get_extractor(format_name: str) -> ContentExtractor | None:
|
||||||
"""
|
"""
|
||||||
根据格式名获取对应的内容提取器实例
|
根据格式名获取对应的内容提取器实例
|
||||||
|
|
||||||
|
|||||||
@@ -14,15 +14,14 @@
|
|||||||
- EndpointCheckOrchestrator: 协调整个流程
|
- EndpointCheckOrchestrator: 协调整个流程
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import Any, AsyncIterator, Dict, Iterable, Optional, Union, List
|
from typing import Any
|
||||||
from abc import ABC, abstractmethod
|
from collections.abc import Iterable
|
||||||
import time
|
import time
|
||||||
import uuid
|
import uuid
|
||||||
import json
|
import json
|
||||||
from functools import lru_cache
|
|
||||||
import asyncio
|
import asyncio
|
||||||
from collections import defaultdict
|
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
|
|
||||||
@@ -31,7 +30,7 @@ from src.core.api_format import CORE_REDACT_HEADERS, merge_headers_with_protecti
|
|||||||
from src.utils.ssl_utils import get_ssl_context
|
from src.utils.ssl_utils import get_ssl_context
|
||||||
|
|
||||||
|
|
||||||
def _redact_headers(headers: Dict[str, str]) -> Dict[str, str]:
|
def _redact_headers(headers: dict[str, str]) -> dict[str, str]:
|
||||||
return redact_headers_for_log(headers, CORE_REDACT_HEADERS)
|
return redact_headers_for_log(headers, CORE_REDACT_HEADERS)
|
||||||
|
|
||||||
|
|
||||||
@@ -46,10 +45,10 @@ def _truncate_repr(value: Any, limit: int = 1200) -> str:
|
|||||||
|
|
||||||
|
|
||||||
def build_safe_headers(
|
def build_safe_headers(
|
||||||
base_headers: Dict[str, str],
|
base_headers: dict[str, str],
|
||||||
extra_headers: Optional[Dict[str, str]],
|
extra_headers: dict[str, str] | None,
|
||||||
protected_keys: Iterable[str],
|
protected_keys: Iterable[str],
|
||||||
) -> Dict[str, str]:
|
) -> dict[str, str]:
|
||||||
"""
|
"""
|
||||||
合并 extra_headers,但防止覆盖 protected_keys(大小写不敏感)。
|
合并 extra_headers,但防止覆盖 protected_keys(大小写不敏感)。
|
||||||
"""
|
"""
|
||||||
@@ -60,16 +59,16 @@ async def run_endpoint_check(
|
|||||||
*,
|
*,
|
||||||
client: httpx.AsyncClient, # 保持兼容性,但内部不使用
|
client: httpx.AsyncClient, # 保持兼容性,但内部不使用
|
||||||
url: str,
|
url: str,
|
||||||
headers: Dict[str, str],
|
headers: dict[str, str],
|
||||||
json_body: Dict[str, Any],
|
json_body: dict[str, Any],
|
||||||
api_format: str,
|
api_format: str,
|
||||||
provider_name: Optional[str] = None,
|
provider_name: str | None = None,
|
||||||
model_name: Optional[str] = None,
|
model_name: str | None = None,
|
||||||
api_key_id: Optional[str] = None,
|
api_key_id: str | None = None,
|
||||||
provider_id: Optional[str] = None,
|
provider_id: str | None = None,
|
||||||
db: Optional[Any] = None, # Session对象,需要时才导入
|
db: Any | None = None, # Session对象,需要时才导入
|
||||||
user: Optional[Any] = None, # User对象
|
user: Any | None = None, # User对象
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
执行端点检查(重构版本,使用新的架构):
|
执行端点检查(重构版本,使用新的架构):
|
||||||
- 使用新的架构类来分离关注点
|
- 使用新的架构类来分离关注点
|
||||||
@@ -123,21 +122,21 @@ async def _calculate_and_record_usage(
|
|||||||
provider_id: str,
|
provider_id: str,
|
||||||
api_key_id: str,
|
api_key_id: str,
|
||||||
model_name: str,
|
model_name: str,
|
||||||
request_data: Dict[str, Any],
|
request_data: dict[str, Any],
|
||||||
response_data: Optional[Dict[str, Any]],
|
response_data: dict[str, Any] | None,
|
||||||
request_id: str,
|
request_id: str,
|
||||||
response_time_ms: int,
|
response_time_ms: int,
|
||||||
request_headers: Dict[str, str],
|
request_headers: dict[str, str],
|
||||||
response_headers: Optional[Dict[str, str]] = None,
|
response_headers: dict[str, str] | None = None,
|
||||||
status_code: int = 0,
|
status_code: int = 0,
|
||||||
error_message: Optional[str] = None,
|
error_message: str | None = None,
|
||||||
# 新增:支持直接传递token数据
|
# 新增:支持直接传递token数据
|
||||||
input_tokens: Optional[int] = None,
|
input_tokens: int | None = None,
|
||||||
output_tokens: Optional[int] = None,
|
output_tokens: int | None = None,
|
||||||
cache_creation_input_tokens: Optional[int] = None,
|
cache_creation_input_tokens: int | None = None,
|
||||||
cache_read_input_tokens: Optional[int] = None,
|
cache_read_input_tokens: int | None = None,
|
||||||
api_format: Optional[str] = None,
|
api_format: str | None = None,
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
计算并记录用量数据(遗留函数)
|
计算并记录用量数据(遗留函数)
|
||||||
|
|
||||||
@@ -149,7 +148,7 @@ async def _calculate_and_record_usage(
|
|||||||
"""
|
"""
|
||||||
from src.services.usage.service import UsageService
|
from src.services.usage.service import UsageService
|
||||||
from src.services.request.candidate import RequestCandidateService
|
from src.services.request.candidate import RequestCandidateService
|
||||||
from src.models.database import ApiKey, ProviderAPIKey, ProviderEndpoint
|
from src.models.database import ApiKey, ProviderAPIKey
|
||||||
|
|
||||||
# 获取Provider API Key对象(不是用户API Key)
|
# 获取Provider API Key对象(不是用户API Key)
|
||||||
provider_api_key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == api_key_id).first()
|
provider_api_key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == api_key_id).first()
|
||||||
@@ -360,7 +359,7 @@ async def _calculate_and_record_usage(
|
|||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
def _extract_tokens_from_response(api_identifier: str, response_data: Optional[Dict[str, Any]]) -> tuple[int, int, int, int]:
|
def _extract_tokens_from_response(api_identifier: str, response_data: dict[str, Any] | None) -> tuple[int, int, int, int]:
|
||||||
"""
|
"""
|
||||||
从响应中提取Token计数信息
|
从响应中提取Token计数信息
|
||||||
|
|
||||||
@@ -446,7 +445,7 @@ def _extract_tokens_from_response(api_identifier: str, response_data: Optional[D
|
|||||||
|
|
||||||
|
|
||||||
|
|
||||||
def _fallback_token_counting(request_data: Dict[str, Any], response_data: Optional[Dict[str, Any]]) -> tuple[int, int, int, int]:
|
def _fallback_token_counting(request_data: dict[str, Any], response_data: dict[str, Any] | None) -> tuple[int, int, int, int]:
|
||||||
"""
|
"""
|
||||||
回退的Token计数方法(简单估算)
|
回退的Token计数方法(简单估算)
|
||||||
|
|
||||||
@@ -508,16 +507,16 @@ def _fallback_token_counting(request_data: Dict[str, Any], response_data: Option
|
|||||||
class EndpointCheckRequest:
|
class EndpointCheckRequest:
|
||||||
"""端点检查请求数据类"""
|
"""端点检查请求数据类"""
|
||||||
url: str
|
url: str
|
||||||
headers: Dict[str, str]
|
headers: dict[str, str]
|
||||||
json_body: Dict[str, Any]
|
json_body: dict[str, Any]
|
||||||
api_format: str
|
api_format: str
|
||||||
provider_name: Optional[str] = None
|
provider_name: str | None = None
|
||||||
model_name: Optional[str] = None
|
model_name: str | None = None
|
||||||
api_key_id: Optional[str] = None
|
api_key_id: str | None = None
|
||||||
provider_id: Optional[str] = None
|
provider_id: str | None = None
|
||||||
db: Optional[Any] = None
|
db: Any | None = None
|
||||||
user: Optional[Any] = None
|
user: Any | None = None
|
||||||
request_id: Optional[str] = None
|
request_id: str | None = None
|
||||||
timeout: float = 30.0
|
timeout: float = 30.0
|
||||||
|
|
||||||
|
|
||||||
@@ -525,12 +524,12 @@ class EndpointCheckRequest:
|
|||||||
class EndpointCheckResult:
|
class EndpointCheckResult:
|
||||||
"""端点检查结果数据类"""
|
"""端点检查结果数据类"""
|
||||||
status_code: int
|
status_code: int
|
||||||
headers: Dict[str, str]
|
headers: dict[str, str]
|
||||||
response_time_ms: int
|
response_time_ms: int
|
||||||
request_id: str
|
request_id: str
|
||||||
response_data: Optional[Dict[str, Any]] = None
|
response_data: dict[str, Any] | None = None
|
||||||
error_message: Optional[str] = None
|
error_message: str | None = None
|
||||||
usage_data: Optional[Dict[str, Any]] = None
|
usage_data: dict[str, Any] | None = None
|
||||||
|
|
||||||
|
|
||||||
class HttpRequestExecutor:
|
class HttpRequestExecutor:
|
||||||
@@ -613,7 +612,7 @@ class UsageCalculator:
|
|||||||
return _extract_tokens_from_response(api_identifier, result.response_data)
|
return _extract_tokens_from_response(api_identifier, result.response_data)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _fallback_token_counting(request_data: Dict[str, Any], response_data: Optional[Dict[str, Any]]) -> tuple[int, int, int, int]:
|
def _fallback_token_counting(request_data: dict[str, Any], response_data: dict[str, Any] | None) -> tuple[int, int, int, int]:
|
||||||
"""回退的Token计数方法(简单估算)"""
|
"""回退的Token计数方法(简单估算)"""
|
||||||
# 估算输入Token
|
# 估算输入Token
|
||||||
messages = request_data.get("messages", request_data.get("contents", []))
|
messages = request_data.get("messages", request_data.get("contents", []))
|
||||||
@@ -665,12 +664,12 @@ class AsyncBatchUsageRecorder:
|
|||||||
def __init__(self, batch_size: int = 10, flush_interval: float = 2.0):
|
def __init__(self, batch_size: int = 10, flush_interval: float = 2.0):
|
||||||
self.batch_size = batch_size
|
self.batch_size = batch_size
|
||||||
self.flush_interval = flush_interval
|
self.flush_interval = flush_interval
|
||||||
self.pending_records: List[Dict[str, Any]] = []
|
self.pending_records: list[dict[str, Any]] = []
|
||||||
self._flush_task: Optional[asyncio.Task] = None
|
self._flush_task: asyncio.Task | None = None
|
||||||
self._lock = asyncio.Lock()
|
self._lock = asyncio.Lock()
|
||||||
self._running = True
|
self._running = True
|
||||||
|
|
||||||
async def add_record(self, usage_data: Dict[str, Any]) -> None:
|
async def add_record(self, usage_data: dict[str, Any]) -> None:
|
||||||
"""添加用量记录到批处理队列"""
|
"""添加用量记录到批处理队列"""
|
||||||
async with self._lock:
|
async with self._lock:
|
||||||
self.pending_records.append(usage_data)
|
self.pending_records.append(usage_data)
|
||||||
@@ -740,7 +739,7 @@ class AsyncBatchUsageRecorder:
|
|||||||
|
|
||||||
|
|
||||||
# 全局批处理器实例(单例)
|
# 全局批处理器实例(单例)
|
||||||
_global_batch_recorder: Optional[AsyncBatchUsageRecorder] = None
|
_global_batch_recorder: AsyncBatchUsageRecorder | None = None
|
||||||
|
|
||||||
def get_batch_recorder() -> AsyncBatchUsageRecorder:
|
def get_batch_recorder() -> AsyncBatchUsageRecorder:
|
||||||
"""获取全局批处理器实例"""
|
"""获取全局批处理器实例"""
|
||||||
@@ -756,7 +755,7 @@ def get_batch_recorder() -> AsyncBatchUsageRecorder:
|
|||||||
|
|
||||||
class EndpointCheckError(Exception):
|
class EndpointCheckError(Exception):
|
||||||
"""端点检查错误基类"""
|
"""端点检查错误基类"""
|
||||||
def __init__(self, message: str, error_type: str, status_code: int = 500, details: Optional[Dict[str, Any]] = None):
|
def __init__(self, message: str, error_type: str, status_code: int = 500, details: dict[str, Any] | None = None):
|
||||||
super().__init__(message)
|
super().__init__(message)
|
||||||
self.message = message
|
self.message = message
|
||||||
self.error_type = error_type
|
self.error_type = error_type
|
||||||
@@ -765,22 +764,22 @@ class EndpointCheckError(Exception):
|
|||||||
|
|
||||||
class NetworkError(EndpointCheckError):
|
class NetworkError(EndpointCheckError):
|
||||||
"""网络请求错误"""
|
"""网络请求错误"""
|
||||||
def __init__(self, message: str, details: Optional[Dict[str, Any]] = None):
|
def __init__(self, message: str, details: dict[str, Any] | None = None):
|
||||||
super().__init__(message, "network_error", 0, details)
|
super().__init__(message, "network_error", 0, details)
|
||||||
|
|
||||||
class AuthenticationError(EndpointCheckError):
|
class AuthenticationError(EndpointCheckError):
|
||||||
"""认证错误"""
|
"""认证错误"""
|
||||||
def __init__(self, message: str, details: Optional[Dict[str, Any]] = None):
|
def __init__(self, message: str, details: dict[str, Any] | None = None):
|
||||||
super().__init__(message, "authentication_error", 401, details)
|
super().__init__(message, "authentication_error", 401, details)
|
||||||
|
|
||||||
class RateLimitError(EndpointCheckError):
|
class RateLimitError(EndpointCheckError):
|
||||||
"""速率限制错误"""
|
"""速率限制错误"""
|
||||||
def __init__(self, message: str, details: Optional[Dict[str, Any]] = None):
|
def __init__(self, message: str, details: dict[str, Any] | None = None):
|
||||||
super().__init__(message, "rate_limit_error", 429, details)
|
super().__init__(message, "rate_limit_error", 429, details)
|
||||||
|
|
||||||
class UpstreamError(EndpointCheckError):
|
class UpstreamError(EndpointCheckError):
|
||||||
"""上游服务错误"""
|
"""上游服务错误"""
|
||||||
def __init__(self, message: str, status_code: int, details: Optional[Dict[str, Any]] = None):
|
def __init__(self, message: str, status_code: int, details: dict[str, Any] | None = None):
|
||||||
super().__init__(message, "upstream_error", status_code, details)
|
super().__init__(message, "upstream_error", status_code, details)
|
||||||
|
|
||||||
|
|
||||||
@@ -982,7 +981,7 @@ class EndpointCheckConfig:
|
|||||||
retry_on_timeouts: bool = True
|
retry_on_timeouts: bool = True
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_env(cls) -> 'EndpointCheckConfig':
|
def from_env(cls) -> EndpointCheckConfig:
|
||||||
"""从环境变量创建配置"""
|
"""从环境变量创建配置"""
|
||||||
import os
|
import os
|
||||||
|
|
||||||
@@ -1004,7 +1003,7 @@ class EndpointCheckConfig:
|
|||||||
)
|
)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_dict(cls, config_dict: Dict[str, Any]) -> 'EndpointCheckConfig':
|
def from_dict(cls, config_dict: dict[str, Any]) -> EndpointCheckConfig:
|
||||||
"""从字典创建配置"""
|
"""从字典创建配置"""
|
||||||
return cls(**{k: v for k, v in config_dict.items() if hasattr(cls, k)})
|
return cls(**{k: v for k, v in config_dict.items() if hasattr(cls, k)})
|
||||||
|
|
||||||
@@ -1012,7 +1011,7 @@ class EndpointCheckConfig:
|
|||||||
class ConfigurableEndpointChecker:
|
class ConfigurableEndpointChecker:
|
||||||
"""可配置的端点检查器"""
|
"""可配置的端点检查器"""
|
||||||
|
|
||||||
def __init__(self, config: Optional[EndpointCheckConfig] = None):
|
def __init__(self, config: EndpointCheckConfig | None = None):
|
||||||
self.config = config or EndpointCheckConfig()
|
self.config = config or EndpointCheckConfig()
|
||||||
self.executor = HttpRequestExecutor(timeout=self.config.timeout)
|
self.executor = HttpRequestExecutor(timeout=self.config.timeout)
|
||||||
self.usage_calculator = UsageCalculator()
|
self.usage_calculator = UsageCalculator()
|
||||||
@@ -1171,9 +1170,9 @@ class ConfigurableEndpointChecker:
|
|||||||
|
|
||||||
|
|
||||||
# 全局配置检查器实例
|
# 全局配置检查器实例
|
||||||
_global_configured_checker: Optional[ConfigurableEndpointChecker] = None
|
_global_configured_checker: ConfigurableEndpointChecker | None = None
|
||||||
|
|
||||||
def get_configured_checker(config: Optional[EndpointCheckConfig] = None) -> ConfigurableEndpointChecker:
|
def get_configured_checker(config: EndpointCheckConfig | None = None) -> ConfigurableEndpointChecker:
|
||||||
"""获取全局配置检查器实例"""
|
"""获取全局配置检查器实例"""
|
||||||
global _global_configured_checker
|
global _global_configured_checker
|
||||||
if _global_configured_checker is None or config is not None:
|
if _global_configured_checker is None or config is not None:
|
||||||
@@ -1186,8 +1185,8 @@ def get_configured_checker(config: Optional[EndpointCheckConfig] = None) -> Conf
|
|||||||
class EndpointCheckOrchestrator:
|
class EndpointCheckOrchestrator:
|
||||||
"""端点检查协调器 - 协调整个流程"""
|
"""端点检查协调器 - 协调整个流程"""
|
||||||
|
|
||||||
def __init__(self, executor: Optional[HttpRequestExecutor] = None,
|
def __init__(self, executor: HttpRequestExecutor | None = None,
|
||||||
usage_calculator: Optional[UsageCalculator] = None):
|
usage_calculator: UsageCalculator | None = None):
|
||||||
self.executor = executor or HttpRequestExecutor()
|
self.executor = executor or HttpRequestExecutor()
|
||||||
self.usage_calculator = usage_calculator or UsageCalculator()
|
self.usage_calculator = usage_calculator or UsageCalculator()
|
||||||
|
|
||||||
|
|||||||
@@ -6,7 +6,7 @@
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import re
|
import re
|
||||||
from typing import Any, Dict, Optional, Tuple, Type
|
from typing import Any
|
||||||
|
|
||||||
from src.api.handlers.base.response_parser import (
|
from src.api.handlers.base.response_parser import (
|
||||||
ParsedChunk,
|
ParsedChunk,
|
||||||
@@ -20,7 +20,7 @@ from src.api.handlers.base.utils import extract_cache_creation_tokens
|
|||||||
from src.core.api_format import is_cli_format
|
from src.core.api_format import is_cli_format
|
||||||
|
|
||||||
|
|
||||||
def _check_nested_error(response: Dict[str, Any]) -> Tuple[bool, Optional[Dict[str, Any]]]:
|
def _check_nested_error(response: dict[str, Any]) -> tuple[bool, dict[str, Any] | None]:
|
||||||
"""
|
"""
|
||||||
检查响应中是否存在嵌套错误(某些代理服务返回 HTTP 200 但在响应体中包含错误)
|
检查响应中是否存在嵌套错误(某些代理服务返回 HTTP 200 但在响应体中包含错误)
|
||||||
|
|
||||||
@@ -62,7 +62,7 @@ def _check_nested_error(response: Dict[str, Any]) -> Tuple[bool, Optional[Dict[s
|
|||||||
return False, None
|
return False, None
|
||||||
|
|
||||||
|
|
||||||
def _extract_embedded_status_code(error_info: Optional[Dict[str, Any]]) -> Optional[int]:
|
def _extract_embedded_status_code(error_info: dict[str, Any] | None) -> int | None:
|
||||||
"""
|
"""
|
||||||
从错误信息中提取嵌套的状态码
|
从错误信息中提取嵌套的状态码
|
||||||
|
|
||||||
@@ -137,7 +137,7 @@ class OpenAIResponseParser(ResponseParser):
|
|||||||
self.name = "OPENAI"
|
self.name = "OPENAI"
|
||||||
self.api_format = "OPENAI"
|
self.api_format = "OPENAI"
|
||||||
|
|
||||||
def parse_sse_line(self, line: str, stats: StreamStats) -> Optional[ParsedChunk]:
|
def parse_sse_line(self, line: str, stats: StreamStats) -> ParsedChunk | None:
|
||||||
if not line or not line.strip():
|
if not line or not line.strip():
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@@ -186,7 +186,7 @@ class OpenAIResponseParser(ResponseParser):
|
|||||||
|
|
||||||
return chunk
|
return chunk
|
||||||
|
|
||||||
def parse_response(self, response: Dict[str, Any], status_code: int) -> ParsedResponse:
|
def parse_response(self, response: dict[str, Any], status_code: int) -> ParsedResponse:
|
||||||
result = ParsedResponse(
|
result = ParsedResponse(
|
||||||
raw_response=response,
|
raw_response=response,
|
||||||
status_code=status_code,
|
status_code=status_code,
|
||||||
@@ -217,7 +217,7 @@ class OpenAIResponseParser(ResponseParser):
|
|||||||
|
|
||||||
return result
|
return result
|
||||||
|
|
||||||
def extract_usage_from_response(self, response: Dict[str, Any]) -> Dict[str, int]:
|
def extract_usage_from_response(self, response: dict[str, Any]) -> dict[str, int]:
|
||||||
usage = response.get("usage") or {}
|
usage = response.get("usage") or {}
|
||||||
return {
|
return {
|
||||||
"input_tokens": usage.get("prompt_tokens", 0),
|
"input_tokens": usage.get("prompt_tokens", 0),
|
||||||
@@ -226,7 +226,7 @@ class OpenAIResponseParser(ResponseParser):
|
|||||||
"cache_read_tokens": 0,
|
"cache_read_tokens": 0,
|
||||||
}
|
}
|
||||||
|
|
||||||
def extract_text_content(self, response: Dict[str, Any]) -> str:
|
def extract_text_content(self, response: dict[str, Any]) -> str:
|
||||||
choices = response.get("choices", [])
|
choices = response.get("choices", [])
|
||||||
if choices:
|
if choices:
|
||||||
message = choices[0].get("message", {})
|
message = choices[0].get("message", {})
|
||||||
@@ -235,7 +235,7 @@ class OpenAIResponseParser(ResponseParser):
|
|||||||
return content
|
return content
|
||||||
return ""
|
return ""
|
||||||
|
|
||||||
def is_error_response(self, response: Dict[str, Any]) -> bool:
|
def is_error_response(self, response: dict[str, Any]) -> bool:
|
||||||
is_error, _ = _check_nested_error(response)
|
is_error, _ = _check_nested_error(response)
|
||||||
return is_error
|
return is_error
|
||||||
|
|
||||||
@@ -259,7 +259,7 @@ class ClaudeResponseParser(ResponseParser):
|
|||||||
self.name = "CLAUDE"
|
self.name = "CLAUDE"
|
||||||
self.api_format = "CLAUDE"
|
self.api_format = "CLAUDE"
|
||||||
|
|
||||||
def parse_sse_line(self, line: str, stats: StreamStats) -> Optional[ParsedChunk]:
|
def parse_sse_line(self, line: str, stats: StreamStats) -> ParsedChunk | None:
|
||||||
if not line or not line.strip():
|
if not line or not line.strip():
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@@ -324,7 +324,7 @@ class ClaudeResponseParser(ResponseParser):
|
|||||||
|
|
||||||
return chunk
|
return chunk
|
||||||
|
|
||||||
def parse_response(self, response: Dict[str, Any], status_code: int) -> ParsedResponse:
|
def parse_response(self, response: dict[str, Any], status_code: int) -> ParsedResponse:
|
||||||
result = ParsedResponse(
|
result = ParsedResponse(
|
||||||
raw_response=response,
|
raw_response=response,
|
||||||
status_code=status_code,
|
status_code=status_code,
|
||||||
@@ -358,7 +358,7 @@ class ClaudeResponseParser(ResponseParser):
|
|||||||
|
|
||||||
return result
|
return result
|
||||||
|
|
||||||
def extract_usage_from_response(self, response: Dict[str, Any]) -> Dict[str, int]:
|
def extract_usage_from_response(self, response: dict[str, Any]) -> dict[str, int]:
|
||||||
# 对于 message_start 事件,usage 在 message.usage 路径下
|
# 对于 message_start 事件,usage 在 message.usage 路径下
|
||||||
# 对于其他响应,usage 在顶层
|
# 对于其他响应,usage 在顶层
|
||||||
usage = response.get("usage") or {}
|
usage = response.get("usage") or {}
|
||||||
@@ -372,7 +372,7 @@ class ClaudeResponseParser(ResponseParser):
|
|||||||
"cache_read_tokens": usage.get("cache_read_input_tokens", 0),
|
"cache_read_tokens": usage.get("cache_read_input_tokens", 0),
|
||||||
}
|
}
|
||||||
|
|
||||||
def extract_text_content(self, response: Dict[str, Any]) -> str:
|
def extract_text_content(self, response: dict[str, Any]) -> str:
|
||||||
content = response.get("content", [])
|
content = response.get("content", [])
|
||||||
if isinstance(content, list):
|
if isinstance(content, list):
|
||||||
text_parts = []
|
text_parts = []
|
||||||
@@ -382,7 +382,7 @@ class ClaudeResponseParser(ResponseParser):
|
|||||||
return "".join(text_parts)
|
return "".join(text_parts)
|
||||||
return ""
|
return ""
|
||||||
|
|
||||||
def is_error_response(self, response: Dict[str, Any]) -> bool:
|
def is_error_response(self, response: dict[str, Any]) -> bool:
|
||||||
is_error, _ = _check_nested_error(response)
|
is_error, _ = _check_nested_error(response)
|
||||||
return is_error
|
return is_error
|
||||||
|
|
||||||
@@ -406,7 +406,7 @@ class GeminiResponseParser(ResponseParser):
|
|||||||
self.name = "GEMINI"
|
self.name = "GEMINI"
|
||||||
self.api_format = "GEMINI"
|
self.api_format = "GEMINI"
|
||||||
|
|
||||||
def parse_sse_line(self, line: str, stats: StreamStats) -> Optional[ParsedChunk]:
|
def parse_sse_line(self, line: str, stats: StreamStats) -> ParsedChunk | None:
|
||||||
"""
|
"""
|
||||||
解析 Gemini SSE 行
|
解析 Gemini SSE 行
|
||||||
|
|
||||||
@@ -473,7 +473,7 @@ class GeminiResponseParser(ResponseParser):
|
|||||||
|
|
||||||
return chunk
|
return chunk
|
||||||
|
|
||||||
def parse_response(self, response: Dict[str, Any], status_code: int) -> ParsedResponse:
|
def parse_response(self, response: dict[str, Any], status_code: int) -> ParsedResponse:
|
||||||
result = ParsedResponse(
|
result = ParsedResponse(
|
||||||
raw_response=response,
|
raw_response=response,
|
||||||
status_code=status_code,
|
status_code=status_code,
|
||||||
@@ -509,7 +509,7 @@ class GeminiResponseParser(ResponseParser):
|
|||||||
|
|
||||||
return result
|
return result
|
||||||
|
|
||||||
def extract_usage_from_response(self, response: Dict[str, Any]) -> Dict[str, int]:
|
def extract_usage_from_response(self, response: dict[str, Any]) -> dict[str, int]:
|
||||||
"""
|
"""
|
||||||
从 Gemini 响应中提取 token 使用量
|
从 Gemini 响应中提取 token 使用量
|
||||||
|
|
||||||
@@ -531,7 +531,7 @@ class GeminiResponseParser(ResponseParser):
|
|||||||
"cache_read_tokens": usage.get("cached_tokens", 0),
|
"cache_read_tokens": usage.get("cached_tokens", 0),
|
||||||
}
|
}
|
||||||
|
|
||||||
def extract_text_content(self, response: Dict[str, Any]) -> str:
|
def extract_text_content(self, response: dict[str, Any]) -> str:
|
||||||
candidates = response.get("candidates", [])
|
candidates = response.get("candidates", [])
|
||||||
if candidates:
|
if candidates:
|
||||||
content = candidates[0].get("content", {})
|
content = candidates[0].get("content", {})
|
||||||
@@ -543,7 +543,7 @@ class GeminiResponseParser(ResponseParser):
|
|||||||
return "".join(text_parts)
|
return "".join(text_parts)
|
||||||
return ""
|
return ""
|
||||||
|
|
||||||
def is_error_response(self, response: Dict[str, Any]) -> bool:
|
def is_error_response(self, response: dict[str, Any]) -> bool:
|
||||||
"""
|
"""
|
||||||
判断响应是否为错误响应
|
判断响应是否为错误响应
|
||||||
|
|
||||||
@@ -562,7 +562,7 @@ class GeminiCliResponseParser(GeminiResponseParser):
|
|||||||
|
|
||||||
|
|
||||||
# 解析器注册表
|
# 解析器注册表
|
||||||
_PARSERS: Dict[str, Type[ResponseParser]] = {
|
_PARSERS: dict[str, type[ResponseParser]] = {
|
||||||
"CLAUDE": ClaudeResponseParser,
|
"CLAUDE": ClaudeResponseParser,
|
||||||
"CLAUDE_CLI": ClaudeCliResponseParser,
|
"CLAUDE_CLI": ClaudeCliResponseParser,
|
||||||
"OPENAI": OpenAIResponseParser,
|
"OPENAI": OpenAIResponseParser,
|
||||||
|
|||||||
@@ -11,12 +11,11 @@
|
|||||||
payload, headers = builder.build(original_body, original_headers, endpoint, key)
|
payload, headers = builder.build(original_body, original_headers, endpoint, key)
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import json
|
import json
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import TYPE_CHECKING, Any, Dict, FrozenSet, Optional, Tuple
|
from typing import TYPE_CHECKING, Any
|
||||||
|
|
||||||
from src.core.crypto import crypto_service
|
from src.core.crypto import crypto_service
|
||||||
from src.core.api_format import HeaderBuilder, UPSTREAM_DROP_HEADERS
|
from src.core.api_format import HeaderBuilder, UPSTREAM_DROP_HEADERS
|
||||||
@@ -37,9 +36,9 @@ class ProviderAuthInfo:
|
|||||||
auth_header: str
|
auth_header: str
|
||||||
auth_value: str
|
auth_value: str
|
||||||
# 解密后的认证配置(用于 URL 构建等场景,避免重复解密)
|
# 解密后的认证配置(用于 URL 构建等场景,避免重复解密)
|
||||||
decrypted_auth_config: Optional[Dict[str, Any]] = None
|
decrypted_auth_config: dict[str, Any] | None = None
|
||||||
|
|
||||||
def as_tuple(self) -> Tuple[str, str]:
|
def as_tuple(self) -> tuple[str, str]:
|
||||||
"""返回 (auth_header, auth_value) 元组"""
|
"""返回 (auth_header, auth_value) 元组"""
|
||||||
return (self.auth_header, self.auth_value)
|
return (self.auth_header, self.auth_value)
|
||||||
|
|
||||||
@@ -48,7 +47,7 @@ class ProviderAuthInfo:
|
|||||||
# ==============================================================================
|
# ==============================================================================
|
||||||
|
|
||||||
# 兼容别名:历史代码使用 SENSITIVE_HEADERS 命名
|
# 兼容别名:历史代码使用 SENSITIVE_HEADERS 命名
|
||||||
SENSITIVE_HEADERS: FrozenSet[str] = UPSTREAM_DROP_HEADERS
|
SENSITIVE_HEADERS: frozenset[str] = UPSTREAM_DROP_HEADERS
|
||||||
|
|
||||||
|
|
||||||
# ==============================================================================
|
# ==============================================================================
|
||||||
@@ -57,14 +56,14 @@ SENSITIVE_HEADERS: FrozenSet[str] = UPSTREAM_DROP_HEADERS
|
|||||||
|
|
||||||
# 标准测试请求体(OpenAI 格式)
|
# 标准测试请求体(OpenAI 格式)
|
||||||
# 用于 check_endpoint 等测试场景,使用简单安全的消息内容避免触发安全过滤
|
# 用于 check_endpoint 等测试场景,使用简单安全的消息内容避免触发安全过滤
|
||||||
DEFAULT_TEST_REQUEST: Dict[str, Any] = {
|
DEFAULT_TEST_REQUEST: dict[str, Any] = {
|
||||||
"messages": [{"role": "user", "content": "Hi"}],
|
"messages": [{"role": "user", "content": "Hi"}],
|
||||||
"max_tokens": 5,
|
"max_tokens": 5,
|
||||||
"temperature": 0,
|
"temperature": 0,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
def get_test_request_data(request_data: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
|
def get_test_request_data(request_data: dict[str, Any] | None = None) -> dict[str, Any]:
|
||||||
"""获取测试请求数据
|
"""获取测试请求数据
|
||||||
|
|
||||||
如果传入 request_data,则合并到默认测试请求中;
|
如果传入 request_data,则合并到默认测试请求中;
|
||||||
@@ -85,8 +84,8 @@ def get_test_request_data(request_data: Optional[Dict[str, Any]] = None) -> Dict
|
|||||||
|
|
||||||
def build_test_request_body(
|
def build_test_request_body(
|
||||||
format_id: str,
|
format_id: str,
|
||||||
request_data: Optional[Dict[str, Any]] = None,
|
request_data: dict[str, Any] | None = None,
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""构建测试请求体,自动处理格式转换
|
"""构建测试请求体,自动处理格式转换
|
||||||
|
|
||||||
使用格式转换注册表将 OpenAI 格式的测试请求转换为目标格式。
|
使用格式转换注册表将 OpenAI 格式的测试请求转换为目标格式。
|
||||||
@@ -127,39 +126,39 @@ class RequestBuilder(ABC):
|
|||||||
@abstractmethod
|
@abstractmethod
|
||||||
def build_payload(
|
def build_payload(
|
||||||
self,
|
self,
|
||||||
original_body: Dict[str, Any],
|
original_body: dict[str, Any],
|
||||||
*,
|
*,
|
||||||
mapped_model: Optional[str] = None,
|
mapped_model: str | None = None,
|
||||||
is_stream: bool = False,
|
is_stream: bool = False,
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""构建请求体"""
|
"""构建请求体"""
|
||||||
pass
|
pass
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def build_headers(
|
def build_headers(
|
||||||
self,
|
self,
|
||||||
original_headers: Dict[str, str],
|
original_headers: dict[str, str],
|
||||||
endpoint: Any,
|
endpoint: Any,
|
||||||
key: Any,
|
key: Any,
|
||||||
*,
|
*,
|
||||||
extra_headers: Optional[Dict[str, str]] = None,
|
extra_headers: dict[str, str] | None = None,
|
||||||
pre_computed_auth: Optional[Tuple[str, str]] = None,
|
pre_computed_auth: tuple[str, str] | None = None,
|
||||||
) -> Dict[str, str]:
|
) -> dict[str, str]:
|
||||||
"""构建请求头"""
|
"""构建请求头"""
|
||||||
pass
|
pass
|
||||||
|
|
||||||
def build(
|
def build(
|
||||||
self,
|
self,
|
||||||
original_body: Dict[str, Any],
|
original_body: dict[str, Any],
|
||||||
original_headers: Dict[str, str],
|
original_headers: dict[str, str],
|
||||||
endpoint: Any,
|
endpoint: Any,
|
||||||
key: Any,
|
key: Any,
|
||||||
*,
|
*,
|
||||||
mapped_model: Optional[str] = None,
|
mapped_model: str | None = None,
|
||||||
is_stream: bool = False,
|
is_stream: bool = False,
|
||||||
extra_headers: Optional[Dict[str, str]] = None,
|
extra_headers: dict[str, str] | None = None,
|
||||||
pre_computed_auth: Optional[Tuple[str, str]] = None,
|
pre_computed_auth: tuple[str, str] | None = None,
|
||||||
) -> Tuple[Dict[str, Any], Dict[str, str]]:
|
) -> tuple[dict[str, Any], dict[str, str]]:
|
||||||
"""
|
"""
|
||||||
构建完整的请求(请求体 + 请求头)
|
构建完整的请求(请求体 + 请求头)
|
||||||
|
|
||||||
@@ -202,11 +201,11 @@ class PassthroughRequestBuilder(RequestBuilder):
|
|||||||
|
|
||||||
def build_payload(
|
def build_payload(
|
||||||
self,
|
self,
|
||||||
original_body: Dict[str, Any],
|
original_body: dict[str, Any],
|
||||||
*,
|
*,
|
||||||
mapped_model: Optional[str] = None, # noqa: ARG002 - 由 apply_mapped_model 处理
|
mapped_model: str | None = None, # noqa: ARG002 - 由 apply_mapped_model 处理
|
||||||
is_stream: bool = False, # noqa: ARG002 - 保留原始值,不自动添加
|
is_stream: bool = False, # noqa: ARG002 - 保留原始值,不自动添加
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
透传请求体 - 原样复制,不做任何修改
|
透传请求体 - 原样复制,不做任何修改
|
||||||
|
|
||||||
@@ -218,13 +217,13 @@ class PassthroughRequestBuilder(RequestBuilder):
|
|||||||
|
|
||||||
def build_headers(
|
def build_headers(
|
||||||
self,
|
self,
|
||||||
original_headers: Dict[str, str],
|
original_headers: dict[str, str],
|
||||||
endpoint: Any,
|
endpoint: Any,
|
||||||
key: Any,
|
key: Any,
|
||||||
*,
|
*,
|
||||||
extra_headers: Optional[Dict[str, str]] = None,
|
extra_headers: dict[str, str] | None = None,
|
||||||
pre_computed_auth: Optional[Tuple[str, str]] = None,
|
pre_computed_auth: tuple[str, str] | None = None,
|
||||||
) -> Dict[str, str]:
|
) -> dict[str, str]:
|
||||||
"""
|
"""
|
||||||
透传请求头 - 清理敏感头部(黑名单),透传其他所有头部
|
透传请求头 - 清理敏感头部(黑名单),透传其他所有头部
|
||||||
|
|
||||||
@@ -289,11 +288,11 @@ class PassthroughRequestBuilder(RequestBuilder):
|
|||||||
|
|
||||||
|
|
||||||
def build_passthrough_request(
|
def build_passthrough_request(
|
||||||
original_body: Dict[str, Any],
|
original_body: dict[str, Any],
|
||||||
original_headers: Dict[str, str],
|
original_headers: dict[str, str],
|
||||||
endpoint: Any,
|
endpoint: Any,
|
||||||
key: Any,
|
key: Any,
|
||||||
) -> Tuple[Dict[str, Any], Dict[str, str]]:
|
) -> tuple[dict[str, Any], dict[str, str]]:
|
||||||
"""
|
"""
|
||||||
构建透传模式的请求
|
构建透传模式的请求
|
||||||
|
|
||||||
|
|||||||
@@ -4,7 +4,7 @@
|
|||||||
|
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from typing import Any, Dict, Optional
|
from typing import Any
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -13,14 +13,14 @@ class ParsedChunk:
|
|||||||
|
|
||||||
# 原始数据
|
# 原始数据
|
||||||
raw_line: str
|
raw_line: str
|
||||||
event_type: Optional[str] = None
|
event_type: str | None = None
|
||||||
data: Optional[Dict[str, Any]] = None
|
data: dict[str, Any] | None = None
|
||||||
|
|
||||||
# 提取的内容
|
# 提取的内容
|
||||||
text_delta: str = ""
|
text_delta: str = ""
|
||||||
is_done: bool = False
|
is_done: bool = False
|
||||||
is_error: bool = False
|
is_error: bool = False
|
||||||
error_message: Optional[str] = None
|
error_message: str | None = None
|
||||||
|
|
||||||
# 使用量信息(通常在最后一个 chunk 中)
|
# 使用量信息(通常在最后一个 chunk 中)
|
||||||
input_tokens: int = 0
|
input_tokens: int = 0
|
||||||
@@ -29,7 +29,7 @@ class ParsedChunk:
|
|||||||
cache_read_tokens: int = 0
|
cache_read_tokens: int = 0
|
||||||
|
|
||||||
# 响应 ID
|
# 响应 ID
|
||||||
response_id: Optional[str] = None
|
response_id: str | None = None
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -48,21 +48,21 @@ class StreamStats:
|
|||||||
|
|
||||||
# 内容
|
# 内容
|
||||||
collected_text: str = ""
|
collected_text: str = ""
|
||||||
response_id: Optional[str] = None
|
response_id: str | None = None
|
||||||
|
|
||||||
# 状态
|
# 状态
|
||||||
has_completion: bool = False
|
has_completion: bool = False
|
||||||
status_code: int = 200
|
status_code: int = 200
|
||||||
error_message: Optional[str] = None
|
error_message: str | None = None
|
||||||
|
|
||||||
# Provider 信息
|
# Provider 信息
|
||||||
provider_name: Optional[str] = None
|
provider_name: str | None = None
|
||||||
endpoint_id: Optional[str] = None
|
endpoint_id: str | None = None
|
||||||
key_id: Optional[str] = None
|
key_id: str | None = None
|
||||||
|
|
||||||
# 响应头和完整响应
|
# 响应头和完整响应
|
||||||
response_headers: Dict[str, str] = field(default_factory=dict)
|
response_headers: dict[str, str] = field(default_factory=dict)
|
||||||
final_response: Optional[Dict[str, Any]] = None
|
final_response: dict[str, Any] | None = None
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -70,12 +70,12 @@ class ParsedResponse:
|
|||||||
"""解析后的非流式响应"""
|
"""解析后的非流式响应"""
|
||||||
|
|
||||||
# 原始响应
|
# 原始响应
|
||||||
raw_response: Dict[str, Any]
|
raw_response: dict[str, Any]
|
||||||
status_code: int
|
status_code: int
|
||||||
|
|
||||||
# 提取的内容
|
# 提取的内容
|
||||||
text_content: str = ""
|
text_content: str = ""
|
||||||
response_id: Optional[str] = None
|
response_id: str | None = None
|
||||||
|
|
||||||
# 使用量
|
# 使用量
|
||||||
input_tokens: int = 0
|
input_tokens: int = 0
|
||||||
@@ -85,10 +85,10 @@ class ParsedResponse:
|
|||||||
|
|
||||||
# 错误信息
|
# 错误信息
|
||||||
is_error: bool = False
|
is_error: bool = False
|
||||||
error_type: Optional[str] = None
|
error_type: str | None = None
|
||||||
error_message: Optional[str] = None
|
error_message: str | None = None
|
||||||
# 从响应体解析出的嵌套状态码(当 HTTP 200 但响应体含错误时使用)
|
# 从响应体解析出的嵌套状态码(当 HTTP 200 但响应体含错误时使用)
|
||||||
embedded_status_code: Optional[int] = None
|
embedded_status_code: int | None = None
|
||||||
|
|
||||||
|
|
||||||
class ResponseParser(ABC):
|
class ResponseParser(ABC):
|
||||||
@@ -106,7 +106,7 @@ class ResponseParser(ABC):
|
|||||||
api_format: str = "UNKNOWN"
|
api_format: str = "UNKNOWN"
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def parse_sse_line(self, line: str, stats: StreamStats) -> Optional[ParsedChunk]:
|
def parse_sse_line(self, line: str, stats: StreamStats) -> ParsedChunk | None:
|
||||||
"""
|
"""
|
||||||
解析单行 SSE 数据
|
解析单行 SSE 数据
|
||||||
|
|
||||||
@@ -120,7 +120,7 @@ class ResponseParser(ABC):
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def parse_response(self, response: Dict[str, Any], status_code: int) -> ParsedResponse:
|
def parse_response(self, response: dict[str, Any], status_code: int) -> ParsedResponse:
|
||||||
"""
|
"""
|
||||||
解析非流式响应
|
解析非流式响应
|
||||||
|
|
||||||
@@ -134,7 +134,7 @@ class ResponseParser(ABC):
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def extract_usage_from_response(self, response: Dict[str, Any]) -> Dict[str, int]:
|
def extract_usage_from_response(self, response: dict[str, Any]) -> dict[str, int]:
|
||||||
"""
|
"""
|
||||||
从响应中提取 token 使用量
|
从响应中提取 token 使用量
|
||||||
|
|
||||||
@@ -147,7 +147,7 @@ class ResponseParser(ABC):
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def extract_text_content(self, response: Dict[str, Any]) -> str:
|
def extract_text_content(self, response: dict[str, Any]) -> str:
|
||||||
"""
|
"""
|
||||||
从响应中提取文本内容
|
从响应中提取文本内容
|
||||||
|
|
||||||
@@ -159,7 +159,7 @@ class ResponseParser(ABC):
|
|||||||
"""
|
"""
|
||||||
pass
|
pass
|
||||||
|
|
||||||
def is_error_response(self, response: Dict[str, Any]) -> bool:
|
def is_error_response(self, response: dict[str, Any]) -> bool:
|
||||||
"""
|
"""
|
||||||
判断响应是否为错误响应
|
判断响应是否为错误响应
|
||||||
|
|
||||||
|
|||||||
@@ -8,9 +8,11 @@
|
|||||||
- 请求/响应数据
|
- 请求/响应数据
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
import time
|
import time
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional
|
from typing import TYPE_CHECKING, Any
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from src.core.api_format.conversion.stream_state import StreamState
|
from src.core.api_format.conversion.stream_state import StreamState
|
||||||
@@ -35,16 +37,16 @@ class StreamContext:
|
|||||||
api_key_id: int = 0
|
api_key_id: int = 0
|
||||||
|
|
||||||
# Provider 信息(在请求执行时填充)
|
# Provider 信息(在请求执行时填充)
|
||||||
provider_name: Optional[str] = None
|
provider_name: str | None = None
|
||||||
provider_id: Optional[str] = None
|
provider_id: str | None = None
|
||||||
endpoint_id: Optional[str] = None
|
endpoint_id: str | None = None
|
||||||
key_id: Optional[str] = None
|
key_id: str | None = None
|
||||||
attempt_id: Optional[str] = None
|
attempt_id: str | None = None
|
||||||
attempt_synced: bool = False
|
attempt_synced: bool = False
|
||||||
provider_api_format: Optional[str] = None # Provider 的响应格式
|
provider_api_format: str | None = None # Provider 的响应格式
|
||||||
|
|
||||||
# 模型映射
|
# 模型映射
|
||||||
mapped_model: Optional[str] = None
|
mapped_model: str | None = None
|
||||||
|
|
||||||
# Token 统计
|
# Token 统计
|
||||||
input_tokens: int = 0
|
input_tokens: int = 0
|
||||||
@@ -53,33 +55,33 @@ class StreamContext:
|
|||||||
cache_creation_tokens: int = 0
|
cache_creation_tokens: int = 0
|
||||||
|
|
||||||
# 响应内容
|
# 响应内容
|
||||||
_collected_text_parts: List[str] = field(default_factory=list, repr=False)
|
_collected_text_parts: list[str] = field(default_factory=list, repr=False)
|
||||||
response_id: Optional[str] = None
|
response_id: str | None = None
|
||||||
final_usage: Optional[Dict[str, Any]] = None
|
final_usage: dict[str, Any] | None = None
|
||||||
final_response: Optional[Dict[str, Any]] = None
|
final_response: dict[str, Any] | None = None
|
||||||
|
|
||||||
# 时间指标
|
# 时间指标
|
||||||
first_byte_time_ms: Optional[int] = None # 首字时间 (TTFB - Time To First Byte)
|
first_byte_time_ms: int | None = None # 首字时间 (TTFB - Time To First Byte)
|
||||||
start_time: float = field(default_factory=time.time)
|
start_time: float = field(default_factory=time.time)
|
||||||
|
|
||||||
# 响应状态
|
# 响应状态
|
||||||
status_code: int = 200
|
status_code: int = 200
|
||||||
error_message: Optional[str] = None # 客户端友好的错误消息
|
error_message: str | None = None # 客户端友好的错误消息
|
||||||
upstream_response: Optional[str] = None # 原始 Provider 响应(用于请求链路追踪)
|
upstream_response: str | None = None # 原始 Provider 响应(用于请求链路追踪)
|
||||||
has_completion: bool = False
|
has_completion: bool = False
|
||||||
|
|
||||||
# 请求/响应数据
|
# 请求/响应数据
|
||||||
response_headers: Dict[str, str] = field(default_factory=dict) # 提供商响应头
|
response_headers: dict[str, str] = field(default_factory=dict) # 提供商响应头
|
||||||
client_response_headers: Dict[str, str] = field(default_factory=dict) # 返回给客户端的响应头
|
client_response_headers: dict[str, str] = field(default_factory=dict) # 返回给客户端的响应头
|
||||||
provider_request_headers: Dict[str, str] = field(default_factory=dict)
|
provider_request_headers: dict[str, str] = field(default_factory=dict)
|
||||||
provider_request_body: Optional[Dict[str, Any]] = None
|
provider_request_body: dict[str, Any] | None = None
|
||||||
|
|
||||||
# 格式转换信息(CLI handler 需要)
|
# 格式转换信息(CLI handler 需要)
|
||||||
client_api_format: str = ""
|
client_api_format: str = ""
|
||||||
needs_conversion: bool = False # 是否需要跨格式转换(由 handler 层设置)
|
needs_conversion: bool = False # 是否需要跨格式转换(由 handler 层设置)
|
||||||
|
|
||||||
# Provider 响应元数据(CLI handler 需要)
|
# Provider 响应元数据(CLI handler 需要)
|
||||||
response_metadata: Dict[str, Any] = field(default_factory=dict)
|
response_metadata: dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
# 整流标记(Thinking Rectifier)
|
# 整流标记(Thinking Rectifier)
|
||||||
rectified: bool = False # 请求是否经过整流(移除 thinking 块后重试)
|
rectified: bool = False # 请求是否经过整流(移除 thinking 块后重试)
|
||||||
@@ -87,10 +89,10 @@ class StreamContext:
|
|||||||
# 流式处理统计
|
# 流式处理统计
|
||||||
data_count: int = 0
|
data_count: int = 0
|
||||||
chunk_count: int = 0
|
chunk_count: int = 0
|
||||||
parsed_chunks: List[Dict[str, Any]] = field(default_factory=list)
|
parsed_chunks: list[dict[str, Any]] = field(default_factory=list)
|
||||||
|
|
||||||
# 流式格式转换状态(跨 chunk 追踪)
|
# 流式格式转换状态(跨 chunk 追踪)
|
||||||
stream_conversion_state: Optional["StreamState"] = None
|
stream_conversion_state: StreamState | None = None
|
||||||
|
|
||||||
def reset_for_retry(self) -> None:
|
def reset_for_retry(self) -> None:
|
||||||
"""
|
"""
|
||||||
@@ -138,7 +140,7 @@ class StreamContext:
|
|||||||
provider_id: str,
|
provider_id: str,
|
||||||
endpoint_id: str,
|
endpoint_id: str,
|
||||||
key_id: str,
|
key_id: str,
|
||||||
provider_api_format: Optional[str] = None,
|
provider_api_format: str | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""更新 Provider 信息"""
|
"""更新 Provider 信息"""
|
||||||
self.provider_name = provider_name
|
self.provider_name = provider_name
|
||||||
@@ -149,10 +151,10 @@ class StreamContext:
|
|||||||
|
|
||||||
def update_usage(
|
def update_usage(
|
||||||
self,
|
self,
|
||||||
input_tokens: Optional[int] = None,
|
input_tokens: int | None = None,
|
||||||
output_tokens: Optional[int] = None,
|
output_tokens: int | None = None,
|
||||||
cached_tokens: Optional[int] = None,
|
cached_tokens: int | None = None,
|
||||||
cache_creation_tokens: Optional[int] = None,
|
cache_creation_tokens: int | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""
|
"""
|
||||||
更新 Token 使用统计
|
更新 Token 使用统计
|
||||||
@@ -194,7 +196,7 @@ class StreamContext:
|
|||||||
self,
|
self,
|
||||||
status_code: int,
|
status_code: int,
|
||||||
error_message: str,
|
error_message: str,
|
||||||
upstream_response: Optional[str] = None,
|
upstream_response: str | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""
|
"""
|
||||||
标记请求失败
|
标记请求失败
|
||||||
@@ -230,7 +232,7 @@ class StreamContext:
|
|||||||
"""检查是否因客户端断开连接而结束"""
|
"""检查是否因客户端断开连接而结束"""
|
||||||
return self.status_code == 499
|
return self.status_code == 499
|
||||||
|
|
||||||
def build_response_body(self, response_time_ms: int) -> Dict[str, Any]:
|
def build_response_body(self, response_time_ms: int) -> dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
构建响应体元数据
|
构建响应体元数据
|
||||||
|
|
||||||
|
|||||||
@@ -13,7 +13,10 @@ import asyncio
|
|||||||
import codecs
|
import codecs
|
||||||
import json
|
import json
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import Any, AsyncGenerator, Callable, Optional
|
from typing import Any
|
||||||
|
|
||||||
|
from collections.abc import Callable
|
||||||
|
from collections.abc import AsyncGenerator
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
|
|
||||||
@@ -65,10 +68,10 @@ class StreamProcessor:
|
|||||||
self,
|
self,
|
||||||
request_id: str,
|
request_id: str,
|
||||||
default_parser: ResponseParser,
|
default_parser: ResponseParser,
|
||||||
on_streaming_start: Optional[Callable[[], None]] = None,
|
on_streaming_start: Callable[[], None] | None = None,
|
||||||
*,
|
*,
|
||||||
collect_text: bool = False,
|
collect_text: bool = False,
|
||||||
smoothing_config: Optional[StreamSmoothingConfig] = None,
|
smoothing_config: StreamSmoothingConfig | None = None,
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
初始化流处理器
|
初始化流处理器
|
||||||
@@ -105,7 +108,7 @@ class StreamProcessor:
|
|||||||
def handle_sse_event(
|
def handle_sse_event(
|
||||||
self,
|
self,
|
||||||
ctx: StreamContext,
|
ctx: StreamContext,
|
||||||
event_name: Optional[str],
|
event_name: str | None,
|
||||||
data_str: str,
|
data_str: str,
|
||||||
*,
|
*,
|
||||||
skip_record: bool = False,
|
skip_record: bool = False,
|
||||||
@@ -363,7 +366,7 @@ class StreamProcessor:
|
|||||||
):
|
):
|
||||||
# 重新抛出可重试的 Provider 异常,触发故障转移
|
# 重新抛出可重试的 Provider 异常,触发故障转移
|
||||||
raise
|
raise
|
||||||
except (OSError, IOError) as e:
|
except OSError as e:
|
||||||
# 网络 I/O 异常:记录警告,可能需要重试
|
# 网络 I/O 异常:记录警告,可能需要重试
|
||||||
logger.warning(f" [{self.request_id}] 预读流时发生网络异常: {type(e).__name__}: {e}")
|
logger.warning(f" [{self.request_id}] 预读流时发生网络异常: {type(e).__name__}: {e}")
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -382,10 +385,10 @@ class StreamProcessor:
|
|||||||
byte_iterator: Any,
|
byte_iterator: Any,
|
||||||
response_ctx: Any,
|
response_ctx: Any,
|
||||||
http_client: httpx.AsyncClient,
|
http_client: httpx.AsyncClient,
|
||||||
prefetched_chunks: Optional[list] = None,
|
prefetched_chunks: list | None = None,
|
||||||
*,
|
*,
|
||||||
start_time: Optional[float] = None,
|
start_time: float | None = None,
|
||||||
) -> AsyncGenerator[bytes, None]:
|
) -> AsyncGenerator[bytes]:
|
||||||
"""
|
"""
|
||||||
创建响应流生成器
|
创建响应流生成器
|
||||||
|
|
||||||
@@ -547,9 +550,7 @@ class StreamProcessor:
|
|||||||
logger.warning(f"[{self.request_id}] 流式格式转换失败: {conv_err}")
|
logger.warning(f"[{self.request_id}] 流式格式转换失败: {conv_err}")
|
||||||
payload = _build_stream_error_payload("响应格式转换失败,请稍后重试")
|
payload = _build_stream_error_payload("响应格式转换失败,请稍后重试")
|
||||||
error_bytes = (
|
error_bytes = (
|
||||||
f"data: {json.dumps(payload, ensure_ascii=False)}\n\n".encode(
|
f"data: {json.dumps(payload, ensure_ascii=False)}\n\n".encode()
|
||||||
"utf-8"
|
|
||||||
)
|
|
||||||
)
|
)
|
||||||
done_bytes = (
|
done_bytes = (
|
||||||
b"data: [DONE]\n\n" if client_format.startswith("OPENAI") else b""
|
b"data: [DONE]\n\n" if client_format.startswith("OPENAI") else b""
|
||||||
@@ -570,7 +571,7 @@ class StreamProcessor:
|
|||||||
# 统一使用 SSE 格式输出(Gemini streamGenerateContent 也使用 SSE)
|
# 统一使用 SSE 格式输出(Gemini streamGenerateContent 也使用 SSE)
|
||||||
# 参考: https://ai.google.dev/api/generate-content
|
# 参考: https://ai.google.dev/api/generate-content
|
||||||
out.append(
|
out.append(
|
||||||
f"data: {json.dumps(evt, ensure_ascii=False)}\n\n".encode("utf-8")
|
f"data: {json.dumps(evt, ensure_ascii=False)}\n\n".encode()
|
||||||
)
|
)
|
||||||
return out
|
return out
|
||||||
|
|
||||||
@@ -769,9 +770,9 @@ class StreamProcessor:
|
|||||||
async def create_monitored_stream(
|
async def create_monitored_stream(
|
||||||
self,
|
self,
|
||||||
ctx: StreamContext,
|
ctx: StreamContext,
|
||||||
stream_generator: AsyncGenerator[bytes, None],
|
stream_generator: AsyncGenerator[bytes],
|
||||||
is_disconnected: Callable[[], Any],
|
is_disconnected: Callable[[], Any],
|
||||||
) -> AsyncGenerator[bytes, None]:
|
) -> AsyncGenerator[bytes]:
|
||||||
"""
|
"""
|
||||||
创建带监控的流生成器
|
创建带监控的流生成器
|
||||||
|
|
||||||
@@ -833,8 +834,8 @@ class StreamProcessor:
|
|||||||
|
|
||||||
async def create_smoothed_stream(
|
async def create_smoothed_stream(
|
||||||
self,
|
self,
|
||||||
stream_generator: AsyncGenerator[bytes, None],
|
stream_generator: AsyncGenerator[bytes],
|
||||||
) -> AsyncGenerator[bytes, None]:
|
) -> AsyncGenerator[bytes]:
|
||||||
"""
|
"""
|
||||||
创建平滑输出的流生成器
|
创建平滑输出的流生成器
|
||||||
|
|
||||||
@@ -933,7 +934,7 @@ class StreamProcessor:
|
|||||||
if buffer:
|
if buffer:
|
||||||
yield buffer
|
yield buffer
|
||||||
|
|
||||||
def _get_extractor(self, format_name: str) -> Optional[ContentExtractor]:
|
def _get_extractor(self, format_name: str) -> ContentExtractor | None:
|
||||||
"""获取或创建格式对应的提取器(带缓存)"""
|
"""获取或创建格式对应的提取器(带缓存)"""
|
||||||
if format_name not in self._extractors:
|
if format_name not in self._extractors:
|
||||||
extractor = get_extractor(format_name)
|
extractor = get_extractor(format_name)
|
||||||
@@ -943,7 +944,7 @@ class StreamProcessor:
|
|||||||
|
|
||||||
def _detect_format_and_extract(
|
def _detect_format_and_extract(
|
||||||
self, data: dict
|
self, data: dict
|
||||||
) -> tuple[Optional[str], Optional[ContentExtractor]]:
|
) -> tuple[str | None, ContentExtractor | None]:
|
||||||
"""
|
"""
|
||||||
检测数据格式并提取内容
|
检测数据格式并提取内容
|
||||||
|
|
||||||
@@ -998,10 +999,10 @@ class StreamProcessor:
|
|||||||
|
|
||||||
|
|
||||||
async def create_smoothed_stream(
|
async def create_smoothed_stream(
|
||||||
stream_generator: AsyncGenerator[bytes, None],
|
stream_generator: AsyncGenerator[bytes],
|
||||||
chunk_size: int = 20,
|
chunk_size: int = 20,
|
||||||
delay_ms: int = 8,
|
delay_ms: int = 8,
|
||||||
) -> AsyncGenerator[bytes, None]:
|
) -> AsyncGenerator[bytes]:
|
||||||
"""
|
"""
|
||||||
独立的平滑流生成函数
|
独立的平滑流生成函数
|
||||||
|
|
||||||
@@ -1032,7 +1033,7 @@ class _LightweightSmoother:
|
|||||||
self.delay_ms = delay_ms
|
self.delay_ms = delay_ms
|
||||||
self._extractors: dict[str, ContentExtractor] = {}
|
self._extractors: dict[str, ContentExtractor] = {}
|
||||||
|
|
||||||
def _get_extractor(self, format_name: str) -> Optional[ContentExtractor]:
|
def _get_extractor(self, format_name: str) -> ContentExtractor | None:
|
||||||
if format_name not in self._extractors:
|
if format_name not in self._extractors:
|
||||||
extractor = get_extractor(format_name)
|
extractor = get_extractor(format_name)
|
||||||
if extractor:
|
if extractor:
|
||||||
@@ -1041,7 +1042,7 @@ class _LightweightSmoother:
|
|||||||
|
|
||||||
def _detect_format_and_extract(
|
def _detect_format_and_extract(
|
||||||
self, data: dict
|
self, data: dict
|
||||||
) -> tuple[Optional[str], Optional[ContentExtractor]]:
|
) -> tuple[str | None, ContentExtractor | None]:
|
||||||
for format_name in get_extractor_formats():
|
for format_name in get_extractor_formats():
|
||||||
extractor = self._get_extractor(format_name)
|
extractor = self._get_extractor(format_name)
|
||||||
if extractor:
|
if extractor:
|
||||||
@@ -1060,8 +1061,8 @@ class _LightweightSmoother:
|
|||||||
return [content[i : i + self.chunk_size] for i in range(0, text_length, self.chunk_size)]
|
return [content[i : i + self.chunk_size] for i in range(0, text_length, self.chunk_size)]
|
||||||
|
|
||||||
async def smooth(
|
async def smooth(
|
||||||
self, stream_generator: AsyncGenerator[bytes, None]
|
self, stream_generator: AsyncGenerator[bytes]
|
||||||
) -> AsyncGenerator[bytes, None]:
|
) -> AsyncGenerator[bytes]:
|
||||||
buffer = b""
|
buffer = b""
|
||||||
is_first_content = True
|
is_first_content = True
|
||||||
|
|
||||||
|
|||||||
@@ -9,7 +9,7 @@
|
|||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import time
|
import time
|
||||||
from typing import Any, Dict, Optional
|
from typing import Any
|
||||||
|
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
@@ -58,8 +58,8 @@ class StreamTelemetryRecorder:
|
|||||||
async def record_stream_stats(
|
async def record_stream_stats(
|
||||||
self,
|
self,
|
||||||
ctx: StreamContext,
|
ctx: StreamContext,
|
||||||
original_headers: Dict[str, str],
|
original_headers: dict[str, str],
|
||||||
original_request_body: Dict[str, Any],
|
original_request_body: dict[str, Any],
|
||||||
start_time: float,
|
start_time: float,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""
|
"""
|
||||||
@@ -144,9 +144,9 @@ class StreamTelemetryRecorder:
|
|||||||
self,
|
self,
|
||||||
writer: TelemetryWriter,
|
writer: TelemetryWriter,
|
||||||
ctx: StreamContext,
|
ctx: StreamContext,
|
||||||
original_headers: Dict[str, str],
|
original_headers: dict[str, str],
|
||||||
actual_request_body: Dict[str, Any],
|
actual_request_body: dict[str, Any],
|
||||||
response_body: Optional[Dict[str, Any]],
|
response_body: dict[str, Any] | None,
|
||||||
response_time_ms: int,
|
response_time_ms: int,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""记录成功的请求"""
|
"""记录成功的请求"""
|
||||||
@@ -193,9 +193,9 @@ class StreamTelemetryRecorder:
|
|||||||
self,
|
self,
|
||||||
writer: TelemetryWriter,
|
writer: TelemetryWriter,
|
||||||
ctx: StreamContext,
|
ctx: StreamContext,
|
||||||
original_headers: Dict[str, str],
|
original_headers: dict[str, str],
|
||||||
actual_request_body: Dict[str, Any],
|
actual_request_body: dict[str, Any],
|
||||||
response_body: Optional[Dict[str, Any]],
|
response_body: dict[str, Any] | None,
|
||||||
response_time_ms: int,
|
response_time_ms: int,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""记录失败的请求"""
|
"""记录失败的请求"""
|
||||||
@@ -236,9 +236,9 @@ class StreamTelemetryRecorder:
|
|||||||
self,
|
self,
|
||||||
writer: TelemetryWriter,
|
writer: TelemetryWriter,
|
||||||
ctx: StreamContext,
|
ctx: StreamContext,
|
||||||
original_headers: Dict[str, str],
|
original_headers: dict[str, str],
|
||||||
actual_request_body: Dict[str, Any],
|
actual_request_body: dict[str, Any],
|
||||||
response_body: Optional[Dict[str, Any]],
|
response_body: dict[str, Any] | None,
|
||||||
response_time_ms: int,
|
response_time_ms: int,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""记录客户端取消的请求"""
|
"""记录客户端取消的请求"""
|
||||||
@@ -285,7 +285,7 @@ class StreamTelemetryRecorder:
|
|||||||
|
|
||||||
from src.services.request.candidate import RequestCandidateService
|
from src.services.request.candidate import RequestCandidateService
|
||||||
|
|
||||||
extra_data: Dict[str, Any] = {
|
extra_data: dict[str, Any] = {
|
||||||
"stream_completed": ctx.is_success(),
|
"stream_completed": ctx.is_success(),
|
||||||
"data_count": ctx.data_count,
|
"data_count": ctx.data_count,
|
||||||
}
|
}
|
||||||
@@ -358,7 +358,7 @@ class StreamTelemetryRecorder:
|
|||||||
status: str,
|
status: str,
|
||||||
response_time_ms: int,
|
response_time_ms: int,
|
||||||
status_code: int = 200,
|
status_code: int = 200,
|
||||||
error_message: Optional[str] = None,
|
error_message: str | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""直接更新 Usage 表的状态字段"""
|
"""直接更新 Usage 表的状态字段"""
|
||||||
try:
|
try:
|
||||||
@@ -378,7 +378,7 @@ class StreamTelemetryRecorder:
|
|||||||
|
|
||||||
async def _get_telemetry_writer(
|
async def _get_telemetry_writer(
|
||||||
self, bg_db: Session, ctx: StreamContext, response_time_ms: int
|
self, bg_db: Session, ctx: StreamContext, response_time_ms: int
|
||||||
) -> Optional[TelemetryWriter]:
|
) -> TelemetryWriter | None:
|
||||||
if config.usage_queue_enabled and self.user_id and self.api_key_id:
|
if config.usage_queue_enabled and self.user_id and self.api_key_id:
|
||||||
return QueueTelemetryWriter(
|
return QueueTelemetryWriter(
|
||||||
request_id=self.request_id,
|
request_id=self.request_id,
|
||||||
@@ -400,9 +400,9 @@ class StreamTelemetryRecorder:
|
|||||||
self,
|
self,
|
||||||
writer: TelemetryWriter,
|
writer: TelemetryWriter,
|
||||||
ctx: StreamContext,
|
ctx: StreamContext,
|
||||||
original_headers: Dict[str, str],
|
original_headers: dict[str, str],
|
||||||
actual_request_body: Dict[str, Any],
|
actual_request_body: dict[str, Any],
|
||||||
response_body: Optional[Dict[str, Any]],
|
response_body: dict[str, Any] | None,
|
||||||
response_time_ms: int,
|
response_time_ms: int,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""根据上下文状态分发到对应的记录方法"""
|
"""根据上下文状态分发到对应的记录方法"""
|
||||||
@@ -430,7 +430,7 @@ class StreamTelemetryRecorder:
|
|||||||
return "cancelled"
|
return "cancelled"
|
||||||
return "failed"
|
return "failed"
|
||||||
|
|
||||||
def _build_db_writer(self, bg_db: Session) -> Optional[DbTelemetryWriter]:
|
def _build_db_writer(self, bg_db: Session) -> DbTelemetryWriter | None:
|
||||||
user = bg_db.query(User).filter(User.id == self.user_id).first()
|
user = bg_db.query(User).filter(User.id == self.user_id).first()
|
||||||
api_key_obj = bg_db.query(ApiKey).filter(ApiKey.id == self.api_key_id).first()
|
api_key_obj = bg_db.query(ApiKey).filter(ApiKey.id == self.api_key_id).first()
|
||||||
|
|
||||||
|
|||||||
@@ -4,8 +4,9 @@ Handler 基础工具函数
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
|
||||||
import json
|
import json
|
||||||
from typing import TYPE_CHECKING, Any, Dict, Optional
|
from typing import TYPE_CHECKING, Any
|
||||||
|
|
||||||
from src.core.exceptions import EmbeddedErrorException, ProviderNotAvailableException
|
from src.core.exceptions import EmbeddedErrorException, ProviderNotAvailableException
|
||||||
from src.core.api_format import filter_response_headers
|
from src.core.api_format import filter_response_headers
|
||||||
@@ -15,7 +16,7 @@ if TYPE_CHECKING:
|
|||||||
from src.core.api_format.conversion.registry import FormatConversionRegistry
|
from src.core.api_format.conversion.registry import FormatConversionRegistry
|
||||||
|
|
||||||
|
|
||||||
def get_format_converter_registry() -> "FormatConversionRegistry":
|
def get_format_converter_registry() -> FormatConversionRegistry:
|
||||||
"""
|
"""
|
||||||
获取格式转换注册表(线程安全)
|
获取格式转换注册表(线程安全)
|
||||||
|
|
||||||
@@ -31,7 +32,7 @@ def get_format_converter_registry() -> "FormatConversionRegistry":
|
|||||||
return format_conversion_registry
|
return format_conversion_registry
|
||||||
|
|
||||||
|
|
||||||
def extract_cache_creation_tokens(usage: Dict[str, Any]) -> int:
|
def extract_cache_creation_tokens(usage: dict[str, Any]) -> int:
|
||||||
"""
|
"""
|
||||||
提取缓存创建 tokens(兼容三种格式)
|
提取缓存创建 tokens(兼容三种格式)
|
||||||
|
|
||||||
@@ -99,7 +100,7 @@ def extract_cache_creation_tokens(usage: Dict[str, Any]) -> int:
|
|||||||
return old_format
|
return old_format
|
||||||
|
|
||||||
|
|
||||||
def build_sse_headers(extra_headers: Optional[Dict[str, str]] = None) -> Dict[str, str]:
|
def build_sse_headers(extra_headers: dict[str, str] | None = None) -> dict[str, str]:
|
||||||
"""
|
"""
|
||||||
构建 SSE(text/event-stream)推荐响应头,用于减少代理缓冲带来的卡顿/成段输出。
|
构建 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 可避免部分代理对流做压缩/改写导致缓冲
|
- Cache-Control: no-transform 可避免部分代理对流做压缩/改写导致缓冲
|
||||||
- X-Accel-Buffering: no 可显式提示 Nginx 关闭缓冲(即使全局已关闭也无害)
|
- X-Accel-Buffering: no 可显式提示 Nginx 关闭缓冲(即使全局已关闭也无害)
|
||||||
"""
|
"""
|
||||||
headers: Dict[str, str] = {
|
headers: dict[str, str] = {
|
||||||
"Cache-Control": "no-cache, no-transform",
|
"Cache-Control": "no-cache, no-transform",
|
||||||
"X-Accel-Buffering": "no",
|
"X-Accel-Buffering": "no",
|
||||||
}
|
}
|
||||||
@@ -116,7 +117,7 @@ def build_sse_headers(extra_headers: Optional[Dict[str, str]] = None) -> Dict[st
|
|||||||
return headers
|
return headers
|
||||||
|
|
||||||
|
|
||||||
def filter_proxy_response_headers(headers: Optional[Dict[str, str]]) -> Dict[str, str]:
|
def filter_proxy_response_headers(headers: dict[str, str] | None) -> dict[str, str]:
|
||||||
"""
|
"""
|
||||||
过滤上游响应头中不应透传给客户端的字段。
|
过滤上游响应头中不应透传给客户端的字段。
|
||||||
|
|
||||||
@@ -148,8 +149,8 @@ def check_prefetched_response_error(
|
|||||||
parser: Any,
|
parser: Any,
|
||||||
request_id: str,
|
request_id: str,
|
||||||
provider_name: str,
|
provider_name: str,
|
||||||
endpoint_id: Optional[str],
|
endpoint_id: str | None,
|
||||||
base_url: Optional[str],
|
base_url: str | None,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""
|
"""
|
||||||
检查预读的响应是否为非 SSE 格式的错误响应(HTML 或纯 JSON 错误)
|
检查预读的响应是否为非 SSE 格式的错误响应(HTML 或纯 JSON 错误)
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ Claude Chat Adapter - 基于 ChatAdapterBase 的 Claude Chat API 适配器
|
|||||||
处理 /v1/messages 端点的 Claude Chat 格式请求。
|
处理 /v1/messages 端点的 Claude Chat 格式请求。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from typing import Any, Dict, Optional, Tuple, Type
|
from typing import Any
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
from fastapi import HTTPException, Request
|
from fastapi import HTTPException, Request
|
||||||
@@ -25,9 +25,9 @@ class ClaudeCapabilityDetector:
|
|||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def detect_from_headers(
|
def detect_from_headers(
|
||||||
headers: Dict[str, str],
|
headers: dict[str, str],
|
||||||
request_body: Optional[Dict[str, Any]] = None,
|
request_body: dict[str, Any] | None = None,
|
||||||
) -> Dict[str, bool]:
|
) -> dict[str, bool]:
|
||||||
"""
|
"""
|
||||||
从 Claude 请求头检测能力需求
|
从 Claude 请求头检测能力需求
|
||||||
|
|
||||||
@@ -38,7 +38,7 @@ class ClaudeCapabilityDetector:
|
|||||||
headers: 请求头字典
|
headers: 请求头字典
|
||||||
request_body: 请求体(Claude 不使用,保留用于接口统一)
|
request_body: 请求体(Claude 不使用,保留用于接口统一)
|
||||||
"""
|
"""
|
||||||
requirements: Dict[str, bool] = {}
|
requirements: dict[str, bool] = {}
|
||||||
|
|
||||||
# 使用统一的大小写不敏感获取
|
# 使用统一的大小写不敏感获取
|
||||||
beta_header = get_header_value(headers, "anthropic-beta")
|
beta_header = get_header_value(headers, "anthropic-beta")
|
||||||
@@ -61,21 +61,21 @@ class ClaudeChatAdapter(ChatAdapterBase):
|
|||||||
name = "claude.chat"
|
name = "claude.chat"
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def HANDLER_CLASS(self) -> Type[ChatHandlerBase]:
|
def HANDLER_CLASS(self) -> type[ChatHandlerBase]:
|
||||||
"""延迟导入 Handler 类避免循环依赖"""
|
"""延迟导入 Handler 类避免循环依赖"""
|
||||||
from src.api.handlers.claude.handler import ClaudeChatHandler
|
from src.api.handlers.claude.handler import ClaudeChatHandler
|
||||||
|
|
||||||
return ClaudeChatHandler
|
return ClaudeChatHandler
|
||||||
|
|
||||||
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
|
def __init__(self, allowed_api_formats: list[str] | None = None):
|
||||||
super().__init__(allowed_api_formats or ["CLAUDE"])
|
super().__init__(allowed_api_formats or ["CLAUDE"])
|
||||||
logger.info(f"[{self.name}] 初始化Chat模式适配器 | API格式: {self.allowed_api_formats}")
|
logger.info(f"[{self.name}] 初始化Chat模式适配器 | API格式: {self.allowed_api_formats}")
|
||||||
|
|
||||||
def detect_capability_requirements(
|
def detect_capability_requirements(
|
||||||
self,
|
self,
|
||||||
headers: Dict[str, str],
|
headers: dict[str, str],
|
||||||
request_body: Optional[Dict[str, Any]] = None,
|
request_body: dict[str, Any] | None = None,
|
||||||
) -> Dict[str, bool]:
|
) -> dict[str, bool]:
|
||||||
"""检测 Claude 请求中隐含的能力需求"""
|
"""检测 Claude 请求中隐含的能力需求"""
|
||||||
return ClaudeCapabilityDetector.detect_from_headers(headers)
|
return ClaudeCapabilityDetector.detect_from_headers(headers)
|
||||||
|
|
||||||
@@ -124,7 +124,7 @@ class ClaudeChatAdapter(ChatAdapterBase):
|
|||||||
)
|
)
|
||||||
return request
|
return request
|
||||||
|
|
||||||
def _build_audit_metadata(self, _payload: Dict[str, Any], request_obj) -> Dict[str, Any]:
|
def _build_audit_metadata(self, _payload: dict[str, Any], request_obj) -> dict[str, Any]:
|
||||||
"""构建 Claude Chat 特定的审计元数据"""
|
"""构建 Claude Chat 特定的审计元数据"""
|
||||||
role_counts: dict[str, int] = {}
|
role_counts: dict[str, int] = {}
|
||||||
for message in request_obj.messages:
|
for message in request_obj.messages:
|
||||||
@@ -153,8 +153,8 @@ class ClaudeChatAdapter(ChatAdapterBase):
|
|||||||
client: httpx.AsyncClient,
|
client: httpx.AsyncClient,
|
||||||
base_url: str,
|
base_url: str,
|
||||||
api_key: str,
|
api_key: str,
|
||||||
extra_headers: Optional[Dict[str, str]] = None,
|
extra_headers: dict[str, str] | None = None,
|
||||||
) -> Tuple[list, Optional[str]]:
|
) -> tuple[list, str | None]:
|
||||||
"""查询 Claude API 支持的模型列表"""
|
"""查询 Claude API 支持的模型列表"""
|
||||||
headers = cls.build_headers_with_extra(api_key, extra_headers)
|
headers = cls.build_headers_with_extra(api_key, extra_headers)
|
||||||
|
|
||||||
@@ -201,7 +201,7 @@ class ClaudeChatAdapter(ChatAdapterBase):
|
|||||||
# build_request_body 使用基类实现,通过 format_conversion_registry 自动转换 OPENAI -> CLAUDE
|
# build_request_body 使用基类实现,通过 format_conversion_registry 自动转换 OPENAI -> CLAUDE
|
||||||
|
|
||||||
|
|
||||||
def build_claude_adapter(x_app_header: Optional[str]):
|
def build_claude_adapter(x_app_header: str | None):
|
||||||
"""根据 x-app 头部构造 Chat 或 Claude Code 适配器。"""
|
"""根据 x-app 头部构造 Chat 或 Claude Code 适配器。"""
|
||||||
if x_app_header and x_app_header.lower() == "cli":
|
if x_app_header and x_app_header.lower() == "cli":
|
||||||
from src.api.handlers.claude_cli.adapter import ClaudeCliAdapter
|
from src.api.handlers.claude_cli.adapter import ClaudeCliAdapter
|
||||||
@@ -216,7 +216,7 @@ class ClaudeTokenCountAdapter(ApiAdapter):
|
|||||||
name = "claude.token_count"
|
name = "claude.token_count"
|
||||||
mode = ApiMode.STANDARD
|
mode = ApiMode.STANDARD
|
||||||
|
|
||||||
def extract_api_key(self, request: Request) -> Optional[str]:
|
def extract_api_key(self, request: Request) -> str | None:
|
||||||
"""从请求中提取 API 密钥 (x-api-key 或 Authorization: Bearer)"""
|
"""从请求中提取 API 密钥 (x-api-key 或 Authorization: Bearer)"""
|
||||||
# 优先检查 x-api-key
|
# 优先检查 x-api-key
|
||||||
api_key = request.headers.get("x-api-key")
|
api_key = request.headers.get("x-api-key")
|
||||||
|
|||||||
@@ -5,7 +5,7 @@ Claude Chat Handler - 基于通用 Chat Handler 基类的简化实现
|
|||||||
代码量从原来的 ~1470 行减少到 ~120 行。
|
代码量从原来的 ~1470 行减少到 ~120 行。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from typing import Any, Dict, Optional
|
from typing import Any
|
||||||
|
|
||||||
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
|
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
|
||||||
from src.api.handlers.base.utils import extract_cache_creation_tokens
|
from src.api.handlers.base.utils import extract_cache_creation_tokens
|
||||||
@@ -25,8 +25,8 @@ class ClaudeChatHandler(ChatHandlerBase):
|
|||||||
|
|
||||||
def extract_model_from_request(
|
def extract_model_from_request(
|
||||||
self,
|
self,
|
||||||
request_body: Dict[str, Any],
|
request_body: dict[str, Any],
|
||||||
path_params: Optional[Dict[str, Any]] = None, # noqa: ARG002
|
path_params: dict[str, Any] | None = None, # noqa: ARG002
|
||||||
) -> str:
|
) -> str:
|
||||||
"""
|
"""
|
||||||
从请求中提取模型名 - Claude 格式实现
|
从请求中提取模型名 - Claude 格式实现
|
||||||
@@ -45,9 +45,9 @@ class ClaudeChatHandler(ChatHandlerBase):
|
|||||||
|
|
||||||
def apply_mapped_model(
|
def apply_mapped_model(
|
||||||
self,
|
self,
|
||||||
request_body: Dict[str, Any],
|
request_body: dict[str, Any],
|
||||||
mapped_model: str,
|
mapped_model: str,
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
将映射后的模型名应用到请求体
|
将映射后的模型名应用到请求体
|
||||||
|
|
||||||
@@ -90,7 +90,7 @@ class ClaudeChatHandler(ChatHandlerBase):
|
|||||||
|
|
||||||
return request
|
return request
|
||||||
|
|
||||||
def _extract_usage(self, response: Dict) -> Dict[str, int]:
|
def _extract_usage(self, response: dict) -> dict[str, int]:
|
||||||
"""
|
"""
|
||||||
从 Claude 响应中提取 token 使用情况
|
从 Claude 响应中提取 token 使用情况
|
||||||
|
|
||||||
@@ -108,7 +108,7 @@ class ClaudeChatHandler(ChatHandlerBase):
|
|||||||
"cache_read_input_tokens": usage.get("cache_read_input_tokens", 0),
|
"cache_read_input_tokens": usage.get("cache_read_input_tokens", 0),
|
||||||
}
|
}
|
||||||
|
|
||||||
def _normalize_response(self, response: Dict[str, Any]) -> Dict[str, Any]:
|
def _normalize_response(self, response: dict[str, Any]) -> dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
规范化 Claude 响应
|
规范化 Claude 响应
|
||||||
|
|
||||||
|
|||||||
@@ -4,10 +4,9 @@ Claude SSE 流解析器
|
|||||||
解析 Claude Messages API 的 Server-Sent Events 流。
|
解析 Claude Messages API 的 Server-Sent Events 流。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import json
|
import json
|
||||||
from typing import Any, Dict, List, Optional
|
from typing import Any
|
||||||
|
|
||||||
from src.api.handlers.base.utils import extract_cache_creation_tokens
|
from src.api.handlers.base.utils import extract_cache_creation_tokens
|
||||||
|
|
||||||
@@ -43,7 +42,7 @@ class ClaudeStreamParser:
|
|||||||
DELTA_TEXT = "text_delta"
|
DELTA_TEXT = "text_delta"
|
||||||
DELTA_INPUT_JSON = "input_json_delta"
|
DELTA_INPUT_JSON = "input_json_delta"
|
||||||
|
|
||||||
def parse_chunk(self, chunk: bytes | str) -> List[Dict[str, Any]]:
|
def parse_chunk(self, chunk: bytes | str) -> list[dict[str, Any]]:
|
||||||
"""
|
"""
|
||||||
解析 SSE 数据块
|
解析 SSE 数据块
|
||||||
|
|
||||||
@@ -58,10 +57,10 @@ class ClaudeStreamParser:
|
|||||||
else:
|
else:
|
||||||
text = chunk
|
text = chunk
|
||||||
|
|
||||||
events: List[Dict[str, Any]] = []
|
events: list[dict[str, Any]] = []
|
||||||
lines = text.strip().split("\n")
|
lines = text.strip().split("\n")
|
||||||
|
|
||||||
current_event_type: Optional[str] = None
|
current_event_type: str | None = None
|
||||||
|
|
||||||
for line in lines:
|
for line in lines:
|
||||||
line = line.strip()
|
line = line.strip()
|
||||||
@@ -96,7 +95,7 @@ class ClaudeStreamParser:
|
|||||||
|
|
||||||
return events
|
return events
|
||||||
|
|
||||||
def parse_line(self, line: str) -> Optional[Dict[str, Any]]:
|
def parse_line(self, line: str) -> dict[str, Any] | None:
|
||||||
"""
|
"""
|
||||||
解析单行 SSE 数据
|
解析单行 SSE 数据
|
||||||
|
|
||||||
@@ -117,7 +116,7 @@ class ClaudeStreamParser:
|
|||||||
except json.JSONDecodeError:
|
except json.JSONDecodeError:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def is_done_event(self, event: Dict[str, Any]) -> bool:
|
def is_done_event(self, event: dict[str, Any]) -> bool:
|
||||||
"""
|
"""
|
||||||
判断是否为结束事件
|
判断是否为结束事件
|
||||||
|
|
||||||
@@ -130,7 +129,7 @@ class ClaudeStreamParser:
|
|||||||
event_type = event.get("type")
|
event_type = event.get("type")
|
||||||
return event_type in (self.EVENT_MESSAGE_STOP, "__done__")
|
return event_type in (self.EVENT_MESSAGE_STOP, "__done__")
|
||||||
|
|
||||||
def is_error_event(self, event: Dict[str, Any]) -> bool:
|
def is_error_event(self, event: dict[str, Any]) -> bool:
|
||||||
"""
|
"""
|
||||||
判断是否为错误事件
|
判断是否为错误事件
|
||||||
|
|
||||||
@@ -142,7 +141,7 @@ class ClaudeStreamParser:
|
|||||||
"""
|
"""
|
||||||
return event.get("type") == self.EVENT_ERROR
|
return event.get("type") == self.EVENT_ERROR
|
||||||
|
|
||||||
def get_event_type(self, event: Dict[str, Any]) -> Optional[str]:
|
def get_event_type(self, event: dict[str, Any]) -> str | None:
|
||||||
"""
|
"""
|
||||||
获取事件类型
|
获取事件类型
|
||||||
|
|
||||||
@@ -155,7 +154,7 @@ class ClaudeStreamParser:
|
|||||||
event_type = event.get("type")
|
event_type = event.get("type")
|
||||||
return str(event_type) if event_type is not None else None
|
return str(event_type) if event_type is not None else None
|
||||||
|
|
||||||
def extract_text_delta(self, event: Dict[str, Any]) -> Optional[str]:
|
def extract_text_delta(self, event: dict[str, Any]) -> str | None:
|
||||||
"""
|
"""
|
||||||
从 content_block_delta 事件中提取文本增量
|
从 content_block_delta 事件中提取文本增量
|
||||||
|
|
||||||
@@ -175,7 +174,7 @@ class ClaudeStreamParser:
|
|||||||
|
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def extract_usage(self, event: Dict[str, Any]) -> Optional[Dict[str, int]]:
|
def extract_usage(self, event: dict[str, Any]) -> dict[str, int] | None:
|
||||||
"""
|
"""
|
||||||
从事件中提取 token 使用量
|
从事件中提取 token 使用量
|
||||||
|
|
||||||
@@ -212,7 +211,7 @@ class ClaudeStreamParser:
|
|||||||
|
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def extract_message_id(self, event: Dict[str, Any]) -> Optional[str]:
|
def extract_message_id(self, event: dict[str, Any]) -> str | None:
|
||||||
"""
|
"""
|
||||||
从 message_start 事件中提取消息 ID
|
从 message_start 事件中提取消息 ID
|
||||||
|
|
||||||
@@ -229,7 +228,7 @@ class ClaudeStreamParser:
|
|||||||
msg_id = message.get("id")
|
msg_id = message.get("id")
|
||||||
return str(msg_id) if msg_id is not None else None
|
return str(msg_id) if msg_id is not None else None
|
||||||
|
|
||||||
def extract_stop_reason(self, event: Dict[str, Any]) -> Optional[str]:
|
def extract_stop_reason(self, event: dict[str, Any]) -> str | None:
|
||||||
"""
|
"""
|
||||||
从 message_delta 事件中提取停止原因
|
从 message_delta 事件中提取停止原因
|
||||||
|
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ Claude CLI Adapter - 基于通用 CLI Adapter 基类的简化实现
|
|||||||
继承 CliAdapterBase,只需配置 FORMAT_ID 和 HANDLER_CLASS。
|
继承 CliAdapterBase,只需配置 FORMAT_ID 和 HANDLER_CLASS。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from typing import Any, Dict, Optional, Tuple, Type
|
from typing import Any
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
|
|
||||||
@@ -27,20 +27,20 @@ class ClaudeCliAdapter(CliAdapterBase):
|
|||||||
name = "claude.cli"
|
name = "claude.cli"
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def HANDLER_CLASS(self) -> Type[CliMessageHandlerBase]:
|
def HANDLER_CLASS(self) -> type[CliMessageHandlerBase]:
|
||||||
"""延迟导入 Handler 类避免循环依赖"""
|
"""延迟导入 Handler 类避免循环依赖"""
|
||||||
from src.api.handlers.claude_cli.handler import ClaudeCliMessageHandler
|
from src.api.handlers.claude_cli.handler import ClaudeCliMessageHandler
|
||||||
|
|
||||||
return ClaudeCliMessageHandler
|
return ClaudeCliMessageHandler
|
||||||
|
|
||||||
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
|
def __init__(self, allowed_api_formats: list[str] | None = None):
|
||||||
super().__init__(allowed_api_formats or ["CLAUDE_CLI"])
|
super().__init__(allowed_api_formats or ["CLAUDE_CLI"])
|
||||||
|
|
||||||
def detect_capability_requirements(
|
def detect_capability_requirements(
|
||||||
self,
|
self,
|
||||||
headers: Dict[str, str],
|
headers: dict[str, str],
|
||||||
request_body: Optional[Dict[str, Any]] = None,
|
request_body: dict[str, Any] | None = None,
|
||||||
) -> Dict[str, bool]:
|
) -> dict[str, bool]:
|
||||||
"""检测 Claude CLI 请求中隐含的能力需求"""
|
"""检测 Claude CLI 请求中隐含的能力需求"""
|
||||||
return ClaudeCapabilityDetector.detect_from_headers(headers)
|
return ClaudeCapabilityDetector.detect_from_headers(headers)
|
||||||
|
|
||||||
@@ -61,16 +61,16 @@ class ClaudeCliAdapter(CliAdapterBase):
|
|||||||
"""
|
"""
|
||||||
return input_tokens + cache_creation_input_tokens + cache_read_input_tokens
|
return input_tokens + cache_creation_input_tokens + cache_read_input_tokens
|
||||||
|
|
||||||
def _extract_message_count(self, payload: Dict[str, Any]) -> int:
|
def _extract_message_count(self, payload: dict[str, Any]) -> int:
|
||||||
"""Claude CLI 使用 messages 字段"""
|
"""Claude CLI 使用 messages 字段"""
|
||||||
messages = payload.get("messages", [])
|
messages = payload.get("messages", [])
|
||||||
return len(messages) if isinstance(messages, list) else 0
|
return len(messages) if isinstance(messages, list) else 0
|
||||||
|
|
||||||
def _build_audit_metadata(
|
def _build_audit_metadata(
|
||||||
self,
|
self,
|
||||||
payload: Dict[str, Any],
|
payload: dict[str, Any],
|
||||||
path_params: Optional[Dict[str, Any]] = None, # noqa: ARG002
|
path_params: dict[str, Any] | None = None, # noqa: ARG002
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""Claude CLI 特定的审计元数据"""
|
"""Claude CLI 特定的审计元数据"""
|
||||||
model = payload.get("model", "unknown")
|
model = payload.get("model", "unknown")
|
||||||
stream = payload.get("stream", False)
|
stream = payload.get("stream", False)
|
||||||
@@ -104,8 +104,8 @@ class ClaudeCliAdapter(CliAdapterBase):
|
|||||||
client: httpx.AsyncClient,
|
client: httpx.AsyncClient,
|
||||||
base_url: str,
|
base_url: str,
|
||||||
api_key: str,
|
api_key: str,
|
||||||
extra_headers: Optional[Dict[str, str]] = None,
|
extra_headers: dict[str, str] | None = None,
|
||||||
) -> Tuple[list, Optional[str]]:
|
) -> tuple[list, str | None]:
|
||||||
"""查询 Claude API 支持的模型列表(带 CLI User-Agent)"""
|
"""查询 Claude API 支持的模型列表(带 CLI User-Agent)"""
|
||||||
# 复用 ClaudeChatAdapter 的实现,添加 CLI User-Agent
|
# 复用 ClaudeChatAdapter 的实现,添加 CLI User-Agent
|
||||||
cli_headers = {"User-Agent": config.internal_user_agent_claude_cli}
|
cli_headers = {"User-Agent": config.internal_user_agent_claude_cli}
|
||||||
@@ -120,7 +120,7 @@ class ClaudeCliAdapter(CliAdapterBase):
|
|||||||
return models, error
|
return models, error
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def build_endpoint_url(cls, base_url: str, request_data: Dict[str, Any], model_name: Optional[str] = None) -> str:
|
def build_endpoint_url(cls, base_url: str, request_data: dict[str, Any], model_name: str | None = None) -> str:
|
||||||
"""构建Claude CLI API端点URL"""
|
"""构建Claude CLI API端点URL"""
|
||||||
base_url = base_url.rstrip("/")
|
base_url = base_url.rstrip("/")
|
||||||
if base_url.endswith("/v1"):
|
if base_url.endswith("/v1"):
|
||||||
@@ -131,12 +131,12 @@ class ClaudeCliAdapter(CliAdapterBase):
|
|||||||
# build_request_body 使用基类实现,通过 format_conversion_registry 自动转换 OPENAI -> CLAUDE_CLI
|
# build_request_body 使用基类实现,通过 format_conversion_registry 自动转换 OPENAI -> CLAUDE_CLI
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def get_cli_user_agent(cls) -> Optional[str]:
|
def get_cli_user_agent(cls) -> str | None:
|
||||||
"""获取Claude CLI User-Agent"""
|
"""获取Claude CLI User-Agent"""
|
||||||
return config.internal_user_agent_claude_cli
|
return config.internal_user_agent_claude_cli
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def get_cli_extra_headers(cls) -> Dict[str, str]:
|
def get_cli_extra_headers(cls) -> dict[str, str]:
|
||||||
"""获取Claude CLI额外请求头,包含 x-app: cli 标识"""
|
"""获取Claude CLI额外请求头,包含 x-app: cli 标识"""
|
||||||
headers = super().get_cli_extra_headers()
|
headers = super().get_cli_extra_headers()
|
||||||
headers["x-app"] = "cli" # 标识 CLI 模式,让上游使用正确的认证方式
|
headers["x-app"] = "cli" # 标识 CLI 模式,让上游使用正确的认证方式
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ Claude CLI Message Handler - 基于通用 CLI Handler 基类的简化实现
|
|||||||
继承 CliMessageHandlerBase,只需覆盖格式特定的配置和事件处理逻辑。
|
继承 CliMessageHandlerBase,只需覆盖格式特定的配置和事件处理逻辑。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from typing import Any, Dict, Optional
|
from typing import Any
|
||||||
|
|
||||||
from src.api.handlers.base.cli_handler_base import (
|
from src.api.handlers.base.cli_handler_base import (
|
||||||
CliMessageHandlerBase,
|
CliMessageHandlerBase,
|
||||||
@@ -33,8 +33,8 @@ class ClaudeCliMessageHandler(CliMessageHandlerBase):
|
|||||||
|
|
||||||
def extract_model_from_request(
|
def extract_model_from_request(
|
||||||
self,
|
self,
|
||||||
request_body: Dict[str, Any],
|
request_body: dict[str, Any],
|
||||||
path_params: Optional[Dict[str, Any]] = None, # noqa: ARG002
|
path_params: dict[str, Any] | None = None, # noqa: ARG002
|
||||||
) -> str:
|
) -> str:
|
||||||
"""
|
"""
|
||||||
从请求中提取模型名 - Claude 格式实现
|
从请求中提取模型名 - Claude 格式实现
|
||||||
@@ -53,9 +53,9 @@ class ClaudeCliMessageHandler(CliMessageHandlerBase):
|
|||||||
|
|
||||||
def apply_mapped_model(
|
def apply_mapped_model(
|
||||||
self,
|
self,
|
||||||
request_body: Dict[str, Any],
|
request_body: dict[str, Any],
|
||||||
mapped_model: str,
|
mapped_model: str,
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
Claude API 的 model 在请求体顶级
|
Claude API 的 model 在请求体顶级
|
||||||
|
|
||||||
@@ -74,7 +74,7 @@ class ClaudeCliMessageHandler(CliMessageHandlerBase):
|
|||||||
self,
|
self,
|
||||||
ctx: StreamContext,
|
ctx: StreamContext,
|
||||||
event_type: str,
|
event_type: str,
|
||||||
data: Dict[str, Any],
|
data: dict[str, Any],
|
||||||
) -> None:
|
) -> None:
|
||||||
"""
|
"""
|
||||||
处理 Claude CLI 格式的 SSE 事件
|
处理 Claude CLI 格式的 SSE 事件
|
||||||
@@ -142,8 +142,8 @@ class ClaudeCliMessageHandler(CliMessageHandlerBase):
|
|||||||
|
|
||||||
def _extract_response_metadata(
|
def _extract_response_metadata(
|
||||||
self,
|
self,
|
||||||
response: Dict[str, Any],
|
response: dict[str, Any],
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
从 Claude 响应中提取元数据
|
从 Claude 响应中提取元数据
|
||||||
|
|
||||||
@@ -155,7 +155,7 @@ class ClaudeCliMessageHandler(CliMessageHandlerBase):
|
|||||||
Returns:
|
Returns:
|
||||||
提取的元数据字典
|
提取的元数据字典
|
||||||
"""
|
"""
|
||||||
metadata: Dict[str, Any] = {}
|
metadata: dict[str, Any] = {}
|
||||||
|
|
||||||
# 提取模型名称(实际使用的模型)
|
# 提取模型名称(实际使用的模型)
|
||||||
if "model" in response:
|
if "model" in response:
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ Gemini Chat Adapter
|
|||||||
处理 Gemini API 格式的请求适配
|
处理 Gemini API 格式的请求适配
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from typing import Any, Dict, Optional, Tuple, Type
|
from typing import Any
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
from fastapi import HTTPException, Request
|
from fastapi import HTTPException, Request
|
||||||
@@ -33,17 +33,17 @@ class GeminiChatAdapter(ChatAdapterBase):
|
|||||||
name = "gemini.chat"
|
name = "gemini.chat"
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def HANDLER_CLASS(self) -> Type[ChatHandlerBase]:
|
def HANDLER_CLASS(self) -> type[ChatHandlerBase]:
|
||||||
"""延迟导入 Handler 类避免循环依赖"""
|
"""延迟导入 Handler 类避免循环依赖"""
|
||||||
from src.api.handlers.gemini.handler import GeminiChatHandler
|
from src.api.handlers.gemini.handler import GeminiChatHandler
|
||||||
|
|
||||||
return GeminiChatHandler
|
return GeminiChatHandler
|
||||||
|
|
||||||
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
|
def __init__(self, allowed_api_formats: list[str] | None = None):
|
||||||
super().__init__(allowed_api_formats or ["GEMINI"])
|
super().__init__(allowed_api_formats or ["GEMINI"])
|
||||||
logger.info(f"[{self.name}] 初始化 Gemini Chat 适配器 | API格式: {self.allowed_api_formats}")
|
logger.info(f"[{self.name}] 初始化 Gemini Chat 适配器 | API格式: {self.allowed_api_formats}")
|
||||||
|
|
||||||
def extract_api_key(self, request: Request) -> Optional[str]:
|
def extract_api_key(self, request: Request) -> str | None:
|
||||||
"""
|
"""
|
||||||
从请求中提取 API 密钥 - Gemini 支持 header 和 query 两种方式
|
从请求中提取 API 密钥 - Gemini 支持 header 和 query 两种方式
|
||||||
|
|
||||||
@@ -68,8 +68,8 @@ class GeminiChatAdapter(ChatAdapterBase):
|
|||||||
return {}
|
return {}
|
||||||
|
|
||||||
def _merge_path_params(
|
def _merge_path_params(
|
||||||
self, original_request_body: Dict[str, Any], path_params: Dict[str, Any] # noqa: ARG002
|
self, original_request_body: dict[str, Any], path_params: dict[str, Any] # noqa: ARG002
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
合并 URL 路径参数到请求体 - Gemini 特化版本
|
合并 URL 路径参数到请求体 - Gemini 特化版本
|
||||||
|
|
||||||
@@ -122,14 +122,14 @@ class GeminiChatAdapter(ChatAdapterBase):
|
|||||||
request.stream = is_stream
|
request.stream = is_stream
|
||||||
return request
|
return request
|
||||||
|
|
||||||
def _extract_message_count(self, payload: Dict[str, Any], request_obj) -> int:
|
def _extract_message_count(self, payload: dict[str, Any], request_obj) -> int:
|
||||||
"""提取消息数量"""
|
"""提取消息数量"""
|
||||||
contents = payload.get("contents", [])
|
contents = payload.get("contents", [])
|
||||||
if hasattr(request_obj, "contents"):
|
if hasattr(request_obj, "contents"):
|
||||||
contents = request_obj.contents
|
contents = request_obj.contents
|
||||||
return len(contents) if isinstance(contents, list) else 0
|
return len(contents) if isinstance(contents, list) else 0
|
||||||
|
|
||||||
def _build_audit_metadata(self, payload: Dict[str, Any], request_obj) -> Dict[str, Any]:
|
def _build_audit_metadata(self, payload: dict[str, Any], request_obj) -> dict[str, Any]:
|
||||||
"""构建 Gemini Chat 特定的审计元数据"""
|
"""构建 Gemini Chat 特定的审计元数据"""
|
||||||
role_counts: dict[str, int] = {}
|
role_counts: dict[str, int] = {}
|
||||||
|
|
||||||
@@ -182,8 +182,8 @@ class GeminiChatAdapter(ChatAdapterBase):
|
|||||||
client: httpx.AsyncClient,
|
client: httpx.AsyncClient,
|
||||||
base_url: str,
|
base_url: str,
|
||||||
api_key: str,
|
api_key: str,
|
||||||
extra_headers: Optional[Dict[str, str]] = None,
|
extra_headers: dict[str, str] | None = None,
|
||||||
) -> Tuple[list, Optional[str]]:
|
) -> tuple[list, str | None]:
|
||||||
"""查询 Gemini API 支持的模型列表"""
|
"""查询 Gemini API 支持的模型列表"""
|
||||||
# Gemini 使用 URL 参数传递 key,不需要 headers 中的认证
|
# Gemini 使用 URL 参数传递 key,不需要 headers 中的认证
|
||||||
base_url_clean = base_url.rstrip("/")
|
base_url_clean = base_url.rstrip("/")
|
||||||
@@ -192,7 +192,7 @@ class GeminiChatAdapter(ChatAdapterBase):
|
|||||||
else:
|
else:
|
||||||
models_url = f"{base_url_clean}/v1beta/models?key={api_key}"
|
models_url = f"{base_url_clean}/v1beta/models?key={api_key}"
|
||||||
|
|
||||||
headers: Dict[str, str] = {}
|
headers: dict[str, str] = {}
|
||||||
if extra_headers:
|
if extra_headers:
|
||||||
headers.update(extra_headers)
|
headers.update(extra_headers)
|
||||||
|
|
||||||
@@ -242,16 +242,16 @@ class GeminiChatAdapter(ChatAdapterBase):
|
|||||||
client: httpx.AsyncClient,
|
client: httpx.AsyncClient,
|
||||||
base_url: str,
|
base_url: str,
|
||||||
api_key: str,
|
api_key: str,
|
||||||
request_data: Dict[str, Any],
|
request_data: dict[str, Any],
|
||||||
extra_headers: Optional[Dict[str, str]] = None,
|
extra_headers: dict[str, str] | None = None,
|
||||||
# 用量计算参数
|
# 用量计算参数
|
||||||
db: Optional[Any] = None,
|
db: Any | None = None,
|
||||||
user: Optional[Any] = None,
|
user: Any | None = None,
|
||||||
provider_name: Optional[str] = None,
|
provider_name: str | None = None,
|
||||||
provider_id: Optional[str] = None,
|
provider_id: str | None = None,
|
||||||
api_key_id: Optional[str] = None,
|
api_key_id: str | None = None,
|
||||||
model_name: Optional[str] = None,
|
model_name: str | None = None,
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""测试 Gemini API 模型连接性(非流式)"""
|
"""测试 Gemini API 模型连接性(非流式)"""
|
||||||
# Gemini需要从request_data或model_name参数获取model名称
|
# Gemini需要从request_data或model_name参数获取model名称
|
||||||
effective_model_name = model_name or request_data.get("model", "")
|
effective_model_name = model_name or request_data.get("model", "")
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ Gemini Chat Handler
|
|||||||
处理 Gemini API 格式的请求
|
处理 Gemini API 格式的请求
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from typing import Any, Dict, Optional
|
from typing import Any
|
||||||
|
|
||||||
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
|
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
|
||||||
|
|
||||||
@@ -76,8 +76,8 @@ class GeminiChatHandler(ChatHandlerBase):
|
|||||||
|
|
||||||
def extract_model_from_request(
|
def extract_model_from_request(
|
||||||
self,
|
self,
|
||||||
request_body: Dict[str, Any],
|
request_body: dict[str, Any],
|
||||||
path_params: Optional[Dict[str, Any]] = None,
|
path_params: dict[str, Any] | None = None,
|
||||||
) -> str:
|
) -> str:
|
||||||
"""
|
"""
|
||||||
从请求中提取模型名 - Gemini Chat 格式实现
|
从请求中提取模型名 - Gemini Chat 格式实现
|
||||||
@@ -126,7 +126,7 @@ class GeminiChatHandler(ChatHandlerBase):
|
|||||||
|
|
||||||
return request
|
return request
|
||||||
|
|
||||||
def _extract_usage(self, response: Dict) -> Dict[str, int]:
|
def _extract_usage(self, response: dict) -> dict[str, int]:
|
||||||
"""
|
"""
|
||||||
从 Gemini 响应中提取 token 使用情况
|
从 Gemini 响应中提取 token 使用情况
|
||||||
|
|
||||||
@@ -151,7 +151,7 @@ class GeminiChatHandler(ChatHandlerBase):
|
|||||||
"cache_read_input_tokens": usage.get("cached_tokens", 0),
|
"cache_read_input_tokens": usage.get("cached_tokens", 0),
|
||||||
}
|
}
|
||||||
|
|
||||||
def _normalize_response(self, response: Dict) -> Dict:
|
def _normalize_response(self, response: dict) -> dict:
|
||||||
"""
|
"""
|
||||||
规范化 Gemini 响应
|
规范化 Gemini 响应
|
||||||
|
|
||||||
|
|||||||
@@ -15,7 +15,7 @@ Gemini streamGenerateContent 常见两种返回(与上游/代理实现有关
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import json
|
import json
|
||||||
from typing import Any, Dict, List, Optional, Union
|
from typing import Any
|
||||||
|
|
||||||
|
|
||||||
class GeminiStreamParser:
|
class GeminiStreamParser:
|
||||||
@@ -43,7 +43,7 @@ class GeminiStreamParser:
|
|||||||
self._in_array = False
|
self._in_array = False
|
||||||
self._brace_depth = 0
|
self._brace_depth = 0
|
||||||
|
|
||||||
def parse_chunk(self, chunk: Union[bytes, str]) -> List[Dict[str, Any]]:
|
def parse_chunk(self, chunk: bytes | str) -> list[dict[str, Any]]:
|
||||||
"""
|
"""
|
||||||
解析流式数据块
|
解析流式数据块
|
||||||
|
|
||||||
@@ -58,7 +58,7 @@ class GeminiStreamParser:
|
|||||||
else:
|
else:
|
||||||
text = chunk
|
text = chunk
|
||||||
|
|
||||||
events: List[Dict[str, Any]] = []
|
events: list[dict[str, Any]] = []
|
||||||
|
|
||||||
for char in text:
|
for char in text:
|
||||||
if char == "[" and not self._in_array:
|
if char == "[" and not self._in_array:
|
||||||
@@ -97,7 +97,7 @@ class GeminiStreamParser:
|
|||||||
|
|
||||||
return events
|
return events
|
||||||
|
|
||||||
def parse_line(self, line: str) -> Optional[Dict[str, Any]]:
|
def parse_line(self, line: str) -> dict[str, Any] | None:
|
||||||
"""
|
"""
|
||||||
解析单行 JSON 数据
|
解析单行 JSON 数据
|
||||||
|
|
||||||
@@ -118,7 +118,7 @@ class GeminiStreamParser:
|
|||||||
except json.JSONDecodeError:
|
except json.JSONDecodeError:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def is_done_event(self, event: Dict[str, Any]) -> bool:
|
def is_done_event(self, event: dict[str, Any]) -> bool:
|
||||||
"""
|
"""
|
||||||
判断是否为结束事件
|
判断是否为结束事件
|
||||||
|
|
||||||
@@ -143,7 +143,7 @@ class GeminiStreamParser:
|
|||||||
|
|
||||||
return False
|
return False
|
||||||
|
|
||||||
def is_error_event(self, event: Dict[str, Any]) -> bool:
|
def is_error_event(self, event: dict[str, Any]) -> bool:
|
||||||
"""
|
"""
|
||||||
判断是否为错误事件
|
判断是否为错误事件
|
||||||
|
|
||||||
@@ -171,7 +171,7 @@ class GeminiStreamParser:
|
|||||||
|
|
||||||
return False
|
return False
|
||||||
|
|
||||||
def extract_error_info(self, event: Dict[str, Any]) -> Optional[Dict[str, Any]]:
|
def extract_error_info(self, event: dict[str, Any]) -> dict[str, Any] | None:
|
||||||
"""
|
"""
|
||||||
从事件中提取错误信息
|
从事件中提取错误信息
|
||||||
|
|
||||||
@@ -208,7 +208,7 @@ class GeminiStreamParser:
|
|||||||
|
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def get_finish_reason(self, event: Dict[str, Any]) -> Optional[str]:
|
def get_finish_reason(self, event: dict[str, Any]) -> str | None:
|
||||||
"""
|
"""
|
||||||
获取结束原因
|
获取结束原因
|
||||||
|
|
||||||
@@ -224,7 +224,7 @@ class GeminiStreamParser:
|
|||||||
return str(reason) if reason is not None else None
|
return str(reason) if reason is not None else None
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def extract_text_delta(self, event: Dict[str, Any]) -> Optional[str]:
|
def extract_text_delta(self, event: dict[str, Any]) -> str | None:
|
||||||
"""
|
"""
|
||||||
从响应中提取文本内容
|
从响应中提取文本内容
|
||||||
|
|
||||||
@@ -248,7 +248,7 @@ class GeminiStreamParser:
|
|||||||
|
|
||||||
return "".join(text_parts) if text_parts else None
|
return "".join(text_parts) if text_parts else None
|
||||||
|
|
||||||
def extract_usage(self, event: Dict[str, Any]) -> Optional[Dict[str, int]]:
|
def extract_usage(self, event: dict[str, Any]) -> dict[str, int] | None:
|
||||||
"""
|
"""
|
||||||
从事件中提取 token 使用量
|
从事件中提取 token 使用量
|
||||||
|
|
||||||
@@ -280,7 +280,7 @@ class GeminiStreamParser:
|
|||||||
"cached_tokens": usage_metadata.get("cachedContentTokenCount", 0),
|
"cached_tokens": usage_metadata.get("cachedContentTokenCount", 0),
|
||||||
}
|
}
|
||||||
|
|
||||||
def extract_model_version(self, event: Dict[str, Any]) -> Optional[str]:
|
def extract_model_version(self, event: dict[str, Any]) -> str | None:
|
||||||
"""
|
"""
|
||||||
从响应中提取模型版本
|
从响应中提取模型版本
|
||||||
|
|
||||||
@@ -293,7 +293,7 @@ class GeminiStreamParser:
|
|||||||
version = event.get("modelVersion")
|
version = event.get("modelVersion")
|
||||||
return str(version) if version is not None else None
|
return str(version) if version is not None else None
|
||||||
|
|
||||||
def extract_safety_ratings(self, event: Dict[str, Any]) -> Optional[List[Dict[str, Any]]]:
|
def extract_safety_ratings(self, event: dict[str, Any]) -> list[dict[str, Any]] | None:
|
||||||
"""
|
"""
|
||||||
从响应中提取安全评级
|
从响应中提取安全评级
|
||||||
|
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ Gemini CLI Adapter - 基于通用 CLI Adapter 基类的实现
|
|||||||
继承 CliAdapterBase,处理 Gemini CLI 格式的请求。
|
继承 CliAdapterBase,处理 Gemini CLI 格式的请求。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from typing import Any, Dict, Optional, Tuple, Type
|
from typing import Any
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
from fastapi import Request
|
from fastapi import Request
|
||||||
@@ -29,16 +29,16 @@ class GeminiCliAdapter(CliAdapterBase):
|
|||||||
name = "gemini.cli"
|
name = "gemini.cli"
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def HANDLER_CLASS(self) -> Type[CliMessageHandlerBase]:
|
def HANDLER_CLASS(self) -> type[CliMessageHandlerBase]:
|
||||||
"""延迟导入 Handler 类避免循环依赖"""
|
"""延迟导入 Handler 类避免循环依赖"""
|
||||||
from src.api.handlers.gemini_cli.handler import GeminiCliMessageHandler
|
from src.api.handlers.gemini_cli.handler import GeminiCliMessageHandler
|
||||||
|
|
||||||
return GeminiCliMessageHandler
|
return GeminiCliMessageHandler
|
||||||
|
|
||||||
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
|
def __init__(self, allowed_api_formats: list[str] | None = None):
|
||||||
super().__init__(allowed_api_formats or ["GEMINI_CLI"])
|
super().__init__(allowed_api_formats or ["GEMINI_CLI"])
|
||||||
|
|
||||||
def extract_api_key(self, request: Request) -> Optional[str]:
|
def extract_api_key(self, request: Request) -> str | None:
|
||||||
"""
|
"""
|
||||||
从请求中提取 API 密钥 - Gemini CLI 支持 header 和 query 两种方式
|
从请求中提取 API 密钥 - Gemini CLI 支持 header 和 query 两种方式
|
||||||
|
|
||||||
@@ -53,8 +53,8 @@ class GeminiCliAdapter(CliAdapterBase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def _merge_path_params(
|
def _merge_path_params(
|
||||||
self, original_request_body: Dict[str, Any], path_params: Dict[str, Any] # noqa: ARG002
|
self, original_request_body: dict[str, Any], path_params: dict[str, Any] # noqa: ARG002
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
合并 URL 路径参数到请求体 - Gemini CLI 特化版本
|
合并 URL 路径参数到请求体 - Gemini CLI 特化版本
|
||||||
|
|
||||||
@@ -74,23 +74,23 @@ class GeminiCliAdapter(CliAdapterBase):
|
|||||||
# Gemini: 不合并任何 path_params 到请求体
|
# Gemini: 不合并任何 path_params 到请求体
|
||||||
return original_request_body.copy()
|
return original_request_body.copy()
|
||||||
|
|
||||||
def _extract_message_count(self, payload: Dict[str, Any]) -> int:
|
def _extract_message_count(self, payload: dict[str, Any]) -> int:
|
||||||
"""Gemini CLI 使用 contents 字段"""
|
"""Gemini CLI 使用 contents 字段"""
|
||||||
contents = payload.get("contents", [])
|
contents = payload.get("contents", [])
|
||||||
return len(contents) if isinstance(contents, list) else 0
|
return len(contents) if isinstance(contents, list) else 0
|
||||||
|
|
||||||
def _build_audit_metadata(
|
def _build_audit_metadata(
|
||||||
self,
|
self,
|
||||||
payload: Dict[str, Any],
|
payload: dict[str, Any],
|
||||||
path_params: Optional[Dict[str, Any]] = None,
|
path_params: dict[str, Any] | None = None,
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""Gemini CLI 特定的审计元数据"""
|
"""Gemini CLI 特定的审计元数据"""
|
||||||
# 从 path_params 获取 model(Gemini 请求体不含 model)
|
# 从 path_params 获取 model(Gemini 请求体不含 model)
|
||||||
model = path_params.get("model", "unknown") if path_params else "unknown"
|
model = path_params.get("model", "unknown") if path_params else "unknown"
|
||||||
contents = payload.get("contents", [])
|
contents = payload.get("contents", [])
|
||||||
generation_config = payload.get("generation_config", {}) or {}
|
generation_config = payload.get("generation_config", {}) or {}
|
||||||
|
|
||||||
role_counts: Dict[str, int] = {}
|
role_counts: dict[str, int] = {}
|
||||||
for content in contents:
|
for content in contents:
|
||||||
role = content.get("role", "unknown") if isinstance(content, dict) else "unknown"
|
role = content.get("role", "unknown") if isinstance(content, dict) else "unknown"
|
||||||
role_counts[role] = role_counts.get(role, 0) + 1
|
role_counts[role] = role_counts.get(role, 0) + 1
|
||||||
@@ -120,8 +120,8 @@ class GeminiCliAdapter(CliAdapterBase):
|
|||||||
client: httpx.AsyncClient,
|
client: httpx.AsyncClient,
|
||||||
base_url: str,
|
base_url: str,
|
||||||
api_key: str,
|
api_key: str,
|
||||||
extra_headers: Optional[Dict[str, str]] = None,
|
extra_headers: dict[str, str] | None = None,
|
||||||
) -> Tuple[list, Optional[str]]:
|
) -> tuple[list, str | None]:
|
||||||
"""查询 Gemini API 支持的模型列表(带 CLI User-Agent)"""
|
"""查询 Gemini API 支持的模型列表(带 CLI User-Agent)"""
|
||||||
# 复用 GeminiChatAdapter 的实现,添加 CLI User-Agent
|
# 复用 GeminiChatAdapter 的实现,添加 CLI User-Agent
|
||||||
cli_headers = {"User-Agent": config.internal_user_agent_gemini_cli}
|
cli_headers = {"User-Agent": config.internal_user_agent_gemini_cli}
|
||||||
@@ -136,7 +136,7 @@ class GeminiCliAdapter(CliAdapterBase):
|
|||||||
return models, error
|
return models, error
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def build_endpoint_url(cls, base_url: str, request_data: Dict[str, Any], model_name: Optional[str] = None) -> str:
|
def build_endpoint_url(cls, base_url: str, request_data: dict[str, Any], model_name: str | None = None) -> str:
|
||||||
"""构建Gemini CLI API端点URL"""
|
"""构建Gemini CLI API端点URL"""
|
||||||
effective_model_name = model_name or request_data.get("model", "")
|
effective_model_name = model_name or request_data.get("model", "")
|
||||||
if not effective_model_name:
|
if not effective_model_name:
|
||||||
@@ -152,12 +152,12 @@ class GeminiCliAdapter(CliAdapterBase):
|
|||||||
# build_request_body 使用基类实现,通过 format_conversion_registry 自动转换 OPENAI -> GEMINI_CLI
|
# build_request_body 使用基类实现,通过 format_conversion_registry 自动转换 OPENAI -> GEMINI_CLI
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def get_cli_user_agent(cls) -> Optional[str]:
|
def get_cli_user_agent(cls) -> str | None:
|
||||||
"""获取Gemini CLI User-Agent"""
|
"""获取Gemini CLI User-Agent"""
|
||||||
return config.internal_user_agent_gemini_cli
|
return config.internal_user_agent_gemini_cli
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def get_cli_extra_headers(cls) -> Dict[str, str]:
|
def get_cli_extra_headers(cls) -> dict[str, str]:
|
||||||
"""获取Gemini CLI额外请求头,包含 x-app: cli 标识"""
|
"""获取Gemini CLI额外请求头,包含 x-app: cli 标识"""
|
||||||
headers = super().get_cli_extra_headers()
|
headers = super().get_cli_extra_headers()
|
||||||
headers["x-app"] = "cli" # 标识 CLI 模式,让上游使用正确的 adapter
|
headers["x-app"] = "cli" # 标识 CLI 模式,让上游使用正确的 adapter
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ Gemini CLI Message Handler - 基于通用 CLI Handler 基类的实现
|
|||||||
继承 CliMessageHandlerBase,处理 Gemini CLI API 格式的请求。
|
继承 CliMessageHandlerBase,处理 Gemini CLI API 格式的请求。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from typing import Any, Dict, Optional
|
from typing import Any
|
||||||
|
|
||||||
from src.api.handlers.base.cli_handler_base import (
|
from src.api.handlers.base.cli_handler_base import (
|
||||||
CliMessageHandlerBase,
|
CliMessageHandlerBase,
|
||||||
@@ -34,8 +34,8 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
|
|||||||
|
|
||||||
def extract_model_from_request(
|
def extract_model_from_request(
|
||||||
self,
|
self,
|
||||||
request_body: Dict[str, Any], # noqa: ARG002 - 基类签名要求
|
request_body: dict[str, Any], # noqa: ARG002 - 基类签名要求
|
||||||
path_params: Optional[Dict[str, Any]] = None,
|
path_params: dict[str, Any] | None = None,
|
||||||
) -> str:
|
) -> str:
|
||||||
"""
|
"""
|
||||||
从请求中提取模型名 - Gemini 格式实现
|
从请求中提取模型名 - Gemini 格式实现
|
||||||
@@ -57,8 +57,8 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
|
|||||||
|
|
||||||
def prepare_provider_request_body(
|
def prepare_provider_request_body(
|
||||||
self,
|
self,
|
||||||
request_body: Dict[str, Any],
|
request_body: dict[str, Any],
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
准备发送给 Gemini API 的请求体 - 移除 model 字段
|
准备发送给 Gemini API 的请求体 - 移除 model 字段
|
||||||
|
|
||||||
@@ -77,9 +77,9 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
|
|||||||
|
|
||||||
def get_model_for_url(
|
def get_model_for_url(
|
||||||
self,
|
self,
|
||||||
request_body: Dict[str, Any],
|
request_body: dict[str, Any],
|
||||||
mapped_model: Optional[str],
|
mapped_model: str | None,
|
||||||
) -> Optional[str]:
|
) -> str | None:
|
||||||
"""
|
"""
|
||||||
Gemini 需要将 model 放入 URL 路径中
|
Gemini 需要将 model 放入 URL 路径中
|
||||||
|
|
||||||
@@ -93,7 +93,7 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
|
|||||||
# 优先使用映射后的模型名,否则使用请求体中的
|
# 优先使用映射后的模型名,否则使用请求体中的
|
||||||
return mapped_model or request_body.get("model")
|
return mapped_model or request_body.get("model")
|
||||||
|
|
||||||
def _extract_usage_from_event(self, event: Dict[str, Any]) -> Dict[str, int]:
|
def _extract_usage_from_event(self, event: dict[str, Any]) -> dict[str, int]:
|
||||||
"""
|
"""
|
||||||
从 Gemini 事件中提取 token 使用情况
|
从 Gemini 事件中提取 token 使用情况
|
||||||
|
|
||||||
@@ -126,7 +126,7 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
|
|||||||
self,
|
self,
|
||||||
ctx: StreamContext,
|
ctx: StreamContext,
|
||||||
_event_type: str,
|
_event_type: str,
|
||||||
data: Dict[str, Any],
|
data: dict[str, Any],
|
||||||
) -> None:
|
) -> None:
|
||||||
"""
|
"""
|
||||||
处理 Gemini CLI 格式的流式事件
|
处理 Gemini CLI 格式的流式事件
|
||||||
@@ -190,8 +190,8 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
|
|||||||
|
|
||||||
def _extract_response_metadata(
|
def _extract_response_metadata(
|
||||||
self,
|
self,
|
||||||
response: Dict[str, Any],
|
response: dict[str, Any],
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
从 Gemini 响应中提取元数据
|
从 Gemini 响应中提取元数据
|
||||||
|
|
||||||
@@ -203,7 +203,7 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
|
|||||||
Returns:
|
Returns:
|
||||||
包含 model_version 的元数据字典
|
包含 model_version 的元数据字典
|
||||||
"""
|
"""
|
||||||
metadata: Dict[str, Any] = {}
|
metadata: dict[str, Any] = {}
|
||||||
model_version = response.get("modelVersion")
|
model_version = response.get("modelVersion")
|
||||||
if model_version:
|
if model_version:
|
||||||
metadata["model_version"] = model_version
|
metadata["model_version"] = model_version
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ OpenAI Chat Adapter - 基于 ChatAdapterBase 的 OpenAI Chat API 适配器
|
|||||||
处理 /v1/chat/completions 端点的 OpenAI Chat 格式请求。
|
处理 /v1/chat/completions 端点的 OpenAI Chat 格式请求。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from typing import Any, Dict, Optional, Tuple, Type
|
from typing import Any
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
from fastapi.responses import JSONResponse
|
from fastapi.responses import JSONResponse
|
||||||
@@ -28,13 +28,13 @@ class OpenAIChatAdapter(ChatAdapterBase):
|
|||||||
name = "openai.chat"
|
name = "openai.chat"
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def HANDLER_CLASS(self) -> Type[ChatHandlerBase]:
|
def HANDLER_CLASS(self) -> type[ChatHandlerBase]:
|
||||||
"""延迟导入 Handler 类避免循环依赖"""
|
"""延迟导入 Handler 类避免循环依赖"""
|
||||||
from src.api.handlers.openai.handler import OpenAIChatHandler
|
from src.api.handlers.openai.handler import OpenAIChatHandler
|
||||||
|
|
||||||
return OpenAIChatHandler
|
return OpenAIChatHandler
|
||||||
|
|
||||||
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
|
def __init__(self, allowed_api_formats: list[str] | None = None):
|
||||||
super().__init__(allowed_api_formats or ["OPENAI"])
|
super().__init__(allowed_api_formats or ["OPENAI"])
|
||||||
|
|
||||||
def _validate_request_body(self, original_request_body: dict, path_params: dict = None):
|
def _validate_request_body(self, original_request_body: dict, path_params: dict = None):
|
||||||
@@ -66,7 +66,7 @@ class OpenAIChatAdapter(ChatAdapterBase):
|
|||||||
max_tokens=original_request_body.get("max_tokens"),
|
max_tokens=original_request_body.get("max_tokens"),
|
||||||
)
|
)
|
||||||
|
|
||||||
def _build_audit_metadata(self, payload: Dict[str, Any], request_obj) -> Dict[str, Any]:
|
def _build_audit_metadata(self, payload: dict[str, Any], request_obj) -> dict[str, Any]:
|
||||||
"""构建 OpenAI Chat 特定的审计元数据"""
|
"""构建 OpenAI Chat 特定的审计元数据"""
|
||||||
role_counts = {}
|
role_counts = {}
|
||||||
for message in request_obj.messages:
|
for message in request_obj.messages:
|
||||||
@@ -105,8 +105,8 @@ class OpenAIChatAdapter(ChatAdapterBase):
|
|||||||
client: httpx.AsyncClient,
|
client: httpx.AsyncClient,
|
||||||
base_url: str,
|
base_url: str,
|
||||||
api_key: str,
|
api_key: str,
|
||||||
extra_headers: Optional[Dict[str, str]] = None,
|
extra_headers: dict[str, str] | None = None,
|
||||||
) -> Tuple[list, Optional[str]]:
|
) -> tuple[list, str | None]:
|
||||||
"""查询 OpenAI 兼容 API 支持的模型列表"""
|
"""查询 OpenAI 兼容 API 支持的模型列表"""
|
||||||
headers = cls.build_headers_with_extra(api_key, extra_headers)
|
headers = cls.build_headers_with_extra(api_key, extra_headers)
|
||||||
|
|
||||||
|
|||||||
@@ -5,7 +5,7 @@ OpenAI Chat Handler - 基于通用 Chat Handler 基类的简化实现
|
|||||||
代码量从原来的 ~1315 行减少到 ~100 行。
|
代码量从原来的 ~1315 行减少到 ~100 行。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from typing import Any, Dict, Optional
|
from typing import Any
|
||||||
|
|
||||||
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
|
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
|
||||||
|
|
||||||
@@ -24,8 +24,8 @@ class OpenAIChatHandler(ChatHandlerBase):
|
|||||||
|
|
||||||
def extract_model_from_request(
|
def extract_model_from_request(
|
||||||
self,
|
self,
|
||||||
request_body: Dict[str, Any],
|
request_body: dict[str, Any],
|
||||||
path_params: Optional[Dict[str, Any]] = None, # noqa: ARG002
|
path_params: dict[str, Any] | None = None, # noqa: ARG002
|
||||||
) -> str:
|
) -> str:
|
||||||
"""
|
"""
|
||||||
从请求中提取模型名 - OpenAI 格式实现
|
从请求中提取模型名 - OpenAI 格式实现
|
||||||
@@ -44,9 +44,9 @@ class OpenAIChatHandler(ChatHandlerBase):
|
|||||||
|
|
||||||
def apply_mapped_model(
|
def apply_mapped_model(
|
||||||
self,
|
self,
|
||||||
request_body: Dict[str, Any],
|
request_body: dict[str, Any],
|
||||||
mapped_model: str,
|
mapped_model: str,
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
将映射后的模型名应用到请求体
|
将映射后的模型名应用到请求体
|
||||||
|
|
||||||
@@ -89,7 +89,7 @@ class OpenAIChatHandler(ChatHandlerBase):
|
|||||||
|
|
||||||
return request
|
return request
|
||||||
|
|
||||||
def _extract_usage(self, response: Dict) -> Dict[str, int]:
|
def _extract_usage(self, response: dict) -> dict[str, int]:
|
||||||
"""
|
"""
|
||||||
从 OpenAI 响应中提取 token 使用情况
|
从 OpenAI 响应中提取 token 使用情况
|
||||||
|
|
||||||
@@ -106,7 +106,7 @@ class OpenAIChatHandler(ChatHandlerBase):
|
|||||||
"cache_read_input_tokens": 0,
|
"cache_read_input_tokens": 0,
|
||||||
}
|
}
|
||||||
|
|
||||||
def _normalize_response(self, response: Dict) -> Dict:
|
def _normalize_response(self, response: dict) -> dict:
|
||||||
"""
|
"""
|
||||||
规范化 OpenAI 响应
|
规范化 OpenAI 响应
|
||||||
|
|
||||||
|
|||||||
@@ -4,10 +4,9 @@ OpenAI SSE 流解析器
|
|||||||
解析 OpenAI Chat Completions API 的 Server-Sent Events 流。
|
解析 OpenAI Chat Completions API 的 Server-Sent Events 流。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import json
|
import json
|
||||||
from typing import Any, Dict, List, Optional
|
from typing import Any
|
||||||
|
|
||||||
|
|
||||||
class OpenAIStreamParser:
|
class OpenAIStreamParser:
|
||||||
@@ -23,7 +22,7 @@ class OpenAIStreamParser:
|
|||||||
- 流结束时发送 data: [DONE]
|
- 流结束时发送 data: [DONE]
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def parse_chunk(self, chunk: bytes | str) -> List[Dict[str, Any]]:
|
def parse_chunk(self, chunk: bytes | str) -> list[dict[str, Any]]:
|
||||||
"""
|
"""
|
||||||
解析 SSE 数据块
|
解析 SSE 数据块
|
||||||
|
|
||||||
@@ -38,7 +37,7 @@ class OpenAIStreamParser:
|
|||||||
else:
|
else:
|
||||||
text = chunk
|
text = chunk
|
||||||
|
|
||||||
chunks: List[Dict[str, Any]] = []
|
chunks: list[dict[str, Any]] = []
|
||||||
lines = text.strip().split("\n")
|
lines = text.strip().split("\n")
|
||||||
|
|
||||||
for line in lines:
|
for line in lines:
|
||||||
@@ -64,7 +63,7 @@ class OpenAIStreamParser:
|
|||||||
|
|
||||||
return chunks
|
return chunks
|
||||||
|
|
||||||
def parse_line(self, line: str) -> Optional[Dict[str, Any]]:
|
def parse_line(self, line: str) -> dict[str, Any] | None:
|
||||||
"""
|
"""
|
||||||
解析单行 SSE 数据
|
解析单行 SSE 数据
|
||||||
|
|
||||||
@@ -85,7 +84,7 @@ class OpenAIStreamParser:
|
|||||||
except json.JSONDecodeError:
|
except json.JSONDecodeError:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def is_done_chunk(self, chunk: Dict[str, Any]) -> bool:
|
def is_done_chunk(self, chunk: dict[str, Any]) -> bool:
|
||||||
"""
|
"""
|
||||||
判断是否为结束 chunk
|
判断是否为结束 chunk
|
||||||
|
|
||||||
@@ -107,7 +106,7 @@ class OpenAIStreamParser:
|
|||||||
|
|
||||||
return False
|
return False
|
||||||
|
|
||||||
def get_finish_reason(self, chunk: Dict[str, Any]) -> Optional[str]:
|
def get_finish_reason(self, chunk: dict[str, Any]) -> str | None:
|
||||||
"""
|
"""
|
||||||
获取结束原因
|
获取结束原因
|
||||||
|
|
||||||
@@ -123,7 +122,7 @@ class OpenAIStreamParser:
|
|||||||
return str(reason) if reason is not None else None
|
return str(reason) if reason is not None else None
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def extract_text_delta(self, chunk: Dict[str, Any]) -> Optional[str]:
|
def extract_text_delta(self, chunk: dict[str, Any]) -> str | None:
|
||||||
"""
|
"""
|
||||||
从 chunk 中提取文本增量
|
从 chunk 中提取文本增量
|
||||||
|
|
||||||
@@ -145,7 +144,7 @@ class OpenAIStreamParser:
|
|||||||
|
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def extract_tool_calls_delta(self, chunk: Dict[str, Any]) -> Optional[List[Dict[str, Any]]]:
|
def extract_tool_calls_delta(self, chunk: dict[str, Any]) -> list[dict[str, Any]] | None:
|
||||||
"""
|
"""
|
||||||
从 chunk 中提取工具调用增量
|
从 chunk 中提取工具调用增量
|
||||||
|
|
||||||
@@ -165,7 +164,7 @@ class OpenAIStreamParser:
|
|||||||
return tool_calls
|
return tool_calls
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def extract_role(self, chunk: Dict[str, Any]) -> Optional[str]:
|
def extract_role(self, chunk: dict[str, Any]) -> str | None:
|
||||||
"""
|
"""
|
||||||
从 chunk 中提取角色
|
从 chunk 中提取角色
|
||||||
|
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ OpenAI CLI Adapter - 基于通用 CLI Adapter 基类的简化实现
|
|||||||
继承 CliAdapterBase,只需配置 FORMAT_ID 和 HANDLER_CLASS。
|
继承 CliAdapterBase,只需配置 FORMAT_ID 和 HANDLER_CLASS。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from typing import Any, Dict, Optional, Tuple, Type
|
from typing import Any
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
|
|
||||||
@@ -27,13 +27,13 @@ class OpenAICliAdapter(CliAdapterBase):
|
|||||||
name = "openai.cli"
|
name = "openai.cli"
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def HANDLER_CLASS(self) -> Type[CliMessageHandlerBase]:
|
def HANDLER_CLASS(self) -> type[CliMessageHandlerBase]:
|
||||||
"""延迟导入 Handler 类避免循环依赖"""
|
"""延迟导入 Handler 类避免循环依赖"""
|
||||||
from src.api.handlers.openai_cli.handler import OpenAICliMessageHandler
|
from src.api.handlers.openai_cli.handler import OpenAICliMessageHandler
|
||||||
|
|
||||||
return OpenAICliMessageHandler
|
return OpenAICliMessageHandler
|
||||||
|
|
||||||
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
|
def __init__(self, allowed_api_formats: list[str] | None = None):
|
||||||
super().__init__(allowed_api_formats or ["OPENAI_CLI"])
|
super().__init__(allowed_api_formats or ["OPENAI_CLI"])
|
||||||
|
|
||||||
# =========================================================================
|
# =========================================================================
|
||||||
@@ -46,8 +46,8 @@ class OpenAICliAdapter(CliAdapterBase):
|
|||||||
client: httpx.AsyncClient,
|
client: httpx.AsyncClient,
|
||||||
base_url: str,
|
base_url: str,
|
||||||
api_key: str,
|
api_key: str,
|
||||||
extra_headers: Optional[Dict[str, str]] = None,
|
extra_headers: dict[str, str] | None = None,
|
||||||
) -> Tuple[list, Optional[str]]:
|
) -> tuple[list, str | None]:
|
||||||
"""查询 OpenAI 兼容 API 支持的模型列表(带 CLI User-Agent)"""
|
"""查询 OpenAI 兼容 API 支持的模型列表(带 CLI User-Agent)"""
|
||||||
# 复用 OpenAIChatAdapter 的实现,添加 CLI User-Agent
|
# 复用 OpenAIChatAdapter 的实现,添加 CLI User-Agent
|
||||||
cli_headers = {"User-Agent": config.internal_user_agent_openai_cli}
|
cli_headers = {"User-Agent": config.internal_user_agent_openai_cli}
|
||||||
@@ -62,7 +62,7 @@ class OpenAICliAdapter(CliAdapterBase):
|
|||||||
return models, error
|
return models, error
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def build_endpoint_url(cls, base_url: str, request_data: Dict[str, Any], model_name: Optional[str] = None) -> str:
|
def build_endpoint_url(cls, base_url: str, request_data: dict[str, Any], model_name: str | None = None) -> str:
|
||||||
"""构建OpenAI CLI API端点URL"""
|
"""构建OpenAI CLI API端点URL"""
|
||||||
base_url = base_url.rstrip("/")
|
base_url = base_url.rstrip("/")
|
||||||
if base_url.endswith("/v1"):
|
if base_url.endswith("/v1"):
|
||||||
@@ -74,7 +74,7 @@ class OpenAICliAdapter(CliAdapterBase):
|
|||||||
# OPENAI -> OPENAI_CLI 无转换器,会直接透传原始请求
|
# OPENAI -> OPENAI_CLI 无转换器,会直接透传原始请求
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def get_cli_user_agent(cls) -> Optional[str]:
|
def get_cli_user_agent(cls) -> str | None:
|
||||||
"""获取OpenAI CLI User-Agent"""
|
"""获取OpenAI CLI User-Agent"""
|
||||||
return config.internal_user_agent_openai_cli
|
return config.internal_user_agent_openai_cli
|
||||||
|
|
||||||
|
|||||||
@@ -5,7 +5,7 @@ OpenAI CLI Message Handler - 基于通用 CLI Handler 基类的简化实现
|
|||||||
代码量从原来的 900+ 行减少到 ~100 行。
|
代码量从原来的 900+ 行减少到 ~100 行。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from typing import Any, Dict, Optional
|
from typing import Any
|
||||||
|
|
||||||
from src.api.handlers.base.cli_handler_base import (
|
from src.api.handlers.base.cli_handler_base import (
|
||||||
CliMessageHandlerBase,
|
CliMessageHandlerBase,
|
||||||
@@ -32,8 +32,8 @@ class OpenAICliMessageHandler(CliMessageHandlerBase):
|
|||||||
|
|
||||||
def extract_model_from_request(
|
def extract_model_from_request(
|
||||||
self,
|
self,
|
||||||
request_body: Dict[str, Any],
|
request_body: dict[str, Any],
|
||||||
path_params: Optional[Dict[str, Any]] = None, # noqa: ARG002
|
path_params: dict[str, Any] | None = None, # noqa: ARG002
|
||||||
) -> str:
|
) -> str:
|
||||||
"""
|
"""
|
||||||
从请求中提取模型名 - OpenAI 格式实现
|
从请求中提取模型名 - OpenAI 格式实现
|
||||||
@@ -52,9 +52,9 @@ class OpenAICliMessageHandler(CliMessageHandlerBase):
|
|||||||
|
|
||||||
def apply_mapped_model(
|
def apply_mapped_model(
|
||||||
self,
|
self,
|
||||||
request_body: Dict[str, Any],
|
request_body: dict[str, Any],
|
||||||
mapped_model: str,
|
mapped_model: str,
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
OpenAI CLI (Responses API) 的 model 在请求体顶级字段。
|
OpenAI CLI (Responses API) 的 model 在请求体顶级字段。
|
||||||
|
|
||||||
@@ -73,7 +73,7 @@ class OpenAICliMessageHandler(CliMessageHandlerBase):
|
|||||||
self,
|
self,
|
||||||
ctx: StreamContext,
|
ctx: StreamContext,
|
||||||
event_type: str,
|
event_type: str,
|
||||||
data: Dict[str, Any],
|
data: dict[str, Any],
|
||||||
) -> None:
|
) -> None:
|
||||||
"""
|
"""
|
||||||
处理 OpenAI CLI 格式的 SSE 事件
|
处理 OpenAI CLI 格式的 SSE 事件
|
||||||
@@ -144,8 +144,8 @@ class OpenAICliMessageHandler(CliMessageHandlerBase):
|
|||||||
|
|
||||||
def _extract_response_metadata(
|
def _extract_response_metadata(
|
||||||
self,
|
self,
|
||||||
response: Dict[str, Any],
|
response: dict[str, Any],
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
从 OpenAI 响应中提取元数据
|
从 OpenAI 响应中提取元数据
|
||||||
|
|
||||||
@@ -157,7 +157,7 @@ class OpenAICliMessageHandler(CliMessageHandlerBase):
|
|||||||
Returns:
|
Returns:
|
||||||
提取的元数据字典
|
提取的元数据字典
|
||||||
"""
|
"""
|
||||||
metadata: Dict[str, Any] = {}
|
metadata: dict[str, Any] = {}
|
||||||
|
|
||||||
# 提取模型名称(实际使用的模型)
|
# 提取模型名称(实际使用的模型)
|
||||||
if "model" in response:
|
if "model" in response:
|
||||||
|
|||||||
@@ -2,7 +2,6 @@
|
|||||||
|
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from datetime import datetime, timedelta, timezone
|
from datetime import datetime, timedelta, timezone
|
||||||
from typing import Optional
|
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
@@ -22,7 +21,7 @@ pipeline = ApiRequestPipeline()
|
|||||||
@router.get("/my-audit-logs")
|
@router.get("/my-audit-logs")
|
||||||
async def get_my_audit_logs(
|
async def get_my_audit_logs(
|
||||||
request: Request,
|
request: Request,
|
||||||
event_type: Optional[str] = Query(None, description="事件类型筛选"),
|
event_type: str | None = Query(None, description="事件类型筛选"),
|
||||||
days: int = Query(30, description="查询天数"),
|
days: int = Query(30, description="查询天数"),
|
||||||
limit: int = Query(50, description="返回数量限制"),
|
limit: int = Query(50, description="返回数量限制"),
|
||||||
offset: int = Query(0, ge=0, description="偏移量"),
|
offset: int = Query(0, ge=0, description="偏移量"),
|
||||||
@@ -86,7 +85,7 @@ class AuthenticatedApiAdapter(ApiAdapter):
|
|||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class UserAuditLogsAdapter(AuthenticatedApiAdapter):
|
class UserAuditLogsAdapter(AuthenticatedApiAdapter):
|
||||||
event_type: Optional[str]
|
event_type: str | None
|
||||||
days: int
|
days: int
|
||||||
limit: int
|
limit: int
|
||||||
offset: int
|
offset: int
|
||||||
|
|||||||
@@ -1,8 +1,7 @@
|
|||||||
"""OAuth 管理端点(管理员)。"""
|
"""OAuth 管理端点(管理员)。"""
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from typing import Any, Dict, List, Optional
|
from typing import Any
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, Request
|
from fastapi import APIRouter, Depends, Request
|
||||||
from pydantic import BaseModel, Field, ValidationError
|
from pydantic import BaseModel, Field, ValidationError
|
||||||
@@ -27,24 +26,24 @@ class SupportedOAuthType(BaseModel):
|
|||||||
default_authorization_url: str
|
default_authorization_url: str
|
||||||
default_token_url: str
|
default_token_url: str
|
||||||
default_userinfo_url: str
|
default_userinfo_url: str
|
||||||
default_scopes: List[str]
|
default_scopes: list[str]
|
||||||
|
|
||||||
|
|
||||||
class OAuthProviderUpsertRequest(BaseModel):
|
class OAuthProviderUpsertRequest(BaseModel):
|
||||||
display_name: str = Field(..., min_length=1, max_length=100)
|
display_name: str = Field(..., min_length=1, max_length=100)
|
||||||
client_id: str = Field(..., min_length=1, max_length=255)
|
client_id: str = Field(..., min_length=1, max_length=255)
|
||||||
client_secret: Optional[str] = Field(None, max_length=2048)
|
client_secret: str | None = Field(None, max_length=2048)
|
||||||
|
|
||||||
authorization_url_override: Optional[str] = Field(None, max_length=500)
|
authorization_url_override: str | None = Field(None, max_length=500)
|
||||||
token_url_override: Optional[str] = Field(None, max_length=500)
|
token_url_override: str | None = Field(None, max_length=500)
|
||||||
userinfo_url_override: Optional[str] = Field(None, max_length=500)
|
userinfo_url_override: str | None = Field(None, max_length=500)
|
||||||
scopes: Optional[List[str]] = None
|
scopes: list[str] | None = None
|
||||||
|
|
||||||
redirect_uri: str = Field(..., min_length=1, max_length=500)
|
redirect_uri: str = Field(..., min_length=1, max_length=500)
|
||||||
frontend_callback_url: str = Field(..., min_length=1, max_length=500)
|
frontend_callback_url: str = Field(..., min_length=1, max_length=500)
|
||||||
|
|
||||||
attribute_mapping: Optional[Dict[str, Any]] = None
|
attribute_mapping: dict[str, Any] | None = None
|
||||||
extra_config: Optional[Dict[str, Any]] = None
|
extra_config: dict[str, Any] | None = None
|
||||||
|
|
||||||
is_enabled: bool = False
|
is_enabled: bool = False
|
||||||
force: bool = False
|
force: bool = False
|
||||||
@@ -55,14 +54,14 @@ class OAuthProviderAdminResponse(BaseModel):
|
|||||||
display_name: str
|
display_name: str
|
||||||
client_id: str
|
client_id: str
|
||||||
has_secret: bool
|
has_secret: bool
|
||||||
authorization_url_override: Optional[str] = None
|
authorization_url_override: str | None = None
|
||||||
token_url_override: Optional[str] = None
|
token_url_override: str | None = None
|
||||||
userinfo_url_override: Optional[str] = None
|
userinfo_url_override: str | None = None
|
||||||
scopes: Optional[List[str]] = None
|
scopes: list[str] | None = None
|
||||||
redirect_uri: str
|
redirect_uri: str
|
||||||
frontend_callback_url: str
|
frontend_callback_url: str
|
||||||
attribute_mapping: Optional[Dict[str, Any]] = None
|
attribute_mapping: dict[str, Any] | None = None
|
||||||
extra_config: Optional[Dict[str, Any]] = None
|
extra_config: dict[str, Any] | None = None
|
||||||
is_enabled: bool
|
is_enabled: bool
|
||||||
|
|
||||||
|
|
||||||
@@ -77,19 +76,19 @@ class OAuthProviderTestRequest(BaseModel):
|
|||||||
"""测试请求,使用表单数据而非数据库配置"""
|
"""测试请求,使用表单数据而非数据库配置"""
|
||||||
|
|
||||||
client_id: str = Field(..., min_length=1)
|
client_id: str = Field(..., min_length=1)
|
||||||
client_secret: Optional[str] = None
|
client_secret: str | None = None
|
||||||
authorization_url_override: Optional[str] = None
|
authorization_url_override: str | None = None
|
||||||
token_url_override: Optional[str] = None
|
token_url_override: str | None = None
|
||||||
redirect_uri: str = Field(..., min_length=1)
|
redirect_uri: str = Field(..., min_length=1)
|
||||||
|
|
||||||
|
|
||||||
@router.get("/supported-types", response_model=List[SupportedOAuthType])
|
@router.get("/supported-types", response_model=list[SupportedOAuthType])
|
||||||
async def get_supported_types(request: Request, db: Session = Depends(get_db)) -> Any:
|
async def get_supported_types(request: Request, db: Session = Depends(get_db)) -> Any:
|
||||||
adapter = GetSupportedTypesAdapter()
|
adapter = GetSupportedTypesAdapter()
|
||||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||||
|
|
||||||
|
|
||||||
@router.get("/providers", response_model=List[OAuthProviderAdminResponse])
|
@router.get("/providers", response_model=list[OAuthProviderAdminResponse])
|
||||||
async def list_provider_configs(request: Request, db: Session = Depends(get_db)) -> Any:
|
async def list_provider_configs(request: Request, db: Session = Depends(get_db)) -> Any:
|
||||||
adapter = ListOAuthProviderConfigsAdapter()
|
adapter = ListOAuthProviderConfigsAdapter()
|
||||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||||
|
|||||||
@@ -1,8 +1,7 @@
|
|||||||
"""OAuth 公开端点(无需登录)。"""
|
"""OAuth 公开端点(无需登录)。"""
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from typing import Any, Optional
|
from typing import Any
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, Query, status
|
from fastapi import APIRouter, Depends, Query, status
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
@@ -38,10 +37,10 @@ async def oauth_authorize(provider_type: str, db: Session = Depends(get_db)) ->
|
|||||||
async def oauth_callback(
|
async def oauth_callback(
|
||||||
provider_type: str,
|
provider_type: str,
|
||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
code: Optional[str] = Query(None),
|
code: str | None = Query(None),
|
||||||
state: Optional[str] = Query(None),
|
state: str | None = Query(None),
|
||||||
error: Optional[str] = Query(None),
|
error: str | None = Query(None),
|
||||||
error_description: Optional[str] = Query(None),
|
error_description: str | None = Query(None),
|
||||||
) -> RedirectResponse:
|
) -> RedirectResponse:
|
||||||
"""
|
"""
|
||||||
OAuth 回调端点。
|
OAuth 回调端点。
|
||||||
|
|||||||
@@ -1,8 +1,7 @@
|
|||||||
"""OAuth 用户端点(需登录)。"""
|
"""OAuth 用户端点(需登录)。"""
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from typing import Any, Optional, cast
|
from typing import Any, cast
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
@@ -53,7 +52,7 @@ async def bind_oauth_provider(
|
|||||||
provider_type: str,
|
provider_type: str,
|
||||||
request: Request,
|
request: Request,
|
||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
bind_token: Optional[str] = None,
|
bind_token: str | None = None,
|
||||||
) -> RedirectResponse:
|
) -> RedirectResponse:
|
||||||
"""发起 OAuth 绑定流程,支持通过 bind_token 参数进行安全认证"""
|
"""发起 OAuth 绑定流程,支持通过 bind_token 参数进行安全认证"""
|
||||||
adapter = BindOAuthProviderAdapter(provider_type=provider_type, bind_token=bind_token)
|
adapter = BindOAuthProviderAdapter(provider_type=provider_type, bind_token=bind_token)
|
||||||
@@ -115,10 +114,10 @@ class BindOAuthProviderAdapter(AuthenticatedApiAdapter):
|
|||||||
2. bind_token 参数 (浏览器跳转场景)
|
2. bind_token 参数 (浏览器跳转场景)
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, provider_type: str, bind_token: Optional[str] = None):
|
def __init__(self, provider_type: str, bind_token: str | None = None):
|
||||||
self.provider_type = provider_type
|
self.provider_type = provider_type
|
||||||
self.bind_token = bind_token
|
self.bind_token = bind_token
|
||||||
self._user_from_bind_token: Optional[User] = None
|
self._user_from_bind_token: User | None = None
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def mode(self) -> ApiMode: # type: ignore[override]
|
def mode(self) -> ApiMode: # type: ignore[override]
|
||||||
@@ -136,7 +135,7 @@ class BindOAuthProviderAdapter(AuthenticatedApiAdapter):
|
|||||||
raise HTTPException(status_code=401, detail="未登录")
|
raise HTTPException(status_code=401, detail="未登录")
|
||||||
|
|
||||||
async def handle(self, context: ApiRequestContext) -> RedirectResponse: # type: ignore[override]
|
async def handle(self, context: ApiRequestContext) -> RedirectResponse: # type: ignore[override]
|
||||||
user: Optional[User] = context.user
|
user: User | None = context.user
|
||||||
|
|
||||||
# 如果使用 bind_token,验证并获取用户
|
# 如果使用 bind_token,验证并获取用户
|
||||||
if self.bind_token:
|
if self.bind_token:
|
||||||
|
|||||||
@@ -6,10 +6,9 @@
|
|||||||
from collections import defaultdict
|
from collections import defaultdict
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from datetime import datetime, timedelta, timezone
|
from datetime import datetime, timedelta, timezone
|
||||||
from typing import Dict, List, Optional
|
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, Query, Request
|
from fastapi import APIRouter, Depends, Query, Request
|
||||||
from sqlalchemy import and_, func, or_
|
from sqlalchemy import and_, or_
|
||||||
from sqlalchemy.orm import Session, joinedload
|
from sqlalchemy.orm import Session, joinedload
|
||||||
|
|
||||||
from src.api.base.adapter import ApiAdapter, ApiMode
|
from src.api.base.adapter import ApiAdapter, ApiMode
|
||||||
@@ -41,10 +40,10 @@ router = APIRouter(prefix="/api/public", tags=["System Catalog"])
|
|||||||
pipeline = ApiRequestPipeline()
|
pipeline = ApiRequestPipeline()
|
||||||
|
|
||||||
|
|
||||||
@router.get("/providers", response_model=List[PublicProviderResponse])
|
@router.get("/providers", response_model=list[PublicProviderResponse])
|
||||||
async def get_public_providers(
|
async def get_public_providers(
|
||||||
request: Request,
|
request: Request,
|
||||||
is_active: Optional[bool] = Query(None, description="过滤活跃状态"),
|
is_active: bool | None = Query(None, description="过滤活跃状态"),
|
||||||
skip: int = Query(0, description="跳过记录数"),
|
skip: int = Query(0, description="跳过记录数"),
|
||||||
limit: int = Query(100, description="返回记录数限制"),
|
limit: int = Query(100, description="返回记录数限制"),
|
||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
@@ -77,11 +76,11 @@ async def get_public_providers(
|
|||||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=ApiMode.PUBLIC)
|
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=ApiMode.PUBLIC)
|
||||||
|
|
||||||
|
|
||||||
@router.get("/models", response_model=List[PublicModelResponse])
|
@router.get("/models", response_model=list[PublicModelResponse])
|
||||||
async def get_public_models(
|
async def get_public_models(
|
||||||
request: Request,
|
request: Request,
|
||||||
provider_id: Optional[str] = Query(None, description="提供商ID过滤"),
|
provider_id: str | None = Query(None, description="提供商ID过滤"),
|
||||||
is_active: Optional[bool] = Query(None, description="过滤活跃状态"),
|
is_active: bool | None = Query(None, description="过滤活跃状态"),
|
||||||
skip: int = Query(0, description="跳过记录数"),
|
skip: int = Query(0, description="跳过记录数"),
|
||||||
limit: int = Query(100, description="返回记录数限制"),
|
limit: int = Query(100, description="返回记录数限制"),
|
||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
@@ -145,7 +144,7 @@ async def get_public_stats(request: Request, db: Session = Depends(get_db)):
|
|||||||
async def search_models(
|
async def search_models(
|
||||||
request: Request,
|
request: Request,
|
||||||
q: str = Query(..., description="搜索关键词"),
|
q: str = Query(..., description="搜索关键词"),
|
||||||
provider_id: Optional[int] = Query(None, description="提供商ID过滤"),
|
provider_id: int | None = Query(None, description="提供商ID过滤"),
|
||||||
limit: int = Query(20, description="返回记录数限制"),
|
limit: int = Query(20, description="返回记录数限制"),
|
||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
):
|
):
|
||||||
@@ -234,8 +233,8 @@ async def get_public_global_models(
|
|||||||
request: Request,
|
request: Request,
|
||||||
skip: int = Query(0, ge=0, description="跳过记录数"),
|
skip: int = Query(0, ge=0, description="跳过记录数"),
|
||||||
limit: int = Query(100, ge=1, le=1000, description="返回记录数限制"),
|
limit: int = Query(100, ge=1, le=1000, description="返回记录数限制"),
|
||||||
is_active: Optional[bool] = Query(None, description="过滤活跃状态"),
|
is_active: bool | None = Query(None, description="过滤活跃状态"),
|
||||||
search: Optional[str] = Query(None, description="搜索关键词"),
|
search: str | None = Query(None, description="搜索关键词"),
|
||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
@@ -283,7 +282,7 @@ class PublicApiAdapter(ApiAdapter):
|
|||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class PublicProvidersAdapter(PublicApiAdapter):
|
class PublicProvidersAdapter(PublicApiAdapter):
|
||||||
is_active: Optional[bool]
|
is_active: bool | None
|
||||||
skip: int
|
skip: int
|
||||||
limit: int
|
limit: int
|
||||||
|
|
||||||
@@ -338,8 +337,8 @@ class PublicProvidersAdapter(PublicApiAdapter):
|
|||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class PublicModelsAdapter(PublicApiAdapter):
|
class PublicModelsAdapter(PublicApiAdapter):
|
||||||
provider_id: Optional[str]
|
provider_id: str | None
|
||||||
is_active: Optional[bool]
|
is_active: bool | None
|
||||||
skip: int
|
skip: int
|
||||||
limit: int
|
limit: int
|
||||||
|
|
||||||
@@ -426,7 +425,7 @@ class PublicStatsAdapter(PublicApiAdapter):
|
|||||||
@dataclass
|
@dataclass
|
||||||
class PublicSearchModelsAdapter(PublicApiAdapter):
|
class PublicSearchModelsAdapter(PublicApiAdapter):
|
||||||
query: str
|
query: str
|
||||||
provider_id: Optional[int]
|
provider_id: int | None
|
||||||
limit: int
|
limit: int
|
||||||
|
|
||||||
async def handle(self, context): # type: ignore[override]
|
async def handle(self, context): # type: ignore[override]
|
||||||
@@ -508,7 +507,7 @@ class PublicApiFormatHealthMonitorAdapter(PublicApiAdapter):
|
|||||||
.all()
|
.all()
|
||||||
)
|
)
|
||||||
|
|
||||||
all_formats: List[str] = []
|
all_formats: list[str] = []
|
||||||
for (api_format_enum,) in active_formats:
|
for (api_format_enum,) in active_formats:
|
||||||
api_format = (
|
api_format = (
|
||||||
api_format_enum.value if hasattr(api_format_enum, "value") else str(api_format_enum)
|
api_format_enum.value if hasattr(api_format_enum, "value") else str(api_format_enum)
|
||||||
@@ -525,7 +524,7 @@ class PublicApiFormatHealthMonitorAdapter(PublicApiAdapter):
|
|||||||
)
|
)
|
||||||
.all()
|
.all()
|
||||||
)
|
)
|
||||||
endpoint_map: Dict[str, List[str]] = defaultdict(list)
|
endpoint_map: dict[str, list[str]] = defaultdict(list)
|
||||||
for api_format_enum, endpoint_id in endpoint_rows:
|
for api_format_enum, endpoint_id in endpoint_rows:
|
||||||
api_format = (
|
api_format = (
|
||||||
api_format_enum.value if hasattr(api_format_enum, "value") else str(api_format_enum)
|
api_format_enum.value if hasattr(api_format_enum, "value") else str(api_format_enum)
|
||||||
@@ -551,7 +550,7 @@ class PublicApiFormatHealthMonitorAdapter(PublicApiAdapter):
|
|||||||
.all()
|
.all()
|
||||||
)
|
)
|
||||||
|
|
||||||
grouped_candidates: Dict[str, List[RequestCandidate]] = {}
|
grouped_candidates: dict[str, list[RequestCandidate]] = {}
|
||||||
|
|
||||||
for candidate, api_format_enum in rows:
|
for candidate, api_format_enum in rows:
|
||||||
api_format = (
|
api_format = (
|
||||||
@@ -564,7 +563,7 @@ class PublicApiFormatHealthMonitorAdapter(PublicApiAdapter):
|
|||||||
grouped_candidates[api_format].append(candidate)
|
grouped_candidates[api_format].append(candidate)
|
||||||
|
|
||||||
# 3. 为所有活跃格式生成监控数据
|
# 3. 为所有活跃格式生成监控数据
|
||||||
monitors: List[PublicApiFormatHealthMonitor] = []
|
monitors: list[PublicApiFormatHealthMonitor] = []
|
||||||
for api_format in all_formats:
|
for api_format in all_formats:
|
||||||
candidates = grouped_candidates.get(api_format, [])
|
candidates = grouped_candidates.get(api_format, [])
|
||||||
|
|
||||||
@@ -579,7 +578,7 @@ class PublicApiFormatHealthMonitorAdapter(PublicApiAdapter):
|
|||||||
success_rate = success_count / actual_completed if actual_completed > 0 else 1.0
|
success_rate = success_count / actual_completed if actual_completed > 0 else 1.0
|
||||||
|
|
||||||
# 转换为公开版事件列表(不含敏感信息如 provider_id, key_id)
|
# 转换为公开版事件列表(不含敏感信息如 provider_id, key_id)
|
||||||
events: List[PublicHealthEvent] = []
|
events: list[PublicHealthEvent] = []
|
||||||
for c in candidates:
|
for c in candidates:
|
||||||
event_time = c.finished_at or c.started_at or c.created_at
|
event_time = c.finished_at or c.started_at or c.created_at
|
||||||
events.append(
|
events.append(
|
||||||
@@ -649,8 +648,8 @@ class PublicGlobalModelsAdapter(PublicApiAdapter):
|
|||||||
|
|
||||||
skip: int
|
skip: int
|
||||||
limit: int
|
limit: int
|
||||||
is_active: Optional[bool]
|
is_active: bool | None
|
||||||
search: Optional[str]
|
search: str | None
|
||||||
|
|
||||||
async def handle(self, context): # type: ignore[override]
|
async def handle(self, context): # type: ignore[override]
|
||||||
db = context.db
|
db = context.db
|
||||||
|
|||||||
@@ -7,8 +7,6 @@
|
|||||||
- Authorization: Bearer (bearer) -> OpenAI 格式
|
- Authorization: Bearer (bearer) -> OpenAI 格式
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from typing import Optional, Tuple, Union
|
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, Query, Request
|
from fastapi import APIRouter, Depends, Query, Request
|
||||||
from fastapi.responses import JSONResponse
|
from fastapi.responses import JSONResponse
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
@@ -35,7 +33,6 @@ from src.core.logger import logger
|
|||||||
from src.database import get_db
|
from src.database import get_db
|
||||||
from src.models.database import ApiKey, User
|
from src.models.database import ApiKey, User
|
||||||
from src.services.auth.service import AuthService
|
from src.services.auth.service import AuthService
|
||||||
from src.services.system.config import SystemConfigService
|
|
||||||
|
|
||||||
router = APIRouter(tags=["System Catalog"])
|
router = APIRouter(tags=["System Catalog"])
|
||||||
|
|
||||||
@@ -54,7 +51,7 @@ _ALL_CHAT_FORMATS = [
|
|||||||
|
|
||||||
def _extract_api_key_from_request(
|
def _extract_api_key_from_request(
|
||||||
request: Request, definition: ApiFormatDefinition
|
request: Request, definition: ApiFormatDefinition
|
||||||
) -> Optional[str]:
|
) -> str | None:
|
||||||
"""根据格式定义从请求中提取 API Key"""
|
"""根据格式定义从请求中提取 API Key"""
|
||||||
auth_header = definition.auth_header.lower()
|
auth_header = definition.auth_header.lower()
|
||||||
auth_type = definition.auth_type
|
auth_type = definition.auth_type
|
||||||
@@ -76,7 +73,7 @@ def _extract_api_key_from_request(
|
|||||||
return header_value
|
return header_value
|
||||||
|
|
||||||
|
|
||||||
def _detect_api_format_and_key(request: Request) -> Tuple[str, Optional[str]]:
|
def _detect_api_format_and_key(request: Request) -> tuple[str, str | None]:
|
||||||
"""
|
"""
|
||||||
根据请求头检测 API 格式并提取 API Key
|
根据请求头检测 API 格式并提取 API Key
|
||||||
|
|
||||||
@@ -163,7 +160,7 @@ def _build_empty_list_response(api_format: str) -> dict:
|
|||||||
|
|
||||||
def _filter_formats_by_restrictions(
|
def _filter_formats_by_restrictions(
|
||||||
formats: list[str], restrictions: AccessRestrictions, api_format: str
|
formats: list[str], restrictions: AccessRestrictions, api_format: str
|
||||||
) -> Tuple[list[str], Optional[dict]]:
|
) -> tuple[list[str], dict | None]:
|
||||||
"""
|
"""
|
||||||
根据访问限制过滤 API 格式
|
根据访问限制过滤 API 格式
|
||||||
|
|
||||||
@@ -182,7 +179,7 @@ def _filter_formats_by_restrictions(
|
|||||||
return filtered, None
|
return filtered, None
|
||||||
|
|
||||||
|
|
||||||
def _authenticate(db: Session, api_key: Optional[str]) -> Tuple[Optional[User], Optional[ApiKey]]:
|
def _authenticate(db: Session, api_key: str | None) -> tuple[User | None, ApiKey | None]:
|
||||||
"""
|
"""
|
||||||
认证 API Key
|
认证 API Key
|
||||||
|
|
||||||
@@ -248,8 +245,8 @@ def _build_auth_error_response(api_format: str) -> JSONResponse:
|
|||||||
|
|
||||||
def _build_claude_list_response(
|
def _build_claude_list_response(
|
||||||
models: list[ModelInfo],
|
models: list[ModelInfo],
|
||||||
before_id: Optional[str],
|
before_id: str | None,
|
||||||
after_id: Optional[str],
|
after_id: str | None,
|
||||||
limit: int,
|
limit: int,
|
||||||
) -> dict:
|
) -> dict:
|
||||||
"""构建 Claude 格式的列表响应"""
|
"""构建 Claude 格式的列表响应"""
|
||||||
@@ -309,7 +306,7 @@ def _build_openai_list_response(models: list[ModelInfo]) -> dict:
|
|||||||
def _build_gemini_list_response(
|
def _build_gemini_list_response(
|
||||||
models: list[ModelInfo],
|
models: list[ModelInfo],
|
||||||
page_size: int,
|
page_size: int,
|
||||||
page_token: Optional[str],
|
page_token: str | None,
|
||||||
) -> dict:
|
) -> dict:
|
||||||
"""构建 Gemini 格式的列表响应"""
|
"""构建 Gemini 格式的列表响应"""
|
||||||
# 处理分页
|
# 处理分页
|
||||||
@@ -435,14 +432,14 @@ def _build_404_response(model_id: str, api_format: str) -> JSONResponse:
|
|||||||
async def list_models(
|
async def list_models(
|
||||||
request: Request,
|
request: Request,
|
||||||
# Claude 分页参数
|
# Claude 分页参数
|
||||||
before_id: Optional[str] = Query(None, description="返回此 ID 之前的结果 (Claude)"),
|
before_id: str | None = Query(None, description="返回此 ID 之前的结果 (Claude)"),
|
||||||
after_id: Optional[str] = Query(None, description="返回此 ID 之后的结果 (Claude)"),
|
after_id: str | None = Query(None, description="返回此 ID 之后的结果 (Claude)"),
|
||||||
limit: int = Query(20, ge=1, le=1000, description="返回数量限制 (Claude)"),
|
limit: int = Query(20, ge=1, le=1000, description="返回数量限制 (Claude)"),
|
||||||
# Gemini 分页参数
|
# Gemini 分页参数
|
||||||
page_size: int = Query(50, alias="pageSize", ge=1, le=1000, description="每页数量 (Gemini)"),
|
page_size: int = Query(50, alias="pageSize", ge=1, le=1000, description="每页数量 (Gemini)"),
|
||||||
page_token: Optional[str] = Query(None, alias="pageToken", description="分页 token (Gemini)"),
|
page_token: str | None = Query(None, alias="pageToken", description="分页 token (Gemini)"),
|
||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
) -> Union[dict, JSONResponse]:
|
) -> dict | JSONResponse:
|
||||||
"""
|
"""
|
||||||
列出可用模型(统一端点)
|
列出可用模型(统一端点)
|
||||||
|
|
||||||
@@ -556,7 +553,7 @@ async def retrieve_model(
|
|||||||
model_id: str,
|
model_id: str,
|
||||||
request: Request,
|
request: Request,
|
||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
) -> Union[dict, JSONResponse]:
|
) -> dict | JSONResponse:
|
||||||
"""
|
"""
|
||||||
获取单个模型详情(统一端点)
|
获取单个模型详情(统一端点)
|
||||||
|
|
||||||
@@ -658,9 +655,9 @@ async def retrieve_model(
|
|||||||
async def list_models_gemini(
|
async def list_models_gemini(
|
||||||
request: Request,
|
request: Request,
|
||||||
page_size: int = Query(50, alias="pageSize", ge=1, le=1000),
|
page_size: int = Query(50, alias="pageSize", ge=1, le=1000),
|
||||||
page_token: Optional[str] = Query(None, alias="pageToken"),
|
page_token: str | None = Query(None, alias="pageToken"),
|
||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
) -> Union[dict, JSONResponse]:
|
) -> dict | JSONResponse:
|
||||||
"""
|
"""
|
||||||
列出可用模型(Gemini v1beta 专用端点)
|
列出可用模型(Gemini v1beta 专用端点)
|
||||||
|
|
||||||
@@ -741,7 +738,7 @@ async def get_model_gemini(
|
|||||||
request: Request,
|
request: Request,
|
||||||
model_name: str,
|
model_name: str,
|
||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
) -> Union[dict, JSONResponse]:
|
) -> dict | JSONResponse:
|
||||||
"""
|
"""
|
||||||
获取单个模型详情(Gemini v1beta 专用端点)
|
获取单个模型详情(Gemini v1beta 专用端点)
|
||||||
|
|
||||||
|
|||||||
@@ -1,12 +1,11 @@
|
|||||||
"""公开模块状态 API(供登录页等使用)"""
|
"""公开模块状态 API(供登录页等使用)"""
|
||||||
|
|
||||||
from typing import List
|
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends
|
from fastapi import APIRouter, Depends
|
||||||
from pydantic import BaseModel
|
from pydantic import BaseModel
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
from src.core.modules import ModuleCategory, get_module_registry
|
from src.core.modules import get_module_registry
|
||||||
from src.database import get_db
|
from src.database import get_db
|
||||||
|
|
||||||
router = APIRouter(prefix="/api/modules", tags=["Modules"])
|
router = APIRouter(prefix="/api/modules", tags=["Modules"])
|
||||||
@@ -20,7 +19,7 @@ class AuthModuleInfo(BaseModel):
|
|||||||
active: bool
|
active: bool
|
||||||
|
|
||||||
|
|
||||||
@router.get("/auth-status", response_model=List[AuthModuleInfo])
|
@router.get("/auth-status", response_model=list[AuthModuleInfo])
|
||||||
async def get_auth_modules_status(db: Session = Depends(get_db)):
|
async def get_auth_modules_status(db: Session = Depends(get_db)):
|
||||||
"""
|
"""
|
||||||
获取认证模块状态(公开接口)
|
获取认证模块状态(公开接口)
|
||||||
|
|||||||
@@ -5,7 +5,7 @@ System Catalog / 健康检查相关端点
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
from datetime import datetime, timezone
|
from datetime import datetime, timezone
|
||||||
from typing import Any, Dict, Optional
|
from typing import Any
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||||
@@ -28,7 +28,7 @@ router = APIRouter(tags=["System Catalog"])
|
|||||||
# ============== 辅助函数 ==============
|
# ============== 辅助函数 ==============
|
||||||
|
|
||||||
|
|
||||||
def _as_bool(value: Optional[str], default: bool) -> bool:
|
def _as_bool(value: str | None, default: bool) -> bool:
|
||||||
"""将字符串转换为布尔值"""
|
"""将字符串转换为布尔值"""
|
||||||
if value is None:
|
if value is None:
|
||||||
return default
|
return default
|
||||||
@@ -39,9 +39,9 @@ def _serialize_provider(
|
|||||||
provider: Provider,
|
provider: Provider,
|
||||||
include_models: bool,
|
include_models: bool,
|
||||||
include_endpoints: bool,
|
include_endpoints: bool,
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""序列化 Provider 对象"""
|
"""序列化 Provider 对象"""
|
||||||
provider_data: Dict[str, Any] = {
|
provider_data: dict[str, Any] = {
|
||||||
"id": provider.id,
|
"id": provider.id,
|
||||||
"name": provider.name,
|
"name": provider.name,
|
||||||
"is_active": provider.is_active,
|
"is_active": provider.is_active,
|
||||||
@@ -81,7 +81,7 @@ def _serialize_provider(
|
|||||||
return provider_data
|
return provider_data
|
||||||
|
|
||||||
|
|
||||||
def _select_provider(db: Session, provider_name: Optional[str]) -> Optional[Provider]:
|
def _select_provider(db: Session, provider_name: str | None) -> Provider | None:
|
||||||
"""选择 Provider(按 provider_priority 优先级选择)"""
|
"""选择 Provider(按 provider_priority 优先级选择)"""
|
||||||
query = db.query(Provider).filter(Provider.is_active == True)
|
query = db.query(Provider).filter(Provider.is_active == True)
|
||||||
if provider_name:
|
if provider_name:
|
||||||
@@ -104,7 +104,7 @@ async def service_health(db: Session = Depends(get_db)):
|
|||||||
)
|
)
|
||||||
active_models = db.query(func.count(Model.id)).filter(Model.is_active == True).scalar() or 0
|
active_models = db.query(func.count(Model.id)).filter(Model.is_active == True).scalar() or 0
|
||||||
|
|
||||||
redis_info: Dict[str, Any] = {"status": "unknown"}
|
redis_info: dict[str, Any] = {"status": "unknown"}
|
||||||
try:
|
try:
|
||||||
redis = await get_redis_client()
|
redis = await get_redis_client()
|
||||||
if redis:
|
if redis:
|
||||||
@@ -245,9 +245,9 @@ async def provider_detail(
|
|||||||
async def test_connection(
|
async def test_connection(
|
||||||
request: Request,
|
request: Request,
|
||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
provider: Optional[str] = Query(None),
|
provider: str | None = Query(None),
|
||||||
model: str = Query("claude-3-haiku-20240307"),
|
model: str = Query("claude-3-haiku-20240307"),
|
||||||
api_format: Optional[str] = Query(None),
|
api_format: str | None = Query(None),
|
||||||
):
|
):
|
||||||
"""测试 Provider 连接"""
|
"""测试 Provider 连接"""
|
||||||
selected_provider = _select_provider(db, provider)
|
selected_provider = _select_provider(db, provider)
|
||||||
|
|||||||
@@ -2,7 +2,6 @@
|
|||||||
|
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
from typing import Optional
|
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||||
from fastapi.responses import JSONResponse
|
from fastapi.responses import JSONResponse
|
||||||
@@ -55,13 +54,13 @@ class CreateManagementTokenRequest(BaseModel):
|
|||||||
"""创建 Management Token 请求"""
|
"""创建 Management Token 请求"""
|
||||||
|
|
||||||
name: str = Field(..., min_length=1, max_length=100, description="Token 名称")
|
name: str = Field(..., min_length=1, max_length=100, description="Token 名称")
|
||||||
description: Optional[str] = Field(None, max_length=500, description="描述")
|
description: str | None = Field(None, max_length=500, description="描述")
|
||||||
allowed_ips: Optional[list[str]] = Field(None, description="IP 白名单")
|
allowed_ips: list[str] | None = Field(None, description="IP 白名单")
|
||||||
expires_at: Optional[datetime] = Field(None, description="过期时间")
|
expires_at: datetime | None = Field(None, description="过期时间")
|
||||||
|
|
||||||
@field_validator("allowed_ips")
|
@field_validator("allowed_ips")
|
||||||
@classmethod
|
@classmethod
|
||||||
def validate_allowed_ips(cls, v: Optional[list[str]]) -> Optional[list[str]]:
|
def validate_allowed_ips(cls, v: list[str] | None) -> list[str] | None:
|
||||||
return validate_ip_list(v)
|
return validate_ip_list(v)
|
||||||
|
|
||||||
@field_validator("expires_at", mode="before")
|
@field_validator("expires_at", mode="before")
|
||||||
@@ -81,10 +80,10 @@ class UpdateManagementTokenRequest(BaseModel):
|
|||||||
|
|
||||||
model_config = {"extra": "allow"} # 允许额外字段以便检测哪些字段被显式提供
|
model_config = {"extra": "allow"} # 允许额外字段以便检测哪些字段被显式提供
|
||||||
|
|
||||||
name: Optional[str] = Field(None, min_length=1, max_length=100)
|
name: str | None = Field(None, min_length=1, max_length=100)
|
||||||
description: Optional[str] = Field(None, max_length=500)
|
description: str | None = Field(None, max_length=500)
|
||||||
allowed_ips: Optional[list[str]] = None
|
allowed_ips: list[str] | None = None
|
||||||
expires_at: Optional[datetime] = None
|
expires_at: datetime | None = None
|
||||||
|
|
||||||
# 用于追踪哪些字段被显式提供(包括显式设为 null 的情况)
|
# 用于追踪哪些字段被显式提供(包括显式设为 null 的情况)
|
||||||
_provided_fields: set[str] = set()
|
_provided_fields: set[str] = set()
|
||||||
@@ -101,7 +100,7 @@ class UpdateManagementTokenRequest(BaseModel):
|
|||||||
|
|
||||||
@field_validator("allowed_ips")
|
@field_validator("allowed_ips")
|
||||||
@classmethod
|
@classmethod
|
||||||
def validate_allowed_ips(cls, v: Optional[list[str]]) -> Optional[list[str]]:
|
def validate_allowed_ips(cls, v: list[str] | None) -> list[str] | None:
|
||||||
# 如果是 None,表示要清空,直接返回
|
# 如果是 None,表示要清空,直接返回
|
||||||
if v is None:
|
if v is None:
|
||||||
return None
|
return None
|
||||||
@@ -122,7 +121,7 @@ class UpdateManagementTokenRequest(BaseModel):
|
|||||||
@router.get("")
|
@router.get("")
|
||||||
async def list_my_management_tokens(
|
async def list_my_management_tokens(
|
||||||
request: Request,
|
request: Request,
|
||||||
is_active: Optional[bool] = Query(None, description="筛选激活状态"),
|
is_active: bool | None = Query(None, description="筛选激活状态"),
|
||||||
skip: int = Query(0, ge=0),
|
skip: int = Query(0, ge=0),
|
||||||
limit: int = Query(50, ge=1, le=100),
|
limit: int = Query(50, ge=1, le=100),
|
||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
@@ -347,7 +346,7 @@ class ListMyManagementTokensAdapter(ManagementTokenApiAdapter):
|
|||||||
"""列出用户的 Management Tokens"""
|
"""列出用户的 Management Tokens"""
|
||||||
|
|
||||||
name: str = "list_my_management_tokens"
|
name: str = "list_my_management_tokens"
|
||||||
is_active: Optional[bool] = None
|
is_active: bool | None = None
|
||||||
skip: int = 0
|
skip: int = 0
|
||||||
limit: int = 50
|
limit: int = 50
|
||||||
|
|
||||||
|
|||||||
@@ -2,14 +2,12 @@
|
|||||||
|
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from datetime import datetime, timezone
|
from datetime import datetime, timezone
|
||||||
from typing import Optional
|
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||||
from pydantic import ValidationError
|
from pydantic import ValidationError
|
||||||
from sqlalchemy import and_, func
|
from sqlalchemy import and_, func
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
from src.api.base.adapter import ApiAdapter, ApiMode
|
|
||||||
from src.api.base.authenticated_adapter import AuthenticatedApiAdapter
|
from src.api.base.authenticated_adapter import AuthenticatedApiAdapter
|
||||||
from src.api.base.pipeline import ApiRequestPipeline
|
from src.api.base.pipeline import ApiRequestPipeline
|
||||||
from src.core.crypto import crypto_service
|
from src.core.crypto import crypto_service
|
||||||
@@ -170,9 +168,9 @@ async def toggle_my_api_key(key_id: str, request: Request, db: Session = Depends
|
|||||||
@router.get("/usage")
|
@router.get("/usage")
|
||||||
async def get_my_usage(
|
async def get_my_usage(
|
||||||
request: Request,
|
request: Request,
|
||||||
start_date: Optional[datetime] = Query(None, description="开始时间(ISO 格式)"),
|
start_date: datetime | None = Query(None, description="开始时间(ISO 格式)"),
|
||||||
end_date: Optional[datetime] = Query(None, description="结束时间(ISO 格式)"),
|
end_date: datetime | None = Query(None, description="结束时间(ISO 格式)"),
|
||||||
search: Optional[str] = Query(None, description="搜索关键词(密钥名、模型名)"),
|
search: str | None = Query(None, description="搜索关键词(密钥名、模型名)"),
|
||||||
limit: int = Query(100, ge=1, le=200, description="每页记录数,默认100,最大200"),
|
limit: int = Query(100, ge=1, le=200, description="每页记录数,默认100,最大200"),
|
||||||
offset: int = Query(0, ge=0, le=2000, description="偏移量,用于分页,最大2000"),
|
offset: int = Query(0, ge=0, le=2000, description="偏移量,用于分页,最大2000"),
|
||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
@@ -200,7 +198,7 @@ async def get_my_usage(
|
|||||||
@router.get("/usage/active")
|
@router.get("/usage/active")
|
||||||
async def get_my_active_requests(
|
async def get_my_active_requests(
|
||||||
request: Request,
|
request: Request,
|
||||||
ids: Optional[str] = Query(None, description="请求 ID 列表,逗号分隔"),
|
ids: str | None = Query(None, description="请求 ID 列表,逗号分隔"),
|
||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
@@ -268,7 +266,7 @@ async def list_available_models(
|
|||||||
request: Request,
|
request: Request,
|
||||||
skip: int = Query(0, ge=0, description="跳过记录数"),
|
skip: int = Query(0, ge=0, description="跳过记录数"),
|
||||||
limit: int = Query(100, ge=1, le=1000, description="返回记录数限制"),
|
limit: int = Query(100, ge=1, le=1000, description="返回记录数限制"),
|
||||||
search: Optional[str] = Query(None, description="搜索关键词"),
|
search: str | None = Query(None, description="搜索关键词"),
|
||||||
db: Session = Depends(get_db),
|
db: Session = Depends(get_db),
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
@@ -721,9 +719,9 @@ class ToggleMyApiKeyAdapter(AuthenticatedApiAdapter):
|
|||||||
class GetUsageAdapter(AuthenticatedApiAdapter):
|
class GetUsageAdapter(AuthenticatedApiAdapter):
|
||||||
"""获取用户使用统计的适配器"""
|
"""获取用户使用统计的适配器"""
|
||||||
|
|
||||||
start_date: Optional[datetime]
|
start_date: datetime | None
|
||||||
end_date: Optional[datetime]
|
end_date: datetime | None
|
||||||
search: Optional[str] = None
|
search: str | None = None
|
||||||
limit: int = 100
|
limit: int = 100
|
||||||
offset: int = 0
|
offset: int = 0
|
||||||
|
|
||||||
@@ -983,7 +981,7 @@ class GetUsageAdapter(AuthenticatedApiAdapter):
|
|||||||
class GetActiveRequestsAdapter(AuthenticatedApiAdapter):
|
class GetActiveRequestsAdapter(AuthenticatedApiAdapter):
|
||||||
"""轻量级活跃请求状态查询适配器(用于用户端轮询)"""
|
"""轻量级活跃请求状态查询适配器(用于用户端轮询)"""
|
||||||
|
|
||||||
ids: Optional[str] = None
|
ids: str | None = None
|
||||||
|
|
||||||
async def handle(self, context): # type: ignore[override]
|
async def handle(self, context): # type: ignore[override]
|
||||||
from src.services.usage import UsageService
|
from src.services.usage import UsageService
|
||||||
@@ -1045,7 +1043,7 @@ class ListAvailableModelsAdapter(AuthenticatedApiAdapter):
|
|||||||
|
|
||||||
skip: int
|
skip: int
|
||||||
limit: int
|
limit: int
|
||||||
search: Optional[str]
|
search: str | None
|
||||||
|
|
||||||
async def handle(self, context): # type: ignore[override]
|
async def handle(self, context): # type: ignore[override]
|
||||||
from sqlalchemy import or_
|
from sqlalchemy import or_
|
||||||
@@ -1220,7 +1218,6 @@ class ListAvailableProvidersAdapter(AuthenticatedApiAdapter):
|
|||||||
async def handle(self, context): # type: ignore[override]
|
async def handle(self, context): # type: ignore[override]
|
||||||
from sqlalchemy.orm import selectinload
|
from sqlalchemy.orm import selectinload
|
||||||
|
|
||||||
from src.models.database import ProviderEndpoint
|
|
||||||
|
|
||||||
db = context.db
|
db = context.db
|
||||||
|
|
||||||
|
|||||||
@@ -8,11 +8,13 @@
|
|||||||
3. 连接池复用:Keep-alive 连接减少 TCP 握手开销
|
3. 连接池复用:Keep-alive 连接减少 TCP 握手开销
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import hashlib
|
import hashlib
|
||||||
import time
|
import time
|
||||||
from contextlib import asynccontextmanager
|
from contextlib import asynccontextmanager
|
||||||
from typing import Any, Dict, Optional, Tuple
|
from typing import Any
|
||||||
from urllib.parse import quote, urlparse
|
from urllib.parse import quote, urlparse
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
@@ -26,7 +28,7 @@ _proxy_clients_lock = asyncio.Lock()
|
|||||||
_default_client_lock = asyncio.Lock()
|
_default_client_lock = asyncio.Lock()
|
||||||
|
|
||||||
|
|
||||||
def _compute_proxy_cache_key(proxy_config: Optional[Dict[str, Any]]) -> str:
|
def _compute_proxy_cache_key(proxy_config: dict[str, Any] | None) -> str:
|
||||||
"""
|
"""
|
||||||
计算代理配置的缓存键
|
计算代理配置的缓存键
|
||||||
|
|
||||||
@@ -48,7 +50,7 @@ def _compute_proxy_cache_key(proxy_config: Optional[Dict[str, Any]]) -> str:
|
|||||||
return f"proxy:{hashlib.md5(proxy_url.encode()).hexdigest()[:16]}"
|
return f"proxy:{hashlib.md5(proxy_url.encode()).hexdigest()[:16]}"
|
||||||
|
|
||||||
|
|
||||||
def build_proxy_url(proxy_config: Dict[str, Any]) -> Optional[str]:
|
def build_proxy_url(proxy_config: dict[str, Any]) -> str | None:
|
||||||
"""
|
"""
|
||||||
根据代理配置构建完整的代理 URL
|
根据代理配置构建完整的代理 URL
|
||||||
|
|
||||||
@@ -103,11 +105,11 @@ class HTTPClientPool:
|
|||||||
3. LRU 淘汰:代理客户端超过上限时淘汰最久未使用的
|
3. LRU 淘汰:代理客户端超过上限时淘汰最久未使用的
|
||||||
"""
|
"""
|
||||||
|
|
||||||
_instance: Optional["HTTPClientPool"] = None
|
_instance: HTTPClientPool | None = None
|
||||||
_default_client: Optional[httpx.AsyncClient] = None
|
_default_client: httpx.AsyncClient | None = None
|
||||||
_clients: Dict[str, httpx.AsyncClient] = {}
|
_clients: dict[str, httpx.AsyncClient] = {}
|
||||||
# 代理客户端缓存:{cache_key: (client, last_used_time)}
|
# 代理客户端缓存:{cache_key: (client, last_used_time)}
|
||||||
_proxy_clients: Dict[str, Tuple[httpx.AsyncClient, float]] = {}
|
_proxy_clients: dict[str, tuple[httpx.AsyncClient, float]] = {}
|
||||||
# 代理客户端缓存上限(避免内存泄漏)
|
# 代理客户端缓存上限(避免内存泄漏)
|
||||||
_max_proxy_clients: int = 50
|
_max_proxy_clients: int = 50
|
||||||
|
|
||||||
@@ -242,7 +244,7 @@ class HTTPClientPool:
|
|||||||
@classmethod
|
@classmethod
|
||||||
async def get_proxy_client(
|
async def get_proxy_client(
|
||||||
cls,
|
cls,
|
||||||
proxy_config: Optional[Dict[str, Any]] = None,
|
proxy_config: dict[str, Any] | None = None,
|
||||||
) -> httpx.AsyncClient:
|
) -> httpx.AsyncClient:
|
||||||
"""
|
"""
|
||||||
获取代理客户端(带缓存复用)
|
获取代理客户端(带缓存复用)
|
||||||
@@ -280,7 +282,7 @@ class HTTPClientPool:
|
|||||||
await cls._evict_lru_proxy_client()
|
await cls._evict_lru_proxy_client()
|
||||||
|
|
||||||
# 创建新客户端(使用默认超时,请求时可覆盖)
|
# 创建新客户端(使用默认超时,请求时可覆盖)
|
||||||
client_config: Dict[str, Any] = {
|
client_config: dict[str, Any] = {
|
||||||
"http2": False,
|
"http2": False,
|
||||||
"verify": get_ssl_context(),
|
"verify": get_ssl_context(),
|
||||||
"follow_redirects": True,
|
"follow_redirects": True,
|
||||||
@@ -370,8 +372,8 @@ class HTTPClientPool:
|
|||||||
@classmethod
|
@classmethod
|
||||||
def create_client_with_proxy(
|
def create_client_with_proxy(
|
||||||
cls,
|
cls,
|
||||||
proxy_config: Optional[Dict[str, Any]] = None,
|
proxy_config: dict[str, Any] | None = None,
|
||||||
timeout: Optional[httpx.Timeout] = None,
|
timeout: httpx.Timeout | None = None,
|
||||||
**kwargs: Any,
|
**kwargs: Any,
|
||||||
) -> httpx.AsyncClient:
|
) -> httpx.AsyncClient:
|
||||||
"""
|
"""
|
||||||
@@ -387,7 +389,7 @@ class HTTPClientPool:
|
|||||||
Returns:
|
Returns:
|
||||||
配置好的 httpx.AsyncClient 实例(调用者需要负责关闭)
|
配置好的 httpx.AsyncClient 实例(调用者需要负责关闭)
|
||||||
"""
|
"""
|
||||||
client_config: Dict[str, Any] = {
|
client_config: dict[str, Any] = {
|
||||||
"http2": False,
|
"http2": False,
|
||||||
"verify": get_ssl_context(),
|
"verify": get_ssl_context(),
|
||||||
"follow_redirects": True,
|
"follow_redirects": True,
|
||||||
@@ -413,7 +415,7 @@ class HTTPClientPool:
|
|||||||
return httpx.AsyncClient(**client_config)
|
return httpx.AsyncClient(**client_config)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def get_pool_stats(cls) -> Dict[str, Any]:
|
def get_pool_stats(cls) -> dict[str, Any]:
|
||||||
"""获取连接池统计信息"""
|
"""获取连接池统计信息"""
|
||||||
return {
|
return {
|
||||||
"default_client_active": cls._default_client is not None,
|
"default_client_active": cls._default_client is not None,
|
||||||
|
|||||||
@@ -9,10 +9,11 @@
|
|||||||
- 调用方可以根据状态决定降级策略
|
- 调用方可以根据状态决定降级策略
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
import os
|
import os
|
||||||
import time
|
import time
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
from typing import Optional
|
|
||||||
|
|
||||||
import redis.asyncio as aioredis
|
import redis.asyncio as aioredis
|
||||||
from src.core.logger import logger
|
from src.core.logger import logger
|
||||||
@@ -35,8 +36,8 @@ class RedisClientManager:
|
|||||||
提供 Redis 连接管理、熔断器保护和状态监控。
|
提供 Redis 连接管理、熔断器保护和状态监控。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
_instance: Optional["RedisClientManager"] = None
|
_instance: RedisClientManager | None = None
|
||||||
_redis: Optional[aioredis.Redis] = None
|
_redis: aioredis.Redis | None = None
|
||||||
|
|
||||||
def __new__(cls):
|
def __new__(cls):
|
||||||
"""单例模式"""
|
"""单例模式"""
|
||||||
@@ -50,11 +51,11 @@ class RedisClientManager:
|
|||||||
return
|
return
|
||||||
|
|
||||||
self._initialized = True
|
self._initialized = True
|
||||||
self._circuit_open_until: Optional[float] = None
|
self._circuit_open_until: float | None = None
|
||||||
self._consecutive_failures: int = 0
|
self._consecutive_failures: int = 0
|
||||||
self._circuit_threshold = int(os.getenv("REDIS_CIRCUIT_BREAKER_THRESHOLD", "3"))
|
self._circuit_threshold = int(os.getenv("REDIS_CIRCUIT_BREAKER_THRESHOLD", "3"))
|
||||||
self._circuit_reset_seconds = int(os.getenv("REDIS_CIRCUIT_BREAKER_RESET_SECONDS", "60"))
|
self._circuit_reset_seconds = int(os.getenv("REDIS_CIRCUIT_BREAKER_RESET_SECONDS", "60"))
|
||||||
self._last_error: Optional[str] = None # 记录最后一次错误
|
self._last_error: str | None = None # 记录最后一次错误
|
||||||
|
|
||||||
def get_state(self) -> RedisState:
|
def get_state(self) -> RedisState:
|
||||||
"""
|
"""
|
||||||
@@ -100,7 +101,7 @@ class RedisClientManager:
|
|||||||
self._consecutive_failures = 0
|
self._consecutive_failures = 0
|
||||||
self._last_error = None
|
self._last_error = None
|
||||||
|
|
||||||
async def initialize(self, require_redis: bool = False) -> Optional[aioredis.Redis]:
|
async def initialize(self, require_redis: bool = False) -> aioredis.Redis | None:
|
||||||
"""
|
"""
|
||||||
初始化Redis连接
|
初始化Redis连接
|
||||||
|
|
||||||
@@ -236,7 +237,7 @@ class RedisClientManager:
|
|||||||
self._redis = None
|
self._redis = None
|
||||||
logger.info("全局Redis客户端已关闭")
|
logger.info("全局Redis客户端已关闭")
|
||||||
|
|
||||||
def get_client(self) -> Optional[aioredis.Redis]:
|
def get_client(self) -> aioredis.Redis | None:
|
||||||
"""
|
"""
|
||||||
获取Redis客户端(非异步)
|
获取Redis客户端(非异步)
|
||||||
|
|
||||||
@@ -249,10 +250,10 @@ class RedisClientManager:
|
|||||||
|
|
||||||
|
|
||||||
# 全局单例
|
# 全局单例
|
||||||
_redis_manager: Optional[RedisClientManager] = None
|
_redis_manager: RedisClientManager | None = None
|
||||||
|
|
||||||
|
|
||||||
async def get_redis_client(require_redis: bool = False) -> Optional[aioredis.Redis]:
|
async def get_redis_client(require_redis: bool = False) -> aioredis.Redis | None:
|
||||||
"""
|
"""
|
||||||
获取全局Redis客户端
|
获取全局Redis客户端
|
||||||
|
|
||||||
@@ -277,7 +278,7 @@ async def get_redis_client(require_redis: bool = False) -> Optional[aioredis.Red
|
|||||||
return _redis_manager.get_client()
|
return _redis_manager.get_client()
|
||||||
|
|
||||||
|
|
||||||
def get_redis_client_sync() -> Optional[aioredis.Redis]:
|
def get_redis_client_sync() -> aioredis.Redis | None:
|
||||||
"""
|
"""
|
||||||
同步获取Redis客户端(不会初始化)
|
同步获取Redis客户端(不会初始化)
|
||||||
|
|
||||||
|
|||||||
@@ -12,8 +12,9 @@
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
from typing import TYPE_CHECKING, Optional, Tuple
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from src.core.api_format.conversion.registry import FormatConversionRegistry
|
from src.core.api_format.conversion.registry import FormatConversionRegistry
|
||||||
@@ -26,11 +27,11 @@ logger = logging.getLogger(__name__)
|
|||||||
def is_format_compatible(
|
def is_format_compatible(
|
||||||
client_format: str,
|
client_format: str,
|
||||||
endpoint_api_format: str,
|
endpoint_api_format: str,
|
||||||
endpoint_format_acceptance_config: Optional[dict],
|
endpoint_format_acceptance_config: dict | None,
|
||||||
is_stream: bool,
|
is_stream: bool,
|
||||||
global_conversion_enabled: bool,
|
global_conversion_enabled: bool,
|
||||||
registry: Optional["FormatConversionRegistry"] = None,
|
registry: FormatConversionRegistry | None = None,
|
||||||
) -> Tuple[bool, bool, Optional[str]]:
|
) -> tuple[bool, bool, str | None]:
|
||||||
"""
|
"""
|
||||||
检查端点是否兼容客户端格式
|
检查端点是否兼容客户端格式
|
||||||
|
|
||||||
|
|||||||
@@ -4,7 +4,6 @@
|
|||||||
用于严格模式转换失败时抛出,让编排器可以尝试下一个候选。
|
用于严格模式转换失败时抛出,让编排器可以尝试下一个候选。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
|
|
||||||
class FormatConversionError(Exception):
|
class FormatConversionError(Exception):
|
||||||
|
|||||||
@@ -9,13 +9,11 @@
|
|||||||
这些应复用 `src/core/api_format/metadata.py`(API_FORMAT_DEFINITIONS)作为单一事实来源。
|
这些应复用 `src/core/api_format/metadata.py`(API_FORMAT_DEFINITIONS)作为单一事实来源。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from typing import Dict, Set
|
|
||||||
|
|
||||||
|
|
||||||
# 角色映射(仅作为辅助;system/tool 的具体落点以 Normalizer 规则为准)
|
# 角色映射(仅作为辅助;system/tool 的具体落点以 Normalizer 规则为准)
|
||||||
ROLE_MAPPINGS: Dict[str, Dict[str, str]] = {
|
ROLE_MAPPINGS: dict[str, dict[str, str]] = {
|
||||||
"OPENAI": {
|
"OPENAI": {
|
||||||
"user": "user",
|
"user": "user",
|
||||||
"assistant": "assistant",
|
"assistant": "assistant",
|
||||||
@@ -29,7 +27,7 @@ ROLE_MAPPINGS: Dict[str, Dict[str, str]] = {
|
|||||||
|
|
||||||
|
|
||||||
# 停止原因映射(internal -> provider),未知值使用 UNKNOWN 并写入 extra/raw
|
# 停止原因映射(internal -> provider),未知值使用 UNKNOWN 并写入 extra/raw
|
||||||
STOP_REASON_MAPPINGS: Dict[str, Dict[str, str]] = {
|
STOP_REASON_MAPPINGS: dict[str, dict[str, str]] = {
|
||||||
"CLAUDE": {
|
"CLAUDE": {
|
||||||
"end_turn": "end_turn",
|
"end_turn": "end_turn",
|
||||||
"max_tokens": "max_tokens",
|
"max_tokens": "max_tokens",
|
||||||
@@ -60,7 +58,7 @@ STOP_REASON_MAPPINGS: Dict[str, Dict[str, str]] = {
|
|||||||
|
|
||||||
|
|
||||||
# 使用量字段映射(provider usage field -> internal UsageInfo field)
|
# 使用量字段映射(provider usage field -> internal UsageInfo field)
|
||||||
USAGE_FIELD_MAPPINGS: Dict[str, Dict[str, str]] = {
|
USAGE_FIELD_MAPPINGS: dict[str, dict[str, str]] = {
|
||||||
"CLAUDE": {
|
"CLAUDE": {
|
||||||
"input_tokens": "input_tokens",
|
"input_tokens": "input_tokens",
|
||||||
"output_tokens": "output_tokens",
|
"output_tokens": "output_tokens",
|
||||||
@@ -82,7 +80,7 @@ USAGE_FIELD_MAPPINGS: Dict[str, Dict[str, str]] = {
|
|||||||
|
|
||||||
|
|
||||||
# 错误类型映射(provider -> internal ErrorType.value)
|
# 错误类型映射(provider -> internal ErrorType.value)
|
||||||
ERROR_TYPE_MAPPINGS: Dict[str, Dict[str, str]] = {
|
ERROR_TYPE_MAPPINGS: dict[str, dict[str, str]] = {
|
||||||
"CLAUDE": {
|
"CLAUDE": {
|
||||||
"invalid_request_error": "invalid_request",
|
"invalid_request_error": "invalid_request",
|
||||||
"authentication_error": "authentication",
|
"authentication_error": "authentication",
|
||||||
@@ -116,7 +114,7 @@ ERROR_TYPE_MAPPINGS: Dict[str, Dict[str, str]] = {
|
|||||||
|
|
||||||
|
|
||||||
# 可重试的错误类型(internal ErrorType.value)
|
# 可重试的错误类型(internal ErrorType.value)
|
||||||
RETRYABLE_ERROR_TYPES: Set[str] = {
|
RETRYABLE_ERROR_TYPES: set[str] = {
|
||||||
"rate_limit",
|
"rate_limit",
|
||||||
"overloaded",
|
"overloaded",
|
||||||
"server_error",
|
"server_error",
|
||||||
|
|||||||
@@ -10,11 +10,10 @@
|
|||||||
- 兼容优先:UnknownBlock 在内部保留,但默认在输出阶段丢弃(可观测、可随时调整策略)
|
- 兼容优先:UnknownBlock 在内部保留,但默认在输出阶段丢弃(可观测、可随时调整策略)
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
from typing import Any, Dict, FrozenSet, List, Optional, Union
|
from typing import Any
|
||||||
|
|
||||||
|
|
||||||
class Role(str, Enum):
|
class Role(str, Enum):
|
||||||
@@ -65,7 +64,7 @@ class TextBlock:
|
|||||||
|
|
||||||
type: ContentType = field(default=ContentType.TEXT, init=False)
|
type: ContentType = field(default=ContentType.TEXT, init=False)
|
||||||
text: str = ""
|
text: str = ""
|
||||||
extra: Dict[str, Any] = field(default_factory=dict)
|
extra: dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -74,11 +73,11 @@ class ImageBlock:
|
|||||||
|
|
||||||
type: ContentType = field(default=ContentType.IMAGE, init=False)
|
type: ContentType = field(default=ContentType.IMAGE, init=False)
|
||||||
# base64 编码的图片数据(二选一)
|
# base64 编码的图片数据(二选一)
|
||||||
data: Optional[str] = None
|
data: str | None = None
|
||||||
media_type: Optional[str] = None
|
media_type: str | None = None
|
||||||
# 或者 URL 引用
|
# 或者 URL 引用
|
||||||
url: Optional[str] = None
|
url: str | None = None
|
||||||
extra: Dict[str, Any] = field(default_factory=dict)
|
extra: dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -88,8 +87,8 @@ class ToolUseBlock:
|
|||||||
type: ContentType = field(default=ContentType.TOOL_USE, init=False)
|
type: ContentType = field(default=ContentType.TOOL_USE, init=False)
|
||||||
tool_id: str = ""
|
tool_id: str = ""
|
||||||
tool_name: str = ""
|
tool_name: str = ""
|
||||||
tool_input: Dict[str, Any] = field(default_factory=dict)
|
tool_input: dict[str, Any] = field(default_factory=dict)
|
||||||
extra: Dict[str, Any] = field(default_factory=dict)
|
extra: dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -100,9 +99,9 @@ class ToolResultBlock:
|
|||||||
tool_use_id: str = "" # 对应的 ToolUseBlock.tool_id
|
tool_use_id: str = "" # 对应的 ToolUseBlock.tool_id
|
||||||
# 工具输出可能是纯文本,也可能是结构化 JSON(Gemini functionResponse 等)
|
# 工具输出可能是纯文本,也可能是结构化 JSON(Gemini functionResponse 等)
|
||||||
output: Any = None
|
output: Any = None
|
||||||
content_text: Optional[str] = None
|
content_text: str | None = None
|
||||||
is_error: bool = False
|
is_error: bool = False
|
||||||
extra: Dict[str, Any] = field(default_factory=dict)
|
extra: dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -111,11 +110,11 @@ class UnknownBlock:
|
|||||||
|
|
||||||
type: ContentType = field(default=ContentType.UNKNOWN, init=False)
|
type: ContentType = field(default=ContentType.UNKNOWN, init=False)
|
||||||
raw_type: str = "" # 原始的类型字符串(各格式不一致)
|
raw_type: str = "" # 原始的类型字符串(各格式不一致)
|
||||||
payload: Dict[str, Any] = field(default_factory=dict) # 原始结构(尽量保持)
|
payload: dict[str, Any] = field(default_factory=dict) # 原始结构(尽量保持)
|
||||||
extra: Dict[str, Any] = field(default_factory=dict)
|
extra: dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
|
|
||||||
ContentBlock = Union[TextBlock, ImageBlock, ToolUseBlock, ToolResultBlock, UnknownBlock]
|
ContentBlock = TextBlock | ImageBlock | ToolUseBlock | ToolResultBlock | UnknownBlock
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -123,8 +122,8 @@ class InternalMessage:
|
|||||||
"""统一的消息表示"""
|
"""统一的消息表示"""
|
||||||
|
|
||||||
role: Role
|
role: Role
|
||||||
content: List[ContentBlock] # 统一使用列表,纯文本用单个 TextBlock
|
content: list[ContentBlock] # 统一使用列表,纯文本用单个 TextBlock
|
||||||
extra: Dict[str, Any] = field(default_factory=dict)
|
extra: dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -132,9 +131,9 @@ class ToolDefinition:
|
|||||||
"""统一的工具定义"""
|
"""统一的工具定义"""
|
||||||
|
|
||||||
name: str
|
name: str
|
||||||
description: Optional[str] = None
|
description: str | None = None
|
||||||
parameters: Optional[Dict[str, Any]] = None # JSON Schema
|
parameters: dict[str, Any] | None = None # JSON Schema
|
||||||
extra: Dict[str, Any] = field(default_factory=dict)
|
extra: dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
|
|
||||||
class ToolChoiceType(str, Enum):
|
class ToolChoiceType(str, Enum):
|
||||||
@@ -149,8 +148,8 @@ class ToolChoice:
|
|||||||
"""统一的工具选择"""
|
"""统一的工具选择"""
|
||||||
|
|
||||||
type: ToolChoiceType
|
type: ToolChoiceType
|
||||||
tool_name: Optional[str] = None
|
tool_name: str | None = None
|
||||||
extra: Dict[str, Any] = field(default_factory=dict)
|
extra: dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -159,7 +158,7 @@ class InstructionSegment:
|
|||||||
|
|
||||||
role: Role # 仅允许 Role.SYSTEM / Role.DEVELOPER
|
role: Role # 仅允许 Role.SYSTEM / Role.DEVELOPER
|
||||||
text: str = ""
|
text: str = ""
|
||||||
extra: Dict[str, Any] = field(default_factory=dict)
|
extra: dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -167,25 +166,25 @@ class InternalRequest:
|
|||||||
"""统一的请求表示"""
|
"""统一的请求表示"""
|
||||||
|
|
||||||
model: str
|
model: str
|
||||||
messages: List[InternalMessage]
|
messages: list[InternalMessage]
|
||||||
|
|
||||||
# 指令层:保留 system/developer 结构与顺序
|
# 指令层:保留 system/developer 结构与顺序
|
||||||
instructions: List[InstructionSegment] = field(default_factory=list)
|
instructions: list[InstructionSegment] = field(default_factory=list)
|
||||||
|
|
||||||
# 兼容字段:instructions 的 join 文本(无 role 标签),用于 Claude/Gemini 这类仅接受字符串 system 的格式
|
# 兼容字段:instructions 的 join 文本(无 role 标签),用于 Claude/Gemini 这类仅接受字符串 system 的格式
|
||||||
system: Optional[str] = None
|
system: str | None = None
|
||||||
|
|
||||||
max_tokens: Optional[int] = None
|
max_tokens: int | None = None
|
||||||
temperature: Optional[float] = None
|
temperature: float | None = None
|
||||||
top_p: Optional[float] = None
|
top_p: float | None = None
|
||||||
top_k: Optional[int] = None
|
top_k: int | None = None
|
||||||
stop_sequences: Optional[List[str]] = None
|
stop_sequences: list[str] | None = None
|
||||||
stream: bool = False
|
stream: bool = False
|
||||||
tools: Optional[List[ToolDefinition]] = None
|
tools: list[ToolDefinition] | None = None
|
||||||
tool_choice: Optional[ToolChoice] = None # auto/none/required 或指定 tool_name
|
tool_choice: ToolChoice | None = None # auto/none/required 或指定 tool_name
|
||||||
extra: Dict[str, Any] = field(default_factory=dict) # 未识别字段透传
|
extra: dict[str, Any] = field(default_factory=dict) # 未识别字段透传
|
||||||
|
|
||||||
def to_debug_dict(self) -> Dict[str, Any]:
|
def to_debug_dict(self) -> dict[str, Any]:
|
||||||
"""用于日志和调试的简化表示"""
|
"""用于日志和调试的简化表示"""
|
||||||
return {
|
return {
|
||||||
"model": self.model,
|
"model": self.model,
|
||||||
@@ -208,7 +207,7 @@ class UsageInfo:
|
|||||||
total_tokens: int = 0
|
total_tokens: int = 0
|
||||||
cache_read_tokens: int = 0
|
cache_read_tokens: int = 0
|
||||||
cache_write_tokens: int = 0
|
cache_write_tokens: int = 0
|
||||||
extra: Dict[str, Any] = field(default_factory=dict)
|
extra: dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -217,12 +216,12 @@ class InternalResponse:
|
|||||||
|
|
||||||
id: str
|
id: str
|
||||||
model: str
|
model: str
|
||||||
content: List[ContentBlock]
|
content: list[ContentBlock]
|
||||||
stop_reason: Optional[StopReason] = None
|
stop_reason: StopReason | None = None
|
||||||
usage: Optional[UsageInfo] = None
|
usage: UsageInfo | None = None
|
||||||
extra: Dict[str, Any] = field(default_factory=dict)
|
extra: dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
def to_debug_dict(self) -> Dict[str, Any]:
|
def to_debug_dict(self) -> dict[str, Any]:
|
||||||
"""用于日志和调试的简化表示"""
|
"""用于日志和调试的简化表示"""
|
||||||
usage = None
|
usage = None
|
||||||
if self.usage:
|
if self.usage:
|
||||||
@@ -246,12 +245,12 @@ class InternalError:
|
|||||||
|
|
||||||
type: ErrorType
|
type: ErrorType
|
||||||
message: str
|
message: str
|
||||||
code: Optional[str] = None # 原始错误码
|
code: str | None = None # 原始错误码
|
||||||
param: Optional[str] = None # 导致错误的参数
|
param: str | None = None # 导致错误的参数
|
||||||
retryable: bool = False # 是否可重试
|
retryable: bool = False # 是否可重试
|
||||||
extra: Dict[str, Any] = field(default_factory=dict)
|
extra: dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
def to_debug_dict(self) -> Dict[str, Any]:
|
def to_debug_dict(self) -> dict[str, Any]:
|
||||||
"""用于日志和调试"""
|
"""用于日志和调试"""
|
||||||
return {
|
return {
|
||||||
"type": self.type.value,
|
"type": self.type.value,
|
||||||
@@ -269,7 +268,7 @@ class FormatCapabilities:
|
|||||||
supports_error_conversion: bool = True
|
supports_error_conversion: bool = True
|
||||||
supports_tools: bool = True
|
supports_tools: bool = True
|
||||||
supports_images: bool = False
|
supports_images: bool = False
|
||||||
supported_features: FrozenSet[str] = field(default_factory=frozenset)
|
supported_features: frozenset[str] = field(default_factory=frozenset)
|
||||||
|
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
|
|||||||
@@ -5,10 +5,9 @@
|
|||||||
再从 internal 输出到目标格式。
|
再从 internal 输出到目标格式。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from typing import Any, Dict, List, Optional
|
from typing import Any
|
||||||
|
|
||||||
from .internal import FormatCapabilities, InternalError, InternalRequest, InternalResponse
|
from .internal import FormatCapabilities, InternalError, InternalRequest, InternalResponse
|
||||||
from .stream_events import InternalStreamEvent
|
from .stream_events import InternalStreamEvent
|
||||||
@@ -24,19 +23,19 @@ class FormatNormalizer(ABC):
|
|||||||
# ============ 请求转换 ============
|
# ============ 请求转换 ============
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def request_to_internal(self, request: Dict[str, Any]) -> InternalRequest:
|
def request_to_internal(self, request: dict[str, Any]) -> InternalRequest:
|
||||||
"""将格式特定请求转换为内部表示"""
|
"""将格式特定请求转换为内部表示"""
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def request_from_internal(self, internal: InternalRequest) -> Dict[str, Any]:
|
def request_from_internal(self, internal: InternalRequest) -> dict[str, Any]:
|
||||||
"""将内部表示转换为格式特定请求"""
|
"""将内部表示转换为格式特定请求"""
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
# ============ 响应转换 ============
|
# ============ 响应转换 ============
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def response_to_internal(self, response: Dict[str, Any]) -> InternalResponse:
|
def response_to_internal(self, response: dict[str, Any]) -> InternalResponse:
|
||||||
"""将格式特定响应转换为内部表示"""
|
"""将格式特定响应转换为内部表示"""
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
@@ -45,8 +44,8 @@ class FormatNormalizer(ABC):
|
|||||||
self,
|
self,
|
||||||
internal: InternalResponse,
|
internal: InternalResponse,
|
||||||
*,
|
*,
|
||||||
requested_model: Optional[str] = None,
|
requested_model: str | None = None,
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""将内部表示转换为格式特定响应
|
"""将内部表示转换为格式特定响应
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
@@ -61,9 +60,9 @@ class FormatNormalizer(ABC):
|
|||||||
|
|
||||||
def stream_chunk_to_internal(
|
def stream_chunk_to_internal(
|
||||||
self,
|
self,
|
||||||
chunk: Dict[str, Any],
|
chunk: dict[str, Any],
|
||||||
state: StreamState,
|
state: StreamState,
|
||||||
) -> List[InternalStreamEvent]:
|
) -> list[InternalStreamEvent]:
|
||||||
"""将格式特定流式块转换为内部事件"""
|
"""将格式特定流式块转换为内部事件"""
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
@@ -71,21 +70,21 @@ class FormatNormalizer(ABC):
|
|||||||
self,
|
self,
|
||||||
event: InternalStreamEvent,
|
event: InternalStreamEvent,
|
||||||
state: StreamState,
|
state: StreamState,
|
||||||
) -> List[Dict[str, Any]]:
|
) -> list[dict[str, Any]]:
|
||||||
"""将内部事件转换为格式特定流式块"""
|
"""将内部事件转换为格式特定流式块"""
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
# ============ 错误转换(可选) ============
|
# ============ 错误转换(可选) ============
|
||||||
|
|
||||||
def is_error_response(self, response: Dict[str, Any]) -> bool:
|
def is_error_response(self, response: dict[str, Any]) -> bool:
|
||||||
"""基于 body 的兜底判断(不可靠),子类可覆盖"""
|
"""基于 body 的兜底判断(不可靠),子类可覆盖"""
|
||||||
return False
|
return False
|
||||||
|
|
||||||
def error_to_internal(self, error_response: Dict[str, Any]) -> InternalError:
|
def error_to_internal(self, error_response: dict[str, Any]) -> InternalError:
|
||||||
"""将格式特定错误转换为内部表示"""
|
"""将格式特定错误转换为内部表示"""
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
def error_from_internal(self, internal: InternalError) -> Dict[str, Any]:
|
def error_from_internal(self, internal: InternalError) -> dict[str, Any]:
|
||||||
"""将内部错误表示转换为格式特定错误"""
|
"""将内部错误表示转换为格式特定错误"""
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
|
|||||||
@@ -6,7 +6,6 @@ Normalizers
|
|||||||
本目录在 Phase 1 仅创建结构;具体实现将在 Phase 2+ 补齐。
|
本目录在 Phase 1 仅创建结构;具体实现将在 Phase 2+ 补齐。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
__all__: list[str] = []
|
__all__: list[str] = []
|
||||||
|
|
||||||
|
|||||||
@@ -7,10 +7,9 @@ Claude Messages API Normalizer
|
|||||||
- 可选:Claude error <-> InternalError
|
- 可选:Claude error <-> InternalError
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import json
|
import json
|
||||||
from typing import Any, Dict, List, Optional, Tuple
|
from typing import Any
|
||||||
|
|
||||||
from src.core.api_format.conversion.field_mappings import (
|
from src.core.api_format.conversion.field_mappings import (
|
||||||
ERROR_TYPE_MAPPINGS,
|
ERROR_TYPE_MAPPINGS,
|
||||||
@@ -63,7 +62,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
supports_images=True,
|
supports_images=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
_CLAUDE_STOP_TO_INTERNAL: Dict[str, StopReason] = {
|
_CLAUDE_STOP_TO_INTERNAL: dict[str, StopReason] = {
|
||||||
"end_turn": StopReason.END_TURN,
|
"end_turn": StopReason.END_TURN,
|
||||||
"max_tokens": StopReason.MAX_TOKENS,
|
"max_tokens": StopReason.MAX_TOKENS,
|
||||||
"stop_sequence": StopReason.STOP_SEQUENCE,
|
"stop_sequence": StopReason.STOP_SEQUENCE,
|
||||||
@@ -73,7 +72,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
"content_filtered": StopReason.CONTENT_FILTERED,
|
"content_filtered": StopReason.CONTENT_FILTERED,
|
||||||
}
|
}
|
||||||
|
|
||||||
_ERROR_TYPE_TO_CLAUDE: Dict[ErrorType, str] = {
|
_ERROR_TYPE_TO_CLAUDE: dict[ErrorType, str] = {
|
||||||
ErrorType.INVALID_REQUEST: "invalid_request_error",
|
ErrorType.INVALID_REQUEST: "invalid_request_error",
|
||||||
ErrorType.AUTHENTICATION: "authentication_error",
|
ErrorType.AUTHENTICATION: "authentication_error",
|
||||||
ErrorType.PERMISSION_DENIED: "permission_error",
|
ErrorType.PERMISSION_DENIED: "permission_error",
|
||||||
@@ -90,11 +89,11 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
# Requests
|
# Requests
|
||||||
# =========================
|
# =========================
|
||||||
|
|
||||||
def request_to_internal(self, request: Dict[str, Any]) -> InternalRequest:
|
def request_to_internal(self, request: dict[str, Any]) -> InternalRequest:
|
||||||
model = str(request.get("model") or "")
|
model = str(request.get("model") or "")
|
||||||
dropped: Dict[str, int] = {}
|
dropped: dict[str, int] = {}
|
||||||
|
|
||||||
instructions: List[InstructionSegment] = []
|
instructions: list[InstructionSegment] = []
|
||||||
|
|
||||||
# 顶层 system 先进入 instructions(保持确定性优先级)
|
# 顶层 system 先进入 instructions(保持确定性优先级)
|
||||||
sys_value = request.get("system")
|
sys_value = request.get("system")
|
||||||
@@ -103,7 +102,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
if sys_text:
|
if sys_text:
|
||||||
instructions.append(InstructionSegment(role=Role.SYSTEM, text=sys_text))
|
instructions.append(InstructionSegment(role=Role.SYSTEM, text=sys_text))
|
||||||
|
|
||||||
messages: List[InternalMessage] = []
|
messages: list[InternalMessage] = []
|
||||||
for msg in request.get("messages") or []:
|
for msg in request.get("messages") or []:
|
||||||
if not isinstance(msg, dict):
|
if not isinstance(msg, dict):
|
||||||
dropped["claude_message_non_dict"] = dropped.get("claude_message_non_dict", 0) + 1
|
dropped["claude_message_non_dict"] = dropped.get("claude_message_non_dict", 0) + 1
|
||||||
@@ -156,15 +155,15 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
|
|
||||||
return internal
|
return internal
|
||||||
|
|
||||||
def request_from_internal(self, internal: InternalRequest) -> Dict[str, Any]:
|
def request_from_internal(self, internal: InternalRequest) -> dict[str, Any]:
|
||||||
system_text = internal.system or self._join_instructions(internal.instructions)
|
system_text = internal.system or self._join_instructions(internal.instructions)
|
||||||
|
|
||||||
# Claude Messages API: messages[] 仅允许 user/assistant,且需要交替;这里做最小修复
|
# Claude Messages API: messages[] 仅允许 user/assistant,且需要交替;这里做最小修复
|
||||||
fixed_messages = self._coerce_claude_message_sequence(internal.messages)
|
fixed_messages = self._coerce_claude_message_sequence(internal.messages)
|
||||||
|
|
||||||
out_messages: List[Dict[str, Any]] = [self._internal_message_to_claude(m) for m in fixed_messages]
|
out_messages: list[dict[str, Any]] = [self._internal_message_to_claude(m) for m in fixed_messages]
|
||||||
|
|
||||||
result: Dict[str, Any] = {
|
result: dict[str, Any] = {
|
||||||
"model": internal.model,
|
"model": internal.model,
|
||||||
"messages": out_messages,
|
"messages": out_messages,
|
||||||
"max_tokens": internal.max_tokens if internal.max_tokens is not None else 4096,
|
"max_tokens": internal.max_tokens if internal.max_tokens is not None else 4096,
|
||||||
@@ -210,20 +209,20 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
# Responses
|
# Responses
|
||||||
# =========================
|
# =========================
|
||||||
|
|
||||||
def response_to_internal(self, response: Dict[str, Any]) -> InternalResponse:
|
def response_to_internal(self, response: dict[str, Any]) -> InternalResponse:
|
||||||
rid = str(response.get("id") or "")
|
rid = str(response.get("id") or "")
|
||||||
model = str(response.get("model") or "")
|
model = str(response.get("model") or "")
|
||||||
|
|
||||||
blocks, dropped = self._claude_content_to_blocks(response.get("content"))
|
blocks, dropped = self._claude_content_to_blocks(response.get("content"))
|
||||||
|
|
||||||
raw_stop = response.get("stop_reason")
|
raw_stop = response.get("stop_reason")
|
||||||
stop_reason: Optional[StopReason] = None
|
stop_reason: StopReason | None = None
|
||||||
if raw_stop is not None:
|
if raw_stop is not None:
|
||||||
stop_reason = self._CLAUDE_STOP_TO_INTERNAL.get(str(raw_stop), StopReason.UNKNOWN)
|
stop_reason = self._CLAUDE_STOP_TO_INTERNAL.get(str(raw_stop), StopReason.UNKNOWN)
|
||||||
|
|
||||||
usage_info = self._claude_usage_to_internal(response.get("usage"))
|
usage_info = self._claude_usage_to_internal(response.get("usage"))
|
||||||
|
|
||||||
extra: Dict[str, Any] = {}
|
extra: dict[str, Any] = {}
|
||||||
if raw_stop is not None:
|
if raw_stop is not None:
|
||||||
extra.setdefault("raw", {})["stop_reason"] = raw_stop
|
extra.setdefault("raw", {})["stop_reason"] = raw_stop
|
||||||
|
|
||||||
@@ -245,13 +244,13 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
self,
|
self,
|
||||||
internal: InternalResponse,
|
internal: InternalResponse,
|
||||||
*,
|
*,
|
||||||
requested_model: Optional[str] = None,
|
requested_model: str | None = None,
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
cid = internal.id or "unknown"
|
cid = internal.id or "unknown"
|
||||||
if not cid.startswith("msg_"):
|
if not cid.startswith("msg_"):
|
||||||
cid = f"msg_{cid}"
|
cid = f"msg_{cid}"
|
||||||
|
|
||||||
content: List[Dict[str, Any]] = []
|
content: list[dict[str, Any]] = []
|
||||||
for b in internal.content:
|
for b in internal.content:
|
||||||
if isinstance(b, TextBlock):
|
if isinstance(b, TextBlock):
|
||||||
if b.text:
|
if b.text:
|
||||||
@@ -288,7 +287,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
if internal.stop_reason is not None:
|
if internal.stop_reason is not None:
|
||||||
stop_reason = STOP_REASON_MAPPINGS.get("CLAUDE", {}).get(internal.stop_reason.value, "end_turn")
|
stop_reason = STOP_REASON_MAPPINGS.get("CLAUDE", {}).get(internal.stop_reason.value, "end_turn")
|
||||||
|
|
||||||
usage: Dict[str, Any] = {"input_tokens": 0, "output_tokens": 0}
|
usage: dict[str, Any] = {"input_tokens": 0, "output_tokens": 0}
|
||||||
if internal.usage:
|
if internal.usage:
|
||||||
usage = {
|
usage = {
|
||||||
"input_tokens": int(internal.usage.input_tokens),
|
"input_tokens": int(internal.usage.input_tokens),
|
||||||
@@ -319,11 +318,11 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
|
|
||||||
def stream_chunk_to_internal(
|
def stream_chunk_to_internal(
|
||||||
self,
|
self,
|
||||||
chunk: Dict[str, Any],
|
chunk: dict[str, Any],
|
||||||
state: StreamState,
|
state: StreamState,
|
||||||
) -> List[InternalStreamEvent]:
|
) -> list[InternalStreamEvent]:
|
||||||
ss = state.substate(self.FORMAT_ID)
|
ss = state.substate(self.FORMAT_ID)
|
||||||
events: List[InternalStreamEvent] = []
|
events: list[InternalStreamEvent] = []
|
||||||
|
|
||||||
event_type = chunk.get("type")
|
event_type = chunk.get("type")
|
||||||
if event_type is None:
|
if event_type is None:
|
||||||
@@ -335,7 +334,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
|
|
||||||
if event_type == "message_start":
|
if event_type == "message_start":
|
||||||
message_raw = chunk.get("message")
|
message_raw = chunk.get("message")
|
||||||
message: Dict[str, Any] = message_raw if isinstance(message_raw, dict) else {}
|
message: dict[str, Any] = message_raw if isinstance(message_raw, dict) else {}
|
||||||
msg_id = str(message.get("id") or "")
|
msg_id = str(message.get("id") or "")
|
||||||
# 保留初始化时设置的 model(客户端请求的模型),仅在空时用上游值
|
# 保留初始化时设置的 model(客户端请求的模型),仅在空时用上游值
|
||||||
model = state.model or str(message.get("model") or "")
|
model = state.model or str(message.get("model") or "")
|
||||||
@@ -350,7 +349,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
if event_type == "content_block_start":
|
if event_type == "content_block_start":
|
||||||
index = int(chunk.get("index") or 0)
|
index = int(chunk.get("index") or 0)
|
||||||
block_raw = chunk.get("content_block")
|
block_raw = chunk.get("content_block")
|
||||||
block: Dict[str, Any] = block_raw if isinstance(block_raw, dict) else {}
|
block: dict[str, Any] = block_raw if isinstance(block_raw, dict) else {}
|
||||||
btype = str(block.get("type") or "unknown")
|
btype = str(block.get("type") or "unknown")
|
||||||
|
|
||||||
if btype == "text":
|
if btype == "text":
|
||||||
@@ -385,7 +384,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
if event_type == "content_block_delta":
|
if event_type == "content_block_delta":
|
||||||
index = int(chunk.get("index") or 0)
|
index = int(chunk.get("index") or 0)
|
||||||
delta_raw = chunk.get("delta")
|
delta_raw = chunk.get("delta")
|
||||||
delta: Dict[str, Any] = delta_raw if isinstance(delta_raw, dict) else {}
|
delta: dict[str, Any] = delta_raw if isinstance(delta_raw, dict) else {}
|
||||||
dtype = str(delta.get("type") or "unknown")
|
dtype = str(delta.get("type") or "unknown")
|
||||||
|
|
||||||
if dtype == "text_delta":
|
if dtype == "text_delta":
|
||||||
@@ -415,7 +414,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
|
|
||||||
if event_type == "message_delta":
|
if event_type == "message_delta":
|
||||||
delta_raw2 = chunk.get("delta")
|
delta_raw2 = chunk.get("delta")
|
||||||
delta2: Dict[str, Any] = delta_raw2 if isinstance(delta_raw2, dict) else {}
|
delta2: dict[str, Any] = delta_raw2 if isinstance(delta_raw2, dict) else {}
|
||||||
raw_stop = delta2.get("stop_reason")
|
raw_stop = delta2.get("stop_reason")
|
||||||
if raw_stop is not None:
|
if raw_stop is not None:
|
||||||
ss["stop_reason"] = str(raw_stop)
|
ss["stop_reason"] = str(raw_stop)
|
||||||
@@ -426,7 +425,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
|
|
||||||
if event_type == "message_stop":
|
if event_type == "message_stop":
|
||||||
raw_stop = ss.get("stop_reason")
|
raw_stop = ss.get("stop_reason")
|
||||||
stop_reason: Optional[StopReason] = None
|
stop_reason: StopReason | None = None
|
||||||
if raw_stop is not None:
|
if raw_stop is not None:
|
||||||
stop_reason = self._CLAUDE_STOP_TO_INTERNAL.get(str(raw_stop), StopReason.UNKNOWN)
|
stop_reason = self._CLAUDE_STOP_TO_INTERNAL.get(str(raw_stop), StopReason.UNKNOWN)
|
||||||
usage_info = self._claude_usage_to_internal(ss.get("usage"))
|
usage_info = self._claude_usage_to_internal(ss.get("usage"))
|
||||||
@@ -444,9 +443,9 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
self,
|
self,
|
||||||
event: InternalStreamEvent,
|
event: InternalStreamEvent,
|
||||||
state: StreamState,
|
state: StreamState,
|
||||||
) -> List[Dict[str, Any]]:
|
) -> list[dict[str, Any]]:
|
||||||
ss = state.substate(self.FORMAT_ID)
|
ss = state.substate(self.FORMAT_ID)
|
||||||
out: List[Dict[str, Any]] = []
|
out: list[dict[str, Any]] = []
|
||||||
|
|
||||||
if isinstance(event, MessageStartEvent):
|
if isinstance(event, MessageStartEvent):
|
||||||
state.message_id = event.message_id or state.message_id
|
state.message_id = event.message_id or state.message_id
|
||||||
@@ -454,7 +453,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
if not state.model:
|
if not state.model:
|
||||||
state.model = event.model or ""
|
state.model = event.model or ""
|
||||||
ss.setdefault("block_index_to_tool_id", {})
|
ss.setdefault("block_index_to_tool_id", {})
|
||||||
message_obj: Dict[str, Any] = {
|
message_obj: dict[str, Any] = {
|
||||||
"id": state.message_id or "msg_stream",
|
"id": state.message_id or "msg_stream",
|
||||||
"type": "message",
|
"type": "message",
|
||||||
"role": "assistant",
|
"role": "assistant",
|
||||||
@@ -531,7 +530,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
if event.stop_reason is not None:
|
if event.stop_reason is not None:
|
||||||
stop_reason = STOP_REASON_MAPPINGS.get("CLAUDE", {}).get(event.stop_reason.value, "end_turn")
|
stop_reason = STOP_REASON_MAPPINGS.get("CLAUDE", {}).get(event.stop_reason.value, "end_turn")
|
||||||
|
|
||||||
msg_delta: Dict[str, Any] = {
|
msg_delta: dict[str, Any] = {
|
||||||
"type": "message_delta",
|
"type": "message_delta",
|
||||||
"delta": {"stop_reason": stop_reason},
|
"delta": {"stop_reason": stop_reason},
|
||||||
}
|
}
|
||||||
@@ -553,15 +552,15 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
# Error conversion
|
# Error conversion
|
||||||
# =========================
|
# =========================
|
||||||
|
|
||||||
def is_error_response(self, response: Dict[str, Any]) -> bool:
|
def is_error_response(self, response: dict[str, Any]) -> bool:
|
||||||
if not isinstance(response, dict):
|
if not isinstance(response, dict):
|
||||||
return False
|
return False
|
||||||
if response.get("type") == "error":
|
if response.get("type") == "error":
|
||||||
return True
|
return True
|
||||||
return "error" in response
|
return "error" in response
|
||||||
|
|
||||||
def error_to_internal(self, error_response: Dict[str, Any]) -> InternalError:
|
def error_to_internal(self, error_response: dict[str, Any]) -> InternalError:
|
||||||
err: Dict[str, Any] = {}
|
err: dict[str, Any] = {}
|
||||||
if isinstance(error_response, dict):
|
if isinstance(error_response, dict):
|
||||||
err_raw = error_response.get("error")
|
err_raw = error_response.get("error")
|
||||||
err = err_raw if isinstance(err_raw, dict) else {}
|
err = err_raw if isinstance(err_raw, dict) else {}
|
||||||
@@ -580,9 +579,9 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
extra={"claude": {"error": err}, "raw": {"type": raw_type}},
|
extra={"claude": {"error": err}, "raw": {"type": raw_type}},
|
||||||
)
|
)
|
||||||
|
|
||||||
def error_from_internal(self, internal: InternalError) -> Dict[str, Any]:
|
def error_from_internal(self, internal: InternalError) -> dict[str, Any]:
|
||||||
type_str = self._ERROR_TYPE_TO_CLAUDE.get(internal.type, "api_error")
|
type_str = self._ERROR_TYPE_TO_CLAUDE.get(internal.type, "api_error")
|
||||||
payload: Dict[str, Any] = {"type": type_str, "message": internal.message}
|
payload: dict[str, Any] = {"type": type_str, "message": internal.message}
|
||||||
if internal.param is not None:
|
if internal.param is not None:
|
||||||
payload["param"] = internal.param
|
payload["param"] = internal.param
|
||||||
if internal.code is not None:
|
if internal.code is not None:
|
||||||
@@ -593,8 +592,8 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
# Helpers
|
# Helpers
|
||||||
# =========================
|
# =========================
|
||||||
|
|
||||||
def _claude_message_to_internal(self, msg: Dict[str, Any]) -> Tuple[Optional[InternalMessage], Dict[str, int]]:
|
def _claude_message_to_internal(self, msg: dict[str, Any]) -> tuple[InternalMessage | None, dict[str, int]]:
|
||||||
dropped: Dict[str, int] = {}
|
dropped: dict[str, int] = {}
|
||||||
role_raw = str(msg.get("role") or "unknown")
|
role_raw = str(msg.get("role") or "unknown")
|
||||||
|
|
||||||
if role_raw == "user":
|
if role_raw == "user":
|
||||||
@@ -616,8 +615,8 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
dropped,
|
dropped,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _claude_content_to_blocks(self, content: Any) -> Tuple[List[ContentBlock], Dict[str, int]]:
|
def _claude_content_to_blocks(self, content: Any) -> tuple[list[ContentBlock], dict[str, int]]:
|
||||||
dropped: Dict[str, int] = {}
|
dropped: dict[str, int] = {}
|
||||||
if content is None:
|
if content is None:
|
||||||
return [], dropped
|
return [], dropped
|
||||||
if isinstance(content, str):
|
if isinstance(content, str):
|
||||||
@@ -626,7 +625,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
dropped["claude_content_non_list"] = dropped.get("claude_content_non_list", 0) + 1
|
dropped["claude_content_non_list"] = dropped.get("claude_content_non_list", 0) + 1
|
||||||
return [], dropped
|
return [], dropped
|
||||||
|
|
||||||
blocks: List[ContentBlock] = []
|
blocks: list[ContentBlock] = []
|
||||||
for block in content:
|
for block in content:
|
||||||
if not isinstance(block, dict):
|
if not isinstance(block, dict):
|
||||||
dropped["claude_block_non_dict"] = dropped.get("claude_block_non_dict", 0) + 1
|
dropped["claude_block_non_dict"] = dropped.get("claude_block_non_dict", 0) + 1
|
||||||
@@ -641,7 +640,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
|
|
||||||
if btype == "image":
|
if btype == "image":
|
||||||
src_raw = block.get("source")
|
src_raw = block.get("source")
|
||||||
src: Dict[str, Any] = src_raw if isinstance(src_raw, dict) else {}
|
src: dict[str, Any] = src_raw if isinstance(src_raw, dict) else {}
|
||||||
stype = src.get("type")
|
stype = src.get("type")
|
||||||
if stype == "base64":
|
if stype == "base64":
|
||||||
data = src.get("data")
|
data = src.get("data")
|
||||||
@@ -687,7 +686,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
tool_use_id: str,
|
tool_use_id: str,
|
||||||
raw_content: Any,
|
raw_content: Any,
|
||||||
is_error: bool,
|
is_error: bool,
|
||||||
raw_block: Dict[str, Any],
|
raw_block: dict[str, Any],
|
||||||
) -> ToolResultBlock:
|
) -> ToolResultBlock:
|
||||||
if raw_content is None:
|
if raw_content is None:
|
||||||
return ToolResultBlock(
|
return ToolResultBlock(
|
||||||
@@ -723,7 +722,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
)
|
)
|
||||||
|
|
||||||
if isinstance(raw_content, list):
|
if isinstance(raw_content, list):
|
||||||
text_parts: List[str] = []
|
text_parts: list[str] = []
|
||||||
for part in raw_content:
|
for part in raw_content:
|
||||||
if isinstance(part, dict) and part.get("type") == "text":
|
if isinstance(part, dict) and part.get("type") == "text":
|
||||||
text = part.get("text")
|
text = part.get("text")
|
||||||
@@ -747,15 +746,15 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
extra={"claude": raw_block},
|
extra={"claude": raw_block},
|
||||||
)
|
)
|
||||||
|
|
||||||
def _collapse_claude_system(self, system_value: Any) -> Tuple[Optional[str], Dict[str, int]]:
|
def _collapse_claude_system(self, system_value: Any) -> tuple[str | None, dict[str, int]]:
|
||||||
dropped: Dict[str, int] = {}
|
dropped: dict[str, int] = {}
|
||||||
if system_value is None:
|
if system_value is None:
|
||||||
return None, dropped
|
return None, dropped
|
||||||
if isinstance(system_value, str):
|
if isinstance(system_value, str):
|
||||||
return (system_value or None), dropped
|
return (system_value or None), dropped
|
||||||
|
|
||||||
if isinstance(system_value, list):
|
if isinstance(system_value, list):
|
||||||
texts: List[str] = []
|
texts: list[str] = []
|
||||||
for item in system_value:
|
for item in system_value:
|
||||||
if not isinstance(item, dict):
|
if not isinstance(item, dict):
|
||||||
dropped["claude_system_item_non_dict"] = dropped.get("claude_system_item_non_dict", 0) + 1
|
dropped["claude_system_item_non_dict"] = dropped.get("claude_system_item_non_dict", 0) + 1
|
||||||
@@ -773,16 +772,16 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
dropped["claude_system_unsupported"] = dropped.get("claude_system_unsupported", 0) + 1
|
dropped["claude_system_unsupported"] = dropped.get("claude_system_unsupported", 0) + 1
|
||||||
return None, dropped
|
return None, dropped
|
||||||
|
|
||||||
def _join_instructions(self, instructions: List[InstructionSegment]) -> Optional[str]:
|
def _join_instructions(self, instructions: list[InstructionSegment]) -> str | None:
|
||||||
parts = [seg.text for seg in instructions if seg.text]
|
parts = [seg.text for seg in instructions if seg.text]
|
||||||
joined = "\n\n".join(parts)
|
joined = "\n\n".join(parts)
|
||||||
return joined or None
|
return joined or None
|
||||||
|
|
||||||
def _claude_tools_to_internal(self, tools: Any) -> Optional[List[ToolDefinition]]:
|
def _claude_tools_to_internal(self, tools: Any) -> list[ToolDefinition] | None:
|
||||||
if not tools or not isinstance(tools, list):
|
if not tools or not isinstance(tools, list):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
out: List[ToolDefinition] = []
|
out: list[ToolDefinition] = []
|
||||||
for tool in tools:
|
for tool in tools:
|
||||||
if not isinstance(tool, dict):
|
if not isinstance(tool, dict):
|
||||||
continue
|
continue
|
||||||
@@ -799,7 +798,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
)
|
)
|
||||||
return out or None
|
return out or None
|
||||||
|
|
||||||
def _claude_tool_choice_to_internal(self, tool_choice: Any) -> Optional[ToolChoice]:
|
def _claude_tool_choice_to_internal(self, tool_choice: Any) -> ToolChoice | None:
|
||||||
if tool_choice is None:
|
if tool_choice is None:
|
||||||
return None
|
return None
|
||||||
if not isinstance(tool_choice, dict):
|
if not isinstance(tool_choice, dict):
|
||||||
@@ -818,7 +817,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
|
|
||||||
return ToolChoice(type=ToolChoiceType.AUTO, extra={"claude": tool_choice})
|
return ToolChoice(type=ToolChoiceType.AUTO, extra={"claude": tool_choice})
|
||||||
|
|
||||||
def _tool_choice_to_claude(self, tool_choice: ToolChoice) -> Dict[str, Any]:
|
def _tool_choice_to_claude(self, tool_choice: ToolChoice) -> dict[str, Any]:
|
||||||
if tool_choice.type == ToolChoiceType.NONE:
|
if tool_choice.type == ToolChoiceType.NONE:
|
||||||
return {"type": "none"}
|
return {"type": "none"}
|
||||||
if tool_choice.type == ToolChoiceType.AUTO:
|
if tool_choice.type == ToolChoiceType.AUTO:
|
||||||
@@ -829,11 +828,11 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
return {"type": "tool_use", "name": tool_choice.tool_name or ""}
|
return {"type": "tool_use", "name": tool_choice.tool_name or ""}
|
||||||
return {"type": "auto"}
|
return {"type": "auto"}
|
||||||
|
|
||||||
def _internal_message_to_claude(self, msg: InternalMessage) -> Dict[str, Any]:
|
def _internal_message_to_claude(self, msg: InternalMessage) -> dict[str, Any]:
|
||||||
role = "user" if msg.role == Role.USER else "assistant"
|
role = "user" if msg.role == Role.USER else "assistant"
|
||||||
|
|
||||||
blocks: List[Dict[str, Any]] = []
|
blocks: list[dict[str, Any]] = []
|
||||||
text_parts: List[str] = []
|
text_parts: list[str] = []
|
||||||
|
|
||||||
for b in msg.content:
|
for b in msg.content:
|
||||||
if isinstance(b, UnknownBlock):
|
if isinstance(b, UnknownBlock):
|
||||||
@@ -902,8 +901,8 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
|
|
||||||
return {"role": role, "content": blocks}
|
return {"role": role, "content": blocks}
|
||||||
|
|
||||||
def _coerce_claude_message_sequence(self, messages: List[InternalMessage]) -> List[InternalMessage]:
|
def _coerce_claude_message_sequence(self, messages: list[InternalMessage]) -> list[InternalMessage]:
|
||||||
normalized: List[InternalMessage] = []
|
normalized: list[InternalMessage] = []
|
||||||
for m in messages:
|
for m in messages:
|
||||||
role = m.role
|
role = m.role
|
||||||
if role not in (Role.USER, Role.ASSISTANT):
|
if role not in (Role.USER, Role.ASSISTANT):
|
||||||
@@ -916,7 +915,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
if normalized[0].role != Role.USER:
|
if normalized[0].role != Role.USER:
|
||||||
normalized = [InternalMessage(role=Role.USER, content=[])] + normalized
|
normalized = [InternalMessage(role=Role.USER, content=[])] + normalized
|
||||||
|
|
||||||
merged: List[InternalMessage] = []
|
merged: list[InternalMessage] = []
|
||||||
for m in normalized:
|
for m in normalized:
|
||||||
if merged and merged[-1].role == m.role:
|
if merged and merged[-1].role == m.role:
|
||||||
merged[-1].content.extend(m.content)
|
merged[-1].content.extend(m.content)
|
||||||
@@ -925,12 +924,12 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
|
|
||||||
return merged
|
return merged
|
||||||
|
|
||||||
def _claude_usage_to_internal(self, usage: Any) -> Optional[UsageInfo]:
|
def _claude_usage_to_internal(self, usage: Any) -> UsageInfo | None:
|
||||||
if not isinstance(usage, dict):
|
if not isinstance(usage, dict):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
mapping = USAGE_FIELD_MAPPINGS.get("CLAUDE", {})
|
mapping = USAGE_FIELD_MAPPINGS.get("CLAUDE", {})
|
||||||
fields: Dict[str, int] = {}
|
fields: dict[str, int] = {}
|
||||||
extra = self._extract_extra(usage, set(mapping.keys()))
|
extra = self._extract_extra(usage, set(mapping.keys()))
|
||||||
|
|
||||||
for provider_key, internal_key in mapping.items():
|
for provider_key, internal_key in mapping.items():
|
||||||
@@ -952,8 +951,8 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
extra={"claude": extra} if extra else {},
|
extra={"claude": extra} if extra else {},
|
||||||
)
|
)
|
||||||
|
|
||||||
def _usage_to_claude(self, usage: UsageInfo) -> Dict[str, Any]:
|
def _usage_to_claude(self, usage: UsageInfo) -> dict[str, Any]:
|
||||||
result: Dict[str, Any] = {
|
result: dict[str, Any] = {
|
||||||
"input_tokens": int(usage.input_tokens),
|
"input_tokens": int(usage.input_tokens),
|
||||||
"output_tokens": int(usage.output_tokens),
|
"output_tokens": int(usage.output_tokens),
|
||||||
}
|
}
|
||||||
@@ -969,7 +968,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
except ValueError:
|
except ValueError:
|
||||||
return ErrorType.UNKNOWN
|
return ErrorType.UNKNOWN
|
||||||
|
|
||||||
def _optional_int(self, value: Any) -> Optional[int]:
|
def _optional_int(self, value: Any) -> int | None:
|
||||||
if value is None:
|
if value is None:
|
||||||
return None
|
return None
|
||||||
try:
|
try:
|
||||||
@@ -977,7 +976,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
except (TypeError, ValueError):
|
except (TypeError, ValueError):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def _optional_float(self, value: Any) -> Optional[float]:
|
def _optional_float(self, value: Any) -> float | None:
|
||||||
if value is None:
|
if value is None:
|
||||||
return None
|
return None
|
||||||
try:
|
try:
|
||||||
@@ -985,7 +984,7 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
except (TypeError, ValueError):
|
except (TypeError, ValueError):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def _coerce_str_list(self, value: Any) -> Optional[List[str]]:
|
def _coerce_str_list(self, value: Any) -> list[str] | None:
|
||||||
if value is None:
|
if value is None:
|
||||||
return None
|
return None
|
||||||
if isinstance(value, str):
|
if isinstance(value, str):
|
||||||
@@ -994,10 +993,10 @@ class ClaudeNormalizer(FormatNormalizer):
|
|||||||
return [str(x) for x in value if x is not None]
|
return [str(x) for x in value if x is not None]
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def _extract_extra(self, payload: Dict[str, Any], known_keys: set[str]) -> Dict[str, Any]:
|
def _extract_extra(self, payload: dict[str, Any], known_keys: set[str]) -> dict[str, Any]:
|
||||||
return {k: v for k, v in payload.items() if k not in known_keys}
|
return {k: v for k, v in payload.items() if k not in known_keys}
|
||||||
|
|
||||||
def _merge_dropped(self, target: Dict[str, int], source: Dict[str, int]) -> None:
|
def _merge_dropped(self, target: dict[str, int], source: dict[str, int]) -> None:
|
||||||
for k, v in source.items():
|
for k, v in source.items():
|
||||||
target[k] = target.get(k, 0) + int(v)
|
target[k] = target.get(k, 0) + int(v)
|
||||||
|
|
||||||
|
|||||||
@@ -7,7 +7,6 @@ CLAUDE_CLI 的请求/响应 body 与 CLAUDE 一致(Anthropic Messages API)
|
|||||||
如需 CLI 特殊处理,可覆盖 request_from_internal / request_to_internal 等方法。
|
如需 CLI 特殊处理,可覆盖 request_from_internal / request_to_internal 等方法。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from src.core.api_format.conversion.normalizers.claude import ClaudeNormalizer
|
from src.core.api_format.conversion.normalizers.claude import ClaudeNormalizer
|
||||||
|
|
||||||
|
|||||||
@@ -11,11 +11,9 @@ Gemini (GenerateContent / streamGenerateContent) Normalizer
|
|||||||
- 响应/流式通常为 camelCase(candidates/finishReason/usageMetadata/modelVersion)。
|
- 响应/流式通常为 camelCase(candidates/finishReason/usageMetadata/modelVersion)。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import json
|
import json
|
||||||
import time
|
from typing import Any
|
||||||
from typing import Any, Dict, List, Optional, Tuple
|
|
||||||
|
|
||||||
from src.core.api_format.conversion.field_mappings import (
|
from src.core.api_format.conversion.field_mappings import (
|
||||||
ERROR_TYPE_MAPPINGS,
|
ERROR_TYPE_MAPPINGS,
|
||||||
@@ -68,7 +66,7 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
supports_images=True,
|
supports_images=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
_FINISH_REASON_TO_STOP: Dict[str, StopReason] = {
|
_FINISH_REASON_TO_STOP: dict[str, StopReason] = {
|
||||||
"STOP": StopReason.END_TURN,
|
"STOP": StopReason.END_TURN,
|
||||||
"MAX_TOKENS": StopReason.MAX_TOKENS,
|
"MAX_TOKENS": StopReason.MAX_TOKENS,
|
||||||
"SAFETY": StopReason.CONTENT_FILTERED,
|
"SAFETY": StopReason.CONTENT_FILTERED,
|
||||||
@@ -77,7 +75,7 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
"OTHER": StopReason.UNKNOWN,
|
"OTHER": StopReason.UNKNOWN,
|
||||||
}
|
}
|
||||||
|
|
||||||
_ERROR_TYPE_TO_GEMINI_STATUS: Dict[ErrorType, str] = {
|
_ERROR_TYPE_TO_GEMINI_STATUS: dict[ErrorType, str] = {
|
||||||
ErrorType.INVALID_REQUEST: "INVALID_ARGUMENT",
|
ErrorType.INVALID_REQUEST: "INVALID_ARGUMENT",
|
||||||
ErrorType.AUTHENTICATION: "UNAUTHENTICATED",
|
ErrorType.AUTHENTICATION: "UNAUTHENTICATED",
|
||||||
ErrorType.PERMISSION_DENIED: "PERMISSION_DENIED",
|
ErrorType.PERMISSION_DENIED: "PERMISSION_DENIED",
|
||||||
@@ -94,11 +92,11 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
# Requests
|
# Requests
|
||||||
# =========================
|
# =========================
|
||||||
|
|
||||||
def request_to_internal(self, request: Dict[str, Any]) -> InternalRequest:
|
def request_to_internal(self, request: dict[str, Any]) -> InternalRequest:
|
||||||
model = str(request.get("model") or "")
|
model = str(request.get("model") or "")
|
||||||
dropped: Dict[str, int] = {}
|
dropped: dict[str, int] = {}
|
||||||
|
|
||||||
instructions: List[InstructionSegment] = []
|
instructions: list[InstructionSegment] = []
|
||||||
system_text, sys_dropped = self._collapse_system_instruction(
|
system_text, sys_dropped = self._collapse_system_instruction(
|
||||||
request.get("system_instruction")
|
request.get("system_instruction")
|
||||||
if "system_instruction" in request
|
if "system_instruction" in request
|
||||||
@@ -108,7 +106,7 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
if system_text:
|
if system_text:
|
||||||
instructions.append(InstructionSegment(role=Role.SYSTEM, text=system_text))
|
instructions.append(InstructionSegment(role=Role.SYSTEM, text=system_text))
|
||||||
|
|
||||||
messages: List[InternalMessage] = []
|
messages: list[InternalMessage] = []
|
||||||
contents = request.get("contents") or []
|
contents = request.get("contents") or []
|
||||||
if isinstance(contents, list):
|
if isinstance(contents, list):
|
||||||
for content in contents:
|
for content in contents:
|
||||||
@@ -152,7 +150,7 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# 构建 extra,保留原始 gemini 字段
|
# 构建 extra,保留原始 gemini 字段
|
||||||
extra: Dict[str, Any] = {"gemini": self._extract_extra(request, {"contents"})}
|
extra: dict[str, Any] = {"gemini": self._extract_extra(request, {"contents"})}
|
||||||
|
|
||||||
# 保留 generationConfig 中的特殊字段(responseModalities, thinkingConfig 等)
|
# 保留 generationConfig 中的特殊字段(responseModalities, thinkingConfig 等)
|
||||||
# 这些字段在 _get_generation_config 中已提取,需要单独存储以便转换时使用
|
# 这些字段在 _get_generation_config 中已提取,需要单独存储以便转换时使用
|
||||||
@@ -160,7 +158,7 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
response_modalities = generation_config.get("response_modalities")
|
response_modalities = generation_config.get("response_modalities")
|
||||||
thinking_config = generation_config.get("thinking_config")
|
thinking_config = generation_config.get("thinking_config")
|
||||||
if response_modalities or thinking_config:
|
if response_modalities or thinking_config:
|
||||||
google_extra: Dict[str, Any] = {}
|
google_extra: dict[str, Any] = {}
|
||||||
if response_modalities:
|
if response_modalities:
|
||||||
google_extra["response_modalities"] = response_modalities
|
google_extra["response_modalities"] = response_modalities
|
||||||
if thinking_config:
|
if thinking_config:
|
||||||
@@ -188,7 +186,7 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
|
|
||||||
return internal
|
return internal
|
||||||
|
|
||||||
def request_from_internal(self, internal: InternalRequest) -> Dict[str, Any]:
|
def request_from_internal(self, internal: InternalRequest) -> dict[str, Any]:
|
||||||
system_text = internal.system or self._join_instructions(internal.instructions)
|
system_text = internal.system or self._join_instructions(internal.instructions)
|
||||||
|
|
||||||
# tools/tool_choice
|
# tools/tool_choice
|
||||||
@@ -212,7 +210,7 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
if internal.tool_choice:
|
if internal.tool_choice:
|
||||||
tool_config = self._tool_choice_to_gemini_tool_config(internal.tool_choice)
|
tool_config = self._tool_choice_to_gemini_tool_config(internal.tool_choice)
|
||||||
|
|
||||||
generation_config: Dict[str, Any] = {}
|
generation_config: dict[str, Any] = {}
|
||||||
if internal.max_tokens is not None:
|
if internal.max_tokens is not None:
|
||||||
generation_config["max_output_tokens"] = internal.max_tokens
|
generation_config["max_output_tokens"] = internal.max_tokens
|
||||||
if internal.temperature is not None:
|
if internal.temperature is not None:
|
||||||
@@ -231,7 +229,7 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
thinking_config = google_extra.get("thinking_config")
|
thinking_config = google_extra.get("thinking_config")
|
||||||
if isinstance(thinking_config, dict):
|
if isinstance(thinking_config, dict):
|
||||||
# snake_case -> camelCase 转换
|
# snake_case -> camelCase 转换
|
||||||
gemini_thinking: Dict[str, Any] = {}
|
gemini_thinking: dict[str, Any] = {}
|
||||||
if "thinking_budget" in thinking_config:
|
if "thinking_budget" in thinking_config:
|
||||||
gemini_thinking["thinkingBudget"] = thinking_config["thinking_budget"]
|
gemini_thinking["thinkingBudget"] = thinking_config["thinking_budget"]
|
||||||
if "include_thoughts" in thinking_config:
|
if "include_thoughts" in thinking_config:
|
||||||
@@ -265,11 +263,11 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
if "thinking_config" in orig_gc and "thinkingConfig" not in generation_config:
|
if "thinking_config" in orig_gc and "thinkingConfig" not in generation_config:
|
||||||
generation_config["thinkingConfig"] = orig_gc["thinking_config"]
|
generation_config["thinkingConfig"] = orig_gc["thinking_config"]
|
||||||
|
|
||||||
contents: List[Dict[str, Any]] = []
|
contents: list[dict[str, Any]] = []
|
||||||
for msg in internal.messages:
|
for msg in internal.messages:
|
||||||
contents.append(self._internal_message_to_content(msg))
|
contents.append(self._internal_message_to_content(msg))
|
||||||
|
|
||||||
result: Dict[str, Any] = {
|
result: dict[str, Any] = {
|
||||||
"contents": contents,
|
"contents": contents,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -295,7 +293,7 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
# Responses
|
# Responses
|
||||||
# =========================
|
# =========================
|
||||||
|
|
||||||
def response_to_internal(self, response: Dict[str, Any]) -> InternalResponse:
|
def response_to_internal(self, response: dict[str, Any]) -> InternalResponse:
|
||||||
rid = str(response.get("id") or "")
|
rid = str(response.get("id") or "")
|
||||||
model = str(response.get("modelVersion") or response.get("model") or "")
|
model = str(response.get("modelVersion") or response.get("model") or "")
|
||||||
|
|
||||||
@@ -315,7 +313,7 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
|
|
||||||
usage_info = self._usage_metadata_to_internal(response.get("usageMetadata"))
|
usage_info = self._usage_metadata_to_internal(response.get("usageMetadata"))
|
||||||
|
|
||||||
extra: Dict[str, Any] = {}
|
extra: dict[str, Any] = {}
|
||||||
if finish_reason is not None:
|
if finish_reason is not None:
|
||||||
extra.setdefault("raw", {})["finishReason"] = finish_reason
|
extra.setdefault("raw", {})["finishReason"] = finish_reason
|
||||||
|
|
||||||
@@ -337,9 +335,9 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
self,
|
self,
|
||||||
internal: InternalResponse,
|
internal: InternalResponse,
|
||||||
*,
|
*,
|
||||||
requested_model: Optional[str] = None,
|
requested_model: str | None = None,
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
parts: List[Dict[str, Any]] = []
|
parts: list[dict[str, Any]] = []
|
||||||
for b in internal.content:
|
for b in internal.content:
|
||||||
if isinstance(b, TextBlock):
|
if isinstance(b, TextBlock):
|
||||||
if b.text:
|
if b.text:
|
||||||
@@ -374,7 +372,7 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
if internal.stop_reason is not None:
|
if internal.stop_reason is not None:
|
||||||
finish_reason = STOP_REASON_MAPPINGS.get("GEMINI", {}).get(internal.stop_reason.value, "OTHER")
|
finish_reason = STOP_REASON_MAPPINGS.get("GEMINI", {}).get(internal.stop_reason.value, "OTHER")
|
||||||
|
|
||||||
usage_metadata: Dict[str, Any] = {}
|
usage_metadata: dict[str, Any] = {}
|
||||||
if internal.usage:
|
if internal.usage:
|
||||||
usage_metadata = {
|
usage_metadata = {
|
||||||
"promptTokenCount": int(internal.usage.input_tokens),
|
"promptTokenCount": int(internal.usage.input_tokens),
|
||||||
@@ -384,7 +382,7 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
if internal.usage.cache_read_tokens:
|
if internal.usage.cache_read_tokens:
|
||||||
usage_metadata["cachedContentTokenCount"] = int(internal.usage.cache_read_tokens)
|
usage_metadata["cachedContentTokenCount"] = int(internal.usage.cache_read_tokens)
|
||||||
|
|
||||||
candidate: Dict[str, Any] = {
|
candidate: dict[str, Any] = {
|
||||||
"content": {"parts": parts, "role": "model"},
|
"content": {"parts": parts, "role": "model"},
|
||||||
"index": 0,
|
"index": 0,
|
||||||
}
|
}
|
||||||
@@ -394,7 +392,7 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
# 优先使用用户请求的原始模型名,回退到上游返回的模型名
|
# 优先使用用户请求的原始模型名,回退到上游返回的模型名
|
||||||
model_name = requested_model if requested_model else (internal.model or "gemini")
|
model_name = requested_model if requested_model else (internal.model or "gemini")
|
||||||
|
|
||||||
out: Dict[str, Any] = {
|
out: dict[str, Any] = {
|
||||||
"candidates": [candidate],
|
"candidates": [candidate],
|
||||||
"modelVersion": model_name,
|
"modelVersion": model_name,
|
||||||
}
|
}
|
||||||
@@ -412,9 +410,9 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
# Streaming
|
# Streaming
|
||||||
# =========================
|
# =========================
|
||||||
|
|
||||||
def stream_chunk_to_internal(self, chunk: Dict[str, Any], state: StreamState) -> List[InternalStreamEvent]:
|
def stream_chunk_to_internal(self, chunk: dict[str, Any], state: StreamState) -> list[InternalStreamEvent]:
|
||||||
ss = state.substate(self.FORMAT_ID)
|
ss = state.substate(self.FORMAT_ID)
|
||||||
events: List[InternalStreamEvent] = []
|
events: list[InternalStreamEvent] = []
|
||||||
|
|
||||||
if not ss.get("message_started"):
|
if not ss.get("message_started"):
|
||||||
# 保留初始化时设置的 model(客户端请求的模型),仅在空时用上游值
|
# 保留初始化时设置的 model(客户端请求的模型),仅在空时用上游值
|
||||||
@@ -544,11 +542,11 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
self,
|
self,
|
||||||
event: InternalStreamEvent,
|
event: InternalStreamEvent,
|
||||||
state: StreamState,
|
state: StreamState,
|
||||||
) -> List[Dict[str, Any]]:
|
) -> list[dict[str, Any]]:
|
||||||
ss = state.substate(self.FORMAT_ID)
|
ss = state.substate(self.FORMAT_ID)
|
||||||
out: List[Dict[str, Any]] = []
|
out: list[dict[str, Any]] = []
|
||||||
|
|
||||||
def base_chunk(parts: List[Dict[str, Any]]) -> Dict[str, Any]:
|
def base_chunk(parts: list[dict[str, Any]]) -> dict[str, Any]:
|
||||||
return {
|
return {
|
||||||
"candidates": [
|
"candidates": [
|
||||||
{
|
{
|
||||||
@@ -622,7 +620,7 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
|
|
||||||
name = str(entry.get("name") or "")
|
name = str(entry.get("name") or "")
|
||||||
raw_json = str(entry.get("json") or "")
|
raw_json = str(entry.get("json") or "")
|
||||||
args: Dict[str, Any] = {}
|
args: dict[str, Any] = {}
|
||||||
if raw_json:
|
if raw_json:
|
||||||
try:
|
try:
|
||||||
parsed = json.loads(raw_json)
|
parsed = json.loads(raw_json)
|
||||||
@@ -639,7 +637,7 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
if event.stop_reason is not None:
|
if event.stop_reason is not None:
|
||||||
finish_reason = STOP_REASON_MAPPINGS.get("GEMINI", {}).get(event.stop_reason.value, "OTHER")
|
finish_reason = STOP_REASON_MAPPINGS.get("GEMINI", {}).get(event.stop_reason.value, "OTHER")
|
||||||
|
|
||||||
chunk: Dict[str, Any] = base_chunk([])
|
chunk: dict[str, Any] = base_chunk([])
|
||||||
if finish_reason is not None:
|
if finish_reason is not None:
|
||||||
chunk["candidates"][0]["finishReason"] = finish_reason
|
chunk["candidates"][0]["finishReason"] = finish_reason
|
||||||
|
|
||||||
@@ -665,10 +663,10 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
# Error conversion
|
# Error conversion
|
||||||
# =========================
|
# =========================
|
||||||
|
|
||||||
def is_error_response(self, response: Dict[str, Any]) -> bool:
|
def is_error_response(self, response: dict[str, Any]) -> bool:
|
||||||
return isinstance(response, dict) and "error" in response
|
return isinstance(response, dict) and "error" in response
|
||||||
|
|
||||||
def error_to_internal(self, error_response: Dict[str, Any]) -> InternalError:
|
def error_to_internal(self, error_response: dict[str, Any]) -> InternalError:
|
||||||
err = error_response.get("error") if isinstance(error_response, dict) else None
|
err = error_response.get("error") if isinstance(error_response, dict) else None
|
||||||
err = err if isinstance(err, dict) else {}
|
err = err if isinstance(err, dict) else {}
|
||||||
|
|
||||||
@@ -691,9 +689,9 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
extra={"gemini": {"error": err}, "raw": {"status": raw_status}},
|
extra={"gemini": {"error": err}, "raw": {"status": raw_status}},
|
||||||
)
|
)
|
||||||
|
|
||||||
def error_from_internal(self, internal: InternalError) -> Dict[str, Any]:
|
def error_from_internal(self, internal: InternalError) -> dict[str, Any]:
|
||||||
status = self._ERROR_TYPE_TO_GEMINI_STATUS.get(internal.type, "INTERNAL")
|
status = self._ERROR_TYPE_TO_GEMINI_STATUS.get(internal.type, "INTERNAL")
|
||||||
payload: Dict[str, Any] = {
|
payload: dict[str, Any] = {
|
||||||
"code": 400 if internal.type == ErrorType.INVALID_REQUEST else 500,
|
"code": 400 if internal.type == ErrorType.INVALID_REQUEST else 500,
|
||||||
"message": internal.message,
|
"message": internal.message,
|
||||||
"status": status,
|
"status": status,
|
||||||
@@ -704,8 +702,8 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
# Helpers
|
# Helpers
|
||||||
# =========================
|
# =========================
|
||||||
|
|
||||||
def _content_to_internal_message(self, content: Dict[str, Any]) -> Tuple[Optional[InternalMessage], Dict[str, int]]:
|
def _content_to_internal_message(self, content: dict[str, Any]) -> tuple[InternalMessage | None, dict[str, int]]:
|
||||||
dropped: Dict[str, int] = {}
|
dropped: dict[str, int] = {}
|
||||||
|
|
||||||
role_raw = str(content.get("role") or "user")
|
role_raw = str(content.get("role") or "user")
|
||||||
if role_raw == "model":
|
if role_raw == "model":
|
||||||
@@ -727,15 +725,15 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
dropped,
|
dropped,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _parts_to_blocks(self, parts: Any) -> Tuple[List[ContentBlock], Dict[str, int]]:
|
def _parts_to_blocks(self, parts: Any) -> tuple[list[ContentBlock], dict[str, int]]:
|
||||||
dropped: Dict[str, int] = {}
|
dropped: dict[str, int] = {}
|
||||||
if parts is None:
|
if parts is None:
|
||||||
return [], dropped
|
return [], dropped
|
||||||
if not isinstance(parts, list):
|
if not isinstance(parts, list):
|
||||||
dropped["gemini_parts_non_list"] = dropped.get("gemini_parts_non_list", 0) + 1
|
dropped["gemini_parts_non_list"] = dropped.get("gemini_parts_non_list", 0) + 1
|
||||||
return [], dropped
|
return [], dropped
|
||||||
|
|
||||||
blocks: List[ContentBlock] = []
|
blocks: list[ContentBlock] = []
|
||||||
for part in parts:
|
for part in parts:
|
||||||
if not isinstance(part, dict):
|
if not isinstance(part, dict):
|
||||||
dropped["gemini_part_non_dict"] = dropped.get("gemini_part_non_dict", 0) + 1
|
dropped["gemini_part_non_dict"] = dropped.get("gemini_part_non_dict", 0) + 1
|
||||||
@@ -785,7 +783,7 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
name = str(func_resp.get("name") or "")
|
name = str(func_resp.get("name") or "")
|
||||||
response = func_resp.get("response")
|
response = func_resp.get("response")
|
||||||
output: Any = None
|
output: Any = None
|
||||||
content_text: Optional[str] = None
|
content_text: str | None = None
|
||||||
|
|
||||||
# 兼容历史:response 常见结构为 {"result": ...}
|
# 兼容历史:response 常见结构为 {"result": ...}
|
||||||
if isinstance(response, dict) and "result" in response:
|
if isinstance(response, dict) and "result" in response:
|
||||||
@@ -815,10 +813,10 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
|
|
||||||
return blocks, dropped
|
return blocks, dropped
|
||||||
|
|
||||||
def _internal_message_to_content(self, msg: InternalMessage) -> Dict[str, Any]:
|
def _internal_message_to_content(self, msg: InternalMessage) -> dict[str, Any]:
|
||||||
role = "model" if msg.role == Role.ASSISTANT else "user"
|
role = "model" if msg.role == Role.ASSISTANT else "user"
|
||||||
|
|
||||||
parts: List[Dict[str, Any]] = []
|
parts: list[dict[str, Any]] = []
|
||||||
for b in msg.content:
|
for b in msg.content:
|
||||||
if isinstance(b, UnknownBlock):
|
if isinstance(b, UnknownBlock):
|
||||||
continue
|
continue
|
||||||
@@ -861,8 +859,8 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
|
|
||||||
return {"role": role, "parts": parts}
|
return {"role": role, "parts": parts}
|
||||||
|
|
||||||
def _collapse_system_instruction(self, system_instruction: Any) -> Tuple[Optional[str], Dict[str, int]]:
|
def _collapse_system_instruction(self, system_instruction: Any) -> tuple[str | None, dict[str, int]]:
|
||||||
dropped: Dict[str, int] = {}
|
dropped: dict[str, int] = {}
|
||||||
if system_instruction is None:
|
if system_instruction is None:
|
||||||
return None, dropped
|
return None, dropped
|
||||||
|
|
||||||
@@ -870,7 +868,7 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
if isinstance(system_instruction, dict):
|
if isinstance(system_instruction, dict):
|
||||||
parts = system_instruction.get("parts")
|
parts = system_instruction.get("parts")
|
||||||
if isinstance(parts, list):
|
if isinstance(parts, list):
|
||||||
texts: List[str] = []
|
texts: list[str] = []
|
||||||
for part in parts:
|
for part in parts:
|
||||||
if isinstance(part, dict) and "text" in part and part.get("text"):
|
if isinstance(part, dict) and "text" in part and part.get("text"):
|
||||||
texts.append(str(part.get("text")))
|
texts.append(str(part.get("text")))
|
||||||
@@ -880,7 +878,7 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
dropped["gemini_system_instruction_unsupported"] = dropped.get("gemini_system_instruction_unsupported", 0) + 1
|
dropped["gemini_system_instruction_unsupported"] = dropped.get("gemini_system_instruction_unsupported", 0) + 1
|
||||||
return None, dropped
|
return None, dropped
|
||||||
|
|
||||||
def _get_generation_config(self, request: Dict[str, Any]) -> Dict[str, Any]:
|
def _get_generation_config(self, request: dict[str, Any]) -> dict[str, Any]:
|
||||||
# 兼容 snake_case 与 camelCase
|
# 兼容 snake_case 与 camelCase
|
||||||
gc = request.get("generation_config") if "generation_config" in request else request.get("generationConfig")
|
gc = request.get("generation_config") if "generation_config" in request else request.get("generationConfig")
|
||||||
if not isinstance(gc, dict):
|
if not isinstance(gc, dict):
|
||||||
@@ -893,7 +891,7 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
return gc.get(k)
|
return gc.get(k)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
normalized: Dict[str, Any] = {}
|
normalized: dict[str, Any] = {}
|
||||||
normalized["max_output_tokens"] = pick("max_output_tokens", "maxOutputTokens")
|
normalized["max_output_tokens"] = pick("max_output_tokens", "maxOutputTokens")
|
||||||
normalized["temperature"] = pick("temperature")
|
normalized["temperature"] = pick("temperature")
|
||||||
normalized["top_p"] = pick("top_p", "topP")
|
normalized["top_p"] = pick("top_p", "topP")
|
||||||
@@ -912,11 +910,11 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
|
|
||||||
return {k: v for k, v in normalized.items() if v is not None}
|
return {k: v for k, v in normalized.items() if v is not None}
|
||||||
|
|
||||||
def _gemini_tools_to_internal(self, tools: Any) -> Optional[List[ToolDefinition]]:
|
def _gemini_tools_to_internal(self, tools: Any) -> list[ToolDefinition] | None:
|
||||||
if not tools or not isinstance(tools, list):
|
if not tools or not isinstance(tools, list):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
out: List[ToolDefinition] = []
|
out: list[ToolDefinition] = []
|
||||||
for tool in tools:
|
for tool in tools:
|
||||||
if not isinstance(tool, dict):
|
if not isinstance(tool, dict):
|
||||||
continue
|
continue
|
||||||
@@ -945,7 +943,7 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
|
|
||||||
return out or None
|
return out or None
|
||||||
|
|
||||||
def _gemini_tool_config_to_tool_choice(self, tool_config: Any) -> Optional[ToolChoice]:
|
def _gemini_tool_config_to_tool_choice(self, tool_config: Any) -> ToolChoice | None:
|
||||||
if tool_config is None:
|
if tool_config is None:
|
||||||
return None
|
return None
|
||||||
if not isinstance(tool_config, dict):
|
if not isinstance(tool_config, dict):
|
||||||
@@ -972,9 +970,9 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
|
|
||||||
return ToolChoice(type=ToolChoiceType.AUTO, extra={"gemini": tool_config})
|
return ToolChoice(type=ToolChoiceType.AUTO, extra={"gemini": tool_config})
|
||||||
|
|
||||||
def _tool_choice_to_gemini_tool_config(self, tool_choice: ToolChoice) -> Dict[str, Any]:
|
def _tool_choice_to_gemini_tool_config(self, tool_choice: ToolChoice) -> dict[str, Any]:
|
||||||
mode = "AUTO"
|
mode = "AUTO"
|
||||||
cfg: Dict[str, Any] = {}
|
cfg: dict[str, Any] = {}
|
||||||
|
|
||||||
if tool_choice.type == ToolChoiceType.NONE:
|
if tool_choice.type == ToolChoiceType.NONE:
|
||||||
mode = "NONE"
|
mode = "NONE"
|
||||||
@@ -987,12 +985,12 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
cfg["mode"] = mode
|
cfg["mode"] = mode
|
||||||
return {"function_calling_config": cfg}
|
return {"function_calling_config": cfg}
|
||||||
|
|
||||||
def _usage_metadata_to_internal(self, usage_metadata: Any) -> Optional[UsageInfo]:
|
def _usage_metadata_to_internal(self, usage_metadata: Any) -> UsageInfo | None:
|
||||||
if not isinstance(usage_metadata, dict):
|
if not isinstance(usage_metadata, dict):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
mapping = USAGE_FIELD_MAPPINGS.get("GEMINI", {})
|
mapping = USAGE_FIELD_MAPPINGS.get("GEMINI", {})
|
||||||
fields: Dict[str, int] = {}
|
fields: dict[str, int] = {}
|
||||||
extra = self._extract_extra(usage_metadata, set(mapping.keys()))
|
extra = self._extract_extra(usage_metadata, set(mapping.keys()))
|
||||||
|
|
||||||
# promptTokenCount/candidatesTokenCount/totalTokenCount/cachedContentTokenCount
|
# promptTokenCount/candidatesTokenCount/totalTokenCount/cachedContentTokenCount
|
||||||
@@ -1023,7 +1021,7 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
extra={"gemini": extra} if extra else {},
|
extra={"gemini": extra} if extra else {},
|
||||||
)
|
)
|
||||||
|
|
||||||
def _join_instructions(self, instructions: List[InstructionSegment]) -> Optional[str]:
|
def _join_instructions(self, instructions: list[InstructionSegment]) -> str | None:
|
||||||
parts = [seg.text for seg in instructions if seg.text]
|
parts = [seg.text for seg in instructions if seg.text]
|
||||||
joined = "\n\n".join(parts)
|
joined = "\n\n".join(parts)
|
||||||
return joined or None
|
return joined or None
|
||||||
@@ -1034,7 +1032,7 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
except ValueError:
|
except ValueError:
|
||||||
return ErrorType.UNKNOWN
|
return ErrorType.UNKNOWN
|
||||||
|
|
||||||
def _optional_int(self, value: Any) -> Optional[int]:
|
def _optional_int(self, value: Any) -> int | None:
|
||||||
if value is None:
|
if value is None:
|
||||||
return None
|
return None
|
||||||
try:
|
try:
|
||||||
@@ -1042,7 +1040,7 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
except (TypeError, ValueError):
|
except (TypeError, ValueError):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def _optional_float(self, value: Any) -> Optional[float]:
|
def _optional_float(self, value: Any) -> float | None:
|
||||||
if value is None:
|
if value is None:
|
||||||
return None
|
return None
|
||||||
try:
|
try:
|
||||||
@@ -1050,7 +1048,7 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
except (TypeError, ValueError):
|
except (TypeError, ValueError):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def _coerce_str_list(self, value: Any) -> Optional[List[str]]:
|
def _coerce_str_list(self, value: Any) -> list[str] | None:
|
||||||
if value is None:
|
if value is None:
|
||||||
return None
|
return None
|
||||||
if isinstance(value, str):
|
if isinstance(value, str):
|
||||||
@@ -1059,10 +1057,10 @@ class GeminiNormalizer(FormatNormalizer):
|
|||||||
return [str(x) for x in value if x is not None]
|
return [str(x) for x in value if x is not None]
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def _extract_extra(self, payload: Dict[str, Any], known_keys: set[str]) -> Dict[str, Any]:
|
def _extract_extra(self, payload: dict[str, Any], known_keys: set[str]) -> dict[str, Any]:
|
||||||
return {k: v for k, v in payload.items() if k not in known_keys}
|
return {k: v for k, v in payload.items() if k not in known_keys}
|
||||||
|
|
||||||
def _merge_dropped(self, target: Dict[str, int], source: Dict[str, int]) -> None:
|
def _merge_dropped(self, target: dict[str, int], source: dict[str, int]) -> None:
|
||||||
for k, v in source.items():
|
for k, v in source.items():
|
||||||
target[k] = target.get(k, 0) + int(v)
|
target[k] = target.get(k, 0) + int(v)
|
||||||
|
|
||||||
|
|||||||
@@ -7,7 +7,6 @@ GEMINI_CLI 的请求/响应 body 与 GEMINI 一致(Google Gemini API),差
|
|||||||
如需 CLI 特殊处理,可覆盖 request_from_internal / request_to_internal 等方法。
|
如需 CLI 特殊处理,可覆盖 request_from_internal / request_to_internal 等方法。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from src.core.api_format.conversion.normalizers.gemini import GeminiNormalizer
|
from src.core.api_format.conversion.normalizers.gemini import GeminiNormalizer
|
||||||
|
|
||||||
|
|||||||
@@ -7,11 +7,10 @@ OpenAI Chat Completions Normalizer
|
|||||||
- 可选:OpenAI error <-> InternalError
|
- 可选:OpenAI error <-> InternalError
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import json
|
import json
|
||||||
import time
|
import time
|
||||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
from typing import Any
|
||||||
|
|
||||||
from src.core.logger import logger
|
from src.core.logger import logger
|
||||||
from src.core.api_format.conversion.field_mappings import (
|
from src.core.api_format.conversion.field_mappings import (
|
||||||
@@ -49,7 +48,6 @@ from src.core.api_format.conversion.stream_events import (
|
|||||||
InternalStreamEvent,
|
InternalStreamEvent,
|
||||||
MessageStartEvent,
|
MessageStartEvent,
|
||||||
MessageStopEvent,
|
MessageStopEvent,
|
||||||
StreamEventType,
|
|
||||||
ToolCallDeltaEvent,
|
ToolCallDeltaEvent,
|
||||||
)
|
)
|
||||||
from src.core.api_format.conversion.stream_state import StreamState
|
from src.core.api_format.conversion.stream_state import StreamState
|
||||||
@@ -65,7 +63,7 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# finish_reason -> StopReason
|
# finish_reason -> StopReason
|
||||||
_FINISH_REASON_TO_STOP: Dict[str, StopReason] = {
|
_FINISH_REASON_TO_STOP: dict[str, StopReason] = {
|
||||||
"stop": StopReason.END_TURN,
|
"stop": StopReason.END_TURN,
|
||||||
"length": StopReason.MAX_TOKENS,
|
"length": StopReason.MAX_TOKENS,
|
||||||
"tool_calls": StopReason.TOOL_USE,
|
"tool_calls": StopReason.TOOL_USE,
|
||||||
@@ -74,7 +72,7 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
}
|
}
|
||||||
|
|
||||||
# StopReason -> finish_reason
|
# StopReason -> finish_reason
|
||||||
_STOP_TO_FINISH_REASON: Dict[StopReason, str] = {
|
_STOP_TO_FINISH_REASON: dict[StopReason, str] = {
|
||||||
StopReason.END_TURN: "stop",
|
StopReason.END_TURN: "stop",
|
||||||
StopReason.MAX_TOKENS: "length",
|
StopReason.MAX_TOKENS: "length",
|
||||||
StopReason.STOP_SEQUENCE: "stop",
|
StopReason.STOP_SEQUENCE: "stop",
|
||||||
@@ -84,7 +82,7 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
}
|
}
|
||||||
|
|
||||||
# InternalError.type -> OpenAI error.type(最佳努力)
|
# InternalError.type -> OpenAI error.type(最佳努力)
|
||||||
_ERROR_TYPE_TO_OPENAI: Dict[ErrorType, str] = {
|
_ERROR_TYPE_TO_OPENAI: dict[ErrorType, str] = {
|
||||||
ErrorType.INVALID_REQUEST: "invalid_request_error",
|
ErrorType.INVALID_REQUEST: "invalid_request_error",
|
||||||
ErrorType.AUTHENTICATION: "invalid_api_key",
|
ErrorType.AUTHENTICATION: "invalid_api_key",
|
||||||
ErrorType.PERMISSION_DENIED: "invalid_request_error",
|
ErrorType.PERMISSION_DENIED: "invalid_request_error",
|
||||||
@@ -101,13 +99,13 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
# Requests
|
# Requests
|
||||||
# =========================
|
# =========================
|
||||||
|
|
||||||
def request_to_internal(self, request: Dict[str, Any]) -> InternalRequest:
|
def request_to_internal(self, request: dict[str, Any]) -> InternalRequest:
|
||||||
model = str(request.get("model") or "")
|
model = str(request.get("model") or "")
|
||||||
|
|
||||||
dropped: Dict[str, int] = {}
|
dropped: dict[str, int] = {}
|
||||||
|
|
||||||
instructions: List[InstructionSegment] = []
|
instructions: list[InstructionSegment] = []
|
||||||
messages: List[InternalMessage] = []
|
messages: list[InternalMessage] = []
|
||||||
|
|
||||||
for msg in request.get("messages") or []:
|
for msg in request.get("messages") or []:
|
||||||
if not isinstance(msg, dict):
|
if not isinstance(msg, dict):
|
||||||
@@ -146,7 +144,7 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# 构建 extra,保留未识别字段
|
# 构建 extra,保留未识别字段
|
||||||
extra: Dict[str, Any] = {"openai": self._extract_extra(request, {"messages"})}
|
extra: dict[str, Any] = {"openai": self._extract_extra(request, {"messages"})}
|
||||||
|
|
||||||
# 处理 extra_body.google (用于 Gemini 特定功能透传,如 thinkingConfig, responseModalities)
|
# 处理 extra_body.google (用于 Gemini 特定功能透传,如 thinkingConfig, responseModalities)
|
||||||
extra_body = request.get("extra_body")
|
extra_body = request.get("extra_body")
|
||||||
@@ -175,8 +173,8 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
|
|
||||||
return internal
|
return internal
|
||||||
|
|
||||||
def request_from_internal(self, internal: InternalRequest) -> Dict[str, Any]:
|
def request_from_internal(self, internal: InternalRequest) -> dict[str, Any]:
|
||||||
out_messages: List[Dict[str, Any]] = []
|
out_messages: list[dict[str, Any]] = []
|
||||||
|
|
||||||
if internal.instructions:
|
if internal.instructions:
|
||||||
for seg in internal.instructions:
|
for seg in internal.instructions:
|
||||||
@@ -189,7 +187,7 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
for msg in internal.messages:
|
for msg in internal.messages:
|
||||||
out_messages.extend(self._internal_message_to_openai_messages(msg))
|
out_messages.extend(self._internal_message_to_openai_messages(msg))
|
||||||
|
|
||||||
result: Dict[str, Any] = {
|
result: dict[str, Any] = {
|
||||||
"model": internal.model,
|
"model": internal.model,
|
||||||
"messages": out_messages,
|
"messages": out_messages,
|
||||||
}
|
}
|
||||||
@@ -233,11 +231,11 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
# Responses
|
# Responses
|
||||||
# =========================
|
# =========================
|
||||||
|
|
||||||
def response_to_internal(self, response: Dict[str, Any]) -> InternalResponse:
|
def response_to_internal(self, response: dict[str, Any]) -> InternalResponse:
|
||||||
rid = str(response.get("id") or "")
|
rid = str(response.get("id") or "")
|
||||||
model = str(response.get("model") or "")
|
model = str(response.get("model") or "")
|
||||||
|
|
||||||
extra: Dict[str, Any] = {}
|
extra: dict[str, Any] = {}
|
||||||
|
|
||||||
choices = response.get("choices") or []
|
choices = response.get("choices") or []
|
||||||
if isinstance(choices, list) and len(choices) > 1:
|
if isinstance(choices, list) and len(choices) > 1:
|
||||||
@@ -285,12 +283,12 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
self,
|
self,
|
||||||
internal: InternalResponse,
|
internal: InternalResponse,
|
||||||
*,
|
*,
|
||||||
requested_model: Optional[str] = None,
|
requested_model: str | None = None,
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
# OpenAI Chat Completions response envelope
|
# OpenAI Chat Completions response envelope
|
||||||
# 优先使用用户请求的原始模型名,回退到上游返回的模型名
|
# 优先使用用户请求的原始模型名,回退到上游返回的模型名
|
||||||
model_name = requested_model if requested_model else internal.model
|
model_name = requested_model if requested_model else internal.model
|
||||||
out: Dict[str, Any] = {
|
out: dict[str, Any] = {
|
||||||
"id": internal.id or "chatcmpl-unknown",
|
"id": internal.id or "chatcmpl-unknown",
|
||||||
"object": "chat.completion",
|
"object": "chat.completion",
|
||||||
"created": int(time.time()),
|
"created": int(time.time()),
|
||||||
@@ -298,7 +296,7 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
"choices": [],
|
"choices": [],
|
||||||
}
|
}
|
||||||
|
|
||||||
message: Dict[str, Any] = {"role": "assistant"}
|
message: dict[str, Any] = {"role": "assistant"}
|
||||||
|
|
||||||
content_blocks, tool_blocks = self._split_blocks(internal.content)
|
content_blocks, tool_blocks = self._split_blocks(internal.content)
|
||||||
content_value = self._blocks_to_openai_content(content_blocks)
|
content_value = self._blocks_to_openai_content(content_blocks)
|
||||||
@@ -336,9 +334,9 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
# Streaming
|
# Streaming
|
||||||
# =========================
|
# =========================
|
||||||
|
|
||||||
def stream_chunk_to_internal(self, chunk: Dict[str, Any], state: StreamState) -> List[InternalStreamEvent]:
|
def stream_chunk_to_internal(self, chunk: dict[str, Any], state: StreamState) -> list[InternalStreamEvent]:
|
||||||
ss = state.substate(self.FORMAT_ID)
|
ss = state.substate(self.FORMAT_ID)
|
||||||
events: List[InternalStreamEvent] = []
|
events: list[InternalStreamEvent] = []
|
||||||
|
|
||||||
# OpenAI streaming error(通常是单个 {"error": {...}})
|
# OpenAI streaming error(通常是单个 {"error": {...}})
|
||||||
if isinstance(chunk, dict) and "error" in chunk:
|
if isinstance(chunk, dict) and "error" in chunk:
|
||||||
@@ -437,11 +435,11 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
self,
|
self,
|
||||||
event: InternalStreamEvent,
|
event: InternalStreamEvent,
|
||||||
state: StreamState,
|
state: StreamState,
|
||||||
) -> List[Dict[str, Any]]:
|
) -> list[dict[str, Any]]:
|
||||||
ss = state.substate(self.FORMAT_ID)
|
ss = state.substate(self.FORMAT_ID)
|
||||||
out: List[Dict[str, Any]] = []
|
out: list[dict[str, Any]] = []
|
||||||
|
|
||||||
def base_chunk(delta: Dict[str, Any], finish_reason: Optional[str] = None) -> Dict[str, Any]:
|
def base_chunk(delta: dict[str, Any], finish_reason: str | None = None) -> dict[str, Any]:
|
||||||
return {
|
return {
|
||||||
"id": state.message_id or "chatcmpl-stream",
|
"id": state.message_id or "chatcmpl-stream",
|
||||||
"object": "chat.completion.chunk",
|
"object": "chat.completion.chunk",
|
||||||
@@ -570,10 +568,10 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
# Error conversion
|
# Error conversion
|
||||||
# =========================
|
# =========================
|
||||||
|
|
||||||
def is_error_response(self, response: Dict[str, Any]) -> bool:
|
def is_error_response(self, response: dict[str, Any]) -> bool:
|
||||||
return isinstance(response, dict) and "error" in response
|
return isinstance(response, dict) and "error" in response
|
||||||
|
|
||||||
def error_to_internal(self, error_response: Dict[str, Any]) -> InternalError:
|
def error_to_internal(self, error_response: dict[str, Any]) -> InternalError:
|
||||||
err = error_response.get("error") if isinstance(error_response, dict) else None
|
err = error_response.get("error") if isinstance(error_response, dict) else None
|
||||||
err = err if isinstance(err, dict) else {}
|
err = err if isinstance(err, dict) else {}
|
||||||
|
|
||||||
@@ -592,9 +590,9 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
extra={"openai": {"error": err}, "raw": {"type": raw_type}},
|
extra={"openai": {"error": err}, "raw": {"type": raw_type}},
|
||||||
)
|
)
|
||||||
|
|
||||||
def error_from_internal(self, internal: InternalError) -> Dict[str, Any]:
|
def error_from_internal(self, internal: InternalError) -> dict[str, Any]:
|
||||||
type_str = self._ERROR_TYPE_TO_OPENAI.get(internal.type, "server_error")
|
type_str = self._ERROR_TYPE_TO_OPENAI.get(internal.type, "server_error")
|
||||||
payload: Dict[str, Any] = {
|
payload: dict[str, Any] = {
|
||||||
"message": internal.message,
|
"message": internal.message,
|
||||||
"type": type_str,
|
"type": type_str,
|
||||||
}
|
}
|
||||||
@@ -608,8 +606,8 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
# Helpers
|
# Helpers
|
||||||
# =========================
|
# =========================
|
||||||
|
|
||||||
def _openai_message_to_internal(self, msg: Dict[str, Any]) -> Tuple[Optional[InternalMessage], Dict[str, int]]:
|
def _openai_message_to_internal(self, msg: dict[str, Any]) -> tuple[InternalMessage | None, dict[str, int]]:
|
||||||
dropped: Dict[str, int] = {}
|
dropped: dict[str, int] = {}
|
||||||
|
|
||||||
role_raw = str(msg.get("role") or "unknown")
|
role_raw = str(msg.get("role") or "unknown")
|
||||||
role = self._role_from_openai(role_raw)
|
role = self._role_from_openai(role_raw)
|
||||||
@@ -656,8 +654,8 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
dropped,
|
dropped,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _openai_content_to_blocks(self, content: Any) -> Tuple[List[ContentBlock], Dict[str, int]]:
|
def _openai_content_to_blocks(self, content: Any) -> tuple[list[ContentBlock], dict[str, int]]:
|
||||||
dropped: Dict[str, int] = {}
|
dropped: dict[str, int] = {}
|
||||||
|
|
||||||
if content is None:
|
if content is None:
|
||||||
return [], dropped
|
return [], dropped
|
||||||
@@ -667,7 +665,7 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
dropped["openai_content_non_list"] = dropped.get("openai_content_non_list", 0) + 1
|
dropped["openai_content_non_list"] = dropped.get("openai_content_non_list", 0) + 1
|
||||||
return [], dropped
|
return [], dropped
|
||||||
|
|
||||||
blocks: List[ContentBlock] = []
|
blocks: list[ContentBlock] = []
|
||||||
for part in content:
|
for part in content:
|
||||||
if not isinstance(part, dict):
|
if not isinstance(part, dict):
|
||||||
dropped["openai_content_part_non_dict"] = dropped.get("openai_content_part_non_dict", 0) + 1
|
dropped["openai_content_part_non_dict"] = dropped.get("openai_content_part_non_dict", 0) + 1
|
||||||
@@ -698,21 +696,21 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
|
|
||||||
return blocks, dropped
|
return blocks, dropped
|
||||||
|
|
||||||
def _collapse_openai_text(self, content: Any) -> Tuple[str, Dict[str, int]]:
|
def _collapse_openai_text(self, content: Any) -> tuple[str, dict[str, int]]:
|
||||||
blocks, dropped = self._openai_content_to_blocks(content)
|
blocks, dropped = self._openai_content_to_blocks(content)
|
||||||
text_parts = [b.text for b in blocks if isinstance(b, TextBlock) and b.text]
|
text_parts = [b.text for b in blocks if isinstance(b, TextBlock) and b.text]
|
||||||
return ("\n\n".join(text_parts), dropped)
|
return ("\n\n".join(text_parts), dropped)
|
||||||
|
|
||||||
def _join_instructions(self, instructions: List[InstructionSegment]) -> Optional[str]:
|
def _join_instructions(self, instructions: list[InstructionSegment]) -> str | None:
|
||||||
parts = [seg.text for seg in instructions if seg.text]
|
parts = [seg.text for seg in instructions if seg.text]
|
||||||
joined = "\n\n".join(parts)
|
joined = "\n\n".join(parts)
|
||||||
return joined or None
|
return joined or None
|
||||||
|
|
||||||
def _openai_tools_to_internal(self, tools: Any) -> Optional[List[ToolDefinition]]:
|
def _openai_tools_to_internal(self, tools: Any) -> list[ToolDefinition] | None:
|
||||||
if not tools or not isinstance(tools, list):
|
if not tools or not isinstance(tools, list):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
out: List[ToolDefinition] = []
|
out: list[ToolDefinition] = []
|
||||||
for tool in tools:
|
for tool in tools:
|
||||||
if not isinstance(tool, dict):
|
if not isinstance(tool, dict):
|
||||||
continue
|
continue
|
||||||
@@ -720,7 +718,7 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
continue
|
continue
|
||||||
|
|
||||||
function_raw = tool.get("function")
|
function_raw = tool.get("function")
|
||||||
function: Dict[str, Any] = function_raw if isinstance(function_raw, dict) else {}
|
function: dict[str, Any] = function_raw if isinstance(function_raw, dict) else {}
|
||||||
name = str(function.get("name") or "")
|
name = str(function.get("name") or "")
|
||||||
if not name:
|
if not name:
|
||||||
continue
|
continue
|
||||||
@@ -740,7 +738,7 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
|
|
||||||
return out or None
|
return out or None
|
||||||
|
|
||||||
def _openai_tool_choice_to_internal(self, tool_choice: Any) -> Optional[ToolChoice]:
|
def _openai_tool_choice_to_internal(self, tool_choice: Any) -> ToolChoice | None:
|
||||||
if tool_choice is None:
|
if tool_choice is None:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@@ -756,13 +754,13 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
|
|
||||||
if isinstance(tool_choice, dict) and tool_choice.get("type") == "function":
|
if isinstance(tool_choice, dict) and tool_choice.get("type") == "function":
|
||||||
fn_raw = tool_choice.get("function")
|
fn_raw = tool_choice.get("function")
|
||||||
fn: Dict[str, Any] = fn_raw if isinstance(fn_raw, dict) else {}
|
fn: dict[str, Any] = fn_raw if isinstance(fn_raw, dict) else {}
|
||||||
name = str(fn.get("name") or "")
|
name = str(fn.get("name") or "")
|
||||||
return ToolChoice(type=ToolChoiceType.TOOL, tool_name=name, extra={"openai": tool_choice})
|
return ToolChoice(type=ToolChoiceType.TOOL, tool_name=name, extra={"openai": tool_choice})
|
||||||
|
|
||||||
return ToolChoice(type=ToolChoiceType.AUTO, extra={"openai": tool_choice})
|
return ToolChoice(type=ToolChoiceType.AUTO, extra={"openai": tool_choice})
|
||||||
|
|
||||||
def _tool_choice_to_openai(self, tool_choice: ToolChoice) -> Union[str, Dict[str, Any]]:
|
def _tool_choice_to_openai(self, tool_choice: ToolChoice) -> str | dict[str, Any]:
|
||||||
if tool_choice.type == ToolChoiceType.NONE:
|
if tool_choice.type == ToolChoiceType.NONE:
|
||||||
return "none"
|
return "none"
|
||||||
if tool_choice.type == ToolChoiceType.AUTO:
|
if tool_choice.type == ToolChoiceType.AUTO:
|
||||||
@@ -773,8 +771,8 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
return {"type": "function", "function": {"name": tool_choice.tool_name or ""}}
|
return {"type": "function", "function": {"name": tool_choice.tool_name or ""}}
|
||||||
return "auto"
|
return "auto"
|
||||||
|
|
||||||
def _openai_tool_call_to_block(self, tool_call: Any) -> Tuple[Optional[ToolUseBlock], Dict[str, int]]:
|
def _openai_tool_call_to_block(self, tool_call: Any) -> tuple[ToolUseBlock | None, dict[str, int]]:
|
||||||
dropped: Dict[str, int] = {}
|
dropped: dict[str, int] = {}
|
||||||
if not isinstance(tool_call, dict):
|
if not isinstance(tool_call, dict):
|
||||||
dropped["openai_tool_call_non_dict"] = dropped.get("openai_tool_call_non_dict", 0) + 1
|
dropped["openai_tool_call_non_dict"] = dropped.get("openai_tool_call_non_dict", 0) + 1
|
||||||
return None, dropped
|
return None, dropped
|
||||||
@@ -785,12 +783,12 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
return None, dropped
|
return None, dropped
|
||||||
|
|
||||||
fn_raw = tool_call.get("function")
|
fn_raw = tool_call.get("function")
|
||||||
fn: Dict[str, Any] = fn_raw if isinstance(fn_raw, dict) else {}
|
fn: dict[str, Any] = fn_raw if isinstance(fn_raw, dict) else {}
|
||||||
name = str(fn.get("name") or "")
|
name = str(fn.get("name") or "")
|
||||||
args_str = str(fn.get("arguments") or "")
|
args_str = str(fn.get("arguments") or "")
|
||||||
tool_id = str(tool_call.get("id") or "")
|
tool_id = str(tool_call.get("id") or "")
|
||||||
|
|
||||||
tool_input: Dict[str, Any]
|
tool_input: dict[str, Any]
|
||||||
if args_str:
|
if args_str:
|
||||||
try:
|
try:
|
||||||
parsed = json.loads(args_str)
|
parsed = json.loads(args_str)
|
||||||
@@ -810,15 +808,15 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
dropped,
|
dropped,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _legacy_function_call_to_block(self, func_call: Dict[str, Any]) -> Tuple[Optional[ToolUseBlock], Dict[str, int]]:
|
def _legacy_function_call_to_block(self, func_call: dict[str, Any]) -> tuple[ToolUseBlock | None, dict[str, int]]:
|
||||||
dropped: Dict[str, int] = {}
|
dropped: dict[str, int] = {}
|
||||||
name = str(func_call.get("name") or "")
|
name = str(func_call.get("name") or "")
|
||||||
args_str = str(func_call.get("arguments") or "")
|
args_str = str(func_call.get("arguments") or "")
|
||||||
if not name:
|
if not name:
|
||||||
dropped["openai_function_call_missing_name"] = dropped.get("openai_function_call_missing_name", 0) + 1
|
dropped["openai_function_call_missing_name"] = dropped.get("openai_function_call_missing_name", 0) + 1
|
||||||
return None, dropped
|
return None, dropped
|
||||||
|
|
||||||
tool_input: Dict[str, Any]
|
tool_input: dict[str, Any]
|
||||||
if args_str:
|
if args_str:
|
||||||
try:
|
try:
|
||||||
parsed = json.loads(args_str)
|
parsed = json.loads(args_str)
|
||||||
@@ -840,10 +838,10 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
|
|
||||||
def _openai_tool_result_message_to_block(
|
def _openai_tool_result_message_to_block(
|
||||||
self,
|
self,
|
||||||
msg: Dict[str, Any],
|
msg: dict[str, Any],
|
||||||
tool_call_id: str,
|
tool_call_id: str,
|
||||||
) -> Tuple[Optional[ToolResultBlock], Dict[str, int]]:
|
) -> tuple[ToolResultBlock | None, dict[str, int]]:
|
||||||
dropped: Dict[str, int] = {}
|
dropped: dict[str, int] = {}
|
||||||
content = msg.get("content")
|
content = msg.get("content")
|
||||||
if content is None:
|
if content is None:
|
||||||
return ToolResultBlock(tool_use_id=tool_call_id, output=None, content_text=None, extra={"openai": msg}), dropped
|
return ToolResultBlock(tool_use_id=tool_call_id, output=None, content_text=None, extra={"openai": msg}), dropped
|
||||||
@@ -887,7 +885,7 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
dropped,
|
dropped,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _openai_usage_to_internal(self, usage: Any) -> Optional[UsageInfo]:
|
def _openai_usage_to_internal(self, usage: Any) -> UsageInfo | None:
|
||||||
if not isinstance(usage, dict):
|
if not isinstance(usage, dict):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@@ -908,9 +906,9 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
extra={"openai": extra} if extra else {},
|
extra={"openai": extra} if extra else {},
|
||||||
)
|
)
|
||||||
|
|
||||||
def _blocks_to_openai_content(self, blocks: List[ContentBlock]) -> Optional[Union[str, List[Dict[str, Any]]]]:
|
def _blocks_to_openai_content(self, blocks: list[ContentBlock]) -> str | list[dict[str, Any]] | None:
|
||||||
parts: List[Dict[str, Any]] = []
|
parts: list[dict[str, Any]] = []
|
||||||
text_parts: List[str] = []
|
text_parts: list[str] = []
|
||||||
|
|
||||||
for b in blocks:
|
for b in blocks:
|
||||||
if isinstance(b, TextBlock):
|
if isinstance(b, TextBlock):
|
||||||
@@ -948,9 +946,9 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
# OpenAI content 可以是空字符串;但作为响应 message.content 通常允许为 ""/None。
|
# OpenAI content 可以是空字符串;但作为响应 message.content 通常允许为 ""/None。
|
||||||
return ""
|
return ""
|
||||||
|
|
||||||
def _split_blocks(self, blocks: List[ContentBlock]) -> Tuple[List[ContentBlock], List[ToolUseBlock]]:
|
def _split_blocks(self, blocks: list[ContentBlock]) -> tuple[list[ContentBlock], list[ToolUseBlock]]:
|
||||||
content_blocks: List[ContentBlock] = []
|
content_blocks: list[ContentBlock] = []
|
||||||
tool_blocks: List[ToolUseBlock] = []
|
tool_blocks: list[ToolUseBlock] = []
|
||||||
for b in blocks:
|
for b in blocks:
|
||||||
if isinstance(b, ToolUseBlock):
|
if isinstance(b, ToolUseBlock):
|
||||||
tool_blocks.append(b)
|
tool_blocks.append(b)
|
||||||
@@ -963,7 +961,7 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
content_blocks.append(b)
|
content_blocks.append(b)
|
||||||
return content_blocks, tool_blocks
|
return content_blocks, tool_blocks
|
||||||
|
|
||||||
def _internal_message_to_openai_messages(self, msg: InternalMessage) -> List[Dict[str, Any]]:
|
def _internal_message_to_openai_messages(self, msg: InternalMessage) -> list[dict[str, Any]]:
|
||||||
if msg.role == Role.USER:
|
if msg.role == Role.USER:
|
||||||
return self._user_message_to_openai(msg)
|
return self._user_message_to_openai(msg)
|
||||||
if msg.role == Role.ASSISTANT:
|
if msg.role == Role.ASSISTANT:
|
||||||
@@ -974,9 +972,9 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
return [{"role": "tool", "content": content_value or ""}]
|
return [{"role": "tool", "content": content_value or ""}]
|
||||||
return [{"role": "user", "content": ""}]
|
return [{"role": "user", "content": ""}]
|
||||||
|
|
||||||
def _user_message_to_openai(self, msg: InternalMessage) -> List[Dict[str, Any]]:
|
def _user_message_to_openai(self, msg: InternalMessage) -> list[dict[str, Any]]:
|
||||||
out: List[Dict[str, Any]] = []
|
out: list[dict[str, Any]] = []
|
||||||
pending: List[ContentBlock] = []
|
pending: list[ContentBlock] = []
|
||||||
|
|
||||||
def flush_user() -> None:
|
def flush_user() -> None:
|
||||||
nonlocal pending
|
nonlocal pending
|
||||||
@@ -1008,9 +1006,9 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
|
|
||||||
return out
|
return out
|
||||||
|
|
||||||
def _assistant_message_to_openai(self, msg: InternalMessage) -> Dict[str, Any]:
|
def _assistant_message_to_openai(self, msg: InternalMessage) -> dict[str, Any]:
|
||||||
content_blocks: List[ContentBlock] = []
|
content_blocks: list[ContentBlock] = []
|
||||||
tool_blocks: List[ToolUseBlock] = []
|
tool_blocks: list[ToolUseBlock] = []
|
||||||
|
|
||||||
for b in msg.content:
|
for b in msg.content:
|
||||||
if isinstance(b, ToolUseBlock):
|
if isinstance(b, ToolUseBlock):
|
||||||
@@ -1022,7 +1020,7 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
continue
|
continue
|
||||||
content_blocks.append(b)
|
content_blocks.append(b)
|
||||||
|
|
||||||
out: Dict[str, Any] = {"role": "assistant"}
|
out: dict[str, Any] = {"role": "assistant"}
|
||||||
content_value = self._blocks_to_openai_content(content_blocks)
|
content_value = self._blocks_to_openai_content(content_blocks)
|
||||||
out["content"] = content_value if content_value is not None else ""
|
out["content"] = content_value if content_value is not None else ""
|
||||||
|
|
||||||
@@ -1031,7 +1029,7 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
|
|
||||||
return out
|
return out
|
||||||
|
|
||||||
def _tool_result_block_to_openai_message(self, block: ToolResultBlock) -> Dict[str, Any]:
|
def _tool_result_block_to_openai_message(self, block: ToolResultBlock) -> dict[str, Any]:
|
||||||
content: str
|
content: str
|
||||||
if block.content_text is not None:
|
if block.content_text is not None:
|
||||||
content = block.content_text
|
content = block.content_text
|
||||||
@@ -1048,7 +1046,7 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
"content": content,
|
"content": content,
|
||||||
}
|
}
|
||||||
|
|
||||||
def _tool_use_block_to_openai_call(self, block: ToolUseBlock, index: int) -> Dict[str, Any]:
|
def _tool_use_block_to_openai_call(self, block: ToolUseBlock, index: int) -> dict[str, Any]:
|
||||||
return {
|
return {
|
||||||
"index": index,
|
"index": index,
|
||||||
"id": block.tool_id or f"call_{index}",
|
"id": block.tool_id or f"call_{index}",
|
||||||
@@ -1085,7 +1083,7 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
except ValueError:
|
except ValueError:
|
||||||
return ErrorType.UNKNOWN
|
return ErrorType.UNKNOWN
|
||||||
|
|
||||||
def _optional_int(self, value: Any) -> Optional[int]:
|
def _optional_int(self, value: Any) -> int | None:
|
||||||
if value is None:
|
if value is None:
|
||||||
return None
|
return None
|
||||||
try:
|
try:
|
||||||
@@ -1093,7 +1091,7 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
except (TypeError, ValueError):
|
except (TypeError, ValueError):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def _optional_float(self, value: Any) -> Optional[float]:
|
def _optional_float(self, value: Any) -> float | None:
|
||||||
if value is None:
|
if value is None:
|
||||||
return None
|
return None
|
||||||
try:
|
try:
|
||||||
@@ -1101,7 +1099,7 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
except (TypeError, ValueError):
|
except (TypeError, ValueError):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def _coerce_str_list(self, value: Any) -> Optional[List[str]]:
|
def _coerce_str_list(self, value: Any) -> list[str] | None:
|
||||||
if value is None:
|
if value is None:
|
||||||
return None
|
return None
|
||||||
if isinstance(value, str):
|
if isinstance(value, str):
|
||||||
@@ -1110,14 +1108,14 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
return [str(x) for x in value if x is not None]
|
return [str(x) for x in value if x is not None]
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def _extract_extra(self, payload: Dict[str, Any], known_keys: set[str]) -> Dict[str, Any]:
|
def _extract_extra(self, payload: dict[str, Any], known_keys: set[str]) -> dict[str, Any]:
|
||||||
return {k: v for k, v in payload.items() if k not in known_keys}
|
return {k: v for k, v in payload.items() if k not in known_keys}
|
||||||
|
|
||||||
def _merge_dropped(self, target: Dict[str, int], source: Dict[str, int]) -> None:
|
def _merge_dropped(self, target: dict[str, int], source: dict[str, int]) -> None:
|
||||||
for k, v in source.items():
|
for k, v in source.items():
|
||||||
target[k] = target.get(k, 0) + int(v)
|
target[k] = target.get(k, 0) + int(v)
|
||||||
|
|
||||||
def _ensure_tool_block_index(self, ss: Dict[str, Any], tool_key: str) -> int:
|
def _ensure_tool_block_index(self, ss: dict[str, Any], tool_key: str) -> int:
|
||||||
mapping = ss.get("tool_id_to_block_index")
|
mapping = ss.get("tool_id_to_block_index")
|
||||||
if not isinstance(mapping, dict):
|
if not isinstance(mapping, dict):
|
||||||
mapping = {}
|
mapping = {}
|
||||||
@@ -1131,7 +1129,7 @@ class OpenAINormalizer(FormatNormalizer):
|
|||||||
ss["next_block_index"] = next_idx + 1
|
ss["next_block_index"] = next_idx + 1
|
||||||
return next_idx
|
return next_idx
|
||||||
|
|
||||||
def _ensure_tool_call_index(self, ss: Dict[str, Any], tool_id: str) -> int:
|
def _ensure_tool_call_index(self, ss: dict[str, Any], tool_id: str) -> int:
|
||||||
mapping = ss.get("tool_id_to_index")
|
mapping = ss.get("tool_id_to_index")
|
||||||
if not isinstance(mapping, dict):
|
if not isinstance(mapping, dict):
|
||||||
mapping = {}
|
mapping = {}
|
||||||
|
|||||||
@@ -10,11 +10,10 @@ OpenAI CLI / Responses Normalizer (OPENAI_CLI)
|
|||||||
- 未识别的字段会进入 extra/raw,未知内容块保留在 internal,但默认输出阶段会丢弃。
|
- 未识别的字段会进入 extra/raw,未知内容块保留在 internal,但默认输出阶段会丢弃。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import json
|
import json
|
||||||
import time
|
import time
|
||||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
from typing import Any
|
||||||
|
|
||||||
from src.core.api_format.conversion.field_mappings import (
|
from src.core.api_format.conversion.field_mappings import (
|
||||||
ERROR_TYPE_MAPPINGS,
|
ERROR_TYPE_MAPPINGS,
|
||||||
@@ -65,7 +64,7 @@ class OpenAICliNormalizer(FormatNormalizer):
|
|||||||
supports_images=True,
|
supports_images=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
_ERROR_TYPE_TO_OPENAI: Dict[ErrorType, str] = {
|
_ERROR_TYPE_TO_OPENAI: dict[ErrorType, str] = {
|
||||||
ErrorType.INVALID_REQUEST: "invalid_request_error",
|
ErrorType.INVALID_REQUEST: "invalid_request_error",
|
||||||
ErrorType.AUTHENTICATION: "invalid_api_key",
|
ErrorType.AUTHENTICATION: "invalid_api_key",
|
||||||
ErrorType.PERMISSION_DENIED: "invalid_request_error",
|
ErrorType.PERMISSION_DENIED: "invalid_request_error",
|
||||||
@@ -82,12 +81,12 @@ class OpenAICliNormalizer(FormatNormalizer):
|
|||||||
# Requests
|
# Requests
|
||||||
# =========================
|
# =========================
|
||||||
|
|
||||||
def request_to_internal(self, request: Dict[str, Any]) -> InternalRequest:
|
def request_to_internal(self, request: dict[str, Any]) -> InternalRequest:
|
||||||
model = str(request.get("model") or "")
|
model = str(request.get("model") or "")
|
||||||
|
|
||||||
instructions_text = request.get("instructions")
|
instructions_text = request.get("instructions")
|
||||||
instructions: List[InstructionSegment] = []
|
instructions: list[InstructionSegment] = []
|
||||||
system_text: Optional[str] = None
|
system_text: str | None = None
|
||||||
if isinstance(instructions_text, str) and instructions_text.strip():
|
if isinstance(instructions_text, str) and instructions_text.strip():
|
||||||
system_text = instructions_text
|
system_text = instructions_text
|
||||||
instructions.append(InstructionSegment(role=Role.SYSTEM, text=instructions_text))
|
instructions.append(InstructionSegment(role=Role.SYSTEM, text=instructions_text))
|
||||||
@@ -118,8 +117,8 @@ class OpenAICliNormalizer(FormatNormalizer):
|
|||||||
|
|
||||||
return internal
|
return internal
|
||||||
|
|
||||||
def request_from_internal(self, internal: InternalRequest) -> Dict[str, Any]:
|
def request_from_internal(self, internal: InternalRequest) -> dict[str, Any]:
|
||||||
result: Dict[str, Any] = {
|
result: dict[str, Any] = {
|
||||||
"model": internal.model,
|
"model": internal.model,
|
||||||
"input": self._internal_messages_to_input(internal.messages),
|
"input": self._internal_messages_to_input(internal.messages),
|
||||||
}
|
}
|
||||||
@@ -164,7 +163,7 @@ class OpenAICliNormalizer(FormatNormalizer):
|
|||||||
# Responses
|
# Responses
|
||||||
# =========================
|
# =========================
|
||||||
|
|
||||||
def response_to_internal(self, response: Dict[str, Any]) -> InternalResponse:
|
def response_to_internal(self, response: dict[str, Any]) -> InternalResponse:
|
||||||
payload = self._unwrap_response_object(response)
|
payload = self._unwrap_response_object(response)
|
||||||
|
|
||||||
rid = str(payload.get("id") or "")
|
rid = str(payload.get("id") or "")
|
||||||
@@ -191,8 +190,8 @@ class OpenAICliNormalizer(FormatNormalizer):
|
|||||||
self,
|
self,
|
||||||
internal: InternalResponse,
|
internal: InternalResponse,
|
||||||
*,
|
*,
|
||||||
requested_model: Optional[str] = None,
|
requested_model: str | None = None,
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
text = self._collapse_internal_text(internal.content)
|
text = self._collapse_internal_text(internal.content)
|
||||||
|
|
||||||
output_message = {
|
output_message = {
|
||||||
@@ -203,7 +202,7 @@ class OpenAICliNormalizer(FormatNormalizer):
|
|||||||
}
|
}
|
||||||
|
|
||||||
usage = internal.usage or UsageInfo()
|
usage = internal.usage or UsageInfo()
|
||||||
usage_obj: Dict[str, Any] = {
|
usage_obj: dict[str, Any] = {
|
||||||
"input_tokens": usage.input_tokens,
|
"input_tokens": usage.input_tokens,
|
||||||
"output_tokens": usage.output_tokens,
|
"output_tokens": usage.output_tokens,
|
||||||
"total_tokens": usage.total_tokens or (usage.input_tokens + usage.output_tokens),
|
"total_tokens": usage.total_tokens or (usage.input_tokens + usage.output_tokens),
|
||||||
@@ -228,11 +227,11 @@ class OpenAICliNormalizer(FormatNormalizer):
|
|||||||
|
|
||||||
def stream_chunk_to_internal(
|
def stream_chunk_to_internal(
|
||||||
self,
|
self,
|
||||||
chunk: Dict[str, Any],
|
chunk: dict[str, Any],
|
||||||
state: StreamState,
|
state: StreamState,
|
||||||
) -> List[InternalStreamEvent]:
|
) -> list[InternalStreamEvent]:
|
||||||
ss = state.substate(self.FORMAT_ID)
|
ss = state.substate(self.FORMAT_ID)
|
||||||
events: List[InternalStreamEvent] = []
|
events: list[InternalStreamEvent] = []
|
||||||
|
|
||||||
# 统一错误结构(最佳努力)
|
# 统一错误结构(最佳努力)
|
||||||
if isinstance(chunk, dict) and "error" in chunk:
|
if isinstance(chunk, dict) and "error" in chunk:
|
||||||
@@ -392,11 +391,11 @@ class OpenAICliNormalizer(FormatNormalizer):
|
|||||||
self,
|
self,
|
||||||
event: InternalStreamEvent,
|
event: InternalStreamEvent,
|
||||||
state: StreamState,
|
state: StreamState,
|
||||||
) -> List[Dict[str, Any]]:
|
) -> list[dict[str, Any]]:
|
||||||
ss = state.substate(self.FORMAT_ID)
|
ss = state.substate(self.FORMAT_ID)
|
||||||
out: List[Dict[str, Any]] = []
|
out: list[dict[str, Any]] = []
|
||||||
|
|
||||||
def event_block(payload: Dict[str, Any]) -> Dict[str, Any]:
|
def event_block(payload: dict[str, Any]) -> dict[str, Any]:
|
||||||
# OpenAI Responses SSE 的 payload 通常自带 type 字段;这里强制保证
|
# OpenAI Responses SSE 的 payload 通常自带 type 字段;这里强制保证
|
||||||
return payload
|
return payload
|
||||||
|
|
||||||
@@ -463,10 +462,10 @@ class OpenAICliNormalizer(FormatNormalizer):
|
|||||||
# Error conversion
|
# Error conversion
|
||||||
# =========================
|
# =========================
|
||||||
|
|
||||||
def is_error_response(self, response: Dict[str, Any]) -> bool:
|
def is_error_response(self, response: dict[str, Any]) -> bool:
|
||||||
return isinstance(response, dict) and "error" in response
|
return isinstance(response, dict) and "error" in response
|
||||||
|
|
||||||
def error_to_internal(self, error_response: Dict[str, Any]) -> InternalError:
|
def error_to_internal(self, error_response: dict[str, Any]) -> InternalError:
|
||||||
err = error_response.get("error") if isinstance(error_response, dict) else None
|
err = error_response.get("error") if isinstance(error_response, dict) else None
|
||||||
err = err if isinstance(err, dict) else {}
|
err = err if isinstance(err, dict) else {}
|
||||||
|
|
||||||
@@ -484,9 +483,9 @@ class OpenAICliNormalizer(FormatNormalizer):
|
|||||||
extra={"openai_cli": {"error": err}, "raw": {"type": raw_type}},
|
extra={"openai_cli": {"error": err}, "raw": {"type": raw_type}},
|
||||||
)
|
)
|
||||||
|
|
||||||
def error_from_internal(self, internal: InternalError) -> Dict[str, Any]:
|
def error_from_internal(self, internal: InternalError) -> dict[str, Any]:
|
||||||
type_str = self._ERROR_TYPE_TO_OPENAI.get(internal.type, "server_error")
|
type_str = self._ERROR_TYPE_TO_OPENAI.get(internal.type, "server_error")
|
||||||
payload: Dict[str, Any] = {"type": type_str, "message": internal.message}
|
payload: dict[str, Any] = {"type": type_str, "message": internal.message}
|
||||||
if internal.code is not None:
|
if internal.code is not None:
|
||||||
payload["code"] = internal.code
|
payload["code"] = internal.code
|
||||||
if internal.param is not None:
|
if internal.param is not None:
|
||||||
@@ -497,7 +496,7 @@ class OpenAICliNormalizer(FormatNormalizer):
|
|||||||
# Helpers
|
# Helpers
|
||||||
# =========================
|
# =========================
|
||||||
|
|
||||||
def _unwrap_response_object(self, response: Dict[str, Any]) -> Dict[str, Any]:
|
def _unwrap_response_object(self, response: dict[str, Any]) -> dict[str, Any]:
|
||||||
if not isinstance(response, dict):
|
if not isinstance(response, dict):
|
||||||
return {}
|
return {}
|
||||||
resp_inner = response.get("response")
|
resp_inner = response.get("response")
|
||||||
@@ -506,8 +505,8 @@ class OpenAICliNormalizer(FormatNormalizer):
|
|||||||
return resp_inner
|
return resp_inner
|
||||||
return response
|
return response
|
||||||
|
|
||||||
def _extract_output_text_blocks(self, payload: Dict[str, Any]) -> Tuple[List[ContentBlock], Dict[str, Any]]:
|
def _extract_output_text_blocks(self, payload: dict[str, Any]) -> tuple[list[ContentBlock], dict[str, Any]]:
|
||||||
text_parts: List[str] = []
|
text_parts: list[str] = []
|
||||||
|
|
||||||
output = payload.get("output")
|
output = payload.get("output")
|
||||||
if isinstance(output, list):
|
if isinstance(output, list):
|
||||||
@@ -532,12 +531,12 @@ class OpenAICliNormalizer(FormatNormalizer):
|
|||||||
if not text_parts and isinstance(payload.get("output_text"), str):
|
if not text_parts and isinstance(payload.get("output_text"), str):
|
||||||
text_parts.append(payload.get("output_text") or "")
|
text_parts.append(payload.get("output_text") or "")
|
||||||
|
|
||||||
blocks: List[ContentBlock] = []
|
blocks: list[ContentBlock] = []
|
||||||
text = "".join(text_parts)
|
text = "".join(text_parts)
|
||||||
if text:
|
if text:
|
||||||
blocks.append(TextBlock(text=text))
|
blocks.append(TextBlock(text=text))
|
||||||
|
|
||||||
extra: Dict[str, Any] = {"raw": {"openai_cli_output": output}} if output is not None else {}
|
extra: dict[str, Any] = {"raw": {"openai_cli_output": output}} if output is not None else {}
|
||||||
return blocks, extra
|
return blocks, extra
|
||||||
|
|
||||||
def _usage_to_internal(self, usage: Any) -> UsageInfo:
|
def _usage_to_internal(self, usage: Any) -> UsageInfo:
|
||||||
@@ -553,14 +552,14 @@ class OpenAICliNormalizer(FormatNormalizer):
|
|||||||
extra={"openai_cli": {"usage": usage}},
|
extra={"openai_cli": {"usage": usage}},
|
||||||
)
|
)
|
||||||
|
|
||||||
def _collapse_internal_text(self, blocks: List[ContentBlock]) -> str:
|
def _collapse_internal_text(self, blocks: list[ContentBlock]) -> str:
|
||||||
parts: List[str] = []
|
parts: list[str] = []
|
||||||
for block in blocks:
|
for block in blocks:
|
||||||
if isinstance(block, TextBlock) and block.text:
|
if isinstance(block, TextBlock) and block.text:
|
||||||
parts.append(block.text)
|
parts.append(block.text)
|
||||||
return "".join(parts)
|
return "".join(parts)
|
||||||
|
|
||||||
def _input_to_internal_messages(self, input_data: Any) -> List[InternalMessage]:
|
def _input_to_internal_messages(self, input_data: Any) -> list[InternalMessage]:
|
||||||
if input_data is None:
|
if input_data is None:
|
||||||
return []
|
return []
|
||||||
|
|
||||||
@@ -575,7 +574,7 @@ class OpenAICliNormalizer(FormatNormalizer):
|
|||||||
if not isinstance(input_data, list):
|
if not isinstance(input_data, list):
|
||||||
return [InternalMessage(role=Role.USER, content=[UnknownBlock(raw_type="input", payload={"input": input_data})])]
|
return [InternalMessage(role=Role.USER, content=[UnknownBlock(raw_type="input", payload={"input": input_data})])]
|
||||||
|
|
||||||
messages: List[InternalMessage] = []
|
messages: list[InternalMessage] = []
|
||||||
for item in input_data:
|
for item in input_data:
|
||||||
if not isinstance(item, dict):
|
if not isinstance(item, dict):
|
||||||
continue
|
continue
|
||||||
@@ -624,7 +623,7 @@ class OpenAICliNormalizer(FormatNormalizer):
|
|||||||
|
|
||||||
# reasoning -> assistant 消息,提取 summary 作为文本
|
# reasoning -> assistant 消息,提取 summary 作为文本
|
||||||
if item_type == "reasoning":
|
if item_type == "reasoning":
|
||||||
summary_parts: List[str] = []
|
summary_parts: list[str] = []
|
||||||
summary = item.get("summary")
|
summary = item.get("summary")
|
||||||
if isinstance(summary, list):
|
if isinstance(summary, list):
|
||||||
for s in summary:
|
for s in summary:
|
||||||
@@ -638,7 +637,7 @@ class OpenAICliNormalizer(FormatNormalizer):
|
|||||||
summary_parts.append(summary)
|
summary_parts.append(summary)
|
||||||
|
|
||||||
# 如果有 summary 文本,创建一个 UnknownBlock 保留原始结构
|
# 如果有 summary 文本,创建一个 UnknownBlock 保留原始结构
|
||||||
reasoning_blocks: List[ContentBlock] = []
|
reasoning_blocks: list[ContentBlock] = []
|
||||||
if summary_parts:
|
if summary_parts:
|
||||||
# 保留 reasoning 的 summary 作为 UnknownBlock,便于输出时决策
|
# 保留 reasoning 的 summary 作为 UnknownBlock,便于输出时决策
|
||||||
reasoning_blocks.append(UnknownBlock(
|
reasoning_blocks.append(UnknownBlock(
|
||||||
@@ -662,7 +661,7 @@ class OpenAICliNormalizer(FormatNormalizer):
|
|||||||
|
|
||||||
return messages
|
return messages
|
||||||
|
|
||||||
def _responses_content_to_blocks(self, content: Any) -> List[ContentBlock]:
|
def _responses_content_to_blocks(self, content: Any) -> list[ContentBlock]:
|
||||||
if content is None:
|
if content is None:
|
||||||
return []
|
return []
|
||||||
if isinstance(content, str):
|
if isinstance(content, str):
|
||||||
@@ -673,7 +672,7 @@ class OpenAICliNormalizer(FormatNormalizer):
|
|||||||
if not isinstance(content, list):
|
if not isinstance(content, list):
|
||||||
return [UnknownBlock(raw_type="content", payload={"content": content})]
|
return [UnknownBlock(raw_type="content", payload={"content": content})]
|
||||||
|
|
||||||
blocks: List[ContentBlock] = []
|
blocks: list[ContentBlock] = []
|
||||||
for part in content:
|
for part in content:
|
||||||
if isinstance(part, str):
|
if isinstance(part, str):
|
||||||
if part:
|
if part:
|
||||||
@@ -690,8 +689,8 @@ class OpenAICliNormalizer(FormatNormalizer):
|
|||||||
blocks.append(UnknownBlock(raw_type=ptype or "unknown", payload=part))
|
blocks.append(UnknownBlock(raw_type=ptype or "unknown", payload=part))
|
||||||
return blocks
|
return blocks
|
||||||
|
|
||||||
def _internal_messages_to_input(self, messages: List[InternalMessage]) -> List[Dict[str, Any]]:
|
def _internal_messages_to_input(self, messages: list[InternalMessage]) -> list[dict[str, Any]]:
|
||||||
out: List[Dict[str, Any]] = []
|
out: list[dict[str, Any]] = []
|
||||||
for msg in messages:
|
for msg in messages:
|
||||||
# ToolUseBlock -> function_call
|
# ToolUseBlock -> function_call
|
||||||
for block in msg.content:
|
for block in msg.content:
|
||||||
@@ -729,7 +728,7 @@ class OpenAICliNormalizer(FormatNormalizer):
|
|||||||
|
|
||||||
# 普通 message(TextBlock)
|
# 普通 message(TextBlock)
|
||||||
role = self._role_to_openai(msg.role)
|
role = self._role_to_openai(msg.role)
|
||||||
content_items: List[Dict[str, Any]] = []
|
content_items: list[dict[str, Any]] = []
|
||||||
has_text = False
|
has_text = False
|
||||||
|
|
||||||
for block in msg.content:
|
for block in msg.content:
|
||||||
@@ -748,10 +747,10 @@ class OpenAICliNormalizer(FormatNormalizer):
|
|||||||
|
|
||||||
return out
|
return out
|
||||||
|
|
||||||
def _tools_to_internal(self, tools: Any) -> Optional[List[ToolDefinition]]:
|
def _tools_to_internal(self, tools: Any) -> list[ToolDefinition] | None:
|
||||||
if not isinstance(tools, list):
|
if not isinstance(tools, list):
|
||||||
return None
|
return None
|
||||||
out: List[ToolDefinition] = []
|
out: list[ToolDefinition] = []
|
||||||
for tool in tools:
|
for tool in tools:
|
||||||
if not isinstance(tool, dict):
|
if not isinstance(tool, dict):
|
||||||
continue
|
continue
|
||||||
@@ -783,7 +782,7 @@ class OpenAICliNormalizer(FormatNormalizer):
|
|||||||
)
|
)
|
||||||
return out or None
|
return out or None
|
||||||
|
|
||||||
def _tool_choice_to_internal(self, tool_choice: Any) -> Optional[ToolChoice]:
|
def _tool_choice_to_internal(self, tool_choice: Any) -> ToolChoice | None:
|
||||||
if tool_choice is None:
|
if tool_choice is None:
|
||||||
return None
|
return None
|
||||||
if isinstance(tool_choice, str):
|
if isinstance(tool_choice, str):
|
||||||
@@ -802,7 +801,7 @@ class OpenAICliNormalizer(FormatNormalizer):
|
|||||||
|
|
||||||
return ToolChoice(type=ToolChoiceType.AUTO, extra={"raw": tool_choice})
|
return ToolChoice(type=ToolChoiceType.AUTO, extra={"raw": tool_choice})
|
||||||
|
|
||||||
def _tool_choice_to_openai(self, tool_choice: ToolChoice) -> Union[str, Dict[str, Any]]:
|
def _tool_choice_to_openai(self, tool_choice: ToolChoice) -> str | dict[str, Any]:
|
||||||
if tool_choice.type == ToolChoiceType.NONE:
|
if tool_choice.type == ToolChoiceType.NONE:
|
||||||
return "none"
|
return "none"
|
||||||
if tool_choice.type == ToolChoiceType.AUTO:
|
if tool_choice.type == ToolChoiceType.AUTO:
|
||||||
@@ -840,7 +839,7 @@ class OpenAICliNormalizer(FormatNormalizer):
|
|||||||
return "tool"
|
return "tool"
|
||||||
return "user"
|
return "user"
|
||||||
|
|
||||||
def _optional_int(self, value: Any) -> Optional[int]:
|
def _optional_int(self, value: Any) -> int | None:
|
||||||
if value is None:
|
if value is None:
|
||||||
return None
|
return None
|
||||||
try:
|
try:
|
||||||
@@ -848,7 +847,7 @@ class OpenAICliNormalizer(FormatNormalizer):
|
|||||||
except (TypeError, ValueError):
|
except (TypeError, ValueError):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def _optional_float(self, value: Any) -> Optional[float]:
|
def _optional_float(self, value: Any) -> float | None:
|
||||||
if value is None:
|
if value is None:
|
||||||
return None
|
return None
|
||||||
try:
|
try:
|
||||||
@@ -856,13 +855,13 @@ class OpenAICliNormalizer(FormatNormalizer):
|
|||||||
except (TypeError, ValueError):
|
except (TypeError, ValueError):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def _coerce_str_list(self, value: Any) -> Optional[List[str]]:
|
def _coerce_str_list(self, value: Any) -> list[str] | None:
|
||||||
if value is None:
|
if value is None:
|
||||||
return None
|
return None
|
||||||
if isinstance(value, str):
|
if isinstance(value, str):
|
||||||
return [value]
|
return [value]
|
||||||
if isinstance(value, list):
|
if isinstance(value, list):
|
||||||
out: List[str] = []
|
out: list[str] = []
|
||||||
for item in value:
|
for item in value:
|
||||||
if item is None:
|
if item is None:
|
||||||
continue
|
continue
|
||||||
@@ -870,14 +869,14 @@ class OpenAICliNormalizer(FormatNormalizer):
|
|||||||
return out
|
return out
|
||||||
return [str(value)]
|
return [str(value)]
|
||||||
|
|
||||||
def _extract_extra(self, payload: Dict[str, Any], keep_keys: set[str]) -> Dict[str, Any]:
|
def _extract_extra(self, payload: dict[str, Any], keep_keys: set[str]) -> dict[str, Any]:
|
||||||
if not isinstance(payload, dict):
|
if not isinstance(payload, dict):
|
||||||
return {}
|
return {}
|
||||||
return {k: v for k, v in payload.items() if k not in keep_keys}
|
return {k: v for k, v in payload.items() if k not in keep_keys}
|
||||||
|
|
||||||
def _join_instructions(self, internal: InternalRequest) -> str:
|
def _join_instructions(self, internal: InternalRequest) -> str:
|
||||||
if internal.instructions:
|
if internal.instructions:
|
||||||
parts: List[str] = []
|
parts: list[str] = []
|
||||||
for seg in internal.instructions:
|
for seg in internal.instructions:
|
||||||
if seg.text:
|
if seg.text:
|
||||||
parts.append(seg.text)
|
parts.append(seg.text)
|
||||||
|
|||||||
@@ -9,12 +9,12 @@ source -> internal -> target
|
|||||||
- 转换失败将抛出 `FormatConversionError`(不再静默回退)。
|
- 转换失败将抛出 `FormatConversionError`(不再静默回退)。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import threading
|
import threading
|
||||||
import time
|
import time
|
||||||
from contextlib import contextmanager
|
from contextlib import contextmanager
|
||||||
from typing import Any, Dict, Generator, List, Optional
|
from typing import Any
|
||||||
|
from collections.abc import Generator
|
||||||
|
|
||||||
from src.core.logger import logger
|
from src.core.logger import logger
|
||||||
from src.core.metrics import format_conversion_duration_seconds, format_conversion_total
|
from src.core.metrics import format_conversion_duration_seconds, format_conversion_total
|
||||||
@@ -29,7 +29,7 @@ def _track_conversion_metrics(
|
|||||||
direction: str,
|
direction: str,
|
||||||
source: str,
|
source: str,
|
||||||
target: str,
|
target: str,
|
||||||
) -> Generator[None, None, None]:
|
) -> Generator[None]:
|
||||||
start = time.perf_counter()
|
start = time.perf_counter()
|
||||||
try:
|
try:
|
||||||
yield
|
yield
|
||||||
@@ -47,13 +47,13 @@ class FormatConversionRegistry:
|
|||||||
"""基于 Normalizer 的格式转换注册表"""
|
"""基于 Normalizer 的格式转换注册表"""
|
||||||
|
|
||||||
def __init__(self) -> None:
|
def __init__(self) -> None:
|
||||||
self._normalizers: Dict[str, FormatNormalizer] = {}
|
self._normalizers: dict[str, FormatNormalizer] = {}
|
||||||
|
|
||||||
def register(self, normalizer: FormatNormalizer) -> None:
|
def register(self, normalizer: FormatNormalizer) -> None:
|
||||||
self._normalizers[str(normalizer.FORMAT_ID).upper()] = normalizer
|
self._normalizers[str(normalizer.FORMAT_ID).upper()] = normalizer
|
||||||
logger.info(f"[FormatConversionRegistry] 注册 normalizer: {normalizer.FORMAT_ID}")
|
logger.info(f"[FormatConversionRegistry] 注册 normalizer: {normalizer.FORMAT_ID}")
|
||||||
|
|
||||||
def get_normalizer(self, format_id: str) -> Optional[FormatNormalizer]:
|
def get_normalizer(self, format_id: str) -> FormatNormalizer | None:
|
||||||
return self._normalizers.get(str(format_id).upper())
|
return self._normalizers.get(str(format_id).upper())
|
||||||
|
|
||||||
def _require_normalizer(self, format_id: str) -> FormatNormalizer:
|
def _require_normalizer(self, format_id: str) -> FormatNormalizer:
|
||||||
@@ -66,10 +66,10 @@ class FormatConversionRegistry:
|
|||||||
|
|
||||||
def convert_request(
|
def convert_request(
|
||||||
self,
|
self,
|
||||||
request: Dict[str, Any],
|
request: dict[str, Any],
|
||||||
source_format: str,
|
source_format: str,
|
||||||
target_format: str,
|
target_format: str,
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
if str(source_format).upper() == str(target_format).upper():
|
if str(source_format).upper() == str(target_format).upper():
|
||||||
return request
|
return request
|
||||||
|
|
||||||
@@ -85,12 +85,12 @@ class FormatConversionRegistry:
|
|||||||
|
|
||||||
def convert_response(
|
def convert_response(
|
||||||
self,
|
self,
|
||||||
response: Dict[str, Any],
|
response: dict[str, Any],
|
||||||
source_format: str,
|
source_format: str,
|
||||||
target_format: str,
|
target_format: str,
|
||||||
*,
|
*,
|
||||||
requested_model: Optional[str] = None,
|
requested_model: str | None = None,
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""转换响应格式
|
"""转换响应格式
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
@@ -124,10 +124,10 @@ class FormatConversionRegistry:
|
|||||||
|
|
||||||
def convert_error_response(
|
def convert_error_response(
|
||||||
self,
|
self,
|
||||||
error_response: Dict[str, Any],
|
error_response: dict[str, Any],
|
||||||
source_format: str,
|
source_format: str,
|
||||||
target_format: str,
|
target_format: str,
|
||||||
) -> Dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
if str(source_format).upper() == str(target_format).upper():
|
if str(source_format).upper() == str(target_format).upper():
|
||||||
return error_response
|
return error_response
|
||||||
|
|
||||||
@@ -152,11 +152,11 @@ class FormatConversionRegistry:
|
|||||||
|
|
||||||
def convert_stream_chunk(
|
def convert_stream_chunk(
|
||||||
self,
|
self,
|
||||||
chunk: Dict[str, Any],
|
chunk: dict[str, Any],
|
||||||
source_format: str,
|
source_format: str,
|
||||||
target_format: str,
|
target_format: str,
|
||||||
state: Optional[StreamState] = None,
|
state: StreamState | None = None,
|
||||||
) -> List[Dict[str, Any]]:
|
) -> list[dict[str, Any]]:
|
||||||
if str(source_format).upper() == str(target_format).upper():
|
if str(source_format).upper() == str(target_format).upper():
|
||||||
return [chunk]
|
return [chunk]
|
||||||
|
|
||||||
@@ -182,7 +182,7 @@ class FormatConversionRegistry:
|
|||||||
with _track_conversion_metrics("stream", str(source_format).upper(), str(target_format).upper()):
|
with _track_conversion_metrics("stream", str(source_format).upper(), str(target_format).upper()):
|
||||||
try:
|
try:
|
||||||
events = src.stream_chunk_to_internal(chunk, state)
|
events = src.stream_chunk_to_internal(chunk, state)
|
||||||
out: List[Dict[str, Any]] = []
|
out: list[dict[str, Any]] = []
|
||||||
for event in events:
|
for event in events:
|
||||||
out.extend(tgt.stream_event_from_internal(event, state))
|
out.extend(tgt.stream_event_from_internal(event, state))
|
||||||
return out
|
return out
|
||||||
@@ -226,10 +226,10 @@ class FormatConversionRegistry:
|
|||||||
return self.can_convert_stream(format_a, format_b) and self.can_convert_stream(format_b, format_a)
|
return self.can_convert_stream(format_a, format_b) and self.can_convert_stream(format_b, format_a)
|
||||||
return True
|
return True
|
||||||
|
|
||||||
def list_normalizers(self) -> List[str]:
|
def list_normalizers(self) -> list[str]:
|
||||||
return sorted(self._normalizers.keys())
|
return sorted(self._normalizers.keys())
|
||||||
|
|
||||||
def get_supported_targets(self, source_format: str) -> List[str]:
|
def get_supported_targets(self, source_format: str) -> list[str]:
|
||||||
src = str(source_format).upper()
|
src = str(source_format).upper()
|
||||||
if src not in self._normalizers:
|
if src not in self._normalizers:
|
||||||
return []
|
return []
|
||||||
|
|||||||
@@ -4,11 +4,10 @@
|
|||||||
用于把 OpenAI/Claude/Gemini 的流式协议映射为统一事件序列,再由目标格式 Normalizer 输出。
|
用于把 OpenAI/Claude/Gemini 的流式协议映射为统一事件序列,再由目标格式 Normalizer 输出。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
from typing import Any, Dict, Optional, Union
|
from typing import Any
|
||||||
|
|
||||||
from .internal import ContentType, InternalError, StopReason, UsageInfo
|
from .internal import ContentType, InternalError, StopReason, UsageInfo
|
||||||
|
|
||||||
@@ -34,8 +33,8 @@ class MessageStartEvent:
|
|||||||
type: StreamEventType = field(default=StreamEventType.MESSAGE_START, init=False)
|
type: StreamEventType = field(default=StreamEventType.MESSAGE_START, init=False)
|
||||||
message_id: str = ""
|
message_id: str = ""
|
||||||
model: str = ""
|
model: str = ""
|
||||||
usage: Optional[UsageInfo] = None # Claude 流式响应的 message_start 可能包含 usage
|
usage: UsageInfo | None = None # Claude 流式响应的 message_start 可能包含 usage
|
||||||
extra: Dict[str, Any] = field(default_factory=dict)
|
extra: dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -46,9 +45,9 @@ class ContentBlockStartEvent:
|
|||||||
block_index: int = 0
|
block_index: int = 0
|
||||||
block_type: ContentType = ContentType.TEXT
|
block_type: ContentType = ContentType.TEXT
|
||||||
# 工具调用时使用(TOOL_USE block)
|
# 工具调用时使用(TOOL_USE block)
|
||||||
tool_id: Optional[str] = None
|
tool_id: str | None = None
|
||||||
tool_name: Optional[str] = None
|
tool_name: str | None = None
|
||||||
extra: Dict[str, Any] = field(default_factory=dict)
|
extra: dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -58,7 +57,7 @@ class ContentDeltaEvent:
|
|||||||
type: StreamEventType = field(default=StreamEventType.CONTENT_DELTA, init=False)
|
type: StreamEventType = field(default=StreamEventType.CONTENT_DELTA, init=False)
|
||||||
block_index: int = 0
|
block_index: int = 0
|
||||||
text_delta: str = ""
|
text_delta: str = ""
|
||||||
extra: Dict[str, Any] = field(default_factory=dict)
|
extra: dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -69,7 +68,7 @@ class ToolCallDeltaEvent:
|
|||||||
block_index: int = 0
|
block_index: int = 0
|
||||||
tool_id: str = ""
|
tool_id: str = ""
|
||||||
input_delta: str = "" # JSON 字符串片段
|
input_delta: str = "" # JSON 字符串片段
|
||||||
extra: Dict[str, Any] = field(default_factory=dict)
|
extra: dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -78,7 +77,7 @@ class ContentBlockStopEvent:
|
|||||||
|
|
||||||
type: StreamEventType = field(default=StreamEventType.CONTENT_BLOCK_STOP, init=False)
|
type: StreamEventType = field(default=StreamEventType.CONTENT_BLOCK_STOP, init=False)
|
||||||
block_index: int = 0
|
block_index: int = 0
|
||||||
extra: Dict[str, Any] = field(default_factory=dict)
|
extra: dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -86,9 +85,9 @@ class MessageStopEvent:
|
|||||||
"""消息结束事件"""
|
"""消息结束事件"""
|
||||||
|
|
||||||
type: StreamEventType = field(default=StreamEventType.MESSAGE_STOP, init=False)
|
type: StreamEventType = field(default=StreamEventType.MESSAGE_STOP, init=False)
|
||||||
stop_reason: Optional[StopReason] = None
|
stop_reason: StopReason | None = None
|
||||||
usage: Optional[UsageInfo] = None
|
usage: UsageInfo | None = None
|
||||||
extra: Dict[str, Any] = field(default_factory=dict)
|
extra: dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -97,7 +96,7 @@ class UsageEvent:
|
|||||||
|
|
||||||
type: StreamEventType = field(default=StreamEventType.USAGE, init=False)
|
type: StreamEventType = field(default=StreamEventType.USAGE, init=False)
|
||||||
usage: UsageInfo = field(default_factory=UsageInfo)
|
usage: UsageInfo = field(default_factory=UsageInfo)
|
||||||
extra: Dict[str, Any] = field(default_factory=dict)
|
extra: dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -106,7 +105,7 @@ class ErrorEvent:
|
|||||||
|
|
||||||
type: StreamEventType = field(default=StreamEventType.ERROR, init=False)
|
type: StreamEventType = field(default=StreamEventType.ERROR, init=False)
|
||||||
error: InternalError
|
error: InternalError
|
||||||
extra: Dict[str, Any] = field(default_factory=dict)
|
extra: dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -115,21 +114,21 @@ class UnknownStreamEvent:
|
|||||||
|
|
||||||
type: StreamEventType = field(default=StreamEventType.UNKNOWN, init=False)
|
type: StreamEventType = field(default=StreamEventType.UNKNOWN, init=False)
|
||||||
raw_type: str = ""
|
raw_type: str = ""
|
||||||
payload: Dict[str, Any] = field(default_factory=dict)
|
payload: dict[str, Any] = field(default_factory=dict)
|
||||||
extra: Dict[str, Any] = field(default_factory=dict)
|
extra: dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
|
|
||||||
InternalStreamEvent = Union[
|
InternalStreamEvent = (
|
||||||
MessageStartEvent,
|
MessageStartEvent
|
||||||
ContentBlockStartEvent,
|
| ContentBlockStartEvent
|
||||||
ContentDeltaEvent,
|
| ContentDeltaEvent
|
||||||
ToolCallDeltaEvent,
|
| ToolCallDeltaEvent
|
||||||
ContentBlockStopEvent,
|
| ContentBlockStopEvent
|
||||||
MessageStopEvent,
|
| MessageStopEvent
|
||||||
UsageEvent,
|
| UsageEvent
|
||||||
ErrorEvent,
|
| ErrorEvent
|
||||||
UnknownStreamEvent,
|
| UnknownStreamEvent
|
||||||
]
|
)
|
||||||
|
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
|
|||||||
@@ -5,10 +5,9 @@
|
|||||||
每个 Normalizer 通过 `substate(format_id)` 获取自己的隔离状态字典。
|
每个 Normalizer 通过 `substate(format_id)` 获取自己的隔离状态字典。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from typing import Any, Dict
|
from typing import Any
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -26,12 +25,12 @@ class StreamState:
|
|||||||
message_id: str = ""
|
message_id: str = ""
|
||||||
|
|
||||||
# Registry/调用层的通用扩展信息(与具体格式无关)
|
# Registry/调用层的通用扩展信息(与具体格式无关)
|
||||||
extra: Dict[str, Any] = field(default_factory=dict)
|
extra: dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
# 各 Normalizer 的隔离状态(key: FORMAT_ID)
|
# 各 Normalizer 的隔离状态(key: FORMAT_ID)
|
||||||
by_format: Dict[str, Dict[str, Any]] = field(default_factory=dict)
|
by_format: dict[str, dict[str, Any]] = field(default_factory=dict)
|
||||||
|
|
||||||
def substate(self, format_id: str) -> Dict[str, Any]:
|
def substate(self, format_id: str) -> dict[str, Any]:
|
||||||
"""获取指定格式的隔离子状态"""
|
"""获取指定格式的隔离子状态"""
|
||||||
key = str(format_id).upper()
|
key = str(format_id).upper()
|
||||||
return self.by_format.setdefault(key, {})
|
return self.by_format.setdefault(key, {})
|
||||||
|
|||||||
@@ -6,7 +6,8 @@ API 格式检测
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from typing import TYPE_CHECKING, Dict, Optional, Tuple
|
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from starlette.requests import Request
|
from starlette.requests import Request
|
||||||
@@ -16,10 +17,10 @@ from src.core.api_format.metadata import API_FORMAT_DEFINITIONS, ApiFormatDefini
|
|||||||
|
|
||||||
|
|
||||||
def _extract_api_key_by_definition(
|
def _extract_api_key_by_definition(
|
||||||
headers: Dict[str, str],
|
headers: dict[str, str],
|
||||||
query_params: Optional[Dict[str, str]],
|
query_params: dict[str, str] | None,
|
||||||
definition: ApiFormatDefinition,
|
definition: ApiFormatDefinition,
|
||||||
) -> Tuple[Optional[str], str]:
|
) -> tuple[str | None, str]:
|
||||||
"""
|
"""
|
||||||
根据格式定义从请求中提取 API Key
|
根据格式定义从请求中提取 API Key
|
||||||
|
|
||||||
@@ -64,9 +65,9 @@ def _extract_api_key_by_definition(
|
|||||||
|
|
||||||
|
|
||||||
def detect_format_from_request(
|
def detect_format_from_request(
|
||||||
headers: Dict[str, str],
|
headers: dict[str, str],
|
||||||
query_params: Optional[Dict[str, str]] = None,
|
query_params: dict[str, str] | None = None,
|
||||||
) -> Tuple[APIFormat, Optional[str], str]:
|
) -> tuple[APIFormat, str | None, str]:
|
||||||
"""
|
"""
|
||||||
从请求头检测 API 格式和 API Key
|
从请求头检测 API 格式和 API Key
|
||||||
|
|
||||||
@@ -107,8 +108,8 @@ def detect_format_from_request(
|
|||||||
|
|
||||||
|
|
||||||
def detect_format_and_key_from_starlette(
|
def detect_format_and_key_from_starlette(
|
||||||
request: "Request",
|
request: Request,
|
||||||
) -> Tuple[str, Optional[str], str]:
|
) -> tuple[str, str | None, str]:
|
||||||
"""
|
"""
|
||||||
从 Starlette Request 对象检测 API 格式和 API Key
|
从 Starlette Request 对象检测 API 格式和 API Key
|
||||||
|
|
||||||
@@ -135,7 +136,7 @@ def detect_format_and_key_from_starlette(
|
|||||||
|
|
||||||
def detect_format_from_response(
|
def detect_format_from_response(
|
||||||
response_data: dict,
|
response_data: dict,
|
||||||
) -> Optional[APIFormat]:
|
) -> APIFormat | None:
|
||||||
"""
|
"""
|
||||||
从响应内容检测 API 格式
|
从响应内容检测 API 格式
|
||||||
|
|
||||||
|
|||||||
@@ -12,7 +12,8 @@
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from typing import AbstractSet, Any, Dict, FrozenSet, Optional, Set
|
from collections.abc import Set as AbstractSet
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
from src.core.api_format.enums import APIFormat
|
from src.core.api_format.enums import APIFormat
|
||||||
from src.core.api_format.metadata import get_auth_config, get_extra_headers, get_protected_keys
|
from src.core.api_format.metadata import get_auth_config, get_extra_headers, get_protected_keys
|
||||||
@@ -23,7 +24,7 @@ from src.core.api_format.metadata import get_auth_config, get_extra_headers, get
|
|||||||
# =============================================================================
|
# =============================================================================
|
||||||
|
|
||||||
# 转发给上游时需要剔除的头部(系统管理 + 认证替换)
|
# 转发给上游时需要剔除的头部(系统管理 + 认证替换)
|
||||||
UPSTREAM_DROP_HEADERS: FrozenSet[str] = frozenset(
|
UPSTREAM_DROP_HEADERS: frozenset[str] = frozenset(
|
||||||
{
|
{
|
||||||
# 认证头 - 会被替换为 Provider 的认证
|
# 认证头 - 会被替换为 Provider 的认证
|
||||||
"authorization",
|
"authorization",
|
||||||
@@ -41,7 +42,7 @@ UPSTREAM_DROP_HEADERS: FrozenSet[str] = frozenset(
|
|||||||
|
|
||||||
# 最小必脱敏集合(编译时常量,用于快速路径)
|
# 最小必脱敏集合(编译时常量,用于快速路径)
|
||||||
# 完整脱敏应使用 SystemConfigService.get_sensitive_headers()
|
# 完整脱敏应使用 SystemConfigService.get_sensitive_headers()
|
||||||
CORE_REDACT_HEADERS: FrozenSet[str] = frozenset(
|
CORE_REDACT_HEADERS: frozenset[str] = frozenset(
|
||||||
{
|
{
|
||||||
"authorization",
|
"authorization",
|
||||||
"x-api-key",
|
"x-api-key",
|
||||||
@@ -50,7 +51,7 @@ CORE_REDACT_HEADERS: FrozenSet[str] = frozenset(
|
|||||||
)
|
)
|
||||||
|
|
||||||
# Hop-by-hop 头部 (RFC 7230)
|
# Hop-by-hop 头部 (RFC 7230)
|
||||||
HOP_BY_HOP_HEADERS: FrozenSet[str] = frozenset(
|
HOP_BY_HOP_HEADERS: frozenset[str] = frozenset(
|
||||||
{
|
{
|
||||||
"connection",
|
"connection",
|
||||||
"keep-alive",
|
"keep-alive",
|
||||||
@@ -64,7 +65,7 @@ HOP_BY_HOP_HEADERS: FrozenSet[str] = frozenset(
|
|||||||
)
|
)
|
||||||
|
|
||||||
# 响应时需要过滤的头部(body-dependent + hop-by-hop)
|
# 响应时需要过滤的头部(body-dependent + hop-by-hop)
|
||||||
RESPONSE_DROP_HEADERS: FrozenSet[str] = (
|
RESPONSE_DROP_HEADERS: frozenset[str] = (
|
||||||
frozenset(
|
frozenset(
|
||||||
{
|
{
|
||||||
"content-length",
|
"content-length",
|
||||||
@@ -82,7 +83,7 @@ RESPONSE_DROP_HEADERS: FrozenSet[str] = (
|
|||||||
# =============================================================================
|
# =============================================================================
|
||||||
|
|
||||||
|
|
||||||
def normalize_headers(headers: Dict[str, str]) -> Dict[str, str]:
|
def normalize_headers(headers: dict[str, str]) -> dict[str, str]:
|
||||||
"""
|
"""
|
||||||
将请求头 key 统一为小写
|
将请求头 key 统一为小写
|
||||||
|
|
||||||
@@ -92,7 +93,7 @@ def normalize_headers(headers: Dict[str, str]) -> Dict[str, str]:
|
|||||||
return {k.lower(): v for k, v in headers.items()}
|
return {k.lower(): v for k, v in headers.items()}
|
||||||
|
|
||||||
|
|
||||||
def get_header_value(headers: Dict[str, str], key: str, default: str = "") -> str:
|
def get_header_value(headers: dict[str, str], key: str, default: str = "") -> str:
|
||||||
"""
|
"""
|
||||||
大小写不敏感地获取请求头值
|
大小写不敏感地获取请求头值
|
||||||
|
|
||||||
@@ -117,7 +118,7 @@ def get_header_value(headers: Dict[str, str], key: str, default: str = "") -> st
|
|||||||
# =============================================================================
|
# =============================================================================
|
||||||
|
|
||||||
|
|
||||||
def extract_client_api_key(headers: Dict[str, str], api_format: APIFormat) -> Optional[str]:
|
def extract_client_api_key(headers: dict[str, str], api_format: APIFormat) -> str | None:
|
||||||
"""
|
"""
|
||||||
从客户端请求头提取 API Key
|
从客户端请求头提取 API Key
|
||||||
|
|
||||||
@@ -147,10 +148,10 @@ def extract_client_api_key(headers: Dict[str, str], api_format: APIFormat) -> Op
|
|||||||
|
|
||||||
|
|
||||||
def extract_client_api_key_with_query(
|
def extract_client_api_key_with_query(
|
||||||
headers: Dict[str, str],
|
headers: dict[str, str],
|
||||||
query_params: Optional[Dict[str, str]],
|
query_params: dict[str, str] | None,
|
||||||
api_format: APIFormat,
|
api_format: APIFormat,
|
||||||
) -> Optional[str]:
|
) -> str | None:
|
||||||
"""
|
"""
|
||||||
从客户端请求头或 URL 参数提取 API Key
|
从客户端请求头或 URL 参数提取 API Key
|
||||||
|
|
||||||
@@ -184,10 +185,10 @@ def extract_client_api_key_with_query(
|
|||||||
|
|
||||||
|
|
||||||
def detect_capabilities(
|
def detect_capabilities(
|
||||||
headers: Dict[str, str],
|
headers: dict[str, str],
|
||||||
api_format: APIFormat,
|
api_format: APIFormat,
|
||||||
request_body: Optional[Dict[str, Any]] = None, # noqa: ARG001 - 预留给部分格式使用
|
request_body: dict[str, Any] | None = None, # noqa: ARG001 - 预留给部分格式使用
|
||||||
) -> Dict[str, bool]:
|
) -> dict[str, bool]:
|
||||||
"""
|
"""
|
||||||
从请求头检测能力需求
|
从请求头检测能力需求
|
||||||
|
|
||||||
@@ -203,7 +204,7 @@ def detect_capabilities(
|
|||||||
能力需求字典,如 {"context_1m": True}
|
能力需求字典,如 {"context_1m": True}
|
||||||
"""
|
"""
|
||||||
|
|
||||||
requirements: Dict[str, bool] = {}
|
requirements: dict[str, bool] = {}
|
||||||
|
|
||||||
if api_format in (APIFormat.CLAUDE, APIFormat.CLAUDE_CLI):
|
if api_format in (APIFormat.CLAUDE, APIFormat.CLAUDE_CLI):
|
||||||
beta_header = get_header_value(headers, "anthropic-beta")
|
beta_header = get_header_value(headers, "anthropic-beta")
|
||||||
@@ -228,20 +229,20 @@ class HeaderBuilder:
|
|||||||
|
|
||||||
def __init__(self) -> None:
|
def __init__(self) -> None:
|
||||||
# key: (original_case_key, value)
|
# key: (original_case_key, value)
|
||||||
self._headers: Dict[str, tuple[str, str]] = {}
|
self._headers: dict[str, tuple[str, str]] = {}
|
||||||
|
|
||||||
def add(self, key: str, value: str) -> "HeaderBuilder":
|
def add(self, key: str, value: str) -> HeaderBuilder:
|
||||||
"""添加单个头部(会覆盖同名头部)"""
|
"""添加单个头部(会覆盖同名头部)"""
|
||||||
self._headers[key.lower()] = (key, value)
|
self._headers[key.lower()] = (key, value)
|
||||||
return self
|
return self
|
||||||
|
|
||||||
def add_many(self, headers: Dict[str, str]) -> "HeaderBuilder":
|
def add_many(self, headers: dict[str, str]) -> HeaderBuilder:
|
||||||
"""批量添加头部"""
|
"""批量添加头部"""
|
||||||
for k, v in headers.items():
|
for k, v in headers.items():
|
||||||
self.add(k, v)
|
self.add(k, v)
|
||||||
return self
|
return self
|
||||||
|
|
||||||
def add_protected(self, headers: Dict[str, str], protected_keys: AbstractSet[str]) -> "HeaderBuilder":
|
def add_protected(self, headers: dict[str, str], protected_keys: AbstractSet[str]) -> HeaderBuilder:
|
||||||
"""
|
"""
|
||||||
添加头部但保护指定的 key 不被覆盖
|
添加头部但保护指定的 key 不被覆盖
|
||||||
|
|
||||||
@@ -253,13 +254,13 @@ class HeaderBuilder:
|
|||||||
self.add(k, v)
|
self.add(k, v)
|
||||||
return self
|
return self
|
||||||
|
|
||||||
def remove(self, keys: FrozenSet[str]) -> "HeaderBuilder":
|
def remove(self, keys: frozenset[str]) -> HeaderBuilder:
|
||||||
"""移除指定的头部"""
|
"""移除指定的头部"""
|
||||||
for k in keys:
|
for k in keys:
|
||||||
self._headers.pop(k.lower(), None)
|
self._headers.pop(k.lower(), None)
|
||||||
return self
|
return self
|
||||||
|
|
||||||
def rename(self, from_key: str, to_key: str) -> "HeaderBuilder":
|
def rename(self, from_key: str, to_key: str) -> HeaderBuilder:
|
||||||
"""
|
"""
|
||||||
重命名头部(保留原值)
|
重命名头部(保留原值)
|
||||||
|
|
||||||
@@ -273,9 +274,9 @@ class HeaderBuilder:
|
|||||||
|
|
||||||
def apply_rules(
|
def apply_rules(
|
||||||
self,
|
self,
|
||||||
rules: list[Dict[str, Any]],
|
rules: list[dict[str, Any]],
|
||||||
protected_keys: Optional[AbstractSet[str]] = None,
|
protected_keys: AbstractSet[str] | None = None,
|
||||||
) -> "HeaderBuilder":
|
) -> HeaderBuilder:
|
||||||
"""
|
"""
|
||||||
应用请求头规则
|
应用请求头规则
|
||||||
|
|
||||||
@@ -314,20 +315,20 @@ class HeaderBuilder:
|
|||||||
|
|
||||||
return self
|
return self
|
||||||
|
|
||||||
def build(self) -> Dict[str, str]:
|
def build(self) -> dict[str, str]:
|
||||||
"""构建最终的头部字典"""
|
"""构建最终的头部字典"""
|
||||||
return {original_key: value for original_key, value in self._headers.values()}
|
return {original_key: value for original_key, value in self._headers.values()}
|
||||||
|
|
||||||
|
|
||||||
def build_upstream_headers(
|
def build_upstream_headers(
|
||||||
original_headers: Dict[str, str],
|
original_headers: dict[str, str],
|
||||||
api_format: APIFormat,
|
api_format: APIFormat,
|
||||||
provider_api_key: str,
|
provider_api_key: str,
|
||||||
*,
|
*,
|
||||||
endpoint_headers: Optional[Dict[str, str]] = None,
|
endpoint_headers: dict[str, str] | None = None,
|
||||||
extra_headers: Optional[Dict[str, str]] = None,
|
extra_headers: dict[str, str] | None = None,
|
||||||
drop_headers: Optional[FrozenSet[str]] = None,
|
drop_headers: frozenset[str] | None = None,
|
||||||
) -> Dict[str, str]:
|
) -> dict[str, str]:
|
||||||
"""
|
"""
|
||||||
构建发送给上游 Provider 的请求头
|
构建发送给上游 Provider 的请求头
|
||||||
|
|
||||||
@@ -386,10 +387,10 @@ def build_upstream_headers(
|
|||||||
|
|
||||||
|
|
||||||
def merge_headers_with_protection(
|
def merge_headers_with_protection(
|
||||||
base_headers: Dict[str, str],
|
base_headers: dict[str, str],
|
||||||
extra_headers: Optional[Dict[str, str]],
|
extra_headers: dict[str, str] | None,
|
||||||
protected_keys: FrozenSet[str] | Set[str],
|
protected_keys: frozenset[str] | set[str],
|
||||||
) -> Dict[str, str]:
|
) -> dict[str, str]:
|
||||||
"""
|
"""
|
||||||
合并头部但保护指定的 key 不被覆盖
|
合并头部但保护指定的 key 不被覆盖
|
||||||
|
|
||||||
@@ -418,9 +419,9 @@ def merge_headers_with_protection(
|
|||||||
|
|
||||||
|
|
||||||
def filter_response_headers(
|
def filter_response_headers(
|
||||||
headers: Optional[Dict[str, str]],
|
headers: dict[str, str] | None,
|
||||||
drop_headers: Optional[FrozenSet[str]] = None,
|
drop_headers: frozenset[str] | None = None,
|
||||||
) -> Dict[str, str]:
|
) -> dict[str, str]:
|
||||||
"""
|
"""
|
||||||
过滤上游响应头中不应透传给客户端的字段
|
过滤上游响应头中不应透传给客户端的字段
|
||||||
|
|
||||||
@@ -446,9 +447,9 @@ def filter_response_headers(
|
|||||||
|
|
||||||
|
|
||||||
def redact_headers_for_log(
|
def redact_headers_for_log(
|
||||||
headers: Dict[str, str],
|
headers: dict[str, str],
|
||||||
redact_keys: Optional[FrozenSet[str]] = None,
|
redact_keys: frozenset[str] | None = None,
|
||||||
) -> Dict[str, str]:
|
) -> dict[str, str]:
|
||||||
"""
|
"""
|
||||||
将敏感头部值替换为 *** 用于日志记录
|
将敏感头部值替换为 *** 用于日志记录
|
||||||
|
|
||||||
@@ -487,7 +488,7 @@ def build_adapter_base_headers(
|
|||||||
api_key: str,
|
api_key: str,
|
||||||
*,
|
*,
|
||||||
include_extra: bool = True,
|
include_extra: bool = True,
|
||||||
) -> Dict[str, str]:
|
) -> dict[str, str]:
|
||||||
"""
|
"""
|
||||||
根据 API 格式构建基础请求头
|
根据 API 格式构建基础请求头
|
||||||
|
|
||||||
@@ -504,7 +505,7 @@ def build_adapter_base_headers(
|
|||||||
auth_header, auth_type = get_auth_config(api_format)
|
auth_header, auth_type = get_auth_config(api_format)
|
||||||
auth_value = f"Bearer {api_key}" if auth_type == "bearer" else api_key
|
auth_value = f"Bearer {api_key}" if auth_type == "bearer" else api_key
|
||||||
|
|
||||||
headers: Dict[str, str] = {
|
headers: dict[str, str] = {
|
||||||
auth_header: auth_value,
|
auth_header: auth_value,
|
||||||
"Content-Type": "application/json",
|
"Content-Type": "application/json",
|
||||||
}
|
}
|
||||||
@@ -520,8 +521,8 @@ def build_adapter_base_headers(
|
|||||||
def build_adapter_headers(
|
def build_adapter_headers(
|
||||||
api_format: APIFormat,
|
api_format: APIFormat,
|
||||||
api_key: str,
|
api_key: str,
|
||||||
extra_headers: Optional[Dict[str, str]] = None,
|
extra_headers: dict[str, str] | None = None,
|
||||||
) -> Dict[str, str]:
|
) -> dict[str, str]:
|
||||||
"""
|
"""
|
||||||
构建完整的 Adapter 请求头
|
构建完整的 Adapter 请求头
|
||||||
|
|
||||||
@@ -565,8 +566,8 @@ def get_adapter_protected_keys(api_format: APIFormat) -> tuple[str, ...]:
|
|||||||
|
|
||||||
|
|
||||||
def extract_set_headers_from_rules(
|
def extract_set_headers_from_rules(
|
||||||
header_rules: Optional[list[Dict[str, Any]]],
|
header_rules: list[dict[str, Any]] | None,
|
||||||
) -> Optional[Dict[str, str]]:
|
) -> dict[str, str] | None:
|
||||||
"""
|
"""
|
||||||
从 header_rules 中提取 set 操作生成的头部字典
|
从 header_rules 中提取 set 操作生成的头部字典
|
||||||
|
|
||||||
@@ -582,7 +583,7 @@ def extract_set_headers_from_rules(
|
|||||||
if not header_rules:
|
if not header_rules:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
headers: Dict[str, str] = {}
|
headers: dict[str, str] = {}
|
||||||
for rule in header_rules:
|
for rule in header_rules:
|
||||||
if rule.get("action") == "set":
|
if rule.get("action") == "set":
|
||||||
key = rule.get("key", "")
|
key = rule.get("key", "")
|
||||||
@@ -593,7 +594,7 @@ def extract_set_headers_from_rules(
|
|||||||
return headers if headers else None
|
return headers if headers else None
|
||||||
|
|
||||||
|
|
||||||
def get_extra_headers_from_endpoint(endpoint: Any) -> Optional[Dict[str, str]]:
|
def get_extra_headers_from_endpoint(endpoint: Any) -> dict[str, str] | None:
|
||||||
"""
|
"""
|
||||||
从 endpoint 提取额外请求头
|
从 endpoint 提取额外请求头
|
||||||
|
|
||||||
|
|||||||
@@ -13,13 +13,12 @@ API 格式元数据定义
|
|||||||
definition = get_api_format_definition(APIFormat.CLAUDE)
|
definition = get_api_format_definition(APIFormat.CLAUDE)
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import re
|
import re
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from functools import lru_cache
|
from functools import lru_cache
|
||||||
from types import MappingProxyType
|
from types import MappingProxyType
|
||||||
from typing import Dict, Iterable, List, Mapping, MutableMapping, Optional, Sequence, Union
|
from collections.abc import Iterable, Mapping, MutableMapping, Sequence
|
||||||
|
|
||||||
from .enums import APIFormat
|
from .enums import APIFormat
|
||||||
|
|
||||||
@@ -64,7 +63,7 @@ class ApiFormatDefinition:
|
|||||||
yield normalized
|
yield normalized
|
||||||
|
|
||||||
|
|
||||||
_DEFINITIONS: Dict[APIFormat, ApiFormatDefinition] = {
|
_DEFINITIONS: dict[APIFormat, ApiFormatDefinition] = {
|
||||||
APIFormat.CLAUDE: ApiFormatDefinition(
|
APIFormat.CLAUDE: ApiFormatDefinition(
|
||||||
api_format=APIFormat.CLAUDE,
|
api_format=APIFormat.CLAUDE,
|
||||||
aliases=("claude", "anthropic", "claude_compatible"),
|
aliases=("claude", "anthropic", "claude_compatible"),
|
||||||
@@ -151,12 +150,12 @@ def get_api_format_definition(api_format: APIFormat) -> ApiFormatDefinition:
|
|||||||
return API_FORMAT_DEFINITIONS[api_format]
|
return API_FORMAT_DEFINITIONS[api_format]
|
||||||
|
|
||||||
|
|
||||||
def list_api_format_definitions() -> List[ApiFormatDefinition]:
|
def list_api_format_definitions() -> list[ApiFormatDefinition]:
|
||||||
"""返回所有定义的浅拷贝列表,供遍历使用。"""
|
"""返回所有定义的浅拷贝列表,供遍历使用。"""
|
||||||
return list(API_FORMAT_DEFINITIONS.values())
|
return list(API_FORMAT_DEFINITIONS.values())
|
||||||
|
|
||||||
|
|
||||||
def build_alias_lookup() -> Dict[str, APIFormat]:
|
def build_alias_lookup() -> dict[str, APIFormat]:
|
||||||
"""
|
"""
|
||||||
构建 alias -> APIFormat 的查找表。
|
构建 alias -> APIFormat 的查找表。
|
||||||
每次调用都会返回新的 dict,避免可变全局引发并发问题。
|
每次调用都会返回新的 dict,避免可变全局引发并发问题。
|
||||||
@@ -237,7 +236,7 @@ def get_protected_keys(api_format: APIFormat) -> frozenset[str]:
|
|||||||
return frozenset({"authorization", "content-type"})
|
return frozenset({"authorization", "content-type"})
|
||||||
|
|
||||||
|
|
||||||
def get_data_format_id(api_format: Union[str, APIFormat]) -> str:
|
def get_data_format_id(api_format: str | APIFormat) -> str:
|
||||||
"""
|
"""
|
||||||
获取格式的数据格式标识。
|
获取格式的数据格式标识。
|
||||||
|
|
||||||
@@ -264,7 +263,7 @@ def get_data_format_id(api_format: Union[str, APIFormat]) -> str:
|
|||||||
return api_format.value.lower()
|
return api_format.value.lower()
|
||||||
|
|
||||||
|
|
||||||
def can_passthrough(client_format: Union[str, APIFormat], endpoint_format: Union[str, APIFormat]) -> bool:
|
def can_passthrough(client_format: str | APIFormat, endpoint_format: str | APIFormat) -> bool:
|
||||||
"""
|
"""
|
||||||
判断两个格式之间是否可以透传(不需要数据转换)。
|
判断两个格式之间是否可以透传(不需要数据转换)。
|
||||||
|
|
||||||
@@ -294,12 +293,12 @@ def can_passthrough(client_format: Union[str, APIFormat], endpoint_format: Union
|
|||||||
|
|
||||||
|
|
||||||
@lru_cache(maxsize=1)
|
@lru_cache(maxsize=1)
|
||||||
def _alias_lookup_cache() -> Dict[str, APIFormat]:
|
def _alias_lookup_cache() -> dict[str, APIFormat]:
|
||||||
"""缓存 alias -> APIFormat 查找表,减少重复构建。"""
|
"""缓存 alias -> APIFormat 查找表,减少重复构建。"""
|
||||||
return build_alias_lookup()
|
return build_alias_lookup()
|
||||||
|
|
||||||
|
|
||||||
def resolve_api_format_alias(value: str) -> Optional[APIFormat]:
|
def resolve_api_format_alias(value: str) -> APIFormat | None:
|
||||||
"""根据别名查找 APIFormat,找不到时返回 None。"""
|
"""根据别名查找 APIFormat,找不到时返回 None。"""
|
||||||
if not value:
|
if not value:
|
||||||
return None
|
return None
|
||||||
@@ -310,9 +309,9 @@ def resolve_api_format_alias(value: str) -> Optional[APIFormat]:
|
|||||||
|
|
||||||
|
|
||||||
def resolve_api_format(
|
def resolve_api_format(
|
||||||
value: Union[str, APIFormat, None],
|
value: str | APIFormat | None,
|
||||||
default: Optional[APIFormat] = None,
|
default: APIFormat | None = None,
|
||||||
) -> Optional[APIFormat]:
|
) -> APIFormat | None:
|
||||||
"""
|
"""
|
||||||
将任意字符串/枚举值解析为 APIFormat。
|
将任意字符串/枚举值解析为 APIFormat。
|
||||||
|
|
||||||
|
|||||||
@@ -6,13 +6,13 @@ API 格式工具函数
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from typing import TYPE_CHECKING, Optional, Union
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from src.core.api_format.enums import APIFormat
|
from src.core.api_format.enums import APIFormat
|
||||||
|
|
||||||
|
|
||||||
def is_cli_format(format_id: Union[str, "APIFormat", None]) -> bool:
|
def is_cli_format(format_id: str | APIFormat | None) -> bool:
|
||||||
"""
|
"""
|
||||||
判断是否为 CLI 透传格式
|
判断是否为 CLI 透传格式
|
||||||
|
|
||||||
@@ -40,7 +40,7 @@ def is_cli_format(format_id: Union[str, "APIFormat", None]) -> bool:
|
|||||||
return str(format_id).upper().endswith("_CLI")
|
return str(format_id).upper().endswith("_CLI")
|
||||||
|
|
||||||
|
|
||||||
def get_base_format(format_id: Union[str, "APIFormat", None]) -> Optional[str]:
|
def get_base_format(format_id: str | APIFormat | None) -> str | None:
|
||||||
"""
|
"""
|
||||||
获取基础格式(去除 _CLI 后缀)
|
获取基础格式(去除 _CLI 后缀)
|
||||||
|
|
||||||
@@ -66,7 +66,7 @@ def get_base_format(format_id: Union[str, "APIFormat", None]) -> Optional[str]:
|
|||||||
return format_str
|
return format_str
|
||||||
|
|
||||||
|
|
||||||
def normalize_format(format_id: Union[str, "APIFormat", None]) -> Optional[str]:
|
def normalize_format(format_id: str | APIFormat | None) -> str | None:
|
||||||
"""
|
"""
|
||||||
规范化格式标识符
|
规范化格式标识符
|
||||||
|
|
||||||
@@ -84,8 +84,8 @@ def normalize_format(format_id: Union[str, "APIFormat", None]) -> Optional[str]:
|
|||||||
|
|
||||||
|
|
||||||
def is_same_format(
|
def is_same_format(
|
||||||
format1: Union[str, "APIFormat", None],
|
format1: str | APIFormat | None,
|
||||||
format2: Union[str, "APIFormat", None],
|
format2: str | APIFormat | None,
|
||||||
) -> bool:
|
) -> bool:
|
||||||
"""
|
"""
|
||||||
判断两个格式是否相同
|
判断两个格式是否相同
|
||||||
@@ -95,7 +95,7 @@ def is_same_format(
|
|||||||
return normalize_format(format1) == normalize_format(format2)
|
return normalize_format(format1) == normalize_format(format2)
|
||||||
|
|
||||||
|
|
||||||
def is_convertible_format(format_id: Union[str, "APIFormat", None]) -> bool:
|
def is_convertible_format(format_id: str | APIFormat | None) -> bool:
|
||||||
"""
|
"""
|
||||||
判断是否为可转换格式
|
判断是否为可转换格式
|
||||||
|
|
||||||
|
|||||||
@@ -8,7 +8,6 @@
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
from typing import Set
|
|
||||||
|
|
||||||
from src.core.logger import logger
|
from src.core.logger import logger
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
@@ -23,7 +22,7 @@ class BatchCommitter:
|
|||||||
interval_seconds: 批量提交间隔(秒)
|
interval_seconds: 批量提交间隔(秒)
|
||||||
"""
|
"""
|
||||||
self.interval_seconds = interval_seconds
|
self.interval_seconds = interval_seconds
|
||||||
self._pending_sessions: Set[Session] = set()
|
self._pending_sessions: set[Session] = set()
|
||||||
self._lock = asyncio.Lock()
|
self._lock = asyncio.Lock()
|
||||||
self._task = None
|
self._task = None
|
||||||
|
|
||||||
|
|||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user