fix: 修复 mypy 类型检查错误并升级到 Python 3.14

主要变更:
- 修复 1483 个 mypy 类型检查错误
- 添加缺失的类型注解 (Any, Callable, Session 等)
- 修复隐式 Optional 类型 (param: Type = None -> param: Type | None = None)
- 修复 __new__ 单例模式返回类型
- 添加 type: ignore 注释处理第三方库类型问题
- 更新 pyproject.toml 依赖到 Python 3.14 兼容版本
- 更新 mypy/black 配置为 Python 3.14
This commit is contained in:
fawney19
2026-01-30 14:30:57 +08:00
parent 7066166757
commit 5603c72f40
142 changed files with 2864 additions and 1853 deletions

View File

@@ -9,6 +9,9 @@
- adaptive_mode 是计算字段,基于 rpm_limit 是否为 NULL
"""
from __future__ import annotations
from typing import Any
from dataclasses import dataclass
from fastapi import APIRouter, Depends, HTTPException, Query, Request
@@ -21,6 +24,7 @@ from src.core.exceptions import InvalidRequestException, translate_pydantic_erro
from src.database import get_db
from src.models.database import ProviderAPIKey
from src.services.rate_limit.adaptive_rpm import get_adaptive_rpm_manager
from src.api.base.context import ApiRequestContext
router = APIRouter(prefix="/api/admin/adaptive", tags=["Adaptive RPM"])
pipeline = ApiRequestPipeline()
@@ -82,7 +86,7 @@ async def list_adaptive_keys(
request: Request,
provider_id: str | None = Query(None, description="按 Provider 过滤"),
db: Session = Depends(get_db),
):
) -> Any:
"""
获取所有启用自适应模式的Key列表
@@ -101,7 +105,7 @@ async def toggle_adaptive_mode(
key_id: str,
request: Request,
db: Session = Depends(get_db),
):
) -> Any:
"""
Toggle the RPM control mode for a specific key
@@ -122,7 +126,7 @@ async def get_adaptive_stats(
key_id: str,
request: Request,
db: Session = Depends(get_db),
):
) -> Any:
"""
获取指定Key的自适应 RPM 统计信息
@@ -144,7 +148,7 @@ async def reset_adaptive_learning(
key_id: str,
request: Request,
db: Session = Depends(get_db),
):
) -> Any:
"""
Reset the adaptive learning state for a specific key
@@ -169,7 +173,7 @@ async def set_rpm_limit(
request: Request,
limit: int = Query(..., ge=1, le=100, description="RPM limit value (1-100)"),
db: Session = Depends(get_db),
):
) -> Any:
"""
Set key to fixed RPM limit mode
@@ -188,7 +192,7 @@ async def set_rpm_limit(
async def get_adaptive_summary(
request: Request,
db: Session = Depends(get_db),
):
) -> Any:
"""
获取自适应 RPM 的全局统计摘要
@@ -208,7 +212,7 @@ async def get_adaptive_summary(
class ListAdaptiveKeysAdapter(AdminApiAdapter):
provider_id: str | None = None
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
# 自适应模式rpm_limit = NULL
query = context.db.query(ProviderAPIKey).filter(ProviderAPIKey.rpm_limit.is_(None))
if self.provider_id:
@@ -240,7 +244,7 @@ class ListAdaptiveKeysAdapter(AdminApiAdapter):
class ToggleAdaptiveModeAdapter(AdminApiAdapter):
key_id: str
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
key = context.db.query(ProviderAPIKey).filter(ProviderAPIKey.id == self.key_id).first()
if not key:
raise HTTPException(status_code=404, detail="Key not found")
@@ -288,7 +292,7 @@ class ToggleAdaptiveModeAdapter(AdminApiAdapter):
class GetAdaptiveStatsAdapter(AdminApiAdapter):
key_id: str
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
key = context.db.query(ProviderAPIKey).filter(ProviderAPIKey.id == self.key_id).first()
if not key:
raise HTTPException(status_code=404, detail="Key not found")
@@ -315,7 +319,7 @@ class GetAdaptiveStatsAdapter(AdminApiAdapter):
class ResetAdaptiveLearningAdapter(AdminApiAdapter):
key_id: str
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
key = context.db.query(ProviderAPIKey).filter(ProviderAPIKey.id == self.key_id).first()
if not key:
raise HTTPException(status_code=404, detail="Key not found")
@@ -330,7 +334,7 @@ class SetRPMLimitAdapter(AdminApiAdapter):
key_id: str
limit: int
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
key = context.db.query(ProviderAPIKey).filter(ProviderAPIKey.id == self.key_id).first()
if not key:
raise HTTPException(status_code=404, detail="Key not found")
@@ -350,7 +354,7 @@ class SetRPMLimitAdapter(AdminApiAdapter):
class AdaptiveSummaryAdapter(AdminApiAdapter):
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
# 自适应模式rpm_limit = NULL
adaptive_keys = (
context.db.query(ProviderAPIKey).filter(ProviderAPIKey.rpm_limit.is_(None)).all()

View File

@@ -3,6 +3,9 @@
独立余额Key不关联用户配额有独立余额限制用于给非注册用户使用。
"""
from __future__ import annotations
from typing import Any
import os
from datetime import datetime, timedelta, timezone
from zoneinfo import ZoneInfo
@@ -18,6 +21,7 @@ from src.database import get_db
from src.models.api import CreateApiKeyRequest
from src.models.database import ApiKey
from src.services.user.apikey import ApiKeyService
from src.api.base.context import ApiRequestContext
# 应用时区配置,默认为 Asia/Shanghai
@@ -71,7 +75,7 @@ async def list_standalone_api_keys(
limit: int = Query(100, ge=1, le=500),
is_active: bool | None = None,
db: Session = Depends(get_db),
):
) -> Any:
"""
列出所有独立余额 API Keys
@@ -101,7 +105,7 @@ async def create_standalone_api_key(
request: Request,
key_data: CreateApiKeyRequest,
db: Session = Depends(get_db),
):
) -> Any:
"""
创建独立余额 API Key
@@ -138,7 +142,7 @@ async def create_standalone_api_key(
@router.put("/{key_id}")
async def update_api_key(
key_id: str, request: Request, key_data: CreateApiKeyRequest, db: Session = Depends(get_db)
):
) -> Any:
"""
更新独立余额 API Key
@@ -174,7 +178,7 @@ async def update_api_key(
@router.patch("/{key_id}")
async def toggle_api_key(key_id: str, request: Request, db: Session = Depends(get_db)):
async def toggle_api_key(key_id: str, request: Request, db: Session = Depends(get_db)) -> Any:
"""
切换 API Key 启用状态
@@ -193,7 +197,7 @@ async def toggle_api_key(key_id: str, request: Request, db: Session = Depends(ge
@router.delete("/{key_id}")
async def delete_api_key(key_id: str, request: Request, db: Session = Depends(get_db)):
async def delete_api_key(key_id: str, request: Request, db: Session = Depends(get_db)) -> None:
"""
删除 API Key
@@ -210,7 +214,7 @@ async def delete_api_key(key_id: str, request: Request, db: Session = Depends(ge
@router.patch("/{key_id}/lock")
async def toggle_lock_api_key(key_id: str, request: Request, db: Session = Depends(get_db)):
async def toggle_lock_api_key(key_id: str, request: Request, db: Session = Depends(get_db)) -> Any:
"""
切换 API Key 锁定状态
@@ -233,7 +237,7 @@ async def add_balance_to_key(
key_id: str,
request: Request,
db: Session = Depends(get_db),
):
) -> Any:
"""
调整独立余额 API Key 的余额
@@ -295,7 +299,7 @@ async def get_api_key_detail(
request: Request,
include_key: bool = Query(False, description="Include full decrypted key in response"),
db: Session = Depends(get_db),
):
) -> Any:
"""
获取 API Key 详情
@@ -335,7 +339,7 @@ class AdminListStandaloneKeysAdapter(AdminApiAdapter):
self.limit = limit
self.is_active = is_active
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
# 只查询独立余额Keys
query = db.query(ApiKey).filter(ApiKey.is_standalone == True)
@@ -396,7 +400,7 @@ class AdminCreateStandaloneKeyAdapter(AdminApiAdapter):
def __init__(self, key_data: CreateApiKeyRequest):
self.key_data = key_data
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
# 独立Key必须设置初始余额
@@ -458,7 +462,7 @@ class AdminUpdateApiKeyAdapter(AdminApiAdapter):
self.key_id = key_id
self.key_data = key_data
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
api_key = db.query(ApiKey).filter(ApiKey.id == self.key_id).first()
if not api_key:
@@ -533,7 +537,7 @@ class AdminToggleApiKeyAdapter(AdminApiAdapter):
def __init__(self, key_id: str):
self.key_id = key_id
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
api_key = db.query(ApiKey).filter(ApiKey.id == self.key_id).first()
if not api_key:
@@ -566,7 +570,7 @@ class AdminToggleLockApiKeyAdapter(AdminApiAdapter):
def __init__(self, key_id: str):
self.key_id = key_id
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
api_key = db.query(ApiKey).filter(ApiKey.id == self.key_id).first()
if not api_key:
@@ -597,7 +601,7 @@ class AdminDeleteApiKeyAdapter(AdminApiAdapter):
def __init__(self, key_id: str):
self.key_id = key_id
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
api_key = db.query(ApiKey).filter(ApiKey.id == self.key_id).first()
if not api_key:
@@ -625,7 +629,7 @@ class AdminAddBalanceAdapter(AdminApiAdapter):
self.key_id = key_id
self.amount_usd = amount_usd
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
# 使用 ApiKeyService 增加余额
@@ -658,7 +662,7 @@ class AdminGetFullKeyAdapter(AdminApiAdapter):
def __init__(self, key_id: str):
self.key_id = key_id
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
from src.core.crypto import crypto_service
db = context.db
@@ -697,7 +701,7 @@ class AdminGetKeyDetailAdapter(AdminApiAdapter):
def __init__(self, key_id: str):
self.key_id = key_id
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
api_key = db.query(ApiKey).filter(ApiKey.id == self.key_id).first()

View File

@@ -2,6 +2,7 @@
Key RPM 限制管理 API
"""
from typing import Any
from dataclasses import dataclass
from fastapi import APIRouter, Depends, Request
@@ -14,6 +15,7 @@ from src.database import get_db
from src.models.database import ProviderAPIKey
from src.models.endpoint_models import KeyRpmStatusResponse
from src.services.rate_limit.concurrency_manager import get_concurrency_manager
from src.api.base.context import ApiRequestContext
router = APIRouter(tags=["RPM Control"])
pipeline = ApiRequestPipeline()
@@ -71,7 +73,7 @@ async def reset_key_rpm(
class AdminKeyRpmAdapter(AdminApiAdapter):
key_id: str
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == self.key_id).first()
if not key:
@@ -91,7 +93,7 @@ class AdminKeyRpmAdapter(AdminApiAdapter):
class AdminResetKeyRpmAdapter(AdminApiAdapter):
key_id: str
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
concurrency_manager = await get_concurrency_manager()
await concurrency_manager.reset_key_rpm(key_id=self.key_id)
return {"message": "RPM 计数已重置"}

View File

@@ -2,6 +2,9 @@
Endpoint 健康监控 API
"""
from __future__ import annotations
from typing import Any
from collections import defaultdict
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
@@ -25,6 +28,7 @@ from src.models.endpoint_models import (
)
from src.services.health.endpoint import EndpointHealthService
from src.services.health.monitor import health_monitor
from src.api.base.context import ApiRequestContext
router = APIRouter(tags=["Endpoint Health"])
pipeline = ApiRequestPipeline()
@@ -58,7 +62,7 @@ async def get_endpoint_health_status(
request: Request,
lookback_hours: int = Query(6, ge=1, le=72, description="回溯的小时数"),
db: Session = Depends(get_db),
):
) -> Any:
"""
获取端点健康状态(简化视图,与用户端点统一)
@@ -218,7 +222,7 @@ async def recover_all_keys_health(
class AdminHealthSummaryAdapter(AdminApiAdapter):
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
summary = health_monitor.get_all_health_status(context.db)
return HealthSummaryResponse(**summary)
@@ -229,7 +233,7 @@ class AdminEndpointHealthStatusAdapter(AdminApiAdapter):
lookback_hours: int
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
from src.services.health.endpoint import EndpointHealthService
db = context.db
@@ -256,7 +260,7 @@ class AdminApiFormatHealthMonitorAdapter(AdminApiAdapter):
lookback_hours: int
per_format_limit: int
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
now = datetime.now(timezone.utc)
since = now - timedelta(hours=self.lookback_hours)
@@ -463,7 +467,7 @@ class AdminKeyHealthAdapter(AdminApiAdapter):
key_id: str
api_format: str | None = None
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
health_data = health_monitor.get_key_health(context.db, self.key_id, self.api_format)
if not health_data:
raise NotFoundException(f"Key {self.key_id} 不存在")
@@ -501,7 +505,7 @@ class AdminRecoverKeyHealthAdapter(AdminApiAdapter):
key_id: str
api_format: str | None = None
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == self.key_id).first()
if not key:
@@ -544,7 +548,7 @@ class AdminRecoverKeyHealthAdapter(AdminApiAdapter):
class AdminRecoverAllKeysHealthAdapter(AdminApiAdapter):
"""批量恢复所有熔断 Key 的健康状态"""
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
# 查找所有有熔断格式的 Key检查 circuit_breaker_by_format JSON 字段)

View File

@@ -2,6 +2,9 @@
Provider API Keys 管理
"""
from __future__ import annotations
from typing import Any
import json
import uuid
from dataclasses import dataclass
@@ -25,6 +28,7 @@ from src.models.endpoint_models import (
EndpointAPIKeyUpdate,
)
from src.services.cache.provider_cache import ProviderCacheService
from src.api.base.context import ApiRequestContext
router = APIRouter(tags=["Provider Keys"])
pipeline = ApiRequestPipeline()
@@ -207,7 +211,7 @@ class AdminUpdateEndpointKeyAdapter(AdminApiAdapter):
key_id: str
key_data: EndpointAPIKeyUpdate
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == self.key_id).first()
if not key:
@@ -378,7 +382,7 @@ class AdminRevealEndpointKeyAdapter(AdminApiAdapter):
key_id: str
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == self.key_id).first()
if not key:
@@ -436,7 +440,7 @@ class AdminRevealEndpointKeyAdapter(AdminApiAdapter):
class AdminDeleteEndpointKeyAdapter(AdminApiAdapter):
key_id: str
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == self.key_id).first()
if not key:
@@ -472,7 +476,7 @@ class AdminDeleteEndpointKeyAdapter(AdminApiAdapter):
class AdminGetKeysGroupedByFormatAdapter(AdminApiAdapter):
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
# Key 属于 Provider按 key.api_formats 分组展示
@@ -677,7 +681,7 @@ class AdminListProviderKeysAdapter(AdminApiAdapter):
skip: int
limit: int
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
provider = db.query(Provider).filter(Provider.id == self.provider_id).first()
if not provider:
@@ -702,7 +706,7 @@ class AdminCreateProviderKeyAdapter(AdminApiAdapter):
provider_id: str
key_data: EndpointAPIKeyCreate
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
provider = db.query(Provider).filter(Provider.id == self.provider_id).first()
if not provider:

View File

