import os from io import BytesIO from datetime import datetime, timezone from sqlalchemy import select from sqlalchemy.ext.asyncio import create_async_engine, async_sessionmaker, AsyncSession from docx import Document from app.tasks.celery_app import celery_app from app.core.config import get_settings from app.models.generation_task import GenerationTask from app.models.generation_point import GenerationPoint from app.models.template import Template from app.services.ai_adapter import call_ai_model from app.services.ref_parser import parse_reference_files from app.services.file_storage import get_storage_dir, get_file_content, RESULTS_DIR 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) def _strip_html(html: str) -> str: from html.parser import HTMLParser class Stripper(HTMLParser): def __init__(self): super().__init__() self.text = "" def handle_data(self, data): self.text += data s = Stripper() s.feed(html) return s.text def _get_selected_text(html_content: str, position: dict) -> str: start = position.get("start", 0) end = position.get("end", 0) plain_text = _strip_html(html_content) if 0 <= start < end <= len(plain_text): return plain_text[start:end].strip() return "" def _replace_text_in_docx(doc: Document, old_text: str, new_text: str) -> bool: if not old_text: return False for paragraph in doc.paragraphs: if old_text in paragraph.text: inline = paragraph.runs for run in inline: if old_text in run.text: run.text = run.text.replace(old_text, new_text) return True full_text = "".join(r.text for r in inline) if old_text in full_text: remaining = old_text for run in inline: if not remaining: break if remaining.startswith(run.text): remaining = remaining[len(run.text):] elif run.text in remaining: idx = remaining.find(run.text) if idx >= 0: remaining = remaining[:idx] + remaining[idx + len(run.text):] if remaining.startswith(run.text): remaining = remaining[len(run.text):] if not remaining: chunk_parts = new_text for run in inline: if chunk_parts: chunk_parts = chunk_parts[len(run.text):] first_run = inline[0] first_run.text = new_text for run in inline[1:]: run.text = "" return True for table in doc.tables: for row in table.rows: for cell in row.cells: for paragraph in cell.paragraphs: if old_text in paragraph.text: for run in paragraph.runs: if old_text in run.text: run.text = run.text.replace(old_text, new_text) return True return False 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: return task.status = "processing" await db.commit() template_result = await db.execute(select(Template).where(Template.id == task.template_id)) template = template_result.scalar_one_or_none() if not template: task.status = "failed" task.error_msg = "模板不存在" task.finished_at = datetime.now(timezone.utc) await db.commit() return if not os.path.exists(template.file_path): task.status = "failed" task.error_msg = f"原始文件不存在: {template.file_path}" task.finished_at = datetime.now(timezone.utc) await db.commit() return docx_content = await get_file_content(template.file_path) doc = Document(BytesIO(docx_content)) html_content = template.html_content or "" points_result = await db.execute( select(GenerationPoint) .where(GenerationPoint.template_id == task.template_id) .order_by(GenerationPoint.order.asc()) ) points = points_result.scalars().all() try: for point in points: selected_text = point.selected_text or _get_selected_text(html_content, point.position) if not selected_text: continue ref_content = "" if point.need_ref_file and point.ref_file_path: ref_content = await parse_reference_files(point.ref_file_path) elif point.ref_file_path: ref_content = await parse_reference_files(point.ref_file_path) model_config = {"provider": "custom", "endpoint": "", "api_key": "", "extra_params": {}} if point.model_id: from app.models.ai_model import AIModel model_result = await db.execute(select(AIModel).where(AIModel.id == point.model_id)) model = model_result.scalar_one_or_none() if model: model_config = { "provider": model.provider, "endpoint": model.endpoint, "api_key": model.api_key, "extra_params": model.extra_params, } ai_result = await call_ai_model(model_config, point.prompt, ref_content) _replace_text_in_docx(doc, selected_text, ai_result) await db.commit() result_dir = get_storage_dir(RESULTS_DIR) result_path = os.path.join(result_dir, f"{task_id}.docx") doc.save(result_path) task.status = "done" task.result_file_path = result_path task.finished_at = datetime.now(timezone.utc) await db.commit() except Exception as e: task.status = "failed" task.error_msg = str(e) 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: import asyncio loop = asyncio.new_event_loop() asyncio.set_event_loop(loop) try: loop.run_until_complete(_generate_document(task_id)) return {"status": "done", "task_id": task_id} except Exception as e: return {"status": "failed", "task_id": task_id, "error": str(e)} finally: loop.close()