mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
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:
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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 计数已重置"}
|
||||
|
||||
@@ -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 字段)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
# 检查模块是否存在
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -4,6 +4,8 @@
|
||||
提供缓存亲和性统计、管理和监控功能
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
"""
|
||||
测试模型连接性
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)}
|
||||
|
||||
@@ -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}' 不存在")
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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存在且属于该用户
|
||||
|
||||
@@ -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": "公告已删除"}
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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="未登录")
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -9,6 +9,8 @@
|
||||
5. 流式平滑输出
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import codecs
|
||||
import json
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -4,6 +4,8 @@ Claude CLI Adapter - 基于通用 CLI Adapter 基类的简化实现
|
||||
继承 CliAdapterBase,只需配置 FORMAT_ID 和 HANDLER_CLASS。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
@@ -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] = {}
|
||||
|
||||
|
||||
@@ -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 对象
|
||||
|
||||
|
||||
@@ -4,6 +4,8 @@ Gemini CLI Adapter - 基于通用 CLI Adapter 基类的实现
|
||||
继承 CliAdapterBase,处理 Gemini CLI 格式的请求。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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 对象
|
||||
|
||||
|
||||
@@ -4,6 +4,8 @@ OpenAI CLI Adapter - 基于通用 CLI Adapter 基类的简化实现
|
||||
继承 CliAdapterBase,只需配置 FORMAT_ID 和 HANDLER_CLASS。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -1,6 +1,8 @@
|
||||
"""OAuth 管理端点(管理员)。"""
|
||||
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, Request
|
||||
|
||||
@@ -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:
|
||||
"""
|
||||
获取指定模型支持的能力列表
|
||||
|
||||
|
||||
@@ -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 列表")
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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 兼容)
|
||||
|
||||
|
||||
@@ -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:
|
||||
"""
|
||||
删除指定文件
|
||||
|
||||
|
||||
@@ -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:
|
||||
"""
|
||||
获取认证模块状态(公开接口)
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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 等异步认证场景)
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -2,6 +2,8 @@
|
||||
缓存服务 - 统一的缓存抽象层
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any
|
||||
|
||||
|
||||
@@ -4,6 +4,8 @@
|
||||
|
||||
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
def extract_error_message(error: Exception, status_code: int | None = None) -> str:
|
||||
"""
|
||||
从异常中提取错误消息,优先使用上游原始响应(用于链路追踪/调试)
|
||||
|
||||
@@ -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异常处理器
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
# ============================================================================
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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)
|
||||
|
||||
26
src/main.py
26
src/main.py
@@ -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"
|
||||
# 日志系统已在导入时自动初始化
|
||||
|
||||
@@ -4,6 +4,8 @@
|
||||
提供完整的输入验证和安全过滤
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
):
|
||||
"""
|
||||
初始化认证插件
|
||||
|
||||
@@ -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:
|
||||
|
||||
12
src/plugins/cache/base.py
vendored
12
src/plugins/cache/base.py
vendored
@@ -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:
|
||||
"""
|
||||
配置插件
|
||||
|
||||
|
||||
16
src/plugins/cache/memory.py
vendored
16
src/plugins/cache/memory.py
vendored
@@ -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()
|
||||
|
||||
@@ -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})>"
|
||||
|
||||
@@ -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:
|
||||
"""
|
||||
记录请求结果(用于动态调整策略)
|
||||
|
||||
|
||||
@@ -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:
|
||||
"""
|
||||
清除粘性提供商缓存
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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})>"
|
||||
|
||||
@@ -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:
|
||||
"""
|
||||
清理资源(别名方法)
|
||||
"""
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -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:
|
||||
"""
|
||||
重置限制
|
||||
|
||||
|
||||
@@ -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:
|
||||
"""
|
||||
配置策略
|
||||
|
||||
|
||||
@@ -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]:
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
from src.core.logger import logger
|
||||
|
||||
10
src/services/cache/affinity_manager.py
vendored
10
src/services/cache/affinity_manager.py
vendored
@@ -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不可用则降级为内存模式)
|
||||
|
||||
|
||||
16
src/services/cache/aware_scheduler.py
vendored
16
src/services/cache/aware_scheduler.py
vendored
@@ -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:
|
||||
|
||||
6
src/services/cache/backend.py
vendored
6
src/services/cache/backend.py
vendored
@@ -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
|
||||
|
||||
11
src/services/cache/invalidation.py
vendored
11
src/services/cache/invalidation.py
vendored
@@ -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()
|
||||
|
||||
2
src/services/cache/provider_cache.py
vendored
2
src/services/cache/provider_cache.py
vendored
@@ -6,6 +6,8 @@ Provider 缓存服务 - 减少 Provider 和 ProviderAPIKey 查询
|
||||
"""
|
||||
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.config.constants import CacheTTL
|
||||
|
||||
23
src/services/cache/sync.py
vendored
23
src/services/cache/sync.py
vendored
@@ -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
|
||||
|
||||
|
||||
5
src/services/cache/user_cache.py
vendored
5
src/services/cache/user_cache.py
vendored
@@ -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:
|
||||
"""
|
||||
清除用户缓存
|
||||
|
||||
|
||||
@@ -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
Reference in New Issue
Block a user