@@ -2,6 +2,9 @@
ProviderEndpoint CRUD 管理 API
"""
from __future__ import annotations
from typing import Any
import uuid
from dataclasses import dataclass
from datetime import datetime, timezone
@@ -18,6 +21,7 @@ from src.core.exceptions import InvalidRequestException, NotFoundException
from src.core.logger import logger
from src.database import get_db
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
from src.api.base.context import ApiRequestContext
from src.models.endpoint_models import (
ProviderEndpointCreate,
ProviderEndpointResponse,
@@ -215,7 +219,7 @@ class AdminListProviderEndpointsAdapter(AdminApiAdapter):
skip: int
limit: int
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
provider = db.query(Provider).filter(Provider.id == self.provider_id).first()
if not provider:
@@ -270,7 +274,7 @@ class AdminCreateProviderEndpointAdapter(AdminApiAdapter):
provider_id: str
endpoint_data: ProviderEndpointCreate
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
provider = db.query(Provider).filter(Provider.id == self.provider_id).first()
if not provider:
@@ -340,7 +344,7 @@ class AdminCreateProviderEndpointAdapter(AdminApiAdapter):
class AdminGetProviderEndpointAdapter(AdminApiAdapter):
endpoint_id: str
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
endpoint = (
db.query(ProviderEndpoint, Provider)
@@ -390,7 +394,7 @@ class AdminUpdateProviderEndpointAdapter(AdminApiAdapter):
endpoint_id: str
endpoint_data: ProviderEndpointUpdate
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
endpoint = (
db.query(ProviderEndpoint).filter(ProviderEndpoint.id == self.endpoint_id).first()
@@ -459,7 +463,7 @@ class AdminUpdateProviderEndpointAdapter(AdminApiAdapter):
class AdminDeleteProviderEndpointAdapter(AdminApiAdapter):
endpoint_id: str
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
endpoint = (
db.query(ProviderEndpoint).filter(ProviderEndpoint.id == self.endpoint_id).first()

View File

@@ -1,5 +1,7 @@
"""LDAP配置管理API端点。"""
from __future__ import annotations
import re
from typing import Any
@@ -15,6 +17,7 @@ from src.core.exceptions import InvalidRequestException, translate_pydantic_erro
from src.core.logger import logger
from src.database import get_db
from src.models.database import AuditEventType, LDAPConfig, User, UserRole
from src.api.base.context import ApiRequestContext
from src.services.system.audit import AuditService
router = APIRouter(prefix="/api/admin/ldap", tags=["Admin - LDAP"])
@@ -263,7 +266,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: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
db = context.db
config = db.query(LDAPConfig).first()
@@ -300,7 +303,7 @@ class AdminGetLDAPConfigAdapter(AdminApiAdapter):
class AdminUpdateLDAPConfigAdapter(AdminApiAdapter):
async def handle(self, context) -> dict[str, str]: # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> dict[str, str]: # type: ignore[override]
db = context.db
payload = context.ensure_json_body()
@@ -421,7 +424,7 @@ class AdminUpdateLDAPConfigAdapter(AdminApiAdapter):
class AdminTestLDAPConnectionAdapter(AdminApiAdapter):
async def handle(self, context) -> dict[str, Any]: # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
from src.services.auth.ldap import LDAPService
db = context.db

View File

@@ -1,5 +1,8 @@
"""管理员 Management Token 管理端点"""
from __future__ import annotations
from typing import Any
from dataclasses import dataclass
from fastapi import APIRouter, Depends, HTTPException, Query, Request
@@ -50,7 +53,7 @@ async def list_all_management_tokens(
skip: int = Query(0, ge=0),
limit: int = Query(50, ge=1, le=100),
db: Session = Depends(get_db),
):
) -> Any:
"""列出所有 Management Tokens管理员
管理员查看所有用户的 Management Tokens支持筛选和分页。
@@ -90,7 +93,7 @@ async def get_management_token(
token_id: str,
request: Request,
db: Session = Depends(get_db),
):
) -> Any:
"""获取 Management Token 详情(管理员)
管理员查看任意 Management Token 的详细信息。
@@ -119,7 +122,7 @@ async def get_management_token(
@router.delete("/{token_id}")
async def delete_management_token(
token_id: str, request: Request, db: Session = Depends(get_db)
):
) -> Any:
"""删除任意 Management Token管理员
管理员可以删除任意用户的 Management Token。
@@ -137,7 +140,7 @@ async def delete_management_token(
@router.patch("/{token_id}/status")
async def toggle_management_token(
token_id: str, request: Request, db: Session = Depends(get_db)
):
) -> Any:
"""切换任意 Management Token 状态(管理员)
管理员可以启用/禁用任意用户的 Management Token。
@@ -178,7 +181,7 @@ class AdminListManagementTokensAdapter(AdminManagementTokenApiAdapter):
skip: int = 0
limit: int = 50
async def handle(self, context: ApiRequestContext):
async def handle(self, context: ApiRequestContext) -> Any:
# 构建查询
query = context.db.query(ManagementToken)
@@ -218,7 +221,7 @@ class AdminGetManagementTokenAdapter(AdminManagementTokenApiAdapter):
name: str = "admin_get_management_token"
token_id: str = ""
async def handle(self, context: ApiRequestContext):
async def handle(self, context: ApiRequestContext) -> Any:
token = ManagementTokenService.get_token_by_id(
db=context.db, token_id=self.token_id
)
@@ -240,7 +243,7 @@ class AdminDeleteManagementTokenAdapter(AdminManagementTokenApiAdapter):
token_id: str = ""
audit_success_event = AuditEventType.MANAGEMENT_TOKEN_DELETED
async def handle(self, context: ApiRequestContext):
async def handle(self, context: ApiRequestContext) -> Any:
# 先获取 token 信息用于审计
token = ManagementTokenService.get_token_by_id(
db=context.db, token_id=self.token_id
@@ -273,7 +276,7 @@ class AdminToggleManagementTokenAdapter(AdminManagementTokenApiAdapter):
token_id: str = ""
audit_success_event = AuditEventType.MANAGEMENT_TOKEN_UPDATED
async def handle(self, context: ApiRequestContext):
async def handle(self, context: ApiRequestContext) -> Any:
token = ManagementTokenService.toggle_status(
db=context.db, token_id=self.token_id
)

View File

@@ -4,6 +4,7 @@
基于 GlobalModel 的聚合视图
"""
from typing import Any
from dataclasses import dataclass
from fastapi import APIRouter, Depends, Request
@@ -13,6 +14,7 @@ from src.api.base.admin_adapter import AdminApiAdapter
from src.api.base.pipeline import ApiRequestPipeline
from src.database import get_db
from src.models.database import GlobalModel, Model
from src.api.base.context import ApiRequestContext
from src.models.pydantic_models import (
ModelCapabilities,
ModelCatalogItem,
@@ -59,7 +61,7 @@ class AdminGetModelCatalogAdapter(AdminApiAdapter):
2. Model 表提供关联提供商和价格
"""
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db: Session = context.db
# 1. 获取所有活跃的 GlobalModel

View File

@@ -4,9 +4,12 @@ GlobalModel Admin API
提供 GlobalModel 的 CRUD 操作接口
"""
from __future__ import annotations
from typing import Any
from dataclasses import dataclass
from fastapi import APIRouter, Depends, Query, Request
from fastapi import APIRouter, Depends, Query, Request, Response
from sqlalchemy.orm import Session
from src.api.base.admin_adapter import AdminApiAdapter
@@ -26,6 +29,7 @@ from src.models.pydantic_models import (
ModelCatalogProviderDetail,
)
from src.services.model.global_model import GlobalModelService
from src.api.base.context import ApiRequestContext
router = APIRouter(prefix="/global", tags=["Admin - Global Models"])
pipeline = ApiRequestPipeline()
@@ -154,12 +158,12 @@ async def update_global_model(
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.delete("/{global_model_id}", status_code=204)
@router.delete("/{global_model_id}", status_code=204, response_class=Response)
async def delete_global_model(
request: Request,
global_model_id: str,
db: Session = Depends(get_db),
):
) -> Response:
"""
删除 GlobalModel
@@ -174,7 +178,7 @@ async def delete_global_model(
"""
adapter = AdminDeleteGlobalModelAdapter(global_model_id=global_model_id)
await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
return None
return Response(status_code=204)
@router.post(
@@ -256,7 +260,7 @@ class AdminListGlobalModelsAdapter(AdminApiAdapter):
is_active: bool | None
search: str | None
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
from sqlalchemy import func
from src.models.database import Model
@@ -302,7 +306,7 @@ class AdminGetGlobalModelAdapter(AdminApiAdapter):
global_model_id: str
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
global_model = GlobalModelService.get_global_model(context.db, self.global_model_id)
stats = GlobalModelService.get_global_model_stats(context.db, self.global_model_id)
@@ -320,7 +324,7 @@ class AdminCreateGlobalModelAdapter(AdminApiAdapter):
payload: GlobalModelCreate
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
from src.core.exceptions import InvalidRequestException
from src.core.model_permissions import validate_and_extract_model_mappings
@@ -359,7 +363,7 @@ class AdminUpdateGlobalModelAdapter(AdminApiAdapter):
global_model_id: str
payload: GlobalModelUpdate
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
from src.core.exceptions import InvalidRequestException
from src.core.model_permissions import validate_and_extract_model_mappings
@@ -418,7 +422,7 @@ class AdminDeleteGlobalModelAdapter(AdminApiAdapter):
global_model_id: str
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
# 使用行级锁获取 GlobalModel 信息,防止并发操作导致的竞态条件
# 设置 2 秒锁超时,允许短暂等待而非立即失败
from sqlalchemy import text
@@ -467,7 +471,7 @@ class AdminBatchAssignToProvidersAdapter(AdminApiAdapter):
global_model_id: str
payload: BatchAssignToProvidersRequest
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
result = GlobalModelService.batch_assign_to_providers(
db=context.db,
global_model_id=self.global_model_id,
@@ -492,7 +496,7 @@ class AdminGetGlobalModelProvidersAdapter(AdminApiAdapter):
global_model_id: str
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
from sqlalchemy.orm import joinedload
from src.models.database import Model

View File

@@ -8,6 +8,8 @@ GlobalModel 请求链路预览 API
- Key 的并发配置和健康状态
"""
from __future__ import annotations
from dataclasses import dataclass
from fastapi import APIRouter, Depends, Request
@@ -30,6 +32,7 @@ from src.models.database import (
ProviderEndpoint,
)
from src.services.cache.aware_scheduler import CacheAwareScheduler
from src.api.base.context import ApiRequestContext
from src.services.system.config import SystemConfigService
router = APIRouter(prefix="/global", tags=["Admin - Global Models"])
@@ -204,7 +207,7 @@ class AdminGetModelRoutingPreviewAdapter(AdminApiAdapter):
global_model_id: str
async def handle(self, context) -> ModelRoutingPreviewResponse: # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> ModelRoutingPreviewResponse: # type: ignore[override]
db = context.db
# 获取 GlobalModel

View File

@@ -12,6 +12,7 @@ from src.api.base.admin_adapter import AdminApiAdapter
from src.api.base.pipeline import ApiRequestPipeline
from src.core.exceptions import InvalidRequestException, NotFoundException
from src.core.modules import ModuleStatus, get_module_registry
from src.api.base.context import ApiRequestContext
from src.database import get_db
router = APIRouter(prefix="/api/admin/modules", tags=["Admin - Modules"])
@@ -69,7 +70,7 @@ class SetModuleEnabledRequest(BaseModel):
@router.get("/status")
async def get_all_modules_status(request: Request, db: Session = Depends(get_db)):
async def get_all_modules_status(request: Request, db: Session = Depends(get_db)) -> Any:
"""
获取所有模块状态
@@ -86,7 +87,7 @@ async def get_all_modules_status(request: Request, db: Session = Depends(get_db)
@router.get("/status/{module_name}")
async def get_module_status(
module_name: str, request: Request, db: Session = Depends(get_db)
):
) -> Any:
"""
获取单个模块状态
@@ -105,7 +106,7 @@ async def get_module_status(
@router.put("/status/{module_name}/enabled")
async def set_module_enabled(
module_name: str, request: Request, db: Session = Depends(get_db)
):
) -> Any:
"""
设置模块启用状态
@@ -131,7 +132,7 @@ async def set_module_enabled(
class AdminGetAllModulesStatusAdapter(AdminApiAdapter):
"""获取所有模块状态"""
async def handle(self, context) -> dict[str, Any]:
async def handle(self, context: ApiRequestContext) -> dict[str, Any]:
registry = get_module_registry()
all_status = await registry.get_all_status_async(context.db)
@@ -147,7 +148,7 @@ class AdminGetModuleStatusAdapter(AdminApiAdapter):
module_name: str
async def handle(self, context) -> dict[str, Any]:
async def handle(self, context: ApiRequestContext) -> dict[str, Any]:
registry = get_module_registry()
status = await registry.get_module_status_async(self.module_name, context.db)
@@ -163,7 +164,7 @@ class AdminSetModuleEnabledAdapter(AdminApiAdapter):
module_name: str
async def handle(self, context) -> dict[str, Any]:
async def handle(self, context: ApiRequestContext) -> dict[str, Any]:
registry = get_module_registry()
# 检查模块是否存在

View File

@@ -1,5 +1,8 @@
"""管理员监控与审计端点。"""
from __future__ import annotations
from typing import Any
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
@@ -23,6 +26,7 @@ from src.models.database import User as DBUser
from src.services.health.monitor import HealthMonitor
from src.services.system.audit import audit_service
from src.utils.database_helpers import escape_like_pattern
from src.api.base.context import ApiRequestContext
router = APIRouter(prefix="/api/admin/monitoring", tags=["Admin - Monitoring"])
@@ -38,7 +42,7 @@ async def get_audit_logs(
limit: int = Query(100, description="返回数量限制"),
offset: int = Query(0, description="偏移量"),
db: Session = Depends(get_db),
):
) -> Any:
"""
获取审计日志
@@ -78,7 +82,7 @@ async def get_audit_logs(
@router.get("/system-status")
async def get_system_status(request: Request, db: Session = Depends(get_db)):
async def get_system_status(request: Request, db: Session = Depends(get_db)) -> Any:
"""
获取系统状态
@@ -101,7 +105,7 @@ async def get_suspicious_activities(
request: Request,
hours: int = Query(24, description="时间范围(小时)"),
db: Session = Depends(get_db),
):
) -> Any:
"""
获取可疑活动记录
@@ -132,7 +136,7 @@ async def analyze_user_behavior(
request: Request,
days: int = Query(30, description="分析天数"),
db: Session = Depends(get_db),
):
) -> Any:
"""
分析用户行为
@@ -152,7 +156,7 @@ async def analyze_user_behavior(
@router.get("/resilience-status")
async def get_resilience_status(request: Request, db: Session = Depends(get_db)):
async def get_resilience_status(request: Request, db: Session = Depends(get_db)) -> Any:
"""
获取韧性系统状态
@@ -171,7 +175,7 @@ async def get_resilience_status(request: Request, db: Session = Depends(get_db))
@router.delete("/resilience/error-stats")
async def reset_error_stats(request: Request, db: Session = Depends(get_db)):
async def reset_error_stats(request: Request, db: Session = Depends(get_db)) -> None:
"""
重置错误统计
@@ -192,7 +196,7 @@ async def get_circuit_history(
request: Request,
limit: int = Query(50, ge=1, le=200),
db: Session = Depends(get_db),
):
) -> Any:
"""
获取熔断器历史记录
@@ -220,7 +224,7 @@ class AdminGetAuditLogsAdapter(AdminApiAdapter):
# 查看审计日志本身不应该产生审计记录,避免刷新页面时产生大量无意义的日志
audit_log_enabled: bool = False
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
cutoff_time = datetime.now(timezone.utc) - timedelta(days=self.days)
@@ -284,7 +288,7 @@ class AdminGetAuditLogsAdapter(AdminApiAdapter):
class AdminSystemStatusAdapter(AdminApiAdapter):
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
total_users = db.query(func.count(DBUser.id)).scalar()
@@ -361,7 +365,7 @@ class AdminSystemStatusAdapter(AdminApiAdapter):
class AdminSuspiciousActivitiesAdapter(AdminApiAdapter):
hours: int
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
activities = audit_service.get_suspicious_activities(db=db, hours=self.hours, limit=100)
response = {
@@ -393,7 +397,7 @@ class AdminUserBehaviorAdapter(AdminApiAdapter):
user_id: str
days: int
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
result = audit_service.analyze_user_behavior(
db=context.db,
user_id=self.user_id,
@@ -409,7 +413,7 @@ class AdminUserBehaviorAdapter(AdminApiAdapter):
class AdminResilienceStatusAdapter(AdminApiAdapter):
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
try:
from src.core.resilience import resilience_manager
except ImportError as exc:
@@ -454,7 +458,7 @@ class AdminResilienceStatusAdapter(AdminApiAdapter):
class AdminResetErrorStatsAdapter(AdminApiAdapter):
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
try:
from src.core.resilience import resilience_manager
except ImportError as exc:
@@ -486,7 +490,7 @@ class AdminCircuitHistoryAdapter(AdminApiAdapter):
super().__init__()
self.limit = limit
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
history = HealthMonitor.get_circuit_history(self.limit)
context.add_audit_metadata(
action="circuit_history",

View File

@@ -4,6 +4,8 @@
提供缓存亲和性统计、管理和监控功能
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import Any

View File

@@ -2,6 +2,9 @@
请求链路追踪 API 端点
"""
from __future__ import annotations
from typing import Any
from dataclasses import dataclass
from datetime import datetime
@@ -15,6 +18,7 @@ from src.database import get_db
from src.models.database import Provider, ProviderEndpoint, ProviderAPIKey
from src.core.crypto import crypto_service
from src.services.request.candidate import RequestCandidateService
from src.api.base.context import ApiRequestContext
router = APIRouter(prefix="/api/admin/monitoring/trace", tags=["Admin - Monitoring: Trace"])
pipeline = ApiRequestPipeline()
@@ -69,7 +73,7 @@ async def get_request_trace(
request_id: str,
request: Request,
db: Session = Depends(get_db),
):
) -> Any:
"""
获取请求的完整追踪信息
@@ -122,7 +126,7 @@ async def get_provider_failure_rate(
request: Request,
limit: int = Query(100, ge=1, le=1000, description="统计最近的尝试数量"),
db: Session = Depends(get_db),
):
) -> Any:
"""
获取提供商的失败率统计
@@ -153,7 +157,7 @@ async def get_provider_failure_rate(
class AdminGetRequestTraceAdapter(AdminApiAdapter):
request_id: str
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
# 只查询 candidates
@@ -326,7 +330,7 @@ class AdminProviderFailureRateAdapter(AdminApiAdapter):
provider_id: str
limit: int
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
result = RequestCandidateService.get_candidate_stats_by_provider(
db=context.db,
provider_id=self.provider_id,

View File

@@ -8,6 +8,8 @@ Provider 操作 API 路由
- 配置管理
"""
from __future__ import annotations
from dataclasses import asdict, is_dataclass
from typing import Any
@@ -146,14 +148,14 @@ def _serialize_data(data: Any) -> Any:
@router.get("/architectures", response_model=list[ArchitectureInfo])
async def list_architectures(_: User = Depends(require_admin)):
async def list_architectures(_: User = Depends(require_admin)) -> Any:
"""获取所有可用的架构"""
registry = get_registry()
return registry.to_dict_list()
@router.get("/architectures/{architecture_id}", response_model=ArchitectureInfo)
async def get_architecture(architecture_id: str, _: User = Depends(require_admin)):
async def get_architecture(architecture_id: str, _: User = Depends(require_admin)) -> Any:
"""获取指定架构的详情"""
registry = get_registry()
arch = registry.get(architecture_id)
@@ -167,7 +169,7 @@ async def get_provider_ops_status(
provider_id: str,
db: Session = Depends(get_db),
_: User = Depends(require_admin),
):
) -> Any:
"""获取 Provider 的操作状态"""
service = ProviderOpsService(db)
@@ -200,7 +202,7 @@ async def get_provider_ops_config(
provider_id: str,
db: Session = Depends(get_db),
_: User = Depends(require_admin),
):
) -> Any:
"""
获取 Provider 的操作配置(脱敏)
@@ -251,7 +253,7 @@ async def save_provider_ops_config(
request: SaveConfigRequest,
db: Session = Depends(get_db),
_: User = Depends(require_admin),
):
) -> Any:
"""保存 Provider 的操作配置"""
service = ProviderOpsService(db)
@@ -287,7 +289,7 @@ async def verify_provider_auth(
request: SaveConfigRequest,
db: Session = Depends(get_db),
_: User = Depends(require_admin),
):
) -> Any:
"""
验证 Provider 认证配置
@@ -343,7 +345,7 @@ async def delete_provider_ops_config(
provider_id: str,
db: Session = Depends(get_db),
_: User = Depends(require_admin),
):
) -> Any:
"""删除 Provider 的操作配置"""
service = ProviderOpsService(db)
success = service.delete_config(provider_id)
@@ -360,7 +362,7 @@ async def connect_provider(
request: ConnectRequest,
db: Session = Depends(get_db),
_: User = Depends(require_admin),
):
) -> Any:
"""建立与 Provider 的连接"""
service = ProviderOpsService(db)
@@ -377,7 +379,7 @@ async def disconnect_provider(
provider_id: str,
db: Session = Depends(get_db),
_: User = Depends(require_admin),
):
) -> Any:
"""断开与 Provider 的连接"""
service = ProviderOpsService(db)
await service.disconnect(provider_id)
@@ -395,7 +397,7 @@ async def execute_action(
request: ExecuteActionRequest,
db: Session = Depends(get_db),
_: User = Depends(require_admin),
):
) -> Any:
"""执行指定操作"""
service = ProviderOpsService(db)
@@ -423,7 +425,7 @@ async def get_balance(
refresh: bool = True,
db: Session = Depends(get_db),
_: User = Depends(require_admin),
):
) -> Any:
"""
获取余额(优先返回缓存,后台异步刷新)
@@ -449,7 +451,7 @@ async def refresh_balance(
provider_id: str,
db: Session = Depends(get_db),
_: User = Depends(require_admin),
):
) -> Any:
"""立即刷新余额(同步等待结果)"""
service = ProviderOpsService(db)
result = await service.query_balance(provider_id)
@@ -470,7 +472,7 @@ async def checkin(
provider_id: str,
db: Session = Depends(get_db),
_: User = Depends(require_admin),
):
) -> Any:
"""签到(快捷方法)"""
service = ProviderOpsService(db)
result = await service.checkin(provider_id)
@@ -491,7 +493,7 @@ async def batch_query_balance(
provider_ids: list[str] | None = None,
db: Session = Depends(get_db),
_: User = Depends(require_admin),
):
) -> Any:
"""批量查询余额"""
service = ProviderOpsService(db)
results = await service.batch_query_balance(provider_ids)

View File

@@ -3,6 +3,9 @@ Provider Query API 端点
用于查询提供商的模型列表等信息
"""
from __future__ import annotations
from typing import Any
import asyncio
import httpx
@@ -63,7 +66,7 @@ async def query_available_models(
request: ModelsQueryRequest,
db: Session = Depends(get_db),
current_user: User = Depends(get_current_user),
):
) -> Any:
"""
查询提供商可用模型
@@ -118,7 +121,7 @@ async def query_available_models(
raise HTTPException(status_code=400, detail="No active API Key found for this provider")
# 并发获取所有 Key 的模型
async def fetch_for_key(api_key):
async def fetch_for_key(api_key: Any) -> Any:
# 非强制刷新时,先检查缓存
if not request.force_refresh:
cached_models = await get_upstream_models_from_cache(
@@ -246,7 +249,7 @@ async def _fetch_models_for_single_key(
api_key_id: str,
format_to_endpoint: dict[str, ProviderEndpoint],
force_refresh: bool,
):
) -> Any:
"""获取单个 Key 的模型列表"""
# 查找指定的 Key
api_key = next(
@@ -305,7 +308,7 @@ async def test_model(
request: TestModelRequest,
db: Session = Depends(get_db),
current_user: User = Depends(get_current_user),
):
) -> Any:
"""
测试模型连接性

View File

@@ -2,6 +2,9 @@
提供商策略管理 API 端点
"""
from __future__ import annotations
from typing import Any
from datetime import datetime, timedelta, timezone
from fastapi import APIRouter, Depends, HTTPException, Request
@@ -16,6 +19,7 @@ from src.core.exceptions import InvalidRequestException, translate_pydantic_erro
from src.core.logger import logger
from src.database import get_db
from src.models.database import Provider
from src.api.base.context import ApiRequestContext
from src.models.database_extensions import ProviderUsageTracking
router = APIRouter(prefix="/api/admin/provider-strategy", tags=["Provider Strategy"])
@@ -37,7 +41,7 @@ async def update_provider_billing(
provider_id: str,
request: Request,
db: Session = Depends(get_db),
):
) -> Any:
"""
更新提供商计费配置
@@ -73,7 +77,7 @@ async def get_provider_stats(
request: Request,
hours: int = 24,
db: Session = Depends(get_db),
):
) -> Any:
"""
获取提供商统计数据
@@ -112,14 +116,14 @@ async def reset_provider_quota(
provider_id: str,
request: Request,
db: Session = Depends(get_db),
):
) -> Any:
"""Reset provider quota usage to zero"""
adapter = AdminProviderResetQuotaAdapter(provider_id=provider_id)
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.get("/strategies")
async def list_available_strategies(request: Request, db: Session = Depends(get_db)):
async def list_available_strategies(request: Request, db: Session = Depends(get_db)) -> Any:
"""
获取可用负载均衡策略列表
@@ -142,7 +146,7 @@ class AdminProviderBillingAdapter(AdminApiAdapter):
def __init__(self, provider_id: str):
self.provider_id = provider_id
async def handle(self, context):
async def handle(self, context: ApiRequestContext) -> Any:
db = context.db
payload = context.ensure_json_body()
try:
@@ -215,7 +219,7 @@ class AdminProviderStatsAdapter(AdminApiAdapter):
self.provider_id = provider_id
self.hours = hours
async def handle(self, context):
async def handle(self, context: ApiRequestContext) -> Any:
db = context.db
provider = db.query(Provider).filter(Provider.id == self.provider_id).first()
if not provider:
@@ -271,7 +275,7 @@ class AdminProviderResetQuotaAdapter(AdminApiAdapter):
def __init__(self, provider_id: str):
self.provider_id = provider_id
async def handle(self, context):
async def handle(self, context: ApiRequestContext) -> Any:
db = context.db
provider = db.query(Provider).filter(Provider.id == self.provider_id).first()
if not provider:
@@ -297,7 +301,7 @@ class AdminProviderResetQuotaAdapter(AdminApiAdapter):
class AdminListStrategiesAdapter(AdminApiAdapter):
async def handle(self, context):
async def handle(self, context: ApiRequestContext) -> Any:
from src.plugins.manager import get_plugin_manager
plugin_manager = get_plugin_manager()

View File

@@ -2,6 +2,8 @@
Provider 模型管理 API
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import Any
@@ -35,6 +37,7 @@ from src.models.database import (
Provider,
)
from src.services.model.service import ModelService
from src.api.base.context import ApiRequestContext
router = APIRouter(tags=["Model Management"])
pipeline = ApiRequestPipeline()
@@ -205,7 +208,7 @@ async def delete_provider_model(
model_id: str,
request: Request,
db: Session = Depends(get_db),
):
) -> Any:
"""
删除模型
@@ -264,7 +267,7 @@ async def get_provider_available_source_models(
provider_id: str,
request: Request,
db: Session = Depends(get_db),
):
) -> Any:
"""
获取提供商支持的可用源模型
@@ -379,7 +382,7 @@ class AdminListProviderModelsAdapter(AdminApiAdapter):
skip: int
limit: int
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
provider = db.query(Provider).filter(Provider.id == self.provider_id).first()
if not provider:
@@ -396,7 +399,7 @@ class AdminCreateProviderModelAdapter(AdminApiAdapter):
provider_id: str
model_data: ModelCreate
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
provider = db.query(Provider).filter(Provider.id == self.provider_id).first()
if not provider:
@@ -416,7 +419,7 @@ class AdminGetProviderModelAdapter(AdminApiAdapter):
provider_id: str
model_id: str
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
model = (
db.query(Model)
@@ -435,7 +438,7 @@ class AdminUpdateProviderModelAdapter(AdminApiAdapter):
model_id: str
model_data: ModelUpdate
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
model = (
db.query(Model)
@@ -459,7 +462,7 @@ class AdminDeleteProviderModelAdapter(AdminApiAdapter):
provider_id: str
model_id: str
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
model = (
db.query(Model)
@@ -484,7 +487,7 @@ class AdminBatchCreateModelsAdapter(AdminApiAdapter):
provider_id: str
models_data: list[ModelCreate]
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
provider = db.query(Provider).filter(Provider.id == self.provider_id).first()
if not provider:
@@ -503,7 +506,7 @@ class AdminBatchCreateModelsAdapter(AdminApiAdapter):
class AdminGetProviderAvailableSourceModelsAdapter(AdminApiAdapter):
provider_id: str
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
"""
返回 Provider 支持的所有 GlobalModel
@@ -571,7 +574,7 @@ class AdminBatchAssignModelsToProviderAdapter(AdminApiAdapter):
provider_id: str
payload: BatchAssignModelsToProviderRequest
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
provider = db.query(Provider).filter(Provider.id == self.provider_id).first()
if not provider:
@@ -654,7 +657,7 @@ class AdminImportFromUpstreamAdapter(AdminApiAdapter):
provider_id: str
payload: ImportFromUpstreamRequest
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
provider = db.query(Provider).filter(Provider.id == self.provider_id).first()
if not provider:

View File

@@ -1,5 +1,8 @@
"""管理员 Provider 管理路由。"""
from __future__ import annotations
from typing import Any
import asyncio
from datetime import datetime, timezone
@@ -19,6 +22,7 @@ from src.database import get_db
from src.models.admin_requests import CreateProviderRequest, UpdateProviderRequest
from src.models.database import GlobalModel, Provider, ProviderAPIKey
from src.services.cache.provider_cache import ProviderCacheService
from src.api.base.context import ApiRequestContext
router = APIRouter(tags=["Provider CRUD"])
pipeline = ApiRequestPipeline()
@@ -96,7 +100,7 @@ async def list_providers(
limit: int = Query(100, ge=1, le=500),
is_active: bool | None = None,
db: Session = Depends(get_db),
):
) -> Any:
"""
获取提供商列表
@@ -123,7 +127,7 @@ async def list_providers(
@router.post("/")
async def create_provider(request: Request, db: Session = Depends(get_db)):
async def create_provider(request: Request, db: Session = Depends(get_db)) -> Any:
"""
创建新提供商
@@ -155,7 +159,7 @@ async def create_provider(request: Request, db: Session = Depends(get_db)):
@router.put("/{provider_id}")
async def update_provider(provider_id: str, request: Request, db: Session = Depends(get_db)):
async def update_provider(provider_id: str, request: Request, db: Session = Depends(get_db)) -> None:
"""
更新提供商配置
@@ -191,7 +195,7 @@ async def update_provider(provider_id: str, request: Request, db: Session = Depe
@router.delete("/{provider_id}")
async def delete_provider(provider_id: str, request: Request, db: Session = Depends(get_db)):
async def delete_provider(provider_id: str, request: Request, db: Session = Depends(get_db)) -> None:
"""
删除提供商
@@ -213,7 +217,7 @@ class AdminListProvidersAdapter(AdminApiAdapter):
self.limit = limit
self.is_active = is_active
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
query = db.query(Provider)
if self.is_active is not None:
@@ -251,7 +255,7 @@ class AdminListProvidersAdapter(AdminApiAdapter):
class AdminCreateProviderAdapter(AdminApiAdapter):
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
payload = context.ensure_json_body()
@@ -333,7 +337,7 @@ class AdminUpdateProviderAdapter(AdminApiAdapter):
def __init__(self, provider_id: str):
self.provider_id = provider_id
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
payload = context.ensure_json_body()
@@ -412,7 +416,7 @@ class AdminDeleteProviderAdapter(AdminApiAdapter):
def __init__(self, provider_id: str):
self.provider_id = provider_id
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
provider = db.query(Provider).filter(Provider.id == self.provider_id).first()
if not provider:
@@ -483,7 +487,7 @@ class AdminGetProviderMappingPreviewAdapter(AdminApiAdapter):
def __init__(self, provider_id: str):
self.provider_id = provider_id
async def handle(self, context) -> ProviderMappingPreviewResponse: # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> ProviderMappingPreviewResponse: # type: ignore[override]
db = context.db
# 获取 Provider

View File

