doc-forge-reds/backend/app/tasks/generate.py

203 lines
7.2 KiB
Python

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_file
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 _get_selected_text(html_content: str, position: dict) -> str:
start = position.get("start", 0)
end = position.get("end", 0)
if 0 <= start < end <= len(html_content):
from html.parser import HTMLParser
class TextStripper(HTMLParser):
def __init__(self):
super().__init__()
self.text = ""
def handle_data(self, data):
self.text += data
stripper = TextStripper()
stripper.feed(html_content[start:end])
return stripper.text.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 = _get_selected_text(html_content, point.position)
if not selected_text:
continue
ref_content = None
if point.ref_file_path:
try:
ref_content = await parse_reference_file(point.ref_file_path)
except Exception:
pass
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()