diff --git a/backend/app/api/router.py b/backend/app/api/router.py index 84c67ee..dd46043 100644 --- a/backend/app/api/router.py +++ b/backend/app/api/router.py @@ -11,4 +11,4 @@ api_router = APIRouter(prefix=settings.API_V1_PREFIX) api_router.include_router(models_router) api_router.include_router(templates_router) api_router.include_router(generation_points_router) -api_router.include_router(tasks_router) +api_router.include_router(tasks_router, prefix="") diff --git a/backend/app/api/tasks.py b/backend/app/api/tasks.py index 4fe82c7..1ecf604 100644 --- a/backend/app/api/tasks.py +++ b/backend/app/api/tasks.py @@ -1,5 +1,153 @@ -from fastapi import APIRouter +import os +import uuid +from datetime import datetime, timezone +from fastapi import APIRouter, Depends, HTTPException, Query +from fastapi.responses import Response +from sqlalchemy.ext.asyncio import AsyncSession +from sqlalchemy import select +from app.core.database import get_db +from app.models.generation_task import GenerationTask +from app.models.generation_point import GenerationPoint +from app.models.template import Template +from app.models.ai_model import AIModel +from app.schemas.generation_task import TaskResponse, GenerateResponse, SingleTestRequest, SingleTestResponse +from app.services.document_processor import docx_to_pdf_bytes +from app.services.file_storage import get_file_content +from app.services.ai_adapter import call_ai_model +from app.services.ref_parser import parse_reference_file +from app.tasks.generate import generate_document -router = APIRouter(prefix="/tasks", tags=["生成任务"]) +router = APIRouter() -# 任务管理 API 将在阶段三实现 + +@router.post("/templates/{template_id}/generate", response_model=GenerateResponse) +async def trigger_generation(template_id: str, db: AsyncSession = Depends(get_db)): + template_result = await db.execute(select(Template).where(Template.id == template_id)) + if not template_result.scalar_one_or_none(): + raise HTTPException(status_code=404, detail="模板不存在") + + points_result = await db.execute( + select(GenerationPoint).where(GenerationPoint.template_id == template_id) + ) + points = points_result.scalars().all() + if not points: + raise HTTPException(status_code=400, detail="模板没有生成点") + + task = GenerationTask( + template_id=uuid.UUID(template_id), + status="pending", + created_at=datetime.now(timezone.utc), + ) + db.add(task) + await db.flush() + + celery_task = generate_document.delay(str(task.id)) + task.celery_task_id = celery_task.id + await db.commit() + + return GenerateResponse(task_id=task.id, status="pending") + + +@router.get("/tasks/{task_id}", response_model=TaskResponse) +async def get_task_status(task_id: str, db: AsyncSession = Depends(get_db)): + result = await db.execute(select(GenerationTask).where(GenerationTask.id == task_id)) + task = result.scalar_one_or_none() + if not task: + raise HTTPException(status_code=404, detail="任务不存在") + return task + + +@router.get("/tasks/{task_id}/download") +async def download_task_result( + task_id: str, + format: str = Query("docx", pattern="^(docx|pdf)$"), + db: AsyncSession = Depends(get_db), +): + result = await db.execute(select(GenerationTask).where(GenerationTask.id == task_id)) + task = result.scalar_one_or_none() + if not task: + raise HTTPException(status_code=404, detail="任务不存在") + + if task.status != "done" or not task.result_file_path: + raise HTTPException(status_code=400, detail="任务未完成或无结果文件") + + file_content = await get_file_content(task.result_file_path) + + if format == "pdf": + file_content = docx_to_pdf_bytes(file_content) + media_type = "application/pdf" + filename = f"result_{task_id}.pdf" + else: + media_type = "application/vnd.openxmlformats-officedocument.wordprocessingml.document" + filename = f"result_{task_id}.docx" + + return Response( + content=file_content, + media_type=media_type, + headers={"Content-Disposition": f'attachment; filename="{filename}"'}, + ) + + +@router.get("/tasks") +async def list_tasks( + skip: int = Query(0, ge=0), + limit: int = Query(20, ge=1, le=100), + db: AsyncSession = Depends(get_db), +): + query = select(GenerationTask).offset(skip).limit(limit).order_by(GenerationTask.created_at.desc()) + result = await db.execute(query) + return result.scalars().all() + + +@router.post("/tasks/{task_id}/cancel") +async def cancel_task(task_id: str, db: AsyncSession = Depends(get_db)): + result = await db.execute(select(GenerationTask).where(GenerationTask.id == task_id)) + task = result.scalar_one_or_none() + if not task: + raise HTTPException(status_code=404, detail="任务不存在") + + if task.status not in ("pending", "processing"): + raise HTTPException(status_code=400, detail="任务无法取消") + + if task.celery_task_id: + from app.tasks.celery_app import celery_app + celery_app.control.revoke(task.celery_task_id, terminate=True) + + task.status = "failed" + task.error_msg = "用户取消" + task.finished_at = datetime.now(timezone.utc) + await db.commit() + return {"detail": "任务已取消"} + + +@router.post("/generation-points/{point_id}/test", response_model=SingleTestResponse) +async def test_single_point( + point_id: str, + db: AsyncSession = Depends(get_db), +): + point_result = await db.execute(select(GenerationPoint).where(GenerationPoint.id == point_id)) + point = point_result.scalar_one_or_none() + if not point: + raise HTTPException(status_code=404, detail="生成点不存在") + + model_config = {"provider": "custom", "endpoint": "", "api_key": "", "extra_params": {}} + if point.model_id: + 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, + } + + ref_content = None + if point.ref_file_path: + ref_content = await parse_reference_file(point.ref_file_path) + + try: + ai_result = await call_ai_model(model_config, point.prompt, ref_content) + return SingleTestResponse(result=ai_result) + except Exception as e: + raise HTTPException(status_code=500, detail=f"AI 调用失败: {str(e)}") diff --git a/backend/app/services/ai_adapter.py b/backend/app/services/ai_adapter.py new file mode 100644 index 0000000..2e8cf38 --- /dev/null +++ b/backend/app/services/ai_adapter.py @@ -0,0 +1,163 @@ +import json +from abc import ABC, abstractmethod +import httpx +from app.core.security import decrypt_api_key + +PROVIDER_OPENAI = "openai" +PROVIDER_AZURE = "azure" +PROVIDER_CUSTOM = "custom" + + +class AIAdapter(ABC): + @abstractmethod + def build_request(self, model_config: dict, prompt: str, ref_content: str | None = None) -> dict: + pass + + @abstractmethod + def parse_response(self, response_data: dict) -> str: + pass + + @property + @abstractmethod + def provider(self) -> str: + pass + + +class OpenAIAdapter(AIAdapter): + @property + def provider(self) -> str: + return PROVIDER_OPENAI + + def build_request(self, model_config: dict, prompt: str, ref_content: str | None = None) -> dict: + extra = model_config.get("extra_params", {}) + temperature = extra.get("temperature", 0.7) + max_tokens = extra.get("max_tokens", 2000) + + messages = [{"role": "system", "content": "你是一个专业的文档内容生成助手。"}] + user_content = prompt + if ref_content: + user_content = f"参考以下内容:\n{ref_content}\n\n任务:{prompt}" + messages.append({"role": "user", "content": user_content}) + + return { + "url": model_config["endpoint"], + "headers": { + "Authorization": f"Bearer {decrypt_api_key(model_config['api_key'])}", + "Content-Type": "application/json", + }, + "json": { + "model": extra.get("model", "gpt-4"), + "messages": messages, + "temperature": temperature, + "max_tokens": max_tokens, + }, + } + + def parse_response(self, response_data: dict) -> str: + return response_data["choices"][0]["message"]["content"] + + +class AzureAdapter(AIAdapter): + @property + def provider(self) -> str: + return PROVIDER_AZURE + + def build_request(self, model_config: dict, prompt: str, ref_content: str | None = None) -> dict: + extra = model_config.get("extra_params", {}) + temperature = extra.get("temperature", 0.7) + max_tokens = extra.get("max_tokens", 2000) + + messages = [{"role": "system", "content": "你是一个专业的文档内容生成助手。"}] + user_content = prompt + if ref_content: + user_content = f"参考以下内容:\n{ref_content}\n\n任务:{prompt}" + messages.append({"role": "user", "content": user_content}) + + api_version = extra.get("api_version", "2024-02-15-preview") + endpoint = model_config["endpoint"] + url = f"{endpoint}?api-version={api_version}" + + return { + "url": url, + "headers": { + "api-key": decrypt_api_key(model_config["api_key"]), + "Content-Type": "application/json", + }, + "json": { + "messages": messages, + "temperature": temperature, + "max_tokens": max_tokens, + }, + } + + def parse_response(self, response_data: dict) -> str: + return response_data["choices"][0]["message"]["content"] + + +class CustomAdapter(AIAdapter): + @property + def provider(self) -> str: + return PROVIDER_CUSTOM + + def build_request(self, model_config: dict, prompt: str, ref_content: str | None = None) -> dict: + extra = model_config.get("extra_params", {}) + user_content = prompt + if ref_content: + user_content = f"参考以下内容:\n{ref_content}\n\n任务:{prompt}" + + body = { + "prompt": user_content, + "max_tokens": extra.get("max_tokens", 2000), + "temperature": extra.get("temperature", 0.7), + } + body.update({k: v for k, v in extra.items() if k not in ("max_tokens", "temperature")}) + + return { + "url": model_config["endpoint"], + "headers": { + "Authorization": f"Bearer {decrypt_api_key(model_config['api_key'])}", + "Content-Type": "application/json", + }, + "json": body, + } + + def parse_response(self, response_data: dict) -> str: + if "choices" in response_data: + return response_data["choices"][0]["message"]["content"] + if "response" in response_data: + return response_data["response"] + if "content" in response_data: + return response_data["content"] + if "text" in response_data: + return response_data["text"] + return json.dumps(response_data) + + +_adapters: dict[str, AIAdapter] = { + PROVIDER_OPENAI: OpenAIAdapter(), + PROVIDER_AZURE: AzureAdapter(), + PROVIDER_CUSTOM: CustomAdapter(), +} + + +def get_adapter(provider: str) -> AIAdapter: + adapter = _adapters.get(provider) + if not adapter: + raise ValueError(f"不支持的供应商: {provider}") + return adapter + + +async def call_ai_model(model_config: dict, prompt: str, ref_content: str | None = None) -> str: + adapter = get_adapter(model_config["provider"]) + request = adapter.build_request(model_config, prompt, ref_content) + + timeout = model_config.get("extra_params", {}).get("timeout", 120) + + async with httpx.AsyncClient(timeout=timeout) as client: + response = await client.post( + request["url"], + headers=request["headers"], + json=request["json"], + ) + response.raise_for_status() + return adapter.parse_response(response.json()) diff --git a/backend/app/services/ref_parser.py b/backend/app/services/ref_parser.py new file mode 100644 index 0000000..4488114 --- /dev/null +++ b/backend/app/services/ref_parser.py @@ -0,0 +1,25 @@ +import os +from docx import Document +from PyPDF2 import PdfReader +from app.services.file_storage import get_file_content + + +async def parse_reference_file(file_path: str) -> str: + ext = os.path.splitext(file_path)[1].lower() + + content = await get_file_content(file_path) + + if ext == ".txt": + return content.decode("utf-8", errors="ignore") + + if ext == ".docx": + from io import BytesIO + doc = Document(BytesIO(content)) + return "\n".join(p.text for p in doc.paragraphs if p.text.strip()) + + if ext == ".pdf": + from io import BytesIO + reader = PdfReader(BytesIO(content)) + return "\n".join(page.extract_text() or "" for page in reader.pages) + + raise ValueError(f"不支持的文件格式: {ext}") diff --git a/backend/app/tasks/__init__.py b/backend/app/tasks/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/backend/app/tasks/celery_app.py b/backend/app/tasks/celery_app.py new file mode 100644 index 0000000..76f876f --- /dev/null +++ b/backend/app/tasks/celery_app.py @@ -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"]) diff --git a/backend/app/tasks/generate.py b/backend/app/tasks/generate.py new file mode 100644 index 0000000..b474c3e --- /dev/null +++ b/backend/app/tasks/generate.py @@ -0,0 +1,107 @@ +import os +import uuid +from datetime import datetime, timezone +from sqlalchemy import select +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 +from app.models.template import Template +from app.services.document_processor import html_to_docx_bytes +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, RESULTS_DIR + + +async def _generate_document(task_id: str) -> None: + settings = get_settings() + + async with async_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 + + 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: + total = len(points) + for idx, point in enumerate(points): + ref_content = None + if point.ref_file_path: + ref_content = await parse_reference_file(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) + + position = point.position + start = position.get("start", 0) + end = position.get("end", 0) + 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) + result_path = os.path.join(result_dir, f"{task_id}.docx") + docx_bytes = html_to_docx_bytes(html_content) + with open(result_path, "wb") as f: + f.write(docx_bytes) + + 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() + + +@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() diff --git a/docs/tasks/task_detail_2026_07_06.md b/docs/tasks/task_detail_2026_07_06.md new file mode 100644 index 0000000..b78af29 --- /dev/null +++ b/docs/tasks/task_detail_2026_07_06.md @@ -0,0 +1,24 @@ +# 任务执行摘要 + +## 会话 ID: 1 +- [2026-07-06 16:00] +- **执行原因**: 用户要求按任务拆解清单逐步实施项目开发,先完成阶段一和阶段二 +- **执行过程**: + 1. 创建项目目录结构(backend/, web/, storage/)。 + 2. 编写 docker-compose.yml(Postgres 18.3 + Redis 7)、.env.example、requirements.txt。 + 3. 实现 backend/app/core/ 核心模块:config(pydantic-settings)、database(SQLAlchemy async engine + session)、security(Fernet 加密)。 + 4. 创建 5 个 ORM 模型:AIModel、Template、GenerationPoint、GenerationTask、SystemConfig + Base/TimestampMixin/UUIDMixin。 + 5. 配置 Alembic(alembic.ini, env.py, script.py.mako)。 + 6. 编写 FastAPI 入口 main.py(CORS + health 端点)。 + 7. 实现文件存储服务 file_storage.py(上传/下载/删除)。 + 8. 初始化前端项目(Vite + React 18 + TypeScript + Ant Design + React Router),侧边栏布局 + 5 个占位页面。 + 9. 编写 4 组 Pydantic Schemas(ai_model, template, generation_point, generation_task)。 + 10. 实现文档处理服务 document_processor.py(Mammoth docx→HTML、python-docx HTML→docx、reportlab docx→PDF)。 + 11. 实现 AI 模型管理 API(CRUD + 启用/禁用切换,6 个端点)。 + 12. 实现模板管理 API(上传/获取HTML/更新HTML/下载docx+pdf,6 个端点)。 + 13. 实现生成点管理 API(CRUD + 参考文件上传 + 批量排序更新,5 个端点)。 + 14. 注册路由(api/router.py),更新 main.py,启动验证通过(/health 返回 OK)。 +- **执行结果**: + - 阶段一(基础设施):45 个文件,后端 5 张表模型导入正常,前端 TypeScript 编译零错误。 + - 阶段二(核心业务 API):17 个 API 端点全部注册,后端启动正常,`/health` 返回 200 OK。 + - Git 提交:2 次提交(b5fa343, aa84521)。