@@ -2,6 +2,7 @@
Provider 摘要与健康监控 API
"""
from typing import Any
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
@@ -22,6 +23,7 @@ from src.models.database import (
ProviderEndpoint,
RequestCandidate,
)
from src.api.base.context import ApiRequestContext
from src.models.endpoint_models import (
EndpointHealthEvent,
EndpointHealthMonitor,
@@ -338,7 +340,7 @@ class AdminProviderHealthMonitorAdapter(AdminApiAdapter):
lookback_hours: int
per_endpoint_limit: int
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
provider = db.query(Provider).filter(Provider.id == self.provider_id).first()
if not provider:
@@ -447,7 +449,7 @@ class AdminProviderHealthMonitorAdapter(AdminApiAdapter):
class AdminProviderSummaryAdapter(AdminApiAdapter):
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
providers = (
db.query(Provider)
@@ -462,7 +464,7 @@ class AdminUpdateProviderSettingsAdapter(AdminApiAdapter):
provider_id: str
update_data: ProviderUpdateRequest
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
provider = db.query(Provider).filter(Provider.id == self.provider_id).first()
if not provider:

View File

@@ -4,6 +4,9 @@ IP 安全管理接口
提供 IP 黑白名单管理和速率限制统计
"""
from __future__ import annotations
from typing import Any
from fastapi import APIRouter, Depends, HTTPException, Request, status
from pydantic import BaseModel, Field, ValidationError
from sqlalchemy.orm import Session
@@ -14,6 +17,7 @@ from src.api.base.pipeline import ApiRequestPipeline
from src.core.exceptions import InvalidRequestException, translate_pydantic_error
from src.database import get_db
from src.services.rate_limit.ip_limiter import IPRateLimiter
from src.api.base.context import ApiRequestContext
router = APIRouter(prefix="/api/admin/security/ip", tags=["Admin - Security"])
pipeline = ApiRequestPipeline()
@@ -52,7 +56,7 @@ class RemoveIPFromWhitelistRequest(BaseModel):
@router.post("/blacklist")
async def add_to_blacklist(request: Request, db: Session = Depends(get_db)):
async def add_to_blacklist(request: Request, db: Session = Depends(get_db)) -> None:
"""
添加 IP 到黑名单
@@ -74,7 +78,7 @@ async def add_to_blacklist(request: Request, db: Session = Depends(get_db)):
@router.delete("/blacklist/{ip_address}")
async def remove_from_blacklist(ip_address: str, request: Request, db: Session = Depends(get_db)):
async def remove_from_blacklist(ip_address: str, request: Request, db: Session = Depends(get_db)) -> None:
"""
从黑名单移除 IP
@@ -92,7 +96,7 @@ async def remove_from_blacklist(ip_address: str, request: Request, db: Session =
@router.get("/blacklist/stats")
async def get_blacklist_stats(request: Request, db: Session = Depends(get_db)):
async def get_blacklist_stats(request: Request, db: Session = Depends(get_db)) -> Any:
"""
获取黑名单统计信息
@@ -111,7 +115,7 @@ async def get_blacklist_stats(request: Request, db: Session = Depends(get_db)):
@router.post("/whitelist")
async def add_to_whitelist(request: Request, db: Session = Depends(get_db)):
async def add_to_whitelist(request: Request, db: Session = Depends(get_db)) -> None:
"""
添加 IP 到白名单
@@ -129,7 +133,7 @@ async def add_to_whitelist(request: Request, db: Session = Depends(get_db)):
@router.delete("/whitelist/{ip_address}")
async def remove_from_whitelist(ip_address: str, request: Request, db: Session = Depends(get_db)):
async def remove_from_whitelist(ip_address: str, request: Request, db: Session = Depends(get_db)) -> None:
"""
从白名单移除 IP
@@ -147,7 +151,7 @@ async def remove_from_whitelist(ip_address: str, request: Request, db: Session =
@router.get("/whitelist")
async def get_whitelist(request: Request, db: Session = Depends(get_db)):
async def get_whitelist(request: Request, db: Session = Depends(get_db)) -> Any:
"""
获取白名单
@@ -167,7 +171,7 @@ async def get_whitelist(request: Request, db: Session = Depends(get_db)):
class AddToBlacklistAdapter(AuthenticatedApiAdapter):
"""添加 IP 到黑名单适配器"""
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
payload = context.ensure_json_body()
try:
req = AddIPToBlacklistRequest.model_validate(payload)
@@ -199,7 +203,7 @@ class RemoveFromBlacklistAdapter(AuthenticatedApiAdapter):
def __init__(self, ip_address: str):
self.ip_address = ip_address
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
success = await IPRateLimiter.remove_from_blacklist(self.ip_address)
if not success:
@@ -213,7 +217,7 @@ class RemoveFromBlacklistAdapter(AuthenticatedApiAdapter):
class GetBlacklistStatsAdapter(AuthenticatedApiAdapter):
"""获取黑名单统计适配器"""
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
stats = await IPRateLimiter.get_blacklist_stats()
return stats
@@ -221,7 +225,7 @@ class GetBlacklistStatsAdapter(AuthenticatedApiAdapter):
class AddToWhitelistAdapter(AuthenticatedApiAdapter):
"""添加 IP 到白名单适配器"""
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
payload = context.ensure_json_body()
try:
req = AddIPToWhitelistRequest.model_validate(payload)
@@ -248,7 +252,7 @@ class RemoveFromWhitelistAdapter(AuthenticatedApiAdapter):
def __init__(self, ip_address: str):
self.ip_address = ip_address
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
success = await IPRateLimiter.remove_from_whitelist(self.ip_address)
if not success:
@@ -262,6 +266,6 @@ class RemoveFromWhitelistAdapter(AuthenticatedApiAdapter):
class GetWhitelistAdapter(AuthenticatedApiAdapter):
"""获取白名单适配器"""
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
whitelist = await IPRateLimiter.get_whitelist()
return {"whitelist": list(whitelist), "total": len(whitelist)}

View File

@@ -1,6 +1,9 @@
"""系统设置API端点。"""
from __future__ import annotations
from typing import Any
import copy
from dataclasses import dataclass
@@ -17,6 +20,7 @@ from src.models.api import SystemSettingsRequest, SystemSettingsResponse
from src.models.database import ApiKey, Provider, Usage, User
from src.services.email.email_template import EmailTemplate
from src.services.system.config import SystemConfigService
from src.api.base.context import ApiRequestContext
router = APIRouter(prefix="/api/admin/system", tags=["Admin - System"])
@@ -79,7 +83,7 @@ def _parse_version(version_str: str) -> tuple:
@router.get("/version")
async def get_system_version():
async def get_system_version() -> Any:
"""
获取系统版本信息
@@ -92,7 +96,7 @@ async def get_system_version():
@router.get("/check-update")
async def check_update():
async def check_update() -> Any:
"""
检查系统更新
@@ -115,7 +119,7 @@ async def check_update():
github_repo = "Aethersailor/Aether"
github_tags_url = f"https://api.github.com/repos/{github_repo}/tags"
def _make_empty_response(error: str | None = None):
def _make_empty_response(error: str | None = None) -> None:
return {
"current_version": current_version,
"latest_version": None,
@@ -238,7 +242,7 @@ pipeline = ApiRequestPipeline()
@router.get("/settings")
async def get_system_settings(request: Request, db: Session = Depends(get_db)):
async def get_system_settings(request: Request, db: Session = Depends(get_db)) -> Any:
"""
获取系统设置
@@ -255,7 +259,7 @@ async def get_system_settings(request: Request, db: Session = Depends(get_db)):
@router.put("/settings")
async def update_system_settings(http_request: Request, db: Session = Depends(get_db)):
async def update_system_settings(http_request: Request, db: Session = Depends(get_db)) -> None:
"""
更新系统设置
@@ -275,7 +279,7 @@ async def update_system_settings(http_request: Request, db: Session = Depends(ge
@router.get("/configs")
async def get_all_system_configs(request: Request, db: Session = Depends(get_db)):
async def get_all_system_configs(request: Request, db: Session = Depends(get_db)) -> Any:
"""
获取所有系统配置
@@ -290,7 +294,7 @@ async def get_all_system_configs(request: Request, db: Session = Depends(get_db)
@router.get("/configs/{key}")
async def get_system_config(key: str, request: Request, db: Session = Depends(get_db)):
async def get_system_config(key: str, request: Request, db: Session = Depends(get_db)) -> Any:
"""
获取特定系统配置
@@ -314,7 +318,7 @@ async def set_system_config(
key: str,
request: Request,
db: Session = Depends(get_db),
):
) -> Any:
"""
设置系统配置
@@ -339,7 +343,7 @@ async def set_system_config(
@router.delete("/configs/{key}")
async def delete_system_config(key: str, request: Request, db: Session = Depends(get_db)):
async def delete_system_config(key: str, request: Request, db: Session = Depends(get_db)) -> None:
"""
删除系统配置
@@ -357,7 +361,7 @@ async def delete_system_config(key: str, request: Request, db: Session = Depends
@router.get("/stats")
async def get_system_stats(request: Request, db: Session = Depends(get_db)):
async def get_system_stats(request: Request, db: Session = Depends(get_db)) -> Any:
"""
获取系统统计信息
@@ -374,7 +378,7 @@ async def get_system_stats(request: Request, db: Session = Depends(get_db)):
@router.post("/cleanup")
async def trigger_cleanup(request: Request, db: Session = Depends(get_db)):
async def trigger_cleanup(request: Request, db: Session = Depends(get_db)) -> Any:
"""
手动触发清理任务
@@ -393,7 +397,7 @@ async def trigger_cleanup(request: Request, db: Session = Depends(get_db)):
@router.get("/api-formats")
async def get_api_formats(request: Request, db: Session = Depends(get_db)):
async def get_api_formats(request: Request, db: Session = Depends(get_db)) -> Any:
"""
获取所有可用的 API 格式列表
@@ -411,35 +415,35 @@ async def get_api_formats(request: Request, db: Session = Depends(get_db)):
@router.get("/config/export")
async def export_config(request: Request, db: Session = Depends(get_db)):
async def export_config(request: Request, db: Session = Depends(get_db)) -> Any:
"""导出提供商和模型配置(管理员)"""
adapter = AdminExportConfigAdapter()
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.post("/config/import")
async def import_config(request: Request, db: Session = Depends(get_db)):
async def import_config(request: Request, db: Session = Depends(get_db)) -> Any:
"""导入提供商和模型配置(管理员)"""
adapter = AdminImportConfigAdapter()
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.get("/users/export")
async def export_users(request: Request, db: Session = Depends(get_db)):
async def export_users(request: Request, db: Session = Depends(get_db)) -> Any:
"""导出用户数据(管理员)"""
adapter = AdminExportUsersAdapter()
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.post("/users/import")
async def import_users(request: Request, db: Session = Depends(get_db)):
async def import_users(request: Request, db: Session = Depends(get_db)) -> Any:
"""导入用户数据(管理员)"""
adapter = AdminImportUsersAdapter()
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.post("/smtp/test")
async def test_smtp(request: Request, db: Session = Depends(get_db)):
async def test_smtp(request: Request, db: Session = Depends(get_db)) -> Any:
"""测试 SMTP 连接(管理员)"""
adapter = AdminTestSmtpAdapter()
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@@ -449,7 +453,7 @@ async def test_smtp(request: Request, db: Session = Depends(get_db)):
@router.get("/email/templates")
async def get_email_templates(request: Request, db: Session = Depends(get_db)):
async def get_email_templates(request: Request, db: Session = Depends(get_db)) -> Any:
"""获取所有邮件模板(管理员)"""
adapter = AdminGetEmailTemplatesAdapter()
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@@ -458,7 +462,7 @@ async def get_email_templates(request: Request, db: Session = Depends(get_db)):
@router.get("/email/templates/{template_type}")
async def get_email_template(
template_type: str, request: Request, db: Session = Depends(get_db)
):
) -> Any:
"""获取指定类型的邮件模板(管理员)"""
adapter = AdminGetEmailTemplateAdapter(template_type=template_type)
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@@ -467,7 +471,7 @@ async def get_email_template(
@router.put("/email/templates/{template_type}")
async def update_email_template(
template_type: str, request: Request, db: Session = Depends(get_db)
):
) -> Any:
"""更新邮件模板(管理员)"""
adapter = AdminUpdateEmailTemplateAdapter(template_type=template_type)
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@@ -476,7 +480,7 @@ async def update_email_template(
@router.post("/email/templates/{template_type}/preview")
async def preview_email_template(
template_type: str, request: Request, db: Session = Depends(get_db)
):
) -> Any:
"""预览邮件模板(管理员)"""
adapter = AdminPreviewEmailTemplateAdapter(template_type=template_type)
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@@ -485,7 +489,7 @@ async def preview_email_template(
@router.post("/email/templates/{template_type}/reset")
async def reset_email_template(
template_type: str, request: Request, db: Session = Depends(get_db)
):
) -> Any:
"""重置邮件模板为默认值(管理员)"""
adapter = AdminResetEmailTemplateAdapter(template_type=template_type)
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@@ -495,7 +499,7 @@ async def reset_email_template(
class AdminGetSystemSettingsAdapter(AdminApiAdapter):
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
default_provider = SystemConfigService.get_default_provider(db)
default_model = SystemConfigService.get_config(db, "default_model")
@@ -511,7 +515,7 @@ class AdminGetSystemSettingsAdapter(AdminApiAdapter):
class AdminUpdateSystemSettingsAdapter(AdminApiAdapter):
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
payload = context.ensure_json_body()
try:
@@ -560,7 +564,7 @@ class AdminUpdateSystemSettingsAdapter(AdminApiAdapter):
class AdminGetAllConfigsAdapter(AdminApiAdapter):
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
return SystemConfigService.get_all_configs(context.db)
@@ -571,7 +575,7 @@ class AdminGetSystemConfigAdapter(AdminApiAdapter):
# 敏感配置项,不返回实际值
SENSITIVE_KEYS = {"smtp_password"}
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
value = SystemConfigService.get_config(context.db, self.key)
if value is None:
raise NotFoundException(f"配置项 '{self.key}' 不存在")
@@ -588,7 +592,7 @@ class AdminSetSystemConfigAdapter(AdminApiAdapter):
# 需要加密存储的配置项
ENCRYPTED_KEYS = {"smtp_password"}
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
payload = context.ensure_json_body()
value = payload.get("value")
@@ -619,7 +623,7 @@ class AdminSetSystemConfigAdapter(AdminApiAdapter):
class AdminDeleteSystemConfigAdapter(AdminApiAdapter):
key: str
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
deleted = SystemConfigService.delete_config(context.db, self.key)
if not deleted:
raise NotFoundException(f"配置项 '{self.key}' 不存在")
@@ -627,7 +631,7 @@ class AdminDeleteSystemConfigAdapter(AdminApiAdapter):
class AdminSystemStatsAdapter(AdminApiAdapter):
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
total_users = db.query(User).count()
active_users = db.query(User).filter(User.is_active.is_(True)).count()
@@ -645,7 +649,7 @@ class AdminSystemStatsAdapter(AdminApiAdapter):
class AdminTriggerCleanupAdapter(AdminApiAdapter):
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
"""手动触发清理任务"""
from datetime import datetime, timezone
@@ -708,7 +712,7 @@ class AdminTriggerCleanupAdapter(AdminApiAdapter):
class AdminGetApiFormatsAdapter(AdminApiAdapter):
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
"""获取所有可用的API格式"""
from src.core.api_format import API_FORMAT_DEFINITIONS, APIFormat
@@ -738,7 +742,7 @@ class AdminExportConfigAdapter(AdminApiAdapter):
"token_cookie", "auth_cookie", "cookie_string", "cookie"
}
def _decrypt_provider_config(self, config: dict, crypto_service) -> dict:
def _decrypt_provider_config(self, config: dict, crypto_service: Any) -> dict:
"""解密 Provider config 中的 provider_ops credentials"""
if not config:
return config
@@ -762,7 +766,7 @@ class AdminExportConfigAdapter(AdminApiAdapter):
return decrypted_config
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
"""导出提供商和模型配置(解密数据)"""
from datetime import datetime, timezone
@@ -976,7 +980,7 @@ class AdminImportConfigAdapter(AdminApiAdapter):
"token_cookie", "auth_cookie", "cookie_string", "cookie"
}
def _encrypt_provider_config(self, config: dict, crypto_service) -> dict:
def _encrypt_provider_config(self, config: dict, crypto_service: Any) -> dict:
"""加密 Provider config 中的 provider_ops credentials"""
if not config:
return config
@@ -998,7 +1002,7 @@ class AdminImportConfigAdapter(AdminApiAdapter):
return encrypted_config
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
"""导入提供商和模型配置"""
import uuid
from datetime import datetime, timezone
@@ -1611,7 +1615,7 @@ class AdminImportConfigAdapter(AdminApiAdapter):
class AdminExportUsersAdapter(AdminApiAdapter):
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
"""导出用户数据(保留加密数据,排除管理员)"""
from datetime import datetime, timezone
@@ -1689,7 +1693,7 @@ class AdminExportUsersAdapter(AdminApiAdapter):
class AdminImportUsersAdapter(AdminApiAdapter):
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
"""导入用户数据"""
import uuid
from datetime import datetime, timezone
@@ -1891,7 +1895,7 @@ class AdminImportUsersAdapter(AdminApiAdapter):
class AdminTestSmtpAdapter(AdminApiAdapter):
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
"""测试 SMTP 连接"""
from src.core.crypto import crypto_service
from src.services.system.config import SystemConfigService
@@ -1968,7 +1972,7 @@ class AdminTestSmtpAdapter(AdminApiAdapter):
class AdminGetEmailTemplatesAdapter(AdminApiAdapter):
"""获取所有邮件模板"""
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
templates = []
@@ -2003,7 +2007,7 @@ class AdminGetEmailTemplateAdapter(AdminApiAdapter):
template_type: str
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
# 验证模板类型
if self.template_type not in EmailTemplate.TEMPLATE_TYPES:
raise NotFoundException(f"模板类型 '{self.template_type}' 不存在")
@@ -2036,7 +2040,7 @@ class AdminUpdateEmailTemplateAdapter(AdminApiAdapter):
template_type: str
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
# 验证模板类型
if self.template_type not in EmailTemplate.TEMPLATE_TYPES:
raise NotFoundException(f"模板类型 '{self.template_type}' 不存在")
@@ -2077,7 +2081,7 @@ class AdminPreviewEmailTemplateAdapter(AdminApiAdapter):
template_type: str
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
# 验证模板类型
if self.template_type not in EmailTemplate.TEMPLATE_TYPES:
raise NotFoundException(f"模板类型 '{self.template_type}' 不存在")
@@ -2123,7 +2127,7 @@ class AdminResetEmailTemplateAdapter(AdminApiAdapter):
template_type: str
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
# 验证模板类型
if self.template_type not in EmailTemplate.TEMPLATE_TYPES:
raise NotFoundException(f"模板类型 '{self.template_type}' 不存在")

View File

@@ -1,5 +1,8 @@
"""管理员使用情况统计路由。"""
from __future__ import annotations
from typing import Any
from collections import defaultdict
from dataclasses import dataclass
from datetime import datetime
@@ -21,6 +24,7 @@ from src.models.database import (
User,
)
from src.services.usage.service import UsageService
from src.api.base.context import ApiRequestContext
router = APIRouter(prefix="/api/admin/usage", tags=["Admin - Usage"])
pipeline = ApiRequestPipeline()
@@ -37,7 +41,7 @@ async def get_usage_aggregation(
end_date: datetime | None = None,
limit: int = Query(20, ge=1, le=100),
db: Session = Depends(get_db),
):
) -> Any:
"""
获取使用情况聚合统计
@@ -77,7 +81,7 @@ async def get_usage_stats(
start_date: datetime | None = None,
end_date: datetime | None = None,
db: Session = Depends(get_db),
):
) -> Any:
"""
获取使用情况总体统计
@@ -105,7 +109,7 @@ async def get_usage_stats(
async def get_activity_heatmap(
request: Request,
db: Session = Depends(get_db),
):
) -> Any:
"""
获取活动热力图数据
@@ -132,7 +136,7 @@ async def get_usage_records(
limit: int = Query(100, ge=1, le=500),
offset: int = Query(0, ge=0),
db: Session = Depends(get_db),
):
) -> Any:
"""
获取使用记录列表
@@ -180,7 +184,7 @@ async def get_active_requests(
request: Request,
ids: str | None = Query(None, description="逗号分隔的请求 ID 列表,用于查询特定请求的状态"),
db: Session = Depends(get_db),
):
) -> Any:
"""
获取活跃请求的状态
@@ -207,7 +211,7 @@ async def get_usage_detail(
usage_id: str,
request: Request,
db: Session = Depends(get_db),
):
) -> Any:
"""
获取使用记录详情
@@ -262,7 +266,7 @@ class AdminUsageStatsAdapter(AdminApiAdapter):
self.start_date = start_date
self.end_date = end_date
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
query = db.query(Usage)
if self.start_date:
@@ -327,7 +331,7 @@ class AdminUsageStatsAdapter(AdminApiAdapter):
class AdminActivityHeatmapAdapter(AdminApiAdapter):
"""Activity heatmap adapter with Redis caching."""
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
result = await UsageService.get_cached_heatmap(
db=context.db,
user_id=None,
@@ -343,7 +347,7 @@ class AdminUsageByModelAdapter(AdminApiAdapter):
self.end_date = end_date
self.limit = limit
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
query = db.query(
Usage.model,
@@ -390,7 +394,7 @@ class AdminUsageByUserAdapter(AdminApiAdapter):
self.end_date = end_date
self.limit = limit
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
query = (
db.query(
@@ -440,7 +444,7 @@ class AdminUsageByProviderAdapter(AdminApiAdapter):
self.end_date = end_date
self.limit = limit
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
# 从 request_candidates 表统计每个 Provider 的尝试次数和成功率
@@ -554,7 +558,7 @@ class AdminUsageByApiFormatAdapter(AdminApiAdapter):
self.end_date = end_date
self.limit = limit
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
query = db.query(
Usage.api_format,
@@ -629,7 +633,7 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
self.limit = limit
self.offset = offset
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
from sqlalchemy import or_
from src.utils.database_helpers import escape_like_pattern, safe_truncate_escaped
@@ -879,7 +883,7 @@ class AdminActiveRequestsAdapter(AdminApiAdapter):
def __init__(self, ids: str | None):
self.ids = ids
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
from src.services.usage import UsageService
db = context.db
@@ -901,7 +905,7 @@ class AdminUsageDetailAdapter(AdminApiAdapter):
usage_id: str
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
usage_record = db.query(Usage).filter(Usage.id == self.usage_id).first()
if not usage_record:
@@ -972,7 +976,7 @@ class AdminUsageDetailAdapter(AdminApiAdapter):
"tiered_pricing": tiered_pricing_info,
}
async def _get_tiered_pricing_info(self, db, usage_record) -> dict | None:
async def _get_tiered_pricing_info(self, db: Session, usage_record: Any) -> dict | None:
"""获取阶梯计费信息"""
from src.services.model.cost import ModelCostService
@@ -1036,7 +1040,7 @@ async def analyze_cache_affinity_ttl(
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),
):
) -> Any:
"""
分析用户请求间隔分布,推荐合适的缓存亲和性 TTL。
@@ -1060,7 +1064,7 @@ async def analyze_cache_hit(
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),
):
) -> Any:
"""
分析缓存命中情况。
@@ -1087,7 +1091,7 @@ class CacheAffinityTTLAnalysisAdapter(AdminApiAdapter):
self.api_key_id = api_key_id
self.hours = hours
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
result = UsageService.analyze_cache_affinity_ttl(
@@ -1121,7 +1125,7 @@ class CacheHitAnalysisAdapter(AdminApiAdapter):
self.api_key_id = api_key_id
self.hours = hours
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
result = UsageService.get_cache_hit_analysis(
@@ -1149,7 +1153,7 @@ async def get_interval_timeline(
user_id: str | None = Query(None, description="指定用户 ID"),
include_user_info: bool = Query(False, description="是否包含用户信息(用于管理员多用户视图)"),
db: Session = Depends(get_db),
):
) -> Any:
"""
获取请求间隔时间线数据,用于散点图展示。
@@ -1184,7 +1188,7 @@ class IntervalTimelineAdapter(AdminApiAdapter):
self.user_id = user_id
self.include_user_info = include_user_info
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
result = UsageService.get_interval_timeline(

View File

@@ -1,5 +1,8 @@
"""用户管理 API 端点。"""
from __future__ import annotations
from typing import Any
from datetime import datetime, timezone
from fastapi import APIRouter, Depends, HTTPException, Query, Request, status
@@ -17,6 +20,7 @@ from src.models.database import ApiKey, User, UserRole
from src.services.system.config import SystemConfigService
from src.services.user.apikey import ApiKeyService
from src.services.user.service import UserService
from src.api.base.context import ApiRequestContext
router = APIRouter(prefix="/api/admin/users", tags=["Admin - Users"])
@@ -25,7 +29,7 @@ pipeline = ApiRequestPipeline()
# 管理员端点
@router.post("")
async def create_user_endpoint(request: Request, db: Session = Depends(get_db)):
async def create_user_endpoint(request: Request, db: Session = Depends(get_db)) -> Any:
"""
创建用户
@@ -50,7 +54,7 @@ async def list_users(
role: str | None = Query(None, description="按角色筛选user/admin"),
is_active: bool | None = Query(None, description="按状态筛选"),
db: Session = Depends(get_db),
):
) -> Any:
"""
获取用户列表
@@ -63,7 +67,7 @@ async def list_users(
@router.get("/{user_id}")
async def get_user(user_id: str, request: Request, db: Session = Depends(get_db)):
async def get_user(user_id: str, request: Request, db: Session = Depends(get_db)) -> Any:
"""
获取用户详情
@@ -81,7 +85,7 @@ async def update_user(
user_id: str,
request: Request,
db: Session = Depends(get_db),
):
) -> Any:
"""
更新用户信息
@@ -104,7 +108,7 @@ async def update_user(
@router.delete("/{user_id}")
async def delete_user(user_id: str, request: Request, db: Session = Depends(get_db)):
async def delete_user(user_id: str, request: Request, db: Session = Depends(get_db)) -> None:
"""
删除用户
@@ -118,7 +122,7 @@ async def delete_user(user_id: str, request: Request, db: Session = Depends(get_
@router.patch("/{user_id}/quota")
async def reset_user_quota(user_id: str, request: Request, db: Session = Depends(get_db)):
async def reset_user_quota(user_id: str, request: Request, db: Session = Depends(get_db)) -> None:
"""
重置用户配额
@@ -137,7 +141,7 @@ async def get_user_api_keys(
request: Request,
is_active: bool | None = Query(None, description="按状态筛选"),
db: Session = Depends(get_db),
):
) -> Any:
"""
获取用户的 API 密钥列表
@@ -155,7 +159,7 @@ async def create_user_api_key(
user_id: str,
request: Request,
db: Session = Depends(get_db),
):
) -> Any:
"""
为用户创建 API 密钥
@@ -183,7 +187,7 @@ async def delete_user_api_key(
key_id: str,
request: Request,
db: Session = Depends(get_db),
):
) -> Any:
"""
删除用户的 API 密钥
@@ -201,7 +205,7 @@ async def delete_user_api_key(
class AdminCreateUserAdapter(AdminApiAdapter):
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
payload = context.ensure_json_body()
try:
@@ -279,7 +283,7 @@ class AdminListUsersAdapter(AdminApiAdapter):
self.role = role
self.is_active = is_active
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
role_enum = UserRole[self.role.upper()] if self.role else None
users = UserService.list_users(db, self.skip, self.limit, role_enum, self.is_active)
@@ -306,7 +310,7 @@ class AdminGetUserAdapter(AdminApiAdapter):
def __init__(self, user_id: str):
self.user_id = user_id
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
user = UserService.get_user(db, self.user_id)
if not user:
@@ -341,7 +345,7 @@ class AdminUpdateUserAdapter(AdminApiAdapter):
def __init__(self, user_id: str):
self.user_id = user_id
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
existing_user = UserService.get_user(db, self.user_id)
if not existing_user:
@@ -406,7 +410,7 @@ class AdminDeleteUserAdapter(AdminApiAdapter):
def __init__(self, user_id: str):
self.user_id = user_id
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
user = UserService.get_user(db, self.user_id)
if not user:
@@ -435,7 +439,7 @@ class AdminResetUserQuotaAdapter(AdminApiAdapter):
def __init__(self, user_id: str):
self.user_id = user_id
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
user = UserService.get_user(db, self.user_id)
if not user:
@@ -470,7 +474,7 @@ class AdminGetUserKeysAdapter(AdminApiAdapter):
self.user_id = user_id
self.is_active = is_active
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
# 验证用户存在
@@ -518,7 +522,7 @@ class AdminCreateUserKeyAdapter(AdminApiAdapter):
def __init__(self, user_id: str):
self.user_id = user_id
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
payload = context.ensure_json_body()
try:
@@ -574,7 +578,7 @@ class AdminDeleteUserKeyAdapter(AdminApiAdapter):
self.user_id = user_id
self.key_id = key_id
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
# 验证Key存在且属于该用户

View File

@@ -1,5 +1,8 @@
"""公告系统 API 端点。"""
from __future__ import annotations
from typing import Any
from dataclasses import dataclass
from fastapi import APIRouter, Depends, Query, Request
@@ -16,6 +19,7 @@ from src.models.api import CreateAnnouncementRequest, UpdateAnnouncementRequest
from src.models.database import User
from src.services.auth.service import AuthService
from src.services.system.announcement import AnnouncementService
from src.api.base.context import ApiRequestContext
router = APIRouter(prefix="/api/announcements", tags=["Announcements"])
@@ -32,7 +36,7 @@ async def list_announcements(
limit: int = Query(50, description="返回数量限制"),
offset: int = Query(0, description="偏移量"),
db: Session = Depends(get_db),
):
) -> Any:
"""
获取公告列表
@@ -67,7 +71,7 @@ async def list_announcements(
async def get_active_announcements(
request: Request,
db: Session = Depends(get_db),
):
) -> Any:
"""
获取当前有效的公告
@@ -87,7 +91,7 @@ async def get_announcement(
announcement_id: str, # UUID
request: Request,
db: Session = Depends(get_db),
):
) -> Any:
"""
获取单个公告详情
@@ -118,7 +122,7 @@ async def mark_announcement_as_read(
announcement_id: str,
request: Request,
db: Session = Depends(get_db),
):
) -> Any:
"""
标记公告为已读
@@ -141,7 +145,7 @@ async def mark_announcement_as_read(
async def create_announcement(
request: Request,
db: Session = Depends(get_db),
):
) -> Any:
"""
创建公告
@@ -170,7 +174,7 @@ async def update_announcement(
announcement_id: str,
request: Request,
db: Session = Depends(get_db),
):
) -> Any:
"""
更新公告
@@ -201,7 +205,7 @@ async def delete_announcement(
announcement_id: str,
request: Request,
db: Session = Depends(get_db),
):
) -> Any:
"""
删除公告
@@ -224,7 +228,7 @@ async def delete_announcement(
async def get_my_unread_announcement_count(
request: Request,
db: Session = Depends(get_db),
):
) -> Any:
"""
获取我的未读公告数量
@@ -245,11 +249,11 @@ class AnnouncementOptionalAuthAdapter(ApiAdapter):
mode = ApiMode.PUBLIC
async def authorize(self, context): # type: ignore[override]
async def authorize(self, context: ApiRequestContext) -> None: # type: ignore[override]
context.extra["optional_user"] = await self._resolve_optional_user(context)
return None
async def _resolve_optional_user(self, context) -> User | None:
async def _resolve_optional_user(self, context: ApiRequestContext) -> User | None:
if context.user:
return context.user
@@ -283,7 +287,7 @@ class AnnouncementOptionalAuthAdapter(ApiAdapter):
except Exception:
return None
def get_optional_user(self, context) -> User | None:
def get_optional_user(self, context: ApiRequestContext) -> User | None:
return context.extra.get("optional_user")
@@ -293,7 +297,7 @@ class ListAnnouncementsAdapter(AnnouncementOptionalAuthAdapter):
limit: int
offset: int
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
optional_user = self.get_optional_user(context)
return AnnouncementService.get_announcements(
db=context.db,
@@ -306,7 +310,7 @@ class ListAnnouncementsAdapter(AnnouncementOptionalAuthAdapter):
class GetActiveAnnouncementsAdapter(AnnouncementOptionalAuthAdapter):
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
optional_user = self.get_optional_user(context)
return AnnouncementService.get_active_announcements(
db=context.db,
@@ -318,7 +322,7 @@ class GetActiveAnnouncementsAdapter(AnnouncementOptionalAuthAdapter):
class GetAnnouncementAdapter(AnnouncementOptionalAuthAdapter):
announcement_id: str
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
announcement = AnnouncementService.get_announcement(context.db, self.announcement_id)
return {
"id": announcement.id,
@@ -345,13 +349,13 @@ class MarkAnnouncementReadAdapter(AnnouncementUserAdapter):
def __init__(self, announcement_id: str):
self.announcement_id = announcement_id
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
AnnouncementService.mark_as_read(context.db, self.announcement_id, context.user.id)
return {"message": "公告已标记为已读"}
class UnreadAnnouncementCountAdapter(AnnouncementUserAdapter):
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
result = AnnouncementService.get_announcements(
db=context.db,
user_id=context.user.id,
@@ -364,7 +368,7 @@ class UnreadAnnouncementCountAdapter(AnnouncementUserAdapter):
class CreateAnnouncementAdapter(AdminApiAdapter):
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
payload = context.ensure_json_body()
try:
req = CreateAnnouncementRequest.model_validate(payload)
@@ -392,7 +396,7 @@ class CreateAnnouncementAdapter(AdminApiAdapter):
class UpdateAnnouncementAdapter(AdminApiAdapter):
announcement_id: str
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
payload = context.ensure_json_body()
try:
req = UpdateAnnouncementRequest.model_validate(payload)
@@ -422,6 +426,6 @@ class UpdateAnnouncementAdapter(AdminApiAdapter):
class DeleteAnnouncementAdapter(AdminApiAdapter):
announcement_id: str
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
AnnouncementService.delete_announcement(context.db, self.announcement_id, context.user.id)
return {"message": "公告已删除"}

View File

@@ -3,6 +3,9 @@
"""
from __future__ import annotations
from typing import Any
from fastapi import APIRouter, Depends, HTTPException, Request, status
from fastapi.security import HTTPBearer
from pydantic import ValidationError
@@ -39,6 +42,7 @@ from src.services.system.config import SystemConfigService
from src.services.user.service import UserService
from src.services.email import EmailSenderService, EmailVerificationService
from src.utils.request_utils import get_client_ip, get_user_agent
from src.api.base.context import ApiRequestContext
def validate_email_suffix(db: Session, email: str) -> tuple[bool, str | None]:
@@ -93,7 +97,7 @@ pipeline = ApiRequestPipeline()
# API端点
@router.get("/registration-settings", response_model=RegistrationSettingsResponse)
async def registration_settings(request: Request, db: Session = Depends(get_db)):
async def registration_settings(request: Request, db: Session = Depends(get_db)) -> Any:
"""
获取注册相关配置
@@ -105,7 +109,7 @@ async def registration_settings(request: Request, db: Session = Depends(get_db))
@router.get("/settings")
async def auth_settings(request: Request, db: Session = Depends(get_db)):
async def auth_settings(request: Request, db: Session = Depends(get_db)) -> Any:
"""
获取认证设置
@@ -117,7 +121,7 @@ async def auth_settings(request: Request, db: Session = Depends(get_db)):
@router.post("/login", response_model=LoginResponse)
async def login(request: Request, db: Session = Depends(get_db)):
async def login(request: Request, db: Session = Depends(get_db)) -> Any:
"""
用户登录
@@ -133,7 +137,7 @@ async def login(request: Request, db: Session = Depends(get_db)):
@router.post("/refresh", response_model=RefreshTokenResponse)
async def refresh_token(request: Request, db: Session = Depends(get_db)):
async def refresh_token(request: Request, db: Session = Depends(get_db)) -> None:
"""
刷新访问令牌
@@ -145,7 +149,7 @@ async def refresh_token(request: Request, db: Session = Depends(get_db)):
@router.post("/register", response_model=RegisterResponse)
async def register(request: Request, db: Session = Depends(get_db)):
async def register(request: Request, db: Session = Depends(get_db)) -> Any:
"""
用户注册
@@ -159,7 +163,7 @@ async def register(request: Request, db: Session = Depends(get_db)):
@router.get("/me")
async def get_current_user_info(request: Request, db: Session = Depends(get_db)):
async def get_current_user_info(request: Request, db: Session = Depends(get_db)) -> Any:
"""
获取当前用户信息
@@ -171,7 +175,7 @@ async def get_current_user_info(request: Request, db: Session = Depends(get_db))
@router.patch("/password")
async def change_password(request: Request, db: Session = Depends(get_db)):
async def change_password(request: Request, db: Session = Depends(get_db)) -> Any:
"""
修改密码
@@ -183,7 +187,7 @@ async def change_password(request: Request, db: Session = Depends(get_db)):
@router.post("/logout", response_model=LogoutResponse)
async def logout(request: Request, db: Session = Depends(get_db)):
async def logout(request: Request, db: Session = Depends(get_db)) -> Any:
"""
用户登出
@@ -194,7 +198,7 @@ async def logout(request: Request, db: Session = Depends(get_db)):
@router.post("/send-verification-code", response_model=SendVerificationCodeResponse)
async def send_verification_code(request: Request, db: Session = Depends(get_db)):
async def send_verification_code(request: Request, db: Session = Depends(get_db)) -> None:
"""
发送邮箱验证码
@@ -208,7 +212,7 @@ async def send_verification_code(request: Request, db: Session = Depends(get_db)
@router.post("/verify-email", response_model=VerifyEmailResponse)
async def verify_email(request: Request, db: Session = Depends(get_db)):
async def verify_email(request: Request, db: Session = Depends(get_db)) -> Any:
"""
验证邮箱验证码
@@ -222,7 +226,7 @@ async def verify_email(request: Request, db: Session = Depends(get_db)):
@router.post("/verification-status", response_model=VerificationStatusResponse)
async def verification_status(request: Request, db: Session = Depends(get_db)):
async def verification_status(request: Request, db: Session = Depends(get_db)) -> Any:
"""
查询邮箱验证状态
@@ -240,12 +244,12 @@ async def verification_status(request: Request, db: Session = Depends(get_db)):
class AuthPublicAdapter(ApiAdapter):
mode = ApiMode.PUBLIC
def authorize(self, context): # type: ignore[override]
def authorize(self, context: ApiRequestContext) -> None: # type: ignore[override]
return None
class AuthLoginAdapter(AuthPublicAdapter):
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
payload = context.ensure_json_body()
@@ -319,7 +323,7 @@ class AuthLoginAdapter(AuthPublicAdapter):
class AuthRefreshAdapter(AuthPublicAdapter):
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
payload = context.ensure_json_body()
refresh_request = RefreshTokenRequest.model_validate(payload)
@@ -374,7 +378,7 @@ class AuthRefreshAdapter(AuthPublicAdapter):
class AuthRegistrationSettingsAdapter(AuthPublicAdapter):
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
"""公开返回注册相关配置"""
db = context.db
@@ -394,7 +398,7 @@ class AuthRegistrationSettingsAdapter(AuthPublicAdapter):
class AuthSettingsAdapter(AuthPublicAdapter):
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
"""公开返回认证设置"""
db = context.db
@@ -409,7 +413,7 @@ class AuthSettingsAdapter(AuthPublicAdapter):
class AuthRegisterAdapter(AuthPublicAdapter):
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
from src.models.database import SystemConfig
db = context.db
@@ -545,7 +549,7 @@ class AuthRegisterAdapter(AuthPublicAdapter):
class AuthCurrentUserAdapter(AuthenticatedApiAdapter):
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
user = context.user
return {
"id": user.id,
@@ -566,7 +570,7 @@ class AuthCurrentUserAdapter(AuthenticatedApiAdapter):
class AuthChangePasswordAdapter(AuthenticatedApiAdapter):
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
payload = context.ensure_json_body()
old_password = payload.get("old_password")
new_password = payload.get("new_password")
@@ -584,7 +588,7 @@ class AuthChangePasswordAdapter(AuthenticatedApiAdapter):
class AuthLogoutAdapter(AuthenticatedApiAdapter):
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
"""用户登出,将 Token 加入黑名单"""
user = context.user
client_ip = get_client_ip(context.request)
@@ -621,7 +625,7 @@ class AuthLogoutAdapter(AuthenticatedApiAdapter):
class AuthSendVerificationCodeAdapter(AuthPublicAdapter):
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
"""发送邮箱验证码"""
db = context.db
payload = context.ensure_json_body()
@@ -702,7 +706,7 @@ class AuthSendVerificationCodeAdapter(AuthPublicAdapter):
class AuthVerifyEmailAdapter(AuthPublicAdapter):
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
"""验证邮箱验证码"""
db = context.db
payload = context.ensure_json_body()
@@ -744,7 +748,7 @@ class AuthVerifyEmailAdapter(AuthPublicAdapter):
class AuthVerificationStatusAdapter(AuthPublicAdapter):
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
"""查询邮箱验证状态"""
payload = context.ensure_json_body()

View File

@@ -1,3 +1,4 @@
from src.api.base.context import ApiRequestContext
from fastapi import HTTPException
from .adapter import ApiAdapter, ApiMode
@@ -8,6 +9,6 @@ class AuthenticatedApiAdapter(ApiAdapter):
mode = ApiMode.USER
def authorize(self, context): # type: ignore[override]
def authorize(self, context: ApiRequestContext) -> None: # type: ignore[override]
if not context.user:
raise HTTPException(status_code=401, detail="未登录")

View File

@@ -54,7 +54,7 @@ class ApiRequestPipeline:
mode: ApiMode = ApiMode.STANDARD,
api_format_hint: str | None = None,
path_params: dict[str, Any] | None = None,
):
) -> Any:
# 高频轮询端点抑制 debug 日志
is_quiet = http_request.url.path in QUIET_POLLING_PATHS
if not is_quiet:
@@ -500,7 +500,7 @@ class ApiRequestPipeline:
return self._sanitize_metadata(metadata)
def _sanitize_metadata(self, value: Any, depth: int = 0):
def _sanitize_metadata(self, value: Any, depth: int = 0) -> None:
if value is None:
return None
if depth > 5:

View File

@@ -1,5 +1,8 @@
"""仪表盘统计 API 端点。"""
from __future__ import annotations
from typing import Any
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
@@ -17,6 +20,7 @@ from src.models.database import ApiKey, Provider, RequestCandidate, StatsDaily,
from src.models.database import User as DBUser
from src.services.system.stats_aggregator import StatsAggregatorService
from src.utils.cache_decorator import cache_result
from src.api.base.context import ApiRequestContext
router = APIRouter(prefix="/api/dashboard", tags=["Dashboard"])
pipeline = ApiRequestPipeline()
@@ -44,7 +48,7 @@ def format_tokens(num: int) -> str:
@router.get("/stats")
async def get_dashboard_stats(request: Request, db: Session = Depends(get_db)):
async def get_dashboard_stats(request: Request, db: Session = Depends(get_db)) -> Any:
"""
获取仪表盘统计数据
@@ -77,7 +81,7 @@ async def get_recent_requests(
request: Request,
limit: int = Query(10, ge=1, le=100),
db: Session = Depends(get_db),
):
) -> Any:
"""
获取最近请求列表
@@ -104,7 +108,7 @@ async def get_recent_requests(
@router.get("/provider-status")
async def get_provider_status(request: Request, db: Session = Depends(get_db)):
async def get_provider_status(request: Request, db: Session = Depends(get_db)) -> Any:
"""
获取提供商状态
@@ -125,7 +129,7 @@ async def get_daily_stats(
request: Request,
days: int = Query(7, ge=1, le=30),
db: Session = Depends(get_db),
):
) -> Any:
"""
获取每日统计数据
@@ -157,13 +161,13 @@ class DashboardAdapter(ApiAdapter):
mode = ApiMode.USER # 普通用户也可访问仪表盘
def authorize(self, context): # type: ignore[override]
def authorize(self, context: ApiRequestContext) -> None: # type: ignore[override]
if not context.user:
raise HTTPException(status_code=401, detail="未登录")
class DashboardStatsAdapter(DashboardAdapter):
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
user = context.user
if not user:
raise HTTPException(status_code=401, detail="未登录")
@@ -178,7 +182,7 @@ class DashboardStatsAdapter(DashboardAdapter):
class AdminDashboardStatsAdapter(AdminApiAdapter):
@cache_result(key_prefix="dashboard:admin:stats", ttl=CacheTTL.DASHBOARD_STATS, user_specific=False)
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
"""管理员仪表盘统计 - 使用预聚合数据优化性能"""
from zoneinfo import ZoneInfo
from src.services.system.stats_aggregator import APP_TIMEZONE
@@ -509,7 +513,7 @@ class AdminDashboardStatsAdapter(AdminApiAdapter):
class UserDashboardStatsAdapter(DashboardAdapter):
@cache_result(key_prefix="dashboard:user:stats", ttl=30, user_specific=True)
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
from zoneinfo import ZoneInfo
from src.services.system.stats_aggregator import APP_TIMEZONE
@@ -724,7 +728,7 @@ class UserDashboardStatsAdapter(DashboardAdapter):
class DashboardRecentRequestsAdapter(DashboardAdapter):
limit: int
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
user = context.user
query = db.query(Usage)
@@ -756,7 +760,7 @@ class DashboardRecentRequestsAdapter(DashboardAdapter):
class DashboardProviderStatusAdapter(DashboardAdapter):
@cache_result(key_prefix="dashboard:provider:status", ttl=60, user_specific=False)
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
user = context.user
providers = db.query(Provider).filter(Provider.is_active.is_(True)).all()
@@ -787,7 +791,7 @@ class DashboardDailyStatsAdapter(DashboardAdapter):
days: int
@cache_result(key_prefix="dashboard:daily:stats", ttl=CacheTTL.DASHBOARD_DAILY, user_specific=True)
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
from zoneinfo import ZoneInfo
from src.services.system.stats_aggregator import APP_TIMEZONE

View File

@@ -88,7 +88,7 @@ _LAZY_IMPORTS = {
}
def __getattr__(name: str):
def __getattr__(name: str) -> None:
"""延迟导入以避免循环依赖"""
if name in _LAZY_IMPORTS:
module_path, attr_name = _LAZY_IMPORTS[name]

View File

@@ -18,6 +18,7 @@ Chat Adapter 通用基类
from __future__ import annotations
from sqlalchemy.orm import Session
import time
import traceback
from abc import abstractmethod
@@ -129,7 +130,7 @@ class ChatAdapterBase(ApiAdapter):
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):
async def handle(self, context: ApiRequestContext) -> Any:
"""处理 Chat API 请求"""
http_request = context.request
user = context.user
@@ -262,14 +263,14 @@ class ChatAdapterBase(ApiAdapter):
def _create_handler(
self,
*,
db,
user,
api_key,
db: Session,
user: Any,
api_key: Any,
request_id: str,
client_ip: str,
user_agent: str,
start_time: float,
):
) -> Any:
"""创建 Handler 实例 - 子类可覆盖"""
return self.HANDLER_CLASS(
db=db,
@@ -305,7 +306,7 @@ class ChatAdapterBase(ApiAdapter):
return merged
@abstractmethod
def _validate_request_body(self, original_request_body: dict, path_params: dict = None):
def _validate_request_body(self, original_request_body: dict, path_params: dict | None = None) -> None:
"""
验证请求体 - 子类必须实现
@@ -318,7 +319,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: Any) -> int:
"""
提取消息数量 - 子类可覆盖
@@ -329,7 +330,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: Any) -> dict[str, Any]:
"""
构建审计日志元数据 - 子类可覆盖
"""
@@ -351,9 +352,9 @@ class ChatAdapterBase(ApiAdapter):
self,
e: Exception,
*,
db,
user,
api_key,
db: Session,
user: Any,
api_key: Any,
model: str,
stream: bool,
start_time: float,
@@ -422,9 +423,9 @@ class ChatAdapterBase(ApiAdapter):
self,
e: Exception,
*,
db,
user,
api_key,
db: Session,
user: Any,
api_key: Any,
model: str,
stream: bool,
start_time: float,
@@ -710,7 +711,7 @@ def register_adapter(adapter_class: type[ChatAdapterBase]) -> type[ChatAdapterBa
return adapter_class
def _ensure_adapters_loaded():
def _ensure_adapters_loaded() -> None:
"""确保所有 Adapter 已被加载(触发注册)"""
global _ADAPTERS_LOADED
if _ADAPTERS_LOADED:

View File

@@ -19,6 +19,8 @@ Chat Handler Base - Chat API 格式的通用基类
- StreamTelemetryRecorder: 统计记录Usage、Audit、Candidate
"""
from __future__ import annotations
import asyncio
import json
from abc import ABC, abstractmethod

