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

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

@@ -8,6 +8,8 @@
4. Redis 缓存优化
"""
from __future__ import annotations
import json
from collections import defaultdict
from datetime import datetime, timedelta, timezone
@@ -25,7 +27,7 @@ CACHE_TTL_SECONDS = 30 # 缓存 30 秒
CACHE_KEY_PREFIX = "health:endpoint:"
def _get_redis_client():
def _get_redis_client() -> Any:
"""获取 Redis 客户端,失败返回 None"""
try:
from src.clients.redis_client import redis_client

View File

@@ -12,6 +12,8 @@
- circuit_breaker_by_format: {"CLAUDE": {"open": false, "open_at": null, ...}, ...}
"""
from __future__ import annotations
import os
from datetime import datetime, timedelta, timezone
from typing import Any

View File

@@ -1,5 +1,8 @@
"""Management Token 服务"""
from __future__ import annotations
from typing import Any
import ipaddress
from datetime import datetime, timezone
@@ -44,7 +47,7 @@ def validate_ip_list(ips: list[str] | None) -> list[str] | None:
return validated
def parse_expires_at(v, allow_past: bool = False) -> datetime | None:
def parse_expires_at(v: Any, allow_past: bool = False) -> datetime | None:
"""解析过期时间,确保时区安全
前端 datetime-local 输入返回本地时间字符串(无时区信息)。

View File

@@ -858,7 +858,7 @@ class ModelCostService:
)
@classmethod
def clear_cache(cls):
def clear_cache(cls) -> None:
"""清理价格相关缓存。"""
cls._price_cache.clear()
cls._cache_price_cache.clear()

View File

@@ -4,6 +4,8 @@ GlobalModel 服务层
提供 GlobalModel 的 CRUD 操作、查询和统计功能
"""
from __future__ import annotations
from typing import cast
from sqlalchemy.orm import Session, joinedload
@@ -122,7 +124,7 @@ class GlobalModelService:
# 按次计费配置
default_price_per_request: float | None = None,
# 阶梯计费配置(必填)
default_tiered_pricing: dict = None,
default_tiered_pricing: dict | None = None,
# Key 能力配置
supported_capabilities: list[str] | None = None,
# 模型配置JSON

View File

@@ -3,6 +3,8 @@
根据数据库中的配置,将用户请求的模型映射到提供商的实际模型
"""
from __future__ import annotations
from sqlalchemy.orm import Session, joinedload
from src.core.cache_utils import SyncLRUCache
@@ -212,12 +214,12 @@ class ModelMapperMiddleware:
return True, None
def clear_cache(self):
def clear_cache(self) -> None:
"""清空缓存"""
self._cache.clear()
logger.debug("Model mapping cache cleared")
def refresh_cache(self, provider_id: str | None = None):
def refresh_cache(self, provider_id: str | None = None) -> None:
"""
刷新缓存

View File

@@ -2,6 +2,8 @@
模型管理服务
"""
from __future__ import annotations
import asyncio
from sqlalchemy import and_
@@ -234,7 +236,7 @@ class ModelService:
raise InvalidRequestException("更新模型失败,请检查输入数据")
@staticmethod
def delete_model(db: Session, model_id: str): # UUID
def delete_model(db: Session, model_id: str) -> None: # UUID
"""删除模型
删除逻辑:

View File

@@ -4,6 +4,8 @@
负责错误分类和处理策略决定
"""
from __future__ import annotations
import json
from enum import Enum
from typing import Any
@@ -100,7 +102,7 @@ class ErrorClassifier:
def __init__(
self,
db: Session,
adaptive_manager: Any = None,
adaptive_manager: Any | None = None,
cache_scheduler: CacheAwareScheduler | None = None,
) -> None:
"""

View File

