feat: 生成点表单重构 - 参考文件改为开关 + 备注字段

- 模型新增 need_ref_file(bool) + remark(text) 字段
- 生成点创建表单移除直接上传,改为「需要参考文件」开关 + 备注文本框
- 测试按钮:检测 need_ref → 弹出文件选择 → 上传后提取文本 + 调AI
- 生成任务:检查 ref_file_path,有则提取文本并入 prompt
- 前后端完整联调,TypeScript 零错误
This commit is contained in:
zwt13703 2026-07-06 18:27:30 +08:00
parent c11ea22e88
commit 57cdb01ec8
10 changed files with 118 additions and 13 deletions

View File

@ -0,0 +1,8 @@
{
"hash": "1e9afdbb",
"configHash": "ac7453e4",
"lockfileHash": "e3b0c442",
"browserHash": "29840d3b",
"optimized": {},
"chunks": {}
}

View File

@ -0,0 +1,3 @@
{
"type": "module"
}

View File

@ -0,0 +1,29 @@
"""Add need_ref_file and remark to generation_points
Revision ID: b4b91d4466e4
Revises: 011bde04d871
Create Date: 2026-07-06 18:27:03.054325
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
revision: str = 'b4b91d4466e4'
down_revision: Union[str, None] = '011bde04d871'
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def upgrade() -> None:
# ### commands auto generated by Alembic - please adjust! ###
op.add_column('generation_points', sa.Column('need_ref_file', sa.Boolean(), nullable=False))
op.add_column('generation_points', sa.Column('remark', sa.Text(), nullable=True))
# ### end Alembic commands ###
def downgrade() -> None:
# ### commands auto generated by Alembic - please adjust! ###
op.drop_column('generation_points', 'remark')
op.drop_column('generation_points', 'need_ref_file')
# ### end Alembic commands ###

View File

@ -26,6 +26,8 @@ async def create_generation_point(
model_id: uuid.UUID | None = Form(None), model_id: uuid.UUID | None = Form(None),
order: int = Form(0), order: int = Form(0),
selected_text: str | None = Form(None), selected_text: str | None = Form(None),
need_ref_file: bool = Form(False),
remark: str | None = Form(None),
ref_file: UploadFile | None = File(None), ref_file: UploadFile | None = File(None),
db: AsyncSession = Depends(get_db), db: AsyncSession = Depends(get_db),
): ):
@ -53,6 +55,8 @@ async def create_generation_point(
order=order, order=order,
ref_file_path=ref_file_path, ref_file_path=ref_file_path,
selected_text=selected_text, selected_text=selected_text,
need_ref_file=need_ref_file,
remark=remark,
) )
db.add(point) db.add(point)
await db.flush() await db.flush()

View File

@ -1,18 +1,19 @@
import os import os
import uuid import uuid
from datetime import datetime, timezone from datetime import datetime, timezone
from fastapi import APIRouter, Depends, HTTPException, Query from fastapi import APIRouter, Depends, HTTPException, Query, UploadFile, File
from fastapi.responses import Response from fastapi.responses import Response
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy import select from sqlalchemy import select
from app.core.database import get_db from app.core.database import get_db
from app.core.security_middleware import validate_file_extension
from app.models.generation_task import GenerationTask from app.models.generation_task import GenerationTask
from app.models.generation_point import GenerationPoint from app.models.generation_point import GenerationPoint
from app.models.template import Template from app.models.template import Template
from app.models.ai_model import AIModel from app.models.ai_model import AIModel
from app.schemas.generation_task import TaskResponse, GenerateResponse, SingleTestRequest, SingleTestResponse from app.schemas.generation_task import TaskResponse, GenerateResponse, SingleTestRequest, SingleTestResponse
from app.services.document_processor import docx_to_pdf_bytes from app.services.document_processor import docx_to_pdf_bytes
from app.services.file_storage import get_file_content from app.services.file_storage import get_file_content, save_upload, REF_FILES_DIR
from app.services.ai_adapter import call_ai_model from app.services.ai_adapter import call_ai_model
from app.services.ref_parser import parse_reference_file from app.services.ref_parser import parse_reference_file
from app.tasks.generate import generate_document from app.tasks.generate import generate_document
@ -123,6 +124,7 @@ async def cancel_task(task_id: str, db: AsyncSession = Depends(get_db)):
@router.post("/generation-points/{point_id}/test", response_model=SingleTestResponse) @router.post("/generation-points/{point_id}/test", response_model=SingleTestResponse)
async def test_single_point( async def test_single_point(
point_id: str, point_id: str,
ref_file: UploadFile | None = File(None),
db: AsyncSession = Depends(get_db), db: AsyncSession = Depends(get_db),
): ):
point_result = await db.execute(select(GenerationPoint).where(GenerationPoint.id == point_id)) point_result = await db.execute(select(GenerationPoint).where(GenerationPoint.id == point_id))
@ -130,6 +132,15 @@ async def test_single_point(
if not point: if not point:
raise HTTPException(status_code=404, detail="生成点不存在") raise HTTPException(status_code=404, detail="生成点不存在")
if point.need_ref_file and not point.ref_file_path and not ref_file:
raise HTTPException(status_code=400, detail="此生成点需要上传参考文件")
if ref_file and ref_file.filename:
validate_file_extension(ref_file.filename)
tmp_path = await save_upload(ref_file, REF_FILES_DIR)
else:
tmp_path = point.ref_file_path
model_config = {"provider": "custom", "endpoint": "", "api_key": "", "extra_params": {}} model_config = {"provider": "custom", "endpoint": "", "api_key": "", "extra_params": {}}
if point.model_id: if point.model_id:
model_result = await db.execute(select(AIModel).where(AIModel.id == point.model_id)) model_result = await db.execute(select(AIModel).where(AIModel.id == point.model_id))
@ -143,8 +154,8 @@ async def test_single_point(
} }
ref_content = None ref_content = None
if point.ref_file_path: if tmp_path:
ref_content = await parse_reference_file(point.ref_file_path) ref_content = await parse_reference_file(tmp_path)
try: try:
ai_result = await call_ai_model(model_config, point.prompt, ref_content) ai_result = await call_ai_model(model_config, point.prompt, ref_content)

View File

@ -18,6 +18,8 @@ class GenerationPoint(Base, UUIDMixin, TimestampMixin):
UUID(as_uuid=True), ForeignKey("models.id", ondelete="SET NULL"), nullable=True UUID(as_uuid=True), ForeignKey("models.id", ondelete="SET NULL"), nullable=True
) )
ref_file_path: Mapped[str | None] = mapped_column(String(500), nullable=True) ref_file_path: Mapped[str | None] = mapped_column(String(500), nullable=True)
need_ref_file: Mapped[bool] = mapped_column(default=False, nullable=False)
remark: Mapped[str | None] = mapped_column(Text, nullable=True)
order: Mapped[int] = mapped_column(Integer, default=0, nullable=False) order: Mapped[int] = mapped_column(Integer, default=0, nullable=False)
selected_text: Mapped[str | None] = mapped_column(Text, nullable=True) selected_text: Mapped[str | None] = mapped_column(Text, nullable=True)

View File

@ -10,6 +10,8 @@ class GenerationPointCreate(BaseModel):
model_id: uuid.UUID | None = None model_id: uuid.UUID | None = None
order: int = 0 order: int = 0
selected_text: str | None = None selected_text: str | None = None
need_ref_file: bool = False
remark: str | None = None
class GenerationPointUpdate(BaseModel): class GenerationPointUpdate(BaseModel):
@ -18,6 +20,8 @@ class GenerationPointUpdate(BaseModel):
model_id: uuid.UUID | None = None model_id: uuid.UUID | None = None
order: int | None = None order: int | None = None
selected_text: str | None = None selected_text: str | None = None
need_ref_file: bool | None = None
remark: str | None = None
class GenerationPointResponse(BaseModel): class GenerationPointResponse(BaseModel):
@ -27,6 +31,8 @@ class GenerationPointResponse(BaseModel):
prompt: str prompt: str
model_id: uuid.UUID | None = None model_id: uuid.UUID | None = None
ref_file_path: str | None = None ref_file_path: str | None = None
need_ref_file: bool = False
remark: str | None = None
order: int order: int
selected_text: str | None = None selected_text: str | None = None
created_at: datetime created_at: datetime

View File

@ -15,6 +15,8 @@ export async function createGenerationPoint(body: {
model_id?: string; model_id?: string;
order?: number; order?: number;
selected_text?: string; selected_text?: string;
need_ref_file?: boolean;
remark?: string;
ref_file?: File; ref_file?: File;
}) { }) {
const formData = new FormData(); const formData = new FormData();
@ -24,6 +26,8 @@ export async function createGenerationPoint(body: {
if (body.model_id) formData.append("model_id", body.model_id); if (body.model_id) formData.append("model_id", body.model_id);
if (body.order !== undefined) formData.append("order", String(body.order)); if (body.order !== undefined) formData.append("order", String(body.order));
if (body.selected_text) formData.append("selected_text", body.selected_text); if (body.selected_text) formData.append("selected_text", body.selected_text);
if (body.need_ref_file) formData.append("need_ref_file", "true");
if (body.remark) formData.append("remark", body.remark);
if (body.ref_file) formData.append("ref_file", body.ref_file); if (body.ref_file) formData.append("ref_file", body.ref_file);
const { data } = await request.post<GenerationPoint>("/generation-points", formData); const { data } = await request.post<GenerationPoint>("/generation-points", formData);
return data; return data;

View File

@ -1,8 +1,8 @@
import { useEffect, useState, useRef, useCallback } from "react"; import { useEffect, useState, useRef, useCallback } from "react";
import { useParams, useNavigate } from "react-router-dom"; import { useParams, useNavigate } from "react-router-dom";
import { import {
Layout, Button, Modal, Form, Input, Select, Upload, Layout, Button, Modal, Form, Input, Select,
Space, message, Spin, Progress, Popconfirm, Card, Collapse, Space, message, Spin, Progress, Popconfirm, Card, Switch,
} from "antd"; } from "antd";
import { import {
PlusOutlined, DeleteOutlined, EditOutlined, SaveOutlined, PlusOutlined, DeleteOutlined, EditOutlined, SaveOutlined,
@ -128,6 +128,8 @@ export default function TemplateEditor() {
pointForm.setFieldsValue({ pointForm.setFieldsValue({
prompt: point.prompt, prompt: point.prompt,
model_id: point.model_id, model_id: point.model_id,
need_ref_file: point.need_ref_file,
remark: point.remark,
}); });
setPointModalOpen(true); setPointModalOpen(true);
}; };
@ -146,6 +148,8 @@ export default function TemplateEditor() {
const updateData: Record<string, unknown> = {}; const updateData: Record<string, unknown> = {};
if (values.prompt) updateData.prompt = values.prompt; if (values.prompt) updateData.prompt = values.prompt;
if (values.model_id) updateData.model_id = values.model_id; if (values.model_id) updateData.model_id = values.model_id;
updateData.need_ref_file = values.need_ref_file || false;
if (values.remark !== undefined) updateData.remark = values.remark;
await pointApi.updateGenerationPoint(editingPoint.id, updateData); await pointApi.updateGenerationPoint(editingPoint.id, updateData);
message.success("更新成功"); message.success("更新成功");
} else if (position) { } else if (position) {
@ -156,6 +160,9 @@ export default function TemplateEditor() {
model_id: values.model_id, model_id: values.model_id,
order: points.length, order: points.length,
selected_text: selectionRef.current?.text, selected_text: selectionRef.current?.text,
need_ref_file: values.need_ref_file || false,
remark: values.remark || "",
});
ref_file: refFile, ref_file: refFile,
}); });
message.success("创建成功"); message.success("创建成功");
@ -213,6 +220,35 @@ export default function TemplateEditor() {
}; };
const handleTestPoint = async (pointId: string) => { const handleTestPoint = async (pointId: string) => {
const point = points.find((p) => p.id === pointId);
if (!point) return;
if (point.need_ref_file) {
const fileInput = document.createElement("input");
fileInput.type = "file";
fileInput.accept = ".txt,.docx,.pdf";
fileInput.onchange = async (e: Event) => {
const file = (e.target as HTMLInputElement).files?.[0];
if (!file) return;
try {
message.loading({ content: "测试中...", key: "test" });
const formData = new FormData();
formData.append("ref_file", file);
const { default: request } = await import("../api/request");
const { data } = await request.post<{ result: string }>(
`/generation-points/${pointId}/test`,
formData
);
message.success({ content: "测试完成", key: "test" });
Modal.info({ title: "AI 生成结果", content: data.result, width: 600 });
} catch {
message.error({ content: "测试失败", key: "test" });
}
};
fileInput.click();
return;
}
try { try {
message.loading({ content: "测试中...", key: "test" }); message.loading({ content: "测试中...", key: "test" });
const result = await pointApi.testGenerationPoint(pointId); const result = await pointApi.testGenerationPoint(pointId);
@ -473,13 +509,12 @@ export default function TemplateEditor() {
options={models.map((m) => ({ label: m.name, value: m.id }))} options={models.map((m) => ({ label: m.name, value: m.id }))}
/> />
</Form.Item> </Form.Item>
{!editingPoint && ( <Form.Item name="need_ref_file" label="需要参考文件" valuePropName="checked">
<Form.Item name="ref_file" label="参考文件" valuePropName="fileList" getValueFromEvent={(e) => e.fileList}> <Switch />
<Upload maxCount={1} beforeUpload={() => false}> </Form.Item>
<Button icon={<PlusOutlined />}></Button> <Form.Item name="remark" label="备注">
</Upload> <Input.TextArea rows={2} placeholder="补充说明(可选)" />
</Form.Item> </Form.Item>
)}
</Form> </Form>
</Modal> </Modal>
</div> </div>

View File

@ -34,7 +34,10 @@ export interface GenerationPoint {
prompt: string; prompt: string;
model_id: string | null; model_id: string | null;
ref_file_path: string | null; ref_file_path: string | null;
need_ref_file: boolean;
remark: string | null;
order: number; order: number;
selected_text: string | null;
created_at: string; created_at: string;
updated_at: string; updated_at: string;
} }