View File

@@ -15,12 +15,15 @@ CLI Adapter 通用基类
- 可选覆盖 compute_total_input_context() 自定义总输入上下文计算
"""
from __future__ import annotations
import time
import traceback
from typing import Any
import httpx
from fastapi import HTTPException, Request
from sqlalchemy.orm import Session
from fastapi.responses import JSONResponse
from src.api.base.adapter import ApiAdapter, ApiMode
@@ -124,7 +127,7 @@ class CliAdapterBase(ApiAdapter):
"""
return get_adapter_protected_keys(cls._get_api_format())
async def handle(self, context: ApiRequestContext):
async def handle(self, context: ApiRequestContext) -> Any:
"""处理 CLI API 请求"""
http_request = context.request
user = context.user
@@ -332,9 +335,9 @@ class CliAdapterBase(ApiAdapter):
self,
e: Exception,
*,
db,
user,
api_key,
db: Session,
user: Any,
api_key: Any,
model: str,
stream: bool,
start_time: float,
@@ -403,9 +406,9 @@ class CliAdapterBase(ApiAdapter):
self,
e: Exception,
*,
db,
user,
api_key,
db: Session,
user: Any,
api_key: Any,
model: str,
stream: bool,
start_time: float,
@@ -748,7 +751,7 @@ def register_cli_adapter(adapter_class: type[CliAdapterBase]) -> type[CliAdapter
return adapter_class
def _ensure_cli_adapters_loaded():
def _ensure_cli_adapters_loaded() -> None:
"""确保所有 CLI Adapter 已被加载(触发注册)"""
global _CLI_ADAPTERS_LOADED
if _CLI_ADAPTERS_LOADED:

View File

@@ -9,6 +9,8 @@
5. 流式平滑输出
"""
from __future__ import annotations
import asyncio
import codecs
import json

View File

@@ -4,6 +4,8 @@ Claude Chat Adapter - 基于 ChatAdapterBase 的 Claude Chat API 适配器
处理 /v1/messages 端点的 Claude Chat 格式请求。
"""
from __future__ import annotations
from typing import Any
import httpx
@@ -96,7 +98,7 @@ class ClaudeChatAdapter(ChatAdapterBase):
"""
return input_tokens + cache_creation_input_tokens + cache_read_input_tokens
def _validate_request_body(self, original_request_body: dict, path_params: dict = None):
def _validate_request_body(self, original_request_body: dict, path_params: dict | None = None) -> None:
"""验证请求体"""
try:
if not isinstance(original_request_body, dict):
@@ -124,7 +126,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: Any) -> dict[str, Any]:
"""构建 Claude Chat 特定的审计元数据"""
role_counts: dict[str, int] = {}
for message in request_obj.messages:
@@ -201,7 +203,7 @@ class ClaudeChatAdapter(ChatAdapterBase):
# build_request_body 使用基类实现,通过 format_conversion_registry 自动转换 OPENAI -> CLAUDE
def build_claude_adapter(x_app_header: str | None):
def build_claude_adapter(x_app_header: str | None) -> Any:
"""根据 x-app 头部构造 Chat 或 Claude Code 适配器。"""
if x_app_header and x_app_header.lower() == "cli":
from src.api.handlers.claude_cli.adapter import ClaudeCliAdapter
@@ -228,7 +230,7 @@ class ClaudeTokenCountAdapter(ApiAdapter):
return authorization.replace("Bearer ", "")
return None
async def handle(self, context: ApiRequestContext):
async def handle(self, context: ApiRequestContext) -> Any:
payload = context.ensure_json_body()
try:

View File

@@ -4,6 +4,8 @@ Claude CLI Adapter - 基于通用 CLI Adapter 基类的简化实现
继承 CliAdapterBase只需配置 FORMAT_ID 和 HANDLER_CLASS。
"""
from __future__ import annotations
from typing import Any
import httpx

View File

@@ -91,7 +91,7 @@ class GeminiChatAdapter(ChatAdapterBase):
"""
return original_request_body.copy()
def _validate_request_body(self, original_request_body: dict, path_params: dict = None):
def _validate_request_body(self, original_request_body: dict, path_params: dict | None = None) -> None:
"""验证请求体"""
path_params = path_params or {}
is_stream = path_params.get("stream", False)
@@ -124,14 +124,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: Any) -> 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: Any) -> dict[str, Any]:
"""构建 Gemini Chat 特定的审计元数据"""
role_counts: dict[str, int] = {}

View File

@@ -6,6 +6,7 @@ Gemini Chat Handler
from __future__ import annotations
from starlette.requests import Request
from typing import Any
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
@@ -102,7 +103,7 @@ class GeminiChatHandler(ChatHandlerBase):
return str(path_params["model"])
return "unknown"
async def _convert_request(self, request):
async def _convert_request(self, request: Request) -> None:
"""
将请求转换为 Gemini 格式的 Pydantic 对象

View File

@@ -4,6 +4,8 @@ Gemini CLI Adapter - 基于通用 CLI Adapter 基类的实现
继承 CliAdapterBase处理 Gemini CLI 格式的请求。
"""
from __future__ import annotations
from typing import Any
import httpx

View File

@@ -4,6 +4,8 @@ OpenAI Chat Adapter - 基于 ChatAdapterBase 的 OpenAI Chat API 适配器
处理 /v1/chat/completions 端点的 OpenAI Chat 格式请求。
"""
from __future__ import annotations
from typing import Any
import httpx
@@ -37,7 +39,7 @@ class OpenAIChatAdapter(ChatAdapterBase):
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):
def _validate_request_body(self, original_request_body: dict, path_params: dict | None = None) -> None:
"""验证请求体"""
if not isinstance(original_request_body, dict):
return self._error_response(
@@ -66,7 +68,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: Any) -> dict[str, Any]:
"""构建 OpenAI Chat 特定的审计元数据"""
role_counts = {}
for message in request_obj.messages:

