diff --git a/backend/.vite/deps/_metadata.json b/backend/.vite/deps/_metadata.json new file mode 100644 index 0000000..4c75da8 --- /dev/null +++ b/backend/.vite/deps/_metadata.json @@ -0,0 +1,8 @@ +{ + "hash": "1e9afdbb", + "configHash": "ac7453e4", + "lockfileHash": "e3b0c442", + "browserHash": "29840d3b", + "optimized": {}, + "chunks": {} +} \ No newline at end of file diff --git a/backend/.vite/deps/package.json b/backend/.vite/deps/package.json new file mode 100644 index 0000000..3dbc1ca --- /dev/null +++ b/backend/.vite/deps/package.json @@ -0,0 +1,3 @@ +{ + "type": "module" +} diff --git a/backend/alembic/versions/b4b91d4466e4_add_need_ref_file_and_remark_to_.py b/backend/alembic/versions/b4b91d4466e4_add_need_ref_file_and_remark_to_.py new file mode 100644 index 0000000..11c9937 --- /dev/null +++ b/backend/alembic/versions/b4b91d4466e4_add_need_ref_file_and_remark_to_.py @@ -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 ### diff --git a/backend/app/api/generation_points.py b/backend/app/api/generation_points.py index c60002b..714a486 100644 --- a/backend/app/api/generation_points.py +++ b/backend/app/api/generation_points.py @@ -26,6 +26,8 @@ async def create_generation_point( model_id: uuid.UUID | None = Form(None), order: int = Form(0), selected_text: str | None = Form(None), + need_ref_file: bool = Form(False), + remark: str | None = Form(None), ref_file: UploadFile | None = File(None), db: AsyncSession = Depends(get_db), ): @@ -53,6 +55,8 @@ async def create_generation_point( order=order, ref_file_path=ref_file_path, selected_text=selected_text, + need_ref_file=need_ref_file, + remark=remark, ) db.add(point) await db.flush() diff --git a/backend/app/api/tasks.py b/backend/app/api/tasks.py index 1ecf604..a485e6a 100644 --- a/backend/app/api/tasks.py +++ b/backend/app/api/tasks.py @@ -1,18 +1,19 @@ import os import uuid 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 sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy import select 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_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.file_storage import get_file_content, save_upload, REF_FILES_DIR 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 @@ -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) async def test_single_point( point_id: str, + ref_file: UploadFile | None = File(None), db: AsyncSession = Depends(get_db), ): point_result = await db.execute(select(GenerationPoint).where(GenerationPoint.id == point_id)) @@ -130,6 +132,15 @@ async def test_single_point( if not point: 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": {}} if 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 - if point.ref_file_path: - ref_content = await parse_reference_file(point.ref_file_path) + if tmp_path: + ref_content = await parse_reference_file(tmp_path) try: ai_result = await call_ai_model(model_config, point.prompt, ref_content) diff --git a/backend/app/models/generation_point.py b/backend/app/models/generation_point.py index 468b556..58cef53 100644 --- a/backend/app/models/generation_point.py +++ b/backend/app/models/generation_point.py @@ -18,6 +18,8 @@ class GenerationPoint(Base, UUIDMixin, TimestampMixin): UUID(as_uuid=True), ForeignKey("models.id", ondelete="SET NULL"), 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) selected_text: Mapped[str | None] = mapped_column(Text, nullable=True) diff --git a/backend/app/schemas/generation_point.py b/backend/app/schemas/generation_point.py index 11ca537..445e904 100644 --- a/backend/app/schemas/generation_point.py +++ b/backend/app/schemas/generation_point.py @@ -10,6 +10,8 @@ class GenerationPointCreate(BaseModel): model_id: uuid.UUID | None = None order: int = 0 selected_text: str | None = None + need_ref_file: bool = False + remark: str | None = None class GenerationPointUpdate(BaseModel): @@ -18,6 +20,8 @@ class GenerationPointUpdate(BaseModel): model_id: uuid.UUID | None = None order: int | None = None selected_text: str | None = None + need_ref_file: bool | None = None + remark: str | None = None class GenerationPointResponse(BaseModel): @@ -27,6 +31,8 @@ class GenerationPointResponse(BaseModel): prompt: str model_id: uuid.UUID | None = None ref_file_path: str | None = None + need_ref_file: bool = False + remark: str | None = None order: int selected_text: str | None = None created_at: datetime diff --git a/web/src/api/generationPoints.ts b/web/src/api/generationPoints.ts index 1a77c58..ffa0d92 100644 --- a/web/src/api/generationPoints.ts +++ b/web/src/api/generationPoints.ts @@ -15,6 +15,8 @@ export async function createGenerationPoint(body: { model_id?: string; order?: number; selected_text?: string; + need_ref_file?: boolean; + remark?: string; ref_file?: File; }) { 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.order !== undefined) formData.append("order", String(body.order)); 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); const { data } = await request.post("/generation-points", formData); return data; diff --git a/web/src/pages/TemplateEditor.tsx b/web/src/pages/TemplateEditor.tsx index abd2849..81c0b9b 100644 --- a/web/src/pages/TemplateEditor.tsx +++ b/web/src/pages/TemplateEditor.tsx @@ -1,8 +1,8 @@ import { useEffect, useState, useRef, useCallback } from "react"; import { useParams, useNavigate } from "react-router-dom"; import { - Layout, Button, Modal, Form, Input, Select, Upload, - Space, message, Spin, Progress, Popconfirm, Card, Collapse, + Layout, Button, Modal, Form, Input, Select, + Space, message, Spin, Progress, Popconfirm, Card, Switch, } from "antd"; import { PlusOutlined, DeleteOutlined, EditOutlined, SaveOutlined, @@ -128,6 +128,8 @@ export default function TemplateEditor() { pointForm.setFieldsValue({ prompt: point.prompt, model_id: point.model_id, + need_ref_file: point.need_ref_file, + remark: point.remark, }); setPointModalOpen(true); }; @@ -146,6 +148,8 @@ export default function TemplateEditor() { const updateData: Record = {}; if (values.prompt) updateData.prompt = values.prompt; 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); message.success("更新成功"); } else if (position) { @@ -156,6 +160,9 @@ export default function TemplateEditor() { model_id: values.model_id, order: points.length, selected_text: selectionRef.current?.text, + need_ref_file: values.need_ref_file || false, + remark: values.remark || "", + }); ref_file: refFile, }); message.success("创建成功"); @@ -213,6 +220,35 @@ export default function TemplateEditor() { }; 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 { message.loading({ content: "测试中...", key: "test" }); const result = await pointApi.testGenerationPoint(pointId); @@ -473,13 +509,12 @@ export default function TemplateEditor() { options={models.map((m) => ({ label: m.name, value: m.id }))} /> - {!editingPoint && ( - e.fileList}> - false}> - - - - )} + + + + + + diff --git a/web/src/types/index.ts b/web/src/types/index.ts index f10451b..d6233db 100644 --- a/web/src/types/index.ts +++ b/web/src/types/index.ts @@ -34,7 +34,10 @@ export interface GenerationPoint { prompt: string; model_id: string | null; ref_file_path: string | null; + need_ref_file: boolean; + remark: string | null; order: number; + selected_text: string | null; created_at: string; updated_at: string; }