@@ -22,6 +22,8 @@
"""
from __future__ import annotations
from typing import Any, NoReturn
from collections.abc import Callable

View File

@@ -4,6 +4,7 @@
"""
from typing import Any
from sqlalchemy.orm import Session
from src.models.database import GlobalModel, Model, Provider
@@ -26,7 +27,7 @@ class ProviderService:
self.router = ModelRoutingMiddleware(db)
self.cost_service = ModelCostService(db)
async def _check_model_availability(self, model_name: str):
async def _check_model_availability(self, model_name: str) -> None:
"""
检查模型是否可用(严格白名单模式)
@@ -55,7 +56,7 @@ class ProviderService:
return None
async def _check_provider_model_availability(self, provider_id: str, model_name: str):
async def _check_provider_model_availability(self, provider_id: str, model_name: str) -> None:
"""
检查特定提供商是否支持特定模型
@@ -115,7 +116,7 @@ class ProviderService:
"""
return self.router.get_available_models()
def select_provider(self, model_name: str, preferred_provider=None):
def select_provider(self, model_name: str, preferred_provider: Any | None = None) -> Any:
"""
选择提供商

View File

@@ -4,6 +4,8 @@ Provider 操作服务
提供操作执行、凭据管理、缓存等业务逻辑。
"""
from __future__ import annotations
import asyncio
import os
from dataclasses import asdict

View File

@@ -335,7 +335,7 @@ def get_adaptive_reservation_manager() -> AdaptiveReservationManager:
return _reservation_manager
def reset_adaptive_reservation_manager():
def reset_adaptive_reservation_manager() -> None:
"""重置全局单例(用于测试)"""
global _reservation_manager
_reservation_manager = None

View File

@@ -13,6 +13,8 @@
3. 安全探测:允许在稳定后尝试更高 RPM
"""
from __future__ import annotations
from datetime import datetime, timezone
from typing import Any, cast

View File

@@ -10,6 +10,7 @@ RPM 限制管理器 - 支持 Redis 或内存的 Key 级别 RPM 限制
from __future__ import annotations
from typing import Any
import asyncio
import math
import os
@@ -30,13 +31,13 @@ class ConcurrencyManager:
_key_rpm_bucket_seconds: int = 60
_key_rpm_key_ttl_seconds: int = 120 # 2 分钟,足够覆盖当前分钟与边界
def __new__(cls):
def __new__(cls) -> "ConcurrencyManager":
"""单例模式"""
if cls._instance is None:
cls._instance = super().__new__(cls)
return cls._instance
def __init__(self):
def __init__(self) -> None:
"""初始化内存后端结构(只执行一次)"""
if hasattr(self, "_memory_initialized"):
return
@@ -90,7 +91,7 @@ class ConcurrencyManager:
if self._cleanup_task is not None:
return # 已经启动
async def cleanup_loop():
async def cleanup_loop() -> None:
"""后台清理循环"""
while True:
try:
@@ -452,7 +453,7 @@ class ConcurrencyManager:
key_rpm_limit: int | None,
is_cached_user: bool = False,
cache_reservation_ratio: float | None = None,
):
) -> Any:
"""
RPM 限制上下文管理器(支持缓存用户优先级)

View File

@@ -4,6 +4,8 @@ IP 级别的速率限制服务
提供基于 IP 地址的速率限制,防止暴力破解和 DDoS 攻击
"""
from __future__ import annotations
import ipaddress
from src.clients.redis_client import get_redis_client

View File

@@ -2,6 +2,8 @@
封装请求执行逻辑,包含并发控制与链路追踪。
"""
from __future__ import annotations
import time
from dataclasses import dataclass
from typing import Any
@@ -48,7 +50,7 @@ class ExecutionError(Exception):
class RequestExecutor:
def __init__(self, db: Session, concurrency_manager, adaptive_manager):
def __init__(self, db: Session, concurrency_manager: Any, adaptive_manager: Any) -> None:
self.db = db
self.concurrency_manager = concurrency_manager
self.adaptive_manager = adaptive_manager
@@ -56,11 +58,11 @@ class RequestExecutor:
async def execute(
self,
*,
candidate,
candidate: Any,
candidate_id: str,
candidate_index: int,
user_api_key,
request_func: Callable,
user_api_key: Any,
request_func: Callable[..., Any],
request_id: str | None,
api_format: str | APIFormat,
model_name: str,

View File

@@ -259,7 +259,7 @@ class RequestResult:
# 尝试从异常中提取 metadata
existing_metadata = getattr(exception, "request_metadata", None)
def get_meta_value(meta, key, default=None):
def get_meta_value(meta: Any, key: str, default: Any | None = None) -> Any:
"""从 metadata 中提取值,支持字典和对象两种形式"""
if meta is None:
return default
@@ -347,7 +347,7 @@ class StreamWithMetadata:
self.response_headers_container = response_headers_container
self._metadata_updated = False
def update_metadata_with_response_headers(self):
def update_metadata_with_response_headers(self) -> None:
"""使用实际的响应头更新元数据"""
if self.response_headers_container and "headers" in self.response_headers_container:
if not self._metadata_updated:
@@ -356,8 +356,8 @@ class StreamWithMetadata:
)
self._metadata_updated = True
def __aiter__(self):
def __aiter__(self) -> None:
return self.stream
async def __anext__(self):
async def __anext__(self) -> None:
return await self.stream.__anext__()

View File

@@ -2,6 +2,8 @@
公告系统服务
"""
from __future__ import annotations
from datetime import datetime, timezone
from sqlalchemy import or_