View File

@@ -5,6 +5,9 @@ OpenAI Chat Handler - 基于通用 Chat Handler 基类的简化实现
代码量从原来的 ~1315 行减少到 ~100 行。
"""
from __future__ import annotations
from starlette.requests import Request
from typing import Any
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
@@ -63,7 +66,7 @@ class OpenAIChatHandler(ChatHandlerBase):
result["model"] = mapped_model
return result
async def _convert_request(self, request):
async def _convert_request(self, request: Request) -> None:
"""
将请求转换为 OpenAI 格式的 Pydantic 对象

View File

@@ -4,6 +4,8 @@ OpenAI CLI Adapter - 基于通用 CLI Adapter 基类的简化实现
继承 CliAdapterBase只需配置 FORMAT_ID 和 HANDLER_CLASS。
"""
from __future__ import annotations
from typing import Any
import httpx

View File

@@ -1,5 +1,8 @@
"""普通用户可访问的监控与审计端点。"""
from __future__ import annotations
from typing import Any
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
@@ -13,6 +16,7 @@ from src.core.logger import logger
from src.database import get_db
from src.models.database import ApiKey, AuditLog
from src.plugins.manager import get_plugin_manager
from src.api.base.context import ApiRequestContext
router = APIRouter(prefix="/api/monitoring", tags=["Monitoring"])
pipeline = ApiRequestPipeline()
@@ -26,7 +30,7 @@ async def get_my_audit_logs(
limit: int = Query(50, description="返回数量限制"),
offset: int = Query(0, ge=0, description="偏移量"),
db: Session = Depends(get_db),
):
) -> Any:
"""
获取我的审计日志
@@ -54,7 +58,7 @@ async def get_my_audit_logs(
@router.get("/rate-limit-status")
async def get_rate_limit_status(request: Request, db: Session = Depends(get_db)):
async def get_rate_limit_status(request: Request, db: Session = Depends(get_db)) -> Any:
"""
获取速率限制状态
@@ -78,7 +82,7 @@ class AuthenticatedApiAdapter(ApiAdapter):
mode = ApiMode.USER
def authorize(self, context): # type: ignore[override]
def authorize(self, context: ApiRequestContext) -> None: # type: ignore[override]
if not context.user:
raise HTTPException(status_code=401, detail="未登录")
@@ -90,7 +94,7 @@ class UserAuditLogsAdapter(AuthenticatedApiAdapter):
limit: int
offset: int
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
user = context.user
if not user:
@@ -136,7 +140,7 @@ class UserAuditLogsAdapter(AuthenticatedApiAdapter):
class UserRateLimitStatusAdapter(AuthenticatedApiAdapter):
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
user = context.user
if not user:
@@ -174,7 +178,7 @@ class UserRateLimitStatusAdapter(AuthenticatedApiAdapter):
return {"user_id": user.id, "api_keys": rate_limit_info}
def _get_rate_limit_plugin():
def _get_rate_limit_plugin() -> Any:
try:
plugin_manager = get_plugin_manager()
return plugin_manager.get_plugin("rate_limit")

View File

@@ -1,6 +1,8 @@
"""OAuth 管理端点(管理员)。"""
from __future__ import annotations
from typing import Any
from fastapi import APIRouter, Depends, Request

View File

@@ -4,6 +4,7 @@
提供系统支持的能力列表,供前端展示和配置使用。
"""
from typing import Any
from fastapi import APIRouter, Depends
from sqlalchemy.orm import Session
@@ -17,7 +18,7 @@ router = APIRouter(prefix="/api/capabilities", tags=["System Catalog"])
@router.get("")
async def list_capabilities():
async def list_capabilities() -> Any:
"""
获取所有能力定义
@@ -49,7 +50,7 @@ async def list_capabilities():
@router.get("/user-configurable")
async def list_user_configurable_capabilities():
async def list_user_configurable_capabilities() -> Any:
"""
获取用户可配置的能力列表
@@ -84,7 +85,7 @@ async def list_user_configurable_capabilities():
async def get_model_supported_capabilities(
model_name: str,
db: Session = Depends(get_db),
):
) -> Any:
"""
获取指定模型支持的能力列表

View File

@@ -3,6 +3,9 @@
不包含敏感信息,普通用户可访问
"""
from __future__ import annotations
from typing import Any
from collections import defaultdict
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
@@ -35,6 +38,7 @@ from src.models.endpoint_models import (
PublicHealthEvent,
)
from src.services.health.endpoint import EndpointHealthService
from src.api.base.context import ApiRequestContext
router = APIRouter(prefix="/api/public", tags=["System Catalog"])
pipeline = ApiRequestPipeline()
@@ -47,7 +51,7 @@ async def get_public_providers(
skip: int = Query(0, description="跳过记录数"),
limit: int = Query(100, description="返回记录数限制"),
db: Session = Depends(get_db),
):
) -> Any:
"""
获取提供商列表(用户视图)
@@ -84,7 +88,7 @@ async def get_public_models(
skip: int = Query(0, description="跳过记录数"),
limit: int = Query(100, description="返回记录数限制"),
db: Session = Depends(get_db),
):
) -> Any:
"""
获取模型列表(用户视图)
@@ -122,7 +126,7 @@ async def get_public_models(
@router.get("/stats", response_model=ProviderStatsResponse)
async def get_public_stats(request: Request, db: Session = Depends(get_db)):
async def get_public_stats(request: Request, db: Session = Depends(get_db)) -> Any:
"""
获取系统统计信息
@@ -147,7 +151,7 @@ async def search_models(
provider_id: int | None = Query(None, description="提供商ID过滤"),
limit: int = Query(20, description="返回记录数限制"),
db: Session = Depends(get_db),
):
) -> Any:
"""
搜索模型
@@ -189,7 +193,7 @@ async def get_public_api_format_health(
lookback_hours: int = Query(6, ge=1, le=168, description="回溯小时数"),
per_format_limit: int = Query(100, ge=10, le=500, description="每个格式的事件数限制"),
db: Session = Depends(get_db),
):
) -> Any:
"""
获取各 API 格式的健康监控数据
@@ -236,7 +240,7 @@ async def get_public_global_models(
is_active: bool | None = Query(None, description="过滤活跃状态"),
search: str | None = Query(None, description="搜索关键词"),
db: Session = Depends(get_db),
):
) -> Any:
"""
获取全局模型GlobalModel列表
@@ -276,7 +280,7 @@ async def get_public_global_models(
class PublicApiAdapter(ApiAdapter):
mode = ApiMode.PUBLIC
def authorize(self, context): # type: ignore[override]
def authorize(self, context: ApiRequestContext) -> None: # type: ignore[override]
return None
@@ -286,7 +290,7 @@ class PublicProvidersAdapter(PublicApiAdapter):
skip: int
limit: int
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
logger.debug("公共API请求提供商列表")
query = db.query(Provider)
@@ -342,7 +346,7 @@ class PublicModelsAdapter(PublicApiAdapter):
skip: int
limit: int
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
logger.debug("公共API请求模型列表")
query = (
@@ -391,7 +395,7 @@ class PublicModelsAdapter(PublicApiAdapter):
class PublicStatsAdapter(PublicApiAdapter):
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
logger.debug("公共API请求系统统计信息")
active_providers = db.query(Provider).filter(Provider.is_active.is_(True)).count()
@@ -428,7 +432,7 @@ class PublicSearchModelsAdapter(PublicApiAdapter):
provider_id: int | None
limit: int
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
logger.debug(f"公共API搜索模型: {self.query}")
query_stmt = (
@@ -490,7 +494,7 @@ class PublicApiFormatHealthMonitorAdapter(PublicApiAdapter):
lookback_hours: int
per_format_limit: int
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
now = datetime.now(timezone.utc)
since = now - timedelta(hours=self.lookback_hours)
@@ -651,7 +655,7 @@ class PublicGlobalModelsAdapter(PublicApiAdapter):
is_active: bool | None
search: str | None
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
logger.debug("公共API请求 GlobalModel 列表")

View File

@@ -7,6 +7,7 @@ Claude API 端点
注意: /v1/models 端点由 models.py 统一处理,根据请求头返回对应格式
"""
from typing import Any
from fastapi import APIRouter, Depends, Request
from sqlalchemy.orm import Session
@@ -27,7 +28,7 @@ pipeline = ApiRequestPipeline()
async def create_message(
http_request: Request,
db: Session = Depends(get_db),
):
) -> Any:
"""
Claude Messages API
@@ -63,7 +64,7 @@ async def create_message(
async def count_tokens(
http_request: Request,
db: Session = Depends(get_db),
):
) -> Any:
"""
Claude Token Count API

View File

@@ -10,6 +10,7 @@ Gemini API 专属端点
- /v1beta/models (列表) 和 /v1beta/models/{model} (详情) 由 models.py 统一处理
"""
from typing import Any
from fastapi import APIRouter, Depends, Request
from sqlalchemy.orm import Session
@@ -53,7 +54,7 @@ async def generate_content(
model: str,
http_request: Request,
db: Session = Depends(get_db),
):
) -> Any:
"""
Gemini generateContent API
@@ -95,7 +96,7 @@ async def stream_generate_content(
model: str,
http_request: Request,
db: Session = Depends(get_db),
):
) -> Any:
"""
Gemini streamGenerateContent API
@@ -133,7 +134,7 @@ async def generate_content_v1(
model: str,
http_request: Request,
db: Session = Depends(get_db),
):
) -> Any:
"""
Gemini generateContent API (v1 兼容)
@@ -147,7 +148,7 @@ async def stream_generate_content_v1(
model: str,
http_request: Request,
db: Session = Depends(get_db),
):
) -> Any:
"""
Gemini streamGenerateContent API (v1 兼容)

View File

@@ -403,7 +403,7 @@ async def _proxy_request(
async def upload_file(
request: Request,
db: Session = Depends(get_db),
):
) -> Any:
"""
上传文件到 Gemini Files API
@@ -463,7 +463,7 @@ async def list_files(
db: Session = Depends(get_db),
pageSize: int | None = None,
pageToken: str | None = None,
):
) -> Any:
"""
列出已上传的文件
@@ -524,7 +524,7 @@ async def get_file(
file_name: str,
request: Request,
db: Session = Depends(get_db),
):
) -> Any:
"""
获取指定文件的元数据
@@ -580,7 +580,7 @@ async def delete_file(
file_name: str,
request: Request,
db: Session = Depends(get_db),
):
) -> Any:
"""
删除指定文件

View File

@@ -1,6 +1,7 @@
"""公开模块状态 API供登录页等使用"""
from typing import Any
from fastapi import APIRouter, Depends
from pydantic import BaseModel
from sqlalchemy.orm import Session
@@ -20,7 +21,7 @@ class AuthModuleInfo(BaseModel):
@router.get("/auth-status", response_model=list[AuthModuleInfo])
async def get_auth_modules_status(db: Session = Depends(get_db)):
async def get_auth_modules_status(db: Session = Depends(get_db)) -> Any:
"""
获取认证模块状态(公开接口)

View File

@@ -7,6 +7,7 @@ OpenAI API 端点
注意: /v1/models 端点由 models.py 统一处理,根据请求头返回对应格式
"""
from typing import Any
from fastapi import APIRouter, Depends, Request
from sqlalchemy.orm import Session
@@ -25,7 +26,7 @@ pipeline = ApiRequestPipeline()
async def create_chat_completion(
http_request: Request,
db: Session = Depends(get_db),
):
) -> Any:
"""
OpenAI Chat Completions API
@@ -58,7 +59,7 @@ async def create_chat_completion(
async def create_responses(
http_request: Request,
db: Session = Depends(get_db),
):
) -> Any:
"""
OpenAI Responses API (CLI)

View File

@@ -4,6 +4,8 @@ System Catalog / 健康检查相关端点
这些是系统工具端点,不需要复杂的 Adapter 抽象。
"""
from __future__ import annotations
from datetime import datetime, timezone
from typing import Any
@@ -97,7 +99,7 @@ def _select_provider(db: Session, provider_name: str | None) -> Provider | None:
@router.get("/v1/health")
async def service_health(db: Session = Depends(get_db)):
async def service_health(db: Session = Depends(get_db)) -> Any:
"""返回服务健康状态与依赖信息"""
active_providers = (
db.query(func.count(Provider.id)).filter(Provider.is_active == True).scalar() or 0
@@ -130,7 +132,7 @@ async def service_health(db: Session = Depends(get_db)):
@router.get("/health")
async def health_check():
async def health_check() -> Any:
"""简单健康检查端点(无需认证)"""
try:
pool_status = get_pool_status()
@@ -156,7 +158,7 @@ async def health_check():
@router.get("/")
async def root(db: Session = Depends(get_db)):
async def root(db: Session = Depends(get_db)) -> Any:
"""Root endpoint - 服务信息概览"""
# 按优先级选择最高优先级的提供商
top_provider = (
@@ -189,7 +191,7 @@ async def list_providers(
include_models: bool = Query(False),
include_endpoints: bool = Query(False),
active_only: bool = Query(True),
):
) -> Any:
"""列出所有 Provider"""
load_options = []
if include_models:
@@ -219,7 +221,7 @@ async def provider_detail(
db: Session = Depends(get_db),
include_models: bool = Query(False),
include_endpoints: bool = Query(False),
):
) -> Any:
"""获取单个 Provider 详情"""
load_options = []
if include_models:
@@ -248,7 +250,7 @@ async def test_connection(
provider: str | None = Query(None),
model: str = Query("claude-3-haiku-20240307"),
api_format: str | None = Query(None),
):
) -> Any:
"""测试 Provider 连接"""
selected_provider = _select_provider(db, provider)
if not selected_provider:
@@ -269,7 +271,7 @@ async def test_connection(
orchestrator = FallbackOrchestrator(db, redis_client)
# 定义请求函数
async def test_request_func(_prov, endpoint, key, _candidate):
async def test_request_func(_prov: Any, endpoint: Any, key: str, _candidate: Any) -> Any:
from src.api.handlers.base.request_builder import get_provider_auth
# 获取认证信息(处理 Service Account 等异步认证场景)

View File

@@ -1,5 +1,8 @@
"""用户 Management Token 管理端点"""
from __future__ import annotations
from typing import Any
from dataclasses import dataclass
from datetime import datetime
@@ -35,7 +38,7 @@ class ManagementTokenApiAdapter(AuthenticatedApiAdapter):
防止用户通过已有的 Token 再创建/修改/删除其他 Token。
"""
def authorize(self, context: ApiRequestContext):
def authorize(self, context: ApiRequestContext) -> Any:
# 先调用父类的认证检查
super().authorize(context)
@@ -65,7 +68,7 @@ class CreateManagementTokenRequest(BaseModel):
@field_validator("expires_at", mode="before")
@classmethod
def parse_expires(cls, v):
def parse_expires(cls, v: Any) -> Any:
return parse_expires_at(v)
@@ -88,7 +91,7 @@ class UpdateManagementTokenRequest(BaseModel):
# 用于追踪哪些字段被显式提供(包括显式设为 null 的情况)
_provided_fields: set[str] = set()
def __init__(self, **data):
def __init__(self, **data: Any) -> None:
# 记录实际传入的字段(包括值为 None 的)
provided = set(data.keys())
super().__init__(**data)
@@ -108,7 +111,7 @@ class UpdateManagementTokenRequest(BaseModel):
@field_validator("expires_at", mode="before")
@classmethod
def parse_expires(cls, v):
def parse_expires(cls, v: Any) -> Any:
# 如果是 None 或空字符串,表示要清空
if v is None or (isinstance(v, str) and not v.strip()):
return None
@@ -125,7 +128,7 @@ async def list_my_management_tokens(
skip: int = Query(0, ge=0),
limit: int = Query(50, ge=1, le=100),
db: Session = Depends(get_db),
):
) -> Any:
"""列出当前用户的 Management Tokens
获取当前登录用户创建的所有 Management Tokens支持按激活状态筛选和分页。
@@ -160,7 +163,7 @@ async def list_my_management_tokens(
@router.post("")
async def create_my_management_token(request: Request, db: Session = Depends(get_db)):
async def create_my_management_token(request: Request, db: Session = Depends(get_db)) -> Any:
"""创建 Management Token
为当前用户创建一个新的 Management Token。
@@ -196,7 +199,7 @@ async def get_my_management_token(
token_id: str,
request: Request,
db: Session = Depends(get_db),
):
) -> Any:
"""获取 Management Token 详情
获取当前用户指定 Token 的详细信息。
@@ -224,7 +227,7 @@ async def get_my_management_token(
@router.put("/{token_id}")
async def update_my_management_token(
token_id: str, request: Request, db: Session = Depends(get_db)
):
) -> Any:
"""更新 Management Token
更新当前用户指定 Token 的信息。支持部分字段更新。
@@ -262,7 +265,7 @@ async def update_my_management_token(
@router.delete("/{token_id}")
async def delete_my_management_token(
token_id: str, request: Request, db: Session = Depends(get_db)
):
) -> Any:
"""删除 Management Token
删除当前用户指定的 Token。
@@ -280,7 +283,7 @@ async def delete_my_management_token(
@router.patch("/{token_id}/status")
async def toggle_my_management_token(
token_id: str, request: Request, db: Session = Depends(get_db)
):
) -> Any:
"""切换 Management Token 状态
启用或禁用当前用户指定的 Token。
@@ -310,7 +313,7 @@ async def toggle_my_management_token(
@router.post("/{token_id}/regenerate")
async def regenerate_my_management_token(
token_id: str, request: Request, db: Session = Depends(get_db)
):
) -> Any:
"""重新生成 Management Token
重新生成当前用户指定 Token 的值,旧 Token 将立即失效。
@@ -350,7 +353,7 @@ class ListMyManagementTokensAdapter(ManagementTokenApiAdapter):
skip: int = 0
limit: int = 50
async def handle(self, context: ApiRequestContext):
async def handle(self, context: ApiRequestContext) -> Any:
from src.config.settings import config
tokens, total = ManagementTokenService.list_tokens(
@@ -385,7 +388,7 @@ class CreateMyManagementTokenAdapter(ManagementTokenApiAdapter):
name: str = "create_my_management_token"
audit_success_event = AuditEventType.MANAGEMENT_TOKEN_CREATED
async def handle(self, context: ApiRequestContext):
async def handle(self, context: ApiRequestContext) -> Any:
body = context.ensure_json_body()
try:
@@ -424,7 +427,7 @@ class GetMyManagementTokenAdapter(ManagementTokenApiAdapter):
name: str = "get_my_management_token"
token_id: str = ""
async def handle(self, context: ApiRequestContext):
async def handle(self, context: ApiRequestContext) -> Any:
token = ManagementTokenService.get_token_by_id(
db=context.db, token_id=self.token_id, user_id=context.user.id
)
@@ -443,7 +446,7 @@ class UpdateMyManagementTokenAdapter(ManagementTokenApiAdapter):
token_id: str = ""
audit_success_event = AuditEventType.MANAGEMENT_TOKEN_UPDATED
async def handle(self, context: ApiRequestContext):
async def handle(self, context: ApiRequestContext) -> Any:
body = context.ensure_json_body()
try:
@@ -496,7 +499,7 @@ class DeleteMyManagementTokenAdapter(ManagementTokenApiAdapter):
token_id: str = ""
audit_success_event = AuditEventType.MANAGEMENT_TOKEN_DELETED
async def handle(self, context: ApiRequestContext):
async def handle(self, context: ApiRequestContext) -> Any:
# 先获取 token 信息用于审计
token = ManagementTokenService.get_token_by_id(
db=context.db, token_id=self.token_id, user_id=context.user.id
@@ -525,7 +528,7 @@ class ToggleMyManagementTokenAdapter(ManagementTokenApiAdapter):
token_id: str = ""
audit_success_event = AuditEventType.MANAGEMENT_TOKEN_UPDATED
async def handle(self, context: ApiRequestContext):
async def handle(self, context: ApiRequestContext) -> Any:
token = ManagementTokenService.toggle_status(
db=context.db, token_id=self.token_id, user_id=context.user.id
)
@@ -553,7 +556,7 @@ class RegenerateMyManagementTokenAdapter(ManagementTokenApiAdapter):
token_id: str = ""
audit_success_event = AuditEventType.MANAGEMENT_TOKEN_UPDATED
async def handle(self, context: ApiRequestContext):
async def handle(self, context: ApiRequestContext) -> Any:
token, raw_token, old_token_hash = ManagementTokenService.regenerate_token(
db=context.db, token_id=self.token_id, user_id=context.user.id
)

View File

@@ -1,5 +1,8 @@
"""用户个人 API 端点。"""
from __future__ import annotations
from typing import Any
from dataclasses import dataclass
from datetime import datetime, timezone
@@ -27,6 +30,7 @@ from src.models.database import ApiKey, GlobalModel, Model, Provider, Usage, Use
from src.services.usage.service import UsageService
from src.services.user.apikey import ApiKeyService
from src.services.user.preference import PreferenceService
from src.api.base.context import ApiRequestContext
router = APIRouter(prefix="/api/users/me", tags=["User Profile"])
@@ -34,7 +38,7 @@ pipeline = ApiRequestPipeline()
@router.get("")
async def get_my_profile(request: Request, db: Session = Depends(get_db)):
async def get_my_profile(request: Request, db: Session = Depends(get_db)) -> Any:
"""
获取当前用户信息
@@ -47,7 +51,7 @@ async def get_my_profile(request: Request, db: Session = Depends(get_db)):
@router.put("")
async def update_my_profile(request: Request, db: Session = Depends(get_db)):
async def update_my_profile(request: Request, db: Session = Depends(get_db)) -> None:
"""
更新个人信息
@@ -62,7 +66,7 @@ async def update_my_profile(request: Request, db: Session = Depends(get_db)):
@router.patch("/password")
async def change_my_password(request: Request, db: Session = Depends(get_db)):
async def change_my_password(request: Request, db: Session = Depends(get_db)) -> Any:
"""
修改密码
@@ -80,7 +84,7 @@ async def change_my_password(request: Request, db: Session = Depends(get_db)):
@router.get("/api-keys")
async def list_my_api_keys(request: Request, db: Session = Depends(get_db)):
async def list_my_api_keys(request: Request, db: Session = Depends(get_db)) -> Any:
"""
获取 API 密钥列表
@@ -94,7 +98,7 @@ async def list_my_api_keys(request: Request, db: Session = Depends(get_db)):
@router.post("/api-keys")
async def create_my_api_key(request: Request, db: Session = Depends(get_db)):
async def create_my_api_key(request: Request, db: Session = Depends(get_db)) -> Any:
"""
创建 API 密钥
@@ -115,7 +119,7 @@ async def get_my_api_key(
request: Request,
include_key: bool = Query(False, description="是否返回完整密钥"),
db: Session = Depends(get_db),
):
) -> Any:
"""
获取 API 密钥详情
@@ -135,7 +139,7 @@ async def get_my_api_key(
@router.delete("/api-keys/{key_id}")
async def delete_my_api_key(key_id: str, request: Request, db: Session = Depends(get_db)):
async def delete_my_api_key(key_id: str, request: Request, db: Session = Depends(get_db)) -> None:
"""
删除 API 密钥
@@ -149,7 +153,7 @@ async def delete_my_api_key(key_id: str, request: Request, db: Session = Depends
@router.patch("/api-keys/{key_id}")
async def toggle_my_api_key(key_id: str, request: Request, db: Session = Depends(get_db)):
async def toggle_my_api_key(key_id: str, request: Request, db: Session = Depends(get_db)) -> Any:
"""
切换 API 密钥状态
@@ -174,7 +178,7 @@ async def get_my_usage(
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),
):
) -> Any:
"""
获取使用统计
@@ -200,7 +204,7 @@ async def get_my_active_requests(
request: Request,
ids: str | None = Query(None, description="请求 ID 列表,逗号分隔"),
db: Session = Depends(get_db),
):
) -> Any:
"""
获取活跃请求状态
@@ -219,7 +223,7 @@ async def get_my_interval_timeline(
hours: int = Query(24, ge=1, le=720, description="分析最近多少小时的数据"),
limit: int = Query(5000, ge=100, le=20000, description="最大返回数据点数量"),
db: Session = Depends(get_db),
):
) -> Any:
"""
获取请求间隔时间线
@@ -235,7 +239,7 @@ async def get_my_interval_timeline(
async def get_my_activity_heatmap(
request: Request,
db: Session = Depends(get_db),
):
) -> Any:
"""
获取活动热力图数据
@@ -249,7 +253,7 @@ async def get_my_activity_heatmap(
@router.get("/providers")
async def list_available_providers(request: Request, db: Session = Depends(get_db)):
async def list_available_providers(request: Request, db: Session = Depends(get_db)) -> Any:
"""
获取可用提供商列表
@@ -268,7 +272,7 @@ async def list_available_models(
limit: int = Query(100, ge=1, le=1000, description="返回记录数限制"),
search: str | None = Query(None, description="搜索关键词"),
db: Session = Depends(get_db),
):
) -> Any:
"""
获取用户可用的模型列表
@@ -290,7 +294,7 @@ async def list_available_models(
@router.get("/endpoint-status")
async def get_endpoint_status(request: Request, db: Session = Depends(get_db)):
async def get_endpoint_status(request: Request, db: Session = Depends(get_db)) -> Any:
"""
获取端点健康状态
@@ -313,7 +317,7 @@ async def update_api_key_providers(
api_key_id: str,
request: Request,
db: Session = Depends(get_db),
):
) -> Any:
"""
更新 API 密钥可用提供商
@@ -334,7 +338,7 @@ async def update_api_key_capabilities(
api_key_id: str,
request: Request,
db: Session = Depends(get_db),
):
) -> Any:
"""
更新 API 密钥能力配置
@@ -354,7 +358,7 @@ async def update_api_key_capabilities(
@router.get("/preferences")
async def get_my_preferences(request: Request, db: Session = Depends(get_db)):
async def get_my_preferences(request: Request, db: Session = Depends(get_db)) -> Any:
"""
获取偏好设置
@@ -367,7 +371,7 @@ async def get_my_preferences(request: Request, db: Session = Depends(get_db)):
@router.put("/preferences")
async def update_my_preferences(request: Request, db: Session = Depends(get_db)):
async def update_my_preferences(request: Request, db: Session = Depends(get_db)) -> None:
"""
更新偏好设置
@@ -386,7 +390,7 @@ async def update_my_preferences(request: Request, db: Session = Depends(get_db))
@router.get("/model-capabilities")
async def get_model_capability_settings(request: Request, db: Session = Depends(get_db)):
async def get_model_capability_settings(request: Request, db: Session = Depends(get_db)) -> Any:
"""
获取模型能力配置
@@ -399,7 +403,7 @@ async def get_model_capability_settings(request: Request, db: Session = Depends(
@router.put("/model-capabilities")
async def update_model_capability_settings(request: Request, db: Session = Depends(get_db)):
async def update_model_capability_settings(request: Request, db: Session = Depends(get_db)) -> None:
"""
更新模型能力配置
@@ -418,14 +422,14 @@ async def update_model_capability_settings(request: Request, db: Session = Depen
class MeProfileAdapter(AuthenticatedApiAdapter):
"""获取当前用户信息的适配器"""
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
return PreferenceService.get_user_with_preferences(context.db, context.user.id)
class UpdateProfileAdapter(AuthenticatedApiAdapter):
"""更新用户个人信息的适配器"""
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
user = context.user
payload = context.ensure_json_body()
@@ -462,7 +466,7 @@ class UpdateProfileAdapter(AuthenticatedApiAdapter):
class ChangePasswordAdapter(AuthenticatedApiAdapter):
"""修改用户密码的适配器"""
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
user = context.user
payload = context.ensure_json_body()
@@ -504,7 +508,7 @@ class ChangePasswordAdapter(AuthenticatedApiAdapter):
class ListMyApiKeysAdapter(AuthenticatedApiAdapter):
"""获取用户 API 密钥列表的适配器"""
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
user = context.user
@@ -576,7 +580,7 @@ class ListMyApiKeysAdapter(AuthenticatedApiAdapter):
class CreateMyApiKeyAdapter(AuthenticatedApiAdapter):
"""创建 API 密钥的适配器"""
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
payload = context.ensure_json_body()
try:
request = CreateMyApiKeyRequest.model_validate(payload)
@@ -609,7 +613,7 @@ class GetMyFullKeyAdapter(AuthenticatedApiAdapter):
key_id: str
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
user = context.user
@@ -643,7 +647,7 @@ class GetMyApiKeyDetailAdapter(AuthenticatedApiAdapter):
key_id: str
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
user = context.user
@@ -674,7 +678,7 @@ class DeleteMyApiKeyAdapter(AuthenticatedApiAdapter):
key_id: str
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
api_key = (
context.db.query(ApiKey)
.filter(ApiKey.id == self.key_id, ApiKey.user_id == context.user.id)
@@ -695,7 +699,7 @@ class ToggleMyApiKeyAdapter(AuthenticatedApiAdapter):
key_id: str
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
api_key = (
context.db.query(ApiKey)
.filter(ApiKey.id == self.key_id, ApiKey.user_id == context.user.id)
@@ -725,7 +729,7 @@ class GetUsageAdapter(AuthenticatedApiAdapter):
limit: int = 100
offset: int = 0
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
from sqlalchemy import or_
from src.models.database import ProviderEndpoint
@@ -983,7 +987,7 @@ class GetActiveRequestsAdapter(AuthenticatedApiAdapter):
ids: str | None = None
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
from src.services.usage import UsageService
db = context.db
@@ -1005,7 +1009,7 @@ class GetMyIntervalTimelineAdapter(AuthenticatedApiAdapter):
hours: int
limit: int
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
user = context.user
@@ -1022,7 +1026,7 @@ class GetMyIntervalTimelineAdapter(AuthenticatedApiAdapter):
class GetMyActivityHeatmapAdapter(AuthenticatedApiAdapter):
"""获取用户活动热力图数据的适配器(带 Redis 缓存)"""
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
user = context.user
result = await UsageService.get_cached_heatmap(
db=context.db,
@@ -1045,7 +1049,7 @@ class ListAvailableModelsAdapter(AuthenticatedApiAdapter):
limit: int
search: str | None
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
from sqlalchemy import or_
from src.api.base.models_service import AccessRestrictions
@@ -1215,7 +1219,7 @@ class ListAvailableModelsAdapter(AuthenticatedApiAdapter):
class ListAvailableProvidersAdapter(AuthenticatedApiAdapter):
"""获取可用提供商列表的适配器"""
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
from sqlalchemy.orm import selectinload
@@ -1288,7 +1292,7 @@ class UpdateApiKeyProvidersAdapter(AuthenticatedApiAdapter):
api_key_id: str
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db
user = context.user
payload = context.ensure_json_body()
@@ -1339,7 +1343,7 @@ class UpdateApiKeyCapabilitiesAdapter(AuthenticatedApiAdapter):
api_key_id: str
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
from src.core.key_capabilities import CAPABILITY_DEFINITIONS, CapabilityConfigMode
from src.models.database import AuditEventType
from src.services.system.audit import audit_service
@@ -1403,7 +1407,7 @@ class UpdateApiKeyCapabilitiesAdapter(AuthenticatedApiAdapter):
class GetPreferencesAdapter(AuthenticatedApiAdapter):
"""获取用户偏好设置的适配器"""
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
preferences = PreferenceService.get_or_create_preferences(context.db, context.user.id)
return {
"avatar_url": preferences.avatar_url,
@@ -1426,7 +1430,7 @@ class GetPreferencesAdapter(AuthenticatedApiAdapter):
class UpdatePreferencesAdapter(AuthenticatedApiAdapter):
"""更新用户偏好设置的适配器"""
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
payload = context.ensure_json_body()
try:
request = UpdatePreferencesRequest.model_validate(payload)
@@ -1455,7 +1459,7 @@ class UpdatePreferencesAdapter(AuthenticatedApiAdapter):
class GetModelCapabilitySettingsAdapter(AuthenticatedApiAdapter):
"""获取用户的模型能力配置"""
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
user = context.user
return {
"model_capability_settings": user.model_capability_settings or {},
@@ -1465,7 +1469,7 @@ class GetModelCapabilitySettingsAdapter(AuthenticatedApiAdapter):
class UpdateModelCapabilitySettingsAdapter(AuthenticatedApiAdapter):
"""更新用户的模型能力配置"""
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
from src.core.key_capabilities import CAPABILITY_DEFINITIONS, CapabilityConfigMode
from src.models.database import AuditEventType
from src.services.cache.user_cache import UserCacheService
@@ -1539,7 +1543,7 @@ class GetEndpointStatusAdapter(AuthenticatedApiAdapter):
_cache_ttl = 60 # 缓存60秒
@classmethod
async def _get_cache(cls):
async def _get_cache(cls) -> Any:
"""获取缓存后端实例(懒加载)"""
if cls._cache_backend is None:
from src.services.cache.backend import get_cache_backend
@@ -1551,7 +1555,7 @@ class GetEndpointStatusAdapter(AuthenticatedApiAdapter):
)
return cls._cache_backend
async def handle(self, context): # type: ignore[override]
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
from src.services.health.endpoint import EndpointHealthService
db = context.db

View File

@@ -68,7 +68,7 @@ def build_proxy_url(proxy_config: dict[str, Any]) -> str | None:
if not proxy_config.get("enabled", True):
return None
proxy_url = proxy_config.get("url")
proxy_url: str | None = proxy_config.get("url")
if not proxy_url:
return None
@@ -113,7 +113,7 @@ class HTTPClientPool:
# 代理客户端缓存上限(避免内存泄漏)
_max_proxy_clients: int = 50
def __new__(cls):
def __new__(cls) -> "HTTPClientPool":
if cls._instance is None:
cls._instance = super().__new__(cls)
return cls._instance
@@ -214,7 +214,7 @@ class HTTPClientPool:
}
default_config.update(kwargs)
cls._clients[name] = httpx.AsyncClient(**default_config)
cls._clients[name] = httpx.AsyncClient(**default_config) # type: ignore[arg-type]
logger.debug(f"创建命名HTTP客户端: {name}")
return cls._clients[name]
@@ -304,7 +304,7 @@ class HTTPClientPool:
if proxy_url:
client_config["proxy"] = proxy_url
client = httpx.AsyncClient(**client_config)
client = httpx.AsyncClient(**client_config) # type: ignore[arg-type]
cls._proxy_clients[cache_key] = (client, time.time())
logger.debug(
@@ -315,7 +315,7 @@ class HTTPClientPool:
return client
@classmethod
async def close_all(cls):
async def close_all(cls) -> None:
"""关闭所有HTTP客户端"""
if cls._default_client is not None:
await cls._default_client.aclose()
@@ -341,7 +341,7 @@ class HTTPClientPool:
@classmethod
@asynccontextmanager
async def get_temp_client(cls, **kwargs: Any):
async def get_temp_client(cls, **kwargs: Any) -> Any:
"""
获取临时HTTP客户端(上下文管理器)
@@ -363,7 +363,7 @@ class HTTPClientPool:
}
default_config.update(kwargs)
client = httpx.AsyncClient(**default_config)
client = httpx.AsyncClient(**default_config) # type: ignore[arg-type]
try:
yield client
finally:
@@ -412,7 +412,7 @@ class HTTPClientPool:
logger.debug(f"创建带代理的HTTP客户端(一次性): {proxy_config.get('url', 'unknown')}")
client_config.update(kwargs)
return httpx.AsyncClient(**client_config)
return httpx.AsyncClient(**client_config) # type: ignore[arg-type]
@classmethod
def get_pool_stats(cls) -> dict[str, Any]:
@@ -431,6 +431,6 @@ def get_http_client() -> httpx.AsyncClient:
return HTTPClientPool.get_default_client()
async def close_http_clients():
async def close_http_clients() -> None:
"""关闭所有HTTP客户端的便捷函数"""
await HTTPClientPool.close_all()

View File

@@ -39,13 +39,13 @@ class RedisClientManager:
_instance: RedisClientManager | None = None
_redis: aioredis.Redis | None = None
def __new__(cls):
def __new__(cls) -> "RedisClientManager":
"""单例模式"""
if cls._instance is None:
cls._instance = super().__new__(cls)
return cls._instance
def __init__(self):
def __init__(self) -> None:
# 避免重复初始化
if getattr(self, "_initialized", False):
return

View File

@@ -3,6 +3,7 @@
从环境变量或 .env 文件加载配置
"""
from typing import Any
import os
from pathlib import Path
@@ -383,7 +384,7 @@ class Config:
return self._database_url
@database_url.setter
def database_url(self, value: str):
def database_url(self, value: str) -> Any:
"""允许在测试中设置数据库 URL"""
self._database_url = value
@@ -452,7 +453,7 @@ class Config:
return errors
def __repr__(self):
def __repr__(self) -> None:
"""配置信息字符串表示"""
return f"""
Configuration:

View File

@@ -7,6 +7,7 @@
- 关键数据(计费)仍然立即 commit
"""
from typing import Any
import asyncio
from src.core.logger import logger
@@ -26,13 +27,13 @@ class BatchCommitter:
self._lock = asyncio.Lock()
self._task = None
async def start(self):
async def start(self) -> Any:
"""启动后台批量提交任务"""
if self._task is None:
self._task = asyncio.create_task(self._batch_commit_loop())
logger.info(f"批量提交器已启动,间隔: {self.interval_seconds}s")
async def stop(self):
async def stop(self) -> Any:
"""停止后台任务"""
if self._task:
self._task.cancel()
@@ -43,7 +44,7 @@ class BatchCommitter:
self._task = None
logger.info("批量提交器已停止")
def mark_dirty(self, session: Session):
def mark_dirty(self, session: Session) -> Any:
"""标记 Session 有待提交的更改"""
# 请求级事务由中间件统一 commit/rollback避免后台任务在请求中途误提交。
if session is None:
@@ -52,7 +53,7 @@ class BatchCommitter:
return
self._pending_sessions.add(session)
async def _batch_commit_loop(self):
async def _batch_commit_loop(self) -> None:
"""后台批量提交循环"""
while True:
try:
@@ -65,7 +66,7 @@ class BatchCommitter:
except Exception as e:
logger.error(f"批量提交出错: {e}")
async def _commit_all(self):
async def _commit_all(self) -> None:
"""提交所有待处理的 Session"""
async with self._lock:
if not self._pending_sessions:
@@ -107,13 +108,13 @@ def get_batch_committer() -> BatchCommitter:
return _batch_committer
async def init_batch_committer():
async def init_batch_committer() -> None:
"""初始化并启动批量提交器"""
committer = get_batch_committer()
await committer.start()
async def shutdown_batch_committer():
async def shutdown_batch_committer() -> None:
"""关闭批量提交器"""
committer = get_batch_committer()
await committer.stop()

View File

@@ -2,6 +2,8 @@
缓存服务 - 统一的缓存抽象层
"""
from __future__ import annotations
import json
from typing import Any

View File

@@ -4,6 +4,8 @@
from __future__ import annotations
def extract_error_message(error: Exception, status_code: int | None = None) -> str:
"""
从异常中提取错误消息,优先使用上游原始响应(用于链路追踪/调试)

View File

@@ -7,6 +7,9 @@
- 开发环境可返回详细信息用于调试
"""
from __future__ import annotations
from starlette.requests import Request
import asyncio
import re
import traceback
@@ -140,7 +143,7 @@ def translate_pydantic_errors(errors: list[dict[str, Any]]) -> str:
# 延迟导入韧性管理器,避免循环导入
def get_resilience_manager():
def get_resilience_manager() -> Any:
try:
from ..core.resilience import resilience_manager
@@ -173,7 +176,7 @@ class ProviderException(ProxyException):
message: str,
provider_name: str | None = None,
request_metadata: Any | None = None,
**kwargs,
**kwargs: Any,
):
self.request_metadata = request_metadata # 保存元数据以便传递
details = {"provider": provider_name} if provider_name else {}
@@ -537,7 +540,7 @@ class ThinkingSignatureException(UpstreamClientException):
message: str,
provider_name: str | None = None,
upstream_error: str | None = None,
request_metadata: Any = None,
request_metadata: Any | None = None,
):
super().__init__(
message=message,
@@ -688,17 +691,17 @@ class ExceptionHandlers:
"""FastAPI异常处理器"""
@staticmethod
async def handle_proxy_exception(request, exc: ProxyException):
async def handle_proxy_exception(request: Request, exc: ProxyException) -> None:
"""处理代理异常"""
return ErrorResponse.from_exception(exc)
@staticmethod
async def handle_http_exception(request, exc: HTTPException):
async def handle_http_exception(request: Request, exc: HTTPException) -> None:
"""处理HTTP异常"""
return ErrorResponse.from_exception(exc)
@staticmethod
async def handle_generic_exception(request, exc: Exception):
async def handle_generic_exception(request: Request, exc: Exception) -> None:
"""处理通用异常 - 集成韧性管理"""
# 首先检查是否为HTTPException如果是则委托给HTTP异常处理器

View File

@@ -21,6 +21,8 @@
logger.exception("异常,带堆栈")
"""
from __future__ import annotations
import logging
import os
import sys
@@ -115,19 +117,19 @@ if not DISABLE_FILE_LOG:
file_log_config["diagnose"] = False
# 主日志文件 - 所有级别
logger.add(
logger.add( # type: ignore[call-overload]
log_dir / "app.log",
level="DEBUG",
**file_log_config, # type: ignore[arg-type]
**file_log_config,
)
# 错误日志文件 - 仅 ERROR 及以上
error_log_config = file_log_config.copy()
error_log_config["rotation"] = "50 MB"
logger.add(
logger.add( # type: ignore[call-overload]
log_dir / "error.log",
level="ERROR",
**error_log_config, # type: ignore[arg-type]
**error_log_config,
)
# ============================================================================

View File

@@ -35,7 +35,7 @@ class ModuleRegistry:
_instance: ModuleRegistry | None = None
def __init__(self):
def __init__(self) -> None:
self._modules: dict[str, ModuleDefinition] = {}
self._initialized: set[str] = set()

View File

@@ -21,11 +21,11 @@ class TokenCounter:
"claude-2": "cl100k_base",
}
def __init__(self):
def __init__(self) -> None:
self._encodings = {}
self._default_encoding = None
def _get_encoding(self, model: str):
def _get_encoding(self, model: str) -> Any:
"""获取模型对应的编码器"""
# 标准化模型名称
model_base = model.lower().split("-")[0]

View File

@@ -3,6 +3,8 @@
提供全局的错误处理、自动恢复、降级策略和用户友好的错误体验
"""
from __future__ import annotations
import asyncio
import functools
import threading
@@ -75,7 +77,7 @@ class CircuitBreaker:
self.state = "closed" # closed, open, half-open
self._lock = threading.Lock()
def call(self, func: Callable, *args, **kwargs):
def call(self, func: Callable[..., Any], *args: Any, **kwargs: Any) -> Any:
"""执行函数调用,应用熔断逻辑"""
with self._lock:
if self.state == "open":
@@ -98,12 +100,12 @@ class CircuitBreaker:
return True
return time.time() - self.last_failure_time >= self.timeout
def _on_success(self):
def _on_success(self) -> None:
"""成功时重置计数器"""
self.failure_count = 0
self.state = "closed"
def _on_failure(self):
def _on_failure(self) -> None:
"""失败时增加计数器"""
self.failure_count += 1
self.last_failure_time = time.time()
@@ -114,14 +116,14 @@ class CircuitBreaker:
class ResilienceManager:
"""系统韧性管理器"""
def __init__(self):
def __init__(self) -> None:
self.error_patterns: list[ErrorPattern] = []
self.circuit_breakers: dict[str, CircuitBreaker] = {}
self.error_stats: dict[str, int] = {}
self.last_errors: list[dict[str, Any]] = []
self._setup_default_patterns()
def _setup_default_patterns(self):
def _setup_default_patterns(self) -> None:
"""设置默认错误处理模式"""
# 数据库连接错误 - 只捕获特定的数据库相关异常
@@ -187,7 +189,7 @@ class ResilienceManager:
)
)
def add_error_pattern(self, pattern: ErrorPattern):
def add_error_pattern(self, pattern: ErrorPattern) -> None:
"""添加错误处理模式"""
self.error_patterns.append(pattern)
@@ -277,12 +279,12 @@ resilience_manager = ResilienceManager()
def resilient_operation(
operation_name: str = None,
max_retries: int = None,
retry_delay: float = None,
circuit_breaker_key: str = None,
operation_name: str | None = None,
max_retries: int | None = None,
retry_delay: float | None = None,
circuit_breaker_key: str | None = None,
context: dict[str, Any] = None,
):
) -> Any:
"""
韧性操作装饰器
自动处理重试、熔断、错误记录等
@@ -290,7 +292,7 @@ def resilient_operation(
def decorator(func: Callable) -> Callable:
@functools.wraps(func)
async def async_wrapper(*args, **kwargs):
async def async_wrapper(*args: Any, **kwargs: Any) -> Any:
op_name = operation_name or f"{func.__module__}.{func.__name__}"
retries = max_retries or 3
delay = retry_delay or 1.0
@@ -342,7 +344,7 @@ def resilient_operation(
raise last_error
@functools.wraps(func)
def sync_wrapper(*args, **kwargs):
def sync_wrapper(*args: Any, **kwargs: Any) -> None:
# 对于同步函数,创建异步包装器并运行
return asyncio.run(async_wrapper(*args, **kwargs))
@@ -356,7 +358,7 @@ def resilient_operation(
@asynccontextmanager
async def safe_operation(operation_name: str, context: dict[str, Any] = None):
async def safe_operation(operation_name: str, context: dict[str, Any] = None) -> Any:
"""
安全操作上下文管理器
自动处理异常并提供用户友好的错误信息
@@ -381,7 +383,7 @@ async def safe_operation(operation_name: str, context: dict[str, Any] = None):
logger.warning(f"操作警告 [{error_result['error_id']}]: {error_result['user_message']}")
def graceful_degradation(fallback_func: Callable = None, fallback_value: Any = None):
def graceful_degradation(fallback_func: Callable | None = None, fallback_value: Any | None = None) -> Any:
"""
优雅降级装饰器
当主要功能失败时,自动切换到备用方案
@@ -389,7 +391,7 @@ def graceful_degradation(fallback_func: Callable = None, fallback_value: Any = N
def decorator(func: Callable) -> Callable:
@functools.wraps(func)
async def async_wrapper(*args, **kwargs):
async def async_wrapper(*args: Any, **kwargs: Any) -> Any:
try:
if asyncio.iscoroutinefunction(func):
return await func(*args, **kwargs)

View File

@@ -3,6 +3,10 @@
采用模块化架构设计
"""
from __future__ import annotations
from typing import Any
from contextlib import asynccontextmanager
import uvicorn
@@ -32,7 +36,7 @@ from src.plugins.manager import get_plugin_manager
async def initialize_providers():
async def initialize_providers() -> None:
"""从数据库初始化提供商(仅用于日志记录)"""
from sqlalchemy.orm import Session, selectinload
@@ -61,9 +65,9 @@ async def initialize_providers():
logger.info(f"从数据库加载了 {len(providers)} 个活跃提供商")
for provider in providers:
# 统计端点信息
endpoint_count = len(provider.endpoints) if provider.endpoints else 0
endpoint_count = len(provider.endpoints) if provider.endpoints else 0 # type: ignore[arg-type]
active_endpoints = (
sum(1 for ep in provider.endpoints if ep.is_active) if provider.endpoints else 0
sum(1 for ep in provider.endpoints if ep.is_active) if provider.endpoints else 0 # type: ignore[misc,attr-defined]
)
logger.info(f"提供商: {provider.name} (端点: {active_endpoints}/{endpoint_count})")
@@ -76,7 +80,7 @@ async def initialize_providers():
@asynccontextmanager
async def lifespan(app: FastAPI):
async def lifespan(app: FastAPI) -> Any:
"""应用生命周期管理"""
# 禁用uvicorn的access日志(在子进程中执行)
import logging
@@ -213,7 +217,7 @@ async def lifespan(app: FastAPI):
await quota_scheduler.start()
else:
logger.info("检测到其他 worker 已运行额度调度器,本实例跳过")
quota_scheduler = None
quota_scheduler = None # type: ignore[assignment]
# 启动维护调度器
maintenance_scheduler_active = await task_coordinator.acquire("maintenance_scheduler")
@@ -222,7 +226,7 @@ async def lifespan(app: FastAPI):
await maintenance_scheduler.start()
else:
logger.info("检测到其他 worker 已运行维护调度器,本实例跳过")
maintenance_scheduler = None
maintenance_scheduler = None # type: ignore[assignment]
# 启动模型自动获取调度器
model_fetch_scheduler_active = await task_coordinator.acquire("model_fetch_scheduler")
@@ -231,7 +235,7 @@ async def lifespan(app: FastAPI):
await model_fetch_scheduler.start()
else:
logger.info("检测到其他 worker 已运行模型获取调度器,本实例跳过")
model_fetch_scheduler = None
model_fetch_scheduler = None # type: ignore[assignment]
# 启动统一的定时任务调度器
from src.services.system.scheduler import get_scheduler
@@ -416,9 +420,9 @@ app = FastAPI(
# - propagate_provider_exceptions=True (默认): 不注册,让异常传播到路由层以记录 provider_request_headers
# - propagate_provider_exceptions=False: 注册全局处理器统一处理
if not config.propagate_provider_exceptions:
app.add_exception_handler(ProxyException, ExceptionHandlers.handle_proxy_exception)
app.add_exception_handler(Exception, ExceptionHandlers.handle_generic_exception)
app.add_exception_handler(HTTPException, ExceptionHandlers.handle_http_exception)
app.add_exception_handler(ProxyException, ExceptionHandlers.handle_proxy_exception) # type: ignore[arg-type]
app.add_exception_handler(Exception, ExceptionHandlers.handle_generic_exception) # type: ignore[arg-type]
app.add_exception_handler(HTTPException, ExceptionHandlers.handle_http_exception) # type: ignore[arg-type]
# 添加插件中间件(包含认证、审计、速率限制等功能)
app.add_middleware(PluginMiddleware)
@@ -455,7 +459,7 @@ app.include_router(monitoring_router) # 监控端点
def main():
def main() -> Any:
# 初始化新日志系统
debug_mode = config.environment == "development"
# 日志系统已在导入时自动初始化

View File

@@ -4,6 +4,8 @@
提供完整的输入验证和安全过滤
"""
from __future__ import annotations
import re
from datetime import datetime
from typing import Any

View File

@@ -2,6 +2,8 @@
API端点请求/响应模型定义
"""
from __future__ import annotations
import re
from datetime import datetime
from typing import Any, Literal
@@ -21,7 +23,7 @@ class LoginRequest(BaseModel):
@classmethod
@field_validator("password")
def validate_password(cls, v):
def validate_password(cls, v: Any) -> Any:
"""验证密码不为空且去除前后空格"""
v = v.strip()
if not v:
@@ -29,7 +31,7 @@ class LoginRequest(BaseModel):
return v
@model_validator(mode="after")
def validate_login(self):
def validate_login(self) -> Any:
"""根据认证类型校验并规范化登录标识"""
identifier = self.email.strip()
@@ -84,7 +86,7 @@ class RegisterRequest(BaseModel):
@field_validator("email")
@classmethod
def validate_email(cls, v):
def validate_email(cls, v: Any) -> Any:
"""验证邮箱格式(如果提供)"""
if v is None:
return None
@@ -98,7 +100,7 @@ class RegisterRequest(BaseModel):
@classmethod
@field_validator("username")
def validate_username(cls, v):
def validate_username(cls, v: Any) -> Any:
"""验证用户名格式"""
v = v.strip()
if not v:
@@ -109,7 +111,7 @@ class RegisterRequest(BaseModel):
@classmethod
@field_validator("password")
def validate_password(cls, v):
def validate_password(cls, v: Any) -> Any:
"""验证密码强度"""
if len(v) < 6:
raise ValueError("密码至少需要6个字符")
@@ -145,7 +147,7 @@ class SendVerificationCodeRequest(BaseModel):
@field_validator("email")
@classmethod
def validate_email(cls, v):
def validate_email(cls, v: Any) -> Any:
"""验证邮箱格式"""
email_pattern = r"^[a-zA-Z0-9._%+-]+@[a-zA-Z0-9.-]+\.[a-zA-Z]{2,}$"
if not re.match(email_pattern, v):
@@ -169,7 +171,7 @@ class VerifyEmailRequest(BaseModel):
@field_validator("email")
@classmethod
def validate_email(cls, v):
def validate_email(cls, v: Any) -> Any:
"""验证邮箱格式"""
email_pattern = r"^[a-zA-Z0-9._%+-]+@[a-zA-Z0-9.-]+\.[a-zA-Z]{2,}$"
if not re.match(email_pattern, v):
@@ -178,7 +180,7 @@ class VerifyEmailRequest(BaseModel):
@field_validator("code")
@classmethod
def validate_code(cls, v):
def validate_code(cls, v: Any) -> Any:
"""验证验证码格式"""
v = v.strip()
if not v.isdigit():
@@ -202,7 +204,7 @@ class VerificationStatusRequest(BaseModel):
@field_validator("email")
@classmethod
def validate_email(cls, v):
def validate_email(cls, v: Any) -> Any:
"""验证邮箱格式"""
v = v.strip().lower()
if not v:
@@ -248,7 +250,7 @@ class CreateUserRequest(BaseModel):
@field_validator("quota_usd", mode="before")
@classmethod
def validate_quota_usd(cls, v):
def validate_quota_usd(cls, v: Any) -> Any:
"""验证配额值null表示使用系统默认配额"""
if v is None:
return None
@@ -274,7 +276,7 @@ class CreateUserRequest(BaseModel):
@field_validator("username")
@classmethod
def validate_username(cls, v):
def validate_username(cls, v: Any) -> Any:
"""验证用户名格式"""
v = v.strip()
if not v:
@@ -285,7 +287,7 @@ class CreateUserRequest(BaseModel):
@classmethod
@field_validator("password")
def validate_password(cls, v):
def validate_password(cls, v: Any) -> Any:
"""验证密码强度"""
if len(v) < 6:
raise ValueError("密码至少需要6个字符")
@@ -313,7 +315,7 @@ class UpdateUserRequest(BaseModel):
@field_validator("quota_usd", mode="before")
@classmethod
def validate_quota_usd(cls, v):
def validate_quota_usd(cls, v: Any) -> Any:
"""验证配额值允许null表示无限制"""
if v is None:
return None

View File

@@ -2,6 +2,9 @@
数据库模型定义
"""
from __future__ import annotations
from typing import Any
import hashlib
import secrets
import uuid
@@ -131,7 +134,7 @@ class User(Base):
)
audit_logs = relationship("AuditLog", back_populates="user", passive_deletes=True)
def set_password(self, password: str):
def set_password(self, password: str) -> None:
"""设置密码"""
self.password_hash = bcrypt.hashpw(password.encode("utf-8"), bcrypt.gensalt()).decode(
"utf-8"
@@ -373,7 +376,7 @@ class Usage(Base):
provider_endpoint = relationship("ProviderEndpoint")
provider_api_key = relationship("ProviderAPIKey")
def get_request_body(self):
def get_request_body(self) -> Any:
"""获取请求体(自动解压)"""
if self.request_body is not None:
return self.request_body
@@ -383,7 +386,7 @@ class Usage(Base):
return decompress_json(self.request_body_compressed)
return None
def get_response_body(self):
def get_response_body(self) -> Any:
"""获取响应体(自动解压)"""
if self.response_body is not None:
return self.response_body

View File

@@ -6,7 +6,7 @@ LDAP 认证模块
from __future__ import annotations
from typing import TYPE_CHECKING
from typing import TYPE_CHECKING, Any
from src.core.modules.base import (
ModuleCategory,
@@ -19,7 +19,7 @@ if TYPE_CHECKING:
from sqlalchemy.orm import Session
def _get_router():
def _get_router() -> Any:
"""延迟导入路由(避免启动时加载重依赖)"""
from src.api.admin.ldap import router

View File

@@ -6,7 +6,7 @@ OAuth 认证模块
from __future__ import annotations
from typing import TYPE_CHECKING
from typing import TYPE_CHECKING, Any
from src.core.modules.base import (
ModuleCategory,
@@ -19,7 +19,7 @@ if TYPE_CHECKING:
from sqlalchemy.orm import Session
def _get_router():
def _get_router() -> Any:
"""延迟导入路由(避免启动时加载重依赖/副作用)。"""
# 延迟 discover避免 alembic/mypy 等场景导入时触发 entry_points 解析
from src.services.auth.oauth.registry import get_oauth_provider_registry

View File

@@ -4,6 +4,8 @@ API Key认证插件
"""
from __future__ import annotations
from fastapi import Request
from sqlalchemy.orm import Session
@@ -21,7 +23,7 @@ class ApiKeyAuthPlugin(AuthPlugin):
支持从x-api-key header或Authorization Bearer token中提取API Key
"""
def __init__(self):
def __init__(self) -> None:
super().__init__(name="api_key", priority=10)
def get_credentials(self, request: Request) -> str | None:

View File

@@ -3,6 +3,8 @@
定义认证插件的接口和认证上下文
"""
from __future__ import annotations
from abc import abstractmethod
from dataclasses import dataclass
from typing import Any
@@ -28,7 +30,7 @@ class AuthContext:
quota_info: dict[str, Any] = None
metadata: dict[str, Any] = None
def __post_init__(self):
def __post_init__(self) -> None:
if self.permissions is None:
self.permissions = {}
if self.metadata is None:
@@ -49,9 +51,9 @@ class AuthPlugin(BasePlugin):
author: str = "Unknown",
description: str = "",
api_version: str = "1.0",
dependencies: list[str] = None,
provides: list[str] = None,
config: dict[str, Any] = None,
dependencies: list[str] | None = None,
provides: list[str] | None = None,
config: dict[str, Any] | None = None,
):
"""
初始化认证插件

View File

@@ -3,6 +3,8 @@ JWT认证插件
支持JWT Bearer token认证
"""
from __future__ import annotations
import hashlib
from fastapi import Request
@@ -22,7 +24,7 @@ class JwtAuthPlugin(AuthPlugin):
支持从Authorization Bearer header中提取JWT token进行认证
"""
def __init__(self):
def __init__(self) -> None:
super().__init__(name="jwt", priority=20) # 高优先级优先于API Key
def get_credentials(self, request: Request) -> str | None:

View File

@@ -3,6 +3,8 @@
定义缓存插件的接口
"""
from __future__ import annotations
import hashlib
import json
from abc import abstractmethod
@@ -25,9 +27,9 @@ class CachePlugin(BasePlugin):
author: str = "Unknown",
description: str = "",
api_version: str = "1.0",
dependencies: list[str] = None,
provides: list[str] = None,
config: dict[str, Any] = None,
dependencies: list[str] | None = None,
provides: list[str] | None = None,
config: dict[str, Any] | None = None,
):
"""
初始化缓存插件
@@ -158,7 +160,7 @@ class CachePlugin(BasePlugin):
"""
pass
def generate_key(self, *args, **kwargs) -> str:
def generate_key(self, *args: Any, **kwargs: Any) -> str:
"""
生成缓存键
@@ -205,7 +207,7 @@ class CachePlugin(BasePlugin):
"""
return json.loads(value)
def configure(self, config: dict[str, Any]):
def configure(self, config: dict[str, Any]) -> Any:
"""
配置插件

View File

@@ -3,6 +3,8 @@
基于Python字典的简单内存缓存实现
"""
from __future__ import annotations
import asyncio
import threading
import time
@@ -18,7 +20,7 @@ class MemoryCachePlugin(CachePlugin):
使用OrderedDict实现LRU缓存
"""
def __init__(self, name: str = "memory", config: dict[str, Any] = None):
def __init__(self, name: str = "memory", config: dict[str, Any] | None = None):
super().__init__(name, config)
self._cache: OrderedDict = OrderedDict()
self._expiry: dict[str, float] = {}
@@ -38,10 +40,10 @@ class MemoryCachePlugin(CachePlugin):
except:
pass # 忽略事件循环错误
def _start_cleanup_task(self):
def _start_cleanup_task(self) -> None:
"""启动后台清理任务"""
async def cleanup_loop():
async def cleanup_loop() -> None:
while self.enabled:
await asyncio.sleep(self._cleanup_interval)
await self._cleanup_expired()
@@ -53,7 +55,7 @@ class MemoryCachePlugin(CachePlugin):
# 没有运行的事件循环,稍后再启动
pass
async def _cleanup_expired(self):
async def _cleanup_expired(self) -> None:
"""清理过期的缓存项"""
now = time.time()
expired_keys = []
@@ -68,7 +70,7 @@ class MemoryCachePlugin(CachePlugin):
self._expiry.pop(key, None)
self._evictions += 1
def _check_size(self):
def _check_size(self) -> None:
"""检查并维护缓存大小限制"""
if len(self._cache) >= self.max_size:
# 删除最老的项LRU
@@ -183,7 +185,7 @@ class MemoryCachePlugin(CachePlugin):
"cleanup_interval": self._cleanup_interval,
}
async def _do_shutdown(self):
async def _do_shutdown(self) -> None:
"""清理资源"""
# 取消清理任务
if self._cleanup_task:
@@ -193,7 +195,7 @@ class MemoryCachePlugin(CachePlugin):
except asyncio.CancelledError:
pass
def __del__(self):
def __del__(self) -> None:
"""清理资源"""
if hasattr(self, "_cleanup_task") and self._cleanup_task:
self._cleanup_task.cancel()

View File

@@ -3,6 +3,8 @@
包含所有插件类型共享的类和接口
"""
from __future__ import annotations
from abc import ABC
from dataclasses import dataclass
from enum import Enum
@@ -28,10 +30,10 @@ class PluginMetadata:
author: str = "Unknown"
description: str = ""
api_version: str = "1.0"
dependencies: list[str] = None
provides: list[str] = None
dependencies: list[str] | None = None
provides: list[str] | None = None
def __post_init__(self):
def __post_init__(self) -> None:
if self.dependencies is None:
self.dependencies = []
if self.provides is None:
@@ -52,9 +54,9 @@ class BasePlugin(ABC):
author: str = "Unknown",
description: str = "",
api_version: str = "1.0",
dependencies: list[str] = None,
provides: list[str] = None,
config: dict[str, Any] = None,
dependencies: list[str] | None = None,
provides: list[str] | None = None,
config: dict[str, Any] | None = None,
):
"""
初始化插件
@@ -103,13 +105,13 @@ class BasePlugin(ABC):
logger.warning(f"Failed to initialize plugin {self.name}: {e}")
return False
async def _do_initialize(self):
async def _do_initialize(self) -> None:
"""
子类可以重写此方法来实现特定的初始化逻辑
"""
pass
async def shutdown(self):
async def shutdown(self) -> Any:
"""
关闭插件,清理资源
"""
@@ -123,7 +125,7 @@ class BasePlugin(ABC):
finally:
self._initialized = False
async def _do_shutdown(self):
async def _do_shutdown(self) -> None:
"""
子类可以重写此方法来实现特定的清理逻辑
"""
@@ -153,7 +155,7 @@ class BasePlugin(ABC):
HealthStatus.HEALTHY if (self._initialized and self.enabled) else HealthStatus.UNHEALTHY
)
def configure(self, config: dict[str, Any]):
def configure(self, config: dict[str, Any]) -> Any:
"""
配置插件
@@ -198,5 +200,5 @@ class BasePlugin(ABC):
missing_deps.append(dep)
return missing_deps
def __repr__(self):
def __repr__(self) -> None:
return f"<{self.__class__.__name__}(name={self.name}, priority={self.priority}, enabled={self.enabled}, version={self.metadata.version})>"

View File

@@ -3,6 +3,8 @@
定义负载均衡策略的接口
"""
from __future__ import annotations
from abc import abstractmethod
from dataclasses import dataclass
from typing import Any
@@ -22,7 +24,7 @@ class ProviderCandidate:
model: Any | None = None # Model 对象(如果需要模型信息)
metadata: dict[str, Any] | None = None # 额外元数据
def __post_init__(self):
def __post_init__(self) -> None:
if self.metadata is None:
self.metadata = {}
@@ -38,7 +40,7 @@ class SelectionResult:
weight: float # 该提供商的权重
selection_metadata: dict[str, Any] | None = None # 选择过程的元数据
def __post_init__(self):
def __post_init__(self) -> None:
if self.selection_metadata is None:
self.selection_metadata = {}
@@ -57,9 +59,9 @@ class LoadBalancerStrategy(BasePlugin):
author: str = "Unknown",
description: str = "",
api_version: str = "1.0",
dependencies: list[str] = None,
provides: list[str] = None,
config: dict[str, Any] = None,
dependencies: list[str] | None = None,
provides: list[str] | None = None,
config: dict[str, Any] | None = None,
):
"""
初始化负载均衡策略
@@ -119,7 +121,7 @@ class LoadBalancerStrategy(BasePlugin):
success: bool,
response_time: float | None = None,
error: Exception | None = None,
):
) -> Any:
"""
记录请求结果(用于动态调整策略)

View File

@@ -19,6 +19,8 @@ WARNING: 多进程环境注意事项
参考src/services/health_monitor.py 中的实现
"""
from __future__ import annotations
import time
from collections import defaultdict
from typing import Any
@@ -49,7 +51,7 @@ class StickyPriorityStrategy(LoadBalancerStrategy):
详见模块文档说明。
"""
def __init__(self, config: dict[str, Any] = None):
def __init__(self, config: dict[str, Any] | None = None):
config = config or {} # 确保 config 不为 None
super().__init__(
name="sticky_priority",
@@ -324,7 +326,7 @@ class StickyPriorityStrategy(LoadBalancerStrategy):
# 选择权重最大的健康提供商
return max(healthy_candidates, key=lambda c: c.weight)
def _record_selection(self, provider: Any, is_sticky: bool = True):
def _record_selection(self, provider: Any, is_sticky: bool = True) -> None:
"""记录选择统计"""
self._stats["total_selections"] += 1
provider_id = str(provider.id)
@@ -342,7 +344,7 @@ class StickyPriorityStrategy(LoadBalancerStrategy):
success: bool,
response_time: float | None = None,
error: Exception | None = None,
):
) -> Any:
"""
记录请求结果,更新健康状态
@@ -422,7 +424,7 @@ class StickyPriorityStrategy(LoadBalancerStrategy):
},
}
async def reset_provider_health(self, provider_id: str):
async def reset_provider_health(self, provider_id: str) -> None:
"""重置指定提供商的健康状态"""
if provider_id in self._provider_health:
self._provider_health[provider_id] = {
@@ -434,7 +436,7 @@ class StickyPriorityStrategy(LoadBalancerStrategy):
}
logger.info(f"Reset health status for provider {provider_id}")
async def clear_sticky_cache(self, cache_key: str | None = None):
async def clear_sticky_cache(self, cache_key: str | None = None) -> None:
"""
清除粘性提供商缓存

View File

@@ -3,6 +3,8 @@
统一管理和协调所有插件系统
"""
from __future__ import annotations
import asyncio
import importlib
import inspect
@@ -82,7 +84,7 @@ class PluginManager:
# 应用配置
self._apply_config()
def _auto_discover_plugins(self):
def _auto_discover_plugins(self) -> None:
"""自动发现和加载插件"""
plugins_dir = Path(__file__).parent
@@ -125,7 +127,7 @@ class PluginManager:
# 解析失败,假设兼容
return True
def _load_plugin_from_module(self, module: Any, plugin_type: str):
def _load_plugin_from_module(self, module: Any, plugin_type: str) -> None:
"""从模块加载插件类"""
base_class = self.PLUGIN_TYPES[plugin_type]
@@ -149,7 +151,7 @@ class PluginManager:
except Exception as e:
logger.error(f"Failed to instantiate plugin {name}: {e}")
def _apply_config(self):
def _apply_config(self) -> None:
"""应用配置到插件"""
for plugin_type, plugins in self.plugins.items():
type_config = self.config.get(plugin_type, {})
@@ -164,7 +166,7 @@ class PluginManager:
if plugin_config:
plugin.configure(plugin_config)
def register_plugin(self, plugin_type: str, plugin: Any, set_as_default: bool = False):
def register_plugin(self, plugin_type: str, plugin: Any, set_as_default: bool = False) -> None:
"""
注册插件
@@ -192,7 +194,7 @@ class PluginManager:
logger.debug(f"Registered {plugin_type} plugin: {plugin.name}")
def unregister_plugin(self, plugin_type: str, plugin_name: str):
def unregister_plugin(self, plugin_type: str, plugin_name: str) -> Any:
"""
注销插件
@@ -267,7 +269,7 @@ class PluginManager:
return [p for p in plugins if getattr(p, "enabled", True)]
async def execute_plugin_chain(
self, plugin_type: str, method_name: str, *args, **kwargs
self, plugin_type: str, method_name: str, *args: Any, **kwargs: Any
) -> Any:
"""
执行插件链(按优先级)
@@ -385,7 +387,7 @@ class PluginManager:
return results
async def shutdown_all(self):
async def shutdown_all(self) -> None:
"""
关闭所有插件
"""
@@ -572,7 +574,7 @@ def get_plugin_manager(config: dict[str, Any] | None = None) -> PluginManager:
return _plugin_manager
def reset_plugin_manager():
def reset_plugin_manager() -> None:
"""重置插件管理器(用于测试)"""
global _plugin_manager
with _plugin_manager_lock:

View File

@@ -3,6 +3,8 @@
定义监控和指标收集的接口
"""
from __future__ import annotations
from abc import abstractmethod
from datetime import datetime, timezone
from enum import Enum
@@ -46,7 +48,7 @@ class MonitorPlugin(BasePlugin):
所有监控插件必须继承此类并实现相关方法
"""
def __init__(self, name: str, config: dict[str, Any] = None):
def __init__(self, name: str, config: dict[str, Any] | None = None):
"""
初始化监控插件
@@ -61,7 +63,7 @@ class MonitorPlugin(BasePlugin):
self.batch_size = self.config.get("batch_size", 100)
@abstractmethod
async def record_metric(self, metric: Metric):
async def record_metric(self, metric: Metric) -> None:
"""
记录单个指标
@@ -71,7 +73,7 @@ class MonitorPlugin(BasePlugin):
pass
@abstractmethod
async def record_batch(self, metrics: list[Metric]):
async def record_batch(self, metrics: list[Metric]) -> None:
"""
批量记录指标
@@ -81,7 +83,7 @@ class MonitorPlugin(BasePlugin):
pass
@abstractmethod
async def increment(self, name: str, value: float = 1, labels: dict[str, str] | None = None):
async def increment(self, name: str, value: float = 1, labels: dict[str, str] | None = None) -> Any:
"""
增加计数器
@@ -93,7 +95,7 @@ class MonitorPlugin(BasePlugin):
pass
@abstractmethod
async def gauge(self, name: str, value: float, labels: dict[str, str] | None = None):
async def gauge(self, name: str, value: float, labels: dict[str, str] | None = None) -> Any:
"""
设置仪表值
@@ -111,7 +113,7 @@ class MonitorPlugin(BasePlugin):
value: float,
labels: dict[str, str] | None = None,
buckets: list[float] | None = None,
):
) -> Any:
"""
记录直方图数据
@@ -124,7 +126,7 @@ class MonitorPlugin(BasePlugin):
pass
@abstractmethod
async def timing(self, name: str, duration: float, labels: dict[str, str] | None = None):
async def timing(self, name: str, duration: float, labels: dict[str, str] | None = None) -> Any:
"""
记录时间指标
@@ -136,7 +138,7 @@ class MonitorPlugin(BasePlugin):
pass
@abstractmethod
async def flush(self):
async def flush(self) -> Any:
"""
刷新缓冲的指标到后端
"""
@@ -160,7 +162,7 @@ class MonitorPlugin(BasePlugin):
duration: float,
provider: str | None = None,
model: str | None = None,
):
) -> Any:
"""
记录API请求指标便捷方法
@@ -209,7 +211,7 @@ class MonitorPlugin(BasePlugin):
input_tokens: int,
output_tokens: int,
cost: float | None = None,
):
) -> Any:
"""
记录Token使用指标便捷方法
@@ -240,7 +242,7 @@ class MonitorPlugin(BasePlugin):
if cost is not None:
loop.create_task(self.increment("usage_cost_total", cost, labels=labels))
def configure(self, config: dict[str, Any]):
def configure(self, config: dict[str, Any]) -> Any:
"""
配置插件
@@ -252,5 +254,5 @@ class MonitorPlugin(BasePlugin):
self.flush_interval = config.get("flush_interval", self.flush_interval)
self.batch_size = config.get("batch_size", self.batch_size)
def __repr__(self):
def __repr__(self) -> None:
return f"<{self.__class__.__name__}(name={self.name}, enabled={self.enabled})>"

View File

@@ -3,6 +3,8 @@ Prometheus监控插件
支持将指标导出到Prometheus
"""
from __future__ import annotations
import asyncio
from typing import Any
@@ -26,7 +28,7 @@ class PrometheusPlugin(MonitorPlugin):
使用prometheus_client库导出指标
"""
def __init__(self, name: str = "prometheus", config: dict[str, Any] = None):
def __init__(self, name: str = "prometheus", config: dict[str, Any] | None = None):
super().__init__(name, config)
# Check if prometheus_client is available
@@ -47,7 +49,7 @@ class PrometheusPlugin(MonitorPlugin):
# 启动刷新任务
self._start_flush_task()
def _init_default_metrics(self):
def _init_default_metrics(self) -> None:
"""初始化默认指标"""
# HTTP请求指标
http_label_names = ["method", "endpoint", "status", "status_class"]
@@ -113,10 +115,10 @@ class PrometheusPlugin(MonitorPlugin):
buckets=(0.1, 0.25, 0.5, 1, 2.5, 5, 10, 30),
)
def _start_flush_task(self):
def _start_flush_task(self) -> None:
"""启动定期刷新任务"""
async def flush_loop():
async def flush_loop() -> Any:
try:
while self.enabled:
await asyncio.sleep(self.flush_interval)
@@ -136,7 +138,7 @@ class PrometheusPlugin(MonitorPlugin):
# 如果没有运行的事件循环,任务将在后续创建
logger.warning("No event loop available for Prometheus flush task")
def _get_or_create_metric(self, name: str, metric_type: MetricType, labels: list[str] = None):
def _get_or_create_metric(self, name: str, metric_type: MetricType, labels: list[str] | None = None) -> Any:
"""获取或创建指标"""
if name not in self._metrics:
labels = labels or []
@@ -151,7 +153,7 @@ class PrometheusPlugin(MonitorPlugin):
return self._metrics[name]
async def record_metric(self, metric: Metric):
async def record_metric(self, metric: Metric) -> None:
"""记录单个指标"""
async with self._lock:
self._buffer.append(metric)
@@ -160,7 +162,7 @@ class PrometheusPlugin(MonitorPlugin):
if len(self._buffer) >= self.batch_size:
await self.flush()
async def record_batch(self, metrics: list[Metric]):
async def record_batch(self, metrics: list[Metric]) -> None:
"""批量记录指标"""
async with self._lock:
self._buffer.extend(metrics)
@@ -169,7 +171,7 @@ class PrometheusPlugin(MonitorPlugin):
if len(self._buffer) >= self.batch_size:
await self.flush()
async def increment(self, name: str, value: float = 1, labels: dict[str, str] | None = None):
async def increment(self, name: str, value: float = 1, labels: dict[str, str] | None = None) -> Any:
"""增加计数器"""
try:
if name in self._metrics:
@@ -192,7 +194,7 @@ class PrometheusPlugin(MonitorPlugin):
# 记录错误但不中断
logger.warning(f"Error recording metric {name}: {e}")
async def gauge(self, name: str, value: float, labels: dict[str, str] | None = None):
async def gauge(self, name: str, value: float, labels: dict[str, str] | None = None) -> Any:
"""设置仪表值"""
try:
if name in self._metrics:
@@ -219,7 +221,7 @@ class PrometheusPlugin(MonitorPlugin):
value: float,
labels: dict[str, str] | None = None,
buckets: list[float] | None = None,
):
) -> Any:
"""记录直方图数据"""
try:
if name in self._metrics:
@@ -247,12 +249,12 @@ class PrometheusPlugin(MonitorPlugin):
except Exception as e:
logger.warning(f"Error recording histogram {name}: {e}")
async def timing(self, name: str, duration: float, labels: dict[str, str] | None = None):
async def timing(self, name: str, duration: float, labels: dict[str, str] | None = None) -> Any:
"""记录时间指标"""
# 使用直方图记录时间
await self.histogram(f"{name}_seconds", duration, labels)
async def flush(self):
async def flush(self) -> Any:
"""刷新缓冲的指标到Prometheus"""
async with self._lock:
if not self._buffer:
@@ -289,7 +291,7 @@ class PrometheusPlugin(MonitorPlugin):
"""
return generate_latest(REGISTRY)
async def shutdown(self):
async def shutdown(self) -> Any:
"""
关闭插件,取消后台任务
@@ -311,7 +313,7 @@ class PrometheusPlugin(MonitorPlugin):
logger.info("Prometheus plugin shutdown complete")
async def cleanup(self):
async def cleanup(self) -> Any:
"""
清理资源(别名方法)
"""

View File

@@ -3,6 +3,8 @@
定义通知的接口和数据结构
"""
from __future__ import annotations
import asyncio
import json
from abc import abstractmethod
@@ -90,7 +92,7 @@ class NotificationPlugin(BasePlugin):
提供统一的重试机制,子类只需实现 _do_send 和 _do_send_batch 方法
"""
def __init__(self, name: str = "notification", config: dict[str, Any] = None):
def __init__(self, name: str = "notification", config: dict[str, Any] | None = None):
# 调用父类初始化设置metadata
super().__init__(
name=name, config=config, description="Notification Plugin", version="1.0.0"

View File

@@ -34,7 +34,7 @@ class WebhookNotificationPlugin(NotificationPlugin):
支持多种Webhook格式Slack, Discord, 通用)
"""
def __init__(self, name: str = "webhook", config: dict[str, Any] = None):
def __init__(self, name: str = "webhook", config: dict[str, Any] | None = None):
super().__init__(name, config)
if not AIOHTTP_AVAILABLE:
@@ -65,10 +65,10 @@ class WebhookNotificationPlugin(NotificationPlugin):
# 启动刷新任务
self._start_flush_task()
def _start_flush_task(self):
def _start_flush_task(self) -> None:
"""启动定时刷新任务"""
async def flush_loop():
async def flush_loop() -> Any:
while self.enabled:
await asyncio.sleep(self.flush_interval)
await self.flush()
@@ -288,11 +288,11 @@ class WebhookNotificationPlugin(NotificationPlugin):
"has_secret": bool(self.secret),
}
async def _do_shutdown(self):
async def _do_shutdown(self) -> None:
"""清理资源"""
await self.close()
async def close(self):
async def close(self) -> Any:
"""关闭插件"""
# 刷新缓冲
await self.flush()
@@ -305,7 +305,7 @@ class WebhookNotificationPlugin(NotificationPlugin):
if self._session:
await self._session.close()
def __del__(self):
def __del__(self) -> None:
"""清理资源"""
try:
asyncio.create_task(self.close())

View File

@@ -3,6 +3,8 @@
定义速率限制策略的接口
"""
from __future__ import annotations
from abc import abstractmethod
from dataclasses import dataclass
from datetime import datetime
@@ -24,7 +26,7 @@ class RateLimitResult:
message: str | None = None
headers: dict[str, str] | None = None
def __post_init__(self):
def __post_init__(self) -> None:
if self.headers is None:
self.headers = {}
if self.remaining is not None:
@@ -49,9 +51,9 @@ class RateLimitStrategy(BasePlugin):
author: str = "Unknown",
description: str = "",
api_version: str = "1.0",
dependencies: list[str] = None,
provides: list[str] = None,
config: dict[str, Any] = None,
dependencies: list[str] | None = None,
provides: list[str] | None = None,
config: dict[str, Any] | None = None,
):
"""
初始化速率限制策略
@@ -80,7 +82,7 @@ class RateLimitStrategy(BasePlugin):
)
@abstractmethod
async def check_limit(self, key: str, **kwargs) -> RateLimitResult:
async def check_limit(self, key: str, **kwargs: Any) -> RateLimitResult:
"""
检查速率限制
@@ -94,7 +96,7 @@ class RateLimitStrategy(BasePlugin):
pass
@abstractmethod
async def consume(self, key: str, amount: int = 1, **kwargs) -> bool:
async def consume(self, key: str, amount: int = 1, **kwargs: Any) -> bool:
"""
消费配额
@@ -109,7 +111,7 @@ class RateLimitStrategy(BasePlugin):
pass
@abstractmethod
async def reset(self, key: str):
async def reset(self, key: str) -> Any:
"""
重置限制

View File

@@ -45,7 +45,7 @@ class SlidingWindow:
self.requests: deque[float] = deque()
self.last_access_time: float = time.time()
def _cleanup(self):
def _cleanup(self) -> None:
"""清理过期的请求记录"""
current_time = time.time()
self.last_access_time = current_time # 更新最后访问时间
@@ -119,7 +119,7 @@ class SlidingWindowStrategy(RateLimitStrategy):
# 默认窗口过期时间(秒)- 超过此时间未访问的窗口将被清理
DEFAULT_WINDOW_EXPIRY = 3600 # 1小时
def __init__(self):
def __init__(self) -> None:
super().__init__("sliding_window")
self.windows: dict[str, SlidingWindow] = {}
self._lock = asyncio.Lock()
@@ -184,7 +184,7 @@ class SlidingWindowStrategy(RateLimitStrategy):
return evicted
async def _maybe_cleanup(self):
async def _maybe_cleanup(self) -> None:
"""检查是否需要执行清理操作"""
current_time = time.time()
@@ -225,7 +225,7 @@ class SlidingWindowStrategy(RateLimitStrategy):
return self.windows[key]
async def check_limit(self, key: str, **kwargs) -> RateLimitResult:
async def check_limit(self, key: str, **kwargs: Any) -> RateLimitResult:
"""
检查速率限制
@@ -264,7 +264,7 @@ class SlidingWindowStrategy(RateLimitStrategy):
),
)
async def consume(self, key: str, amount: int = 1, **kwargs) -> bool:
async def consume(self, key: str, amount: int = 1, **kwargs: Any) -> bool:
"""
消费配额
@@ -286,7 +286,7 @@ class SlidingWindowStrategy(RateLimitStrategy):
return success
async def reset(self, key: str):
async def reset(self, key: str) -> Any:
"""
重置滑动窗口
@@ -324,7 +324,7 @@ class SlidingWindowStrategy(RateLimitStrategy):
"reset_at": window.get_reset_time().isoformat(),
}
def configure(self, config: dict[str, Any]):
def configure(self, config: dict[str, Any]) -> Any:
"""
配置策略

View File

@@ -1,5 +1,7 @@
"""令牌桶速率限制策略,支持 Redis 分布式后端"""
from __future__ import annotations
import asyncio
import os
import time
@@ -28,7 +30,7 @@ class TokenBucket:
self.tokens = capacity
self.last_refill = time.time()
def _refill(self):
def _refill(self) -> None:
"""补充令牌"""
now = time.time()
time_passed = now - self.last_refill
@@ -80,7 +82,7 @@ class TokenBucketStrategy(RateLimitStrategy):
- 适合处理不均匀的流量模式
"""
def __init__(self):
def __init__(self) -> None:
super().__init__("token_bucket")
self.buckets: dict[str, TokenBucket] = {}
self._lock = asyncio.Lock()
@@ -129,7 +131,7 @@ class TokenBucketStrategy(RateLimitStrategy):
def _want_redis_backend(self) -> bool:
return self._backend_mode in {"auto", "redis"}
async def _ensure_backend(self):
async def _ensure_backend(self) -> None:
if self._redis_checked:
return
self._redis_checked = True
@@ -142,7 +144,7 @@ class TokenBucketStrategy(RateLimitStrategy):
elif self._backend_mode == "redis":
logger.warning("RATE_LIMIT_BACKEND=redis 但 Redis 客户端不可用,回退到内存桶")
async def check_limit(self, key: str, **kwargs) -> RateLimitResult:
async def check_limit(self, key: str, **kwargs: Any) -> RateLimitResult:
"""
检查速率限制
@@ -190,7 +192,7 @@ class TokenBucketStrategy(RateLimitStrategy):
),
)
async def consume(self, key: str, amount: int = 1, **kwargs) -> bool:
async def consume(self, key: str, amount: int = 1, **kwargs: Any) -> bool:
"""
消费令牌
@@ -227,7 +229,7 @@ class TokenBucketStrategy(RateLimitStrategy):
return success
async def reset(self, key: str):
async def reset(self, key: str) -> Any:
"""
重置令牌桶
@@ -278,7 +280,7 @@ class TokenBucketStrategy(RateLimitStrategy):
"reset_at": bucket.get_reset_time().isoformat(),
}
def configure(self, config: dict[str, Any]):
def configure(self, config: dict[str, Any]) -> Any:
"""
配置策略
@@ -349,7 +351,7 @@ class RedisTokenBucketBackend:
return {allowed, tokens, retry_after}
"""
def __init__(self, redis_client):
def __init__(self, redis_client: Any) -> None:
self.redis = redis_client
self._consume_script = self.redis.register_script(self._SCRIPT)
@@ -413,7 +415,7 @@ class RedisTokenBucketBackend:
remaining = int(float(result[1]))
return allowed, remaining
async def reset(self, key: str):
async def reset(self, key: str) -> Any:
await self.redis.delete(self._redis_key(key))
async def get_stats(self, key: str, capacity: int, refill_rate: float) -> dict[str, Any]:

View File

@@ -51,7 +51,7 @@ class TokenCounterPlugin(BasePlugin):
支持不同模型的Token计数
"""
def __init__(self, name: str = "token_counter", config: dict[str, Any] = None):
def __init__(self, name: str = "token_counter", config: dict[str, Any] | None = None):
# 调用父类初始化设置metadata
super().__init__(
name=name, config=config, description="Token Counter Plugin", version="1.0.0"

View File

@@ -3,6 +3,8 @@ Claude Token计数插件
专门为Claude模型设计的Token计数器
"""
from __future__ import annotations
import json
import re
from typing import Any
@@ -61,7 +63,7 @@ class ClaudeTokenCounterPlugin(TokenCounterPlugin):
},
}
def __init__(self, name: str = "claude", config: dict[str, Any] = None):
def __init__(self, name: str = "claude", config: dict[str, Any] | None = None):
super().__init__(name, config)
# 价格表每1M tokens的价格 USD

View File

@@ -3,6 +3,8 @@ Tiktoken Token计数插件
支持OpenAI和其他使用tiktoken的模型
"""
from __future__ import annotations
from typing import Any
from src.core.logger import logger
@@ -56,7 +58,7 @@ class TiktokenCounterPlugin(TokenCounterPlugin):
"gpt-4o-mini": 3,
}
def __init__(self, name: str = "tiktoken", config: dict[str, Any] = None):
def __init__(self, name: str = "tiktoken", config: dict[str, Any] | None = None):
super().__init__(name, config)
if not TIKTOKEN_AVAILABLE:

View File

@@ -1,3 +1,5 @@
from __future__ import annotations
from dataclasses import dataclass
from src.core.logger import logger

View File

@@ -18,6 +18,8 @@
- 这样可以支持"独立余额Key"场景每个Key有自己的缓存亲和性
"""
from __future__ import annotations
import asyncio
import json
import os
@@ -76,7 +78,7 @@ class CacheAffinityManager:
# 默认缓存TTL- 使用统一常量
DEFAULT_CACHE_TTL = CacheTTL.CACHE_AFFINITY
def __init__(self, redis_client=None, default_ttl: int = DEFAULT_CACHE_TTL):
def __init__(self, redis_client: Any | None = None, default_ttl: int = DEFAULT_CACHE_TTL) -> None:
"""
初始化缓存亲和性管理器
@@ -149,7 +151,7 @@ class CacheAffinityManager:
return None
return dict(payload)
async def _set_l1_entry(self, cache_key: str, payload: dict[str, Any] | None):
async def _set_l1_entry(self, cache_key: str, payload: dict[str, Any] | None) -> None:
async with self._l1_lock:
if not payload:
self._l1_cache.pop(cache_key, None)
@@ -194,7 +196,7 @@ class CacheAffinityManager:
return len(expired_keys)
@asynccontextmanager
async def _acquire_request_lock(self, cache_key: str):
async def _acquire_request_lock(self, cache_key: str) -> None:
lock = self._request_locks.get(cache_key)
if lock is None:
lock = asyncio.Lock()
@@ -647,7 +649,7 @@ class CacheAffinityManager:
_affinity_manager: CacheAffinityManager | None = None
async def get_affinity_manager(redis_client=None) -> CacheAffinityManager:
async def get_affinity_manager(redis_client: Any | None = None) -> CacheAffinityManager:
"""
获取全局CacheAffinityManager实例若Redis不可用则降级为内存模式

