153 lines
5.0 KiB
Python
153 lines
5.0 KiB
Python
import logging
|
||
from contextlib import asynccontextmanager
|
||
from datetime import datetime, timedelta
|
||
from zoneinfo import ZoneInfo
|
||
|
||
from sqlalchemy import event
|
||
from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession, async_sessionmaker
|
||
from sqlalchemy.orm import declarative_base
|
||
|
||
from app.config import settings
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
|
||
def _get_utc_offset_str(tz_name: str) -> str:
|
||
"""将 IANA 时区名转为 MySQL 兼容的 UTC 偏移字符串,如 '+08:00'"""
|
||
now = datetime.now(ZoneInfo(tz_name))
|
||
offset = now.utcoffset() or timedelta()
|
||
total_seconds = int(offset.total_seconds())
|
||
sign = "+" if total_seconds >= 0 else "-"
|
||
hours, remainder = divmod(abs(total_seconds), 3600)
|
||
minutes = remainder // 60
|
||
return f"{sign}{hours:02d}:{minutes:02d}"
|
||
|
||
|
||
def _build_connect_args() -> dict:
|
||
"""根据数据库类型构建 connect_args,在连接级别设置时区"""
|
||
if settings.DB_TYPE == "postgresql":
|
||
return {"server_settings": {"timezone": settings.TIMEZONE}}
|
||
return {}
|
||
|
||
|
||
# 创建异步引擎
|
||
engine = create_async_engine(
|
||
settings.DATABASE_URL,
|
||
echo=settings.DEBUG,
|
||
connect_args=_build_connect_args(),
|
||
pool_size=5, # 连接池常驻连接数(减少以避免连接耗尽)
|
||
max_overflow=10, # 超出pool_size后可创建的连接数(总计最多15)
|
||
pool_timeout=30, # 获取连接的超时时间(秒)
|
||
pool_recycle=600, # 连接回收时间(秒),更积极地回收空闲连接
|
||
pool_pre_ping=True, # 使用前检查连接是否有效,自动重连
|
||
pool_reset_on_return="rollback", # 连接归还时重置状态
|
||
)
|
||
|
||
|
||
# MySQL / SQL Server 通过 connect 事件设置会话时区
|
||
if settings.DB_TYPE in ("mysql", "sqlserver"):
|
||
@event.listens_for(engine.sync_engine, "connect")
|
||
def _set_session_timezone(dbapi_conn, connection_rec):
|
||
offset_str = _get_utc_offset_str(settings.TIMEZONE)
|
||
cursor = dbapi_conn.cursor()
|
||
try:
|
||
if settings.DB_TYPE == "mysql":
|
||
cursor.execute(f"SET time_zone = '{offset_str}'")
|
||
elif settings.DB_TYPE == "sqlserver":
|
||
pass # SQL Server 不支持会话级时区设置,需在应用层处理
|
||
finally:
|
||
cursor.close()
|
||
logger.debug(f"[DB] Session timezone set to {settings.TIMEZONE} ({offset_str})")
|
||
|
||
|
||
# 连接池监控事件
|
||
@event.listens_for(engine.sync_engine, "checkout")
|
||
def _on_checkout(dbapi_conn, connection_rec, connection_proxy):
|
||
pool = engine.sync_engine.pool
|
||
logger.debug(
|
||
f"[Pool] checkout: size={pool.size()}, checkedin={pool.checkedin()}, "
|
||
f"checkedout={pool.checkedout()}, overflow={pool.overflow()}"
|
||
)
|
||
|
||
|
||
@event.listens_for(engine.sync_engine, "checkin")
|
||
def _on_checkin(dbapi_conn, connection_rec):
|
||
pool = engine.sync_engine.pool
|
||
logger.debug(
|
||
f"[Pool] checkin: size={pool.size()}, checkedin={pool.checkedin()}, "
|
||
f"checkedout={pool.checkedout()}, overflow={pool.overflow()}"
|
||
)
|
||
|
||
logger.info(f"[DB] Engine created with timezone: {settings.TIMEZONE} (db_type={settings.DB_TYPE})")
|
||
|
||
# 创建异步会话工厂
|
||
AsyncSessionLocal = async_sessionmaker(
|
||
bind=engine,
|
||
class_=AsyncSession,
|
||
expire_on_commit=False,
|
||
autocommit=False,
|
||
autoflush=False,
|
||
)
|
||
|
||
# 声明基类
|
||
Base = declarative_base()
|
||
|
||
|
||
async def get_db() -> AsyncSession:
|
||
"""获取数据库会话的依赖函数"""
|
||
async with AsyncSessionLocal() as session:
|
||
try:
|
||
yield session
|
||
finally:
|
||
await session.close()
|
||
|
||
|
||
async def get_db_transaction() -> AsyncSession:
|
||
"""
|
||
获取带事务的数据库会话依赖函数
|
||
|
||
使用方式:
|
||
@router.post("/")
|
||
async def create_something(db: AsyncSession = Depends(get_db_transaction)):
|
||
# 所有数据库操作在同一事务中
|
||
# 如果发生异常,自动回滚
|
||
# 如果成功完成,自动提交
|
||
...
|
||
|
||
注意:使用此依赖时,Service层的方法不应调用commit(),
|
||
因为事务会在API结束时统一提交或回滚。
|
||
可以使用BaseService的_no_commit版本方法,或手动控制。
|
||
"""
|
||
async with AsyncSessionLocal() as session:
|
||
try:
|
||
yield session
|
||
await session.commit()
|
||
except Exception:
|
||
await session.rollback()
|
||
raise
|
||
finally:
|
||
await session.close()
|
||
|
||
|
||
@asynccontextmanager
|
||
async def transaction(db: AsyncSession):
|
||
"""
|
||
事务上下文管理器,用于在API中包装多个操作
|
||
|
||
使用方式:
|
||
@router.post("/")
|
||
async def create_something(db: AsyncSession = Depends(get_db)):
|
||
async with transaction(db):
|
||
# 所有操作在同一事务中
|
||
await SomeService.create_no_commit(db, data1)
|
||
await OtherService.create_no_commit(db, data2)
|
||
# 如果发生异常,自动回滚
|
||
# 如果成功完成,自动提交
|
||
"""
|
||
try:
|
||
yield db
|
||
await db.commit()
|
||
except Exception:
|
||
await db.rollback()
|
||
raise
|