View File

@@ -3,6 +3,8 @@
记录所有重要操作和安全事件
"""
from __future__ import annotations
from datetime import datetime, timedelta, timezone
from typing import Any
@@ -108,7 +110,7 @@ class AuditService:
user_agent: str,
user_id: str | None = None, # UUID
error_reason: str | None = None,
):
) -> Any:
"""
记录登录尝试
@@ -151,7 +153,7 @@ class AuditService:
input_tokens: int | None = None,
output_tokens: int | None = None,
cost_usd: float | None = None,
):
) -> Any:
"""
记录API请求
@@ -204,7 +206,7 @@ class AuditService:
user_id: str | None = None, # UUID
severity: str = "medium",
details: dict[str, Any] | None = None,
):
) -> Any:
"""
记录安全事件

View File

@@ -2,6 +2,8 @@
系统配置服务
"""
from __future__ import annotations
import json
import time
from enum import Enum
@@ -174,7 +176,7 @@ class SystemConfigService:
}
@classmethod
def get_config(cls, db: Session, key: str, default: Any = None) -> Any | None:
def get_config(cls, db: Session, key: str, default: Any | None = None) -> Any | None:
"""获取系统配置值(带进程内缓存)"""
# 1. 检查进程内缓存
hit, cached_value = _get_cached_config(key)
@@ -225,7 +227,7 @@ class SystemConfigService:
return result
@staticmethod
def set_config(db: Session, key: str, value: Any, description: str = None) -> SystemConfig:
def set_config(db: Session, key: str, value: Any, description: str | None = None) -> SystemConfig:
"""设置系统配置值"""
config = db.query(SystemConfig).filter(SystemConfig.key == key).first()
@@ -309,7 +311,7 @@ class SystemConfigService:
return False
@classmethod
def init_default_configs(cls, db: Session):
def init_default_configs(cls, db: Session) -> None:
"""初始化默认配置"""
for key, default_config in cls.DEFAULT_CONFIGS.items():
if not db.query(SystemConfig).filter(SystemConfig.key == key).first():

View File

@@ -12,6 +12,9 @@
使用 APScheduler 进行任务调度,支持时区配置。
"""
from __future__ import annotations
from typing import Any
import asyncio
from datetime import datetime, timedelta, timezone
@@ -32,12 +35,12 @@ from src.utils.compression import compress_json
class MaintenanceScheduler:
"""系统维护任务调度器"""
def __init__(self):
def __init__(self) -> None:
self.running = False
self._interval_tasks = []
self._stats_aggregation_lock = asyncio.Lock()
async def start(self):
async def start(self) -> Any:
"""启动调度器"""
if self.running:
logger.warning("Maintenance scheduler already running")
@@ -112,7 +115,7 @@ class MaintenanceScheduler:
# 启动时执行一次初始化任务
asyncio.create_task(self._run_startup_tasks())
async def _run_startup_tasks(self):
async def _run_startup_tasks(self) -> None:
"""启动时执行的初始化任务"""
# 延迟一点执行,确保系统完全启动
await asyncio.sleep(2)
@@ -129,7 +132,7 @@ class MaintenanceScheduler:
except Exception as e:
logger.exception(f"启动时统计聚合任务出错: {e}")
async def stop(self):
async def stop(self) -> Any:
"""停止调度器"""
if not self.running:
return
@@ -142,15 +145,15 @@ class MaintenanceScheduler:
# ========== 任务函数APScheduler 直接调用异步函数) ==========
async def _scheduled_stats_aggregation(self, backfill: bool = False):
async def _scheduled_stats_aggregation(self, backfill: bool = False) -> None:
"""统计聚合任务(定时调用)"""
await self._perform_stats_aggregation(backfill=backfill)
async def _scheduled_cleanup(self):
async def _scheduled_cleanup(self) -> None:
"""清理任务(定时调用)"""
await self._perform_cleanup()
async def _scheduled_monitor(self):
async def _scheduled_monitor(self) -> None:
"""监控任务(定时调用)"""
try:
from src.database import log_pool_status
@@ -159,21 +162,21 @@ class MaintenanceScheduler:
except Exception as e:
logger.exception(f"连接池监控任务出错: {e}")
async def _scheduled_pending_cleanup(self):
async def _scheduled_pending_cleanup(self) -> None:
"""Pending 清理任务(定时调用)"""
await self._perform_pending_cleanup()
async def _scheduled_audit_cleanup(self):
async def _scheduled_audit_cleanup(self) -> None:
"""审计日志清理任务(定时调用)"""
await self._perform_audit_cleanup()
async def _scheduled_provider_checkin(self):
async def _scheduled_provider_checkin(self) -> None:
"""Provider 签到任务(定时调用)"""
await self._perform_provider_checkin()
# ========== 实际任务实现 ==========
async def _perform_stats_aggregation(self, backfill: bool = False):
async def _perform_stats_aggregation(self, backfill: bool = False) -> None:
"""执行统计聚合任务
Args:
@@ -389,7 +392,7 @@ class MaintenanceScheduler:
finally:
db.close()
async def _perform_pending_cleanup(self):
async def _perform_pending_cleanup(self) -> None:
"""执行 pending 状态清理"""
db = create_session()
try:
@@ -414,7 +417,7 @@ class MaintenanceScheduler:
finally:
db.close()
async def _perform_audit_cleanup(self):
async def _perform_audit_cleanup(self) -> None:
"""执行审计日志清理任务"""
db = create_session()
try:
@@ -478,7 +481,7 @@ class MaintenanceScheduler:
finally:
db.close()
async def _perform_provider_checkin(self):
async def _perform_provider_checkin(self) -> None:
"""执行 Provider 签到任务
遍历所有已配置 provider_ops 的 Provider触发签到。
@@ -561,7 +564,7 @@ class MaintenanceScheduler:
finally:
db.close()
async def _perform_cleanup(self):
async def _perform_cleanup(self) -> None:
"""执行清理任务"""
db = create_session()
try:

View File

@@ -8,6 +8,7 @@
from __future__ import annotations
from typing import Any, Callable
import os
from datetime import datetime
@@ -26,7 +27,7 @@ class TaskScheduler:
_instance: TaskScheduler | None = None
def __init__(self):
def __init__(self) -> None:
self.scheduler = AsyncIOScheduler(timezone=APP_TIMEZONE)
self._started = False
@@ -39,13 +40,13 @@ class TaskScheduler:
def add_cron_job(
self,
func,
func: Callable[..., Any],
hour: int,
minute: int = 0,
job_id: str = None,
name: str = None,
**kwargs,
):
job_id: str | None = None,
name: str | None = None,
**kwargs: Any,
) -> Any:
"""
添加 cron 定时任务
@@ -78,14 +79,14 @@ class TaskScheduler:
def add_interval_job(
self,
func,
seconds: int = None,
minutes: int = None,
hours: int = None,
job_id: str = None,
name: str = None,
**kwargs,
):
func: Callable[..., Any],
seconds: int | None = None,
minutes: int | None = None,
hours: int | None = None,
job_id: str | None = None,
name: str | None = None,
**kwargs: Any,
) -> Any:
"""
添加间隔执行任务
@@ -133,7 +134,7 @@ class TaskScheduler:
logger.info(f"已注册间隔任务: {display_name}, 执行间隔: {interval_desc}")
def start(self):
def start(self) -> Any:
"""启动调度器"""
if self._started:
logger.warning("调度器已在运行中")
@@ -146,7 +147,7 @@ class TaskScheduler:
# 打印下次执行时间
self._log_next_run_times()
def stop(self):
def stop(self) -> Any:
"""停止调度器"""
if not self._started:
return
@@ -155,7 +156,7 @@ class TaskScheduler:
self._started = False
logger.info("定时任务调度器已停止")
def _log_next_run_times(self):
def _log_next_run_times(self) -> None:
"""记录所有任务的下次执行时间"""
jobs = self.scheduler.get_jobs()
if not jobs:

