fix: Celery 异步任务 event loop 冲突修复

- 原因: 全局 async_session_factory 绑定在 Celery 主进程 event loop
- Celery fork worker 后创建新 loop,旧 connection 绑定冲突
- 修复: 每个任务内创建独立 async engine + session factory
- 任务完成后 dispose engine 释放连接
This commit is contained in:
zwt13703 2026-07-06 17:45:05 +08:00
parent 40b861cfcc
commit 3db8208c55
1 changed files with 18 additions and 7 deletions

View File

@ -1,9 +1,8 @@
import os import os
import uuid
from datetime import datetime, timezone from datetime import datetime, timezone
from sqlalchemy import select from sqlalchemy import select
from sqlalchemy.ext.asyncio import create_async_engine, async_sessionmaker, AsyncSession
from app.tasks.celery_app import celery_app from app.tasks.celery_app import celery_app
from app.core.database import async_session_factory
from app.core.config import get_settings from app.core.config import get_settings
from app.models.generation_task import GenerationTask from app.models.generation_task import GenerationTask
from app.models.generation_point import GenerationPoint from app.models.generation_point import GenerationPoint
@ -14,10 +13,22 @@ from app.services.ref_parser import parse_reference_file
from app.services.file_storage import get_storage_dir, RESULTS_DIR from app.services.file_storage import get_storage_dir, RESULTS_DIR
async def _generate_document(task_id: str) -> None: def _create_db_session() -> async_sessionmaker[AsyncSession]:
settings = get_settings() settings = get_settings()
engine = create_async_engine(
settings.DATABASE_URL,
echo=False,
pool_size=5,
max_overflow=5,
pool_pre_ping=True,
)
return async_sessionmaker(engine, class_=AsyncSession, expire_on_commit=False)
async with async_session_factory() as db:
async def _generate_document(task_id: str) -> None:
session_factory = _create_db_session()
async with session_factory() as db:
task_result = await db.execute(select(GenerationTask).where(GenerationTask.id == task_id)) task_result = await db.execute(select(GenerationTask).where(GenerationTask.id == task_id))
task = task_result.scalar_one_or_none() task = task_result.scalar_one_or_none()
if not task: if not task:
@ -45,8 +56,7 @@ async def _generate_document(task_id: str) -> None:
points = points_result.scalars().all() points = points_result.scalars().all()
try: try:
total = len(points) for point in points:
for idx, point in enumerate(points):
ref_content = None ref_content = None
if point.ref_file_path: if point.ref_file_path:
ref_content = await parse_reference_file(point.ref_file_path) ref_content = await parse_reference_file(point.ref_file_path)
@ -72,7 +82,6 @@ async def _generate_document(task_id: str) -> None:
if 0 <= start < end <= len(html_content): if 0 <= start < end <= len(html_content):
html_content = html_content[:start] + ai_result + html_content[end:] html_content = html_content[:start] + ai_result + html_content[end:]
task.status = "processing"
await db.commit() await db.commit()
result_dir = get_storage_dir(RESULTS_DIR) result_dir = get_storage_dir(RESULTS_DIR)
@ -92,6 +101,8 @@ async def _generate_document(task_id: str) -> None:
task.finished_at = datetime.now(timezone.utc) task.finished_at = datetime.now(timezone.utc)
await db.commit() await db.commit()
await session_factory.engine.dispose()
@celery_app.task(bind=True, name="generate_document") @celery_app.task(bind=True, name="generate_document")
def generate_document(self, task_id: str) -> dict: def generate_document(self, task_id: str) -> dict: