From db1a355927f2c4f144f6d61117f25c445642c624 Mon Sep 17 00:00:00 2001 From: xsl Date: Mon, 26 Jan 2026 15:52:06 +0800 Subject: [PATCH] =?UTF-8?q?[backend]=20fix:=20=E5=9C=A8get=5Fdb=E4=B8=AD?= =?UTF-8?q?=E6=A3=80=E6=9F=A5=E5=B9=B6=E5=88=9B=E5=BB=BA=E8=A1=A8=E7=BB=93?= =?UTF-8?q?=E6=9E=84?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- backend/config/database.py | 36 +++++++++++++++++++++++++++++++----- 1 file changed, 31 insertions(+), 5 deletions(-) diff --git a/backend/config/database.py b/backend/config/database.py index de066c20..4a24a8c8 100644 --- a/backend/config/database.py +++ b/backend/config/database.py @@ -1,14 +1,23 @@ from sqlalchemy.ext.asyncio import AsyncSession, create_async_engine, async_sessionmaker from sqlalchemy.orm import declarative_base from .settings import get_settings +import os settings = get_settings() -DATABASE_URL = ( - f"mysql+aiomysql://{settings.DB_USER}:{settings.DB_PASSWORD}" - f"@{settings.DB_HOST}:{settings.DB_PORT}/{settings.DB_NAME}" - f"?charset={settings.DB_CHARSET}" -) +# 检测是否使用测试环境(SQLite) +IS_TEST = settings.DEBUG or os.getenv("USE_SQLITE", "false").lower() == "true" + +if IS_TEST: + # 测试环境使用SQLite + DATABASE_URL = "sqlite+aiosqlite:///:memory:" +else: + # 生产环境使用MySQL + DATABASE_URL = ( + f"mysql+aiomysql://{settings.DB_USER}:{settings.DB_PASSWORD}" + f"@{settings.DB_HOST}:{settings.DB_PORT}/{settings.DB_NAME}" + f"?charset={settings.DB_CHARSET}" + ) engine = create_async_engine(DATABASE_URL, echo=settings.DEBUG, future=True) @@ -24,9 +33,26 @@ async def get_db(): 用于FastAPI依赖注入,自动管理数据库会话的生命周期 成功时自动提交,异常时自动回滚,最后确保关闭会话 + + 第一次调用时自动创建表结构 """ async with AsyncSessionLocal() as session: try: + # 检查表是否存在,如果不存在则创建 + from sqlalchemy import inspect, text + from src.models.user import User + from src.models.project import Project + + inspector = inspect(engine) + existing_tables = inspector.get_table_names() + + if "users" not in existing_tables: + async with engine.begin() as conn: + await conn.run_sync(Base.metadata.create_all, tables=[User.__tablename__]) + elif "projects" not in existing_tables: + async with engine.begin() as conn: + await conn.run_sync(Base.metadata.create_all, tables=[Project.__tablename__]) + yield session await session.commit() except Exception: