diff --git a/apps/worker/worker_app/db.py b/apps/worker/worker_app/db.py index 32c503773..ec59f96a5 100644 --- a/apps/worker/worker_app/db.py +++ b/apps/worker/worker_app/db.py @@ -8,9 +8,10 @@ from packages.adapters.sqlalchemy_impl import ( from packages.adapters.sqlalchemy_impl.schema_guard import assert_auto_create_schema_allowed settings = get_settings() -ensure_database_exists(settings.database_url) +_db_url = settings.effective_database_url +ensure_database_exists(_db_url) engine, SessionLocal = build_session_factory( - settings.database_url, + _db_url, pool_size=settings.database_pool_size, max_overflow=settings.database_max_overflow, pool_timeout=settings.database_pool_timeout, diff --git a/packages/adapters/sqlalchemy_impl/session.py b/packages/adapters/sqlalchemy_impl/session.py index c48d33d72..89f541368 100644 --- a/packages/adapters/sqlalchemy_impl/session.py +++ b/packages/adapters/sqlalchemy_impl/session.py @@ -48,12 +48,20 @@ def build_session_factory( return engine, session_factory +def _is_sqlite(database_url: str) -> bool: + """检测是否为 SQLite 数据库 URL.""" + return database_url.startswith("sqlite") + + def _build_admin_url(database_url: str) -> URL: url = make_url(database_url) return url.set(database="postgres") def ensure_database_exists(database_url: str) -> None: + """确保数据库存在(仅 PostgreSQL 需要,SQLite 自动创建).""" + if _is_sqlite(database_url): + return target_url = make_url(database_url) admin_engine = create_engine(_build_admin_url(database_url), isolation_level="AUTOCOMMIT") try: @@ -70,6 +78,14 @@ def ensure_database_exists(database_url: str) -> None: def initialize_database(engine) -> None: + """初始化数据库 schema。 + + PostgreSQL 使用 advisory lock 防止并发初始化冲突; + SQLite 直接 create_all(单文件,无并发风险)。 + """ + if _is_sqlite(str(engine.url)): + Base.metadata.create_all(bind=engine) + return with engine.connect() as connection: connection.execute(text("SELECT pg_advisory_lock(:lock_id)"), {"lock_id": SCHEMA_INIT_LOCK_ID}) try: diff --git a/packages/config/base.py b/packages/config/base.py index 905c09667..d0ef1d5ac 100755 --- a/packages/config/base.py +++ b/packages/config/base.py @@ -34,6 +34,9 @@ class SharedSettings(BaseSettings): database_pool_timeout: int = 30 database_pool_recycle: int = 3600 + # 测试用:使用 SQLite 内存数据库(CI 环境无需 PostgreSQL) + use_in_memory_db: bool = False + # ── Redis ──────────────────────────────────────────────────────────── redis_url: str = "redis://localhost:6379/0" @@ -66,6 +69,16 @@ class SharedSettings(BaseSettings): doubao_timeout: int = 30 doubao_max_retries: int = 2 + @property + def effective_database_url(self) -> str: + """返回实际使用的数据库 URL。 + + 当 USE_IN_MEMORY_DB=True 时返回 SQLite 内存 URL,否则返回 database_url。 + """ + if self.use_in_memory_db: + return "sqlite:///./test.db" + return self.database_url + model_config = SettingsConfigDict( env_file=".env", env_file_encoding="utf-8", diff --git a/pytest.ini b/pytest.ini index 92a1f25d3..a2b3f4486 100755 --- a/pytest.ini +++ b/pytest.ini @@ -1,6 +1,8 @@ [pytest] pythonpath = . apps/api apps/worker packages testpaths = tests +# importlib 模式避免同名测试文件的模块名冲突 +addopts = --import-mode=importlib # ===== 覆盖率配置 ===== # 覆盖率统计范围(供 --cov 使用时的默认源)