init: 初始化项目
This commit is contained in:
@@ -0,0 +1,25 @@
|
||||
from celery import Celery
|
||||
from app.core.config import get_settings
|
||||
|
||||
settings = get_settings()
|
||||
|
||||
celery_app = Celery(
|
||||
"doc_forge_reds",
|
||||
broker=settings.CELERY_BROKER_URL,
|
||||
backend=settings.CELERY_RESULT_BACKEND,
|
||||
)
|
||||
|
||||
celery_app.conf.update(
|
||||
task_serializer="json",
|
||||
accept_content=["json"],
|
||||
result_serializer="json",
|
||||
timezone="Asia/Shanghai",
|
||||
enable_utc=True,
|
||||
task_track_started=True,
|
||||
task_acks_late=True,
|
||||
worker_prefetch_multiplier=1,
|
||||
task_soft_time_limit=600,
|
||||
task_time_limit=900,
|
||||
)
|
||||
|
||||
celery_app.autodiscover_tasks(["app.tasks.generate"])
|
||||
@@ -0,0 +1,225 @@
|
||||
import asyncio
|
||||
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:
|
||||
# 先收集所有 AI 调用参数
|
||||
point_data = []
|
||||
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.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,
|
||||
}
|
||||
|
||||
point_data.append({
|
||||
"point": point,
|
||||
"selected_text": selected_text,
|
||||
"model_config": model_config,
|
||||
"ref_content": ref_content,
|
||||
})
|
||||
|
||||
# 并发调用所有 AI 模型
|
||||
async def _call_one(pd):
|
||||
try:
|
||||
return await call_ai_model(pd["model_config"], pd["point"].prompt, pd["ref_content"])
|
||||
except Exception as e:
|
||||
return f"[生成失败: {e}]"
|
||||
|
||||
coros = [_call_one(pd) for pd in point_data]
|
||||
ai_results = await asyncio.gather(*coros)
|
||||
|
||||
# 按顺序替换文本
|
||||
for pd, ai_result in zip(point_data, ai_results):
|
||||
_replace_text_in_docx(doc, pd["selected_text"], str(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()
|
||||
Reference in New Issue
Block a user