View File

@@ -3,6 +3,8 @@
实现预聚合统计,避免每次请求都全表扫描。
"""
from __future__ import annotations
import os
import uuid
from datetime import datetime, timedelta, timezone

View File

@@ -4,6 +4,8 @@ API密钥统计同步服务
"""
from __future__ import annotations
from sqlalchemy import func
from sqlalchemy.orm import Session

View File

@@ -5,6 +5,8 @@ Usage Redis Streams consumer.
"""
from __future__ import annotations
import asyncio
import json
import os
@@ -178,7 +180,7 @@ class UsageQueueConsumer:
logger.exception(f"[usage-queue] Consumer loop error: {exc}")
await asyncio.sleep(1)
async def _maybe_claim_pending(self, redis_client) -> None:
async def _maybe_claim_pending(self, redis_client: Any) -> None:
now = time.time()
if now - self._last_claim < self._claim_interval:
return
@@ -200,7 +202,7 @@ class UsageQueueConsumer:
_, messages = result[:2]
await self._process_messages(redis_client, messages)
async def _read_new(self, redis_client) -> None:
async def _read_new(self, redis_client: Any) -> None:
result = await redis_client.xreadgroup(
groupname=self._stream_group,
consumername=self._consumer,
@@ -213,7 +215,7 @@ class UsageQueueConsumer:
for _stream, messages in result:
await self._process_messages(redis_client, messages)
async def _process_messages(self, redis_client, messages: list) -> None:
async def _process_messages(self, redis_client: Any, messages: list) -> None:
"""批量处理消息,区分 STREAMING状态更新和其他事件记录写入"""
if not messages:
return
@@ -247,7 +249,7 @@ class UsageQueueConsumer:
async def _process_streaming_batch(
self,
redis_client,
redis_client: Any,
messages: list[tuple[str, UsageEvent]],
) -> None:
"""批量处理 STREAMING 事件(状态更新)"""
@@ -269,8 +271,8 @@ class UsageQueueConsumer:
async def _process_record_batch(
self,
redis_client,
messages: list[tuple[str, dict[str, Any], UsageEvent]],
redis_client: Any,
messages: list[tuple[str, dict[str, Any], UsageEvent]],
) -> None:
"""批量处理记录类型的事件"""
db = create_session()
@@ -335,10 +337,10 @@ class UsageQueueConsumer:
async def _handle_processing_error(
self,
redis_client,
message_id: str,
fields: dict[str, Any],
error: Exception,
redis_client: Any,
message_id: str,
fields: dict[str, Any],
error: Exception,
) -> None:
retries = await self._get_delivery_count(redis_client, message_id)
if retries >= self._max_retries:
@@ -366,7 +368,7 @@ class UsageQueueConsumer:
f"[usage-queue] Processing failed (attempt {retries}): {message_id} error={error}"
)
async def _get_delivery_count(self, redis_client, message_id: str) -> int:
async def _get_delivery_count(self, redis_client: Any, message_id: str) -> int:
try:
pending = await redis_client.xpending_range(
self._stream_key,
@@ -481,7 +483,7 @@ class UsageQueueConsumer:
else:
await self._apply_record_event(event)
async def _log_metrics(self, redis_client) -> None:
async def _log_metrics(self, redis_client: Any) -> None:
now = time.time()
if now - self._last_metrics_log < self._metrics_interval:
return

View File

@@ -8,6 +8,9 @@
使用统一的 TaskScheduler 进行调度。
"""
from __future__ import annotations
from typing import Any
from datetime import datetime, timezone
from src.core.enums import ProviderBillingType
@@ -20,10 +23,10 @@ from src.services.system.scheduler import get_scheduler
class QuotaScheduler:
"""额度周期重置调度器"""
def __init__(self):
def __init__(self) -> None:
self.running = False
async def start(self):
async def start(self) -> Any:
"""启动调度器"""
if self.running:
logger.warning("Quota scheduler already running")
@@ -45,7 +48,7 @@ class QuotaScheduler:
# 启动时立即执行一次检查
await self._check_and_reset_quotas()
async def stop(self):
async def stop(self) -> Any:
"""停止调度器"""
if not self.running:
return
@@ -53,11 +56,11 @@ class QuotaScheduler:
self.running = False
logger.info("Quota scheduler stopped")
async def _scheduled_quota_check(self):
async def _scheduled_quota_check(self) -> None:
"""额度检查任务(定时调用)"""
await self._check_and_reset_quotas()
async def _check_and_reset_quotas(self):
async def _check_and_reset_quotas(self) -> None:
"""检查并重置周期额度"""
db = create_session()
@@ -114,7 +117,7 @@ class QuotaScheduler:
finally:
db.close()
async def force_reset(self, provider_id: str = None):
async def force_reset(self, provider_id: str | None = None) -> Any:
"""手动强制重置额度"""
db = create_session()
try:

