feat: 阶段三 - 异步任务与 AI 集成

- AI 调用适配器:工厂模式,支持 openai/azure/custom 三种供应商
- 参考文件解析器:提取 txt/docx/pdf 文本内容
- Celery 配置与生成任务:模板加载→生成点排序→参考解析→AI调用→HTML插入→文档保存
- 任务管理 API:触发/状态查询/结果下载(docx+pdf)/取消/列表
- 单点测试 API:生成点预览不保存文档
- 23 个 API 端点全部注册
This commit is contained in:
zwt13703 2026-07-06 16:28:00 +08:00
parent aa84521449
commit afabb6446a
8 changed files with 496 additions and 4 deletions

View File

@ -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="")

View File

@ -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)}")

View File

@ -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())

View File

@ -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}")

View File

View File

@ -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"])

View File

@ -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()

View File

@ -0,0 +1,24 @@
# 任务执行摘要
## 会话 ID: 1
- [2026-07-06 16:00]
- **执行原因**: 用户要求按任务拆解清单逐步实施项目开发,先完成阶段一和阶段二
- **执行过程**:
1. 创建项目目录结构backend/, web/, storage/)。
2. 编写 docker-compose.ymlPostgres 18.3 + Redis 7、.env.example、requirements.txt。
3. 实现 backend/app/core/ 核心模块configpydantic-settings、databaseSQLAlchemy async engine + session、securityFernet 加密)。
4. 创建 5 个 ORM 模型AIModel、Template、GenerationPoint、GenerationTask、SystemConfig + Base/TimestampMixin/UUIDMixin。
5. 配置 Alembicalembic.ini, env.py, script.py.mako
6. 编写 FastAPI 入口 main.pyCORS + health 端点)。
7. 实现文件存储服务 file_storage.py上传/下载/删除)。
8. 初始化前端项目Vite + React 18 + TypeScript + Ant Design + React Router侧边栏布局 + 5 个占位页面。
9. 编写 4 组 Pydantic Schemasai_model, template, generation_point, generation_task
10. 实现文档处理服务 document_processor.pyMammoth docx→HTML、python-docx HTML→docx、reportlab docx→PDF
11. 实现 AI 模型管理 APICRUD + 启用/禁用切换6 个端点)。
12. 实现模板管理 API上传/获取HTML/更新HTML/下载docx+pdf6 个端点)。
13. 实现生成点管理 APICRUD + 参考文件上传 + 批量排序更新5 个端点)。
14. 注册路由api/router.py更新 main.py启动验证通过/health 返回 OK
- **执行结果**:
- 阶段一基础设施45 个文件,后端 5 张表模型导入正常,前端 TypeScript 编译零错误。
- 阶段二(核心业务 API17 个 API 端点全部注册,后端启动正常,`/health` 返回 200 OK。
- Git 提交2 次提交b5fa343, aa84521