PetAgent/backend/utils/database.py
2026-05-25 15:51:16 +08:00

74 lines
2.2 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""
数据库连接和会话管理
"""
import logging
from sqlalchemy import create_engine, select
from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession, async_sessionmaker
from sqlalchemy.orm import sessionmaker
from config.settings import settings
logger = logging.getLogger(__name__)
# 创建异步引擎
# 将数据库URL转换为asyncmy驱动用于异步操作
# 首先尝试pymysql到asyncmy的转换如果失败则尝试mysql到mysql+asyncmy的转换
if 'mysql+pymysql' in settings.DATABASE_URL:
async_db_url = settings.DATABASE_URL.replace('mysql+pymysql', 'mysql+asyncmy')
else:
async_db_url = settings.DATABASE_URL.replace('mysql', 'mysql+asyncmy')
async_engine = create_async_engine(
async_db_url,
echo=settings.DEBUG,
pool_pre_ping=True
)
# 创建会话工厂
AsyncSessionLocal = async_sessionmaker(
bind=async_engine,
class_=AsyncSession,
expire_on_commit=False
)
async def get_db():
"""获取数据库会话的依赖函数"""
async with AsyncSessionLocal() as session:
try:
yield session
await session.commit()
except Exception:
await session.rollback()
raise
finally:
await session.close()
async def resolve_user_id(user_id_str: str, db_session: AsyncSession) -> int:
"""
将用户ID字符串转换为整数ID。
对于guest用户"guest_xxx"),在数据库中查找或创建对应的用户记录。
"""
try:
return int(user_id_str)
except (TypeError, ValueError):
from models.database import User
result = await db_session.execute(
select(User).where(User.uid == user_id_str)
)
user = result.scalar_one_or_none()
if user:
return user.id
# 创建新的guest用户并立即提交确保后续连接能查到
new_user = User(
uid=user_id_str,
username=user_id_str,
nickname=user_id_str,
role=1
)
db_session.add(new_user)
await db_session.flush()
await db_session.commit()
logger.info(f"Created guest user: uid={user_id_str}, id={new_user.id}")
return new_user.id