mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 09:20:22 +08:00
chore: 升级到 Python 3.14 并现代化代码
- 升级 Docker 基础镜像从 Python 3.12 到 3.14 - 更新 pyproject.toml 支持 Python 3.13/3.14 - 移除 Python 3.8/3.9/3.10/3.11 分类器 - 更新 black 和 mypy 配置目标版本 - 将 get_event_loop() 替换为 get_running_loop() 加上 RuntimeError 处理 - 简化 compute_cost_sync 中的 asyncio.run 使用 - Dict/List/Tuple/Set → dict/list/tuple/set (PEP 585) - Optional[T] → T | None (PEP 604) - Union[A, B] → A | B (PEP 604) - 移除废弃的 typing 导入 - 移除不必要的字符串引号注解
This commit is contained in:
@@ -10,7 +10,6 @@
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from pydantic import BaseModel, Field, ValidationError
|
||||
@@ -34,7 +33,7 @@ class EnableAdaptiveRequest(BaseModel):
|
||||
"""启用自适应模式请求"""
|
||||
|
||||
enabled: bool = Field(..., description="是否启用自适应模式(true=自适应,false=固定限制)")
|
||||
fixed_limit: Optional[int] = Field(
|
||||
fixed_limit: int | None = Field(
|
||||
None, ge=1, le=100, description="固定 RPM 限制(仅当 enabled=false 时生效,1-100)"
|
||||
)
|
||||
|
||||
@@ -43,30 +42,30 @@ class AdaptiveStatsResponse(BaseModel):
|
||||
"""自适应统计响应"""
|
||||
|
||||
adaptive_mode: bool = Field(..., description="是否为自适应模式(rpm_limit=NULL)")
|
||||
rpm_limit: Optional[int] = Field(None, description="用户配置的固定限制(NULL=自适应)")
|
||||
effective_limit: Optional[int] = Field(
|
||||
rpm_limit: int | None = Field(None, description="用户配置的固定限制(NULL=自适应)")
|
||||
effective_limit: int | None = Field(
|
||||
None, description="当前有效限制(自适应使用学习值,固定使用配置值)"
|
||||
)
|
||||
learned_limit: Optional[int] = Field(None, description="学习到的 RPM 限制")
|
||||
learned_limit: int | None = Field(None, description="学习到的 RPM 限制")
|
||||
concurrent_429_count: int
|
||||
rpm_429_count: int
|
||||
last_429_at: Optional[str]
|
||||
last_429_type: Optional[str]
|
||||
last_429_at: str | None
|
||||
last_429_type: str | None
|
||||
adjustment_count: int
|
||||
recent_adjustments: List[dict]
|
||||
recent_adjustments: list[dict]
|
||||
|
||||
|
||||
class KeyListItem(BaseModel):
|
||||
"""Key 列表项"""
|
||||
|
||||
id: str
|
||||
name: Optional[str]
|
||||
name: str | None
|
||||
provider_id: str
|
||||
api_formats: List[str] = Field(default_factory=list)
|
||||
api_formats: list[str] = Field(default_factory=list)
|
||||
is_adaptive: bool = Field(..., description="是否为自适应模式(rpm_limit=NULL)")
|
||||
rpm_limit: Optional[int] = Field(None, description="固定 RPM 限制(NULL=自适应)")
|
||||
effective_limit: Optional[int] = Field(None, description="当前有效限制")
|
||||
learned_rpm_limit: Optional[int] = Field(None, description="学习到的 RPM 限制")
|
||||
rpm_limit: int | None = Field(None, description="固定 RPM 限制(NULL=自适应)")
|
||||
effective_limit: int | None = Field(None, description="当前有效限制")
|
||||
learned_rpm_limit: int | None = Field(None, description="学习到的 RPM 限制")
|
||||
concurrent_429_count: int
|
||||
rpm_429_count: int
|
||||
|
||||
@@ -76,12 +75,12 @@ class KeyListItem(BaseModel):
|
||||
|
||||
@router.get(
|
||||
"/keys",
|
||||
response_model=List[KeyListItem],
|
||||
response_model=list[KeyListItem],
|
||||
summary="获取所有启用自适应模式的Key",
|
||||
)
|
||||
async def list_adaptive_keys(
|
||||
request: Request,
|
||||
provider_id: Optional[str] = Query(None, description="按 Provider 过滤"),
|
||||
provider_id: str | None = Query(None, description="按 Provider 过滤"),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
"""
|
||||
@@ -207,7 +206,7 @@ async def get_adaptive_summary(
|
||||
|
||||
@dataclass
|
||||
class ListAdaptiveKeysAdapter(AdminApiAdapter):
|
||||
provider_id: Optional[str] = None
|
||||
provider_id: str | None = None
|
||||
|
||||
async def handle(self, context): # type: ignore[override]
|
||||
# 自适应模式:rpm_limit = NULL
|
||||
|
||||
@@ -5,7 +5,6 @@
|
||||
|
||||
import os
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Optional
|
||||
from zoneinfo import ZoneInfo
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
@@ -25,7 +24,7 @@ from src.services.user.apikey import ApiKeyService
|
||||
APP_TIMEZONE = ZoneInfo(os.getenv("APP_TIMEZONE", "Asia/Shanghai"))
|
||||
|
||||
|
||||
def parse_expiry_date(date_str: Optional[str]) -> Optional[datetime]:
|
||||
def parse_expiry_date(date_str: str | None) -> datetime | None:
|
||||
"""解析过期日期字符串为 datetime 对象。
|
||||
|
||||
Args:
|
||||
@@ -70,7 +69,7 @@ async def list_standalone_api_keys(
|
||||
request: Request,
|
||||
skip: int = Query(0, ge=0),
|
||||
limit: int = Query(100, ge=1, le=500),
|
||||
is_active: Optional[bool] = None,
|
||||
is_active: bool | None = None,
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
"""
|
||||
@@ -330,7 +329,7 @@ class AdminListStandaloneKeysAdapter(AdminApiAdapter):
|
||||
self,
|
||||
skip: int,
|
||||
limit: int,
|
||||
is_active: Optional[bool],
|
||||
is_active: bool | None,
|
||||
):
|
||||
self.skip = skip
|
||||
self.limit = limit
|
||||
|
||||
@@ -5,7 +5,6 @@ Endpoint 健康监控 API
|
||||
from collections import defaultdict
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, Request
|
||||
from sqlalchemy import func
|
||||
@@ -128,7 +127,7 @@ async def get_api_format_health_monitor(
|
||||
async def get_key_health(
|
||||
key_id: str,
|
||||
request: Request,
|
||||
api_format: Optional[str] = Query(None, description="API 格式(可选,如 CLAUDE、OPENAI)"),
|
||||
api_format: str | None = Query(None, description="API 格式(可选,如 CLAUDE、OPENAI)"),
|
||||
db: Session = Depends(get_db),
|
||||
) -> HealthStatusResponse:
|
||||
"""
|
||||
@@ -161,7 +160,7 @@ async def get_key_health(
|
||||
async def recover_key_health(
|
||||
key_id: str,
|
||||
request: Request,
|
||||
api_format: Optional[str] = Query(None, description="API 格式(可选,不指定则恢复所有格式)"),
|
||||
api_format: str | None = Query(None, description="API 格式(可选,不指定则恢复所有格式)"),
|
||||
db: Session = Depends(get_db),
|
||||
) -> dict:
|
||||
"""
|
||||
@@ -278,7 +277,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
||||
)
|
||||
|
||||
# 构建所有格式的 provider_count 映射
|
||||
all_formats: Dict[str, int] = {}
|
||||
all_formats: dict[str, int] = {}
|
||||
for api_format_enum, provider_count in active_formats:
|
||||
api_format = (
|
||||
api_format_enum.value if hasattr(api_format_enum, "value") else str(api_format_enum)
|
||||
@@ -295,7 +294,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
||||
)
|
||||
.all()
|
||||
)
|
||||
endpoint_map: Dict[str, List[str]] = defaultdict(list)
|
||||
endpoint_map: dict[str, list[str]] = defaultdict(list)
|
||||
active_provider_formats: set[tuple[str, str]] = set()
|
||||
for api_format_enum, endpoint_id, provider_id in endpoint_rows:
|
||||
api_format = (
|
||||
@@ -305,7 +304,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
||||
active_provider_formats.add((str(provider_id), api_format))
|
||||
|
||||
# 1.2 统计每个 API 格式可用的活跃 Key 数量(Key 属于 Provider,通过 api_formats 关联格式)
|
||||
key_counts: Dict[str, int] = {}
|
||||
key_counts: dict[str, int] = {}
|
||||
if active_provider_formats:
|
||||
active_provider_keys = (
|
||||
db.query(ProviderAPIKey.provider_id, ProviderAPIKey.api_formats)
|
||||
@@ -342,7 +341,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
||||
)
|
||||
|
||||
# 构建每个格式的状态统计
|
||||
status_counts: Dict[str, Dict[str, int]] = {}
|
||||
status_counts: dict[str, dict[str, int]] = {}
|
||||
for api_format_enum, status, count in status_counts_query:
|
||||
api_format = (
|
||||
api_format_enum.value if hasattr(api_format_enum, "value") else str(api_format_enum)
|
||||
@@ -370,7 +369,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
||||
.all()
|
||||
)
|
||||
|
||||
grouped_attempts: Dict[str, List[RequestCandidate]] = {}
|
||||
grouped_attempts: dict[str, list[RequestCandidate]] = {}
|
||||
|
||||
for attempt, api_format_enum, provider_id in rows:
|
||||
api_format = (
|
||||
@@ -384,7 +383,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
||||
grouped_attempts[api_format].append(attempt)
|
||||
|
||||
# 4. 为所有活跃格式生成监控数据(包括没有请求记录的)
|
||||
monitors: List[ApiFormatHealthMonitor] = []
|
||||
monitors: list[ApiFormatHealthMonitor] = []
|
||||
for api_format in all_formats:
|
||||
attempts = grouped_attempts.get(api_format, [])
|
||||
# 获取窗口内的真实统计数据
|
||||
@@ -399,7 +398,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
||||
|
||||
# 时间线按时间正序
|
||||
attempts_sorted = list(reversed(attempts))
|
||||
events: List[EndpointHealthEvent] = []
|
||||
events: list[EndpointHealthEvent] = []
|
||||
for attempt in attempts_sorted:
|
||||
event_timestamp = attempt.finished_at or attempt.started_at or attempt.created_at
|
||||
events.append(
|
||||
@@ -462,7 +461,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
|
||||
@dataclass
|
||||
class AdminKeyHealthAdapter(AdminApiAdapter):
|
||||
key_id: str
|
||||
api_format: Optional[str] = None
|
||||
api_format: str | None = None
|
||||
|
||||
async def handle(self, context): # type: ignore[override]
|
||||
health_data = health_monitor.get_key_health(context.db, self.key_id, self.api_format)
|
||||
@@ -500,7 +499,7 @@ class AdminKeyHealthAdapter(AdminApiAdapter):
|
||||
@dataclass
|
||||
class AdminRecoverKeyHealthAdapter(AdminApiAdapter):
|
||||
key_id: str
|
||||
api_format: Optional[str] = None
|
||||
api_format: str | None = None
|
||||
|
||||
async def handle(self, context): # type: ignore[override]
|
||||
db = context.db
|
||||
|
||||
@@ -5,7 +5,6 @@ Provider API Keys 管理
|
||||
import uuid
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from typing import Dict, List
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, Request
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -145,14 +144,14 @@ async def delete_endpoint_key(
|
||||
# ========== Provider Keys API ==========
|
||||
|
||||
|
||||
@router.get("/providers/{provider_id}/keys", response_model=List[EndpointAPIKeyResponse])
|
||||
@router.get("/providers/{provider_id}/keys", response_model=list[EndpointAPIKeyResponse])
|
||||
async def list_provider_keys(
|
||||
provider_id: str,
|
||||
request: Request,
|
||||
skip: int = Query(0, ge=0, description="跳过的记录数"),
|
||||
limit: int = Query(100, ge=1, le=1000, description="返回的最大记录数"),
|
||||
db: Session = Depends(get_db),
|
||||
) -> List[EndpointAPIKeyResponse]:
|
||||
) -> list[EndpointAPIKeyResponse]:
|
||||
"""
|
||||
获取 Provider 的所有 Keys
|
||||
|
||||
@@ -439,12 +438,12 @@ class AdminGetKeysGroupedByFormatAdapter(AdminApiAdapter):
|
||||
)
|
||||
.all()
|
||||
)
|
||||
endpoint_base_url_map: Dict[tuple[str, str], str] = {}
|
||||
endpoint_base_url_map: dict[tuple[str, str], str] = {}
|
||||
for provider_id, api_format, base_url in endpoints:
|
||||
fmt = api_format.value if hasattr(api_format, "value") else str(api_format)
|
||||
endpoint_base_url_map[(str(provider_id), fmt)] = base_url
|
||||
|
||||
grouped: Dict[str, List[dict]] = {}
|
||||
grouped: dict[str, list[dict]] = {}
|
||||
for key, provider in keys:
|
||||
api_formats = key.api_formats or []
|
||||
|
||||
|
||||
@@ -5,10 +5,9 @@ ProviderEndpoint CRUD 管理 API
|
||||
import uuid
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from typing import List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, Request
|
||||
from sqlalchemy import and_, func
|
||||
from sqlalchemy import and_
|
||||
from sqlalchemy.orm import Session
|
||||
from sqlalchemy.orm.attributes import flag_modified
|
||||
|
||||
@@ -29,7 +28,7 @@ router = APIRouter(tags=["Endpoint Management"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
|
||||
|
||||
def mask_proxy_password(proxy_config: Optional[dict]) -> Optional[dict]:
|
||||
def mask_proxy_password(proxy_config: dict | None) -> dict | None:
|
||||
"""对代理配置中的密码进行脱敏处理"""
|
||||
if not proxy_config:
|
||||
return None
|
||||
@@ -39,14 +38,14 @@ def mask_proxy_password(proxy_config: Optional[dict]) -> Optional[dict]:
|
||||
return masked
|
||||
|
||||
|
||||
@router.get("/providers/{provider_id}/endpoints", response_model=List[ProviderEndpointResponse])
|
||||
@router.get("/providers/{provider_id}/endpoints", response_model=list[ProviderEndpointResponse])
|
||||
async def list_provider_endpoints(
|
||||
provider_id: str,
|
||||
request: Request,
|
||||
skip: int = Query(0, ge=0, description="跳过的记录数"),
|
||||
limit: int = Query(100, ge=1, le=1000, description="返回的最大记录数"),
|
||||
db: Session = Depends(get_db),
|
||||
) -> List[ProviderEndpointResponse]:
|
||||
) -> list[ProviderEndpointResponse]:
|
||||
"""
|
||||
获取指定 Provider 的所有 Endpoints
|
||||
|
||||
@@ -245,7 +244,7 @@ class AdminListProviderEndpointsAdapter(AdminApiAdapter):
|
||||
if is_active:
|
||||
active_keys_map[fmt] = active_keys_map.get(fmt, 0) + 1
|
||||
|
||||
result: List[ProviderEndpointResponse] = []
|
||||
result: list[ProviderEndpointResponse] = []
|
||||
for endpoint in endpoints:
|
||||
endpoint_format = (
|
||||
endpoint.api_format
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
"""LDAP配置管理API端点。"""
|
||||
|
||||
import re
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, Request
|
||||
from pydantic import BaseModel, Field, ValidationError, field_validator
|
||||
@@ -30,9 +30,9 @@ BCRYPT_HASH_PATTERN = re.compile(r"^\$2[aby]\$\d{2}\$.{53}$")
|
||||
class LDAPConfigResponse(BaseModel):
|
||||
"""LDAP配置响应(不返回密码)"""
|
||||
|
||||
server_url: Optional[str] = None
|
||||
bind_dn: Optional[str] = None
|
||||
base_dn: Optional[str] = None
|
||||
server_url: str | None = None
|
||||
bind_dn: str | None = None
|
||||
base_dn: str | None = None
|
||||
has_bind_password: bool = False
|
||||
user_search_filter: str
|
||||
username_attr: str
|
||||
@@ -50,7 +50,7 @@ class LDAPConfigUpdate(BaseModel):
|
||||
server_url: str = Field(..., min_length=1, max_length=255)
|
||||
bind_dn: str = Field(..., min_length=1, max_length=255)
|
||||
# 允许空字符串表示"清除密码";非空时自动 strip 并校验不能为空
|
||||
bind_password: Optional[str] = Field(None, max_length=1024)
|
||||
bind_password: str | None = Field(None, max_length=1024)
|
||||
base_dn: str = Field(..., min_length=1, max_length=255)
|
||||
user_search_filter: str = Field(default="(uid={username})", max_length=500)
|
||||
username_attr: str = Field(default="uid", max_length=50)
|
||||
@@ -63,7 +63,7 @@ class LDAPConfigUpdate(BaseModel):
|
||||
|
||||
@field_validator("bind_password")
|
||||
@classmethod
|
||||
def validate_bind_password(cls, v: Optional[str]) -> Optional[str]:
|
||||
def validate_bind_password(cls, v: str | None) -> str | None:
|
||||
if v is None or v == "":
|
||||
return v
|
||||
v = v.strip()
|
||||
@@ -114,22 +114,22 @@ class LDAPTestResponse(BaseModel):
|
||||
class LDAPConfigTest(BaseModel):
|
||||
"""LDAP配置测试请求(全部可选,用于临时覆盖)"""
|
||||
|
||||
server_url: Optional[str] = Field(None, min_length=1, max_length=255)
|
||||
bind_dn: Optional[str] = Field(None, min_length=1, max_length=255)
|
||||
bind_password: Optional[str] = Field(None, min_length=1)
|
||||
base_dn: Optional[str] = Field(None, min_length=1, max_length=255)
|
||||
user_search_filter: Optional[str] = Field(None, max_length=500)
|
||||
username_attr: Optional[str] = Field(None, max_length=50)
|
||||
email_attr: Optional[str] = Field(None, max_length=50)
|
||||
display_name_attr: Optional[str] = Field(None, max_length=50)
|
||||
is_enabled: Optional[bool] = None
|
||||
is_exclusive: Optional[bool] = None
|
||||
use_starttls: Optional[bool] = None
|
||||
connect_timeout: Optional[int] = Field(None, ge=1, le=60)
|
||||
server_url: str | None = Field(None, min_length=1, max_length=255)
|
||||
bind_dn: str | None = Field(None, min_length=1, max_length=255)
|
||||
bind_password: str | None = Field(None, min_length=1)
|
||||
base_dn: str | None = Field(None, min_length=1, max_length=255)
|
||||
user_search_filter: str | None = Field(None, max_length=500)
|
||||
username_attr: str | None = Field(None, max_length=50)
|
||||
email_attr: str | None = Field(None, max_length=50)
|
||||
display_name_attr: str | None = Field(None, max_length=50)
|
||||
is_enabled: bool | None = None
|
||||
is_exclusive: bool | None = None
|
||||
use_starttls: bool | None = None
|
||||
connect_timeout: int | None = Field(None, ge=1, le=60)
|
||||
|
||||
@field_validator("user_search_filter")
|
||||
@classmethod
|
||||
def validate_search_filter(cls, v: Optional[str]) -> Optional[str]:
|
||||
def validate_search_filter(cls, v: str | None) -> str | None:
|
||||
if v is None:
|
||||
return v
|
||||
if "{username}" not in v:
|
||||
@@ -263,7 +263,7 @@ async def test_ldap_connection(request: Request, db: Session = Depends(get_db))
|
||||
|
||||
|
||||
class AdminGetLDAPConfigAdapter(AdminApiAdapter):
|
||||
async def handle(self, context) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context) -> dict[str, Any]: # type: ignore[override]
|
||||
db = context.db
|
||||
config = db.query(LDAPConfig).first()
|
||||
|
||||
@@ -300,7 +300,7 @@ class AdminGetLDAPConfigAdapter(AdminApiAdapter):
|
||||
|
||||
|
||||
class AdminUpdateLDAPConfigAdapter(AdminApiAdapter):
|
||||
async def handle(self, context) -> Dict[str, str]: # type: ignore[override]
|
||||
async def handle(self, context) -> dict[str, str]: # type: ignore[override]
|
||||
db = context.db
|
||||
payload = context.ensure_json_body()
|
||||
|
||||
@@ -421,7 +421,7 @@ class AdminUpdateLDAPConfigAdapter(AdminApiAdapter):
|
||||
|
||||
|
||||
class AdminTestLDAPConnectionAdapter(AdminApiAdapter):
|
||||
async def handle(self, context) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context) -> dict[str, Any]: # type: ignore[override]
|
||||
from src.services.auth.ldap import LDAPService
|
||||
|
||||
db = context.db
|
||||
@@ -442,7 +442,7 @@ class AdminTestLDAPConnectionAdapter(AdminApiAdapter):
|
||||
raise InvalidRequestException(translate_pydantic_error(errors[0]))
|
||||
raise InvalidRequestException("请求数据验证失败")
|
||||
|
||||
config_data: Dict[str, Any] = {}
|
||||
config_data: dict[str, Any] = {}
|
||||
|
||||
if saved_config:
|
||||
config_data = {
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
"""管理员 Management Token 管理端点"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from fastapi.responses import JSONResponse
|
||||
@@ -46,8 +45,8 @@ class AdminManagementTokenApiAdapter(AdminApiAdapter):
|
||||
@router.get("")
|
||||
async def list_all_management_tokens(
|
||||
request: Request,
|
||||
user_id: Optional[str] = Query(None, description="筛选用户 ID"),
|
||||
is_active: Optional[bool] = Query(None, description="筛选激活状态"),
|
||||
user_id: str | None = Query(None, description="筛选用户 ID"),
|
||||
is_active: bool | None = Query(None, description="筛选激活状态"),
|
||||
skip: int = Query(0, ge=0),
|
||||
limit: int = Query(50, ge=1, le=100),
|
||||
db: Session = Depends(get_db),
|
||||
@@ -174,8 +173,8 @@ class AdminListManagementTokensAdapter(AdminManagementTokenApiAdapter):
|
||||
"""列出所有 Management Tokens"""
|
||||
|
||||
name: str = "admin_list_management_tokens"
|
||||
user_id: Optional[str] = None
|
||||
is_active: Optional[bool] = None
|
||||
user_id: str | None = None
|
||||
is_active: bool | None = None
|
||||
skip: int = 0
|
||||
limit: int = 50
|
||||
|
||||
@@ -197,7 +196,7 @@ class AdminListManagementTokensAdapter(AdminManagementTokenApiAdapter):
|
||||
)
|
||||
|
||||
# 预加载用户信息
|
||||
user_ids = list(set(t.user_id for t in tokens))
|
||||
user_ids = list({t.user_id for t in tokens})
|
||||
users = {u.id: u for u in context.db.query(User).filter(User.id.in_(user_ids)).all()}
|
||||
for token in tokens:
|
||||
token.user = users.get(token.user_id)
|
||||
|
||||
@@ -5,7 +5,6 @@
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Dict, List
|
||||
|
||||
from fastapi import APIRouter, Depends, Request
|
||||
from sqlalchemy.orm import Session, joinedload
|
||||
@@ -64,12 +63,12 @@ class AdminGetModelCatalogAdapter(AdminApiAdapter):
|
||||
db: Session = context.db
|
||||
|
||||
# 1. 获取所有活跃的 GlobalModel
|
||||
global_models: List[GlobalModel] = (
|
||||
global_models: list[GlobalModel] = (
|
||||
db.query(GlobalModel).filter(GlobalModel.is_active == True).all()
|
||||
)
|
||||
|
||||
# 2. 获取所有活跃的 Model 实现(包含 global_model 以便计算有效价格)
|
||||
models: List[Model] = (
|
||||
models: list[Model] = (
|
||||
db.query(Model)
|
||||
.options(joinedload(Model.provider), joinedload(Model.global_model))
|
||||
.filter(Model.is_active == True)
|
||||
@@ -77,17 +76,17 @@ class AdminGetModelCatalogAdapter(AdminApiAdapter):
|
||||
)
|
||||
|
||||
# 按 GlobalModel ID 组织关联提供商
|
||||
models_by_global_model: Dict[str, List[Model]] = {}
|
||||
models_by_global_model: dict[str, list[Model]] = {}
|
||||
for model in models:
|
||||
if model.global_model_id:
|
||||
models_by_global_model.setdefault(model.global_model_id, []).append(model)
|
||||
|
||||
# 3. 为每个 GlobalModel 构建 catalog item
|
||||
catalog_items: List[ModelCatalogItem] = []
|
||||
catalog_items: list[ModelCatalogItem] = []
|
||||
|
||||
for gm in global_models:
|
||||
gm_id = gm.id
|
||||
provider_entries: List[ModelCatalogProviderDetail] = []
|
||||
provider_entries: list[ModelCatalogProviderDetail] = []
|
||||
# 从 config JSON 读取能力标志
|
||||
gm_config = gm.config or {}
|
||||
capability_flags = {
|
||||
|
||||
@@ -3,8 +3,7 @@ models.dev 外部模型数据代理
|
||||
"""
|
||||
|
||||
import json
|
||||
|
||||
from typing import Any, Optional
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
@@ -43,7 +42,7 @@ OFFICIAL_PROVIDERS = {
|
||||
}
|
||||
|
||||
|
||||
async def _get_cached_data() -> Optional[dict[str, Any]]:
|
||||
async def _get_cached_data() -> dict[str, Any] | None:
|
||||
"""从 Redis 获取缓存数据"""
|
||||
redis = await get_redis_client()
|
||||
if redis is None:
|
||||
|
||||
@@ -5,7 +5,6 @@ GlobalModel Admin API
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, Request
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -37,8 +36,8 @@ async def list_global_models(
|
||||
request: Request,
|
||||
skip: int = Query(0, ge=0),
|
||||
limit: int = Query(100, ge=1, le=1000),
|
||||
is_active: Optional[bool] = Query(None),
|
||||
search: Optional[str] = Query(None),
|
||||
is_active: bool | None = Query(None),
|
||||
search: str | None = Query(None),
|
||||
db: Session = Depends(get_db),
|
||||
) -> GlobalModelListResponse:
|
||||
"""
|
||||
@@ -254,8 +253,8 @@ class AdminListGlobalModelsAdapter(AdminApiAdapter):
|
||||
|
||||
skip: int
|
||||
limit: int
|
||||
is_active: Optional[bool]
|
||||
search: Optional[str]
|
||||
is_active: bool | None
|
||||
search: str | None
|
||||
|
||||
async def handle(self, context): # type: ignore[override]
|
||||
from sqlalchemy import func
|
||||
|
||||
@@ -9,7 +9,6 @@ GlobalModel 请求链路预览 API
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, Request
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
@@ -47,22 +46,22 @@ class RoutingKeyInfo(BaseModel):
|
||||
name: str
|
||||
masked_key: str = Field("", description="脱敏的 API Key")
|
||||
internal_priority: int = Field(..., description="Key 内部优先级")
|
||||
global_priority_by_format: Optional[Dict[str, int]] = Field(None, description="按 API 格式的全局优先级")
|
||||
rpm_limit: Optional[int] = Field(None, description="RPM 限制,null 表示自适应")
|
||||
global_priority_by_format: dict[str, int] | None = Field(None, description="按 API 格式的全局优先级")
|
||||
rpm_limit: int | None = Field(None, description="RPM 限制,null 表示自适应")
|
||||
is_adaptive: bool = Field(False, description="是否为自适应 RPM 模式")
|
||||
effective_rpm: Optional[int] = Field(None, description="有效 RPM 限制")
|
||||
effective_rpm: int | None = Field(None, description="有效 RPM 限制")
|
||||
cache_ttl_minutes: int = Field(0, description="缓存 TTL(分钟)")
|
||||
health_score: float = Field(1.0, description="健康度分数(0-1 小数格式)")
|
||||
is_active: bool
|
||||
api_formats: List[str] = Field(default_factory=list, description="支持的 API 格式")
|
||||
api_formats: list[str] = Field(default_factory=list, description="支持的 API 格式")
|
||||
# 模型白名单
|
||||
allowed_models: Optional[List[str]] = Field(None, description="允许的模型列表,null 表示不限制")
|
||||
allowed_models: list[str] | None = Field(None, description="允许的模型列表,null 表示不限制")
|
||||
# 熔断状态
|
||||
circuit_breaker_open: bool = Field(False, description="熔断器是否打开")
|
||||
circuit_breaker_formats: List[str] = Field(
|
||||
circuit_breaker_formats: list[str] = Field(
|
||||
default_factory=list, description="熔断的 API 格式列表"
|
||||
)
|
||||
next_probe_at: Optional[str] = Field(None, description="下次探测时间(ISO格式)")
|
||||
next_probe_at: str | None = Field(None, description="下次探测时间(ISO格式)")
|
||||
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
@@ -73,9 +72,9 @@ class RoutingEndpointInfo(BaseModel):
|
||||
id: str
|
||||
api_format: str
|
||||
base_url: str
|
||||
custom_path: Optional[str] = None
|
||||
custom_path: str | None = None
|
||||
is_active: bool
|
||||
keys: List[RoutingKeyInfo] = Field(default_factory=list)
|
||||
keys: list[RoutingKeyInfo] = Field(default_factory=list)
|
||||
total_keys: int = 0
|
||||
active_keys: int = 0
|
||||
|
||||
@@ -87,7 +86,7 @@ class RoutingModelMapping(BaseModel):
|
||||
|
||||
name: str = Field(..., description="映射名称")
|
||||
priority: int = Field(..., description="优先级(数字越小优先级越高)")
|
||||
api_formats: Optional[List[str]] = Field(None, description="作用域(适用的 API 格式)")
|
||||
api_formats: list[str] | None = Field(None, description="作用域(适用的 API 格式)")
|
||||
|
||||
|
||||
class RoutingProviderInfo(BaseModel):
|
||||
@@ -97,18 +96,18 @@ class RoutingProviderInfo(BaseModel):
|
||||
name: str
|
||||
model_id: str = Field(..., description="Model ID(GlobalModel 与 Provider 的关联记录 ID)")
|
||||
provider_priority: int = Field(..., description="提供商优先级(数字越小优先级越高)")
|
||||
billing_type: Optional[str] = Field(None, description="计费类型")
|
||||
monthly_quota_usd: Optional[float] = Field(None, description="月额度(美元)")
|
||||
monthly_used_usd: Optional[float] = Field(None, description="已用额度(美元)")
|
||||
billing_type: str | None = Field(None, description="计费类型")
|
||||
monthly_quota_usd: float | None = Field(None, description="月额度(美元)")
|
||||
monthly_used_usd: float | None = Field(None, description="已用额度(美元)")
|
||||
is_active: bool
|
||||
# 模型映射信息
|
||||
provider_model_name: str = Field(..., description="提供商侧的模型名称")
|
||||
model_mappings: List[RoutingModelMapping] = Field(
|
||||
model_mappings: list[RoutingModelMapping] = Field(
|
||||
default_factory=list, description="模型名称映射列表"
|
||||
)
|
||||
model_is_active: bool = Field(True, description="Model 是否活跃")
|
||||
# Endpoint 和 Key 信息
|
||||
endpoints: List[RoutingEndpointInfo] = Field(default_factory=list)
|
||||
endpoints: list[RoutingEndpointInfo] = Field(default_factory=list)
|
||||
total_endpoints: int = 0
|
||||
active_endpoints: int = 0
|
||||
|
||||
@@ -123,7 +122,7 @@ class GlobalKeyWhitelistItem(BaseModel):
|
||||
masked_key: str = Field(..., description="脱敏的 API Key")
|
||||
provider_id: str = Field(..., description="Provider ID")
|
||||
provider_name: str = Field(..., description="Provider 名称")
|
||||
allowed_models: List[str] = Field(default_factory=list, description="Key 白名单模型列表")
|
||||
allowed_models: list[str] = Field(default_factory=list, description="Key 白名单模型列表")
|
||||
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
@@ -136,11 +135,11 @@ class ModelRoutingPreviewResponse(BaseModel):
|
||||
display_name: str
|
||||
is_active: bool
|
||||
# GlobalModel 的模型映射(用于前端匹配 Key 白名单)
|
||||
global_model_mappings: List[str] = Field(
|
||||
global_model_mappings: list[str] = Field(
|
||||
default_factory=list, description="GlobalModel 的模型映射规则(正则模式)"
|
||||
)
|
||||
# 链路信息
|
||||
providers: List[RoutingProviderInfo] = Field(
|
||||
providers: list[RoutingProviderInfo] = Field(
|
||||
default_factory=list, description="按优先级排序的提供商列表"
|
||||
)
|
||||
total_providers: int = 0
|
||||
@@ -149,7 +148,7 @@ class ModelRoutingPreviewResponse(BaseModel):
|
||||
scheduling_mode: str = Field("cache_affinity", description="调度模式")
|
||||
priority_mode: str = Field("provider", description="优先级模式")
|
||||
# 全局 Key 白名单数据(供前端实时匹配,包含所有 Provider 的 Key)
|
||||
all_keys_whitelist: List[GlobalKeyWhitelistItem] = Field(
|
||||
all_keys_whitelist: list[GlobalKeyWhitelistItem] = Field(
|
||||
default_factory=list, description="所有 Provider 的 Key 白名单数据"
|
||||
)
|
||||
|
||||
@@ -227,7 +226,7 @@ class AdminGetModelRoutingPreviewAdapter(AdminApiAdapter):
|
||||
provider_ids = [m.provider_id for m in models if m.provider_id]
|
||||
|
||||
# 批量获取 Provider 的 Endpoints
|
||||
endpoints_by_provider: Dict[str, List[ProviderEndpoint]] = {}
|
||||
endpoints_by_provider: dict[str, list[ProviderEndpoint]] = {}
|
||||
if provider_ids:
|
||||
endpoints = (
|
||||
db.query(ProviderEndpoint)
|
||||
@@ -240,7 +239,7 @@ class AdminGetModelRoutingPreviewAdapter(AdminApiAdapter):
|
||||
endpoints_by_provider[ep.provider_id].append(ep)
|
||||
|
||||
# 批量获取 Provider 的 Keys
|
||||
keys_by_provider: Dict[str, List[ProviderAPIKey]] = {}
|
||||
keys_by_provider: dict[str, list[ProviderAPIKey]] = {}
|
||||
if provider_ids:
|
||||
keys = (
|
||||
db.query(ProviderAPIKey).filter(ProviderAPIKey.provider_id.in_(provider_ids)).all()
|
||||
@@ -251,14 +250,14 @@ class AdminGetModelRoutingPreviewAdapter(AdminApiAdapter):
|
||||
keys_by_provider[key.provider_id].append(key)
|
||||
|
||||
# 提取 GlobalModel 的 model_mappings(用于 Key 白名单匹配)
|
||||
global_model_mappings: List[str] = []
|
||||
global_model_mappings: list[str] = []
|
||||
if global_model.config and isinstance(global_model.config, dict):
|
||||
mappings = global_model.config.get("model_mappings")
|
||||
if isinstance(mappings, list):
|
||||
global_model_mappings = [m for m in mappings if isinstance(m, str)]
|
||||
|
||||
# 构建 Provider 路由信息
|
||||
provider_infos: List[RoutingProviderInfo] = []
|
||||
provider_infos: list[RoutingProviderInfo] = []
|
||||
for model in models:
|
||||
provider = model.provider
|
||||
if not provider:
|
||||
@@ -281,7 +280,7 @@ class AdminGetModelRoutingPreviewAdapter(AdminApiAdapter):
|
||||
provider_keys = keys_by_provider.get(provider.id, [])
|
||||
|
||||
# 按 api_format 组织 Keys
|
||||
keys_by_endpoint: Dict[str, List[ProviderAPIKey]] = {}
|
||||
keys_by_endpoint: dict[str, list[ProviderAPIKey]] = {}
|
||||
for key in provider_keys:
|
||||
# 每个 Key 可能支持多个 api_formats
|
||||
for fmt in key.api_formats or []:
|
||||
@@ -355,8 +354,8 @@ class AdminGetModelRoutingPreviewAdapter(AdminApiAdapter):
|
||||
|
||||
# 检查熔断状态
|
||||
circuit_breaker_open = False
|
||||
circuit_breaker_formats: List[str] = []
|
||||
next_probe_at: Optional[str] = None
|
||||
circuit_breaker_formats: list[str] = []
|
||||
next_probe_at: str | None = None
|
||||
if key.circuit_breaker_by_format:
|
||||
for fmt, cb_state in key.circuit_breaker_by_format.items():
|
||||
if isinstance(cb_state, dict) and cb_state.get("open"):
|
||||
@@ -462,7 +461,7 @@ class AdminGetModelRoutingPreviewAdapter(AdminApiAdapter):
|
||||
)
|
||||
|
||||
# 获取所有活跃 Provider 的 Key 白名单数据(供前端实时匹配)
|
||||
all_keys_whitelist: List[GlobalKeyWhitelistItem] = []
|
||||
all_keys_whitelist: list[GlobalKeyWhitelistItem] = []
|
||||
crypto = CryptoService()
|
||||
|
||||
# 获取所有活跃的 Key(带白名单),使用 selectinload 避免 N+1 查询
|
||||
|
||||
@@ -1,7 +1,8 @@
|
||||
"""模块管理 API 端点"""
|
||||
|
||||
from __future__ import annotations
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, Request
|
||||
from pydantic import BaseModel
|
||||
@@ -28,18 +29,18 @@ class ModuleStatusResponse(BaseModel):
|
||||
enabled: bool
|
||||
active: bool
|
||||
config_validated: bool
|
||||
config_error: Optional[str]
|
||||
config_error: str | None
|
||||
display_name: str
|
||||
description: str
|
||||
category: str
|
||||
admin_route: Optional[str]
|
||||
admin_menu_icon: Optional[str]
|
||||
admin_menu_group: Optional[str]
|
||||
admin_route: str | None
|
||||
admin_menu_icon: str | None
|
||||
admin_menu_group: str | None
|
||||
admin_menu_order: int
|
||||
health: str
|
||||
|
||||
@classmethod
|
||||
def from_status(cls, status: ModuleStatus) -> "ModuleStatusResponse":
|
||||
def from_status(cls, status: ModuleStatus) -> ModuleStatusResponse:
|
||||
return cls(
|
||||
name=status.name,
|
||||
available=status.available,
|
||||
@@ -130,7 +131,7 @@ async def set_module_enabled(
|
||||
class AdminGetAllModulesStatusAdapter(AdminApiAdapter):
|
||||
"""获取所有模块状态"""
|
||||
|
||||
async def handle(self, context) -> Dict[str, Any]:
|
||||
async def handle(self, context) -> dict[str, Any]:
|
||||
registry = get_module_registry()
|
||||
all_status = await registry.get_all_status_async(context.db)
|
||||
|
||||
@@ -146,7 +147,7 @@ class AdminGetModuleStatusAdapter(AdminApiAdapter):
|
||||
|
||||
module_name: str
|
||||
|
||||
async def handle(self, context) -> Dict[str, Any]:
|
||||
async def handle(self, context) -> dict[str, Any]:
|
||||
registry = get_module_registry()
|
||||
status = await registry.get_module_status_async(self.module_name, context.db)
|
||||
|
||||
@@ -162,7 +163,7 @@ class AdminSetModuleEnabledAdapter(AdminApiAdapter):
|
||||
|
||||
module_name: str
|
||||
|
||||
async def handle(self, context) -> Dict[str, Any]:
|
||||
async def handle(self, context) -> dict[str, Any]:
|
||||
registry = get_module_registry()
|
||||
|
||||
# 检查模块是否存在
|
||||
|
||||
@@ -2,7 +2,6 @@
|
||||
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from sqlalchemy import func
|
||||
@@ -33,8 +32,8 @@ pipeline = ApiRequestPipeline()
|
||||
@router.get("/audit-logs")
|
||||
async def get_audit_logs(
|
||||
request: Request,
|
||||
username: Optional[str] = Query(None, description="用户名筛选 (模糊匹配)"),
|
||||
event_type: Optional[str] = Query(None, description="事件类型筛选"),
|
||||
username: str | None = Query(None, description="用户名筛选 (模糊匹配)"),
|
||||
event_type: str | None = Query(None, description="事件类型筛选"),
|
||||
days: int = Query(7, description="查询天数"),
|
||||
limit: int = Query(100, description="返回数量限制"),
|
||||
offset: int = Query(0, description="偏移量"),
|
||||
@@ -212,8 +211,8 @@ async def get_circuit_history(
|
||||
|
||||
@dataclass
|
||||
class AdminGetAuditLogsAdapter(AdminApiAdapter):
|
||||
username: Optional[str]
|
||||
event_type: Optional[str]
|
||||
username: str | None
|
||||
event_type: str | None
|
||||
days: int
|
||||
limit: int
|
||||
offset: int
|
||||
@@ -497,8 +496,8 @@ class AdminCircuitHistoryAdapter(AdminApiAdapter):
|
||||
return {"items": history, "count": len(history)}
|
||||
|
||||
|
||||
def _get_health_recommendations(error_stats: dict, health_score: int) -> List[str]:
|
||||
recommendations: List[str] = []
|
||||
def _get_health_recommendations(error_stats: dict, health_score: int) -> list[str]:
|
||||
recommendations: list[str] = []
|
||||
if health_score < 50:
|
||||
recommendations.append("系统健康状况严重,请立即检查错误日志")
|
||||
if error_stats.get("total_errors", 0) > 100:
|
||||
|
||||
@@ -5,7 +5,7 @@
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from fastapi.responses import PlainTextResponse
|
||||
@@ -13,7 +13,7 @@ from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pagination import PaginationMeta, build_pagination_payload, paginate_sequence
|
||||
from src.api.base.pagination import build_pagination_payload, paginate_sequence
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.clients.redis_client import get_redis_client_sync
|
||||
from src.core.crypto import crypto_service
|
||||
@@ -28,7 +28,7 @@ router = APIRouter(prefix="/api/admin/monitoring/cache", tags=["Admin - Monitori
|
||||
pipeline = ApiRequestPipeline()
|
||||
|
||||
|
||||
def mask_api_key(api_key: Optional[str], prefix_len: int = 8, suffix_len: int = 4) -> Optional[str]:
|
||||
def mask_api_key(api_key: str | None, prefix_len: int = 8, suffix_len: int = 4) -> str | None:
|
||||
"""
|
||||
脱敏 API Key,显示前缀 + 星号 + 后缀
|
||||
例如: sk-jhiId-xxxxxxxxxxxAABB -> sk-jhiId-********AABB
|
||||
@@ -47,7 +47,7 @@ def mask_api_key(api_key: Optional[str], prefix_len: int = 8, suffix_len: int =
|
||||
return f"{api_key[:prefix_len]}********{api_key[-suffix_len:]}"
|
||||
|
||||
|
||||
def decrypt_and_mask(encrypted_key: Optional[str], prefix_len: int = 8) -> Optional[str]:
|
||||
def decrypt_and_mask(encrypted_key: str | None, prefix_len: int = 8) -> str | None:
|
||||
"""
|
||||
解密 API Key 后脱敏显示
|
||||
|
||||
@@ -65,7 +65,7 @@ def decrypt_and_mask(encrypted_key: Optional[str], prefix_len: int = 8) -> Optio
|
||||
return None
|
||||
|
||||
|
||||
def resolve_user_identifier(db: Session, identifier: str) -> Optional[str]:
|
||||
def resolve_user_identifier(db: Session, identifier: str) -> str | None:
|
||||
"""
|
||||
将用户标识符(username/email/user_id/api_key_id)解析为 user_id
|
||||
|
||||
@@ -181,7 +181,7 @@ async def get_user_affinity(
|
||||
@router.get("/affinities")
|
||||
async def list_affinities(
|
||||
request: Request,
|
||||
keyword: Optional[str] = None,
|
||||
keyword: str | None = None,
|
||||
limit: int = Query(100, ge=1, le=1000, description="返回数量限制"),
|
||||
offset: int = Query(0, ge=0, description="偏移量"),
|
||||
db: Session = Depends(get_db),
|
||||
@@ -421,7 +421,7 @@ async def get_cache_metrics(
|
||||
|
||||
|
||||
class AdminCacheStatsAdapter(AdminApiAdapter):
|
||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
try:
|
||||
redis_client = get_redis_client_sync()
|
||||
# 读取系统配置,确保监控接口与编排器使用一致的模式
|
||||
@@ -487,14 +487,14 @@ class AdminCacheMetricsAdapter(AdminApiAdapter):
|
||||
logger.exception(f"导出缓存指标失败: {exc}")
|
||||
raise HTTPException(status_code=500, detail=f"导出缓存指标失败: {exc}")
|
||||
|
||||
def _format_prometheus(self, stats: Dict[str, Any]) -> str:
|
||||
def _format_prometheus(self, stats: dict[str, Any]) -> str:
|
||||
"""
|
||||
将 scheduler/affinity 指标转换为 Prometheus 文本格式。
|
||||
"""
|
||||
scheduler_metrics = stats.get("scheduler_metrics", {})
|
||||
affinity_stats = stats.get("affinity_stats", {})
|
||||
|
||||
metric_map: List[Tuple[str, str, float]] = [
|
||||
metric_map: list[tuple[str, str, float]] = [
|
||||
(
|
||||
"cache_scheduler_total_batches",
|
||||
"Total batches pulled from provider list",
|
||||
@@ -542,7 +542,7 @@ class AdminCacheMetricsAdapter(AdminApiAdapter):
|
||||
),
|
||||
]
|
||||
|
||||
affinity_map: List[Tuple[str, str, float]] = [
|
||||
affinity_map: list[tuple[str, str, float]] = [
|
||||
(
|
||||
"cache_affinity_total",
|
||||
"Total cache affinities stored",
|
||||
@@ -596,7 +596,7 @@ class AdminCacheMetricsAdapter(AdminApiAdapter):
|
||||
class AdminGetUserAffinityAdapter(AdminApiAdapter):
|
||||
user_identifier: str
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
db = context.db
|
||||
try:
|
||||
user_id = resolve_user_identifier(db, self.user_identifier)
|
||||
@@ -673,11 +673,11 @@ class AdminGetUserAffinityAdapter(AdminApiAdapter):
|
||||
|
||||
@dataclass
|
||||
class AdminListAffinitiesAdapter(AdminApiAdapter):
|
||||
keyword: Optional[str]
|
||||
keyword: str | None
|
||||
limit: int
|
||||
offset: int
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
db = context.db
|
||||
redis_client = get_redis_client_sync()
|
||||
if not redis_client:
|
||||
@@ -686,7 +686,7 @@ class AdminListAffinitiesAdapter(AdminApiAdapter):
|
||||
affinity_mgr = await get_affinity_manager(redis_client)
|
||||
matched_user_id = None
|
||||
matched_api_key_id = None
|
||||
raw_affinities: List[Dict[str, Any]] = []
|
||||
raw_affinities: list[dict[str, Any]] = []
|
||||
|
||||
if self.keyword:
|
||||
# 首先检查是否是 API Key ID(affinity_key)
|
||||
@@ -724,14 +724,14 @@ class AdminListAffinitiesAdapter(AdminApiAdapter):
|
||||
}
|
||||
|
||||
# 批量查询用户 API Key 信息
|
||||
user_api_key_map: Dict[str, ApiKey] = {}
|
||||
user_api_key_map: dict[str, ApiKey] = {}
|
||||
if affinity_keys:
|
||||
user_api_keys = db.query(ApiKey).filter(ApiKey.id.in_(list(affinity_keys))).all()
|
||||
user_api_key_map = {str(k.id): k for k in user_api_keys}
|
||||
|
||||
# 收集所有 user_id
|
||||
user_ids = {str(k.user_id) for k in user_api_key_map.values()}
|
||||
user_map: Dict[str, User] = {}
|
||||
user_map: dict[str, User] = {}
|
||||
if user_ids:
|
||||
users = db.query(User).filter(User.id.in_(list(user_ids))).all()
|
||||
user_map = {str(user.id): user for user in users}
|
||||
@@ -771,7 +771,7 @@ class AdminListAffinitiesAdapter(AdminApiAdapter):
|
||||
global_model_ids = {
|
||||
item.get("model_name") for item in raw_affinities if item.get("model_name")
|
||||
}
|
||||
global_model_map: Dict[str, GlobalModel] = {}
|
||||
global_model_map: dict[str, GlobalModel] = {}
|
||||
if global_model_ids:
|
||||
# model_name 可能是 UUID 格式的 global_model_id,也可能是原始模型名称
|
||||
global_models = db.query(GlobalModel).filter(
|
||||
@@ -885,7 +885,7 @@ class AdminListAffinitiesAdapter(AdminApiAdapter):
|
||||
class AdminClearUserCacheAdapter(AdminApiAdapter):
|
||||
user_identifier: str
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
db = context.db
|
||||
try:
|
||||
redis_client = get_redis_client_sync()
|
||||
@@ -995,7 +995,7 @@ class AdminClearSingleAffinityAdapter(AdminApiAdapter):
|
||||
model_id: str
|
||||
api_format: str
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
db = context.db
|
||||
try:
|
||||
redis_client = get_redis_client_sync()
|
||||
@@ -1048,7 +1048,7 @@ class AdminClearSingleAffinityAdapter(AdminApiAdapter):
|
||||
|
||||
|
||||
class AdminClearAllCacheAdapter(AdminApiAdapter):
|
||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
try:
|
||||
redis_client = get_redis_client_sync()
|
||||
affinity_mgr = await get_affinity_manager(redis_client)
|
||||
@@ -1068,7 +1068,7 @@ class AdminClearAllCacheAdapter(AdminApiAdapter):
|
||||
class AdminClearProviderCacheAdapter(AdminApiAdapter):
|
||||
provider_id: str
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
try:
|
||||
redis_client = get_redis_client_sync()
|
||||
affinity_mgr = await get_affinity_manager(redis_client)
|
||||
@@ -1091,7 +1091,7 @@ class AdminClearProviderCacheAdapter(AdminApiAdapter):
|
||||
|
||||
|
||||
class AdminCacheConfigAdapter(AdminApiAdapter):
|
||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
from src.services.cache.affinity_manager import CacheAffinityManager
|
||||
from src.config.constants import ConcurrencyDefaults
|
||||
from src.services.rate_limit.adaptive_reservation import get_adaptive_reservation_manager
|
||||
@@ -1260,7 +1260,7 @@ async def clear_provider_model_mapping_cache(
|
||||
|
||||
|
||||
class AdminModelMappingCacheStatsAdapter(AdminApiAdapter):
|
||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
import json
|
||||
|
||||
from src.clients.redis_client import get_redis_client
|
||||
@@ -1510,7 +1510,7 @@ class AdminModelMappingCacheStatsAdapter(AdminApiAdapter):
|
||||
|
||||
|
||||
class AdminClearAllModelMappingCacheAdapter(AdminApiAdapter):
|
||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
from src.clients.redis_client import get_redis_client
|
||||
|
||||
try:
|
||||
@@ -1552,7 +1552,7 @@ class AdminClearAllModelMappingCacheAdapter(AdminApiAdapter):
|
||||
class AdminClearModelMappingCacheByNameAdapter(AdminApiAdapter):
|
||||
model_name: str
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
from src.clients.redis_client import get_redis_client
|
||||
|
||||
try:
|
||||
@@ -1599,7 +1599,7 @@ class AdminClearProviderModelMappingCacheAdapter(AdminApiAdapter):
|
||||
provider_id: str
|
||||
global_model_id: str
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Dict[str, Any]: # type: ignore[override]
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
from src.clients.redis_client import get_redis_client
|
||||
|
||||
try:
|
||||
|
||||
@@ -4,7 +4,6 @@
|
||||
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from typing import List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
@@ -28,29 +27,29 @@ class CandidateResponse(BaseModel):
|
||||
request_id: str
|
||||
candidate_index: int
|
||||
retry_index: int = 0 # 重试序号(从0开始)
|
||||
provider_id: Optional[str] = None
|
||||
provider_name: Optional[str] = None
|
||||
provider_website: Optional[str] = None # Provider 官网
|
||||
endpoint_id: Optional[str] = None
|
||||
endpoint_name: Optional[str] = None # 端点显示名称(api_format)
|
||||
key_id: Optional[str] = None
|
||||
key_name: Optional[str] = None # 密钥名称
|
||||
key_preview: Optional[str] = None # 密钥脱敏预览(如 sk-***abc)
|
||||
key_capabilities: Optional[dict] = None # Key 支持的能力
|
||||
required_capabilities: Optional[dict] = None # 请求实际需要的能力标签
|
||||
provider_id: str | None = None
|
||||
provider_name: str | None = None
|
||||
provider_website: str | None = None # Provider 官网
|
||||
endpoint_id: str | None = None
|
||||
endpoint_name: str | None = None # 端点显示名称(api_format)
|
||||
key_id: str | None = None
|
||||
key_name: str | None = None # 密钥名称
|
||||
key_preview: str | None = None # 密钥脱敏预览(如 sk-***abc)
|
||||
key_capabilities: dict | None = None # Key 支持的能力
|
||||
required_capabilities: dict | None = None # 请求实际需要的能力标签
|
||||
status: str # 'pending', 'success', 'failed', 'skipped'
|
||||
skip_reason: Optional[str] = None
|
||||
skip_reason: str | None = None
|
||||
is_cached: bool = False
|
||||
# 执行结果字段
|
||||
status_code: Optional[int] = None
|
||||
error_type: Optional[str] = None
|
||||
error_message: Optional[str] = None
|
||||
latency_ms: Optional[int] = None
|
||||
concurrent_requests: Optional[int] = None
|
||||
extra_data: Optional[dict] = None
|
||||
status_code: int | None = None
|
||||
error_type: str | None = None
|
||||
error_message: str | None = None
|
||||
latency_ms: int | None = None
|
||||
concurrent_requests: int | None = None
|
||||
extra_data: dict | None = None
|
||||
created_at: datetime
|
||||
started_at: Optional[datetime] = None
|
||||
finished_at: Optional[datetime] = None
|
||||
started_at: datetime | None = None
|
||||
finished_at: datetime | None = None
|
||||
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
@@ -62,7 +61,7 @@ class RequestTraceResponse(BaseModel):
|
||||
total_candidates: int
|
||||
final_status: str # 'success', 'failed', 'cancelled', 'streaming', 'pending'
|
||||
total_latency_ms: int
|
||||
candidates: List[CandidateResponse]
|
||||
candidates: list[CandidateResponse]
|
||||
|
||||
|
||||
@router.get("/{request_id}", response_model=RequestTraceResponse)
|
||||
@@ -253,7 +252,7 @@ class AdminGetRequestTraceAdapter(AdminApiAdapter):
|
||||
key_preview_map[k.id] = "***"
|
||||
|
||||
# 构建 candidate 响应列表
|
||||
candidate_responses: List[CandidateResponse] = []
|
||||
candidate_responses: list[CandidateResponse] = []
|
||||
for candidate in candidates:
|
||||
provider_name = (
|
||||
provider_map.get(candidate.provider_id) if candidate.provider_id else None
|
||||
|
||||
@@ -9,7 +9,7 @@ Provider 操作 API 路由
|
||||
"""
|
||||
|
||||
from dataclasses import asdict, is_dataclass
|
||||
from typing import Any, Dict, List, Optional
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
from pydantic import BaseModel, Field
|
||||
@@ -18,9 +18,7 @@ from sqlalchemy.orm import Session
|
||||
from src.database import get_db
|
||||
from src.models.database import Provider, User
|
||||
from src.services.provider_ops import (
|
||||
ActionStatus,
|
||||
ConnectorAuthType,
|
||||
ConnectorStatus,
|
||||
ProviderActionType,
|
||||
ProviderOpsConfig,
|
||||
ProviderOpsService,
|
||||
@@ -40,46 +38,46 @@ class ArchitectureInfo(BaseModel):
|
||||
architecture_id: str
|
||||
display_name: str
|
||||
description: str
|
||||
supported_auth_types: List[Dict[str, str]]
|
||||
supported_actions: List[Dict[str, Any]]
|
||||
default_connector: Optional[str]
|
||||
supported_auth_types: list[dict[str, str]]
|
||||
supported_actions: list[dict[str, Any]]
|
||||
default_connector: str | None
|
||||
|
||||
|
||||
class ConnectorConfigRequest(BaseModel):
|
||||
"""连接器配置请求"""
|
||||
|
||||
auth_type: str = Field(..., description="认证类型")
|
||||
config: Dict[str, Any] = Field(default_factory=dict, description="连接器配置")
|
||||
credentials: Dict[str, Any] = Field(default_factory=dict, description="凭据信息")
|
||||
config: dict[str, Any] = Field(default_factory=dict, description="连接器配置")
|
||||
credentials: dict[str, Any] = Field(default_factory=dict, description="凭据信息")
|
||||
|
||||
|
||||
class ActionConfigRequest(BaseModel):
|
||||
"""操作配置请求"""
|
||||
|
||||
enabled: bool = Field(True, description="是否启用")
|
||||
config: Dict[str, Any] = Field(default_factory=dict, description="操作配置")
|
||||
config: dict[str, Any] = Field(default_factory=dict, description="操作配置")
|
||||
|
||||
|
||||
class SaveConfigRequest(BaseModel):
|
||||
"""保存配置请求"""
|
||||
|
||||
architecture_id: str = Field("generic_api", description="架构 ID")
|
||||
base_url: Optional[str] = Field(None, description="API 基础地址")
|
||||
base_url: str | None = Field(None, description="API 基础地址")
|
||||
connector: ConnectorConfigRequest
|
||||
actions: Dict[str, ActionConfigRequest] = Field(default_factory=dict)
|
||||
schedule: Dict[str, str] = Field(default_factory=dict, description="定时任务配置")
|
||||
actions: dict[str, ActionConfigRequest] = Field(default_factory=dict)
|
||||
schedule: dict[str, str] = Field(default_factory=dict, description="定时任务配置")
|
||||
|
||||
|
||||
class ConnectRequest(BaseModel):
|
||||
"""连接请求"""
|
||||
|
||||
credentials: Optional[Dict[str, Any]] = Field(None, description="凭据(可选,使用已保存的)")
|
||||
credentials: dict[str, Any] | None = Field(None, description="凭据(可选,使用已保存的)")
|
||||
|
||||
|
||||
class ExecuteActionRequest(BaseModel):
|
||||
"""执行操作请求"""
|
||||
|
||||
config: Optional[Dict[str, Any]] = Field(None, description="操作配置(覆盖默认)")
|
||||
config: dict[str, Any] | None = Field(None, description="操作配置(覆盖默认)")
|
||||
|
||||
|
||||
class ConnectionStatusResponse(BaseModel):
|
||||
@@ -87,9 +85,9 @@ class ConnectionStatusResponse(BaseModel):
|
||||
|
||||
status: str
|
||||
auth_type: str
|
||||
connected_at: Optional[str]
|
||||
expires_at: Optional[str]
|
||||
last_error: Optional[str]
|
||||
connected_at: str | None
|
||||
expires_at: str | None
|
||||
last_error: str | None
|
||||
|
||||
|
||||
class ActionResultResponse(BaseModel):
|
||||
@@ -97,10 +95,10 @@ class ActionResultResponse(BaseModel):
|
||||
|
||||
status: str
|
||||
action_type: str
|
||||
data: Optional[Any]
|
||||
message: Optional[str]
|
||||
data: Any | None
|
||||
message: str | None
|
||||
executed_at: str
|
||||
response_time_ms: Optional[int]
|
||||
response_time_ms: int | None
|
||||
cache_ttl_seconds: int
|
||||
|
||||
|
||||
@@ -109,9 +107,9 @@ class ProviderOpsStatusResponse(BaseModel):
|
||||
|
||||
provider_id: str
|
||||
is_configured: bool
|
||||
architecture_id: Optional[str]
|
||||
architecture_id: str | None
|
||||
connection_status: ConnectionStatusResponse
|
||||
enabled_actions: List[str]
|
||||
enabled_actions: list[str]
|
||||
|
||||
|
||||
class ProviderOpsConfigResponse(BaseModel):
|
||||
@@ -119,17 +117,17 @@ class ProviderOpsConfigResponse(BaseModel):
|
||||
|
||||
provider_id: str
|
||||
is_configured: bool
|
||||
architecture_id: Optional[str] = None
|
||||
base_url: Optional[str] = None
|
||||
connector: Optional[Dict[str, Any]] = None # 脱敏后的连接器配置
|
||||
architecture_id: str | None = None
|
||||
base_url: str | None = None
|
||||
connector: dict[str, Any] | None = None # 脱敏后的连接器配置
|
||||
|
||||
|
||||
class VerifyAuthResponse(BaseModel):
|
||||
"""验证认证响应"""
|
||||
|
||||
success: bool
|
||||
message: Optional[str] = None
|
||||
data: Optional[Dict[str, Any]] = None
|
||||
message: str | None = None
|
||||
data: dict[str, Any] | None = None
|
||||
|
||||
|
||||
# ==================== Helper Functions ====================
|
||||
@@ -147,7 +145,7 @@ def _serialize_data(data: Any) -> Any:
|
||||
# ==================== Routes ====================
|
||||
|
||||
|
||||
@router.get("/architectures", response_model=List[ArchitectureInfo])
|
||||
@router.get("/architectures", response_model=list[ArchitectureInfo])
|
||||
async def list_architectures(_: User = Depends(require_admin)):
|
||||
"""获取所有可用的架构"""
|
||||
registry = get_registry()
|
||||
@@ -490,7 +488,7 @@ async def checkin(
|
||||
|
||||
@router.post("/batch/balance")
|
||||
async def batch_query_balance(
|
||||
provider_ids: Optional[List[str]] = None,
|
||||
provider_ids: list[str] | None = None,
|
||||
db: Session = Depends(get_db),
|
||||
_: User = Depends(require_admin),
|
||||
):
|
||||
|
||||
@@ -4,7 +4,6 @@ Provider Query API 端点
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from typing import Optional
|
||||
|
||||
import httpx
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
@@ -40,7 +39,7 @@ class ModelsQueryRequest(BaseModel):
|
||||
"""模型列表查询请求"""
|
||||
|
||||
provider_id: str
|
||||
api_key_id: Optional[str] = None
|
||||
api_key_id: str | None = None
|
||||
force_refresh: bool = False # 强制刷新,跳过缓存
|
||||
|
||||
|
||||
@@ -49,11 +48,11 @@ class TestModelRequest(BaseModel):
|
||||
|
||||
provider_id: str
|
||||
model_name: str
|
||||
api_key_id: Optional[str] = None
|
||||
endpoint_id: Optional[str] = None # 指定使用的端点ID
|
||||
api_key_id: str | None = None
|
||||
endpoint_id: str | None = None # 指定使用的端点ID
|
||||
stream: bool = False
|
||||
message: Optional[str] = "你好"
|
||||
api_format: Optional[str] = None # 指定使用的API格式,如果不指定则使用端点的默认格式
|
||||
message: str | None = "你好"
|
||||
api_format: str | None = None # 指定使用的API格式,如果不指定则使用端点的默认格式
|
||||
|
||||
|
||||
# ============ API Endpoints ============
|
||||
|
||||
@@ -3,7 +3,6 @@
|
||||
"""
|
||||
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request
|
||||
from fastapi.responses import JSONResponse
|
||||
@@ -25,11 +24,11 @@ pipeline = ApiRequestPipeline()
|
||||
|
||||
class ProviderBillingUpdate(BaseModel):
|
||||
billing_type: ProviderBillingType
|
||||
monthly_quota_usd: Optional[float] = None
|
||||
monthly_quota_usd: float | None = None
|
||||
quota_reset_day: int = Field(default=30, ge=1, le=365) # 重置周期(天数)
|
||||
quota_last_reset_at: Optional[str] = None # 当前周期开始时间
|
||||
quota_expires_at: Optional[str] = None
|
||||
rpm_limit: Optional[int] = Field(default=None, ge=0)
|
||||
quota_last_reset_at: str | None = None # 当前周期开始时间
|
||||
quota_expires_at: str | None = None
|
||||
rpm_limit: int | None = Field(default=None, ge=0)
|
||||
provider_priority: int = Field(default=100, ge=0, le=200)
|
||||
|
||||
|
||||
@@ -163,13 +162,12 @@ class AdminProviderBillingAdapter(AdminApiAdapter):
|
||||
provider.quota_reset_day = config.quota_reset_day
|
||||
provider.provider_priority = config.provider_priority
|
||||
|
||||
from dateutil import parser
|
||||
from sqlalchemy import func
|
||||
|
||||
from src.models.database import Usage
|
||||
|
||||
if config.quota_last_reset_at:
|
||||
new_reset_at = parser.parse(config.quota_last_reset_at)
|
||||
new_reset_at = datetime.fromisoformat(config.quota_last_reset_at)
|
||||
# 确保有时区信息,如果没有则假设为 UTC
|
||||
if new_reset_at.tzinfo is None:
|
||||
new_reset_at = new_reset_at.replace(tzinfo=timezone.utc)
|
||||
@@ -188,7 +186,7 @@ class AdminProviderBillingAdapter(AdminApiAdapter):
|
||||
logger.info(f"Synced usage for provider {provider.name}: ${period_usage:.4f} since {new_reset_at}")
|
||||
|
||||
if config.quota_expires_at:
|
||||
expires_at = parser.parse(config.quota_expires_at)
|
||||
expires_at = datetime.fromisoformat(config.quota_expires_at)
|
||||
# 确保有时区信息,如果没有则假设为 UTC
|
||||
if expires_at.tzinfo is None:
|
||||
expires_at = expires_at.replace(tzinfo=timezone.utc)
|
||||
|
||||
@@ -3,7 +3,7 @@ Provider 模型管理 API
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, List, Optional
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, Request
|
||||
from sqlalchemy.orm import Session, joinedload
|
||||
@@ -40,15 +40,15 @@ router = APIRouter(tags=["Model Management"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
|
||||
|
||||
@router.get("/{provider_id}/models", response_model=List[ModelResponse])
|
||||
@router.get("/{provider_id}/models", response_model=list[ModelResponse])
|
||||
async def list_provider_models(
|
||||
provider_id: str,
|
||||
request: Request,
|
||||
is_active: Optional[bool] = None,
|
||||
is_active: bool | None = None,
|
||||
skip: int = 0,
|
||||
limit: int = 100,
|
||||
db: Session = Depends(get_db),
|
||||
) -> List[ModelResponse]:
|
||||
) -> list[ModelResponse]:
|
||||
"""
|
||||
获取提供商的所有模型
|
||||
|
||||
@@ -222,13 +222,13 @@ async def delete_provider_model(
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.post("/{provider_id}/models/batch", response_model=List[ModelResponse])
|
||||
@router.post("/{provider_id}/models/batch", response_model=list[ModelResponse])
|
||||
async def batch_create_provider_models(
|
||||
provider_id: str,
|
||||
models_data: List[ModelCreate],
|
||||
models_data: list[ModelCreate],
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
) -> List[ModelResponse]:
|
||||
) -> list[ModelResponse]:
|
||||
"""
|
||||
批量创建模型
|
||||
|
||||
@@ -375,7 +375,7 @@ async def import_models_from_upstream(
|
||||
@dataclass
|
||||
class AdminListProviderModelsAdapter(AdminApiAdapter):
|
||||
provider_id: str
|
||||
is_active: Optional[bool]
|
||||
is_active: bool | None
|
||||
skip: int
|
||||
limit: int
|
||||
|
||||
@@ -482,7 +482,7 @@ class AdminDeleteProviderModelAdapter(AdminApiAdapter):
|
||||
@dataclass
|
||||
class AdminBatchCreateModelsAdapter(AdminApiAdapter):
|
||||
provider_id: str
|
||||
models_data: List[ModelCreate]
|
||||
models_data: list[ModelCreate]
|
||||
|
||||
async def handle(self, context): # type: ignore[override]
|
||||
db = context.db
|
||||
@@ -525,7 +525,7 @@ class AdminGetProviderAvailableSourceModelsAdapter(AdminApiAdapter):
|
||||
)
|
||||
|
||||
# 2. 构建以 GlobalModel 为主键的字典
|
||||
global_models_dict: Dict[str, Dict[str, Any]] = {}
|
||||
global_models_dict: dict[str, dict[str, Any]] = {}
|
||||
|
||||
for model in models:
|
||||
global_model = model.global_model
|
||||
|
||||
@@ -2,7 +2,6 @@
|
||||
|
||||
import asyncio
|
||||
from datetime import datetime, timezone
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, Request
|
||||
from pydantic import BaseModel, ConfigDict, Field, ValidationError
|
||||
@@ -48,7 +47,7 @@ class MappingMatchingGlobalModel(BaseModel):
|
||||
global_model_name: str
|
||||
display_name: str
|
||||
is_active: bool
|
||||
matched_models: List[MappingMatchedModel] = Field(
|
||||
matched_models: list[MappingMatchedModel] = Field(
|
||||
default_factory=list, description="匹配到的模型列表"
|
||||
)
|
||||
|
||||
@@ -62,8 +61,8 @@ class MappingMatchingKey(BaseModel):
|
||||
key_name: str
|
||||
masked_key: str
|
||||
is_active: bool
|
||||
allowed_models: List[str] = Field(default_factory=list, description="Key 的模型白名单")
|
||||
matching_global_models: List[MappingMatchingGlobalModel] = Field(
|
||||
allowed_models: list[str] = Field(default_factory=list, description="Key 的模型白名单")
|
||||
matching_global_models: list[MappingMatchingGlobalModel] = Field(
|
||||
default_factory=list, description="匹配到的 GlobalModel 列表"
|
||||
)
|
||||
|
||||
@@ -75,7 +74,7 @@ class ProviderMappingPreviewResponse(BaseModel):
|
||||
|
||||
provider_id: str
|
||||
provider_name: str
|
||||
keys: List[MappingMatchingKey] = Field(
|
||||
keys: list[MappingMatchingKey] = Field(
|
||||
default_factory=list, description="有白名单配置且匹配到映射的 Key 列表"
|
||||
)
|
||||
total_keys: int = Field(0, description="有匹配结果的 Key 数量")
|
||||
@@ -95,7 +94,7 @@ async def list_providers(
|
||||
request: Request,
|
||||
skip: int = Query(0, ge=0),
|
||||
limit: int = Query(100, ge=1, le=500),
|
||||
is_active: Optional[bool] = None,
|
||||
is_active: bool | None = None,
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
"""
|
||||
@@ -209,7 +208,7 @@ async def delete_provider(provider_id: str, request: Request, db: Session = Depe
|
||||
|
||||
|
||||
class AdminListProvidersAdapter(AdminApiAdapter):
|
||||
def __init__(self, skip: int, limit: int, is_active: Optional[bool]):
|
||||
def __init__(self, skip: int, limit: int, is_active: bool | None):
|
||||
self.skip = skip
|
||||
self.limit = limit
|
||||
self.is_active = is_active
|
||||
@@ -473,7 +472,7 @@ async def get_provider_mapping_preview(
|
||||
pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode),
|
||||
timeout=MAPPING_PREVIEW_TIMEOUT_SECONDS,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
except TimeoutError:
|
||||
logger.warning(f"映射预览超时: provider_id={provider_id}")
|
||||
raise InvalidRequestException("映射预览超时,请简化配置或稍后重试")
|
||||
|
||||
@@ -565,7 +564,7 @@ class AdminGetProviderMappingPreviewAdapter(AdminApiAdapter):
|
||||
truncated_models = total_models_with_mappings - MAPPING_PREVIEW_MAX_MODELS
|
||||
|
||||
# 构建有映射配置的 GlobalModel 映射
|
||||
models_with_mappings: Dict[str, tuple] = {} # id -> (model_info, mappings)
|
||||
models_with_mappings: dict[str, tuple] = {} # id -> (model_info, mappings)
|
||||
for gm in global_models:
|
||||
config = gm.config or {}
|
||||
mappings = config.get("model_mappings", [])
|
||||
@@ -585,7 +584,7 @@ class AdminGetProviderMappingPreviewAdapter(AdminApiAdapter):
|
||||
truncated_models=0,
|
||||
)
|
||||
|
||||
key_infos: List[MappingMatchingKey] = []
|
||||
key_infos: list[MappingMatchingKey] = []
|
||||
total_matches = 0
|
||||
|
||||
# 创建 CryptoService 实例
|
||||
@@ -611,10 +610,10 @@ class AdminGetProviderMappingPreviewAdapter(AdminApiAdapter):
|
||||
pass
|
||||
|
||||
# 查找匹配的 GlobalModel
|
||||
matching_global_models: List[MappingMatchingGlobalModel] = []
|
||||
matching_global_models: list[MappingMatchingGlobalModel] = []
|
||||
|
||||
for gm_id, (gm, mappings) in models_with_mappings.items():
|
||||
matched_models: List[MappingMatchedModel] = []
|
||||
matched_models: list[MappingMatchedModel] = []
|
||||
|
||||
for allowed_model in allowed_models_list:
|
||||
for mapping_pattern in mappings:
|
||||
|
||||
@@ -4,7 +4,6 @@ Provider 摘要与健康监控 API
|
||||
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Dict, List
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, Request
|
||||
from sqlalchemy import case, func
|
||||
@@ -35,11 +34,11 @@ router = APIRouter(tags=["Provider Summary"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
|
||||
|
||||
@router.get("/summary", response_model=List[ProviderWithEndpointsSummary])
|
||||
@router.get("/summary", response_model=list[ProviderWithEndpointsSummary])
|
||||
async def get_providers_summary(
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
) -> List[ProviderWithEndpointsSummary]:
|
||||
) -> list[ProviderWithEndpointsSummary]:
|
||||
"""
|
||||
获取所有提供商摘要信息
|
||||
|
||||
@@ -381,8 +380,8 @@ class AdminProviderHealthMonitorAdapter(AdminApiAdapter):
|
||||
)
|
||||
attempts = attempts_query.limit(limit_rows).all()
|
||||
|
||||
buffered_attempts: Dict[str, List[RequestCandidate]] = {eid: [] for eid in endpoint_ids}
|
||||
counters: Dict[str, int] = {eid: 0 for eid in endpoint_ids}
|
||||
buffered_attempts: dict[str, list[RequestCandidate]] = {eid: [] for eid in endpoint_ids}
|
||||
counters: dict[str, int] = {eid: 0 for eid in endpoint_ids}
|
||||
|
||||
for attempt in attempts:
|
||||
if not attempt.endpoint_id or attempt.endpoint_id not in buffered_attempts:
|
||||
@@ -392,10 +391,10 @@ class AdminProviderHealthMonitorAdapter(AdminApiAdapter):
|
||||
buffered_attempts[attempt.endpoint_id].append(attempt)
|
||||
counters[attempt.endpoint_id] += 1
|
||||
|
||||
endpoint_monitors: List[EndpointHealthMonitor] = []
|
||||
endpoint_monitors: list[EndpointHealthMonitor] = []
|
||||
for endpoint in endpoints:
|
||||
attempt_list = list(reversed(buffered_attempts.get(endpoint.id, [])))
|
||||
events: List[EndpointHealthEvent] = []
|
||||
events: list[EndpointHealthEvent] = []
|
||||
for attempt in attempt_list:
|
||||
event_timestamp = attempt.finished_at or attempt.started_at or attempt.created_at
|
||||
events.append(
|
||||
|
||||
@@ -4,8 +4,6 @@ IP 安全管理接口
|
||||
提供 IP 黑白名单管理和速率限制统计
|
||||
"""
|
||||
|
||||
from typing import List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from pydantic import BaseModel, Field, ValidationError
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -14,7 +12,6 @@ from src.api.base.adapter import ApiMode
|
||||
from src.api.base.authenticated_adapter import AuthenticatedApiAdapter
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.core.exceptions import InvalidRequestException, translate_pydantic_error
|
||||
from src.core.logger import logger
|
||||
from src.database import get_db
|
||||
from src.services.rate_limit.ip_limiter import IPRateLimiter
|
||||
|
||||
@@ -30,7 +27,7 @@ class AddIPToBlacklistRequest(BaseModel):
|
||||
|
||||
ip_address: str = Field(..., description="IP 地址")
|
||||
reason: str = Field(..., min_length=1, max_length=200, description="加入黑名单的原因")
|
||||
ttl: Optional[int] = Field(None, gt=0, description="过期时间(秒),None 表示永久")
|
||||
ttl: int | None = Field(None, gt=0, description="过期时间(秒),None 表示永久")
|
||||
|
||||
|
||||
class RemoveIPFromBlacklistRequest(BaseModel):
|
||||
|
||||
@@ -1,10 +1,8 @@
|
||||
"""系统设置API端点。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import copy
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from pydantic import ValidationError
|
||||
@@ -649,9 +647,8 @@ class AdminSystemStatsAdapter(AdminApiAdapter):
|
||||
class AdminTriggerCleanupAdapter(AdminApiAdapter):
|
||||
async def handle(self, context): # type: ignore[override]
|
||||
"""手动触发清理任务"""
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from sqlalchemy import func
|
||||
|
||||
from src.services.system.maintenance_scheduler import get_maintenance_scheduler
|
||||
|
||||
|
||||
@@ -3,7 +3,6 @@
|
||||
from collections import defaultdict
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from sqlalchemy import func
|
||||
@@ -34,8 +33,8 @@ pipeline = ApiRequestPipeline()
|
||||
async def get_usage_aggregation(
|
||||
request: Request,
|
||||
group_by: str = Query(..., description="Aggregation dimension: model, user, provider, or api_format"),
|
||||
start_date: Optional[datetime] = None,
|
||||
end_date: Optional[datetime] = None,
|
||||
start_date: datetime | None = None,
|
||||
end_date: datetime | None = None,
|
||||
limit: int = Query(20, ge=1, le=100),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
@@ -75,8 +74,8 @@ async def get_usage_aggregation(
|
||||
@router.get("/stats")
|
||||
async def get_usage_stats(
|
||||
request: Request,
|
||||
start_date: Optional[datetime] = None,
|
||||
end_date: Optional[datetime] = None,
|
||||
start_date: datetime | None = None,
|
||||
end_date: datetime | None = None,
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
"""
|
||||
@@ -122,14 +121,14 @@ async def get_activity_heatmap(
|
||||
@router.get("/records")
|
||||
async def get_usage_records(
|
||||
request: Request,
|
||||
start_date: Optional[datetime] = None,
|
||||
end_date: Optional[datetime] = None,
|
||||
search: Optional[str] = None, # 通用搜索:用户名、密钥名、模型名、提供商名
|
||||
user_id: Optional[str] = None,
|
||||
username: Optional[str] = None,
|
||||
model: Optional[str] = None,
|
||||
provider: Optional[str] = None,
|
||||
status: Optional[str] = None, # stream, standard, error
|
||||
start_date: datetime | None = None,
|
||||
end_date: datetime | None = None,
|
||||
search: str | None = None, # 通用搜索:用户名、密钥名、模型名、提供商名
|
||||
user_id: str | None = None,
|
||||
username: str | None = None,
|
||||
model: str | None = None,
|
||||
provider: str | None = None,
|
||||
status: str | None = None, # stream, standard, error
|
||||
limit: int = Query(100, ge=1, le=500),
|
||||
offset: int = Query(0, ge=0),
|
||||
db: Session = Depends(get_db),
|
||||
@@ -179,7 +178,7 @@ async def get_usage_records(
|
||||
@router.get("/active")
|
||||
async def get_active_requests(
|
||||
request: Request,
|
||||
ids: Optional[str] = Query(None, description="逗号分隔的请求 ID 列表,用于查询特定请求的状态"),
|
||||
ids: str | None = Query(None, description="逗号分隔的请求 ID 列表,用于查询特定请求的状态"),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
"""
|
||||
@@ -259,7 +258,7 @@ async def get_usage_detail(
|
||||
|
||||
|
||||
class AdminUsageStatsAdapter(AdminApiAdapter):
|
||||
def __init__(self, start_date: Optional[datetime], end_date: Optional[datetime]):
|
||||
def __init__(self, start_date: datetime | None, end_date: datetime | None):
|
||||
self.start_date = start_date
|
||||
self.end_date = end_date
|
||||
|
||||
@@ -339,7 +338,7 @@ class AdminActivityHeatmapAdapter(AdminApiAdapter):
|
||||
|
||||
|
||||
class AdminUsageByModelAdapter(AdminApiAdapter):
|
||||
def __init__(self, start_date: Optional[datetime], end_date: Optional[datetime], limit: int):
|
||||
def __init__(self, start_date: datetime | None, end_date: datetime | None, limit: int):
|
||||
self.start_date = start_date
|
||||
self.end_date = end_date
|
||||
self.limit = limit
|
||||
@@ -386,7 +385,7 @@ class AdminUsageByModelAdapter(AdminApiAdapter):
|
||||
|
||||
|
||||
class AdminUsageByUserAdapter(AdminApiAdapter):
|
||||
def __init__(self, start_date: Optional[datetime], end_date: Optional[datetime], limit: int):
|
||||
def __init__(self, start_date: datetime | None, end_date: datetime | None, limit: int):
|
||||
self.start_date = start_date
|
||||
self.end_date = end_date
|
||||
self.limit = limit
|
||||
@@ -436,7 +435,7 @@ class AdminUsageByUserAdapter(AdminApiAdapter):
|
||||
|
||||
|
||||
class AdminUsageByProviderAdapter(AdminApiAdapter):
|
||||
def __init__(self, start_date: Optional[datetime], end_date: Optional[datetime], limit: int):
|
||||
def __init__(self, start_date: datetime | None, end_date: datetime | None, limit: int):
|
||||
self.start_date = start_date
|
||||
self.end_date = end_date
|
||||
self.limit = limit
|
||||
@@ -446,7 +445,7 @@ class AdminUsageByProviderAdapter(AdminApiAdapter):
|
||||
|
||||
# 从 request_candidates 表统计每个 Provider 的尝试次数和成功率
|
||||
# 这样可以正确统计 Fallback 场景(一个请求可能尝试多个 Provider)
|
||||
from sqlalchemy import case, Integer
|
||||
from sqlalchemy import case
|
||||
|
||||
attempt_query = db.query(
|
||||
RequestCandidate.provider_id,
|
||||
@@ -550,7 +549,7 @@ class AdminUsageByProviderAdapter(AdminApiAdapter):
|
||||
|
||||
|
||||
class AdminUsageByApiFormatAdapter(AdminApiAdapter):
|
||||
def __init__(self, start_date: Optional[datetime], end_date: Optional[datetime], limit: int):
|
||||
def __init__(self, start_date: datetime | None, end_date: datetime | None, limit: int):
|
||||
self.start_date = start_date
|
||||
self.end_date = end_date
|
||||
self.limit = limit
|
||||
@@ -608,14 +607,14 @@ class AdminUsageByApiFormatAdapter(AdminApiAdapter):
|
||||
class AdminUsageRecordsAdapter(AdminApiAdapter):
|
||||
def __init__(
|
||||
self,
|
||||
start_date: Optional[datetime],
|
||||
end_date: Optional[datetime],
|
||||
search: Optional[str],
|
||||
user_id: Optional[str],
|
||||
username: Optional[str],
|
||||
model: Optional[str],
|
||||
provider: Optional[str],
|
||||
status: Optional[str],
|
||||
start_date: datetime | None,
|
||||
end_date: datetime | None,
|
||||
search: str | None,
|
||||
user_id: str | None,
|
||||
username: str | None,
|
||||
model: str | None,
|
||||
provider: str | None,
|
||||
status: str | None,
|
||||
limit: int,
|
||||
offset: int,
|
||||
):
|
||||
@@ -744,7 +743,7 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
|
||||
|
||||
for req_id, candidates in request_candidates.items():
|
||||
# 提取所有不同的 candidate_index
|
||||
unique_candidates = set(c[0] for c in candidates)
|
||||
unique_candidates = {c[0] for c in candidates}
|
||||
# 如果有多个不同的 candidate_index,说明发生了 Fallback(Provider 切换)
|
||||
fallback_map[req_id] = len(unique_candidates) > 1
|
||||
|
||||
@@ -877,7 +876,7 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
|
||||
class AdminActiveRequestsAdapter(AdminApiAdapter):
|
||||
"""轻量级活跃请求状态查询适配器"""
|
||||
|
||||
def __init__(self, ids: Optional[str]):
|
||||
def __init__(self, ids: str | None):
|
||||
self.ids = ids
|
||||
|
||||
async def handle(self, context): # type: ignore[override]
|
||||
@@ -1033,8 +1032,8 @@ class AdminUsageDetailAdapter(AdminApiAdapter):
|
||||
@router.get("/cache-affinity/ttl-analysis")
|
||||
async def analyze_cache_affinity_ttl(
|
||||
request: Request,
|
||||
user_id: Optional[str] = Query(None, description="指定用户 ID"),
|
||||
api_key_id: Optional[str] = Query(None, description="指定 API Key ID"),
|
||||
user_id: str | None = Query(None, description="指定用户 ID"),
|
||||
api_key_id: str | None = Query(None, description="指定 API Key ID"),
|
||||
hours: int = Query(168, ge=1, le=720, description="分析最近多少小时的数据"),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
@@ -1057,8 +1056,8 @@ async def analyze_cache_affinity_ttl(
|
||||
@router.get("/cache-affinity/hit-analysis")
|
||||
async def analyze_cache_hit(
|
||||
request: Request,
|
||||
user_id: Optional[str] = Query(None, description="指定用户 ID"),
|
||||
api_key_id: Optional[str] = Query(None, description="指定 API Key ID"),
|
||||
user_id: str | None = Query(None, description="指定用户 ID"),
|
||||
api_key_id: str | None = Query(None, description="指定 API Key ID"),
|
||||
hours: int = Query(168, ge=1, le=720, description="分析最近多少小时的数据"),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
@@ -1080,8 +1079,8 @@ class CacheAffinityTTLAnalysisAdapter(AdminApiAdapter):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
user_id: Optional[str],
|
||||
api_key_id: Optional[str],
|
||||
user_id: str | None,
|
||||
api_key_id: str | None,
|
||||
hours: int,
|
||||
):
|
||||
self.user_id = user_id
|
||||
@@ -1114,8 +1113,8 @@ class CacheHitAnalysisAdapter(AdminApiAdapter):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
user_id: Optional[str],
|
||||
api_key_id: Optional[str],
|
||||
user_id: str | None,
|
||||
api_key_id: str | None,
|
||||
hours: int,
|
||||
):
|
||||
self.user_id = user_id
|
||||
@@ -1147,7 +1146,7 @@ async def get_interval_timeline(
|
||||
request: Request,
|
||||
hours: int = Query(24, ge=1, le=720, description="分析最近多少小时的数据"),
|
||||
limit: int = Query(10000, ge=100, le=50000, description="最大返回数据点数量"),
|
||||
user_id: Optional[str] = Query(None, description="指定用户 ID"),
|
||||
user_id: str | None = Query(None, description="指定用户 ID"),
|
||||
include_user_info: bool = Query(False, description="是否包含用户信息(用于管理员多用户视图)"),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
@@ -1177,7 +1176,7 @@ class IntervalTimelineAdapter(AdminApiAdapter):
|
||||
self,
|
||||
hours: int,
|
||||
limit: int,
|
||||
user_id: Optional[str] = None,
|
||||
user_id: str | None = None,
|
||||
include_user_info: bool = False,
|
||||
):
|
||||
self.hours = hours
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
"""用户管理 API 端点。"""
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request, status
|
||||
from pydantic import ValidationError
|
||||
@@ -48,8 +47,8 @@ async def list_users(
|
||||
request: Request,
|
||||
skip: int = Query(0, ge=0, description="跳过记录数"),
|
||||
limit: int = Query(100, ge=1, le=1000, description="返回记录数"),
|
||||
role: Optional[str] = Query(None, description="按角色筛选(user/admin)"),
|
||||
is_active: Optional[bool] = Query(None, description="按状态筛选"),
|
||||
role: str | None = Query(None, description="按角色筛选(user/admin)"),
|
||||
is_active: bool | None = Query(None, description="按状态筛选"),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
"""
|
||||
@@ -136,7 +135,7 @@ async def reset_user_quota(user_id: str, request: Request, db: Session = Depends
|
||||
async def get_user_api_keys(
|
||||
user_id: str,
|
||||
request: Request,
|
||||
is_active: Optional[bool] = Query(None, description="按状态筛选"),
|
||||
is_active: bool | None = Query(None, description="按状态筛选"),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
"""
|
||||
@@ -274,7 +273,7 @@ class AdminCreateUserAdapter(AdminApiAdapter):
|
||||
|
||||
|
||||
class AdminListUsersAdapter(AdminApiAdapter):
|
||||
def __init__(self, skip: int, limit: int, role: Optional[str], is_active: Optional[bool]):
|
||||
def __init__(self, skip: int, limit: int, role: str | None, is_active: bool | None):
|
||||
self.skip = skip
|
||||
self.limit = limit
|
||||
self.role = role
|
||||
@@ -467,7 +466,7 @@ class AdminResetUserQuotaAdapter(AdminApiAdapter):
|
||||
class AdminGetUserKeysAdapter(AdminApiAdapter):
|
||||
"""获取用户的API Keys"""
|
||||
|
||||
def __init__(self, user_id: str, is_active: Optional[bool]):
|
||||
def __init__(self, user_id: str, is_active: bool | None):
|
||||
self.user_id = user_id
|
||||
self.is_active = is_active
|
||||
|
||||
|
||||
@@ -1,9 +1,8 @@
|
||||
"""公告系统 API 端点。"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request, status
|
||||
from fastapi import APIRouter, Depends, Query, Request
|
||||
from pydantic import ValidationError
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
@@ -12,7 +11,6 @@ from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.authenticated_adapter import AuthenticatedApiAdapter
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.core.exceptions import InvalidRequestException, translate_pydantic_error
|
||||
from src.core.logger import logger
|
||||
from src.database import get_db
|
||||
from src.models.api import CreateAnnouncementRequest, UpdateAnnouncementRequest
|
||||
from src.models.database import User
|
||||
@@ -251,7 +249,7 @@ class AnnouncementOptionalAuthAdapter(ApiAdapter):
|
||||
context.extra["optional_user"] = await self._resolve_optional_user(context)
|
||||
return None
|
||||
|
||||
async def _resolve_optional_user(self, context) -> Optional[User]:
|
||||
async def _resolve_optional_user(self, context) -> User | None:
|
||||
if context.user:
|
||||
return context.user
|
||||
|
||||
@@ -285,7 +283,7 @@ class AnnouncementOptionalAuthAdapter(ApiAdapter):
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def get_optional_user(self, context) -> Optional[User]:
|
||||
def get_optional_user(self, context) -> User | None:
|
||||
return context.extra.get("optional_user")
|
||||
|
||||
|
||||
|
||||
@@ -2,10 +2,9 @@
|
||||
认证相关API端点
|
||||
"""
|
||||
|
||||
from typing import Optional, Tuple
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
|
||||
from fastapi.security import HTTPBearer
|
||||
from pydantic import ValidationError
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
@@ -42,7 +41,7 @@ from src.services.email import EmailSenderService, EmailVerificationService
|
||||
from src.utils.request_utils import get_client_ip, get_user_agent
|
||||
|
||||
|
||||
def validate_email_suffix(db: Session, email: str) -> Tuple[bool, Optional[str]]:
|
||||
def validate_email_suffix(db: Session, email: str) -> tuple[bool, str | None]:
|
||||
"""
|
||||
验证邮箱后缀是否允许注册
|
||||
|
||||
|
||||
@@ -1,8 +1,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from enum import Enum
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import Any
|
||||
|
||||
from fastapi import Request, Response
|
||||
|
||||
@@ -23,7 +21,7 @@ class ApiAdapter(ABC):
|
||||
|
||||
name: str = "base"
|
||||
mode: ApiMode = ApiMode.STANDARD
|
||||
api_format: Optional[str] = None # 对应 Provider API 格式提示
|
||||
api_format: str | None = None # 对应 Provider API 格式提示
|
||||
audit_log_enabled: bool = True
|
||||
audit_success_event = None
|
||||
audit_failure_event = None
|
||||
@@ -36,7 +34,7 @@ class ApiAdapter(ABC):
|
||||
"""可选的授权钩子,默认允许通过。"""
|
||||
return None
|
||||
|
||||
def extract_api_key(self, request: Request) -> Optional[str]:
|
||||
def extract_api_key(self, request: Request) -> str | None:
|
||||
"""
|
||||
从请求中提取客户端 API 密钥。
|
||||
|
||||
@@ -55,17 +53,17 @@ class ApiAdapter(ABC):
|
||||
context: ApiRequestContext,
|
||||
*,
|
||||
success: bool,
|
||||
status_code: Optional[int],
|
||||
error: Optional[str] = None,
|
||||
) -> Dict[str, Any]:
|
||||
status_code: int | None,
|
||||
error: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""允许适配器在审计日志中追加自定义字段。"""
|
||||
return {}
|
||||
|
||||
def detect_capability_requirements(
|
||||
self,
|
||||
headers: Dict[str, str],
|
||||
request_body: Optional[Dict[str, Any]] = None,
|
||||
) -> Dict[str, bool]:
|
||||
headers: dict[str, str],
|
||||
request_body: dict[str, Any] | None = None,
|
||||
) -> dict[str, bool]:
|
||||
"""
|
||||
检测请求中隐含的能力需求(子类可覆盖)
|
||||
|
||||
|
||||
@@ -1,5 +1,3 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
from src.models.database import UserRole
|
||||
|
||||
@@ -1,10 +1,9 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import time
|
||||
import uuid
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import Any
|
||||
|
||||
from fastapi import HTTPException, Request
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -21,34 +20,34 @@ class ApiRequestContext:
|
||||
|
||||
request: Request
|
||||
db: Session
|
||||
user: Optional[User]
|
||||
api_key: Optional[ApiKey]
|
||||
user: User | None
|
||||
api_key: ApiKey | None
|
||||
request_id: str
|
||||
start_time: float
|
||||
client_ip: str
|
||||
user_agent: str
|
||||
original_headers: Dict[str, str]
|
||||
query_params: Dict[str, str]
|
||||
original_headers: dict[str, str]
|
||||
query_params: dict[str, str]
|
||||
raw_body: bytes | None = None
|
||||
json_body: Optional[Dict[str, Any]] = None
|
||||
quota_remaining: Optional[float] = None
|
||||
json_body: dict[str, Any] | None = None
|
||||
quota_remaining: float | None = None
|
||||
mode: str = "standard" # standard / proxy
|
||||
api_format_hint: Optional[str] = None
|
||||
api_format_hint: str | None = None
|
||||
|
||||
# URL 路径参数(如 Gemini API 的 /v1beta/models/{model}:generateContent)
|
||||
path_params: Dict[str, Any] = field(default_factory=dict)
|
||||
path_params: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
# Management Token(用于管理 API 认证)
|
||||
management_token: Optional[ManagementToken] = None
|
||||
management_token: ManagementToken | None = None
|
||||
|
||||
# 供适配器扩展的状态存储
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
audit_metadata: Dict[str, Any] = field(default_factory=dict)
|
||||
extra: dict[str, Any] = field(default_factory=dict)
|
||||
audit_metadata: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
# 高频轮询端点日志抑制标志
|
||||
quiet_logging: bool = False
|
||||
|
||||
def ensure_json_body(self) -> Dict[str, Any]:
|
||||
def ensure_json_body(self) -> dict[str, Any]:
|
||||
"""确保请求体已解析为JSON并返回。"""
|
||||
if self.json_body is not None:
|
||||
return self.json_body
|
||||
@@ -70,7 +69,7 @@ class ApiRequestContext:
|
||||
if value is not None:
|
||||
self.audit_metadata[key] = value
|
||||
|
||||
def extend_audit_metadata(self, data: Dict[str, Any]) -> None:
|
||||
def extend_audit_metadata(self, data: dict[str, Any]) -> None:
|
||||
"""批量附加审计字段。"""
|
||||
for key, value in data.items():
|
||||
if value is not None:
|
||||
@@ -81,13 +80,13 @@ class ApiRequestContext:
|
||||
cls,
|
||||
request: Request,
|
||||
db: Session,
|
||||
user: Optional[User],
|
||||
api_key: Optional[ApiKey],
|
||||
raw_body: Optional[bytes] = None,
|
||||
user: User | None,
|
||||
api_key: ApiKey | None,
|
||||
raw_body: bytes | None = None,
|
||||
mode: str = "standard",
|
||||
api_format_hint: Optional[str] = None,
|
||||
path_params: Optional[Dict[str, Any]] = None,
|
||||
) -> "ApiRequestContext":
|
||||
api_format_hint: str | None = None,
|
||||
path_params: dict[str, Any] | None = None,
|
||||
) -> ApiRequestContext:
|
||||
"""创建上下文实例并提前读取必要的元数据。"""
|
||||
request_id = getattr(request.state, "request_id", None) or str(uuid.uuid4())[:8]
|
||||
setattr(request.state, "request_id", request_id)
|
||||
|
||||
@@ -10,8 +10,9 @@
|
||||
4. Key 的 allowed_models 允许该模型(null = 允许所有)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
from dataclasses import asdict, dataclass
|
||||
from typing import Any, Optional
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
@@ -27,7 +28,7 @@ _CACHE_KEY_PREFIX = "models:list"
|
||||
_CACHE_TTL = CacheTTL.MODEL # 300 秒
|
||||
|
||||
|
||||
def _get_cache_key(api_formats: list[str], client_format: Optional[str] = None) -> str:
|
||||
def _get_cache_key(api_formats: list[str], client_format: str | None = None) -> str:
|
||||
"""生成缓存 key"""
|
||||
formats_str = ",".join(sorted(api_formats))
|
||||
format_key = (client_format or "any").lower()
|
||||
@@ -35,8 +36,8 @@ def _get_cache_key(api_formats: list[str], client_format: Optional[str] = None)
|
||||
|
||||
|
||||
async def _get_cached_models(
|
||||
api_formats: list[str], client_format: Optional[str] = None
|
||||
) -> Optional[list["ModelInfo"]]:
|
||||
api_formats: list[str], client_format: str | None = None
|
||||
) -> list[ModelInfo] | None:
|
||||
"""从缓存获取模型列表"""
|
||||
cache_key = _get_cache_key(api_formats, client_format)
|
||||
try:
|
||||
@@ -51,8 +52,8 @@ async def _get_cached_models(
|
||||
|
||||
async def _set_cached_models(
|
||||
api_formats: list[str],
|
||||
models: list["ModelInfo"],
|
||||
client_format: Optional[str] = None,
|
||||
models: list[ModelInfo],
|
||||
client_format: str | None = None,
|
||||
) -> None:
|
||||
"""将模型列表写入缓存"""
|
||||
cache_key = _get_cache_key(api_formats, client_format)
|
||||
@@ -87,8 +88,8 @@ class ModelInfo:
|
||||
|
||||
id: str # 模型 ID (GlobalModel.name 或 provider_model_name)
|
||||
display_name: str
|
||||
description: Optional[str]
|
||||
created_at: Optional[str] # ISO 格式
|
||||
description: str | None
|
||||
created_at: str | None # ISO 格式
|
||||
created_timestamp: int # Unix 时间戳
|
||||
provider_name: str
|
||||
provider_id: str = "" # Provider ID,用于权限过滤
|
||||
@@ -100,27 +101,27 @@ class ModelInfo:
|
||||
image_generation: bool = False
|
||||
structured_output: bool = False
|
||||
# 规格参数
|
||||
context_limit: Optional[int] = None
|
||||
output_limit: Optional[int] = None
|
||||
context_limit: int | None = None
|
||||
output_limit: int | None = None
|
||||
# 元信息
|
||||
family: Optional[str] = None
|
||||
knowledge_cutoff: Optional[str] = None
|
||||
input_modalities: Optional[list[str]] = None
|
||||
output_modalities: Optional[list[str]] = None
|
||||
family: str | None = None
|
||||
knowledge_cutoff: str | None = None
|
||||
input_modalities: list[str] | None = None
|
||||
output_modalities: list[str] | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class AccessRestrictions:
|
||||
"""API Key 或 User 的访问限制"""
|
||||
|
||||
allowed_providers: Optional[list[str]] = None # 允许的 Provider ID 列表
|
||||
allowed_models: Optional[list[str]] = None # 允许的模型名称列表
|
||||
allowed_api_formats: Optional[list[str]] = None # 允许的 API 格式列表
|
||||
allowed_providers: list[str] | None = None # 允许的 Provider ID 列表
|
||||
allowed_models: list[str] | None = None # 允许的模型名称列表
|
||||
allowed_api_formats: list[str] | None = None # 允许的 API 格式列表
|
||||
|
||||
@classmethod
|
||||
def from_api_key_and_user(
|
||||
cls, api_key: Optional[ApiKey], user: Optional[User]
|
||||
) -> "AccessRestrictions":
|
||||
cls, api_key: ApiKey | None, user: User | None
|
||||
) -> AccessRestrictions:
|
||||
"""
|
||||
从 API Key 和 User 合并访问限制
|
||||
|
||||
@@ -130,9 +131,9 @@ class AccessRestrictions:
|
||||
- 如果 API Key 无限制但 User 有限制,使用 User 的限制
|
||||
- 两者都无限制则返回空限制
|
||||
"""
|
||||
allowed_providers: Optional[list[str]] = None
|
||||
allowed_models: Optional[list[str]] = None
|
||||
allowed_api_formats: Optional[list[str]] = None
|
||||
allowed_providers: list[str] | None = None
|
||||
allowed_models: list[str] | None = None
|
||||
allowed_api_formats: list[str] | None = None
|
||||
|
||||
# 优先使用 API Key 的限制
|
||||
if api_key:
|
||||
@@ -197,8 +198,8 @@ class AccessRestrictions:
|
||||
|
||||
|
||||
def _normalize_api_formats(
|
||||
api_formats: Optional[list[str]],
|
||||
provider_to_formats: Optional[dict[str, set[str]]] = None,
|
||||
api_formats: list[str] | None,
|
||||
provider_to_formats: dict[str, set[str]] | None = None,
|
||||
) -> list[str]:
|
||||
"""规范化 API 格式列表(大写),必要时从 provider_to_formats 兜底"""
|
||||
if api_formats:
|
||||
@@ -212,7 +213,7 @@ def _normalize_api_formats(
|
||||
|
||||
|
||||
def _get_provider_model_names_for_formats(
|
||||
model: Model, usable_formats: Optional[set[str]] = None
|
||||
model: Model, usable_formats: set[str] | None = None
|
||||
) -> set[str]:
|
||||
"""
|
||||
获取模型在指定格式下支持的 Provider 模型名称集合
|
||||
@@ -305,7 +306,7 @@ def get_compatible_provider_formats(
|
||||
def get_available_provider_ids(
|
||||
db: Session,
|
||||
api_formats: list[str],
|
||||
provider_to_formats: Optional[dict[str, set[str]]] = None,
|
||||
provider_to_formats: dict[str, set[str]] | None = None,
|
||||
) -> set[str]:
|
||||
"""
|
||||
返回有可用端点的 Provider IDs
|
||||
@@ -334,7 +335,7 @@ def get_available_provider_ids(
|
||||
def _get_available_model_ids_for_format(
|
||||
db: Session,
|
||||
api_formats: list[str],
|
||||
provider_to_formats: Optional[dict[str, set[str]]] = None,
|
||||
provider_to_formats: dict[str, set[str]] | None = None,
|
||||
) -> set[str]:
|
||||
"""
|
||||
获取指定格式下真正可用的模型 ID 集合
|
||||
@@ -410,7 +411,7 @@ def _get_available_model_ids_for_format(
|
||||
return available_model_ids
|
||||
|
||||
|
||||
def _extract_model_info(model: Any) -> Optional[ModelInfo]:
|
||||
def _extract_model_info(model: Any) -> ModelInfo | None:
|
||||
"""
|
||||
从 Model 对象提取 ModelInfo
|
||||
|
||||
@@ -424,7 +425,7 @@ def _extract_model_info(model: Any) -> Optional[ModelInfo]:
|
||||
|
||||
model_id: str = global_model.name
|
||||
display_name: str = global_model.display_name
|
||||
created_at: Optional[str] = (
|
||||
created_at: str | None = (
|
||||
model.created_at.strftime("%Y-%m-%dT%H:%M:%SZ") if model.created_at else None
|
||||
)
|
||||
created_timestamp: int = int(model.created_at.timestamp()) if model.created_at else 0
|
||||
@@ -433,7 +434,7 @@ def _extract_model_info(model: Any) -> Optional[ModelInfo]:
|
||||
|
||||
# 从 GlobalModel.config 提取配置信息
|
||||
config: dict = global_model.config or {}
|
||||
description: Optional[str] = config.get("description")
|
||||
description: str | None = config.get("description")
|
||||
|
||||
return ModelInfo(
|
||||
id=model_id,
|
||||
@@ -464,10 +465,10 @@ def _extract_model_info(model: Any) -> Optional[ModelInfo]:
|
||||
async def list_available_models(
|
||||
db: Session,
|
||||
available_provider_ids: set[str],
|
||||
api_formats: Optional[list[str]] = None,
|
||||
restrictions: Optional[AccessRestrictions] = None,
|
||||
provider_to_formats: Optional[dict[str, set[str]]] = None,
|
||||
client_format: Optional[str] = None,
|
||||
api_formats: list[str] | None = None,
|
||||
restrictions: AccessRestrictions | None = None,
|
||||
provider_to_formats: dict[str, set[str]] | None = None,
|
||||
client_format: str | None = None,
|
||||
) -> list[ModelInfo]:
|
||||
"""
|
||||
获取可用模型列表(已去重,带缓存)
|
||||
@@ -503,7 +504,7 @@ async def list_available_models(
|
||||
return cached
|
||||
|
||||
# 如果提供了 api_formats,获取真正可用的模型 ID
|
||||
available_model_ids: Optional[set[str]] = None
|
||||
available_model_ids: set[str] | None = None
|
||||
if normalized_formats:
|
||||
available_model_ids = _get_available_model_ids_for_format(
|
||||
db, normalized_formats, provider_to_formats
|
||||
@@ -551,10 +552,10 @@ def find_model_by_id(
|
||||
db: Session,
|
||||
model_id: str,
|
||||
available_provider_ids: set[str],
|
||||
api_formats: Optional[list[str]] = None,
|
||||
restrictions: Optional[AccessRestrictions] = None,
|
||||
provider_to_formats: Optional[dict[str, set[str]]] = None,
|
||||
) -> Optional[ModelInfo]:
|
||||
api_formats: list[str] | None = None,
|
||||
restrictions: AccessRestrictions | None = None,
|
||||
provider_to_formats: dict[str, set[str]] | None = None,
|
||||
) -> ModelInfo | None:
|
||||
"""
|
||||
按 ID 查找模型(仅支持 GlobalModel.name)
|
||||
|
||||
@@ -575,7 +576,7 @@ def find_model_by_id(
|
||||
normalized_formats = _normalize_api_formats(api_formats, provider_to_formats)
|
||||
|
||||
# 如果提供了 api_formats,获取真正可用的模型 ID
|
||||
available_model_ids: Optional[set[str]] = None
|
||||
available_model_ids: set[str] | None = None
|
||||
if normalized_formats:
|
||||
available_model_ids = _get_available_model_ids_for_format(
|
||||
db, normalized_formats, provider_to_formats
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import asdict, dataclass
|
||||
from typing import Any, List, Sequence, Tuple, TypeVar
|
||||
from typing import Any, TypeVar
|
||||
from collections.abc import Sequence
|
||||
|
||||
from sqlalchemy.orm import Query
|
||||
|
||||
@@ -19,7 +18,7 @@ class PaginationMeta:
|
||||
return asdict(self)
|
||||
|
||||
|
||||
def paginate_query(query: Query, limit: int, offset: int) -> Tuple[int, List[T]]:
|
||||
def paginate_query(query: Query, limit: int, offset: int) -> tuple[int, list[T]]:
|
||||
"""
|
||||
对 SQLAlchemy 查询应用 limit/offset,并返回总数与结果列表。
|
||||
"""
|
||||
@@ -30,7 +29,7 @@ def paginate_query(query: Query, limit: int, offset: int) -> Tuple[int, List[T]]
|
||||
|
||||
def paginate_sequence(
|
||||
items: Sequence[T], limit: int, offset: int
|
||||
) -> Tuple[List[T], PaginationMeta]:
|
||||
) -> tuple[list[T], PaginationMeta]:
|
||||
"""
|
||||
对内存序列应用分页,返回切片和元数据。
|
||||
"""
|
||||
@@ -40,7 +39,7 @@ def paginate_sequence(
|
||||
return sliced, meta
|
||||
|
||||
|
||||
def build_pagination_payload(items: List[dict], meta: PaginationMeta, **extra: Any) -> dict:
|
||||
def build_pagination_payload(items: list[dict], meta: PaginationMeta, **extra: Any) -> dict:
|
||||
"""
|
||||
构建标准分页响应 payload。
|
||||
"""
|
||||
|
||||
@@ -1,8 +1,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from enum import Enum
|
||||
from typing import TYPE_CHECKING, Any, Optional, Tuple
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from fastapi import HTTPException, Request
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -52,8 +50,8 @@ class ApiRequestPipeline:
|
||||
db: Session,
|
||||
*,
|
||||
mode: ApiMode = ApiMode.STANDARD,
|
||||
api_format_hint: Optional[str] = None,
|
||||
path_params: Optional[dict[str, Any]] = None,
|
||||
api_format_hint: str | None = None,
|
||||
path_params: dict[str, Any] | None = None,
|
||||
):
|
||||
# 高频轮询端点抑制 debug 日志
|
||||
is_quiet = http_request.url.path in QUIET_POLLING_PATHS
|
||||
@@ -95,7 +93,7 @@ class ApiRequestPipeline:
|
||||
)
|
||||
if not is_quiet:
|
||||
logger.debug("[Pipeline] Raw body读取完成 | size=%d bytes", len(raw_body) if raw_body is not None else 0)
|
||||
except asyncio.TimeoutError:
|
||||
except TimeoutError:
|
||||
timeout_sec = int(config.request_body_timeout)
|
||||
logger.error(f"读取请求体超时({timeout_sec}s),可能客户端未发送完整请求体")
|
||||
raise HTTPException(
|
||||
@@ -166,7 +164,7 @@ class ApiRequestPipeline:
|
||||
|
||||
def _authenticate_client(
|
||||
self, request: Request, db: Session, adapter: ApiAdapter, *, quiet: bool = False
|
||||
) -> Tuple[User, ApiKey]:
|
||||
) -> tuple[User, ApiKey]:
|
||||
if not quiet:
|
||||
logger.debug("[Pipeline._authenticate_client] 开始")
|
||||
# 使用 adapter 的 extract_api_key 方法,支持不同 API 格式的认证头
|
||||
@@ -215,7 +213,7 @@ class ApiRequestPipeline:
|
||||
|
||||
async def _authenticate_admin(
|
||||
self, request: Request, db: Session
|
||||
) -> Tuple[User, Optional["ManagementToken"]]:
|
||||
) -> tuple[User, ManagementToken | None]:
|
||||
"""管理员认证,支持 JWT 和 Management Token 两种方式"""
|
||||
from src.models.database import ManagementToken
|
||||
from src.utils.request_utils import get_client_ip
|
||||
@@ -278,7 +276,7 @@ class ApiRequestPipeline:
|
||||
|
||||
async def _authenticate_user(
|
||||
self, request: Request, db: Session
|
||||
) -> Tuple[User, Optional["ManagementToken"]]:
|
||||
) -> tuple[User, ManagementToken | None]:
|
||||
"""用户认证,支持 JWT 和 Management Token 两种方式"""
|
||||
from src.models.database import ManagementToken
|
||||
from src.utils.request_utils import get_client_ip
|
||||
@@ -329,7 +327,7 @@ class ApiRequestPipeline:
|
||||
|
||||
async def _authenticate_management(
|
||||
self, request: Request, db: Session
|
||||
) -> Tuple[User, "ManagementToken"]:
|
||||
) -> tuple[User, ManagementToken]:
|
||||
"""Management Token 认证"""
|
||||
from src.models.database import ManagementToken
|
||||
from src.utils.request_utils import get_client_ip
|
||||
@@ -362,7 +360,7 @@ class ApiRequestPipeline:
|
||||
|
||||
return user, management_token
|
||||
|
||||
def _calculate_quota_remaining(self, user: Optional[User]) -> Optional[float]:
|
||||
def _calculate_quota_remaining(self, user: User | None) -> float | None:
|
||||
if not user:
|
||||
return None
|
||||
if user.quota_usd is None or user.quota_usd < 0:
|
||||
@@ -375,8 +373,8 @@ class ApiRequestPipeline:
|
||||
adapter: ApiAdapter,
|
||||
*,
|
||||
success: bool,
|
||||
status_code: Optional[int] = None,
|
||||
error: Optional[str] = None,
|
||||
status_code: int | None = None,
|
||||
error: str | None = None,
|
||||
) -> None:
|
||||
"""记录审计事件
|
||||
|
||||
@@ -432,8 +430,8 @@ class ApiRequestPipeline:
|
||||
adapter: ApiAdapter,
|
||||
*,
|
||||
success: bool,
|
||||
status_code: Optional[int],
|
||||
error: Optional[str],
|
||||
status_code: int | None,
|
||||
error: str | None,
|
||||
) -> dict:
|
||||
duration_ms = max((time.time() - context.start_time) * 1000, 0.0)
|
||||
request = context.request
|
||||
|
||||
@@ -2,7 +2,6 @@
|
||||
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import List
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from sqlalchemy import and_, func
|
||||
@@ -952,7 +951,7 @@ class DashboardDailyStatsAdapter(DashboardAdapter):
|
||||
# 构建完整日期序列(使用业务时区日期)
|
||||
current_date = start_date_local.date()
|
||||
end_date_date = end_date_local.date()
|
||||
formatted: List[dict] = []
|
||||
formatted: list[dict] = []
|
||||
while current_date <= end_date_date:
|
||||
date_str = current_date.isoformat()
|
||||
stat = stats_map.get(date_str)
|
||||
|
||||
@@ -25,23 +25,20 @@
|
||||
) -> JSONResponse: ...
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import time
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
Awaitable,
|
||||
Callable,
|
||||
Coroutine,
|
||||
Dict,
|
||||
Optional,
|
||||
Protocol,
|
||||
TypeVar,
|
||||
runtime_checkable,
|
||||
)
|
||||
|
||||
from collections.abc import Callable
|
||||
from collections.abc import Awaitable, Coroutine
|
||||
|
||||
from fastapi import Request
|
||||
from fastapi.responses import JSONResponse, StreamingResponse
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -57,6 +54,9 @@ from src.services.usage.service import UsageService
|
||||
if TYPE_CHECKING:
|
||||
from src.api.handlers.base.stream_context import StreamContext
|
||||
|
||||
# Adapter 检测器类型:接受 headers 和可选的 request_body,返回能力需求字典
|
||||
type AdapterDetectorType = Callable[[dict[str, str], dict[str, Any] | None], dict[str, bool]]
|
||||
|
||||
|
||||
class MessageTelemetry:
|
||||
"""
|
||||
@@ -105,29 +105,29 @@ class MessageTelemetry:
|
||||
output_tokens: int,
|
||||
response_time_ms: int,
|
||||
status_code: int,
|
||||
request_body: Dict[str, Any],
|
||||
request_headers: Dict[str, Any],
|
||||
request_body: dict[str, Any],
|
||||
request_headers: dict[str, Any],
|
||||
response_body: Any,
|
||||
response_headers: Dict[str, Any],
|
||||
client_response_headers: Optional[Dict[str, Any]] = None,
|
||||
response_headers: dict[str, Any],
|
||||
client_response_headers: dict[str, Any] | None = None,
|
||||
cache_creation_tokens: int = 0,
|
||||
cache_read_tokens: int = 0,
|
||||
is_stream: bool = False,
|
||||
provider_request_headers: Optional[Dict[str, Any]] = None,
|
||||
provider_request_headers: dict[str, Any] | None = None,
|
||||
# 时间指标
|
||||
first_byte_time_ms: Optional[int] = None, # 首字时间/TTFB
|
||||
first_byte_time_ms: int | None = None, # 首字时间/TTFB
|
||||
# Provider 侧追踪信息(用于记录真实成本)
|
||||
provider_id: Optional[str] = None,
|
||||
provider_endpoint_id: Optional[str] = None,
|
||||
provider_api_key_id: Optional[str] = None,
|
||||
api_format: Optional[str] = None,
|
||||
provider_id: str | None = None,
|
||||
provider_endpoint_id: str | None = None,
|
||||
provider_api_key_id: str | None = None,
|
||||
api_format: str | None = None,
|
||||
# 格式转换追踪
|
||||
endpoint_api_format: Optional[str] = None, # 端点原生 API 格式
|
||||
endpoint_api_format: str | None = None, # 端点原生 API 格式
|
||||
has_format_conversion: bool = False, # 是否发生了格式转换
|
||||
# 模型映射信息
|
||||
target_model: Optional[str] = None,
|
||||
target_model: str | None = None,
|
||||
# Provider 响应元数据(如 Gemini 的 modelVersion)
|
||||
response_metadata: Optional[Dict[str, Any]] = None,
|
||||
response_metadata: dict[str, Any] | None = None,
|
||||
) -> float:
|
||||
total_cost = await self.calculate_cost(
|
||||
provider,
|
||||
@@ -199,24 +199,24 @@ class MessageTelemetry:
|
||||
response_time_ms: int,
|
||||
status_code: int,
|
||||
error_message: str,
|
||||
request_body: Dict[str, Any],
|
||||
request_headers: Dict[str, Any],
|
||||
request_body: dict[str, Any],
|
||||
request_headers: dict[str, Any],
|
||||
is_stream: bool,
|
||||
api_format: Optional[str] = None,
|
||||
provider_request_headers: Optional[Dict[str, Any]] = None,
|
||||
api_format: str | None = None,
|
||||
provider_request_headers: dict[str, Any] | None = None,
|
||||
# 预估 token 信息(来自 message_start 事件,用于中断请求的成本估算)
|
||||
input_tokens: int = 0,
|
||||
output_tokens: int = 0,
|
||||
cache_creation_tokens: int = 0,
|
||||
cache_read_tokens: int = 0,
|
||||
response_body: Optional[Dict[str, Any]] = None,
|
||||
response_headers: Optional[Dict[str, Any]] = None,
|
||||
client_response_headers: Optional[Dict[str, Any]] = None,
|
||||
response_body: dict[str, Any] | None = None,
|
||||
response_headers: dict[str, Any] | None = None,
|
||||
client_response_headers: dict[str, Any] | None = None,
|
||||
# 格式转换追踪
|
||||
endpoint_api_format: Optional[str] = None,
|
||||
endpoint_api_format: str | None = None,
|
||||
has_format_conversion: bool = False,
|
||||
# 模型映射信息
|
||||
target_model: Optional[str] = None,
|
||||
target_model: str | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
记录失败请求
|
||||
@@ -273,24 +273,24 @@ class MessageTelemetry:
|
||||
provider: str,
|
||||
model: str,
|
||||
response_time_ms: int,
|
||||
first_byte_time_ms: Optional[int],
|
||||
first_byte_time_ms: int | None,
|
||||
status_code: int,
|
||||
request_body: Dict[str, Any],
|
||||
request_headers: Dict[str, Any],
|
||||
request_body: dict[str, Any],
|
||||
request_headers: dict[str, Any],
|
||||
is_stream: bool,
|
||||
api_format: Optional[str] = None,
|
||||
provider_request_headers: Optional[Dict[str, Any]] = None,
|
||||
api_format: str | None = None,
|
||||
provider_request_headers: dict[str, Any] | None = None,
|
||||
input_tokens: int = 0,
|
||||
output_tokens: int = 0,
|
||||
cache_creation_tokens: int = 0,
|
||||
cache_read_tokens: int = 0,
|
||||
response_body: Optional[Dict[str, Any]] = None,
|
||||
response_headers: Optional[Dict[str, Any]] = None,
|
||||
client_response_headers: Optional[Dict[str, Any]] = None,
|
||||
response_body: dict[str, Any] | None = None,
|
||||
response_headers: dict[str, Any] | None = None,
|
||||
client_response_headers: dict[str, Any] | None = None,
|
||||
# 格式转换追踪
|
||||
endpoint_api_format: Optional[str] = None,
|
||||
endpoint_api_format: str | None = None,
|
||||
has_format_conversion: bool = False,
|
||||
target_model: Optional[str] = None,
|
||||
target_model: str | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
记录客户端取消的请求
|
||||
@@ -341,9 +341,9 @@ class MessageHandlerProtocol(Protocol):
|
||||
self,
|
||||
request: Any,
|
||||
http_request: Request,
|
||||
original_headers: Dict[str, str],
|
||||
original_request_body: Dict[str, Any],
|
||||
query_params: Optional[Dict[str, str]] = None,
|
||||
original_headers: dict[str, str],
|
||||
original_request_body: dict[str, Any],
|
||||
query_params: dict[str, str] | None = None,
|
||||
) -> StreamingResponse:
|
||||
"""处理流式请求"""
|
||||
...
|
||||
@@ -352,9 +352,9 @@ class MessageHandlerProtocol(Protocol):
|
||||
self,
|
||||
request: Any,
|
||||
http_request: Request,
|
||||
original_headers: Dict[str, str],
|
||||
original_request_body: Dict[str, Any],
|
||||
query_params: Optional[Dict[str, str]] = None,
|
||||
original_headers: dict[str, str],
|
||||
original_request_body: dict[str, Any],
|
||||
query_params: dict[str, str] | None = None,
|
||||
) -> JSONResponse:
|
||||
"""处理非流式请求"""
|
||||
...
|
||||
@@ -371,9 +371,6 @@ class BaseMessageHandler:
|
||||
推荐使用 MessageHandlerProtocol 中定义的签名。
|
||||
"""
|
||||
|
||||
# Adapter 检测器类型
|
||||
AdapterDetectorType = Callable[[Dict[str, str], Optional[Dict[str, Any]]], Dict[str, bool]]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
@@ -384,8 +381,8 @@ class BaseMessageHandler:
|
||||
client_ip: str,
|
||||
user_agent: str,
|
||||
start_time: float,
|
||||
allowed_api_formats: Optional[list[str]] = None,
|
||||
adapter_detector: Optional[AdapterDetectorType] = None,
|
||||
allowed_api_formats: list[str] | None = None,
|
||||
adapter_detector: AdapterDetectorType | None = None,
|
||||
) -> None:
|
||||
self.db = db
|
||||
self.user = user
|
||||
@@ -408,9 +405,9 @@ class BaseMessageHandler:
|
||||
def _resolve_capability_requirements(
|
||||
self,
|
||||
model_name: str,
|
||||
request_headers: Optional[Dict[str, str]] = None,
|
||||
request_body: Optional[Dict[str, Any]] = None,
|
||||
) -> Dict[str, bool]:
|
||||
request_headers: dict[str, str] | None = None,
|
||||
request_body: dict[str, Any] | None = None,
|
||||
) -> dict[str, bool]:
|
||||
"""
|
||||
解析请求的能力需求
|
||||
|
||||
@@ -439,7 +436,7 @@ class BaseMessageHandler:
|
||||
adapter_detector=self.adapter_detector,
|
||||
)
|
||||
|
||||
def get_api_format(self, provider_type: Optional[str] = None) -> APIFormat:
|
||||
def get_api_format(self, provider_type: str | None = None) -> APIFormat:
|
||||
"""根据 provider_type 解析 API 格式,未知类型默认 OPENAI"""
|
||||
if provider_type:
|
||||
result = resolve_api_format(provider_type, default=APIFormat.OPENAI)
|
||||
@@ -448,17 +445,17 @@ class BaseMessageHandler:
|
||||
|
||||
def build_provider_payload(
|
||||
self,
|
||||
original_body: Dict[str, Any],
|
||||
original_body: dict[str, Any],
|
||||
*,
|
||||
mapped_model: Optional[str] = None,
|
||||
) -> Dict[str, Any]:
|
||||
mapped_model: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""构建发送给 Provider 的请求体,替换 model 名称"""
|
||||
payload = dict(original_body)
|
||||
if mapped_model:
|
||||
payload["model"] = mapped_model
|
||||
return payload
|
||||
|
||||
def _update_usage_to_streaming(self, request_id: Optional[str] = None) -> None:
|
||||
def _update_usage_to_streaming(self, request_id: str | None = None) -> None:
|
||||
"""更新 Usage 状态为 streaming(流式传输开始时调用)
|
||||
|
||||
使用 asyncio 后台任务执行数据库更新,避免阻塞流式传输
|
||||
@@ -492,7 +489,7 @@ class BaseMessageHandler:
|
||||
# 创建后台任务,不阻塞当前流
|
||||
asyncio.create_task(_do_update())
|
||||
|
||||
def _update_usage_to_streaming_with_ctx(self, ctx: "StreamContext") -> None:
|
||||
def _update_usage_to_streaming_with_ctx(self, ctx: StreamContext) -> None:
|
||||
"""更新 Usage 状态为 streaming,同时更新 provider 相关信息
|
||||
|
||||
使用 asyncio 后台任务执行数据库更新,避免阻塞流式传输
|
||||
|
||||
@@ -19,7 +19,7 @@ Chat Adapter 通用基类
|
||||
import time
|
||||
import traceback
|
||||
from abc import abstractmethod
|
||||
from typing import Any, Dict, Optional, Tuple, Type
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from fastapi import HTTPException, Request
|
||||
@@ -65,7 +65,7 @@ class ChatAdapterBase(ApiAdapter):
|
||||
|
||||
# 子类必须覆盖
|
||||
FORMAT_ID: str = "UNKNOWN"
|
||||
HANDLER_CLASS: Type[ChatHandlerBase]
|
||||
HANDLER_CLASS: type[ChatHandlerBase]
|
||||
|
||||
# 适配器配置
|
||||
name: str = "chat.base"
|
||||
@@ -90,7 +90,7 @@ class ChatAdapterBase(ApiAdapter):
|
||||
return base_url
|
||||
|
||||
@classmethod
|
||||
def build_base_headers(cls, api_key: str) -> Dict[str, str]:
|
||||
def build_base_headers(cls, api_key: str) -> dict[str, str]:
|
||||
"""构建基础请求头,使用统一的 headers.py 实现"""
|
||||
return build_adapter_base_headers(cls._get_api_format(), api_key)
|
||||
|
||||
@@ -101,13 +101,13 @@ class ChatAdapterBase(ApiAdapter):
|
||||
|
||||
@classmethod
|
||||
def build_headers_with_extra(
|
||||
cls, api_key: str, extra_headers: Optional[Dict[str, str]] = None
|
||||
) -> Dict[str, str]:
|
||||
cls, api_key: str, extra_headers: dict[str, str] | None = None
|
||||
) -> dict[str, str]:
|
||||
"""构建完整请求头(包含 extra_headers),使用统一的 headers.py 实现"""
|
||||
return build_adapter_headers(cls._get_api_format(), api_key, extra_headers)
|
||||
|
||||
@classmethod
|
||||
def build_request_body(cls, request_data: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
|
||||
def build_request_body(cls, request_data: dict[str, Any] | None = None) -> dict[str, Any]:
|
||||
"""构建测试请求体,使用转换器注册表自动处理格式转换
|
||||
|
||||
Args:
|
||||
@@ -120,11 +120,11 @@ class ChatAdapterBase(ApiAdapter):
|
||||
|
||||
return build_test_request_body(cls.FORMAT_ID, request_data)
|
||||
|
||||
def extract_api_key(self, request: Request) -> Optional[str]:
|
||||
def extract_api_key(self, request: Request) -> str | None:
|
||||
"""从请求中提取 API 密钥,使用统一的 headers.py 实现"""
|
||||
return extract_client_api_key(dict(request.headers), self._get_api_format())
|
||||
|
||||
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
|
||||
def __init__(self, allowed_api_formats: list[str] | None = None):
|
||||
self.allowed_api_formats = allowed_api_formats or [self.FORMAT_ID]
|
||||
|
||||
async def handle(self, context: ApiRequestContext):
|
||||
@@ -282,8 +282,8 @@ class ChatAdapterBase(ApiAdapter):
|
||||
)
|
||||
|
||||
def _merge_path_params(
|
||||
self, original_request_body: Dict[str, Any], path_params: Dict[str, Any]
|
||||
) -> Dict[str, Any]:
|
||||
self, original_request_body: dict[str, Any], path_params: dict[str, Any]
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
合并 URL 路径参数到请求体 - 子类可覆盖
|
||||
|
||||
@@ -316,7 +316,7 @@ class ChatAdapterBase(ApiAdapter):
|
||||
"""
|
||||
pass
|
||||
|
||||
def _extract_message_count(self, payload: Dict[str, Any], request_obj) -> int:
|
||||
def _extract_message_count(self, payload: dict[str, Any], request_obj) -> int:
|
||||
"""
|
||||
提取消息数量 - 子类可覆盖
|
||||
|
||||
@@ -327,7 +327,7 @@ class ChatAdapterBase(ApiAdapter):
|
||||
messages = request_obj.messages
|
||||
return len(messages) if isinstance(messages, list) else 0
|
||||
|
||||
def _build_audit_metadata(self, payload: Dict[str, Any], request_obj) -> Dict[str, Any]:
|
||||
def _build_audit_metadata(self, payload: dict[str, Any], request_obj) -> dict[str, Any]:
|
||||
"""
|
||||
构建审计日志元数据 - 子类可覆盖
|
||||
"""
|
||||
@@ -355,8 +355,8 @@ class ChatAdapterBase(ApiAdapter):
|
||||
model: str,
|
||||
stream: bool,
|
||||
start_time: float,
|
||||
original_headers: Dict[str, str],
|
||||
original_request_body: Dict[str, Any],
|
||||
original_headers: dict[str, str],
|
||||
original_request_body: dict[str, Any],
|
||||
client_ip: str,
|
||||
request_id: str,
|
||||
) -> JSONResponse:
|
||||
@@ -426,8 +426,8 @@ class ChatAdapterBase(ApiAdapter):
|
||||
model: str,
|
||||
stream: bool,
|
||||
start_time: float,
|
||||
original_headers: Dict[str, str],
|
||||
original_request_body: Dict[str, Any],
|
||||
original_headers: dict[str, str],
|
||||
original_request_body: dict[str, Any],
|
||||
client_ip: str,
|
||||
request_id: str,
|
||||
) -> JSONResponse:
|
||||
@@ -527,12 +527,12 @@ class ChatAdapterBase(ApiAdapter):
|
||||
cache_read_input_tokens: int,
|
||||
input_price_per_1m: float,
|
||||
output_price_per_1m: float,
|
||||
cache_creation_price_per_1m: Optional[float],
|
||||
cache_read_price_per_1m: Optional[float],
|
||||
price_per_request: Optional[float],
|
||||
tiered_pricing: Optional[dict] = None,
|
||||
cache_ttl_minutes: Optional[int] = None,
|
||||
) -> Dict[str, Any]:
|
||||
cache_creation_price_per_1m: float | None,
|
||||
cache_read_price_per_1m: float | None,
|
||||
price_per_request: float | None,
|
||||
tiered_pricing: dict | None = None,
|
||||
cache_ttl_minutes: int | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
计算请求成本
|
||||
|
||||
@@ -597,8 +597,8 @@ class ChatAdapterBase(ApiAdapter):
|
||||
client: httpx.AsyncClient,
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
) -> Tuple[list, Optional[str]]:
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
) -> tuple[list, str | None]:
|
||||
"""
|
||||
查询上游 API 支持的模型列表
|
||||
|
||||
@@ -626,16 +626,16 @@ class ChatAdapterBase(ApiAdapter):
|
||||
client: httpx.AsyncClient,
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
request_data: Dict[str, Any],
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
request_data: dict[str, Any],
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
# 用量计算参数(现在强制记录)
|
||||
db: Optional[Any] = None,
|
||||
user: Optional[Any] = None,
|
||||
provider_name: Optional[str] = None,
|
||||
provider_id: Optional[str] = None,
|
||||
api_key_id: Optional[str] = None,
|
||||
model_name: Optional[str] = None,
|
||||
) -> Dict[str, Any]:
|
||||
db: Any | None = None,
|
||||
user: Any | None = None,
|
||||
provider_name: str | None = None,
|
||||
provider_id: str | None = None,
|
||||
api_key_id: str | None = None,
|
||||
model_name: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
测试模型连接性(非流式)
|
||||
|
||||
@@ -682,11 +682,11 @@ class ChatAdapterBase(ApiAdapter):
|
||||
# Adapter 注册表 - 用于根据 API format 获取 Adapter 实例
|
||||
# =========================================================================
|
||||
|
||||
_ADAPTER_REGISTRY: Dict[str, Type["ChatAdapterBase"]] = {}
|
||||
_ADAPTER_REGISTRY: dict[str, type[ChatAdapterBase]] = {}
|
||||
_ADAPTERS_LOADED = False
|
||||
|
||||
|
||||
def register_adapter(adapter_class: Type["ChatAdapterBase"]) -> Type["ChatAdapterBase"]:
|
||||
def register_adapter(adapter_class: type[ChatAdapterBase]) -> type[ChatAdapterBase]:
|
||||
"""
|
||||
注册 Adapter 类到注册表
|
||||
|
||||
@@ -731,7 +731,7 @@ def _ensure_adapters_loaded():
|
||||
_ADAPTERS_LOADED = True
|
||||
|
||||
|
||||
def get_adapter_class(api_format: str) -> Optional[Type["ChatAdapterBase"]]:
|
||||
def get_adapter_class(api_format: str) -> type[ChatAdapterBase] | None:
|
||||
"""
|
||||
根据 API format 获取 Adapter 类
|
||||
|
||||
@@ -745,7 +745,7 @@ def get_adapter_class(api_format: str) -> Optional[Type["ChatAdapterBase"]]:
|
||||
return _ADAPTER_REGISTRY.get(api_format.upper()) if api_format else None
|
||||
|
||||
|
||||
def get_adapter_instance(api_format: str) -> Optional["ChatAdapterBase"]:
|
||||
def get_adapter_instance(api_format: str) -> ChatAdapterBase | None:
|
||||
"""
|
||||
根据 API format 获取 Adapter 实例
|
||||
|
||||
|
||||
@@ -22,7 +22,10 @@ Chat Handler Base - Chat API 格式的通用基类
|
||||
import asyncio
|
||||
import json
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Any, AsyncGenerator, Awaitable, Callable, Dict, Optional, Union
|
||||
from typing import Any
|
||||
|
||||
from collections.abc import Callable
|
||||
from collections.abc import AsyncGenerator, Awaitable
|
||||
|
||||
import httpx
|
||||
from fastapi import BackgroundTasks, Request
|
||||
@@ -75,10 +78,10 @@ def _get_error_status_code(e: Exception, default: int = 400) -> int:
|
||||
|
||||
|
||||
def _convert_error_response_best_effort(
|
||||
error_response: Dict[str, Any],
|
||||
error_response: dict[str, Any],
|
||||
source_format: str,
|
||||
target_format: str,
|
||||
) -> Dict[str, Any]:
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
将上游错误响应 best-effort 转换为客户端格式。
|
||||
|
||||
@@ -97,7 +100,7 @@ def _convert_error_response_best_effort(
|
||||
def _build_client_error_response_best_effort(
|
||||
message: str,
|
||||
target_format: str,
|
||||
) -> Dict[str, Any]:
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
当无法解析上游错误 body 时,构造一个目标格式的错误响应(best-effort)。
|
||||
"""
|
||||
@@ -117,11 +120,11 @@ def _build_client_error_response_best_effort(
|
||||
|
||||
|
||||
def _build_error_json_payload(
|
||||
e: Union[ThinkingSignatureException, UpstreamClientException],
|
||||
e: ThinkingSignatureException | UpstreamClientException,
|
||||
client_format: str,
|
||||
provider_format: str,
|
||||
needs_conversion: bool = True,
|
||||
) -> Dict[str, Any]:
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
构建错误 JSON 响应 payload(公共逻辑)。
|
||||
|
||||
@@ -185,10 +188,10 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
client_ip: str,
|
||||
user_agent: str,
|
||||
start_time: float,
|
||||
allowed_api_formats: Optional[list] = None,
|
||||
adapter_detector: Optional[
|
||||
Callable[[Dict[str, str], Optional[Dict[str, Any]]], Dict[str, bool]]
|
||||
] = None,
|
||||
allowed_api_formats: list | None = None,
|
||||
adapter_detector: None | (
|
||||
Callable[[dict[str, str], dict[str, Any] | None], dict[str, bool]]
|
||||
) = None,
|
||||
):
|
||||
allowed = allowed_api_formats or [self.FORMAT_ID]
|
||||
super().__init__(
|
||||
@@ -202,7 +205,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
allowed_api_formats=allowed,
|
||||
adapter_detector=adapter_detector,
|
||||
)
|
||||
self._parser: Optional[ResponseParser] = None
|
||||
self._parser: ResponseParser | None = None
|
||||
self._request_builder = PassthroughRequestBuilder()
|
||||
|
||||
@property
|
||||
@@ -228,7 +231,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def _extract_usage(self, response: Dict) -> Dict[str, int]:
|
||||
def _extract_usage(self, response: dict) -> dict[str, int]:
|
||||
"""
|
||||
从响应中提取 token 使用情况
|
||||
|
||||
@@ -241,7 +244,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
"""
|
||||
pass
|
||||
|
||||
def _normalize_response(self, response: Dict) -> Dict:
|
||||
def _normalize_response(self, response: dict) -> dict:
|
||||
"""
|
||||
规范化响应(可选覆盖)
|
||||
|
||||
@@ -257,8 +260,8 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
|
||||
def extract_model_from_request(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
path_params: Optional[Dict[str, Any]] = None, # noqa: ARG002 - 子类使用
|
||||
request_body: dict[str, Any],
|
||||
path_params: dict[str, Any] | None = None, # noqa: ARG002 - 子类使用
|
||||
) -> str:
|
||||
"""
|
||||
从请求中提取模型名 - 子类可覆盖
|
||||
@@ -282,9 +285,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
|
||||
def apply_mapped_model(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
request_body: dict[str, Any],
|
||||
mapped_model: str, # noqa: ARG002 - 子类使用
|
||||
) -> Dict[str, Any]:
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
将映射后的模型名应用到请求体
|
||||
|
||||
@@ -303,9 +306,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
|
||||
def get_model_for_url(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
mapped_model: Optional[str],
|
||||
) -> Optional[str]:
|
||||
request_body: dict[str, Any],
|
||||
mapped_model: str | None,
|
||||
) -> str | None:
|
||||
"""
|
||||
获取用于 URL 路径的模型名
|
||||
|
||||
@@ -323,8 +326,8 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
|
||||
def prepare_provider_request_body(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
) -> Dict[str, Any]:
|
||||
request_body: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
准备发送给 Provider 的请求体 - 子类可覆盖
|
||||
|
||||
@@ -341,9 +344,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
|
||||
def _set_model_after_conversion(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
request_body: dict[str, Any],
|
||||
provider_api_format: str,
|
||||
mapped_model: Optional[str],
|
||||
mapped_model: str | None,
|
||||
fallback_model: str,
|
||||
) -> None:
|
||||
"""
|
||||
@@ -372,7 +375,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
|
||||
def _set_stream_after_conversion(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
request_body: dict[str, Any],
|
||||
client_api_format: str,
|
||||
provider_api_format: str,
|
||||
is_stream: bool,
|
||||
@@ -414,8 +417,8 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
self,
|
||||
source_model: str,
|
||||
provider_id: str,
|
||||
api_format: Optional[str] = None,
|
||||
) -> Optional[str]:
|
||||
api_format: str | None = None,
|
||||
) -> str | None:
|
||||
"""
|
||||
获取模型映射后的实际模型名
|
||||
|
||||
@@ -452,10 +455,10 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
self,
|
||||
request: Any,
|
||||
http_request: Request,
|
||||
original_headers: Dict[str, Any],
|
||||
original_request_body: Dict[str, Any],
|
||||
query_params: Optional[Dict[str, str]] = None,
|
||||
) -> Union[StreamingResponse, JSONResponse]:
|
||||
original_headers: dict[str, Any],
|
||||
original_request_body: dict[str, Any],
|
||||
query_params: dict[str, str] | None = None,
|
||||
) -> StreamingResponse | JSONResponse:
|
||||
"""处理流式响应"""
|
||||
logger.debug(f"开始流式响应处理 ({self.FORMAT_ID})")
|
||||
|
||||
@@ -466,7 +469,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
|
||||
# 可变请求体容器:允许 orchestrator 在遇到 Thinking 签名错误时整流请求体后重试
|
||||
# 结构: {"body": 实际请求体, "_rectified": 是否已整流, "_rectified_this_turn": 本轮是否整流}
|
||||
request_body_ref: Dict[str, Any] = {"body": original_request_body}
|
||||
request_body_ref: dict[str, Any] = {"body": original_request_body}
|
||||
|
||||
# 创建类型安全的流式上下文
|
||||
ctx = StreamContext(model=model, api_format=api_format)
|
||||
@@ -492,7 +495,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
endpoint: ProviderEndpoint,
|
||||
key: ProviderAPIKey,
|
||||
candidate: ProviderCandidate,
|
||||
) -> AsyncGenerator[bytes, None]:
|
||||
) -> AsyncGenerator[bytes]:
|
||||
return await self._execute_stream_request(
|
||||
ctx,
|
||||
stream_processor,
|
||||
@@ -609,12 +612,12 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
provider: Provider,
|
||||
endpoint: ProviderEndpoint,
|
||||
key: ProviderAPIKey,
|
||||
original_request_body: Dict[str, Any],
|
||||
original_headers: Dict[str, str],
|
||||
query_params: Optional[Dict[str, str]] = None,
|
||||
candidate: Optional[ProviderCandidate] = None,
|
||||
is_disconnected: Optional[Callable[[], Awaitable[bool]]] = None,
|
||||
) -> AsyncGenerator[bytes, None]:
|
||||
original_request_body: dict[str, Any],
|
||||
original_headers: dict[str, str],
|
||||
query_params: dict[str, str] | None = None,
|
||||
candidate: ProviderCandidate | None = None,
|
||||
is_disconnected: Callable[[], Awaitable[bool]] | None = None,
|
||||
) -> AsyncGenerator[bytes]:
|
||||
"""执行流式请求并返回流生成器"""
|
||||
# 重置上下文状态(重试时清除之前的数据)
|
||||
ctx.reset_for_retry()
|
||||
@@ -787,7 +790,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
ctx.error_message = "client_disconnected_during_prefetch"
|
||||
raise
|
||||
|
||||
except asyncio.TimeoutError:
|
||||
except TimeoutError:
|
||||
# 整体请求超时(建立连接 + 获取首字节)
|
||||
# 清理可能已建立的连接上下文
|
||||
if response_ctx is not None:
|
||||
@@ -844,8 +847,8 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
self,
|
||||
ctx: StreamContext,
|
||||
error: Exception,
|
||||
original_headers: Dict[str, str],
|
||||
original_request_body: Dict[str, Any],
|
||||
original_headers: dict[str, str],
|
||||
original_request_body: dict[str, Any],
|
||||
) -> None:
|
||||
"""记录流式请求失败"""
|
||||
response_time_ms = self.elapsed_ms()
|
||||
@@ -892,9 +895,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
self,
|
||||
request: Any,
|
||||
http_request: Request,
|
||||
original_headers: Dict[str, Any],
|
||||
original_request_body: Dict[str, Any],
|
||||
query_params: Optional[Dict[str, str]] = None,
|
||||
original_headers: dict[str, Any],
|
||||
original_request_body: dict[str, Any],
|
||||
query_params: dict[str, str] | None = None,
|
||||
) -> JSONResponse:
|
||||
"""处理非流式响应"""
|
||||
logger.debug(f"开始非流式响应处理 ({self.FORMAT_ID})")
|
||||
@@ -906,29 +909,29 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
|
||||
# 可变请求体容器:允许 orchestrator 在遇到 Thinking 签名错误时整流请求体后重试
|
||||
# 结构: {"body": 实际请求体, "_rectified": 是否已整流, "_rectified_this_turn": 本轮是否整流}
|
||||
request_body_ref: Dict[str, Any] = {"body": original_request_body}
|
||||
request_body_ref: dict[str, Any] = {"body": original_request_body}
|
||||
|
||||
# 用于跟踪的变量
|
||||
provider_name: Optional[str] = None
|
||||
response_json: Optional[Dict[str, Any]] = None
|
||||
provider_name: str | None = None
|
||||
response_json: dict[str, Any] | None = None
|
||||
status_code = 200
|
||||
response_headers: Dict[str, str] = {}
|
||||
provider_request_headers: Dict[str, str] = {}
|
||||
provider_request_body: Optional[Dict[str, Any]] = None
|
||||
provider_api_format_for_error: Optional[str] = None
|
||||
client_api_format_for_error: Optional[str] = None
|
||||
response_headers: dict[str, str] = {}
|
||||
provider_request_headers: dict[str, str] = {}
|
||||
provider_request_body: dict[str, Any] | None = None
|
||||
provider_api_format_for_error: str | None = None
|
||||
client_api_format_for_error: str | None = None
|
||||
needs_conversion_for_error: bool = False
|
||||
provider_id: Optional[str] = None # Provider ID(用于失败记录)
|
||||
endpoint_id: Optional[str] = None # Endpoint ID(用于失败记录)
|
||||
key_id: Optional[str] = None # Key ID(用于失败记录)
|
||||
mapped_model_result: Optional[str] = None # 映射后的目标模型名(用于 Usage 记录)
|
||||
provider_id: str | None = None # Provider ID(用于失败记录)
|
||||
endpoint_id: str | None = None # Endpoint ID(用于失败记录)
|
||||
key_id: str | None = None # Key ID(用于失败记录)
|
||||
mapped_model_result: str | None = None # 映射后的目标模型名(用于 Usage 记录)
|
||||
|
||||
async def sync_request_func(
|
||||
provider: Provider,
|
||||
endpoint: ProviderEndpoint,
|
||||
key: ProviderAPIKey,
|
||||
candidate: ProviderCandidate,
|
||||
) -> Dict[str, Any]:
|
||||
) -> dict[str, Any]:
|
||||
nonlocal provider_name, response_json, status_code, response_headers
|
||||
nonlocal provider_request_headers, provider_request_body, mapped_model_result
|
||||
nonlocal provider_api_format_for_error, client_api_format_for_error, needs_conversion_for_error
|
||||
@@ -1269,7 +1272,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
actual_request_body = provider_request_body or original_request_body
|
||||
|
||||
# 尝试从异常中提取响应头
|
||||
error_response_headers: Dict[str, str] = {}
|
||||
error_response_headers: dict[str, str] = {}
|
||||
if isinstance(e, ProviderRateLimitException) and e.response_headers:
|
||||
error_response_headers = e.response_headers
|
||||
elif isinstance(e, httpx.HTTPStatusError) and hasattr(e, "response"):
|
||||
|
||||
@@ -17,7 +17,7 @@ CLI Adapter 通用基类
|
||||
|
||||
import time
|
||||
import traceback
|
||||
from typing import Any, Dict, Optional, Tuple, Type
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from fastapi import HTTPException, Request
|
||||
@@ -63,7 +63,7 @@ class CliAdapterBase(ApiAdapter):
|
||||
|
||||
# 子类必须覆盖
|
||||
FORMAT_ID: str = "UNKNOWN"
|
||||
HANDLER_CLASS: Type[CliMessageHandlerBase]
|
||||
HANDLER_CLASS: type[CliMessageHandlerBase]
|
||||
|
||||
# 适配器配置
|
||||
name: str = "cli.base"
|
||||
@@ -72,7 +72,7 @@ class CliAdapterBase(ApiAdapter):
|
||||
# 计费模板配置(子类可覆盖,如 "claude", "openai", "gemini")
|
||||
BILLING_TEMPLATE: str = "claude"
|
||||
|
||||
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
|
||||
def __init__(self, allowed_api_formats: list[str] | None = None):
|
||||
self.allowed_api_formats = allowed_api_formats or [self.FORMAT_ID]
|
||||
|
||||
# =========================================================================
|
||||
@@ -87,7 +87,7 @@ class CliAdapterBase(ApiAdapter):
|
||||
except KeyError:
|
||||
return APIFormat.OPENAI
|
||||
|
||||
def extract_api_key(self, request: Request) -> Optional[str]:
|
||||
def extract_api_key(self, request: Request) -> str | None:
|
||||
"""
|
||||
从请求中提取 API 密钥
|
||||
|
||||
@@ -96,7 +96,7 @@ class CliAdapterBase(ApiAdapter):
|
||||
return extract_client_api_key(dict(request.headers), self._get_api_format())
|
||||
|
||||
@classmethod
|
||||
def build_base_headers(cls, api_key: str) -> Dict[str, str]:
|
||||
def build_base_headers(cls, api_key: str) -> dict[str, str]:
|
||||
"""
|
||||
构建 CLI API 认证头
|
||||
|
||||
@@ -106,8 +106,8 @@ class CliAdapterBase(ApiAdapter):
|
||||
|
||||
@classmethod
|
||||
def build_headers_with_extra(
|
||||
cls, api_key: str, extra_headers: Optional[Dict[str, str]] = None
|
||||
) -> Dict[str, str]:
|
||||
cls, api_key: str, extra_headers: dict[str, str] | None = None
|
||||
) -> dict[str, str]:
|
||||
"""
|
||||
构建带额外头部的完整请求头
|
||||
|
||||
@@ -260,8 +260,8 @@ class CliAdapterBase(ApiAdapter):
|
||||
)
|
||||
|
||||
def _merge_path_params(
|
||||
self, original_request_body: Dict[str, Any], path_params: Dict[str, Any]
|
||||
) -> Dict[str, Any]:
|
||||
self, original_request_body: dict[str, Any], path_params: dict[str, Any]
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
合并 URL 路径参数到请求体 - 子类可覆盖
|
||||
|
||||
@@ -280,7 +280,7 @@ class CliAdapterBase(ApiAdapter):
|
||||
merged[key] = value
|
||||
return merged
|
||||
|
||||
def _extract_message_count(self, payload: Dict[str, Any]) -> int:
|
||||
def _extract_message_count(self, payload: dict[str, Any]) -> int:
|
||||
"""
|
||||
提取消息数量 - 子类可覆盖
|
||||
|
||||
@@ -297,9 +297,9 @@ class CliAdapterBase(ApiAdapter):
|
||||
|
||||
def _build_audit_metadata(
|
||||
self,
|
||||
payload: Dict[str, Any],
|
||||
path_params: Optional[Dict[str, Any]] = None,
|
||||
) -> Dict[str, Any]:
|
||||
payload: dict[str, Any],
|
||||
path_params: dict[str, Any] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
构建审计日志元数据 - 子类可覆盖
|
||||
|
||||
@@ -338,8 +338,8 @@ class CliAdapterBase(ApiAdapter):
|
||||
model: str,
|
||||
stream: bool,
|
||||
start_time: float,
|
||||
original_headers: Dict[str, str],
|
||||
original_request_body: Dict[str, Any],
|
||||
original_headers: dict[str, str],
|
||||
original_request_body: dict[str, Any],
|
||||
client_ip: str,
|
||||
request_id: str,
|
||||
) -> JSONResponse:
|
||||
@@ -409,8 +409,8 @@ class CliAdapterBase(ApiAdapter):
|
||||
model: str,
|
||||
stream: bool,
|
||||
start_time: float,
|
||||
original_headers: Dict[str, str],
|
||||
original_request_body: Dict[str, Any],
|
||||
original_headers: dict[str, str],
|
||||
original_request_body: dict[str, Any],
|
||||
client_ip: str,
|
||||
request_id: str,
|
||||
) -> JSONResponse:
|
||||
@@ -507,12 +507,12 @@ class CliAdapterBase(ApiAdapter):
|
||||
cache_read_input_tokens: int,
|
||||
input_price_per_1m: float,
|
||||
output_price_per_1m: float,
|
||||
cache_creation_price_per_1m: Optional[float],
|
||||
cache_read_price_per_1m: Optional[float],
|
||||
price_per_request: Optional[float],
|
||||
tiered_pricing: Optional[dict] = None,
|
||||
cache_ttl_minutes: Optional[int] = None,
|
||||
) -> Dict[str, Any]:
|
||||
cache_creation_price_per_1m: float | None,
|
||||
cache_read_price_per_1m: float | None,
|
||||
price_per_request: float | None,
|
||||
tiered_pricing: dict | None = None,
|
||||
cache_ttl_minutes: int | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
计算请求成本
|
||||
|
||||
@@ -567,8 +567,8 @@ class CliAdapterBase(ApiAdapter):
|
||||
client: httpx.AsyncClient,
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
) -> Tuple[list, Optional[str]]:
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
) -> tuple[list, str | None]:
|
||||
"""
|
||||
查询上游 API 支持的模型列表
|
||||
|
||||
@@ -596,16 +596,16 @@ class CliAdapterBase(ApiAdapter):
|
||||
client: httpx.AsyncClient,
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
request_data: Dict[str, Any],
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
request_data: dict[str, Any],
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
# 用量计算参数
|
||||
db: Optional[Any] = None,
|
||||
user: Optional[Any] = None,
|
||||
provider_name: Optional[str] = None,
|
||||
provider_id: Optional[str] = None,
|
||||
api_key_id: Optional[str] = None,
|
||||
model_name: Optional[str] = None,
|
||||
) -> Dict[str, Any]:
|
||||
db: Any | None = None,
|
||||
user: Any | None = None,
|
||||
provider_name: str | None = None,
|
||||
provider_id: str | None = None,
|
||||
api_key_id: str | None = None,
|
||||
model_name: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
测试模型连接性(非流式)
|
||||
|
||||
@@ -669,7 +669,7 @@ class CliAdapterBase(ApiAdapter):
|
||||
# =========================================================================
|
||||
|
||||
@classmethod
|
||||
def build_endpoint_url(cls, base_url: str, request_data: Dict[str, Any], model_name: Optional[str] = None) -> str:
|
||||
def build_endpoint_url(cls, base_url: str, request_data: dict[str, Any], model_name: str | None = None) -> str:
|
||||
"""
|
||||
构建CLI API端点URL - 子类应覆盖
|
||||
|
||||
@@ -684,7 +684,7 @@ class CliAdapterBase(ApiAdapter):
|
||||
raise NotImplementedError(f"{cls.FORMAT_ID} adapter must implement build_endpoint_url")
|
||||
|
||||
@classmethod
|
||||
def build_request_body(cls, request_data: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
|
||||
def build_request_body(cls, request_data: dict[str, Any] | None = None) -> dict[str, Any]:
|
||||
"""构建测试请求体,使用转换器注册表自动处理格式转换
|
||||
|
||||
Args:
|
||||
@@ -698,7 +698,7 @@ class CliAdapterBase(ApiAdapter):
|
||||
return build_test_request_body(cls.FORMAT_ID, request_data)
|
||||
|
||||
@classmethod
|
||||
def get_cli_user_agent(cls) -> Optional[str]:
|
||||
def get_cli_user_agent(cls) -> str | None:
|
||||
"""
|
||||
获取CLI User-Agent - 子类可覆盖
|
||||
|
||||
@@ -708,7 +708,7 @@ class CliAdapterBase(ApiAdapter):
|
||||
return None
|
||||
|
||||
@classmethod
|
||||
def get_cli_extra_headers(cls) -> Dict[str, str]:
|
||||
def get_cli_extra_headers(cls) -> dict[str, str]:
|
||||
"""
|
||||
获取CLI额外请求头 - 子类可覆盖
|
||||
|
||||
@@ -718,7 +718,7 @@ class CliAdapterBase(ApiAdapter):
|
||||
Returns:
|
||||
额外请求头字典
|
||||
"""
|
||||
headers: Dict[str, str] = {}
|
||||
headers: dict[str, str] = {}
|
||||
cli_user_agent = cls.get_cli_user_agent()
|
||||
if cli_user_agent:
|
||||
headers["User-Agent"] = cli_user_agent
|
||||
@@ -728,11 +728,11 @@ class CliAdapterBase(ApiAdapter):
|
||||
# CLI Adapter 注册表 - 用于根据 API format 获取 CLI Adapter 实例
|
||||
# =========================================================================
|
||||
|
||||
_CLI_ADAPTER_REGISTRY: Dict[str, Type["CliAdapterBase"]] = {}
|
||||
_CLI_ADAPTER_REGISTRY: dict[str, type[CliAdapterBase]] = {}
|
||||
_CLI_ADAPTERS_LOADED = False
|
||||
|
||||
|
||||
def register_cli_adapter(adapter_class: Type["CliAdapterBase"]) -> Type["CliAdapterBase"]:
|
||||
def register_cli_adapter(adapter_class: type[CliAdapterBase]) -> type[CliAdapterBase]:
|
||||
"""
|
||||
注册 CLI Adapter 类到注册表
|
||||
|
||||
@@ -771,13 +771,13 @@ def _ensure_cli_adapters_loaded():
|
||||
_CLI_ADAPTERS_LOADED = True
|
||||
|
||||
|
||||
def get_cli_adapter_class(api_format: str) -> Optional[Type["CliAdapterBase"]]:
|
||||
def get_cli_adapter_class(api_format: str) -> type[CliAdapterBase] | None:
|
||||
"""根据 API format 获取 CLI Adapter 类"""
|
||||
_ensure_cli_adapters_loaded()
|
||||
return _CLI_ADAPTER_REGISTRY.get(api_format.upper()) if api_format else None
|
||||
|
||||
|
||||
def get_cli_adapter_instance(api_format: str) -> Optional["CliAdapterBase"]:
|
||||
def get_cli_adapter_instance(api_format: str) -> CliAdapterBase | None:
|
||||
"""根据 API format 获取 CLI Adapter 实例"""
|
||||
adapter_class = get_cli_adapter_class(api_format)
|
||||
if adapter_class:
|
||||
|
||||
@@ -17,14 +17,11 @@ import time
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
AsyncGenerator,
|
||||
Callable,
|
||||
Dict,
|
||||
List,
|
||||
Optional,
|
||||
Tuple,
|
||||
)
|
||||
|
||||
from collections.abc import Callable
|
||||
from collections.abc import AsyncGenerator
|
||||
|
||||
import httpx
|
||||
from fastapi import BackgroundTasks, Request
|
||||
from fastapi.responses import JSONResponse, StreamingResponse
|
||||
@@ -45,7 +42,6 @@ from src.api.handlers.base.request_builder import PassthroughRequestBuilder
|
||||
# 直接从具体模块导入,避免循环依赖
|
||||
from src.api.handlers.base.response_parser import (
|
||||
ResponseParser,
|
||||
StreamStats,
|
||||
)
|
||||
from src.api.handlers.base.stream_context import StreamContext
|
||||
from src.api.handlers.base.utils import (
|
||||
@@ -86,7 +82,7 @@ from src.utils.timeout import read_first_chunk_with_ttfb_timeout
|
||||
# ==============================================================================
|
||||
|
||||
|
||||
def _parse_sse_data_line(line: str) -> Tuple[Optional[Any], str]:
|
||||
def _parse_sse_data_line(line: str) -> tuple[Any | None, str]:
|
||||
"""
|
||||
解析标准 SSE data 行
|
||||
|
||||
@@ -108,7 +104,7 @@ def _parse_sse_data_line(line: str) -> Tuple[Optional[Any], str]:
|
||||
return None, "invalid"
|
||||
|
||||
|
||||
def _parse_sse_event_data_line(line: str) -> Tuple[Optional[Any], str]:
|
||||
def _parse_sse_event_data_line(line: str) -> tuple[Any | None, str]:
|
||||
"""
|
||||
解析 event + data 同行格式(如 "event: xxx data: {...}")
|
||||
|
||||
@@ -126,7 +122,7 @@ def _parse_sse_event_data_line(line: str) -> Tuple[Optional[Any], str]:
|
||||
return None, "invalid"
|
||||
|
||||
|
||||
def _parse_gemini_json_array_line(line: str) -> Tuple[Optional[Any], str]:
|
||||
def _parse_gemini_json_array_line(line: str) -> tuple[Any | None, str]:
|
||||
"""
|
||||
解析 Gemini JSON-array 格式的裸 JSON 行
|
||||
|
||||
@@ -151,9 +147,9 @@ def _parse_gemini_json_array_line(line: str) -> Tuple[Optional[Any], str]:
|
||||
|
||||
|
||||
def _format_converted_events_to_sse(
|
||||
converted_events: List[Dict[str, Any]],
|
||||
converted_events: list[dict[str, Any]],
|
||||
client_format: str,
|
||||
) -> List[str]:
|
||||
) -> list[str]:
|
||||
"""
|
||||
将转换后的事件格式化为 SSE 行
|
||||
|
||||
@@ -164,7 +160,7 @@ def _format_converted_events_to_sse(
|
||||
Returns:
|
||||
SSE 行列表(每个元素是完整的 SSE 事件,包含尾部空行)
|
||||
"""
|
||||
result: List[str] = []
|
||||
result: list[str] = []
|
||||
needs_event_line = client_format.upper() in ("CLAUDE", "CLAUDE_CLI")
|
||||
|
||||
for evt in converted_events:
|
||||
@@ -213,10 +209,10 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
client_ip: str,
|
||||
user_agent: str,
|
||||
start_time: float,
|
||||
allowed_api_formats: Optional[list] = None,
|
||||
adapter_detector: Optional[
|
||||
Callable[[Dict[str, str], Optional[Dict[str, Any]]], Dict[str, bool]]
|
||||
] = None,
|
||||
allowed_api_formats: list | None = None,
|
||||
adapter_detector: None | (
|
||||
Callable[[dict[str, str], dict[str, Any] | None], dict[str, bool]]
|
||||
) = None,
|
||||
):
|
||||
allowed = allowed_api_formats or [self.FORMAT_ID]
|
||||
super().__init__(
|
||||
@@ -230,7 +226,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
allowed_api_formats=allowed,
|
||||
adapter_detector=adapter_detector,
|
||||
)
|
||||
self._parser: Optional[ResponseParser] = None
|
||||
self._parser: ResponseParser | None = None
|
||||
self._request_builder = PassthroughRequestBuilder()
|
||||
|
||||
@property
|
||||
@@ -253,7 +249,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
self,
|
||||
source_model: str,
|
||||
provider_id: str,
|
||||
) -> Optional[str]:
|
||||
) -> str | None:
|
||||
"""
|
||||
获取模型映射后的实际模型名
|
||||
|
||||
@@ -296,8 +292,8 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
|
||||
def extract_model_from_request(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
path_params: Optional[Dict[str, Any]] = None, # noqa: ARG002 - 子类使用
|
||||
request_body: dict[str, Any],
|
||||
path_params: dict[str, Any] | None = None, # noqa: ARG002 - 子类使用
|
||||
) -> str:
|
||||
"""
|
||||
从请求中提取模型名 - 子类可覆盖
|
||||
@@ -321,9 +317,9 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
|
||||
def apply_mapped_model(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
request_body: dict[str, Any],
|
||||
mapped_model: str, # noqa: ARG002 - 子类使用
|
||||
) -> Dict[str, Any]:
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
将映射后的模型名应用到请求体
|
||||
|
||||
@@ -342,8 +338,8 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
|
||||
def prepare_provider_request_body(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
) -> Dict[str, Any]:
|
||||
request_body: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
准备发送给 Provider 的请求体 - 子类可覆盖
|
||||
|
||||
@@ -359,7 +355,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
return request_body
|
||||
|
||||
@staticmethod
|
||||
def _get_format_metadata(format_id: str) -> Optional["ApiFormatDefinition"]:
|
||||
def _get_format_metadata(format_id: str) -> ApiFormatDefinition | None:
|
||||
"""获取格式元数据(解析失败返回 None)"""
|
||||
from src.core.api_format import APIFormat
|
||||
from src.core.api_format.metadata import API_FORMAT_DEFINITIONS
|
||||
@@ -372,10 +368,10 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
|
||||
def _finalize_converted_request(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
request_body: dict[str, Any],
|
||||
client_api_format: str,
|
||||
provider_api_format: str,
|
||||
mapped_model: Optional[str],
|
||||
mapped_model: str | None,
|
||||
fallback_model: str,
|
||||
is_stream: bool,
|
||||
) -> None:
|
||||
@@ -418,13 +414,13 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
|
||||
def _convert_request_for_cross_format(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
request_body: dict[str, Any],
|
||||
client_api_format: str,
|
||||
provider_api_format: str,
|
||||
mapped_model: Optional[str],
|
||||
mapped_model: str | None,
|
||||
fallback_model: str,
|
||||
is_stream: bool,
|
||||
) -> Tuple[Dict[str, Any], str]:
|
||||
) -> tuple[dict[str, Any], str]:
|
||||
"""
|
||||
跨格式请求转换的公共逻辑
|
||||
|
||||
@@ -465,9 +461,9 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
|
||||
def get_model_for_url(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
mapped_model: Optional[str],
|
||||
) -> Optional[str]:
|
||||
request_body: dict[str, Any],
|
||||
mapped_model: str | None,
|
||||
) -> str | None:
|
||||
"""
|
||||
获取用于 URL 路径的模型名
|
||||
|
||||
@@ -485,8 +481,8 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
|
||||
def _extract_response_metadata(
|
||||
self,
|
||||
response: Dict[str, Any],
|
||||
) -> Dict[str, Any]:
|
||||
response: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
从响应中提取 Provider 特有的元数据 - 子类可覆盖
|
||||
|
||||
@@ -503,11 +499,11 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
|
||||
async def process_stream(
|
||||
self,
|
||||
original_request_body: Dict[str, Any],
|
||||
original_headers: Dict[str, str],
|
||||
query_params: Optional[Dict[str, str]] = None,
|
||||
path_params: Optional[Dict[str, Any]] = None,
|
||||
http_request: Optional[Request] = None,
|
||||
original_request_body: dict[str, Any],
|
||||
original_headers: dict[str, str],
|
||||
query_params: dict[str, str] | None = None,
|
||||
path_params: dict[str, Any] | None = None,
|
||||
http_request: Request | None = None,
|
||||
) -> StreamingResponse:
|
||||
"""
|
||||
处理流式请求
|
||||
@@ -529,7 +525,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
|
||||
# 可变请求体容器:允许 orchestrator 在遇到 Thinking 签名错误时整流请求体后重试
|
||||
# 结构: {"body": 实际请求体, "_rectified": 是否已整流, "_rectified_this_turn": 本轮是否整流}
|
||||
request_body_ref: Dict[str, Any] = {"body": original_request_body}
|
||||
request_body_ref: dict[str, Any] = {"body": original_request_body}
|
||||
|
||||
# 使用子类实现的方法提取 model(不同 API 格式的 model 位置不同)
|
||||
# 注意:使用 original_request_body,因为整流只修改 messages,不影响 model 字段
|
||||
@@ -550,7 +546,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
endpoint: ProviderEndpoint,
|
||||
key: ProviderAPIKey,
|
||||
candidate: ProviderCandidate,
|
||||
) -> AsyncGenerator[bytes, None]:
|
||||
) -> AsyncGenerator[bytes]:
|
||||
return await self._execute_stream_request(
|
||||
ctx,
|
||||
provider,
|
||||
@@ -647,12 +643,12 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
provider: Provider,
|
||||
endpoint: ProviderEndpoint,
|
||||
key: ProviderAPIKey,
|
||||
original_request_body: Dict[str, Any],
|
||||
original_headers: Dict[str, str],
|
||||
query_params: Optional[Dict[str, str]] = None,
|
||||
candidate: Optional[ProviderCandidate] = None,
|
||||
http_request: Optional[Request] = None,
|
||||
) -> AsyncGenerator[bytes, None]:
|
||||
original_request_body: dict[str, Any],
|
||||
original_headers: dict[str, str],
|
||||
query_params: dict[str, str] | None = None,
|
||||
candidate: ProviderCandidate | None = None,
|
||||
http_request: Request | None = None,
|
||||
) -> AsyncGenerator[bytes]:
|
||||
"""执行流式请求并返回流生成器"""
|
||||
# 重置上下文状态(重试时清除之前的数据,避免累积)
|
||||
ctx.parsed_chunks = []
|
||||
@@ -812,7 +808,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
else:
|
||||
await asyncio.wait_for(_connect_and_prefetch(), timeout=request_timeout)
|
||||
|
||||
except asyncio.TimeoutError:
|
||||
except TimeoutError:
|
||||
# 整体请求超时(建立连接 + 获取首字节)
|
||||
# 清理可能已建立的连接上下文
|
||||
if response_ctx is not None:
|
||||
@@ -886,7 +882,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
stream_response: httpx.Response,
|
||||
response_ctx: Any,
|
||||
http_client: httpx.AsyncClient,
|
||||
) -> AsyncGenerator[bytes, None]:
|
||||
) -> AsyncGenerator[bytes]:
|
||||
"""创建响应流生成器(使用字节流)"""
|
||||
try:
|
||||
sse_parser = SSEEventParser()
|
||||
@@ -949,9 +945,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
},
|
||||
}
|
||||
self._mark_first_output(ctx, output_state)
|
||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode(
|
||||
"utf-8"
|
||||
)
|
||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
|
||||
return # 结束生成器
|
||||
|
||||
# 格式转换或直接透传
|
||||
@@ -1003,7 +997,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
"message": ctx.error_message,
|
||||
},
|
||||
}
|
||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode("utf-8")
|
||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
|
||||
else:
|
||||
logger.debug("流式数据转发完成")
|
||||
# 为 OpenAI 客户端补齐 [DONE] 标记(非 CLI 格式)
|
||||
@@ -1028,7 +1022,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
"message": ctx.error_message,
|
||||
},
|
||||
}
|
||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode("utf-8")
|
||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
|
||||
except httpx.RemoteProtocolError:
|
||||
if ctx.data_count > 0:
|
||||
error_event = {
|
||||
@@ -1038,7 +1032,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
"message": "上游连接意外关闭,部分响应已成功传输",
|
||||
},
|
||||
}
|
||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode("utf-8")
|
||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
|
||||
else:
|
||||
raise
|
||||
finally:
|
||||
@@ -1229,7 +1223,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
except (EmbeddedErrorException, ProviderTimeoutException, ProviderNotAvailableException):
|
||||
# 重新抛出可重试的 Provider 异常,触发故障转移
|
||||
raise
|
||||
except (OSError, IOError) as e:
|
||||
except OSError as e:
|
||||
# 网络 I/O 异常:记录警告,可能需要重试
|
||||
logger.warning(f" [{self.request_id}] 预读流时发生网络异常: {type(e).__name__}: {e}")
|
||||
except Exception as e:
|
||||
@@ -1249,7 +1243,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
response_ctx: Any,
|
||||
http_client: httpx.AsyncClient,
|
||||
prefetched_chunks: list,
|
||||
) -> AsyncGenerator[bytes, None]:
|
||||
) -> AsyncGenerator[bytes]:
|
||||
"""创建响应流生成器(带预读数据,使用字节流)"""
|
||||
try:
|
||||
sse_parser = SSEEventParser()
|
||||
@@ -1370,9 +1364,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
},
|
||||
}
|
||||
self._mark_first_output(ctx, output_state)
|
||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode(
|
||||
"utf-8"
|
||||
)
|
||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
|
||||
return
|
||||
|
||||
# 格式转换或直接透传
|
||||
@@ -1427,7 +1419,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
"message": ctx.error_message,
|
||||
},
|
||||
}
|
||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode("utf-8")
|
||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
|
||||
else:
|
||||
logger.debug("流式数据转发完成")
|
||||
# 为 OpenAI 客户端补齐 [DONE] 标记(非 CLI 格式)
|
||||
@@ -1451,7 +1443,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
"message": ctx.error_message,
|
||||
},
|
||||
}
|
||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode("utf-8")
|
||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
|
||||
except httpx.RemoteProtocolError:
|
||||
if ctx.data_count > 0:
|
||||
error_event = {
|
||||
@@ -1461,7 +1453,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
"message": "上游连接意外关闭,部分响应已成功传输",
|
||||
},
|
||||
}
|
||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode("utf-8")
|
||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
|
||||
else:
|
||||
raise
|
||||
finally:
|
||||
@@ -1477,7 +1469,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
def _handle_sse_event(
|
||||
self,
|
||||
ctx: StreamContext,
|
||||
event_name: Optional[str],
|
||||
event_name: str | None,
|
||||
data_str: str,
|
||||
record_chunk: bool = False,
|
||||
) -> None:
|
||||
@@ -1526,7 +1518,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
self,
|
||||
ctx: StreamContext,
|
||||
event_type: str,
|
||||
data: Dict[str, Any],
|
||||
data: dict[str, Any],
|
||||
) -> None:
|
||||
"""
|
||||
处理解析后的事件数据 - 子类应覆盖此方法
|
||||
@@ -1600,7 +1592,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
def _record_converted_chunks(
|
||||
self,
|
||||
ctx: StreamContext,
|
||||
converted_events: List[Dict[str, Any]],
|
||||
converted_events: list[dict[str, Any]],
|
||||
) -> None:
|
||||
"""
|
||||
记录转换后的 chunk 数据到 parsed_chunks,并更新统计信息
|
||||
@@ -1644,7 +1636,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
def _extract_usage_from_converted_event(
|
||||
self,
|
||||
ctx: StreamContext,
|
||||
evt: Dict[str, Any],
|
||||
evt: dict[str, Any],
|
||||
event_type: str,
|
||||
) -> None:
|
||||
"""
|
||||
@@ -1660,7 +1652,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
evt: 转换后的事件
|
||||
event_type: 事件类型
|
||||
"""
|
||||
usage: Optional[Dict[str, Any]] = None
|
||||
usage: dict[str, Any] | None = None
|
||||
|
||||
# Claude 格式: message_delta 或 message_start
|
||||
if event_type == "message_delta":
|
||||
@@ -1725,9 +1717,9 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
async def _create_monitored_stream(
|
||||
self,
|
||||
ctx: StreamContext,
|
||||
stream_generator: AsyncGenerator[bytes, None],
|
||||
http_request: Optional[Request] = None,
|
||||
) -> AsyncGenerator[bytes, None]:
|
||||
stream_generator: AsyncGenerator[bytes],
|
||||
http_request: Request | None = None,
|
||||
) -> AsyncGenerator[bytes]:
|
||||
"""
|
||||
创建带监控的流生成器
|
||||
|
||||
@@ -1821,8 +1813,8 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
async def _record_stream_stats(
|
||||
self,
|
||||
ctx: StreamContext,
|
||||
original_headers: Dict[str, str],
|
||||
original_request_body: Dict[str, Any],
|
||||
original_headers: dict[str, str],
|
||||
original_request_body: dict[str, Any],
|
||||
) -> None:
|
||||
"""在流完成后记录统计信息"""
|
||||
try:
|
||||
@@ -1984,7 +1976,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
from src.services.request.candidate import RequestCandidateService
|
||||
|
||||
# 计算候选自身的 TTFB
|
||||
candidate_first_byte_time_ms: Optional[int] = None
|
||||
candidate_first_byte_time_ms: int | None = None
|
||||
if ctx.first_byte_time_ms is not None:
|
||||
candidate_first_byte_time_ms = (
|
||||
RequestCandidateService.calculate_candidate_ttfb(
|
||||
@@ -2049,8 +2041,8 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
self,
|
||||
ctx: StreamContext,
|
||||
error: Exception,
|
||||
original_headers: Dict[str, str],
|
||||
original_request_body: Dict[str, Any],
|
||||
original_headers: dict[str, str],
|
||||
original_request_body: dict[str, Any],
|
||||
) -> None:
|
||||
"""记录流式请求失败"""
|
||||
# 使用 self.start_time 作为时间基准,与首字时间保持一致
|
||||
@@ -2099,10 +2091,10 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
|
||||
async def process_sync(
|
||||
self,
|
||||
original_request_body: Dict[str, Any],
|
||||
original_headers: Dict[str, str],
|
||||
query_params: Optional[Dict[str, str]] = None,
|
||||
path_params: Optional[Dict[str, Any]] = None,
|
||||
original_request_body: dict[str, Any],
|
||||
original_headers: dict[str, str],
|
||||
query_params: dict[str, str] | None = None,
|
||||
path_params: dict[str, Any] | None = None,
|
||||
) -> JSONResponse:
|
||||
"""
|
||||
处理非流式请求
|
||||
@@ -2130,19 +2122,19 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
endpoint_id = None # Endpoint ID(用于失败记录)
|
||||
key_id = None # Key ID(用于失败记录)
|
||||
mapped_model_result = None # 映射后的目标模型名(用于 Usage 记录)
|
||||
response_metadata_result: Dict[str, Any] = {} # Provider 响应元数据
|
||||
response_metadata_result: dict[str, Any] = {} # Provider 响应元数据
|
||||
needs_conversion = False # 是否需要格式转换(由 candidate 决定)
|
||||
|
||||
# 可变请求体容器:允许 orchestrator 在遇到 Thinking 签名错误时整流请求体后重试
|
||||
# 结构: {"body": 实际请求体, "_rectified": 是否已整流, "_rectified_this_turn": 本轮是否整流}
|
||||
request_body_ref: Dict[str, Any] = {"body": original_request_body}
|
||||
request_body_ref: dict[str, Any] = {"body": original_request_body}
|
||||
|
||||
async def sync_request_func(
|
||||
provider: Provider,
|
||||
endpoint: ProviderEndpoint,
|
||||
key: ProviderAPIKey,
|
||||
candidate: ProviderCandidate,
|
||||
) -> Dict[str, Any]:
|
||||
) -> dict[str, Any]:
|
||||
nonlocal provider_name, response_json, status_code, response_headers, provider_api_format, provider_request_headers, provider_request_body, mapped_model_result, response_metadata_result, needs_conversion
|
||||
provider_name = str(provider.name)
|
||||
provider_api_format = str(endpoint.api_format) if endpoint.api_format else ""
|
||||
@@ -2446,7 +2438,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
actual_request_body = provider_request_body or original_request_body
|
||||
|
||||
# 尝试从异常中提取响应头
|
||||
error_response_headers: Dict[str, str] = {}
|
||||
error_response_headers: dict[str, str] = {}
|
||||
if isinstance(e, ProviderRateLimitException) and e.response_headers:
|
||||
error_response_headers = e.response_headers
|
||||
elif isinstance(e, httpx.HTTPStatusError) and hasattr(e, "response"):
|
||||
@@ -2557,7 +2549,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
)
|
||||
return True
|
||||
|
||||
def _mark_first_output(self, ctx: StreamContext, state: Dict[str, bool]) -> None:
|
||||
def _mark_first_output(self, ctx: StreamContext, state: dict[str, bool]) -> None:
|
||||
"""
|
||||
标记首次输出:记录 TTFB 并更新 streaming 状态
|
||||
|
||||
@@ -2581,7 +2573,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
ctx: StreamContext,
|
||||
line: str,
|
||||
events: list, # noqa: ARG002 - 预留给上下文感知转换
|
||||
) -> Tuple[List[str], List[Dict[str, Any]]]:
|
||||
) -> tuple[list[str], list[dict[str, Any]]]:
|
||||
"""
|
||||
将 SSE 行从 Provider 格式转换为客户端格式
|
||||
|
||||
@@ -2666,7 +2658,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
|
||||
def _parse_sse_line_to_json(
|
||||
self, line: str, provider_format: str
|
||||
) -> Tuple[Optional[Any], str]:
|
||||
) -> tuple[Any | None, str]:
|
||||
"""
|
||||
解析 SSE 行为 JSON 对象
|
||||
|
||||
|
||||
@@ -8,7 +8,6 @@ StreamSmoother 使用这些提取器来处理不同格式的 SSE 事件。
|
||||
import copy
|
||||
import json
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Optional
|
||||
|
||||
|
||||
class ContentExtractor(ABC):
|
||||
@@ -20,7 +19,7 @@ class ContentExtractor(ABC):
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def extract_content(self, data: dict) -> Optional[str]:
|
||||
def extract_content(self, data: dict) -> str | None:
|
||||
"""
|
||||
从 SSE 数据中提取可拆分的文本内容
|
||||
|
||||
@@ -64,7 +63,7 @@ class OpenAIContentExtractor(ContentExtractor):
|
||||
- 只在 delta 仅包含 role/content 时允许拆分,避免破坏 tool_calls 等结构
|
||||
"""
|
||||
|
||||
def extract_content(self, data: dict) -> Optional[str]:
|
||||
def extract_content(self, data: dict) -> str | None:
|
||||
if not isinstance(data, dict):
|
||||
return None
|
||||
|
||||
@@ -115,7 +114,7 @@ class OpenAIContentExtractor(ContentExtractor):
|
||||
new_choices.append(new_choice)
|
||||
new_data["choices"] = new_choices
|
||||
|
||||
return f"data: {json.dumps(new_data, ensure_ascii=False)}\n\n".encode("utf-8")
|
||||
return f"data: {json.dumps(new_data, ensure_ascii=False)}\n\n".encode()
|
||||
|
||||
|
||||
class ClaudeContentExtractor(ContentExtractor):
|
||||
@@ -127,7 +126,7 @@ class ClaudeContentExtractor(ContentExtractor):
|
||||
- 数据结构: delta.type=text_delta, delta.text
|
||||
"""
|
||||
|
||||
def extract_content(self, data: dict) -> Optional[str]:
|
||||
def extract_content(self, data: dict) -> str | None:
|
||||
if not isinstance(data, dict):
|
||||
return None
|
||||
|
||||
@@ -165,9 +164,7 @@ class ClaudeContentExtractor(ContentExtractor):
|
||||
|
||||
# Claude 格式需要 event: 前缀
|
||||
event_name = event_type or "content_block_delta"
|
||||
return f"event: {event_name}\ndata: {json.dumps(new_data, ensure_ascii=False)}\n\n".encode(
|
||||
"utf-8"
|
||||
)
|
||||
return f"event: {event_name}\ndata: {json.dumps(new_data, ensure_ascii=False)}\n\n".encode()
|
||||
|
||||
|
||||
class GeminiContentExtractor(ContentExtractor):
|
||||
@@ -179,7 +176,7 @@ class GeminiContentExtractor(ContentExtractor):
|
||||
- 只有纯文本块才拆分
|
||||
"""
|
||||
|
||||
def extract_content(self, data: dict) -> Optional[str]:
|
||||
def extract_content(self, data: dict) -> str | None:
|
||||
if not isinstance(data, dict):
|
||||
return None
|
||||
|
||||
@@ -226,7 +223,7 @@ class GeminiContentExtractor(ContentExtractor):
|
||||
if "parts" in content and content["parts"]:
|
||||
content["parts"][0]["text"] = new_content
|
||||
|
||||
return f"data: {json.dumps(new_data, ensure_ascii=False)}\n\n".encode("utf-8")
|
||||
return f"data: {json.dumps(new_data, ensure_ascii=False)}\n\n".encode()
|
||||
|
||||
|
||||
# 提取器注册表
|
||||
@@ -237,7 +234,7 @@ _EXTRACTORS: dict[str, type[ContentExtractor]] = {
|
||||
}
|
||||
|
||||
|
||||
def get_extractor(format_name: str) -> Optional[ContentExtractor]:
|
||||
def get_extractor(format_name: str) -> ContentExtractor | None:
|
||||
"""
|
||||
根据格式名获取对应的内容提取器实例
|
||||
|
||||
|
||||
@@ -14,15 +14,14 @@
|
||||
- EndpointCheckOrchestrator: 协调整个流程
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, AsyncIterator, Dict, Iterable, Optional, Union, List
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Any
|
||||
from collections.abc import Iterable
|
||||
import time
|
||||
import uuid
|
||||
import json
|
||||
from functools import lru_cache
|
||||
import asyncio
|
||||
from collections import defaultdict
|
||||
|
||||
import httpx
|
||||
|
||||
@@ -31,7 +30,7 @@ from src.core.api_format import CORE_REDACT_HEADERS, merge_headers_with_protecti
|
||||
from src.utils.ssl_utils import get_ssl_context
|
||||
|
||||
|
||||
def _redact_headers(headers: Dict[str, str]) -> Dict[str, str]:
|
||||
def _redact_headers(headers: dict[str, str]) -> dict[str, str]:
|
||||
return redact_headers_for_log(headers, CORE_REDACT_HEADERS)
|
||||
|
||||
|
||||
@@ -46,10 +45,10 @@ def _truncate_repr(value: Any, limit: int = 1200) -> str:
|
||||
|
||||
|
||||
def build_safe_headers(
|
||||
base_headers: Dict[str, str],
|
||||
extra_headers: Optional[Dict[str, str]],
|
||||
base_headers: dict[str, str],
|
||||
extra_headers: dict[str, str] | None,
|
||||
protected_keys: Iterable[str],
|
||||
) -> Dict[str, str]:
|
||||
) -> dict[str, str]:
|
||||
"""
|
||||
合并 extra_headers,但防止覆盖 protected_keys(大小写不敏感)。
|
||||
"""
|
||||
@@ -60,16 +59,16 @@ async def run_endpoint_check(
|
||||
*,
|
||||
client: httpx.AsyncClient, # 保持兼容性,但内部不使用
|
||||
url: str,
|
||||
headers: Dict[str, str],
|
||||
json_body: Dict[str, Any],
|
||||
headers: dict[str, str],
|
||||
json_body: dict[str, Any],
|
||||
api_format: str,
|
||||
provider_name: Optional[str] = None,
|
||||
model_name: Optional[str] = None,
|
||||
api_key_id: Optional[str] = None,
|
||||
provider_id: Optional[str] = None,
|
||||
db: Optional[Any] = None, # Session对象,需要时才导入
|
||||
user: Optional[Any] = None, # User对象
|
||||
) -> Dict[str, Any]:
|
||||
provider_name: str | None = None,
|
||||
model_name: str | None = None,
|
||||
api_key_id: str | None = None,
|
||||
provider_id: str | None = None,
|
||||
db: Any | None = None, # Session对象,需要时才导入
|
||||
user: Any | None = None, # User对象
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
执行端点检查(重构版本,使用新的架构):
|
||||
- 使用新的架构类来分离关注点
|
||||
@@ -123,21 +122,21 @@ async def _calculate_and_record_usage(
|
||||
provider_id: str,
|
||||
api_key_id: str,
|
||||
model_name: str,
|
||||
request_data: Dict[str, Any],
|
||||
response_data: Optional[Dict[str, Any]],
|
||||
request_data: dict[str, Any],
|
||||
response_data: dict[str, Any] | None,
|
||||
request_id: str,
|
||||
response_time_ms: int,
|
||||
request_headers: Dict[str, str],
|
||||
response_headers: Optional[Dict[str, str]] = None,
|
||||
request_headers: dict[str, str],
|
||||
response_headers: dict[str, str] | None = None,
|
||||
status_code: int = 0,
|
||||
error_message: Optional[str] = None,
|
||||
error_message: str | None = None,
|
||||
# 新增:支持直接传递token数据
|
||||
input_tokens: Optional[int] = None,
|
||||
output_tokens: Optional[int] = None,
|
||||
cache_creation_input_tokens: Optional[int] = None,
|
||||
cache_read_input_tokens: Optional[int] = None,
|
||||
api_format: Optional[str] = None,
|
||||
) -> Dict[str, Any]:
|
||||
input_tokens: int | None = None,
|
||||
output_tokens: int | None = None,
|
||||
cache_creation_input_tokens: int | None = None,
|
||||
cache_read_input_tokens: int | None = None,
|
||||
api_format: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
计算并记录用量数据(遗留函数)
|
||||
|
||||
@@ -149,7 +148,7 @@ async def _calculate_and_record_usage(
|
||||
"""
|
||||
from src.services.usage.service import UsageService
|
||||
from src.services.request.candidate import RequestCandidateService
|
||||
from src.models.database import ApiKey, ProviderAPIKey, ProviderEndpoint
|
||||
from src.models.database import ApiKey, ProviderAPIKey
|
||||
|
||||
# 获取Provider API Key对象(不是用户API Key)
|
||||
provider_api_key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == api_key_id).first()
|
||||
@@ -360,7 +359,7 @@ async def _calculate_and_record_usage(
|
||||
}
|
||||
|
||||
|
||||
def _extract_tokens_from_response(api_identifier: str, response_data: Optional[Dict[str, Any]]) -> tuple[int, int, int, int]:
|
||||
def _extract_tokens_from_response(api_identifier: str, response_data: dict[str, Any] | None) -> tuple[int, int, int, int]:
|
||||
"""
|
||||
从响应中提取Token计数信息
|
||||
|
||||
@@ -446,7 +445,7 @@ def _extract_tokens_from_response(api_identifier: str, response_data: Optional[D
|
||||
|
||||
|
||||
|
||||
def _fallback_token_counting(request_data: Dict[str, Any], response_data: Optional[Dict[str, Any]]) -> tuple[int, int, int, int]:
|
||||
def _fallback_token_counting(request_data: dict[str, Any], response_data: dict[str, Any] | None) -> tuple[int, int, int, int]:
|
||||
"""
|
||||
回退的Token计数方法(简单估算)
|
||||
|
||||
@@ -508,16 +507,16 @@ def _fallback_token_counting(request_data: Dict[str, Any], response_data: Option
|
||||
class EndpointCheckRequest:
|
||||
"""端点检查请求数据类"""
|
||||
url: str
|
||||
headers: Dict[str, str]
|
||||
json_body: Dict[str, Any]
|
||||
headers: dict[str, str]
|
||||
json_body: dict[str, Any]
|
||||
api_format: str
|
||||
provider_name: Optional[str] = None
|
||||
model_name: Optional[str] = None
|
||||
api_key_id: Optional[str] = None
|
||||
provider_id: Optional[str] = None
|
||||
db: Optional[Any] = None
|
||||
user: Optional[Any] = None
|
||||
request_id: Optional[str] = None
|
||||
provider_name: str | None = None
|
||||
model_name: str | None = None
|
||||
api_key_id: str | None = None
|
||||
provider_id: str | None = None
|
||||
db: Any | None = None
|
||||
user: Any | None = None
|
||||
request_id: str | None = None
|
||||
timeout: float = 30.0
|
||||
|
||||
|
||||
@@ -525,12 +524,12 @@ class EndpointCheckRequest:
|
||||
class EndpointCheckResult:
|
||||
"""端点检查结果数据类"""
|
||||
status_code: int
|
||||
headers: Dict[str, str]
|
||||
headers: dict[str, str]
|
||||
response_time_ms: int
|
||||
request_id: str
|
||||
response_data: Optional[Dict[str, Any]] = None
|
||||
error_message: Optional[str] = None
|
||||
usage_data: Optional[Dict[str, Any]] = None
|
||||
response_data: dict[str, Any] | None = None
|
||||
error_message: str | None = None
|
||||
usage_data: dict[str, Any] | None = None
|
||||
|
||||
|
||||
class HttpRequestExecutor:
|
||||
@@ -613,7 +612,7 @@ class UsageCalculator:
|
||||
return _extract_tokens_from_response(api_identifier, result.response_data)
|
||||
|
||||
@staticmethod
|
||||
def _fallback_token_counting(request_data: Dict[str, Any], response_data: Optional[Dict[str, Any]]) -> tuple[int, int, int, int]:
|
||||
def _fallback_token_counting(request_data: dict[str, Any], response_data: dict[str, Any] | None) -> tuple[int, int, int, int]:
|
||||
"""回退的Token计数方法(简单估算)"""
|
||||
# 估算输入Token
|
||||
messages = request_data.get("messages", request_data.get("contents", []))
|
||||
@@ -665,12 +664,12 @@ class AsyncBatchUsageRecorder:
|
||||
def __init__(self, batch_size: int = 10, flush_interval: float = 2.0):
|
||||
self.batch_size = batch_size
|
||||
self.flush_interval = flush_interval
|
||||
self.pending_records: List[Dict[str, Any]] = []
|
||||
self._flush_task: Optional[asyncio.Task] = None
|
||||
self.pending_records: list[dict[str, Any]] = []
|
||||
self._flush_task: asyncio.Task | None = None
|
||||
self._lock = asyncio.Lock()
|
||||
self._running = True
|
||||
|
||||
async def add_record(self, usage_data: Dict[str, Any]) -> None:
|
||||
async def add_record(self, usage_data: dict[str, Any]) -> None:
|
||||
"""添加用量记录到批处理队列"""
|
||||
async with self._lock:
|
||||
self.pending_records.append(usage_data)
|
||||
@@ -740,7 +739,7 @@ class AsyncBatchUsageRecorder:
|
||||
|
||||
|
||||
# 全局批处理器实例(单例)
|
||||
_global_batch_recorder: Optional[AsyncBatchUsageRecorder] = None
|
||||
_global_batch_recorder: AsyncBatchUsageRecorder | None = None
|
||||
|
||||
def get_batch_recorder() -> AsyncBatchUsageRecorder:
|
||||
"""获取全局批处理器实例"""
|
||||
@@ -756,7 +755,7 @@ def get_batch_recorder() -> AsyncBatchUsageRecorder:
|
||||
|
||||
class EndpointCheckError(Exception):
|
||||
"""端点检查错误基类"""
|
||||
def __init__(self, message: str, error_type: str, status_code: int = 500, details: Optional[Dict[str, Any]] = None):
|
||||
def __init__(self, message: str, error_type: str, status_code: int = 500, details: dict[str, Any] | None = None):
|
||||
super().__init__(message)
|
||||
self.message = message
|
||||
self.error_type = error_type
|
||||
@@ -765,22 +764,22 @@ class EndpointCheckError(Exception):
|
||||
|
||||
class NetworkError(EndpointCheckError):
|
||||
"""网络请求错误"""
|
||||
def __init__(self, message: str, details: Optional[Dict[str, Any]] = None):
|
||||
def __init__(self, message: str, details: dict[str, Any] | None = None):
|
||||
super().__init__(message, "network_error", 0, details)
|
||||
|
||||
class AuthenticationError(EndpointCheckError):
|
||||
"""认证错误"""
|
||||
def __init__(self, message: str, details: Optional[Dict[str, Any]] = None):
|
||||
def __init__(self, message: str, details: dict[str, Any] | None = None):
|
||||
super().__init__(message, "authentication_error", 401, details)
|
||||
|
||||
class RateLimitError(EndpointCheckError):
|
||||
"""速率限制错误"""
|
||||
def __init__(self, message: str, details: Optional[Dict[str, Any]] = None):
|
||||
def __init__(self, message: str, details: dict[str, Any] | None = None):
|
||||
super().__init__(message, "rate_limit_error", 429, details)
|
||||
|
||||
class UpstreamError(EndpointCheckError):
|
||||
"""上游服务错误"""
|
||||
def __init__(self, message: str, status_code: int, details: Optional[Dict[str, Any]] = None):
|
||||
def __init__(self, message: str, status_code: int, details: dict[str, Any] | None = None):
|
||||
super().__init__(message, "upstream_error", status_code, details)
|
||||
|
||||
|
||||
@@ -982,7 +981,7 @@ class EndpointCheckConfig:
|
||||
retry_on_timeouts: bool = True
|
||||
|
||||
@classmethod
|
||||
def from_env(cls) -> 'EndpointCheckConfig':
|
||||
def from_env(cls) -> EndpointCheckConfig:
|
||||
"""从环境变量创建配置"""
|
||||
import os
|
||||
|
||||
@@ -1004,7 +1003,7 @@ class EndpointCheckConfig:
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, config_dict: Dict[str, Any]) -> 'EndpointCheckConfig':
|
||||
def from_dict(cls, config_dict: dict[str, Any]) -> EndpointCheckConfig:
|
||||
"""从字典创建配置"""
|
||||
return cls(**{k: v for k, v in config_dict.items() if hasattr(cls, k)})
|
||||
|
||||
@@ -1012,7 +1011,7 @@ class EndpointCheckConfig:
|
||||
class ConfigurableEndpointChecker:
|
||||
"""可配置的端点检查器"""
|
||||
|
||||
def __init__(self, config: Optional[EndpointCheckConfig] = None):
|
||||
def __init__(self, config: EndpointCheckConfig | None = None):
|
||||
self.config = config or EndpointCheckConfig()
|
||||
self.executor = HttpRequestExecutor(timeout=self.config.timeout)
|
||||
self.usage_calculator = UsageCalculator()
|
||||
@@ -1171,9 +1170,9 @@ class ConfigurableEndpointChecker:
|
||||
|
||||
|
||||
# 全局配置检查器实例
|
||||
_global_configured_checker: Optional[ConfigurableEndpointChecker] = None
|
||||
_global_configured_checker: ConfigurableEndpointChecker | None = None
|
||||
|
||||
def get_configured_checker(config: Optional[EndpointCheckConfig] = None) -> ConfigurableEndpointChecker:
|
||||
def get_configured_checker(config: EndpointCheckConfig | None = None) -> ConfigurableEndpointChecker:
|
||||
"""获取全局配置检查器实例"""
|
||||
global _global_configured_checker
|
||||
if _global_configured_checker is None or config is not None:
|
||||
@@ -1186,8 +1185,8 @@ def get_configured_checker(config: Optional[EndpointCheckConfig] = None) -> Conf
|
||||
class EndpointCheckOrchestrator:
|
||||
"""端点检查协调器 - 协调整个流程"""
|
||||
|
||||
def __init__(self, executor: Optional[HttpRequestExecutor] = None,
|
||||
usage_calculator: Optional[UsageCalculator] = None):
|
||||
def __init__(self, executor: HttpRequestExecutor | None = None,
|
||||
usage_calculator: UsageCalculator | None = None):
|
||||
self.executor = executor or HttpRequestExecutor()
|
||||
self.usage_calculator = usage_calculator or UsageCalculator()
|
||||
|
||||
|
||||
@@ -6,7 +6,7 @@
|
||||
"""
|
||||
|
||||
import re
|
||||
from typing import Any, Dict, Optional, Tuple, Type
|
||||
from typing import Any
|
||||
|
||||
from src.api.handlers.base.response_parser import (
|
||||
ParsedChunk,
|
||||
@@ -20,7 +20,7 @@ from src.api.handlers.base.utils import extract_cache_creation_tokens
|
||||
from src.core.api_format import is_cli_format
|
||||
|
||||
|
||||
def _check_nested_error(response: Dict[str, Any]) -> Tuple[bool, Optional[Dict[str, Any]]]:
|
||||
def _check_nested_error(response: dict[str, Any]) -> tuple[bool, dict[str, Any] | None]:
|
||||
"""
|
||||
检查响应中是否存在嵌套错误(某些代理服务返回 HTTP 200 但在响应体中包含错误)
|
||||
|
||||
@@ -62,7 +62,7 @@ def _check_nested_error(response: Dict[str, Any]) -> Tuple[bool, Optional[Dict[s
|
||||
return False, None
|
||||
|
||||
|
||||
def _extract_embedded_status_code(error_info: Optional[Dict[str, Any]]) -> Optional[int]:
|
||||
def _extract_embedded_status_code(error_info: dict[str, Any] | None) -> int | None:
|
||||
"""
|
||||
从错误信息中提取嵌套的状态码
|
||||
|
||||
@@ -137,7 +137,7 @@ class OpenAIResponseParser(ResponseParser):
|
||||
self.name = "OPENAI"
|
||||
self.api_format = "OPENAI"
|
||||
|
||||
def parse_sse_line(self, line: str, stats: StreamStats) -> Optional[ParsedChunk]:
|
||||
def parse_sse_line(self, line: str, stats: StreamStats) -> ParsedChunk | None:
|
||||
if not line or not line.strip():
|
||||
return None
|
||||
|
||||
@@ -186,7 +186,7 @@ class OpenAIResponseParser(ResponseParser):
|
||||
|
||||
return chunk
|
||||
|
||||
def parse_response(self, response: Dict[str, Any], status_code: int) -> ParsedResponse:
|
||||
def parse_response(self, response: dict[str, Any], status_code: int) -> ParsedResponse:
|
||||
result = ParsedResponse(
|
||||
raw_response=response,
|
||||
status_code=status_code,
|
||||
@@ -217,7 +217,7 @@ class OpenAIResponseParser(ResponseParser):
|
||||
|
||||
return result
|
||||
|
||||
def extract_usage_from_response(self, response: Dict[str, Any]) -> Dict[str, int]:
|
||||
def extract_usage_from_response(self, response: dict[str, Any]) -> dict[str, int]:
|
||||
usage = response.get("usage") or {}
|
||||
return {
|
||||
"input_tokens": usage.get("prompt_tokens", 0),
|
||||
@@ -226,7 +226,7 @@ class OpenAIResponseParser(ResponseParser):
|
||||
"cache_read_tokens": 0,
|
||||
}
|
||||
|
||||
def extract_text_content(self, response: Dict[str, Any]) -> str:
|
||||
def extract_text_content(self, response: dict[str, Any]) -> str:
|
||||
choices = response.get("choices", [])
|
||||
if choices:
|
||||
message = choices[0].get("message", {})
|
||||
@@ -235,7 +235,7 @@ class OpenAIResponseParser(ResponseParser):
|
||||
return content
|
||||
return ""
|
||||
|
||||
def is_error_response(self, response: Dict[str, Any]) -> bool:
|
||||
def is_error_response(self, response: dict[str, Any]) -> bool:
|
||||
is_error, _ = _check_nested_error(response)
|
||||
return is_error
|
||||
|
||||
@@ -259,7 +259,7 @@ class ClaudeResponseParser(ResponseParser):
|
||||
self.name = "CLAUDE"
|
||||
self.api_format = "CLAUDE"
|
||||
|
||||
def parse_sse_line(self, line: str, stats: StreamStats) -> Optional[ParsedChunk]:
|
||||
def parse_sse_line(self, line: str, stats: StreamStats) -> ParsedChunk | None:
|
||||
if not line or not line.strip():
|
||||
return None
|
||||
|
||||
@@ -324,7 +324,7 @@ class ClaudeResponseParser(ResponseParser):
|
||||
|
||||
return chunk
|
||||
|
||||
def parse_response(self, response: Dict[str, Any], status_code: int) -> ParsedResponse:
|
||||
def parse_response(self, response: dict[str, Any], status_code: int) -> ParsedResponse:
|
||||
result = ParsedResponse(
|
||||
raw_response=response,
|
||||
status_code=status_code,
|
||||
@@ -358,7 +358,7 @@ class ClaudeResponseParser(ResponseParser):
|
||||
|
||||
return result
|
||||
|
||||
def extract_usage_from_response(self, response: Dict[str, Any]) -> Dict[str, int]:
|
||||
def extract_usage_from_response(self, response: dict[str, Any]) -> dict[str, int]:
|
||||
# 对于 message_start 事件,usage 在 message.usage 路径下
|
||||
# 对于其他响应,usage 在顶层
|
||||
usage = response.get("usage") or {}
|
||||
@@ -372,7 +372,7 @@ class ClaudeResponseParser(ResponseParser):
|
||||
"cache_read_tokens": usage.get("cache_read_input_tokens", 0),
|
||||
}
|
||||
|
||||
def extract_text_content(self, response: Dict[str, Any]) -> str:
|
||||
def extract_text_content(self, response: dict[str, Any]) -> str:
|
||||
content = response.get("content", [])
|
||||
if isinstance(content, list):
|
||||
text_parts = []
|
||||
@@ -382,7 +382,7 @@ class ClaudeResponseParser(ResponseParser):
|
||||
return "".join(text_parts)
|
||||
return ""
|
||||
|
||||
def is_error_response(self, response: Dict[str, Any]) -> bool:
|
||||
def is_error_response(self, response: dict[str, Any]) -> bool:
|
||||
is_error, _ = _check_nested_error(response)
|
||||
return is_error
|
||||
|
||||
@@ -406,7 +406,7 @@ class GeminiResponseParser(ResponseParser):
|
||||
self.name = "GEMINI"
|
||||
self.api_format = "GEMINI"
|
||||
|
||||
def parse_sse_line(self, line: str, stats: StreamStats) -> Optional[ParsedChunk]:
|
||||
def parse_sse_line(self, line: str, stats: StreamStats) -> ParsedChunk | None:
|
||||
"""
|
||||
解析 Gemini SSE 行
|
||||
|
||||
@@ -473,7 +473,7 @@ class GeminiResponseParser(ResponseParser):
|
||||
|
||||
return chunk
|
||||
|
||||
def parse_response(self, response: Dict[str, Any], status_code: int) -> ParsedResponse:
|
||||
def parse_response(self, response: dict[str, Any], status_code: int) -> ParsedResponse:
|
||||
result = ParsedResponse(
|
||||
raw_response=response,
|
||||
status_code=status_code,
|
||||
@@ -509,7 +509,7 @@ class GeminiResponseParser(ResponseParser):
|
||||
|
||||
return result
|
||||
|
||||
def extract_usage_from_response(self, response: Dict[str, Any]) -> Dict[str, int]:
|
||||
def extract_usage_from_response(self, response: dict[str, Any]) -> dict[str, int]:
|
||||
"""
|
||||
从 Gemini 响应中提取 token 使用量
|
||||
|
||||
@@ -531,7 +531,7 @@ class GeminiResponseParser(ResponseParser):
|
||||
"cache_read_tokens": usage.get("cached_tokens", 0),
|
||||
}
|
||||
|
||||
def extract_text_content(self, response: Dict[str, Any]) -> str:
|
||||
def extract_text_content(self, response: dict[str, Any]) -> str:
|
||||
candidates = response.get("candidates", [])
|
||||
if candidates:
|
||||
content = candidates[0].get("content", {})
|
||||
@@ -543,7 +543,7 @@ class GeminiResponseParser(ResponseParser):
|
||||
return "".join(text_parts)
|
||||
return ""
|
||||
|
||||
def is_error_response(self, response: Dict[str, Any]) -> bool:
|
||||
def is_error_response(self, response: dict[str, Any]) -> bool:
|
||||
"""
|
||||
判断响应是否为错误响应
|
||||
|
||||
@@ -562,7 +562,7 @@ class GeminiCliResponseParser(GeminiResponseParser):
|
||||
|
||||
|
||||
# 解析器注册表
|
||||
_PARSERS: Dict[str, Type[ResponseParser]] = {
|
||||
_PARSERS: dict[str, type[ResponseParser]] = {
|
||||
"CLAUDE": ClaudeResponseParser,
|
||||
"CLAUDE_CLI": ClaudeCliResponseParser,
|
||||
"OPENAI": OpenAIResponseParser,
|
||||
|
||||
@@ -11,10 +11,9 @@
|
||||
payload, headers = builder.build(original_body, original_headers, endpoint, key)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Any, Dict, FrozenSet, Optional, Tuple
|
||||
from typing import Any
|
||||
|
||||
from src.core.crypto import crypto_service
|
||||
from src.core.api_format import HeaderBuilder, UPSTREAM_DROP_HEADERS
|
||||
@@ -24,7 +23,7 @@ from src.core.api_format import HeaderBuilder, UPSTREAM_DROP_HEADERS
|
||||
# ==============================================================================
|
||||
|
||||
# 兼容别名:历史代码使用 SENSITIVE_HEADERS 命名
|
||||
SENSITIVE_HEADERS: FrozenSet[str] = UPSTREAM_DROP_HEADERS
|
||||
SENSITIVE_HEADERS: frozenset[str] = UPSTREAM_DROP_HEADERS
|
||||
|
||||
|
||||
# ==============================================================================
|
||||
@@ -33,14 +32,14 @@ SENSITIVE_HEADERS: FrozenSet[str] = UPSTREAM_DROP_HEADERS
|
||||
|
||||
# 标准测试请求体(OpenAI 格式)
|
||||
# 用于 check_endpoint 等测试场景,使用简单安全的消息内容避免触发安全过滤
|
||||
DEFAULT_TEST_REQUEST: Dict[str, Any] = {
|
||||
DEFAULT_TEST_REQUEST: dict[str, Any] = {
|
||||
"messages": [{"role": "user", "content": "Hi"}],
|
||||
"max_tokens": 5,
|
||||
"temperature": 0,
|
||||
}
|
||||
|
||||
|
||||
def get_test_request_data(request_data: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
|
||||
def get_test_request_data(request_data: dict[str, Any] | None = None) -> dict[str, Any]:
|
||||
"""获取测试请求数据
|
||||
|
||||
如果传入 request_data,则合并到默认测试请求中;
|
||||
@@ -61,8 +60,8 @@ def get_test_request_data(request_data: Optional[Dict[str, Any]] = None) -> Dict
|
||||
|
||||
def build_test_request_body(
|
||||
format_id: str,
|
||||
request_data: Optional[Dict[str, Any]] = None,
|
||||
) -> Dict[str, Any]:
|
||||
request_data: dict[str, Any] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""构建测试请求体,自动处理格式转换
|
||||
|
||||
使用格式转换注册表将 OpenAI 格式的测试请求转换为目标格式。
|
||||
@@ -103,37 +102,37 @@ class RequestBuilder(ABC):
|
||||
@abstractmethod
|
||||
def build_payload(
|
||||
self,
|
||||
original_body: Dict[str, Any],
|
||||
original_body: dict[str, Any],
|
||||
*,
|
||||
mapped_model: Optional[str] = None,
|
||||
mapped_model: str | None = None,
|
||||
is_stream: bool = False,
|
||||
) -> Dict[str, Any]:
|
||||
) -> dict[str, Any]:
|
||||
"""构建请求体"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def build_headers(
|
||||
self,
|
||||
original_headers: Dict[str, str],
|
||||
original_headers: dict[str, str],
|
||||
endpoint: Any,
|
||||
key: Any,
|
||||
*,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
) -> Dict[str, str]:
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
) -> dict[str, str]:
|
||||
"""构建请求头"""
|
||||
pass
|
||||
|
||||
def build(
|
||||
self,
|
||||
original_body: Dict[str, Any],
|
||||
original_headers: Dict[str, str],
|
||||
original_body: dict[str, Any],
|
||||
original_headers: dict[str, str],
|
||||
endpoint: Any,
|
||||
key: Any,
|
||||
*,
|
||||
mapped_model: Optional[str] = None,
|
||||
mapped_model: str | None = None,
|
||||
is_stream: bool = False,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
) -> Tuple[Dict[str, Any], Dict[str, str]]:
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
) -> tuple[dict[str, Any], dict[str, str]]:
|
||||
"""
|
||||
构建完整的请求(请求体 + 请求头)
|
||||
|
||||
@@ -174,11 +173,11 @@ class PassthroughRequestBuilder(RequestBuilder):
|
||||
|
||||
def build_payload(
|
||||
self,
|
||||
original_body: Dict[str, Any],
|
||||
original_body: dict[str, Any],
|
||||
*,
|
||||
mapped_model: Optional[str] = None, # noqa: ARG002 - 由 apply_mapped_model 处理
|
||||
mapped_model: str | None = None, # noqa: ARG002 - 由 apply_mapped_model 处理
|
||||
is_stream: bool = False, # noqa: ARG002 - 保留原始值,不自动添加
|
||||
) -> Dict[str, Any]:
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
透传请求体 - 原样复制,不做任何修改
|
||||
|
||||
@@ -190,12 +189,12 @@ class PassthroughRequestBuilder(RequestBuilder):
|
||||
|
||||
def build_headers(
|
||||
self,
|
||||
original_headers: Dict[str, str],
|
||||
original_headers: dict[str, str],
|
||||
endpoint: Any,
|
||||
key: Any,
|
||||
*,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
) -> Dict[str, str]:
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
) -> dict[str, str]:
|
||||
"""
|
||||
透传请求头 - 清理敏感头部(黑名单),透传其他所有头部
|
||||
|
||||
@@ -254,11 +253,11 @@ class PassthroughRequestBuilder(RequestBuilder):
|
||||
|
||||
|
||||
def build_passthrough_request(
|
||||
original_body: Dict[str, Any],
|
||||
original_headers: Dict[str, str],
|
||||
original_body: dict[str, Any],
|
||||
original_headers: dict[str, str],
|
||||
endpoint: Any,
|
||||
key: Any,
|
||||
) -> Tuple[Dict[str, Any], Dict[str, str]]:
|
||||
) -> tuple[dict[str, Any], dict[str, str]]:
|
||||
"""
|
||||
构建透传模式的请求
|
||||
|
||||
|
||||
@@ -4,7 +4,7 @@
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import Any
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -13,14 +13,14 @@ class ParsedChunk:
|
||||
|
||||
# 原始数据
|
||||
raw_line: str
|
||||
event_type: Optional[str] = None
|
||||
data: Optional[Dict[str, Any]] = None
|
||||
event_type: str | None = None
|
||||
data: dict[str, Any] | None = None
|
||||
|
||||
# 提取的内容
|
||||
text_delta: str = ""
|
||||
is_done: bool = False
|
||||
is_error: bool = False
|
||||
error_message: Optional[str] = None
|
||||
error_message: str | None = None
|
||||
|
||||
# 使用量信息(通常在最后一个 chunk 中)
|
||||
input_tokens: int = 0
|
||||
@@ -29,7 +29,7 @@ class ParsedChunk:
|
||||
cache_read_tokens: int = 0
|
||||
|
||||
# 响应 ID
|
||||
response_id: Optional[str] = None
|
||||
response_id: str | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -48,21 +48,21 @@ class StreamStats:
|
||||
|
||||
# 内容
|
||||
collected_text: str = ""
|
||||
response_id: Optional[str] = None
|
||||
response_id: str | None = None
|
||||
|
||||
# 状态
|
||||
has_completion: bool = False
|
||||
status_code: int = 200
|
||||
error_message: Optional[str] = None
|
||||
error_message: str | None = None
|
||||
|
||||
# Provider 信息
|
||||
provider_name: Optional[str] = None
|
||||
endpoint_id: Optional[str] = None
|
||||
key_id: Optional[str] = None
|
||||
provider_name: str | None = None
|
||||
endpoint_id: str | None = None
|
||||
key_id: str | None = None
|
||||
|
||||
# 响应头和完整响应
|
||||
response_headers: Dict[str, str] = field(default_factory=dict)
|
||||
final_response: Optional[Dict[str, Any]] = None
|
||||
response_headers: dict[str, str] = field(default_factory=dict)
|
||||
final_response: dict[str, Any] | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -70,12 +70,12 @@ class ParsedResponse:
|
||||
"""解析后的非流式响应"""
|
||||
|
||||
# 原始响应
|
||||
raw_response: Dict[str, Any]
|
||||
raw_response: dict[str, Any]
|
||||
status_code: int
|
||||
|
||||
# 提取的内容
|
||||
text_content: str = ""
|
||||
response_id: Optional[str] = None
|
||||
response_id: str | None = None
|
||||
|
||||
# 使用量
|
||||
input_tokens: int = 0
|
||||
@@ -85,10 +85,10 @@ class ParsedResponse:
|
||||
|
||||
# 错误信息
|
||||
is_error: bool = False
|
||||
error_type: Optional[str] = None
|
||||
error_message: Optional[str] = None
|
||||
error_type: str | None = None
|
||||
error_message: str | None = None
|
||||
# 从响应体解析出的嵌套状态码(当 HTTP 200 但响应体含错误时使用)
|
||||
embedded_status_code: Optional[int] = None
|
||||
embedded_status_code: int | None = None
|
||||
|
||||
|
||||
class ResponseParser(ABC):
|
||||
@@ -106,7 +106,7 @@ class ResponseParser(ABC):
|
||||
api_format: str = "UNKNOWN"
|
||||
|
||||
@abstractmethod
|
||||
def parse_sse_line(self, line: str, stats: StreamStats) -> Optional[ParsedChunk]:
|
||||
def parse_sse_line(self, line: str, stats: StreamStats) -> ParsedChunk | None:
|
||||
"""
|
||||
解析单行 SSE 数据
|
||||
|
||||
@@ -120,7 +120,7 @@ class ResponseParser(ABC):
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def parse_response(self, response: Dict[str, Any], status_code: int) -> ParsedResponse:
|
||||
def parse_response(self, response: dict[str, Any], status_code: int) -> ParsedResponse:
|
||||
"""
|
||||
解析非流式响应
|
||||
|
||||
@@ -134,7 +134,7 @@ class ResponseParser(ABC):
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def extract_usage_from_response(self, response: Dict[str, Any]) -> Dict[str, int]:
|
||||
def extract_usage_from_response(self, response: dict[str, Any]) -> dict[str, int]:
|
||||
"""
|
||||
从响应中提取 token 使用量
|
||||
|
||||
@@ -147,7 +147,7 @@ class ResponseParser(ABC):
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def extract_text_content(self, response: Dict[str, Any]) -> str:
|
||||
def extract_text_content(self, response: dict[str, Any]) -> str:
|
||||
"""
|
||||
从响应中提取文本内容
|
||||
|
||||
@@ -159,7 +159,7 @@ class ResponseParser(ABC):
|
||||
"""
|
||||
pass
|
||||
|
||||
def is_error_response(self, response: Dict[str, Any]) -> bool:
|
||||
def is_error_response(self, response: dict[str, Any]) -> bool:
|
||||
"""
|
||||
判断响应是否为错误响应
|
||||
|
||||
|
||||
@@ -10,7 +10,7 @@
|
||||
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.core.api_format.conversion.stream_state import StreamState
|
||||
@@ -35,16 +35,16 @@ class StreamContext:
|
||||
api_key_id: int = 0
|
||||
|
||||
# Provider 信息(在请求执行时填充)
|
||||
provider_name: Optional[str] = None
|
||||
provider_id: Optional[str] = None
|
||||
endpoint_id: Optional[str] = None
|
||||
key_id: Optional[str] = None
|
||||
attempt_id: Optional[str] = None
|
||||
provider_name: str | None = None
|
||||
provider_id: str | None = None
|
||||
endpoint_id: str | None = None
|
||||
key_id: str | None = None
|
||||
attempt_id: str | None = None
|
||||
attempt_synced: bool = False
|
||||
provider_api_format: Optional[str] = None # Provider 的响应格式
|
||||
provider_api_format: str | None = None # Provider 的响应格式
|
||||
|
||||
# 模型映射
|
||||
mapped_model: Optional[str] = None
|
||||
mapped_model: str | None = None
|
||||
|
||||
# Token 统计
|
||||
input_tokens: int = 0
|
||||
@@ -53,33 +53,33 @@ class StreamContext:
|
||||
cache_creation_tokens: int = 0
|
||||
|
||||
# 响应内容
|
||||
_collected_text_parts: List[str] = field(default_factory=list, repr=False)
|
||||
response_id: Optional[str] = None
|
||||
final_usage: Optional[Dict[str, Any]] = None
|
||||
final_response: Optional[Dict[str, Any]] = None
|
||||
_collected_text_parts: list[str] = field(default_factory=list, repr=False)
|
||||
response_id: str | None = None
|
||||
final_usage: dict[str, Any] | None = None
|
||||
final_response: dict[str, Any] | None = None
|
||||
|
||||
# 时间指标
|
||||
first_byte_time_ms: Optional[int] = None # 首字时间 (TTFB - Time To First Byte)
|
||||
first_byte_time_ms: int | None = None # 首字时间 (TTFB - Time To First Byte)
|
||||
start_time: float = field(default_factory=time.time)
|
||||
|
||||
# 响应状态
|
||||
status_code: int = 200
|
||||
error_message: Optional[str] = None # 客户端友好的错误消息
|
||||
upstream_response: Optional[str] = None # 原始 Provider 响应(用于请求链路追踪)
|
||||
error_message: str | None = None # 客户端友好的错误消息
|
||||
upstream_response: str | None = None # 原始 Provider 响应(用于请求链路追踪)
|
||||
has_completion: bool = False
|
||||
|
||||
# 请求/响应数据
|
||||
response_headers: Dict[str, str] = field(default_factory=dict) # 提供商响应头
|
||||
client_response_headers: Dict[str, str] = field(default_factory=dict) # 返回给客户端的响应头
|
||||
provider_request_headers: Dict[str, str] = field(default_factory=dict)
|
||||
provider_request_body: Optional[Dict[str, Any]] = None
|
||||
response_headers: dict[str, str] = field(default_factory=dict) # 提供商响应头
|
||||
client_response_headers: dict[str, str] = field(default_factory=dict) # 返回给客户端的响应头
|
||||
provider_request_headers: dict[str, str] = field(default_factory=dict)
|
||||
provider_request_body: dict[str, Any] | None = None
|
||||
|
||||
# 格式转换信息(CLI handler 需要)
|
||||
client_api_format: str = ""
|
||||
needs_conversion: bool = False # 是否需要跨格式转换(由 handler 层设置)
|
||||
|
||||
# Provider 响应元数据(CLI handler 需要)
|
||||
response_metadata: Dict[str, Any] = field(default_factory=dict)
|
||||
response_metadata: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
# 整流标记(Thinking Rectifier)
|
||||
rectified: bool = False # 请求是否经过整流(移除 thinking 块后重试)
|
||||
@@ -87,10 +87,10 @@ class StreamContext:
|
||||
# 流式处理统计
|
||||
data_count: int = 0
|
||||
chunk_count: int = 0
|
||||
parsed_chunks: List[Dict[str, Any]] = field(default_factory=list)
|
||||
parsed_chunks: list[dict[str, Any]] = field(default_factory=list)
|
||||
|
||||
# 流式格式转换状态(跨 chunk 追踪)
|
||||
stream_conversion_state: Optional["StreamState"] = None
|
||||
stream_conversion_state: StreamState | None = None
|
||||
|
||||
def reset_for_retry(self) -> None:
|
||||
"""
|
||||
@@ -138,7 +138,7 @@ class StreamContext:
|
||||
provider_id: str,
|
||||
endpoint_id: str,
|
||||
key_id: str,
|
||||
provider_api_format: Optional[str] = None,
|
||||
provider_api_format: str | None = None,
|
||||
) -> None:
|
||||
"""更新 Provider 信息"""
|
||||
self.provider_name = provider_name
|
||||
@@ -149,10 +149,10 @@ class StreamContext:
|
||||
|
||||
def update_usage(
|
||||
self,
|
||||
input_tokens: Optional[int] = None,
|
||||
output_tokens: Optional[int] = None,
|
||||
cached_tokens: Optional[int] = None,
|
||||
cache_creation_tokens: Optional[int] = None,
|
||||
input_tokens: int | None = None,
|
||||
output_tokens: int | None = None,
|
||||
cached_tokens: int | None = None,
|
||||
cache_creation_tokens: int | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
更新 Token 使用统计
|
||||
@@ -194,7 +194,7 @@ class StreamContext:
|
||||
self,
|
||||
status_code: int,
|
||||
error_message: str,
|
||||
upstream_response: Optional[str] = None,
|
||||
upstream_response: str | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
标记请求失败
|
||||
@@ -230,7 +230,7 @@ class StreamContext:
|
||||
"""检查是否因客户端断开连接而结束"""
|
||||
return self.status_code == 499
|
||||
|
||||
def build_response_body(self, response_time_ms: int) -> Dict[str, Any]:
|
||||
def build_response_body(self, response_time_ms: int) -> dict[str, Any]:
|
||||
"""
|
||||
构建响应体元数据
|
||||
|
||||
|
||||
@@ -13,7 +13,10 @@ import asyncio
|
||||
import codecs
|
||||
import json
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, AsyncGenerator, Callable, Optional
|
||||
from typing import Any
|
||||
|
||||
from collections.abc import Callable
|
||||
from collections.abc import AsyncGenerator
|
||||
|
||||
import httpx
|
||||
|
||||
@@ -65,10 +68,10 @@ class StreamProcessor:
|
||||
self,
|
||||
request_id: str,
|
||||
default_parser: ResponseParser,
|
||||
on_streaming_start: Optional[Callable[[], None]] = None,
|
||||
on_streaming_start: Callable[[], None] | None = None,
|
||||
*,
|
||||
collect_text: bool = False,
|
||||
smoothing_config: Optional[StreamSmoothingConfig] = None,
|
||||
smoothing_config: StreamSmoothingConfig | None = None,
|
||||
):
|
||||
"""
|
||||
初始化流处理器
|
||||
@@ -105,7 +108,7 @@ class StreamProcessor:
|
||||
def handle_sse_event(
|
||||
self,
|
||||
ctx: StreamContext,
|
||||
event_name: Optional[str],
|
||||
event_name: str | None,
|
||||
data_str: str,
|
||||
*,
|
||||
skip_record: bool = False,
|
||||
@@ -363,7 +366,7 @@ class StreamProcessor:
|
||||
):
|
||||
# 重新抛出可重试的 Provider 异常,触发故障转移
|
||||
raise
|
||||
except (OSError, IOError) as e:
|
||||
except OSError as e:
|
||||
# 网络 I/O 异常:记录警告,可能需要重试
|
||||
logger.warning(f" [{self.request_id}] 预读流时发生网络异常: {type(e).__name__}: {e}")
|
||||
except Exception as e:
|
||||
@@ -382,10 +385,10 @@ class StreamProcessor:
|
||||
byte_iterator: Any,
|
||||
response_ctx: Any,
|
||||
http_client: httpx.AsyncClient,
|
||||
prefetched_chunks: Optional[list] = None,
|
||||
prefetched_chunks: list | None = None,
|
||||
*,
|
||||
start_time: Optional[float] = None,
|
||||
) -> AsyncGenerator[bytes, None]:
|
||||
start_time: float | None = None,
|
||||
) -> AsyncGenerator[bytes]:
|
||||
"""
|
||||
创建响应流生成器
|
||||
|
||||
@@ -547,9 +550,7 @@ class StreamProcessor:
|
||||
logger.warning(f"[{self.request_id}] 流式格式转换失败: {conv_err}")
|
||||
payload = _build_stream_error_payload("响应格式转换失败,请稍后重试")
|
||||
error_bytes = (
|
||||
f"data: {json.dumps(payload, ensure_ascii=False)}\n\n".encode(
|
||||
"utf-8"
|
||||
)
|
||||
f"data: {json.dumps(payload, ensure_ascii=False)}\n\n".encode()
|
||||
)
|
||||
done_bytes = (
|
||||
b"data: [DONE]\n\n" if client_format.startswith("OPENAI") else b""
|
||||
@@ -570,7 +571,7 @@ class StreamProcessor:
|
||||
# 统一使用 SSE 格式输出(Gemini streamGenerateContent 也使用 SSE)
|
||||
# 参考: https://ai.google.dev/api/generate-content
|
||||
out.append(
|
||||
f"data: {json.dumps(evt, ensure_ascii=False)}\n\n".encode("utf-8")
|
||||
f"data: {json.dumps(evt, ensure_ascii=False)}\n\n".encode()
|
||||
)
|
||||
return out
|
||||
|
||||
@@ -769,9 +770,9 @@ class StreamProcessor:
|
||||
async def create_monitored_stream(
|
||||
self,
|
||||
ctx: StreamContext,
|
||||
stream_generator: AsyncGenerator[bytes, None],
|
||||
stream_generator: AsyncGenerator[bytes],
|
||||
is_disconnected: Callable[[], Any],
|
||||
) -> AsyncGenerator[bytes, None]:
|
||||
) -> AsyncGenerator[bytes]:
|
||||
"""
|
||||
创建带监控的流生成器
|
||||
|
||||
@@ -833,8 +834,8 @@ class StreamProcessor:
|
||||
|
||||
async def create_smoothed_stream(
|
||||
self,
|
||||
stream_generator: AsyncGenerator[bytes, None],
|
||||
) -> AsyncGenerator[bytes, None]:
|
||||
stream_generator: AsyncGenerator[bytes],
|
||||
) -> AsyncGenerator[bytes]:
|
||||
"""
|
||||
创建平滑输出的流生成器
|
||||
|
||||
@@ -933,7 +934,7 @@ class StreamProcessor:
|
||||
if buffer:
|
||||
yield buffer
|
||||
|
||||
def _get_extractor(self, format_name: str) -> Optional[ContentExtractor]:
|
||||
def _get_extractor(self, format_name: str) -> ContentExtractor | None:
|
||||
"""获取或创建格式对应的提取器(带缓存)"""
|
||||
if format_name not in self._extractors:
|
||||
extractor = get_extractor(format_name)
|
||||
@@ -943,7 +944,7 @@ class StreamProcessor:
|
||||
|
||||
def _detect_format_and_extract(
|
||||
self, data: dict
|
||||
) -> tuple[Optional[str], Optional[ContentExtractor]]:
|
||||
) -> tuple[str | None, ContentExtractor | None]:
|
||||
"""
|
||||
检测数据格式并提取内容
|
||||
|
||||
@@ -998,10 +999,10 @@ class StreamProcessor:
|
||||
|
||||
|
||||
async def create_smoothed_stream(
|
||||
stream_generator: AsyncGenerator[bytes, None],
|
||||
stream_generator: AsyncGenerator[bytes],
|
||||
chunk_size: int = 20,
|
||||
delay_ms: int = 8,
|
||||
) -> AsyncGenerator[bytes, None]:
|
||||
) -> AsyncGenerator[bytes]:
|
||||
"""
|
||||
独立的平滑流生成函数
|
||||
|
||||
@@ -1032,7 +1033,7 @@ class _LightweightSmoother:
|
||||
self.delay_ms = delay_ms
|
||||
self._extractors: dict[str, ContentExtractor] = {}
|
||||
|
||||
def _get_extractor(self, format_name: str) -> Optional[ContentExtractor]:
|
||||
def _get_extractor(self, format_name: str) -> ContentExtractor | None:
|
||||
if format_name not in self._extractors:
|
||||
extractor = get_extractor(format_name)
|
||||
if extractor:
|
||||
@@ -1041,7 +1042,7 @@ class _LightweightSmoother:
|
||||
|
||||
def _detect_format_and_extract(
|
||||
self, data: dict
|
||||
) -> tuple[Optional[str], Optional[ContentExtractor]]:
|
||||
) -> tuple[str | None, ContentExtractor | None]:
|
||||
for format_name in get_extractor_formats():
|
||||
extractor = self._get_extractor(format_name)
|
||||
if extractor:
|
||||
@@ -1060,8 +1061,8 @@ class _LightweightSmoother:
|
||||
return [content[i : i + self.chunk_size] for i in range(0, text_length, self.chunk_size)]
|
||||
|
||||
async def smooth(
|
||||
self, stream_generator: AsyncGenerator[bytes, None]
|
||||
) -> AsyncGenerator[bytes, None]:
|
||||
self, stream_generator: AsyncGenerator[bytes]
|
||||
) -> AsyncGenerator[bytes]:
|
||||
buffer = b""
|
||||
is_first_content = True
|
||||
|
||||
|
||||
@@ -9,7 +9,7 @@
|
||||
|
||||
import asyncio
|
||||
import time
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
@@ -58,8 +58,8 @@ class StreamTelemetryRecorder:
|
||||
async def record_stream_stats(
|
||||
self,
|
||||
ctx: StreamContext,
|
||||
original_headers: Dict[str, str],
|
||||
original_request_body: Dict[str, Any],
|
||||
original_headers: dict[str, str],
|
||||
original_request_body: dict[str, Any],
|
||||
start_time: float,
|
||||
) -> None:
|
||||
"""
|
||||
@@ -144,9 +144,9 @@ class StreamTelemetryRecorder:
|
||||
self,
|
||||
writer: TelemetryWriter,
|
||||
ctx: StreamContext,
|
||||
original_headers: Dict[str, str],
|
||||
actual_request_body: Dict[str, Any],
|
||||
response_body: Optional[Dict[str, Any]],
|
||||
original_headers: dict[str, str],
|
||||
actual_request_body: dict[str, Any],
|
||||
response_body: dict[str, Any] | None,
|
||||
response_time_ms: int,
|
||||
) -> None:
|
||||
"""记录成功的请求"""
|
||||
@@ -193,9 +193,9 @@ class StreamTelemetryRecorder:
|
||||
self,
|
||||
writer: TelemetryWriter,
|
||||
ctx: StreamContext,
|
||||
original_headers: Dict[str, str],
|
||||
actual_request_body: Dict[str, Any],
|
||||
response_body: Optional[Dict[str, Any]],
|
||||
original_headers: dict[str, str],
|
||||
actual_request_body: dict[str, Any],
|
||||
response_body: dict[str, Any] | None,
|
||||
response_time_ms: int,
|
||||
) -> None:
|
||||
"""记录失败的请求"""
|
||||
@@ -236,9 +236,9 @@ class StreamTelemetryRecorder:
|
||||
self,
|
||||
writer: TelemetryWriter,
|
||||
ctx: StreamContext,
|
||||
original_headers: Dict[str, str],
|
||||
actual_request_body: Dict[str, Any],
|
||||
response_body: Optional[Dict[str, Any]],
|
||||
original_headers: dict[str, str],
|
||||
actual_request_body: dict[str, Any],
|
||||
response_body: dict[str, Any] | None,
|
||||
response_time_ms: int,
|
||||
) -> None:
|
||||
"""记录客户端取消的请求"""
|
||||
@@ -285,7 +285,7 @@ class StreamTelemetryRecorder:
|
||||
|
||||
from src.services.request.candidate import RequestCandidateService
|
||||
|
||||
extra_data: Dict[str, Any] = {
|
||||
extra_data: dict[str, Any] = {
|
||||
"stream_completed": ctx.is_success(),
|
||||
"data_count": ctx.data_count,
|
||||
}
|
||||
@@ -358,7 +358,7 @@ class StreamTelemetryRecorder:
|
||||
status: str,
|
||||
response_time_ms: int,
|
||||
status_code: int = 200,
|
||||
error_message: Optional[str] = None,
|
||||
error_message: str | None = None,
|
||||
) -> None:
|
||||
"""直接更新 Usage 表的状态字段"""
|
||||
try:
|
||||
@@ -378,7 +378,7 @@ class StreamTelemetryRecorder:
|
||||
|
||||
async def _get_telemetry_writer(
|
||||
self, bg_db: Session, ctx: StreamContext, response_time_ms: int
|
||||
) -> Optional[TelemetryWriter]:
|
||||
) -> TelemetryWriter | None:
|
||||
if config.usage_queue_enabled and self.user_id and self.api_key_id:
|
||||
return QueueTelemetryWriter(
|
||||
request_id=self.request_id,
|
||||
@@ -400,9 +400,9 @@ class StreamTelemetryRecorder:
|
||||
self,
|
||||
writer: TelemetryWriter,
|
||||
ctx: StreamContext,
|
||||
original_headers: Dict[str, str],
|
||||
actual_request_body: Dict[str, Any],
|
||||
response_body: Optional[Dict[str, Any]],
|
||||
original_headers: dict[str, str],
|
||||
actual_request_body: dict[str, Any],
|
||||
response_body: dict[str, Any] | None,
|
||||
response_time_ms: int,
|
||||
) -> None:
|
||||
"""根据上下文状态分发到对应的记录方法"""
|
||||
@@ -430,7 +430,7 @@ class StreamTelemetryRecorder:
|
||||
return "cancelled"
|
||||
return "failed"
|
||||
|
||||
def _build_db_writer(self, bg_db: Session) -> Optional[DbTelemetryWriter]:
|
||||
def _build_db_writer(self, bg_db: Session) -> DbTelemetryWriter | None:
|
||||
user = bg_db.query(User).filter(User.id == self.user_id).first()
|
||||
api_key_obj = bg_db.query(ApiKey).filter(ApiKey.id == self.api_key_id).first()
|
||||
|
||||
|
||||
@@ -2,10 +2,9 @@
|
||||
Handler 基础工具函数
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import TYPE_CHECKING, Any, Dict, Optional
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from src.core.exceptions import EmbeddedErrorException, ProviderNotAvailableException
|
||||
from src.core.api_format import filter_response_headers
|
||||
@@ -15,7 +14,7 @@ if TYPE_CHECKING:
|
||||
from src.core.api_format.conversion.registry import FormatConversionRegistry
|
||||
|
||||
|
||||
def get_format_converter_registry() -> "FormatConversionRegistry":
|
||||
def get_format_converter_registry() -> FormatConversionRegistry:
|
||||
"""
|
||||
获取格式转换注册表(线程安全)
|
||||
|
||||
@@ -31,7 +30,7 @@ def get_format_converter_registry() -> "FormatConversionRegistry":
|
||||
return format_conversion_registry
|
||||
|
||||
|
||||
def extract_cache_creation_tokens(usage: Dict[str, Any]) -> int:
|
||||
def extract_cache_creation_tokens(usage: dict[str, Any]) -> int:
|
||||
"""
|
||||
提取缓存创建 tokens(兼容三种格式)
|
||||
|
||||
@@ -99,7 +98,7 @@ def extract_cache_creation_tokens(usage: Dict[str, Any]) -> int:
|
||||
return old_format
|
||||
|
||||
|
||||
def build_sse_headers(extra_headers: Optional[Dict[str, str]] = None) -> Dict[str, str]:
|
||||
def build_sse_headers(extra_headers: dict[str, str] | None = None) -> dict[str, str]:
|
||||
"""
|
||||
构建 SSE(text/event-stream)推荐响应头,用于减少代理缓冲带来的卡顿/成段输出。
|
||||
|
||||
@@ -107,7 +106,7 @@ def build_sse_headers(extra_headers: Optional[Dict[str, str]] = None) -> Dict[st
|
||||
- Cache-Control: no-transform 可避免部分代理对流做压缩/改写导致缓冲
|
||||
- X-Accel-Buffering: no 可显式提示 Nginx 关闭缓冲(即使全局已关闭也无害)
|
||||
"""
|
||||
headers: Dict[str, str] = {
|
||||
headers: dict[str, str] = {
|
||||
"Cache-Control": "no-cache, no-transform",
|
||||
"X-Accel-Buffering": "no",
|
||||
}
|
||||
@@ -116,7 +115,7 @@ def build_sse_headers(extra_headers: Optional[Dict[str, str]] = None) -> Dict[st
|
||||
return headers
|
||||
|
||||
|
||||
def filter_proxy_response_headers(headers: Optional[Dict[str, str]]) -> Dict[str, str]:
|
||||
def filter_proxy_response_headers(headers: dict[str, str] | None) -> dict[str, str]:
|
||||
"""
|
||||
过滤上游响应头中不应透传给客户端的字段。
|
||||
|
||||
@@ -148,8 +147,8 @@ def check_prefetched_response_error(
|
||||
parser: Any,
|
||||
request_id: str,
|
||||
provider_name: str,
|
||||
endpoint_id: Optional[str],
|
||||
base_url: Optional[str],
|
||||
endpoint_id: str | None,
|
||||
base_url: str | None,
|
||||
) -> None:
|
||||
"""
|
||||
检查预读的响应是否为非 SSE 格式的错误响应(HTML 或纯 JSON 错误)
|
||||
|
||||
@@ -4,7 +4,7 @@ Claude Chat Adapter - 基于 ChatAdapterBase 的 Claude Chat API 适配器
|
||||
处理 /v1/messages 端点的 Claude Chat 格式请求。
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, Optional, Tuple, Type
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from fastapi import HTTPException, Request
|
||||
@@ -25,9 +25,9 @@ class ClaudeCapabilityDetector:
|
||||
|
||||
@staticmethod
|
||||
def detect_from_headers(
|
||||
headers: Dict[str, str],
|
||||
request_body: Optional[Dict[str, Any]] = None,
|
||||
) -> Dict[str, bool]:
|
||||
headers: dict[str, str],
|
||||
request_body: dict[str, Any] | None = None,
|
||||
) -> dict[str, bool]:
|
||||
"""
|
||||
从 Claude 请求头检测能力需求
|
||||
|
||||
@@ -38,7 +38,7 @@ class ClaudeCapabilityDetector:
|
||||
headers: 请求头字典
|
||||
request_body: 请求体(Claude 不使用,保留用于接口统一)
|
||||
"""
|
||||
requirements: Dict[str, bool] = {}
|
||||
requirements: dict[str, bool] = {}
|
||||
|
||||
# 使用统一的大小写不敏感获取
|
||||
beta_header = get_header_value(headers, "anthropic-beta")
|
||||
@@ -61,21 +61,21 @@ class ClaudeChatAdapter(ChatAdapterBase):
|
||||
name = "claude.chat"
|
||||
|
||||
@property
|
||||
def HANDLER_CLASS(self) -> Type[ChatHandlerBase]:
|
||||
def HANDLER_CLASS(self) -> type[ChatHandlerBase]:
|
||||
"""延迟导入 Handler 类避免循环依赖"""
|
||||
from src.api.handlers.claude.handler import ClaudeChatHandler
|
||||
|
||||
return ClaudeChatHandler
|
||||
|
||||
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
|
||||
def __init__(self, allowed_api_formats: list[str] | None = None):
|
||||
super().__init__(allowed_api_formats or ["CLAUDE"])
|
||||
logger.info(f"[{self.name}] 初始化Chat模式适配器 | API格式: {self.allowed_api_formats}")
|
||||
|
||||
def detect_capability_requirements(
|
||||
self,
|
||||
headers: Dict[str, str],
|
||||
request_body: Optional[Dict[str, Any]] = None,
|
||||
) -> Dict[str, bool]:
|
||||
headers: dict[str, str],
|
||||
request_body: dict[str, Any] | None = None,
|
||||
) -> dict[str, bool]:
|
||||
"""检测 Claude 请求中隐含的能力需求"""
|
||||
return ClaudeCapabilityDetector.detect_from_headers(headers)
|
||||
|
||||
@@ -124,7 +124,7 @@ class ClaudeChatAdapter(ChatAdapterBase):
|
||||
)
|
||||
return request
|
||||
|
||||
def _build_audit_metadata(self, _payload: Dict[str, Any], request_obj) -> Dict[str, Any]:
|
||||
def _build_audit_metadata(self, _payload: dict[str, Any], request_obj) -> dict[str, Any]:
|
||||
"""构建 Claude Chat 特定的审计元数据"""
|
||||
role_counts: dict[str, int] = {}
|
||||
for message in request_obj.messages:
|
||||
@@ -153,8 +153,8 @@ class ClaudeChatAdapter(ChatAdapterBase):
|
||||
client: httpx.AsyncClient,
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
) -> Tuple[list, Optional[str]]:
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
) -> tuple[list, str | None]:
|
||||
"""查询 Claude API 支持的模型列表"""
|
||||
headers = cls.build_headers_with_extra(api_key, extra_headers)
|
||||
|
||||
@@ -201,7 +201,7 @@ class ClaudeChatAdapter(ChatAdapterBase):
|
||||
# build_request_body 使用基类实现,通过 format_conversion_registry 自动转换 OPENAI -> CLAUDE
|
||||
|
||||
|
||||
def build_claude_adapter(x_app_header: Optional[str]):
|
||||
def build_claude_adapter(x_app_header: str | None):
|
||||
"""根据 x-app 头部构造 Chat 或 Claude Code 适配器。"""
|
||||
if x_app_header and x_app_header.lower() == "cli":
|
||||
from src.api.handlers.claude_cli.adapter import ClaudeCliAdapter
|
||||
@@ -216,7 +216,7 @@ class ClaudeTokenCountAdapter(ApiAdapter):
|
||||
name = "claude.token_count"
|
||||
mode = ApiMode.STANDARD
|
||||
|
||||
def extract_api_key(self, request: Request) -> Optional[str]:
|
||||
def extract_api_key(self, request: Request) -> str | None:
|
||||
"""从请求中提取 API 密钥 (x-api-key 或 Authorization: Bearer)"""
|
||||
# 优先检查 x-api-key
|
||||
api_key = request.headers.get("x-api-key")
|
||||
|
||||
@@ -5,7 +5,7 @@ Claude Chat Handler - 基于通用 Chat Handler 基类的简化实现
|
||||
代码量从原来的 ~1470 行减少到 ~120 行。
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import Any
|
||||
|
||||
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
|
||||
from src.api.handlers.base.utils import extract_cache_creation_tokens
|
||||
@@ -25,8 +25,8 @@ class ClaudeChatHandler(ChatHandlerBase):
|
||||
|
||||
def extract_model_from_request(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
path_params: Optional[Dict[str, Any]] = None, # noqa: ARG002
|
||||
request_body: dict[str, Any],
|
||||
path_params: dict[str, Any] | None = None, # noqa: ARG002
|
||||
) -> str:
|
||||
"""
|
||||
从请求中提取模型名 - Claude 格式实现
|
||||
@@ -45,9 +45,9 @@ class ClaudeChatHandler(ChatHandlerBase):
|
||||
|
||||
def apply_mapped_model(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
request_body: dict[str, Any],
|
||||
mapped_model: str,
|
||||
) -> Dict[str, Any]:
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
将映射后的模型名应用到请求体
|
||||
|
||||
@@ -90,7 +90,7 @@ class ClaudeChatHandler(ChatHandlerBase):
|
||||
|
||||
return request
|
||||
|
||||
def _extract_usage(self, response: Dict) -> Dict[str, int]:
|
||||
def _extract_usage(self, response: dict) -> dict[str, int]:
|
||||
"""
|
||||
从 Claude 响应中提取 token 使用情况
|
||||
|
||||
@@ -108,7 +108,7 @@ class ClaudeChatHandler(ChatHandlerBase):
|
||||
"cache_read_input_tokens": usage.get("cache_read_input_tokens", 0),
|
||||
}
|
||||
|
||||
def _normalize_response(self, response: Dict[str, Any]) -> Dict[str, Any]:
|
||||
def _normalize_response(self, response: dict[str, Any]) -> dict[str, Any]:
|
||||
"""
|
||||
规范化 Claude 响应
|
||||
|
||||
|
||||
@@ -4,10 +4,9 @@ Claude SSE 流解析器
|
||||
解析 Claude Messages API 的 Server-Sent Events 流。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any, Dict, List, Optional
|
||||
from typing import Any
|
||||
|
||||
from src.api.handlers.base.utils import extract_cache_creation_tokens
|
||||
|
||||
@@ -43,7 +42,7 @@ class ClaudeStreamParser:
|
||||
DELTA_TEXT = "text_delta"
|
||||
DELTA_INPUT_JSON = "input_json_delta"
|
||||
|
||||
def parse_chunk(self, chunk: bytes | str) -> List[Dict[str, Any]]:
|
||||
def parse_chunk(self, chunk: bytes | str) -> list[dict[str, Any]]:
|
||||
"""
|
||||
解析 SSE 数据块
|
||||
|
||||
@@ -58,10 +57,10 @@ class ClaudeStreamParser:
|
||||
else:
|
||||
text = chunk
|
||||
|
||||
events: List[Dict[str, Any]] = []
|
||||
events: list[dict[str, Any]] = []
|
||||
lines = text.strip().split("\n")
|
||||
|
||||
current_event_type: Optional[str] = None
|
||||
current_event_type: str | None = None
|
||||
|
||||
for line in lines:
|
||||
line = line.strip()
|
||||
@@ -96,7 +95,7 @@ class ClaudeStreamParser:
|
||||
|
||||
return events
|
||||
|
||||
def parse_line(self, line: str) -> Optional[Dict[str, Any]]:
|
||||
def parse_line(self, line: str) -> dict[str, Any] | None:
|
||||
"""
|
||||
解析单行 SSE 数据
|
||||
|
||||
@@ -117,7 +116,7 @@ class ClaudeStreamParser:
|
||||
except json.JSONDecodeError:
|
||||
return None
|
||||
|
||||
def is_done_event(self, event: Dict[str, Any]) -> bool:
|
||||
def is_done_event(self, event: dict[str, Any]) -> bool:
|
||||
"""
|
||||
判断是否为结束事件
|
||||
|
||||
@@ -130,7 +129,7 @@ class ClaudeStreamParser:
|
||||
event_type = event.get("type")
|
||||
return event_type in (self.EVENT_MESSAGE_STOP, "__done__")
|
||||
|
||||
def is_error_event(self, event: Dict[str, Any]) -> bool:
|
||||
def is_error_event(self, event: dict[str, Any]) -> bool:
|
||||
"""
|
||||
判断是否为错误事件
|
||||
|
||||
@@ -142,7 +141,7 @@ class ClaudeStreamParser:
|
||||
"""
|
||||
return event.get("type") == self.EVENT_ERROR
|
||||
|
||||
def get_event_type(self, event: Dict[str, Any]) -> Optional[str]:
|
||||
def get_event_type(self, event: dict[str, Any]) -> str | None:
|
||||
"""
|
||||
获取事件类型
|
||||
|
||||
@@ -155,7 +154,7 @@ class ClaudeStreamParser:
|
||||
event_type = event.get("type")
|
||||
return str(event_type) if event_type is not None else None
|
||||
|
||||
def extract_text_delta(self, event: Dict[str, Any]) -> Optional[str]:
|
||||
def extract_text_delta(self, event: dict[str, Any]) -> str | None:
|
||||
"""
|
||||
从 content_block_delta 事件中提取文本增量
|
||||
|
||||
@@ -175,7 +174,7 @@ class ClaudeStreamParser:
|
||||
|
||||
return None
|
||||
|
||||
def extract_usage(self, event: Dict[str, Any]) -> Optional[Dict[str, int]]:
|
||||
def extract_usage(self, event: dict[str, Any]) -> dict[str, int] | None:
|
||||
"""
|
||||
从事件中提取 token 使用量
|
||||
|
||||
@@ -212,7 +211,7 @@ class ClaudeStreamParser:
|
||||
|
||||
return None
|
||||
|
||||
def extract_message_id(self, event: Dict[str, Any]) -> Optional[str]:
|
||||
def extract_message_id(self, event: dict[str, Any]) -> str | None:
|
||||
"""
|
||||
从 message_start 事件中提取消息 ID
|
||||
|
||||
@@ -229,7 +228,7 @@ class ClaudeStreamParser:
|
||||
msg_id = message.get("id")
|
||||
return str(msg_id) if msg_id is not None else None
|
||||
|
||||
def extract_stop_reason(self, event: Dict[str, Any]) -> Optional[str]:
|
||||
def extract_stop_reason(self, event: dict[str, Any]) -> str | None:
|
||||
"""
|
||||
从 message_delta 事件中提取停止原因
|
||||
|
||||
|
||||
@@ -4,7 +4,7 @@ Claude CLI Adapter - 基于通用 CLI Adapter 基类的简化实现
|
||||
继承 CliAdapterBase,只需配置 FORMAT_ID 和 HANDLER_CLASS。
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, Optional, Tuple, Type
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
@@ -27,20 +27,20 @@ class ClaudeCliAdapter(CliAdapterBase):
|
||||
name = "claude.cli"
|
||||
|
||||
@property
|
||||
def HANDLER_CLASS(self) -> Type[CliMessageHandlerBase]:
|
||||
def HANDLER_CLASS(self) -> type[CliMessageHandlerBase]:
|
||||
"""延迟导入 Handler 类避免循环依赖"""
|
||||
from src.api.handlers.claude_cli.handler import ClaudeCliMessageHandler
|
||||
|
||||
return ClaudeCliMessageHandler
|
||||
|
||||
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
|
||||
def __init__(self, allowed_api_formats: list[str] | None = None):
|
||||
super().__init__(allowed_api_formats or ["CLAUDE_CLI"])
|
||||
|
||||
def detect_capability_requirements(
|
||||
self,
|
||||
headers: Dict[str, str],
|
||||
request_body: Optional[Dict[str, Any]] = None,
|
||||
) -> Dict[str, bool]:
|
||||
headers: dict[str, str],
|
||||
request_body: dict[str, Any] | None = None,
|
||||
) -> dict[str, bool]:
|
||||
"""检测 Claude CLI 请求中隐含的能力需求"""
|
||||
return ClaudeCapabilityDetector.detect_from_headers(headers)
|
||||
|
||||
@@ -61,16 +61,16 @@ class ClaudeCliAdapter(CliAdapterBase):
|
||||
"""
|
||||
return input_tokens + cache_creation_input_tokens + cache_read_input_tokens
|
||||
|
||||
def _extract_message_count(self, payload: Dict[str, Any]) -> int:
|
||||
def _extract_message_count(self, payload: dict[str, Any]) -> int:
|
||||
"""Claude CLI 使用 messages 字段"""
|
||||
messages = payload.get("messages", [])
|
||||
return len(messages) if isinstance(messages, list) else 0
|
||||
|
||||
def _build_audit_metadata(
|
||||
self,
|
||||
payload: Dict[str, Any],
|
||||
path_params: Optional[Dict[str, Any]] = None, # noqa: ARG002
|
||||
) -> Dict[str, Any]:
|
||||
payload: dict[str, Any],
|
||||
path_params: dict[str, Any] | None = None, # noqa: ARG002
|
||||
) -> dict[str, Any]:
|
||||
"""Claude CLI 特定的审计元数据"""
|
||||
model = payload.get("model", "unknown")
|
||||
stream = payload.get("stream", False)
|
||||
@@ -104,8 +104,8 @@ class ClaudeCliAdapter(CliAdapterBase):
|
||||
client: httpx.AsyncClient,
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
) -> Tuple[list, Optional[str]]:
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
) -> tuple[list, str | None]:
|
||||
"""查询 Claude API 支持的模型列表(带 CLI User-Agent)"""
|
||||
# 复用 ClaudeChatAdapter 的实现,添加 CLI User-Agent
|
||||
cli_headers = {"User-Agent": config.internal_user_agent_claude_cli}
|
||||
@@ -120,7 +120,7 @@ class ClaudeCliAdapter(CliAdapterBase):
|
||||
return models, error
|
||||
|
||||
@classmethod
|
||||
def build_endpoint_url(cls, base_url: str, request_data: Dict[str, Any], model_name: Optional[str] = None) -> str:
|
||||
def build_endpoint_url(cls, base_url: str, request_data: dict[str, Any], model_name: str | None = None) -> str:
|
||||
"""构建Claude CLI API端点URL"""
|
||||
base_url = base_url.rstrip("/")
|
||||
if base_url.endswith("/v1"):
|
||||
@@ -131,12 +131,12 @@ class ClaudeCliAdapter(CliAdapterBase):
|
||||
# build_request_body 使用基类实现,通过 format_conversion_registry 自动转换 OPENAI -> CLAUDE_CLI
|
||||
|
||||
@classmethod
|
||||
def get_cli_user_agent(cls) -> Optional[str]:
|
||||
def get_cli_user_agent(cls) -> str | None:
|
||||
"""获取Claude CLI User-Agent"""
|
||||
return config.internal_user_agent_claude_cli
|
||||
|
||||
@classmethod
|
||||
def get_cli_extra_headers(cls) -> Dict[str, str]:
|
||||
def get_cli_extra_headers(cls) -> dict[str, str]:
|
||||
"""获取Claude CLI额外请求头,包含 x-app: cli 标识"""
|
||||
headers = super().get_cli_extra_headers()
|
||||
headers["x-app"] = "cli" # 标识 CLI 模式,让上游使用正确的认证方式
|
||||
|
||||
@@ -4,7 +4,7 @@ Claude CLI Message Handler - 基于通用 CLI Handler 基类的简化实现
|
||||
继承 CliMessageHandlerBase,只需覆盖格式特定的配置和事件处理逻辑。
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import Any
|
||||
|
||||
from src.api.handlers.base.cli_handler_base import (
|
||||
CliMessageHandlerBase,
|
||||
@@ -33,8 +33,8 @@ class ClaudeCliMessageHandler(CliMessageHandlerBase):
|
||||
|
||||
def extract_model_from_request(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
path_params: Optional[Dict[str, Any]] = None, # noqa: ARG002
|
||||
request_body: dict[str, Any],
|
||||
path_params: dict[str, Any] | None = None, # noqa: ARG002
|
||||
) -> str:
|
||||
"""
|
||||
从请求中提取模型名 - Claude 格式实现
|
||||
@@ -53,9 +53,9 @@ class ClaudeCliMessageHandler(CliMessageHandlerBase):
|
||||
|
||||
def apply_mapped_model(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
request_body: dict[str, Any],
|
||||
mapped_model: str,
|
||||
) -> Dict[str, Any]:
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Claude API 的 model 在请求体顶级
|
||||
|
||||
@@ -74,7 +74,7 @@ class ClaudeCliMessageHandler(CliMessageHandlerBase):
|
||||
self,
|
||||
ctx: StreamContext,
|
||||
event_type: str,
|
||||
data: Dict[str, Any],
|
||||
data: dict[str, Any],
|
||||
) -> None:
|
||||
"""
|
||||
处理 Claude CLI 格式的 SSE 事件
|
||||
@@ -142,8 +142,8 @@ class ClaudeCliMessageHandler(CliMessageHandlerBase):
|
||||
|
||||
def _extract_response_metadata(
|
||||
self,
|
||||
response: Dict[str, Any],
|
||||
) -> Dict[str, Any]:
|
||||
response: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
从 Claude 响应中提取元数据
|
||||
|
||||
@@ -155,7 +155,7 @@ class ClaudeCliMessageHandler(CliMessageHandlerBase):
|
||||
Returns:
|
||||
提取的元数据字典
|
||||
"""
|
||||
metadata: Dict[str, Any] = {}
|
||||
metadata: dict[str, Any] = {}
|
||||
|
||||
# 提取模型名称(实际使用的模型)
|
||||
if "model" in response:
|
||||
|
||||
@@ -4,7 +4,7 @@ Gemini Chat Adapter
|
||||
处理 Gemini API 格式的请求适配
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, Optional, Tuple, Type
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from fastapi import HTTPException, Request
|
||||
@@ -32,17 +32,17 @@ class GeminiChatAdapter(ChatAdapterBase):
|
||||
name = "gemini.chat"
|
||||
|
||||
@property
|
||||
def HANDLER_CLASS(self) -> Type[ChatHandlerBase]:
|
||||
def HANDLER_CLASS(self) -> type[ChatHandlerBase]:
|
||||
"""延迟导入 Handler 类避免循环依赖"""
|
||||
from src.api.handlers.gemini.handler import GeminiChatHandler
|
||||
|
||||
return GeminiChatHandler
|
||||
|
||||
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
|
||||
def __init__(self, allowed_api_formats: list[str] | None = None):
|
||||
super().__init__(allowed_api_formats or ["GEMINI"])
|
||||
logger.info(f"[{self.name}] 初始化 Gemini Chat 适配器 | API格式: {self.allowed_api_formats}")
|
||||
|
||||
def extract_api_key(self, request: Request) -> Optional[str]:
|
||||
def extract_api_key(self, request: Request) -> str | None:
|
||||
"""
|
||||
从请求中提取 API 密钥 - Gemini 支持 header 和 query 两种方式
|
||||
|
||||
@@ -57,8 +57,8 @@ class GeminiChatAdapter(ChatAdapterBase):
|
||||
)
|
||||
|
||||
def _merge_path_params(
|
||||
self, original_request_body: Dict[str, Any], path_params: Dict[str, Any] # noqa: ARG002
|
||||
) -> Dict[str, Any]:
|
||||
self, original_request_body: dict[str, Any], path_params: dict[str, Any] # noqa: ARG002
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
合并 URL 路径参数到请求体 - Gemini 特化版本
|
||||
|
||||
@@ -111,14 +111,14 @@ class GeminiChatAdapter(ChatAdapterBase):
|
||||
request.stream = is_stream
|
||||
return request
|
||||
|
||||
def _extract_message_count(self, payload: Dict[str, Any], request_obj) -> int:
|
||||
def _extract_message_count(self, payload: dict[str, Any], request_obj) -> int:
|
||||
"""提取消息数量"""
|
||||
contents = payload.get("contents", [])
|
||||
if hasattr(request_obj, "contents"):
|
||||
contents = request_obj.contents
|
||||
return len(contents) if isinstance(contents, list) else 0
|
||||
|
||||
def _build_audit_metadata(self, payload: Dict[str, Any], request_obj) -> Dict[str, Any]:
|
||||
def _build_audit_metadata(self, payload: dict[str, Any], request_obj) -> dict[str, Any]:
|
||||
"""构建 Gemini Chat 特定的审计元数据"""
|
||||
role_counts: dict[str, int] = {}
|
||||
|
||||
@@ -171,8 +171,8 @@ class GeminiChatAdapter(ChatAdapterBase):
|
||||
client: httpx.AsyncClient,
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
) -> Tuple[list, Optional[str]]:
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
) -> tuple[list, str | None]:
|
||||
"""查询 Gemini API 支持的模型列表"""
|
||||
# Gemini 使用 URL 参数传递 key,不需要 headers 中的认证
|
||||
base_url_clean = base_url.rstrip("/")
|
||||
@@ -181,7 +181,7 @@ class GeminiChatAdapter(ChatAdapterBase):
|
||||
else:
|
||||
models_url = f"{base_url_clean}/v1beta/models?key={api_key}"
|
||||
|
||||
headers: Dict[str, str] = {}
|
||||
headers: dict[str, str] = {}
|
||||
if extra_headers:
|
||||
headers.update(extra_headers)
|
||||
|
||||
@@ -231,16 +231,16 @@ class GeminiChatAdapter(ChatAdapterBase):
|
||||
client: httpx.AsyncClient,
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
request_data: Dict[str, Any],
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
request_data: dict[str, Any],
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
# 用量计算参数
|
||||
db: Optional[Any] = None,
|
||||
user: Optional[Any] = None,
|
||||
provider_name: Optional[str] = None,
|
||||
provider_id: Optional[str] = None,
|
||||
api_key_id: Optional[str] = None,
|
||||
model_name: Optional[str] = None,
|
||||
) -> Dict[str, Any]:
|
||||
db: Any | None = None,
|
||||
user: Any | None = None,
|
||||
provider_name: str | None = None,
|
||||
provider_id: str | None = None,
|
||||
api_key_id: str | None = None,
|
||||
model_name: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""测试 Gemini API 模型连接性(非流式)"""
|
||||
# Gemini需要从request_data或model_name参数获取model名称
|
||||
effective_model_name = model_name or request_data.get("model", "")
|
||||
|
||||
@@ -4,7 +4,7 @@ Gemini Chat Handler
|
||||
处理 Gemini API 格式的请求
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import Any
|
||||
|
||||
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
|
||||
|
||||
@@ -24,8 +24,8 @@ class GeminiChatHandler(ChatHandlerBase):
|
||||
|
||||
def extract_model_from_request(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
path_params: Optional[Dict[str, Any]] = None,
|
||||
request_body: dict[str, Any],
|
||||
path_params: dict[str, Any] | None = None,
|
||||
) -> str:
|
||||
"""
|
||||
从请求中提取模型名 - Gemini Chat 格式实现
|
||||
@@ -74,7 +74,7 @@ class GeminiChatHandler(ChatHandlerBase):
|
||||
|
||||
return request
|
||||
|
||||
def _extract_usage(self, response: Dict) -> Dict[str, int]:
|
||||
def _extract_usage(self, response: dict) -> dict[str, int]:
|
||||
"""
|
||||
从 Gemini 响应中提取 token 使用情况
|
||||
|
||||
@@ -99,7 +99,7 @@ class GeminiChatHandler(ChatHandlerBase):
|
||||
"cache_read_input_tokens": usage.get("cached_tokens", 0),
|
||||
}
|
||||
|
||||
def _normalize_response(self, response: Dict) -> Dict:
|
||||
def _normalize_response(self, response: dict) -> dict:
|
||||
"""
|
||||
规范化 Gemini 响应
|
||||
|
||||
|
||||
@@ -15,7 +15,7 @@ Gemini streamGenerateContent 常见两种返回(与上游/代理实现有关
|
||||
"""
|
||||
|
||||
import json
|
||||
from typing import Any, Dict, List, Optional, Union
|
||||
from typing import Any
|
||||
|
||||
|
||||
class GeminiStreamParser:
|
||||
@@ -43,7 +43,7 @@ class GeminiStreamParser:
|
||||
self._in_array = False
|
||||
self._brace_depth = 0
|
||||
|
||||
def parse_chunk(self, chunk: Union[bytes, str]) -> List[Dict[str, Any]]:
|
||||
def parse_chunk(self, chunk: bytes | str) -> list[dict[str, Any]]:
|
||||
"""
|
||||
解析流式数据块
|
||||
|
||||
@@ -58,7 +58,7 @@ class GeminiStreamParser:
|
||||
else:
|
||||
text = chunk
|
||||
|
||||
events: List[Dict[str, Any]] = []
|
||||
events: list[dict[str, Any]] = []
|
||||
|
||||
for char in text:
|
||||
if char == "[" and not self._in_array:
|
||||
@@ -97,7 +97,7 @@ class GeminiStreamParser:
|
||||
|
||||
return events
|
||||
|
||||
def parse_line(self, line: str) -> Optional[Dict[str, Any]]:
|
||||
def parse_line(self, line: str) -> dict[str, Any] | None:
|
||||
"""
|
||||
解析单行 JSON 数据
|
||||
|
||||
@@ -118,7 +118,7 @@ class GeminiStreamParser:
|
||||
except json.JSONDecodeError:
|
||||
return None
|
||||
|
||||
def is_done_event(self, event: Dict[str, Any]) -> bool:
|
||||
def is_done_event(self, event: dict[str, Any]) -> bool:
|
||||
"""
|
||||
判断是否为结束事件
|
||||
|
||||
@@ -143,7 +143,7 @@ class GeminiStreamParser:
|
||||
|
||||
return False
|
||||
|
||||
def is_error_event(self, event: Dict[str, Any]) -> bool:
|
||||
def is_error_event(self, event: dict[str, Any]) -> bool:
|
||||
"""
|
||||
判断是否为错误事件
|
||||
|
||||
@@ -171,7 +171,7 @@ class GeminiStreamParser:
|
||||
|
||||
return False
|
||||
|
||||
def extract_error_info(self, event: Dict[str, Any]) -> Optional[Dict[str, Any]]:
|
||||
def extract_error_info(self, event: dict[str, Any]) -> dict[str, Any] | None:
|
||||
"""
|
||||
从事件中提取错误信息
|
||||
|
||||
@@ -208,7 +208,7 @@ class GeminiStreamParser:
|
||||
|
||||
return None
|
||||
|
||||
def get_finish_reason(self, event: Dict[str, Any]) -> Optional[str]:
|
||||
def get_finish_reason(self, event: dict[str, Any]) -> str | None:
|
||||
"""
|
||||
获取结束原因
|
||||
|
||||
@@ -224,7 +224,7 @@ class GeminiStreamParser:
|
||||
return str(reason) if reason is not None else None
|
||||
return None
|
||||
|
||||
def extract_text_delta(self, event: Dict[str, Any]) -> Optional[str]:
|
||||
def extract_text_delta(self, event: dict[str, Any]) -> str | None:
|
||||
"""
|
||||
从响应中提取文本内容
|
||||
|
||||
@@ -248,7 +248,7 @@ class GeminiStreamParser:
|
||||
|
||||
return "".join(text_parts) if text_parts else None
|
||||
|
||||
def extract_usage(self, event: Dict[str, Any]) -> Optional[Dict[str, int]]:
|
||||
def extract_usage(self, event: dict[str, Any]) -> dict[str, int] | None:
|
||||
"""
|
||||
从事件中提取 token 使用量
|
||||
|
||||
@@ -280,7 +280,7 @@ class GeminiStreamParser:
|
||||
"cached_tokens": usage_metadata.get("cachedContentTokenCount", 0),
|
||||
}
|
||||
|
||||
def extract_model_version(self, event: Dict[str, Any]) -> Optional[str]:
|
||||
def extract_model_version(self, event: dict[str, Any]) -> str | None:
|
||||
"""
|
||||
从响应中提取模型版本
|
||||
|
||||
@@ -293,7 +293,7 @@ class GeminiStreamParser:
|
||||
version = event.get("modelVersion")
|
||||
return str(version) if version is not None else None
|
||||
|
||||
def extract_safety_ratings(self, event: Dict[str, Any]) -> Optional[List[Dict[str, Any]]]:
|
||||
def extract_safety_ratings(self, event: dict[str, Any]) -> list[dict[str, Any]] | None:
|
||||
"""
|
||||
从响应中提取安全评级
|
||||
|
||||
|
||||
@@ -4,7 +4,7 @@ Gemini CLI Adapter - 基于通用 CLI Adapter 基类的实现
|
||||
继承 CliAdapterBase,处理 Gemini CLI 格式的请求。
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, Optional, Tuple, Type
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from fastapi import Request
|
||||
@@ -29,16 +29,16 @@ class GeminiCliAdapter(CliAdapterBase):
|
||||
name = "gemini.cli"
|
||||
|
||||
@property
|
||||
def HANDLER_CLASS(self) -> Type[CliMessageHandlerBase]:
|
||||
def HANDLER_CLASS(self) -> type[CliMessageHandlerBase]:
|
||||
"""延迟导入 Handler 类避免循环依赖"""
|
||||
from src.api.handlers.gemini_cli.handler import GeminiCliMessageHandler
|
||||
|
||||
return GeminiCliMessageHandler
|
||||
|
||||
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
|
||||
def __init__(self, allowed_api_formats: list[str] | None = None):
|
||||
super().__init__(allowed_api_formats or ["GEMINI_CLI"])
|
||||
|
||||
def extract_api_key(self, request: Request) -> Optional[str]:
|
||||
def extract_api_key(self, request: Request) -> str | None:
|
||||
"""
|
||||
从请求中提取 API 密钥 - Gemini CLI 支持 header 和 query 两种方式
|
||||
|
||||
@@ -53,8 +53,8 @@ class GeminiCliAdapter(CliAdapterBase):
|
||||
)
|
||||
|
||||
def _merge_path_params(
|
||||
self, original_request_body: Dict[str, Any], path_params: Dict[str, Any] # noqa: ARG002
|
||||
) -> Dict[str, Any]:
|
||||
self, original_request_body: dict[str, Any], path_params: dict[str, Any] # noqa: ARG002
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
合并 URL 路径参数到请求体 - Gemini CLI 特化版本
|
||||
|
||||
@@ -74,23 +74,23 @@ class GeminiCliAdapter(CliAdapterBase):
|
||||
# Gemini: 不合并任何 path_params 到请求体
|
||||
return original_request_body.copy()
|
||||
|
||||
def _extract_message_count(self, payload: Dict[str, Any]) -> int:
|
||||
def _extract_message_count(self, payload: dict[str, Any]) -> int:
|
||||
"""Gemini CLI 使用 contents 字段"""
|
||||
contents = payload.get("contents", [])
|
||||
return len(contents) if isinstance(contents, list) else 0
|
||||
|
||||
def _build_audit_metadata(
|
||||
self,
|
||||
payload: Dict[str, Any],
|
||||
path_params: Optional[Dict[str, Any]] = None,
|
||||
) -> Dict[str, Any]:
|
||||
payload: dict[str, Any],
|
||||
path_params: dict[str, Any] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""Gemini CLI 特定的审计元数据"""
|
||||
# 从 path_params 获取 model(Gemini 请求体不含 model)
|
||||
model = path_params.get("model", "unknown") if path_params else "unknown"
|
||||
contents = payload.get("contents", [])
|
||||
generation_config = payload.get("generation_config", {}) or {}
|
||||
|
||||
role_counts: Dict[str, int] = {}
|
||||
role_counts: dict[str, int] = {}
|
||||
for content in contents:
|
||||
role = content.get("role", "unknown") if isinstance(content, dict) else "unknown"
|
||||
role_counts[role] = role_counts.get(role, 0) + 1
|
||||
@@ -120,8 +120,8 @@ class GeminiCliAdapter(CliAdapterBase):
|
||||
client: httpx.AsyncClient,
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
) -> Tuple[list, Optional[str]]:
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
) -> tuple[list, str | None]:
|
||||
"""查询 Gemini API 支持的模型列表(带 CLI User-Agent)"""
|
||||
# 复用 GeminiChatAdapter 的实现,添加 CLI User-Agent
|
||||
cli_headers = {"User-Agent": config.internal_user_agent_gemini_cli}
|
||||
@@ -136,7 +136,7 @@ class GeminiCliAdapter(CliAdapterBase):
|
||||
return models, error
|
||||
|
||||
@classmethod
|
||||
def build_endpoint_url(cls, base_url: str, request_data: Dict[str, Any], model_name: Optional[str] = None) -> str:
|
||||
def build_endpoint_url(cls, base_url: str, request_data: dict[str, Any], model_name: str | None = None) -> str:
|
||||
"""构建Gemini CLI API端点URL"""
|
||||
effective_model_name = model_name or request_data.get("model", "")
|
||||
if not effective_model_name:
|
||||
@@ -152,12 +152,12 @@ class GeminiCliAdapter(CliAdapterBase):
|
||||
# build_request_body 使用基类实现,通过 format_conversion_registry 自动转换 OPENAI -> GEMINI_CLI
|
||||
|
||||
@classmethod
|
||||
def get_cli_user_agent(cls) -> Optional[str]:
|
||||
def get_cli_user_agent(cls) -> str | None:
|
||||
"""获取Gemini CLI User-Agent"""
|
||||
return config.internal_user_agent_gemini_cli
|
||||
|
||||
@classmethod
|
||||
def get_cli_extra_headers(cls) -> Dict[str, str]:
|
||||
def get_cli_extra_headers(cls) -> dict[str, str]:
|
||||
"""获取Gemini CLI额外请求头,包含 x-app: cli 标识"""
|
||||
headers = super().get_cli_extra_headers()
|
||||
headers["x-app"] = "cli" # 标识 CLI 模式,让上游使用正确的 adapter
|
||||
|
||||
@@ -4,7 +4,7 @@ Gemini CLI Message Handler - 基于通用 CLI Handler 基类的实现
|
||||
继承 CliMessageHandlerBase,处理 Gemini CLI API 格式的请求。
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import Any
|
||||
|
||||
from src.api.handlers.base.cli_handler_base import (
|
||||
CliMessageHandlerBase,
|
||||
@@ -34,8 +34,8 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
|
||||
|
||||
def extract_model_from_request(
|
||||
self,
|
||||
request_body: Dict[str, Any], # noqa: ARG002 - 基类签名要求
|
||||
path_params: Optional[Dict[str, Any]] = None,
|
||||
request_body: dict[str, Any], # noqa: ARG002 - 基类签名要求
|
||||
path_params: dict[str, Any] | None = None,
|
||||
) -> str:
|
||||
"""
|
||||
从请求中提取模型名 - Gemini 格式实现
|
||||
@@ -57,8 +57,8 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
|
||||
|
||||
def prepare_provider_request_body(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
) -> Dict[str, Any]:
|
||||
request_body: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
准备发送给 Gemini API 的请求体 - 移除 model 字段
|
||||
|
||||
@@ -77,9 +77,9 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
|
||||
|
||||
def get_model_for_url(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
mapped_model: Optional[str],
|
||||
) -> Optional[str]:
|
||||
request_body: dict[str, Any],
|
||||
mapped_model: str | None,
|
||||
) -> str | None:
|
||||
"""
|
||||
Gemini 需要将 model 放入 URL 路径中
|
||||
|
||||
@@ -93,7 +93,7 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
|
||||
# 优先使用映射后的模型名,否则使用请求体中的
|
||||
return mapped_model or request_body.get("model")
|
||||
|
||||
def _extract_usage_from_event(self, event: Dict[str, Any]) -> Dict[str, int]:
|
||||
def _extract_usage_from_event(self, event: dict[str, Any]) -> dict[str, int]:
|
||||
"""
|
||||
从 Gemini 事件中提取 token 使用情况
|
||||
|
||||
@@ -126,7 +126,7 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
|
||||
self,
|
||||
ctx: StreamContext,
|
||||
_event_type: str,
|
||||
data: Dict[str, Any],
|
||||
data: dict[str, Any],
|
||||
) -> None:
|
||||
"""
|
||||
处理 Gemini CLI 格式的流式事件
|
||||
@@ -190,8 +190,8 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
|
||||
|
||||
def _extract_response_metadata(
|
||||
self,
|
||||
response: Dict[str, Any],
|
||||
) -> Dict[str, Any]:
|
||||
response: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
从 Gemini 响应中提取元数据
|
||||
|
||||
@@ -203,7 +203,7 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
|
||||
Returns:
|
||||
包含 model_version 的元数据字典
|
||||
"""
|
||||
metadata: Dict[str, Any] = {}
|
||||
metadata: dict[str, Any] = {}
|
||||
model_version = response.get("modelVersion")
|
||||
if model_version:
|
||||
metadata["model_version"] = model_version
|
||||
|
||||
@@ -4,7 +4,7 @@ OpenAI Chat Adapter - 基于 ChatAdapterBase 的 OpenAI Chat API 适配器
|
||||
处理 /v1/chat/completions 端点的 OpenAI Chat 格式请求。
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, Optional, Tuple, Type
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from fastapi.responses import JSONResponse
|
||||
@@ -28,13 +28,13 @@ class OpenAIChatAdapter(ChatAdapterBase):
|
||||
name = "openai.chat"
|
||||
|
||||
@property
|
||||
def HANDLER_CLASS(self) -> Type[ChatHandlerBase]:
|
||||
def HANDLER_CLASS(self) -> type[ChatHandlerBase]:
|
||||
"""延迟导入 Handler 类避免循环依赖"""
|
||||
from src.api.handlers.openai.handler import OpenAIChatHandler
|
||||
|
||||
return OpenAIChatHandler
|
||||
|
||||
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
|
||||
def __init__(self, allowed_api_formats: list[str] | None = None):
|
||||
super().__init__(allowed_api_formats or ["OPENAI"])
|
||||
|
||||
def _validate_request_body(self, original_request_body: dict, path_params: dict = None):
|
||||
@@ -66,7 +66,7 @@ class OpenAIChatAdapter(ChatAdapterBase):
|
||||
max_tokens=original_request_body.get("max_tokens"),
|
||||
)
|
||||
|
||||
def _build_audit_metadata(self, payload: Dict[str, Any], request_obj) -> Dict[str, Any]:
|
||||
def _build_audit_metadata(self, payload: dict[str, Any], request_obj) -> dict[str, Any]:
|
||||
"""构建 OpenAI Chat 特定的审计元数据"""
|
||||
role_counts = {}
|
||||
for message in request_obj.messages:
|
||||
@@ -105,8 +105,8 @@ class OpenAIChatAdapter(ChatAdapterBase):
|
||||
client: httpx.AsyncClient,
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
) -> Tuple[list, Optional[str]]:
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
) -> tuple[list, str | None]:
|
||||
"""查询 OpenAI 兼容 API 支持的模型列表"""
|
||||
headers = cls.build_headers_with_extra(api_key, extra_headers)
|
||||
|
||||
|
||||
@@ -5,7 +5,7 @@ OpenAI Chat Handler - 基于通用 Chat Handler 基类的简化实现
|
||||
代码量从原来的 ~1315 行减少到 ~100 行。
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import Any
|
||||
|
||||
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
|
||||
|
||||
@@ -24,8 +24,8 @@ class OpenAIChatHandler(ChatHandlerBase):
|
||||
|
||||
def extract_model_from_request(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
path_params: Optional[Dict[str, Any]] = None, # noqa: ARG002
|
||||
request_body: dict[str, Any],
|
||||
path_params: dict[str, Any] | None = None, # noqa: ARG002
|
||||
) -> str:
|
||||
"""
|
||||
从请求中提取模型名 - OpenAI 格式实现
|
||||
@@ -44,9 +44,9 @@ class OpenAIChatHandler(ChatHandlerBase):
|
||||
|
||||
def apply_mapped_model(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
request_body: dict[str, Any],
|
||||
mapped_model: str,
|
||||
) -> Dict[str, Any]:
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
将映射后的模型名应用到请求体
|
||||
|
||||
@@ -89,7 +89,7 @@ class OpenAIChatHandler(ChatHandlerBase):
|
||||
|
||||
return request
|
||||
|
||||
def _extract_usage(self, response: Dict) -> Dict[str, int]:
|
||||
def _extract_usage(self, response: dict) -> dict[str, int]:
|
||||
"""
|
||||
从 OpenAI 响应中提取 token 使用情况
|
||||
|
||||
@@ -106,7 +106,7 @@ class OpenAIChatHandler(ChatHandlerBase):
|
||||
"cache_read_input_tokens": 0,
|
||||
}
|
||||
|
||||
def _normalize_response(self, response: Dict) -> Dict:
|
||||
def _normalize_response(self, response: dict) -> dict:
|
||||
"""
|
||||
规范化 OpenAI 响应
|
||||
|
||||
|
||||
@@ -4,10 +4,9 @@ OpenAI SSE 流解析器
|
||||
解析 OpenAI Chat Completions API 的 Server-Sent Events 流。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any, Dict, List, Optional
|
||||
from typing import Any
|
||||
|
||||
|
||||
class OpenAIStreamParser:
|
||||
@@ -23,7 +22,7 @@ class OpenAIStreamParser:
|
||||
- 流结束时发送 data: [DONE]
|
||||
"""
|
||||
|
||||
def parse_chunk(self, chunk: bytes | str) -> List[Dict[str, Any]]:
|
||||
def parse_chunk(self, chunk: bytes | str) -> list[dict[str, Any]]:
|
||||
"""
|
||||
解析 SSE 数据块
|
||||
|
||||
@@ -38,7 +37,7 @@ class OpenAIStreamParser:
|
||||
else:
|
||||
text = chunk
|
||||
|
||||
chunks: List[Dict[str, Any]] = []
|
||||
chunks: list[dict[str, Any]] = []
|
||||
lines = text.strip().split("\n")
|
||||
|
||||
for line in lines:
|
||||
@@ -64,7 +63,7 @@ class OpenAIStreamParser:
|
||||
|
||||
return chunks
|
||||
|
||||
def parse_line(self, line: str) -> Optional[Dict[str, Any]]:
|
||||
def parse_line(self, line: str) -> dict[str, Any] | None:
|
||||
"""
|
||||
解析单行 SSE 数据
|
||||
|
||||
@@ -85,7 +84,7 @@ class OpenAIStreamParser:
|
||||
except json.JSONDecodeError:
|
||||
return None
|
||||
|
||||
def is_done_chunk(self, chunk: Dict[str, Any]) -> bool:
|
||||
def is_done_chunk(self, chunk: dict[str, Any]) -> bool:
|
||||
"""
|
||||
判断是否为结束 chunk
|
||||
|
||||
@@ -107,7 +106,7 @@ class OpenAIStreamParser:
|
||||
|
||||
return False
|
||||
|
||||
def get_finish_reason(self, chunk: Dict[str, Any]) -> Optional[str]:
|
||||
def get_finish_reason(self, chunk: dict[str, Any]) -> str | None:
|
||||
"""
|
||||
获取结束原因
|
||||
|
||||
@@ -123,7 +122,7 @@ class OpenAIStreamParser:
|
||||
return str(reason) if reason is not None else None
|
||||
return None
|
||||
|
||||
def extract_text_delta(self, chunk: Dict[str, Any]) -> Optional[str]:
|
||||
def extract_text_delta(self, chunk: dict[str, Any]) -> str | None:
|
||||
"""
|
||||
从 chunk 中提取文本增量
|
||||
|
||||
@@ -145,7 +144,7 @@ class OpenAIStreamParser:
|
||||
|
||||
return None
|
||||
|
||||
def extract_tool_calls_delta(self, chunk: Dict[str, Any]) -> Optional[List[Dict[str, Any]]]:
|
||||
def extract_tool_calls_delta(self, chunk: dict[str, Any]) -> list[dict[str, Any]] | None:
|
||||
"""
|
||||
从 chunk 中提取工具调用增量
|
||||
|
||||
@@ -165,7 +164,7 @@ class OpenAIStreamParser:
|
||||
return tool_calls
|
||||
return None
|
||||
|
||||
def extract_role(self, chunk: Dict[str, Any]) -> Optional[str]:
|
||||
def extract_role(self, chunk: dict[str, Any]) -> str | None:
|
||||
"""
|
||||
从 chunk 中提取角色
|
||||
|
||||
|
||||
@@ -4,7 +4,7 @@ OpenAI CLI Adapter - 基于通用 CLI Adapter 基类的简化实现
|
||||
继承 CliAdapterBase,只需配置 FORMAT_ID 和 HANDLER_CLASS。
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, Optional, Tuple, Type
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
@@ -27,13 +27,13 @@ class OpenAICliAdapter(CliAdapterBase):
|
||||
name = "openai.cli"
|
||||
|
||||
@property
|
||||
def HANDLER_CLASS(self) -> Type[CliMessageHandlerBase]:
|
||||
def HANDLER_CLASS(self) -> type[CliMessageHandlerBase]:
|
||||
"""延迟导入 Handler 类避免循环依赖"""
|
||||
from src.api.handlers.openai_cli.handler import OpenAICliMessageHandler
|
||||
|
||||
return OpenAICliMessageHandler
|
||||
|
||||
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
|
||||
def __init__(self, allowed_api_formats: list[str] | None = None):
|
||||
super().__init__(allowed_api_formats or ["OPENAI_CLI"])
|
||||
|
||||
# =========================================================================
|
||||
@@ -46,8 +46,8 @@ class OpenAICliAdapter(CliAdapterBase):
|
||||
client: httpx.AsyncClient,
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
) -> Tuple[list, Optional[str]]:
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
) -> tuple[list, str | None]:
|
||||
"""查询 OpenAI 兼容 API 支持的模型列表(带 CLI User-Agent)"""
|
||||
# 复用 OpenAIChatAdapter 的实现,添加 CLI User-Agent
|
||||
cli_headers = {"User-Agent": config.internal_user_agent_openai_cli}
|
||||
@@ -62,7 +62,7 @@ class OpenAICliAdapter(CliAdapterBase):
|
||||
return models, error
|
||||
|
||||
@classmethod
|
||||
def build_endpoint_url(cls, base_url: str, request_data: Dict[str, Any], model_name: Optional[str] = None) -> str:
|
||||
def build_endpoint_url(cls, base_url: str, request_data: dict[str, Any], model_name: str | None = None) -> str:
|
||||
"""构建OpenAI CLI API端点URL"""
|
||||
base_url = base_url.rstrip("/")
|
||||
if base_url.endswith("/v1"):
|
||||
@@ -74,7 +74,7 @@ class OpenAICliAdapter(CliAdapterBase):
|
||||
# OPENAI -> OPENAI_CLI 无转换器,会直接透传原始请求
|
||||
|
||||
@classmethod
|
||||
def get_cli_user_agent(cls) -> Optional[str]:
|
||||
def get_cli_user_agent(cls) -> str | None:
|
||||
"""获取OpenAI CLI User-Agent"""
|
||||
return config.internal_user_agent_openai_cli
|
||||
|
||||
|
||||
@@ -5,7 +5,7 @@ OpenAI CLI Message Handler - 基于通用 CLI Handler 基类的简化实现
|
||||
代码量从原来的 900+ 行减少到 ~100 行。
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import Any
|
||||
|
||||
from src.api.handlers.base.cli_handler_base import (
|
||||
CliMessageHandlerBase,
|
||||
@@ -32,8 +32,8 @@ class OpenAICliMessageHandler(CliMessageHandlerBase):
|
||||
|
||||
def extract_model_from_request(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
path_params: Optional[Dict[str, Any]] = None, # noqa: ARG002
|
||||
request_body: dict[str, Any],
|
||||
path_params: dict[str, Any] | None = None, # noqa: ARG002
|
||||
) -> str:
|
||||
"""
|
||||
从请求中提取模型名 - OpenAI 格式实现
|
||||
@@ -52,9 +52,9 @@ class OpenAICliMessageHandler(CliMessageHandlerBase):
|
||||
|
||||
def apply_mapped_model(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
request_body: dict[str, Any],
|
||||
mapped_model: str,
|
||||
) -> Dict[str, Any]:
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
OpenAI CLI (Responses API) 的 model 在请求体顶级字段。
|
||||
|
||||
@@ -73,7 +73,7 @@ class OpenAICliMessageHandler(CliMessageHandlerBase):
|
||||
self,
|
||||
ctx: StreamContext,
|
||||
event_type: str,
|
||||
data: Dict[str, Any],
|
||||
data: dict[str, Any],
|
||||
) -> None:
|
||||
"""
|
||||
处理 OpenAI CLI 格式的 SSE 事件
|
||||
@@ -144,8 +144,8 @@ class OpenAICliMessageHandler(CliMessageHandlerBase):
|
||||
|
||||
def _extract_response_metadata(
|
||||
self,
|
||||
response: Dict[str, Any],
|
||||
) -> Dict[str, Any]:
|
||||
response: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
从 OpenAI 响应中提取元数据
|
||||
|
||||
@@ -157,7 +157,7 @@ class OpenAICliMessageHandler(CliMessageHandlerBase):
|
||||
Returns:
|
||||
提取的元数据字典
|
||||
"""
|
||||
metadata: Dict[str, Any] = {}
|
||||
metadata: dict[str, Any] = {}
|
||||
|
||||
# 提取模型名称(实际使用的模型)
|
||||
if "model" in response:
|
||||
|
||||
@@ -2,7 +2,6 @@
|
||||
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -22,7 +21,7 @@ pipeline = ApiRequestPipeline()
|
||||
@router.get("/my-audit-logs")
|
||||
async def get_my_audit_logs(
|
||||
request: Request,
|
||||
event_type: Optional[str] = Query(None, description="事件类型筛选"),
|
||||
event_type: str | None = Query(None, description="事件类型筛选"),
|
||||
days: int = Query(30, description="查询天数"),
|
||||
limit: int = Query(50, description="返回数量限制"),
|
||||
offset: int = Query(0, ge=0, description="偏移量"),
|
||||
@@ -86,7 +85,7 @@ class AuthenticatedApiAdapter(ApiAdapter):
|
||||
|
||||
@dataclass
|
||||
class UserAuditLogsAdapter(AuthenticatedApiAdapter):
|
||||
event_type: Optional[str]
|
||||
event_type: str | None
|
||||
days: int
|
||||
limit: int
|
||||
offset: int
|
||||
|
||||
@@ -1,8 +1,7 @@
|
||||
"""OAuth 管理端点(管理员)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Dict, List, Optional
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, Request
|
||||
from pydantic import BaseModel, Field, ValidationError
|
||||
@@ -27,24 +26,24 @@ class SupportedOAuthType(BaseModel):
|
||||
default_authorization_url: str
|
||||
default_token_url: str
|
||||
default_userinfo_url: str
|
||||
default_scopes: List[str]
|
||||
default_scopes: list[str]
|
||||
|
||||
|
||||
class OAuthProviderUpsertRequest(BaseModel):
|
||||
display_name: str = Field(..., min_length=1, max_length=100)
|
||||
client_id: str = Field(..., min_length=1, max_length=255)
|
||||
client_secret: Optional[str] = Field(None, max_length=2048)
|
||||
client_secret: str | None = Field(None, max_length=2048)
|
||||
|
||||
authorization_url_override: Optional[str] = Field(None, max_length=500)
|
||||
token_url_override: Optional[str] = Field(None, max_length=500)
|
||||
userinfo_url_override: Optional[str] = Field(None, max_length=500)
|
||||
scopes: Optional[List[str]] = None
|
||||
authorization_url_override: str | None = Field(None, max_length=500)
|
||||
token_url_override: str | None = Field(None, max_length=500)
|
||||
userinfo_url_override: str | None = Field(None, max_length=500)
|
||||
scopes: list[str] | None = None
|
||||
|
||||
redirect_uri: str = Field(..., min_length=1, max_length=500)
|
||||
frontend_callback_url: str = Field(..., min_length=1, max_length=500)
|
||||
|
||||
attribute_mapping: Optional[Dict[str, Any]] = None
|
||||
extra_config: Optional[Dict[str, Any]] = None
|
||||
attribute_mapping: dict[str, Any] | None = None
|
||||
extra_config: dict[str, Any] | None = None
|
||||
|
||||
is_enabled: bool = False
|
||||
force: bool = False
|
||||
@@ -55,14 +54,14 @@ class OAuthProviderAdminResponse(BaseModel):
|
||||
display_name: str
|
||||
client_id: str
|
||||
has_secret: bool
|
||||
authorization_url_override: Optional[str] = None
|
||||
token_url_override: Optional[str] = None
|
||||
userinfo_url_override: Optional[str] = None
|
||||
scopes: Optional[List[str]] = None
|
||||
authorization_url_override: str | None = None
|
||||
token_url_override: str | None = None
|
||||
userinfo_url_override: str | None = None
|
||||
scopes: list[str] | None = None
|
||||
redirect_uri: str
|
||||
frontend_callback_url: str
|
||||
attribute_mapping: Optional[Dict[str, Any]] = None
|
||||
extra_config: Optional[Dict[str, Any]] = None
|
||||
attribute_mapping: dict[str, Any] | None = None
|
||||
extra_config: dict[str, Any] | None = None
|
||||
is_enabled: bool
|
||||
|
||||
|
||||
@@ -77,19 +76,19 @@ class OAuthProviderTestRequest(BaseModel):
|
||||
"""测试请求,使用表单数据而非数据库配置"""
|
||||
|
||||
client_id: str = Field(..., min_length=1)
|
||||
client_secret: Optional[str] = None
|
||||
authorization_url_override: Optional[str] = None
|
||||
token_url_override: Optional[str] = None
|
||||
client_secret: str | None = None
|
||||
authorization_url_override: str | None = None
|
||||
token_url_override: str | None = None
|
||||
redirect_uri: str = Field(..., min_length=1)
|
||||
|
||||
|
||||
@router.get("/supported-types", response_model=List[SupportedOAuthType])
|
||||
@router.get("/supported-types", response_model=list[SupportedOAuthType])
|
||||
async def get_supported_types(request: Request, db: Session = Depends(get_db)) -> Any:
|
||||
adapter = GetSupportedTypesAdapter()
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.get("/providers", response_model=List[OAuthProviderAdminResponse])
|
||||
@router.get("/providers", response_model=list[OAuthProviderAdminResponse])
|
||||
async def list_provider_configs(request: Request, db: Session = Depends(get_db)) -> Any:
|
||||
adapter = ListOAuthProviderConfigsAdapter()
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
@@ -1,8 +1,7 @@
|
||||
"""OAuth 公开端点(无需登录)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Optional
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, status
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -38,10 +37,10 @@ async def oauth_authorize(provider_type: str, db: Session = Depends(get_db)) ->
|
||||
async def oauth_callback(
|
||||
provider_type: str,
|
||||
db: Session = Depends(get_db),
|
||||
code: Optional[str] = Query(None),
|
||||
state: Optional[str] = Query(None),
|
||||
error: Optional[str] = Query(None),
|
||||
error_description: Optional[str] = Query(None),
|
||||
code: str | None = Query(None),
|
||||
state: str | None = Query(None),
|
||||
error: str | None = Query(None),
|
||||
error_description: str | None = Query(None),
|
||||
) -> RedirectResponse:
|
||||
"""
|
||||
OAuth 回调端点。
|
||||
|
||||
@@ -1,8 +1,7 @@
|
||||
"""OAuth 用户端点(需登录)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Optional, cast
|
||||
from typing import Any, cast
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -53,7 +52,7 @@ async def bind_oauth_provider(
|
||||
provider_type: str,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
bind_token: Optional[str] = None,
|
||||
bind_token: str | None = None,
|
||||
) -> RedirectResponse:
|
||||
"""发起 OAuth 绑定流程,支持通过 bind_token 参数进行安全认证"""
|
||||
adapter = BindOAuthProviderAdapter(provider_type=provider_type, bind_token=bind_token)
|
||||
@@ -115,10 +114,10 @@ class BindOAuthProviderAdapter(AuthenticatedApiAdapter):
|
||||
2. bind_token 参数 (浏览器跳转场景)
|
||||
"""
|
||||
|
||||
def __init__(self, provider_type: str, bind_token: Optional[str] = None):
|
||||
def __init__(self, provider_type: str, bind_token: str | None = None):
|
||||
self.provider_type = provider_type
|
||||
self.bind_token = bind_token
|
||||
self._user_from_bind_token: Optional[User] = None
|
||||
self._user_from_bind_token: User | None = None
|
||||
|
||||
@property
|
||||
def mode(self) -> ApiMode: # type: ignore[override]
|
||||
@@ -136,7 +135,7 @@ class BindOAuthProviderAdapter(AuthenticatedApiAdapter):
|
||||
raise HTTPException(status_code=401, detail="未登录")
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> RedirectResponse: # type: ignore[override]
|
||||
user: Optional[User] = context.user
|
||||
user: User | None = context.user
|
||||
|
||||
# 如果使用 bind_token,验证并获取用户
|
||||
if self.bind_token:
|
||||
|
||||
@@ -6,10 +6,9 @@
|
||||
from collections import defaultdict
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, Request
|
||||
from sqlalchemy import and_, func, or_
|
||||
from sqlalchemy import and_, or_
|
||||
from sqlalchemy.orm import Session, joinedload
|
||||
|
||||
from src.api.base.adapter import ApiAdapter, ApiMode
|
||||
@@ -41,10 +40,10 @@ router = APIRouter(prefix="/api/public", tags=["System Catalog"])
|
||||
pipeline = ApiRequestPipeline()
|
||||
|
||||
|
||||
@router.get("/providers", response_model=List[PublicProviderResponse])
|
||||
@router.get("/providers", response_model=list[PublicProviderResponse])
|
||||
async def get_public_providers(
|
||||
request: Request,
|
||||
is_active: Optional[bool] = Query(None, description="过滤活跃状态"),
|
||||
is_active: bool | None = Query(None, description="过滤活跃状态"),
|
||||
skip: int = Query(0, description="跳过记录数"),
|
||||
limit: int = Query(100, description="返回记录数限制"),
|
||||
db: Session = Depends(get_db),
|
||||
@@ -77,11 +76,11 @@ async def get_public_providers(
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=ApiMode.PUBLIC)
|
||||
|
||||
|
||||
@router.get("/models", response_model=List[PublicModelResponse])
|
||||
@router.get("/models", response_model=list[PublicModelResponse])
|
||||
async def get_public_models(
|
||||
request: Request,
|
||||
provider_id: Optional[str] = Query(None, description="提供商ID过滤"),
|
||||
is_active: Optional[bool] = Query(None, description="过滤活跃状态"),
|
||||
provider_id: str | None = Query(None, description="提供商ID过滤"),
|
||||
is_active: bool | None = Query(None, description="过滤活跃状态"),
|
||||
skip: int = Query(0, description="跳过记录数"),
|
||||
limit: int = Query(100, description="返回记录数限制"),
|
||||
db: Session = Depends(get_db),
|
||||
@@ -145,7 +144,7 @@ async def get_public_stats(request: Request, db: Session = Depends(get_db)):
|
||||
async def search_models(
|
||||
request: Request,
|
||||
q: str = Query(..., description="搜索关键词"),
|
||||
provider_id: Optional[int] = Query(None, description="提供商ID过滤"),
|
||||
provider_id: int | None = Query(None, description="提供商ID过滤"),
|
||||
limit: int = Query(20, description="返回记录数限制"),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
@@ -234,8 +233,8 @@ async def get_public_global_models(
|
||||
request: Request,
|
||||
skip: int = Query(0, ge=0, description="跳过记录数"),
|
||||
limit: int = Query(100, ge=1, le=1000, description="返回记录数限制"),
|
||||
is_active: Optional[bool] = Query(None, description="过滤活跃状态"),
|
||||
search: Optional[str] = Query(None, description="搜索关键词"),
|
||||
is_active: bool | None = Query(None, description="过滤活跃状态"),
|
||||
search: str | None = Query(None, description="搜索关键词"),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
"""
|
||||
@@ -283,7 +282,7 @@ class PublicApiAdapter(ApiAdapter):
|
||||
|
||||
@dataclass
|
||||
class PublicProvidersAdapter(PublicApiAdapter):
|
||||
is_active: Optional[bool]
|
||||
is_active: bool | None
|
||||
skip: int
|
||||
limit: int
|
||||
|
||||
@@ -338,8 +337,8 @@ class PublicProvidersAdapter(PublicApiAdapter):
|
||||
|
||||
@dataclass
|
||||
class PublicModelsAdapter(PublicApiAdapter):
|
||||
provider_id: Optional[str]
|
||||
is_active: Optional[bool]
|
||||
provider_id: str | None
|
||||
is_active: bool | None
|
||||
skip: int
|
||||
limit: int
|
||||
|
||||
@@ -426,7 +425,7 @@ class PublicStatsAdapter(PublicApiAdapter):
|
||||
@dataclass
|
||||
class PublicSearchModelsAdapter(PublicApiAdapter):
|
||||
query: str
|
||||
provider_id: Optional[int]
|
||||
provider_id: int | None
|
||||
limit: int
|
||||
|
||||
async def handle(self, context): # type: ignore[override]
|
||||
@@ -508,7 +507,7 @@ class PublicApiFormatHealthMonitorAdapter(PublicApiAdapter):
|
||||
.all()
|
||||
)
|
||||
|
||||
all_formats: List[str] = []
|
||||
all_formats: list[str] = []
|
||||
for (api_format_enum,) in active_formats:
|
||||
api_format = (
|
||||
api_format_enum.value if hasattr(api_format_enum, "value") else str(api_format_enum)
|
||||
@@ -525,7 +524,7 @@ class PublicApiFormatHealthMonitorAdapter(PublicApiAdapter):
|
||||
)
|
||||
.all()
|
||||
)
|
||||
endpoint_map: Dict[str, List[str]] = defaultdict(list)
|
||||
endpoint_map: dict[str, list[str]] = defaultdict(list)
|
||||
for api_format_enum, endpoint_id in endpoint_rows:
|
||||
api_format = (
|
||||
api_format_enum.value if hasattr(api_format_enum, "value") else str(api_format_enum)
|
||||
@@ -551,7 +550,7 @@ class PublicApiFormatHealthMonitorAdapter(PublicApiAdapter):
|
||||
.all()
|
||||
)
|
||||
|
||||
grouped_candidates: Dict[str, List[RequestCandidate]] = {}
|
||||
grouped_candidates: dict[str, list[RequestCandidate]] = {}
|
||||
|
||||
for candidate, api_format_enum in rows:
|
||||
api_format = (
|
||||
@@ -564,7 +563,7 @@ class PublicApiFormatHealthMonitorAdapter(PublicApiAdapter):
|
||||
grouped_candidates[api_format].append(candidate)
|
||||
|
||||
# 3. 为所有活跃格式生成监控数据
|
||||
monitors: List[PublicApiFormatHealthMonitor] = []
|
||||
monitors: list[PublicApiFormatHealthMonitor] = []
|
||||
for api_format in all_formats:
|
||||
candidates = grouped_candidates.get(api_format, [])
|
||||
|
||||
@@ -579,7 +578,7 @@ class PublicApiFormatHealthMonitorAdapter(PublicApiAdapter):
|
||||
success_rate = success_count / actual_completed if actual_completed > 0 else 1.0
|
||||
|
||||
# 转换为公开版事件列表(不含敏感信息如 provider_id, key_id)
|
||||
events: List[PublicHealthEvent] = []
|
||||
events: list[PublicHealthEvent] = []
|
||||
for c in candidates:
|
||||
event_time = c.finished_at or c.started_at or c.created_at
|
||||
events.append(
|
||||
@@ -649,8 +648,8 @@ class PublicGlobalModelsAdapter(PublicApiAdapter):
|
||||
|
||||
skip: int
|
||||
limit: int
|
||||
is_active: Optional[bool]
|
||||
search: Optional[str]
|
||||
is_active: bool | None
|
||||
search: str | None
|
||||
|
||||
async def handle(self, context): # type: ignore[override]
|
||||
db = context.db
|
||||
|
||||
@@ -7,8 +7,6 @@
|
||||
- Authorization: Bearer (bearer) -> OpenAI 格式
|
||||
"""
|
||||
|
||||
from typing import Optional, Tuple, Union
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, Request
|
||||
from fastapi.responses import JSONResponse
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -35,7 +33,6 @@ from src.core.logger import logger
|
||||
from src.database import get_db
|
||||
from src.models.database import ApiKey, User
|
||||
from src.services.auth.service import AuthService
|
||||
from src.services.system.config import SystemConfigService
|
||||
|
||||
router = APIRouter(tags=["System Catalog"])
|
||||
|
||||
@@ -54,7 +51,7 @@ _ALL_CHAT_FORMATS = [
|
||||
|
||||
def _extract_api_key_from_request(
|
||||
request: Request, definition: ApiFormatDefinition
|
||||
) -> Optional[str]:
|
||||
) -> str | None:
|
||||
"""根据格式定义从请求中提取 API Key"""
|
||||
auth_header = definition.auth_header.lower()
|
||||
auth_type = definition.auth_type
|
||||
@@ -76,7 +73,7 @@ def _extract_api_key_from_request(
|
||||
return header_value
|
||||
|
||||
|
||||
def _detect_api_format_and_key(request: Request) -> Tuple[str, Optional[str]]:
|
||||
def _detect_api_format_and_key(request: Request) -> tuple[str, str | None]:
|
||||
"""
|
||||
根据请求头检测 API 格式并提取 API Key
|
||||
|
||||
@@ -163,7 +160,7 @@ def _build_empty_list_response(api_format: str) -> dict:
|
||||
|
||||
def _filter_formats_by_restrictions(
|
||||
formats: list[str], restrictions: AccessRestrictions, api_format: str
|
||||
) -> Tuple[list[str], Optional[dict]]:
|
||||
) -> tuple[list[str], dict | None]:
|
||||
"""
|
||||
根据访问限制过滤 API 格式
|
||||
|
||||
@@ -182,7 +179,7 @@ def _filter_formats_by_restrictions(
|
||||
return filtered, None
|
||||
|
||||
|
||||
def _authenticate(db: Session, api_key: Optional[str]) -> Tuple[Optional[User], Optional[ApiKey]]:
|
||||
def _authenticate(db: Session, api_key: str | None) -> tuple[User | None, ApiKey | None]:
|
||||
"""
|
||||
认证 API Key
|
||||
|
||||
@@ -248,8 +245,8 @@ def _build_auth_error_response(api_format: str) -> JSONResponse:
|
||||
|
||||
def _build_claude_list_response(
|
||||
models: list[ModelInfo],
|
||||
before_id: Optional[str],
|
||||
after_id: Optional[str],
|
||||
before_id: str | None,
|
||||
after_id: str | None,
|
||||
limit: int,
|
||||
) -> dict:
|
||||
"""构建 Claude 格式的列表响应"""
|
||||
@@ -309,7 +306,7 @@ def _build_openai_list_response(models: list[ModelInfo]) -> dict:
|
||||
def _build_gemini_list_response(
|
||||
models: list[ModelInfo],
|
||||
page_size: int,
|
||||
page_token: Optional[str],
|
||||
page_token: str | None,
|
||||
) -> dict:
|
||||
"""构建 Gemini 格式的列表响应"""
|
||||
# 处理分页
|
||||
@@ -435,14 +432,14 @@ def _build_404_response(model_id: str, api_format: str) -> JSONResponse:
|
||||
async def list_models(
|
||||
request: Request,
|
||||
# Claude 分页参数
|
||||
before_id: Optional[str] = Query(None, description="返回此 ID 之前的结果 (Claude)"),
|
||||
after_id: Optional[str] = Query(None, description="返回此 ID 之后的结果 (Claude)"),
|
||||
before_id: str | None = Query(None, description="返回此 ID 之前的结果 (Claude)"),
|
||||
after_id: str | None = Query(None, description="返回此 ID 之后的结果 (Claude)"),
|
||||
limit: int = Query(20, ge=1, le=1000, description="返回数量限制 (Claude)"),
|
||||
# Gemini 分页参数
|
||||
page_size: int = Query(50, alias="pageSize", ge=1, le=1000, description="每页数量 (Gemini)"),
|
||||
page_token: Optional[str] = Query(None, alias="pageToken", description="分页 token (Gemini)"),
|
||||
page_token: str | None = Query(None, alias="pageToken", description="分页 token (Gemini)"),
|
||||
db: Session = Depends(get_db),
|
||||
) -> Union[dict, JSONResponse]:
|
||||
) -> dict | JSONResponse:
|
||||
"""
|
||||
列出可用模型(统一端点)
|
||||
|
||||
@@ -556,7 +553,7 @@ async def retrieve_model(
|
||||
model_id: str,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
) -> Union[dict, JSONResponse]:
|
||||
) -> dict | JSONResponse:
|
||||
"""
|
||||
获取单个模型详情(统一端点)
|
||||
|
||||
@@ -658,9 +655,9 @@ async def retrieve_model(
|
||||
async def list_models_gemini(
|
||||
request: Request,
|
||||
page_size: int = Query(50, alias="pageSize", ge=1, le=1000),
|
||||
page_token: Optional[str] = Query(None, alias="pageToken"),
|
||||
page_token: str | None = Query(None, alias="pageToken"),
|
||||
db: Session = Depends(get_db),
|
||||
) -> Union[dict, JSONResponse]:
|
||||
) -> dict | JSONResponse:
|
||||
"""
|
||||
列出可用模型(Gemini v1beta 专用端点)
|
||||
|
||||
@@ -741,7 +738,7 @@ async def get_model_gemini(
|
||||
request: Request,
|
||||
model_name: str,
|
||||
db: Session = Depends(get_db),
|
||||
) -> Union[dict, JSONResponse]:
|
||||
) -> dict | JSONResponse:
|
||||
"""
|
||||
获取单个模型详情(Gemini v1beta 专用端点)
|
||||
|
||||
|
||||
@@ -1,12 +1,11 @@
|
||||
"""公开模块状态 API(供登录页等使用)"""
|
||||
|
||||
from typing import List
|
||||
|
||||
from fastapi import APIRouter, Depends
|
||||
from pydantic import BaseModel
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.core.modules import ModuleCategory, get_module_registry
|
||||
from src.core.modules import get_module_registry
|
||||
from src.database import get_db
|
||||
|
||||
router = APIRouter(prefix="/api/modules", tags=["Modules"])
|
||||
@@ -20,7 +19,7 @@ class AuthModuleInfo(BaseModel):
|
||||
active: bool
|
||||
|
||||
|
||||
@router.get("/auth-status", response_model=List[AuthModuleInfo])
|
||||
@router.get("/auth-status", response_model=list[AuthModuleInfo])
|
||||
async def get_auth_modules_status(db: Session = Depends(get_db)):
|
||||
"""
|
||||
获取认证模块状态(公开接口)
|
||||
|
||||
@@ -5,7 +5,7 @@ System Catalog / 健康检查相关端点
|
||||
"""
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
@@ -28,7 +28,7 @@ router = APIRouter(tags=["System Catalog"])
|
||||
# ============== 辅助函数 ==============
|
||||
|
||||
|
||||
def _as_bool(value: Optional[str], default: bool) -> bool:
|
||||
def _as_bool(value: str | None, default: bool) -> bool:
|
||||
"""将字符串转换为布尔值"""
|
||||
if value is None:
|
||||
return default
|
||||
@@ -39,9 +39,9 @@ def _serialize_provider(
|
||||
provider: Provider,
|
||||
include_models: bool,
|
||||
include_endpoints: bool,
|
||||
) -> Dict[str, Any]:
|
||||
) -> dict[str, Any]:
|
||||
"""序列化 Provider 对象"""
|
||||
provider_data: Dict[str, Any] = {
|
||||
provider_data: dict[str, Any] = {
|
||||
"id": provider.id,
|
||||
"name": provider.name,
|
||||
"is_active": provider.is_active,
|
||||
@@ -81,7 +81,7 @@ def _serialize_provider(
|
||||
return provider_data
|
||||
|
||||
|
||||
def _select_provider(db: Session, provider_name: Optional[str]) -> Optional[Provider]:
|
||||
def _select_provider(db: Session, provider_name: str | None) -> Provider | None:
|
||||
"""选择 Provider(按 provider_priority 优先级选择)"""
|
||||
query = db.query(Provider).filter(Provider.is_active == True)
|
||||
if provider_name:
|
||||
@@ -104,7 +104,7 @@ async def service_health(db: Session = Depends(get_db)):
|
||||
)
|
||||
active_models = db.query(func.count(Model.id)).filter(Model.is_active == True).scalar() or 0
|
||||
|
||||
redis_info: Dict[str, Any] = {"status": "unknown"}
|
||||
redis_info: dict[str, Any] = {"status": "unknown"}
|
||||
try:
|
||||
redis = await get_redis_client()
|
||||
if redis:
|
||||
@@ -245,9 +245,9 @@ async def provider_detail(
|
||||
async def test_connection(
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
provider: Optional[str] = Query(None),
|
||||
provider: str | None = Query(None),
|
||||
model: str = Query("claude-3-haiku-20240307"),
|
||||
api_format: Optional[str] = Query(None),
|
||||
api_format: str | None = Query(None),
|
||||
):
|
||||
"""测试 Provider 连接"""
|
||||
selected_provider = _select_provider(db, provider)
|
||||
|
||||
@@ -2,7 +2,6 @@
|
||||
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from fastapi.responses import JSONResponse
|
||||
@@ -55,13 +54,13 @@ class CreateManagementTokenRequest(BaseModel):
|
||||
"""创建 Management Token 请求"""
|
||||
|
||||
name: str = Field(..., min_length=1, max_length=100, description="Token 名称")
|
||||
description: Optional[str] = Field(None, max_length=500, description="描述")
|
||||
allowed_ips: Optional[list[str]] = Field(None, description="IP 白名单")
|
||||
expires_at: Optional[datetime] = Field(None, description="过期时间")
|
||||
description: str | None = Field(None, max_length=500, description="描述")
|
||||
allowed_ips: list[str] | None = Field(None, description="IP 白名单")
|
||||
expires_at: datetime | None = Field(None, description="过期时间")
|
||||
|
||||
@field_validator("allowed_ips")
|
||||
@classmethod
|
||||
def validate_allowed_ips(cls, v: Optional[list[str]]) -> Optional[list[str]]:
|
||||
def validate_allowed_ips(cls, v: list[str] | None) -> list[str] | None:
|
||||
return validate_ip_list(v)
|
||||
|
||||
@field_validator("expires_at", mode="before")
|
||||
@@ -81,10 +80,10 @@ class UpdateManagementTokenRequest(BaseModel):
|
||||
|
||||
model_config = {"extra": "allow"} # 允许额外字段以便检测哪些字段被显式提供
|
||||
|
||||
name: Optional[str] = Field(None, min_length=1, max_length=100)
|
||||
description: Optional[str] = Field(None, max_length=500)
|
||||
allowed_ips: Optional[list[str]] = None
|
||||
expires_at: Optional[datetime] = None
|
||||
name: str | None = Field(None, min_length=1, max_length=100)
|
||||
description: str | None = Field(None, max_length=500)
|
||||
allowed_ips: list[str] | None = None
|
||||
expires_at: datetime | None = None
|
||||
|
||||
# 用于追踪哪些字段被显式提供(包括显式设为 null 的情况)
|
||||
_provided_fields: set[str] = set()
|
||||
@@ -101,7 +100,7 @@ class UpdateManagementTokenRequest(BaseModel):
|
||||
|
||||
@field_validator("allowed_ips")
|
||||
@classmethod
|
||||
def validate_allowed_ips(cls, v: Optional[list[str]]) -> Optional[list[str]]:
|
||||
def validate_allowed_ips(cls, v: list[str] | None) -> list[str] | None:
|
||||
# 如果是 None,表示要清空,直接返回
|
||||
if v is None:
|
||||
return None
|
||||
@@ -122,7 +121,7 @@ class UpdateManagementTokenRequest(BaseModel):
|
||||
@router.get("")
|
||||
async def list_my_management_tokens(
|
||||
request: Request,
|
||||
is_active: Optional[bool] = Query(None, description="筛选激活状态"),
|
||||
is_active: bool | None = Query(None, description="筛选激活状态"),
|
||||
skip: int = Query(0, ge=0),
|
||||
limit: int = Query(50, ge=1, le=100),
|
||||
db: Session = Depends(get_db),
|
||||
@@ -347,7 +346,7 @@ class ListMyManagementTokensAdapter(ManagementTokenApiAdapter):
|
||||
"""列出用户的 Management Tokens"""
|
||||
|
||||
name: str = "list_my_management_tokens"
|
||||
is_active: Optional[bool] = None
|
||||
is_active: bool | None = None
|
||||
skip: int = 0
|
||||
limit: int = 50
|
||||
|
||||
|
||||
@@ -2,14 +2,12 @@
|
||||
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from pydantic import ValidationError
|
||||
from sqlalchemy import and_, func
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.base.adapter import ApiAdapter, ApiMode
|
||||
from src.api.base.authenticated_adapter import AuthenticatedApiAdapter
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.core.crypto import crypto_service
|
||||
@@ -170,9 +168,9 @@ async def toggle_my_api_key(key_id: str, request: Request, db: Session = Depends
|
||||
@router.get("/usage")
|
||||
async def get_my_usage(
|
||||
request: Request,
|
||||
start_date: Optional[datetime] = Query(None, description="开始时间(ISO 格式)"),
|
||||
end_date: Optional[datetime] = Query(None, description="结束时间(ISO 格式)"),
|
||||
search: Optional[str] = Query(None, description="搜索关键词(密钥名、模型名)"),
|
||||
start_date: datetime | None = Query(None, description="开始时间(ISO 格式)"),
|
||||
end_date: datetime | None = Query(None, description="结束时间(ISO 格式)"),
|
||||
search: str | None = Query(None, description="搜索关键词(密钥名、模型名)"),
|
||||
limit: int = Query(100, ge=1, le=200, description="每页记录数,默认100,最大200"),
|
||||
offset: int = Query(0, ge=0, le=2000, description="偏移量,用于分页,最大2000"),
|
||||
db: Session = Depends(get_db),
|
||||
@@ -200,7 +198,7 @@ async def get_my_usage(
|
||||
@router.get("/usage/active")
|
||||
async def get_my_active_requests(
|
||||
request: Request,
|
||||
ids: Optional[str] = Query(None, description="请求 ID 列表,逗号分隔"),
|
||||
ids: str | None = Query(None, description="请求 ID 列表,逗号分隔"),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
"""
|
||||
@@ -268,7 +266,7 @@ async def list_available_models(
|
||||
request: Request,
|
||||
skip: int = Query(0, ge=0, description="跳过记录数"),
|
||||
limit: int = Query(100, ge=1, le=1000, description="返回记录数限制"),
|
||||
search: Optional[str] = Query(None, description="搜索关键词"),
|
||||
search: str | None = Query(None, description="搜索关键词"),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
"""
|
||||
@@ -721,9 +719,9 @@ class ToggleMyApiKeyAdapter(AuthenticatedApiAdapter):
|
||||
class GetUsageAdapter(AuthenticatedApiAdapter):
|
||||
"""获取用户使用统计的适配器"""
|
||||
|
||||
start_date: Optional[datetime]
|
||||
end_date: Optional[datetime]
|
||||
search: Optional[str] = None
|
||||
start_date: datetime | None
|
||||
end_date: datetime | None
|
||||
search: str | None = None
|
||||
limit: int = 100
|
||||
offset: int = 0
|
||||
|
||||
@@ -983,7 +981,7 @@ class GetUsageAdapter(AuthenticatedApiAdapter):
|
||||
class GetActiveRequestsAdapter(AuthenticatedApiAdapter):
|
||||
"""轻量级活跃请求状态查询适配器(用于用户端轮询)"""
|
||||
|
||||
ids: Optional[str] = None
|
||||
ids: str | None = None
|
||||
|
||||
async def handle(self, context): # type: ignore[override]
|
||||
from src.services.usage import UsageService
|
||||
@@ -1045,7 +1043,7 @@ class ListAvailableModelsAdapter(AuthenticatedApiAdapter):
|
||||
|
||||
skip: int
|
||||
limit: int
|
||||
search: Optional[str]
|
||||
search: str | None
|
||||
|
||||
async def handle(self, context): # type: ignore[override]
|
||||
from sqlalchemy import or_
|
||||
@@ -1220,7 +1218,6 @@ class ListAvailableProvidersAdapter(AuthenticatedApiAdapter):
|
||||
async def handle(self, context): # type: ignore[override]
|
||||
from sqlalchemy.orm import selectinload
|
||||
|
||||
from src.models.database import ProviderEndpoint
|
||||
|
||||
db = context.db
|
||||
|
||||
|
||||
Reference in New Issue
Block a user