diff --git a/backend/app/tasks/generate.py b/backend/app/tasks/generate.py index b474c3e..b3aeb10 100644 --- a/backend/app/tasks/generate.py +++ b/backend/app/tasks/generate.py @@ -1,9 +1,8 @@ import os -import uuid from datetime import datetime, timezone 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.core.database import async_session_factory from app.core.config import get_settings from app.models.generation_task import GenerationTask 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 -async def _generate_document(task_id: str) -> None: +def _create_db_session() -> async_sessionmaker[AsyncSession]: 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 = task_result.scalar_one_or_none() if not task: @@ -45,8 +56,7 @@ async def _generate_document(task_id: str) -> None: points = points_result.scalars().all() try: - total = len(points) - for idx, point in enumerate(points): + for point in points: ref_content = None if 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): html_content = html_content[:start] + ai_result + html_content[end:] - task.status = "processing" await db.commit() 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) await db.commit() + await session_factory.engine.dispose() + @celery_app.task(bind=True, name="generate_document") def generate_document(self, task_id: str) -> dict: