108 lines
3.1 KiB
Python
108 lines
3.1 KiB
Python
# shared/database/database.py
|
|
import sys
|
|
import os
|
|
from pathlib import Path
|
|
|
|
from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession, async_sessionmaker
|
|
from sqlalchemy import text
|
|
from sqlalchemy.orm import sessionmaker
|
|
from shared.config.settings import settings
|
|
import asyncio
|
|
from contextlib import asynccontextmanager
|
|
from shared.utils.logger import get_logger
|
|
|
|
logger = get_logger(__name__)
|
|
|
|
class DatabaseManager:
|
|
"""数据库管理器"""
|
|
|
|
def __init__(self):
|
|
self.engine = None
|
|
self.async_session = None
|
|
self.is_connected = False
|
|
|
|
async def connect(self):
|
|
"""连接数据库"""
|
|
if not settings.DATABASE_URL:
|
|
logger.warning("未配置数据库连接,跳过数据库初始化")
|
|
self.is_connected = False
|
|
return
|
|
|
|
try:
|
|
# 创建异步引擎
|
|
self.engine = create_async_engine(
|
|
settings.DATABASE_URL,
|
|
echo=settings.DEBUG,
|
|
pool_size=20,
|
|
max_overflow=30,
|
|
pool_recycle=3600
|
|
)
|
|
|
|
# 创建异步会话工厂
|
|
self.async_session = async_sessionmaker(
|
|
self.engine,
|
|
class_=AsyncSession,
|
|
expire_on_commit=False
|
|
)
|
|
|
|
# 测试连接
|
|
async with self.engine.begin() as conn:
|
|
await conn.execute(text("SELECT 1"))
|
|
|
|
self.is_connected = True
|
|
logger.info("数据库连接成功")
|
|
|
|
except Exception as e:
|
|
logger.error(f"数据库连接失败: {e}")
|
|
self.is_connected = False
|
|
raise
|
|
|
|
async def disconnect(self):
|
|
"""断开数据库连接"""
|
|
if self.engine:
|
|
await self.engine.dispose()
|
|
self.is_connected = False
|
|
logger.info("数据库连接已断开")
|
|
|
|
@asynccontextmanager
|
|
async def session(self):
|
|
"""获取数据库会话的异步上下文管理器"""
|
|
if not self.is_connected:
|
|
await self.connect()
|
|
|
|
session = self.async_session()
|
|
try:
|
|
yield session
|
|
finally:
|
|
await session.close()
|
|
|
|
async def get_session(self) -> AsyncSession:
|
|
"""获取数据库会话"""
|
|
if not self.is_connected:
|
|
await self.connect()
|
|
|
|
return self.async_session()
|
|
|
|
async def create_tables(self):
|
|
"""创建数据库表"""
|
|
from shared.models.database import Base
|
|
|
|
try:
|
|
async with self.engine.begin() as conn:
|
|
await conn.run_sync(Base.metadata.create_all)
|
|
logger.info("数据库表创建成功")
|
|
except Exception as e:
|
|
logger.error(f"数据库表创建失败: {e}")
|
|
raise
|
|
|
|
# 全局数据库管理器实例
|
|
db_manager = DatabaseManager()
|
|
|
|
# 数据库依赖注入
|
|
async def get_db_session():
|
|
"""获取数据库会话的依赖函数"""
|
|
session = await db_manager.get_session()
|
|
try:
|
|
yield session
|
|
finally:
|
|
await session.close() |