View File

@@ -133,10 +133,10 @@ class CacheAwareScheduler:
def __init__(
self,
redis_client=None,
redis_client: Any | None = None,
priority_mode: str | None = None,
scheduling_mode: str | None = None,
):
) -> None:
"""
初始化调度器
@@ -182,7 +182,7 @@ class CacheAwareScheduler:
"last_reservation_result": None,
}
async def _ensure_initialized(self):
async def _ensure_initialized(self) -> None:
"""确保所有异步组件已初始化"""
if self._affinity_manager is None:
self._affinity_manager = await get_affinity_manager(self.redis)
@@ -512,7 +512,7 @@ class CacheAwareScheduler:
f"User.allowed_models={user.allowed_models if user else 'N/A'}"
)
def merge_restrictions(key_restriction, user_restriction):
def merge_restrictions(key_restriction: Any, user_restriction: Any) -> Any:
"""合并两个限制列表,返回有效的限制集合"""
key_set = set(key_restriction) if key_restriction else None
user_set = set(user_restriction) if user_restriction else None
@@ -1405,7 +1405,7 @@ class CacheAwareScheduler:
result.extend(sorted_group)
else:
# 单个候选或没有 affinity_key按次要排序条件排序
def secondary_sort(c: ProviderCandidate):
def secondary_sort(c: ProviderCandidate) -> Any:
return (
c.provider.provider_priority,
c.key.internal_priority if c.key else 999999,
@@ -1541,7 +1541,7 @@ class CacheAwareScheduler:
endpoint_id: str | None = None,
key_id: str | None = None,
provider_id: str | None = None,
):
) -> Any:
"""
失效指定亲和性标识符对特定API格式和模型的缓存亲和性
@@ -1572,7 +1572,7 @@ class CacheAwareScheduler:
api_format: str,
global_model_id: str,
ttl: int | None = None,
):
) -> Any:
"""
记录缓存亲和性(供编排器调用)
@@ -1645,7 +1645,7 @@ _scheduler: CacheAwareScheduler | None = None
async def get_cache_aware_scheduler(
redis_client=None,
redis_client: Any | None = None,
priority_mode: str | None = None,
scheduling_mode: str | None = None,
) -> CacheAwareScheduler:

View File

@@ -10,6 +10,8 @@
- 其他需要缓存的服务
"""
from __future__ import annotations
import asyncio
import json
import time
@@ -87,7 +89,7 @@ class LocalCache(BaseCacheBackend):
self._cache.move_to_end(key)
return self._cache[key]
async def set(self, key: str, value: Any, ttl: int = None) -> None:
async def set(self, key: str, value: Any, ttl: int | None = None) -> None:
"""设置缓存值(线程安全)"""
async with self._lock:
if ttl is None:
@@ -195,7 +197,7 @@ class RedisCache(BaseCacheBackend):
logger.error(f"[RedisCache] 获取缓存失败: {key}, 错误: {e}")
return None
async def set(self, key: str, value: Any, ttl: int = None) -> None:
async def set(self, key: str, value: Any, ttl: int | None = None) -> None:
"""设置缓存值"""
if ttl is None:
ttl = self._default_ttl

View File

@@ -5,16 +5,19 @@
"""
from __future__ import annotations
from typing import Any
from src.core.logger import logger
class CacheInvalidationService:
"""缓存失效服务"""
def __init__(self):
def __init__(self) -> None:
self._model_mappers = []
def register_model_mapper(self, model_mapper):
def register_model_mapper(self, model_mapper: Any) -> None:
"""注册 ModelMapper 实例"""
if model_mapper not in self._model_mappers:
self._model_mappers.append(model_mapper)
@@ -58,7 +61,7 @@ class CacheInvalidationService:
except Exception as e:
logger.error(f"[CacheInvalidation] 失效 models list 缓存失败: {e}")
def on_model_changed(self, provider_id: str, global_model_id: str):
def on_model_changed(self, provider_id: str, global_model_id: str) -> Any:
"""Model 变更时的缓存失效"""
self._refresh_provider_cache(provider_id)
@@ -88,7 +91,7 @@ class CacheInvalidationService:
for mapper in self._model_mappers:
mapper.refresh_cache(provider_id)
def clear_all_caches(self):
def clear_all_caches(self) -> None:
"""清空所有缓存"""
for mapper in self._model_mappers:
mapper.clear_cache()

View File

@@ -6,6 +6,8 @@ Provider 缓存服务 - 减少 Provider 和 ProviderAPIKey 查询
"""
from __future__ import annotations
from sqlalchemy.orm import Session
from src.config.constants import CacheTTL

View File

@@ -9,6 +9,9 @@
2. GlobalModel/Model 变更时,同步失效所有实例的缓存
"""
from __future__ import annotations
from typing import Any
import asyncio
import json
@@ -46,7 +49,7 @@ class CacheSyncService:
self._handlers: dict[str, Callable] = {}
self._running = False
async def start(self):
async def start(self) -> Any:
"""启动缓存同步服务(订阅 Redis 频道)"""
if self._running:
logger.warning("[CacheSync] 服务已在运行")
@@ -73,7 +76,7 @@ class CacheSyncService:
logger.error(f"[CacheSync] 启动失败: {e}")
raise
async def stop(self):
async def stop(self) -> Any:
"""停止缓存同步服务"""
if not self._running:
return
@@ -95,7 +98,7 @@ class CacheSyncService:
logger.info("[CacheSync] 缓存同步服务已停止")
def register_handler(self, channel: str, handler: Callable):
def register_handler(self, channel: str, handler: Callable) -> None:
"""
注册缓存失效处理器
@@ -106,7 +109,7 @@ class CacheSyncService:
self._handlers[channel] = handler
logger.debug(f"[CacheSync] 注册处理器: {channel}")
async def _listen(self):
async def _listen(self) -> None:
"""监听 Redis pub/sub 消息"""
logger.info("[CacheSync] 开始监听缓存失效消息")
@@ -136,21 +139,21 @@ class CacheSyncService:
except Exception as e:
logger.error(f"[CacheSync] 监听失败: {e}")
async def publish_global_model_changed(self, model_name: str):
async def publish_global_model_changed(self, model_name: str) -> Any:
"""发布 GlobalModel 变更通知"""
await self._publish(self.CHANNEL_GLOBAL_MODEL, {"model_name": model_name})
async def publish_model_changed(self, provider_id: str, global_model_id: str):
async def publish_model_changed(self, provider_id: str, global_model_id: str) -> Any:
"""发布 Model 变更通知"""
await self._publish(
self.CHANNEL_MODEL, {"provider_id": provider_id, "global_model_id": global_model_id}
)
async def publish_clear_all(self):
async def publish_clear_all(self) -> Any:
"""发布清空所有缓存通知"""
await self._publish(self.CHANNEL_CLEAR_ALL, {})
async def _publish(self, channel: str, data: dict):
async def _publish(self, channel: str, data: dict) -> None:
"""发布消息到 Redis 频道"""
try:
message = json.dumps(data)
@@ -164,7 +167,7 @@ class CacheSyncService:
_cache_sync_service: CacheSyncService | None = None
async def get_cache_sync_service(redis_client: aioredis.Redis = None) -> CacheSyncService | None:
async def get_cache_sync_service(redis_client: aioredis.Redis | None = None) -> CacheSyncService | None:
"""
获取缓存同步服务实例
@@ -191,7 +194,7 @@ async def get_cache_sync_service(redis_client: aioredis.Redis = None) -> CacheSy
return _cache_sync_service
async def close_cache_sync_service():
async def close_cache_sync_service() -> None:
"""关闭缓存同步服务"""
global _cache_sync_service

View File

@@ -20,6 +20,9 @@
"""
from __future__ import annotations
from typing import Any
from sqlalchemy.orm import Session
from src.config.constants import CacheTTL
@@ -102,7 +105,7 @@ class UserCacheService:
return user
@staticmethod
async def invalidate_user_cache(user_id: str, email: str | None = None):
async def invalidate_user_cache(user_id: str, email: str | None = None) -> Any:
"""
清除用户缓存

View File

@@ -3,6 +3,8 @@
提供验证码邮件的 HTML 和纯文本模板,支持从数据库加载自定义模板
"""
from __future__ import annotations
import html
import re
from html.parser import HTMLParser
@@ -16,12 +18,12 @@ from src.services.system.config import SystemConfigService
class HTMLToTextParser(HTMLParser):
"""HTML 转纯文本解析器"""
def __init__(self):
def __init__(self) -> None:
super().__init__()
self.text_parts = []
self.skip_data = False
def handle_starttag(self, tag, attrs): # noqa: ARG002
def handle_starttag(self, tag: Any, attrs: Any) -> None: # noqa: ARG002
if tag in ("script", "style", "head"):
self.skip_data = True
elif tag == "br":
@@ -29,13 +31,13 @@ class HTMLToTextParser(HTMLParser):
elif tag in ("p", "div", "tr", "h1", "h2", "h3", "h4", "h5", "h6"):
self.text_parts.append("\n")
def handle_endtag(self, tag):
def handle_endtag(self, tag: Any) -> None:
if tag in ("script", "style", "head"):
self.skip_data = False
elif tag in ("p", "div", "tr", "h1", "h2", "h3", "h4", "h5", "h6", "td"):
self.text_parts.append("\n")
def handle_data(self, data):
def handle_data(self, data: Any) -> None:
if not self.skip_data:
text = data.strip()
if text:
@@ -310,7 +312,7 @@ class EmailTemplate:
@staticmethod
def get_verification_code_html(
code: str, expire_minutes: int = 5, db: Session | None = None, **kwargs
code: str, expire_minutes: int = 5, db: Session | None = None, **kwargs: Any
) -> str:
"""
获取验证码邮件 HTML
@@ -345,7 +347,7 @@ class EmailTemplate:
@staticmethod
def get_verification_code_text(
code: str, expire_minutes: int = 5, db: Session | None = None, **kwargs
code: str, expire_minutes: int = 5, db: Session | None = None, **kwargs: Any
) -> str:
"""
获取验证码邮件纯文本(从 HTML 自动生成)
@@ -364,7 +366,7 @@ class EmailTemplate:
@staticmethod
def get_password_reset_html(
reset_link: str, expire_minutes: int = 30, db: Session | None = None, **kwargs
reset_link: str, expire_minutes: int = 30, db: Session | None = None, **kwargs: Any
) -> str:
"""
获取密码重置邮件 HTML
@@ -399,7 +401,7 @@ class EmailTemplate:
@staticmethod
def get_password_reset_text(
reset_link: str, expire_minutes: int = 30, db: Session | None = None, **kwargs
reset_link: str, expire_minutes: int = 30, db: Session | None = None, **kwargs: Any
) -> str:
"""
获取密码重置邮件纯文本(从 HTML 自动生成)

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