mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 17:30:23 +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:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user