View File

@@ -22,6 +22,8 @@ await recorder.record(result)
```
"""
from __future__ import annotations
from typing import Any
from sqlalchemy.orm import Session

View File

@@ -2,6 +2,8 @@
用量统计和配额管理服务
"""
from __future__ import annotations
import uuid
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone

View File

@@ -3,6 +3,8 @@
处理流式响应的token计算和使用量记录
"""
from __future__ import annotations
import json
import re
from typing import Any
@@ -132,7 +134,7 @@ class StreamUsageTracker:
self.error_message = None # 错误消息(如果有)
self.attempt_id = attempt_id
def set_error_status(self, status_code: int, error_message: str):
def set_error_status(self, status_code: int, error_message: str) -> None:
"""
设置错误状态
@@ -144,7 +146,7 @@ class StreamUsageTracker:
self.error_message = error_message
logger.debug(f"ID:{self.request_id} | 流式响应错误状态已设置 | 状态码:{status_code} | 错误:{error_message[:100]}")
def _update_complete_response(self, chunk: dict[str, Any]):
def _update_complete_response(self, chunk: dict[str, Any]) -> None:
"""根据响应块更新完整响应结构"""
try:
# 更新响应ID
@@ -590,7 +592,7 @@ class StreamUsageTracker:
# 如果连最简单的日志都失败了,放弃
pass
async def _record_usage(self):
async def _record_usage(self) -> None:
"""记录最终的使用量"""
try:
if self.request_start_time and self.end_time:
@@ -839,7 +841,7 @@ class EnhancedStreamUsageTracker(StreamUsageTracker):
# 继承父类的SSE解析缓冲区
# 这些已经在父类中初始化了
def _init_tokenizer(self):
def _init_tokenizer(self) -> None:
"""初始化分词器(如果可用)"""
try:
# 尝试导入tiktoken用于更准确的token计算

View File

@@ -3,6 +3,8 @@ Telemetry writer abstraction for stream usage.
"""
from __future__ import annotations
import json
from abc import ABC, abstractmethod
from typing import Any

View File

@@ -2,6 +2,8 @@
API密钥管理服务
"""
from __future__ import annotations
from datetime import datetime, timedelta, timezone
from typing import Any
@@ -124,7 +126,7 @@ class ApiKeyService:
return query.order_by(ApiKey.created_at.desc()).all()
@staticmethod
def update_api_key(db: Session, key_id: str, **kwargs) -> ApiKey | None: # UUID
def update_api_key(db: Session, key_id: str, **kwargs: Any) -> ApiKey | None: # UUID
"""更新API密钥"""
api_key = db.query(ApiKey).filter(ApiKey.id == key_id).first()
if not api_key:

View File

@@ -3,6 +3,8 @@
"""
from __future__ import annotations
from sqlalchemy.orm import Session
from src.core.exceptions import NotFoundException

View File

@@ -2,6 +2,8 @@
用户管理服务
"""
from __future__ import annotations
import asyncio
from datetime import datetime, timezone
from typing import Any
@@ -176,7 +178,7 @@ class UserService:
@staticmethod
@transactional()
def update_user(db: Session, user_id: str, **kwargs) -> User | None:
def update_user(db: Session, user_id: str, **kwargs: Any) -> User | None:
"""更新用户信息"""
user = db.query(User).filter(User.id == user_id).first()
if not user: