100 lines
2.8 KiB
Python
100 lines
2.8 KiB
Python
"""SQLAlchemy 2.x 同步引擎 + 会话管理。
|
|
|
|
支持 MySQL / SQLite,通过 DATABASE_URL 自动切换。
|
|
"""
|
|
|
|
from collections.abc import Generator
|
|
|
|
from sqlalchemy import create_engine, event, text
|
|
from sqlalchemy.orm import Session, sessionmaker
|
|
|
|
from app.core.config import get_settings
|
|
from app.core.logging import get_logger
|
|
|
|
logger = get_logger(__name__)
|
|
|
|
_engine = None
|
|
_session_factory: sessionmaker[Session] | None = None
|
|
|
|
|
|
def get_engine():
|
|
global _engine
|
|
if _engine is None:
|
|
settings = get_settings()
|
|
|
|
engine_kwargs = {
|
|
"pool_pre_ping": True,
|
|
"echo": False,
|
|
}
|
|
|
|
if settings.is_mysql:
|
|
# MySQL 配置
|
|
engine_kwargs.update({
|
|
"pool_size": 10,
|
|
"max_overflow": 20,
|
|
"pool_recycle": 3600, # 1 小时回收连接,防止 MySQL 超时断开
|
|
"connect_args": {
|
|
"charset": "utf8mb4",
|
|
},
|
|
})
|
|
elif settings.is_sqlite:
|
|
# SQLite 配置
|
|
engine_kwargs["connect_args"] = {"check_same_thread": False}
|
|
|
|
_engine = create_engine(settings.database_url, **engine_kwargs)
|
|
|
|
# SQLite: 启用 WAL 模式 + 外键约束
|
|
if settings.is_sqlite:
|
|
@event.listens_for(_engine, "connect")
|
|
def _set_sqlite_pragma(dbapi_conn, _):
|
|
cursor = dbapi_conn.cursor()
|
|
cursor.execute("PRAGMA journal_mode=WAL")
|
|
cursor.execute("PRAGMA foreign_keys=ON")
|
|
cursor.close()
|
|
|
|
# MySQL: 设置字符集和 SQL 模式
|
|
if settings.is_mysql:
|
|
@event.listens_for(_engine, "connect")
|
|
def _set_mysql_session(dbapi_conn, _):
|
|
cursor = dbapi_conn.cursor()
|
|
cursor.execute("SET NAMES utf8mb4")
|
|
cursor.execute("SET SESSION sql_mode='STRICT_TRANS_TABLES,NO_ZERO_DATE,NO_ZERO_IN_DATE,ERROR_FOR_DIVISION_BY_ZERO'")
|
|
cursor.close()
|
|
|
|
logger.info("Database engine created: %s", "MySQL" if settings.is_mysql else "SQLite")
|
|
|
|
return _engine
|
|
|
|
|
|
def get_session_factory() -> sessionmaker[Session]:
|
|
global _session_factory
|
|
if _session_factory is None:
|
|
_session_factory = sessionmaker(
|
|
get_engine(),
|
|
class_=Session,
|
|
expire_on_commit=False,
|
|
)
|
|
return _session_factory
|
|
|
|
|
|
def get_session() -> Generator[Session, None, None]:
|
|
"""FastAPI 依赖:请求级会话。"""
|
|
factory = get_session_factory()
|
|
session = factory()
|
|
try:
|
|
yield session
|
|
session.commit()
|
|
except Exception:
|
|
session.rollback()
|
|
raise
|
|
finally:
|
|
session.close()
|
|
|
|
|
|
def dispose_engine() -> None:
|
|
global _engine, _session_factory
|
|
if _engine is not None:
|
|
_engine.dispose()
|
|
_engine = None
|
|
_session_factory = None
|