refactor: 提取正则工具模块并清理废弃代码

前端:
- 新增 model-mapping-regex.ts 工具模块,统一正则验证和 LRU 缓存
- ModelMappingsTab/RoutingTab 重构为使用 computed 缓存匹配结果
- 移除多个组件中未使用的变量和函数
- 添加 HTMLImageElement/HTMLIFrameElement 到 ESLint 全局类型
- 修复 vitest 需要 --experimental-require-module 的问题

后端:
- 移除废弃的异步数据库支持 (get_async_db, AsyncSession)
- 将 async_utils 从 database/ 迁移到 utils/
- 改进 database.py 类型标注
- 修复 email 模块的 aiosmtplib 可选导入类型问题
This commit is contained in:
fawney19
2026-01-15 17:03:19 +08:00
parent bc16c0eec8
commit a223819dd7
25 changed files with 526 additions and 688 deletions

View File

@@ -3,7 +3,7 @@
"""
from ..models.database import ApiKey, Base, Usage, User, UserQuota
from .database import create_session, get_async_db, get_db, get_db_url, init_db, log_pool_status
from .database import create_session, get_db, get_db_url, init_db, log_pool_status
__all__ = [
"Base",
@@ -12,7 +12,6 @@ __all__ = [
"Usage",
"UserQuota",
"get_db",
"get_async_db",
"init_db",
"create_session",
"get_db_url",

View File

@@ -1,41 +0,0 @@
"""
异步数据库工具
提供在异步上下文中安全使用同步数据库操作的工具
"""
import asyncio
from functools import wraps
from typing import Any, Callable, Coroutine, TypeVar
T = TypeVar("T")
async def run_in_executor(func: Callable[..., T], *args: Any, **kwargs: Any) -> T:
"""
在线程池中运行同步函数,避免阻塞事件循环
用法:
result = await run_in_executor(some_sync_function, arg1, arg2)
"""
loop = asyncio.get_event_loop()
return await loop.run_in_executor(None, lambda: func(*args, **kwargs))
def async_wrap_sync_db(func: Callable[..., T]) -> Callable[..., Coroutine[Any, Any, T]]:
"""
装饰器:包装同步数据库函数为异步函数
用法:
@async_wrap_sync_db
def get_user(db: Session, user_id: int):
return db.query(User).filter(User.id == user_id).first()
# 现在可以在异步上下文中调用
user = await get_user(db, 123)
"""
@wraps(func)
async def wrapper(*args: Any, **kwargs: Any) -> T:
return await run_in_executor(func, *args, **kwargs)
return wrapper

View File

@@ -3,19 +3,13 @@
"""
import time
from typing import AsyncGenerator, Generator, Optional
from typing import Any, Generator, Optional, cast
from starlette.requests import Request
from sqlalchemy import create_engine, event
from sqlalchemy.engine import Engine
from sqlalchemy.ext.asyncio import (
AsyncEngine,
AsyncSession,
async_sessionmaker,
create_async_engine,
)
from sqlalchemy.orm import Session, sessionmaker
from sqlalchemy.pool import Pool, QueuePool
from sqlalchemy.pool import QueuePool
from ..config import config
from src.core.logger import logger
@@ -24,29 +18,27 @@ from ..models.database import Base, SystemConfig, User, UserRole
# 延迟初始化的数据库引擎和会话工厂
_engine: Optional[Engine] = None
_SessionLocal: Optional[sessionmaker] = None
_async_engine: Optional[AsyncEngine] = None
_AsyncSessionLocal: Optional[async_sessionmaker] = None
_SessionLocal: Optional[sessionmaker[Session]] = None
# 连接池监控
_last_pool_warning: float = 0.0
POOL_WARNING_INTERVAL = 60 # 每60秒最多警告一次
def _setup_pool_monitoring(engine: Engine):
def _setup_pool_monitoring(engine: Engine) -> None:
"""设置连接池监控事件"""
@event.listens_for(engine, "connect")
def receive_connect(dbapi_conn, connection_record):
def receive_connect(dbapi_conn: Any, connection_record: Any) -> None:
"""连接创建时的监控"""
pass
@event.listens_for(engine, "checkout")
def receive_checkout(dbapi_conn, connection_record, connection_proxy):
def receive_checkout(dbapi_conn: Any, connection_record: Any, connection_proxy: Any) -> None:
"""从连接池检出连接时的监控"""
global _last_pool_warning
pool = engine.pool
pool = cast(QueuePool, engine.pool)
# 获取连接池状态
checked_out = pool.checkedout()
pool_size = pool.size()
@@ -70,10 +62,10 @@ def _setup_pool_monitoring(engine: Engine):
)
def get_pool_status() -> dict:
def get_pool_status() -> dict[str, Any]:
"""获取连接池状态"""
engine = _ensure_engine()
pool = engine.pool
pool = cast(QueuePool, engine.pool)
return {
"checked_out": pool.checkedout(),
@@ -84,7 +76,7 @@ def get_pool_status() -> dict:
}
def log_pool_status():
def log_pool_status() -> None:
"""记录连接池状态到日志(用于监控)"""
try:
status = get_pool_status()
@@ -147,7 +139,7 @@ def _ensure_engine() -> Engine:
return _engine
def _log_pool_capacity():
def _log_pool_capacity() -> None:
theoretical = config.db_pool_size + config.db_max_overflow
workers = max(1, config.worker_processes)
total_estimated = theoretical * workers
@@ -169,82 +161,6 @@ def _log_pool_capacity():
)
def _ensure_async_engine() -> AsyncEngine:
"""
确保异步数据库引擎已创建(延迟加载)
这允许异步路由使用非阻塞的数据库访问
"""
global _async_engine, _AsyncSessionLocal
if _async_engine is not None:
return _async_engine
# 获取数据库配置并转换为异步URL
DATABASE_URL = config.database_url
# 转换同步URL为异步URLpostgresql:// -> postgresql+asyncpg://
if DATABASE_URL.startswith("postgresql://"):
ASYNC_DATABASE_URL = DATABASE_URL.replace("postgresql://", "postgresql+asyncpg://", 1)
elif DATABASE_URL.startswith("sqlite:///"):
ASYNC_DATABASE_URL = DATABASE_URL.replace("sqlite:///", "sqlite+aiosqlite:///", 1)
else:
raise ValueError(f"不支持的数据库类型: {DATABASE_URL}")
# 验证数据库类型(生产环境要求 PostgreSQL
is_production = config.environment == "production"
if is_production and not ASYNC_DATABASE_URL.startswith("postgresql+asyncpg://"):
raise ValueError("生产环境只支持 PostgreSQL 数据库,请配置正确的 DATABASE_URL")
# 创建异步引擎
_async_engine = create_async_engine(
ASYNC_DATABASE_URL,
# AsyncEngine 不能使用 QueuePool默认使用 AsyncAdaptedQueuePool
pool_size=config.db_pool_size,
max_overflow=config.db_max_overflow,
pool_timeout=config.db_pool_timeout,
pool_recycle=config.db_pool_recycle,
pool_pre_ping=True,
echo=False,
)
# 创建异步会话工厂
_AsyncSessionLocal = async_sessionmaker(
_async_engine,
class_=AsyncSession,
expire_on_commit=False,
autocommit=False,
autoflush=False,
)
logger.debug(f"异步数据库引擎已初始化: {ASYNC_DATABASE_URL.split('@')[-1] if '@' in ASYNC_DATABASE_URL else 'local'}")
return _async_engine
async def get_async_db() -> AsyncGenerator[AsyncSession, None]:
"""获取异步数据库会话
.. deprecated::
此方法已废弃,项目统一使用同步 Session。
未来版本可能移除此方法。请使用 get_db() 代替。
"""
import warnings
warnings.warn(
"get_async_db() 已废弃,项目统一使用同步 Session。请使用 get_db() 代替。",
DeprecationWarning,
stacklevel=2,
)
# 确保异步引擎已初始化
_ensure_async_engine()
async with _AsyncSessionLocal() as session:
try:
yield session
finally:
await session.close()
def get_db(request: Request = None) -> Generator[Session, None, None]: # type: ignore[assignment]
"""获取数据库会话
@@ -298,6 +214,7 @@ def get_db(request: Request = None) -> Generator[Session, None, None]: # type:
# 确保引擎已初始化
_ensure_engine()
assert _SessionLocal is not None
db = _SessionLocal()
@@ -347,6 +264,7 @@ def create_session() -> Session:
db.close()
"""
_ensure_engine()
assert _SessionLocal is not None
return _SessionLocal()
@@ -355,7 +273,7 @@ def get_db_url() -> str:
return config.database_url
def init_db():
def init_db() -> None:
"""初始化数据库
注意:数据库表结构由 Alembic 管理,部署时请运行 ./migrate.sh
@@ -367,6 +285,7 @@ def init_db():
# 确保引擎已创建
_ensure_engine()
assert _SessionLocal is not None
# 数据库表结构由 Alembic 迁移管理
@@ -424,7 +343,7 @@ def init_db():
db.close()
def init_admin_user(db: Session):
def init_admin_user(db: Session) -> None:
"""从环境变量创建管理员账户"""
# 检查是否使用默认凭据
if config.admin_email == "admin@localhost" and config.admin_password == "admin123":
@@ -447,9 +366,9 @@ def init_admin_user(db: Session):
email=config.admin_email,
username=config.admin_username,
role=UserRole.ADMIN,
quota_usd=1000.0,
is_active=True,
)
admin.quota_usd = cast(Any, 1000.0)
admin.set_password(config.admin_password)
db.add(admin)
@@ -461,7 +380,7 @@ def init_admin_user(db: Session):
raise
def init_default_models(db: Session):
def init_default_models(db: Session) -> None:
"""初始化默认模型配置"""
# 注意:作为中转代理服务,不再预设模型配置
@@ -470,10 +389,10 @@ def init_default_models(db: Session):
pass
def init_system_configs(db: Session):
def init_system_configs(db: Session) -> None:
"""初始化系统配置"""
configs = [
configs: list[dict[str, Any]] = [
{"key": "default_user_quota_usd", "value": 10.0, "description": "新用户默认美元配额"},
{"key": "rate_limit_per_minute", "value": 60, "description": "每分钟请求限制"},
{"key": "enable_registration", "value": False, "description": "是否开放用户注册"},
@@ -484,12 +403,15 @@ def init_system_configs(db: Session):
for config_data in configs:
existing = db.query(SystemConfig).filter_by(key=config_data["key"]).first()
if not existing:
config = SystemConfig(**config_data)
db.add(config)
row = SystemConfig()
row.key = config_data["key"]
row.value = config_data["value"]
row.description = config_data["description"]
db.add(row)
logger.info(f"添加系统配置: {config_data['key']}")
def reset_db():
def reset_db() -> None:
"""重置数据库(仅用于开发)"""
logger.warning("重置数据库...")

View File

@@ -11,28 +11,30 @@ from email.mime.multipart import MIMEMultipart
from email.mime.text import MIMEText
from typing import Any, Dict, List, Optional
aiosmtplib: Any
try:
import aiosmtplib
AIOSMTPLIB_AVAILABLE = True
import aiosmtplib as _aiosmtplib
except ImportError:
AIOSMTPLIB_AVAILABLE = False
aiosmtplib = None
else:
AIOSMTPLIB_AVAILABLE = True
aiosmtplib = _aiosmtplib
from src.core.logger import logger
from src.utils.async_utils import run_in_executor
from .base import Notification, NotificationLevel, NotificationPlugin
class EmailNotificationPlugin(NotificationPlugin):
"""
邮件通知插件
支持HTML和纯文本邮件
"""
def __init__(self, name: str = "email", config: Dict[str, Any] = None):
super().__init__(name, config)
def __init__(self, name: str = "email", config: Optional[Dict[str, Any]] = None):
super().__init__(name, config or {})
# SMTP配置
self.smtp_host = config.get("smtp_host") if config else None
@@ -60,7 +62,7 @@ class EmailNotificationPlugin(NotificationPlugin):
# 缓冲配置
self._buffer: List[Notification] = []
self._lock = asyncio.Lock()
self._flush_task = None
self._flush_task: Optional[asyncio.Task[None]] = None
# 验证配置
config_errors = []
@@ -97,10 +99,10 @@ class EmailNotificationPlugin(NotificationPlugin):
return True
def _start_flush_task(self):
def _start_flush_task(self) -> None:
"""启动定时刷新任务"""
async def flush_loop():
async def flush_loop() -> None:
while self.enabled:
await asyncio.sleep(self.flush_interval)
await self.flush()
@@ -235,9 +237,7 @@ class EmailNotificationPlugin(NotificationPlugin):
return False
else:
# 使用同步SMTP在线程中运行
return await asyncio.get_event_loop().run_in_executor(
None, self._send_email_sync, subject, body, is_html
)
return await run_in_executor(self._send_email_sync, subject, body, is_html)
def _send_email_sync(self, subject: str, body: str, is_html: bool = True) -> bool:
"""同步发送邮件"""
@@ -256,11 +256,15 @@ class EmailNotificationPlugin(NotificationPlugin):
message.attach(MIMEText(body, "plain"))
try:
smtp_host = self.smtp_host
assert smtp_host is not None
# 连接SMTP服务器
server: smtplib.SMTP
if self.use_ssl:
server = smtplib.SMTP_SSL(self.smtp_host, self.smtp_port)
server = smtplib.SMTP_SSL(smtp_host, self.smtp_port)
else:
server = smtplib.SMTP(self.smtp_host, self.smtp_port)
server = smtplib.SMTP(smtp_host, self.smtp_port)
if self.use_tls:
server.starttls()
@@ -299,7 +303,7 @@ class EmailNotificationPlugin(NotificationPlugin):
return True
async def _do_send_batch(self, notifications: List[Notification]) -> Dict[str, Any]:
async def _do_send_batch(self, notifications: List[Notification]) -> Dict[str, int]:
"""实际批量发送通知"""
if not notifications:
return {"total": 0, "sent": 0, "failed": 0}
@@ -357,7 +361,7 @@ class EmailNotificationPlugin(NotificationPlugin):
"aiosmtplib_available": AIOSMTPLIB_AVAILABLE,
}
async def close(self):
async def close(self) -> None:
"""关闭插件"""
# 刷新缓冲
await self.flush()
@@ -366,7 +370,7 @@ class EmailNotificationPlugin(NotificationPlugin):
if self._flush_task:
self._flush_task.cancel()
def __del__(self):
def __del__(self) -> None:
"""清理资源"""
try:
asyncio.create_task(self.close())

View File

@@ -3,20 +3,21 @@
提供 SMTP 邮件发送功能
"""
import asyncio
import smtplib
import ssl
from email.mime.multipart import MIMEMultipart
from email.mime.text import MIMEText
from typing import Optional, Tuple
from typing import Any, Optional, Tuple, Union
aiosmtplib: Any
try:
import aiosmtplib
AIOSMTPLIB_AVAILABLE = True
import aiosmtplib as _aiosmtplib
except ImportError:
AIOSMTPLIB_AVAILABLE = False
aiosmtplib = None
else:
AIOSMTPLIB_AVAILABLE = True
aiosmtplib = _aiosmtplib
def _create_ssl_context() -> ssl.SSLContext:
@@ -34,6 +35,7 @@ from sqlalchemy.orm import Session
from src.core.crypto import crypto_service
from src.core.logger import logger
from src.services.system.config import SystemConfigService
from src.utils.async_utils import run_in_executor
from .email_template import EmailTemplate
@@ -260,9 +262,13 @@ class EmailSenderService:
Returns:
(是否发送成功, 错误信息)
"""
loop = asyncio.get_event_loop()
return await loop.run_in_executor(
None, EmailSenderService._send_email_sync, config, to_email, subject, html_body, text_body
return await run_in_executor(
EmailSenderService._send_email_sync,
config,
to_email,
subject,
html_body,
text_body,
)
@staticmethod
@@ -302,7 +308,7 @@ class EmailSenderService:
message.attach(MIMEText(html_body, "html", "utf-8"))
# 连接 SMTP 服务器
server = None
server: Optional[smtplib.SMTP] = None
ssl_context = _create_ssl_context()
try:
if config["smtp_use_ssl"]:
@@ -321,6 +327,8 @@ class EmailSenderService:
if config["smtp_use_tls"]:
server.starttls(context=ssl_context)
assert server is not None
# 登录
if config["smtp_user"] and config["smtp_password"]:
server.login(config["smtp_user"], config["smtp_password"])
@@ -393,6 +401,7 @@ class EmailSenderService:
await smtp.quit()
else:
# 使用同步方式测试
server: Union[smtplib.SMTP, smtplib.SMTP_SSL]
if config["smtp_use_ssl"]:
server = smtplib.SMTP_SSL(
config["smtp_host"],

44
src/utils/async_utils.py Normal file
View File

@@ -0,0 +1,44 @@
"""
异步工具函数
提供在异步上下文中安全执行同步函数的工具,避免阻塞事件循环。
"""
from __future__ import annotations
import asyncio
from functools import partial, wraps
from typing import Any, Callable, Coroutine, TypeVar
T = TypeVar("T")
async def run_in_executor(func: Callable[..., T], *args: Any, **kwargs: Any) -> T:
"""
在线程池中运行同步函数,避免阻塞事件循环。
用法:
result = await run_in_executor(some_sync_function, arg1, arg2)
"""
loop = asyncio.get_running_loop()
bound = partial(func, *args, **kwargs)
return await loop.run_in_executor(None, bound)
def async_wrap_sync(func: Callable[..., T]) -> Callable[..., Coroutine[Any, Any, T]]:
"""
装饰器:将同步函数包装成异步函数(在线程池中执行)。
用法:
@async_wrap_sync
def do_sync(...): ...
result = await do_sync(...)
"""
@wraps(func)
async def wrapper(*args: Any, **kwargs: Any) -> T:
return await run_in_executor(func, *args, **kwargs)
return wrapper