init: 初始化项目

This commit is contained in:
zwt13703
2026-07-08 20:02:29 +08:00
parent 22590ae7b8
commit 1bb84df4ca
98 changed files with 8535 additions and 2 deletions
View File
+128
View File
@@ -0,0 +1,128 @@
import uuid
from fastapi import APIRouter, Depends, HTTPException, UploadFile, File, Form, Query
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_point import GenerationPoint
from app.models.template import Template
from app.models.ai_model import AIModel
from app.schemas.generation_point import (
GenerationPointCreate,
GenerationPointUpdate,
GenerationPointResponse,
BatchOrderUpdate,
)
from app.services.file_storage import save_upload, REF_FILES_DIR
router = APIRouter(prefix="/generation-points", tags=["生成点管理"])
@router.post("", response_model=GenerationPointResponse)
async def create_generation_point(
template_id: uuid.UUID = Form(...),
position: str = Form(...),
prompt: str = Form(...),
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),
):
import json
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="模板不存在")
if model_id:
model_result = await db.execute(select(AIModel).where(AIModel.id == model_id))
if not model_result.scalar_one_or_none():
raise HTTPException(status_code=404, detail="模型不存在")
ref_file_path = None
if ref_file and ref_file.filename:
validate_file_extension(ref_file.filename)
ref_file_path = await save_upload(ref_file, REF_FILES_DIR)
point = GenerationPoint(
template_id=template_id,
position=json.loads(position),
prompt=prompt,
model_id=model_id,
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()
await db.refresh(point)
return point
@router.get("", response_model=list[GenerationPointResponse])
async def list_generation_points(
template_id: uuid.UUID = Query(...),
db: AsyncSession = Depends(get_db),
):
query = (
select(GenerationPoint)
.where(GenerationPoint.template_id == template_id)
.order_by(GenerationPoint.order.asc(), GenerationPoint.created_at.asc())
)
result = await db.execute(query)
return result.scalars().all()
@router.get("/{point_id}", response_model=GenerationPointResponse)
async def get_generation_point(point_id: str, db: AsyncSession = Depends(get_db)):
result = await db.execute(select(GenerationPoint).where(GenerationPoint.id == point_id))
point = result.scalar_one_or_none()
if not point:
raise HTTPException(status_code=404, detail="生成点不存在")
return point
@router.put("/{point_id}", response_model=GenerationPointResponse)
async def update_generation_point(
point_id: str,
data: GenerationPointUpdate,
db: AsyncSession = Depends(get_db),
):
result = await db.execute(select(GenerationPoint).where(GenerationPoint.id == point_id))
point = result.scalar_one_or_none()
if not point:
raise HTTPException(status_code=404, detail="生成点不存在")
update_data = data.model_dump(exclude_unset=True)
for key, value in update_data.items():
setattr(point, key, value)
await db.flush()
await db.refresh(point)
return point
@router.delete("/{point_id}")
async def delete_generation_point(point_id: str, db: AsyncSession = Depends(get_db)):
result = await db.execute(select(GenerationPoint).where(GenerationPoint.id == point_id))
point = result.scalar_one_or_none()
if not point:
raise HTTPException(status_code=404, detail="生成点不存在")
await db.delete(point)
return {"detail": "删除成功"}
@router.post("/batch-order")
async def batch_update_order(data: BatchOrderUpdate, db: AsyncSession = Depends(get_db)):
for item in data.points:
result = await db.execute(select(GenerationPoint).where(GenerationPoint.id == item["id"]))
point = result.scalar_one_or_none()
if point:
point.order = item["order"]
await db.flush()
return {"detail": "排序更新成功"}
+115
View File
@@ -0,0 +1,115 @@
from fastapi import APIRouter, Depends, HTTPException, Query
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy import select
from app.core.database import get_db
from app.core.security import encrypt_api_key
from app.models.ai_model import AIModel
from app.schemas.ai_model import AIModelCreate, AIModelUpdate, AIModelToggle, AIModelResponse
from app.services.ai_adapter import call_ai_model
router = APIRouter(prefix="/models", tags=["AI模型管理"])
@router.post("", response_model=AIModelResponse)
async def create_model(data: AIModelCreate, db: AsyncSession = Depends(get_db)):
encrypted_key = encrypt_api_key(data.api_key)
model = AIModel(
name=data.name,
provider=data.provider,
endpoint=data.endpoint,
api_key=encrypted_key,
extra_params=data.extra_params,
is_enabled=data.is_enabled,
remark=data.remark,
)
db.add(model)
await db.flush()
await db.refresh(model)
return model
@router.get("", response_model=list[AIModelResponse])
async def list_models(
enabled: bool | None = Query(None, description="过滤启用/禁用"),
skip: int = Query(0, ge=0),
limit: int = Query(20, ge=1, le=100),
db: AsyncSession = Depends(get_db),
):
query = select(AIModel)
if enabled is not None:
query = query.where(AIModel.is_enabled == enabled)
query = query.offset(skip).limit(limit).order_by(AIModel.created_at.desc())
result = await db.execute(query)
models = result.scalars().all()
return models
@router.get("/{model_id}", response_model=AIModelResponse)
async def get_model(model_id: str, db: AsyncSession = Depends(get_db)):
result = await db.execute(select(AIModel).where(AIModel.id == model_id))
model = result.scalar_one_or_none()
if not model:
raise HTTPException(status_code=404, detail="模型不存在")
return model
@router.put("/{model_id}", response_model=AIModelResponse)
async def update_model(model_id: str, data: AIModelUpdate, db: AsyncSession = Depends(get_db)):
result = await db.execute(select(AIModel).where(AIModel.id == model_id))
model = result.scalar_one_or_none()
if not model:
raise HTTPException(status_code=404, detail="模型不存在")
update_data = data.model_dump(exclude_unset=True)
if "api_key" in update_data and update_data["api_key"] is not None:
update_data["api_key"] = encrypt_api_key(update_data["api_key"])
for key, value in update_data.items():
setattr(model, key, value)
await db.flush()
await db.refresh(model)
return model
@router.delete("/{model_id}")
async def delete_model(model_id: str, db: AsyncSession = Depends(get_db)):
result = await db.execute(select(AIModel).where(AIModel.id == model_id))
model = result.scalar_one_or_none()
if not model:
raise HTTPException(status_code=404, detail="模型不存在")
await db.delete(model)
return {"detail": "删除成功"}
@router.patch("/{model_id}/toggle", response_model=AIModelResponse)
async def toggle_model(model_id: str, data: AIModelToggle, db: AsyncSession = Depends(get_db)):
result = await db.execute(select(AIModel).where(AIModel.id == model_id))
model = result.scalar_one_or_none()
if not model:
raise HTTPException(status_code=404, detail="模型不存在")
model.is_enabled = data.is_enabled
await db.flush()
await db.refresh(model)
return model
@router.post("/{model_id}/test")
async def test_model(model_id: str, db: AsyncSession = Depends(get_db)):
result = await db.execute(select(AIModel).where(AIModel.id == model_id))
model = result.scalar_one_or_none()
if not model:
raise HTTPException(status_code=404, detail="模型不存在")
model_config = {
"provider": model.provider,
"endpoint": model.endpoint,
"api_key": model.api_key,
"extra_params": model.extra_params,
}
try:
response = await call_ai_model(model_config, "请用一句话介绍你自己。")
return {"success": True, "result": response}
except Exception as e:
return {"success": False, "error": str(e)}
+16
View File
@@ -0,0 +1,16 @@
from fastapi import APIRouter
from app.api.models import router as models_router
from app.api.templates import router as templates_router
from app.api.generation_points import router as generation_points_router
from app.api.tasks import router as tasks_router
from app.api.settings import router as settings_router
from app.core.config import get_settings
settings = get_settings()
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, prefix="")
api_router.include_router(settings_router)
+34
View File
@@ -0,0 +1,34 @@
from fastapi import APIRouter, Depends
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy import select
from app.core.database import get_db
from app.models.system_config import SystemConfig
router = APIRouter(prefix="/settings", tags=["系统设置"])
@router.get("")
async def get_all_settings(db: AsyncSession = Depends(get_db)):
result = await db.execute(select(SystemConfig))
configs = result.scalars().all()
return {c.key: c.value for c in configs}
@router.put("/{key}")
async def update_setting(key: str, body: dict, db: AsyncSession = Depends(get_db)):
result = await db.execute(select(SystemConfig).where(SystemConfig.key == key))
config = result.scalar_one_or_none()
value = body.get("value", "")
description = body.get("description", "")
if config:
config.value = value
if description:
config.description = description
else:
config = SystemConfig(key=key, value=value, description=description)
db.add(config)
await db.flush()
return {"key": key, "value": value}
+233
View File
@@ -0,0 +1,233 @@
import json
import os
import uuid
from datetime import datetime, timezone
from fastapi import APIRouter, Depends, HTTPException, Query, UploadFile, File
from fastapi.responses import Response
from typing import List
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, TaskDetailResponse
from app.services.document_processor import docx_to_pdf_bytes
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, parse_reference_files
from app.tasks.generate import generate_document
router = APIRouter()
@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}")
async def get_task_status(task_id: str, detail: bool = Query(False), 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 detail:
tpl_result = await db.execute(select(Template).where(Template.id == task.template_id))
tpl = tpl_result.scalar_one_or_none()
points_result = await db.execute(
select(GenerationPoint).where(GenerationPoint.template_id == task.template_id).order_by(GenerationPoint.order.asc())
)
points = points_result.scalars().all()
return TaskDetailResponse(
id=task.id,
template_id=task.template_id,
template_name=tpl.name if tpl else "",
status=task.status,
result_file_path=task.result_file_path,
error_msg=task.error_msg,
created_at=task.created_at,
finished_at=task.finished_at,
points=[{
"id": str(p.id),
"prompt": p.prompt,
"position": p.position,
"model_id": str(p.model_id) if p.model_id else None,
"ref_file_path": p.ref_file_path,
"selected_text": p.selected_text,
"need_ref_file": p.need_ref_file,
"remark": p.remark,
"order": p.order,
} for p in points],
)
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,
ref_files: List[UploadFile] = File(default=[]),
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="生成点不存在")
tmp_paths = []
if point.need_ref_file:
if not ref_files and not point.ref_file_path:
raise HTTPException(status_code=400, detail="此生成点需要上传参考文件")
for f in ref_files:
if f.filename:
validate_file_extension(f.filename)
path = await save_upload(f, REF_FILES_DIR)
tmp_paths.append(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))
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 = ""
if tmp_paths:
ref_content = await parse_reference_files(json.dumps(tmp_paths))
elif point.ref_file_path:
ref_content = await parse_reference_files(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)}")
@router.post("/generation-points/{point_id}/upload-ref")
async def upload_ref_files(
point_id: str,
ref_files: List[UploadFile] = File(...),
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="生成点不存在")
paths = []
for f in ref_files:
if f.filename:
validate_file_extension(f.filename)
path = await save_upload(f, REF_FILES_DIR)
paths.append(path)
existing = []
if point.ref_file_path:
try:
existing = json.loads(point.ref_file_path)
if not isinstance(existing, list):
existing = [point.ref_file_path]
except json.JSONDecodeError:
existing = [point.ref_file_path] if point.ref_file_path else []
all_paths = existing + paths
point.ref_file_path = json.dumps(all_paths)
await db.commit()
return {"success": True, "count": len(all_paths)}
+150
View File
@@ -0,0 +1,150 @@
import uuid
import bleach
from fastapi import APIRouter, Depends, HTTPException, UploadFile, File, Form, 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.core.security_middleware import validate_file_extension
from app.models.template import Template
from app.schemas.template import TemplateResponse, TemplateListItem, HTMLContentResponse, HTMLUpdateRequest
from app.services.file_storage import save_upload, get_file_content, delete_file, TEMPLATES_DIR
from app.services.document_processor import docx_to_html, html_to_docx_bytes, docx_to_pdf_bytes
ALLOWED_TAGS = [
"p", "div", "span", "br", "hr",
"h1", "h2", "h3", "h4", "h5", "h6",
"ul", "ol", "li",
"a", "img", "table", "thead", "tbody", "tr", "td", "th",
"b", "i", "u", "strong", "em", "del", "sub", "sup",
"pre", "code", "blockquote",
]
ALLOWED_ATTRS = {
"a": ["href", "title", "target"],
"img": ["src", "alt", "width", "height"],
"td": ["colspan", "rowspan"],
"th": ["colspan", "rowspan"],
"p": ["style"],
"span": ["style"],
"div": ["style"],
"table": ["style"],
}
router = APIRouter(prefix="/templates", tags=["模板管理"])
@router.post("", response_model=TemplateResponse)
async def upload_template(
file: UploadFile = File(...),
name: str | None = Form(None),
db: AsyncSession = Depends(get_db),
):
if not file.filename or not file.filename.endswith(".docx"):
raise HTTPException(status_code=400, detail="仅支持 .docx 文件")
validate_file_extension(file.filename)
file_path = await save_upload(file, TEMPLATES_DIR)
file_content = await get_file_content(file_path)
html_content = await docx_to_html(file_content)
template = Template(
name=name or file.filename.replace(".docx", ""),
file_path=file_path,
html_content=html_content,
)
db.add(template)
await db.flush()
await db.refresh(template)
return template
@router.get("", response_model=list[TemplateListItem])
async def list_templates(
skip: int = Query(0, ge=0),
limit: int = Query(20, ge=1, le=100),
db: AsyncSession = Depends(get_db),
):
query = select(Template).offset(skip).limit(limit).order_by(Template.created_at.desc())
result = await db.execute(query)
return result.scalars().all()
@router.get("/{template_id}", response_model=TemplateResponse)
async def get_template(template_id: str, db: AsyncSession = Depends(get_db)):
result = await db.execute(select(Template).where(Template.id == template_id))
template = result.scalar_one_or_none()
if not template:
raise HTTPException(status_code=404, detail="模板不存在")
return template
@router.get("/{template_id}/html", response_model=HTMLContentResponse)
async def get_template_html(template_id: str, db: AsyncSession = Depends(get_db)):
result = await db.execute(select(Template).where(Template.id == template_id))
template = result.scalar_one_or_none()
if not template:
raise HTTPException(status_code=404, detail="模板不存在")
return HTMLContentResponse(html_content=template.html_content or "")
@router.put("/{template_id}/html")
async def update_template_html(
template_id: str,
data: HTMLUpdateRequest,
db: AsyncSession = Depends(get_db),
):
result = await db.execute(select(Template).where(Template.id == template_id))
template = result.scalar_one_or_none()
if not template:
raise HTTPException(status_code=404, detail="模板不存在")
template.html_content = bleach.clean(
data.html_content, tags=ALLOWED_TAGS, attributes=ALLOWED_ATTRS, strip=True
)
docx_bytes = html_to_docx_bytes(template.html_content)
with open(template.file_path, "wb") as f:
f.write(docx_bytes)
await db.flush()
return {"detail": "保存成功"}
@router.get("/{template_id}/download")
async def download_template(
template_id: str,
format: str = Query("docx", pattern="^(docx|pdf)$"),
db: AsyncSession = Depends(get_db),
):
result = await db.execute(select(Template).where(Template.id == template_id))
template = result.scalar_one_or_none()
if not template:
raise HTTPException(status_code=404, detail="模板不存在")
file_content = await get_file_content(template.file_path)
if format == "pdf":
file_content = docx_to_pdf_bytes(file_content)
media_type = "application/pdf"
filename = f"{template.name}.pdf"
else:
media_type = "application/vnd.openxmlformats-officedocument.wordprocessingml.document"
filename = f"{template.name}.docx"
return Response(
content=file_content,
media_type=media_type,
headers={"Content-Disposition": f'attachment; filename="{filename}"'},
)
@router.delete("/{template_id}")
async def delete_template(template_id: str, db: AsyncSession = Depends(get_db)):
result = await db.execute(select(Template).where(Template.id == template_id))
template = result.scalar_one_or_none()
if not template:
raise HTTPException(status_code=404, detail="模板不存在")
delete_file(template.file_path)
await db.delete(template)
return {"detail": "删除成功"}