diff --git a/backend/app/tasks/generate.py b/backend/app/tasks/generate.py index aa368a9..89afd59 100644 --- a/backend/app/tasks/generate.py +++ b/backend/app/tasks/generate.py @@ -1,3 +1,4 @@ +import asyncio import os from io import BytesIO from datetime import datetime, timezone @@ -145,15 +146,15 @@ async def _generate_document(task_id: str) -> None: 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.need_ref_file and point.ref_file_path: - ref_content = await parse_reference_files(point.ref_file_path) - elif point.ref_file_path: + if point.ref_file_path: ref_content = await parse_reference_files(point.ref_file_path) model_config = {"provider": "custom", "endpoint": "", "api_key": "", "extra_params": {}} @@ -169,10 +170,28 @@ async def _generate_document(task_id: str) -> None: "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) + point_data.append({ + "point": point, + "selected_text": selected_text, + "model_config": model_config, + "ref_content": ref_content, + }) - await db.commit() + # 并发调用所有 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") diff --git a/web/src/pages/TemplateEditor.tsx b/web/src/pages/TemplateEditor.tsx index 1acbbac..2da8808 100644 --- a/web/src/pages/TemplateEditor.tsx +++ b/web/src/pages/TemplateEditor.tsx @@ -16,6 +16,7 @@ import type { GenerationPoint, AIModel } from "../types"; import * as templateApi from "../api/templates"; import * as pointApi from "../api/generationPoints"; import * as modelApi from "../api/models"; +import request from "../api/request"; import GenerateModal from "../components/GenerateModal"; const { Sider, Content } = Layout; @@ -315,25 +316,32 @@ export default function TemplateEditor() { {showSelectionBtn && ( - +
+ +
)} { + handleEditorInit(null, editor); + if (htmlContent) editor.setContent(htmlContent); + }} init={{ height: "100%", menubar: true, @@ -390,7 +398,23 @@ export default function TemplateEditor() { type="text" size="small" icon={} - onClick={() => uploadRefForPoints([point.id])} + onClick={async () => { + const input = document.createElement("input"); + input.type = "file"; + input.multiple = true; + input.accept = ".txt,.docx,.pdf"; + input.onchange = async (e) => { + const files = (e.target as HTMLInputElement).files; + if (!files?.length) return; + const fd = new FormData(); + Array.from(files).forEach((f) => fd.append("ref_files", f)); + await request.post(`/generation-points/${point.id}/upload-ref`, fd); + message.success("上传成功"); + const updated = await pointApi.listGenerationPoints(id!); + setPoints(updated); + }; + input.click(); + }} /> )}