mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-03 01:40:21 +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:
@@ -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 自动生成)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 输入返回本地时间字符串(无时区信息)。
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
"""
|
||||
刷新缓存
|
||||
|
||||
|
||||
@@ -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
|
||||
"""删除模型
|
||||
|
||||
删除逻辑:
|
||||
|
||||
@@ -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:
|
||||
"""
|
||||
|
||||
@@ -22,6 +22,8 @@
|
||||
"""
|
||||
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, NoReturn
|
||||
|
||||
from collections.abc import Callable
|
||||
|
||||
@@ -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:
|
||||
"""
|
||||
选择提供商
|
||||
|
||||
|
||||
@@ -4,6 +4,8 @@ Provider 操作服务
|
||||
提供操作执行、凭据管理、缓存等业务逻辑。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
from dataclasses import asdict
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -13,6 +13,8 @@
|
||||
3. 安全探测:允许在稳定后尝试更高 RPM
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, cast
|
||||
|
||||
|
||||
@@ -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 限制上下文管理器(支持缓存用户优先级)
|
||||
|
||||
|
||||
@@ -4,6 +4,8 @@ IP 级别的速率限制服务
|
||||
提供基于 IP 地址的速率限制,防止暴力破解和 DDoS 攻击
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import ipaddress
|
||||
|
||||
from src.clients.redis_client import get_redis_client
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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__()
|
||||
|
||||
@@ -2,6 +2,8 @@
|
||||
公告系统服务
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from sqlalchemy import or_
|
||||
|
||||
@@ -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:
|
||||
"""
|
||||
记录安全事件
|
||||
|
||||
|
||||
@@ -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():
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -3,6 +3,8 @@
|
||||
实现预聚合统计,避免每次请求都全表扫描。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import uuid
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
||||
@@ -4,6 +4,8 @@ API密钥统计同步服务
|
||||
"""
|
||||
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from sqlalchemy import func
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -22,6 +22,8 @@ await recorder.record(result)
|
||||
```
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
@@ -2,6 +2,8 @@
|
||||
用量统计和配额管理服务
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
||||
@@ -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计算
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -3,6 +3,8 @@
|
||||
"""
|
||||
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.core.exceptions import NotFoundException
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user