51 lines
1.3 KiB
Python
51 lines
1.3 KiB
Python
"""数据库引擎与会话工厂(SQLite 本地过渡 / PostgreSQL 生产)。"""
|
|
from __future__ import annotations
|
|
|
|
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
|
|
from sqlalchemy.orm import DeclarativeBase
|
|
|
|
from config import get_settings
|
|
|
|
|
|
class Base(DeclarativeBase):
|
|
pass
|
|
|
|
|
|
_engine = None
|
|
_session_factory = None
|
|
|
|
|
|
def get_engine():
|
|
global _engine
|
|
if _engine is None:
|
|
settings = get_settings()
|
|
connect_args: dict = {}
|
|
# SQLite 需允许多线程/多协程访问同一文件
|
|
if settings.database_url.startswith("sqlite"):
|
|
connect_args["check_same_thread"] = False
|
|
_engine = create_async_engine(
|
|
settings.database_url,
|
|
echo=False,
|
|
future=True,
|
|
connect_args=connect_args,
|
|
)
|
|
return _engine
|
|
|
|
|
|
def get_session_factory() -> async_sessionmaker[AsyncSession]:
|
|
global _session_factory
|
|
if _session_factory is None:
|
|
_session_factory = async_sessionmaker(
|
|
get_engine(),
|
|
class_=AsyncSession,
|
|
expire_on_commit=False,
|
|
)
|
|
return _session_factory
|
|
|
|
|
|
async def get_db():
|
|
"""FastAPI 依赖:请求级 AsyncSession。"""
|
|
factory = get_session_factory()
|
|
async with factory() as session:
|
|
yield session
|