652 lines
24 KiB
Python
652 lines
24 KiB
Python
import asyncio
|
|
import json
|
|
import os
|
|
import time
|
|
import uuid
|
|
from datetime import datetime
|
|
|
|
from fastapi import APIRouter, Depends, File, HTTPException, Query, UploadFile
|
|
from sqlalchemy import func, select
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
from sse_starlette.sse import EventSourceResponse
|
|
|
|
from config import settings
|
|
from database import get_db
|
|
from models.ai_model import AiModel
|
|
from models.document import Document
|
|
from models.generation_log import GenerationLog
|
|
from models.paragraph import Paragraph
|
|
from models.reference_file import ReferenceFile
|
|
from models.template import Template
|
|
from models.template_block import TemplateBlock
|
|
from schemas.schemas import GenerateFullRequest, GenerateTestRequest, ReferenceFileUpdate, Response
|
|
from services.ai_service import call_ai, stream_ai_preview
|
|
from services.file_summary import summarize_minio_files
|
|
from services.generation_runtime import (
|
|
build_mock_content,
|
|
build_mock_content_for_block,
|
|
generation_progress,
|
|
request_cancel,
|
|
run_generation,
|
|
run_generation_for_blocks,
|
|
update_progress,
|
|
)
|
|
from services.minio_client import delete_object, split_bucket_path, upload_bytes, get_presigned_url
|
|
|
|
router = APIRouter()
|
|
|
|
|
|
def _serialize_document(document: Document) -> dict:
|
|
request_payload = {}
|
|
if document.request_payload_json:
|
|
try:
|
|
request_payload = json.loads(document.request_payload_json)
|
|
except Exception:
|
|
request_payload = {}
|
|
return {
|
|
"id": document.id,
|
|
"template_id": document.template_id,
|
|
"name": document.name,
|
|
"para_count_done": document.para_count_done,
|
|
"para_count_total": document.para_count_total,
|
|
"status": document.status,
|
|
"file_path": document.file_path,
|
|
"error": document.error,
|
|
"request_payload": request_payload,
|
|
"created_at": document.created_at,
|
|
"updated_at": document.updated_at,
|
|
}
|
|
|
|
|
|
def _serialize_reference_file(file: ReferenceFile) -> dict:
|
|
return {
|
|
"id": file.id,
|
|
"file_name": file.file_name,
|
|
"file_path": file.file_path,
|
|
"file_size": file.file_size,
|
|
"content_type": file.content_type,
|
|
"created_at": file.created_at.isoformat() if file.created_at else None,
|
|
}
|
|
|
|
|
|
async def _build_reference_name_mapping(db: AsyncSession, file_paths: list[str]) -> dict[str, str]:
|
|
if not file_paths:
|
|
return {}
|
|
result = await db.execute(
|
|
select(ReferenceFile.file_path, ReferenceFile.file_name).where(ReferenceFile.file_path.in_(file_paths))
|
|
)
|
|
return {file_path: file_name for file_path, file_name in result.all()}
|
|
|
|
|
|
async def _build_reference_records_mapping(db: AsyncSession, file_paths: list[str]) -> dict[str, ReferenceFile]:
|
|
if not file_paths:
|
|
return {}
|
|
result = await db.execute(select(ReferenceFile).where(ReferenceFile.file_path.in_(file_paths)))
|
|
return {item.file_path: item for item in result.scalars().all()}
|
|
|
|
@router.post("/test")
|
|
async def generate_test(body: GenerateTestRequest, db: AsyncSession = Depends(get_db)):
|
|
paragraph = await db.get(Paragraph, body.paragraph_id)
|
|
if paragraph is None or paragraph.template_id != body.template_id:
|
|
raise HTTPException(status_code=404, detail="段落不存在")
|
|
if body.prompt_text:
|
|
paragraph.prompt_text = body.prompt_text
|
|
|
|
model = None
|
|
if body.model_id:
|
|
model = await db.get(AiModel, body.model_id)
|
|
elif paragraph.model_id:
|
|
model = await db.get(AiModel, paragraph.model_id)
|
|
|
|
file_name_mapping = await _build_reference_name_mapping(db, body.file_paths or [])
|
|
file_summaries = (
|
|
await asyncio.to_thread(summarize_minio_files, body.file_paths or [], file_name_mapping)
|
|
if body.file_paths
|
|
else []
|
|
)
|
|
if model is None or model.status != "enabled":
|
|
content = build_mock_content(paragraph)
|
|
message = "当前未找到可用模型,返回本地模拟生成结果。"
|
|
else:
|
|
setattr(paragraph, "enable_reasoning", bool(model.enable_reasoning))
|
|
result = await call_ai(paragraph, model, file_summaries)
|
|
content = result.content
|
|
message = f"已通过模型 {result.used_model} 生成。"
|
|
return Response(
|
|
data={
|
|
"paragraph_id": paragraph.id,
|
|
"content": content,
|
|
"message": message,
|
|
"file_summaries": file_summaries,
|
|
}
|
|
)
|
|
|
|
|
|
@router.post("/test-stream")
|
|
async def generate_test_stream(body: GenerateTestRequest, db: AsyncSession = Depends(get_db)):
|
|
paragraph = await db.get(Paragraph, body.paragraph_id)
|
|
if paragraph is None or paragraph.template_id != body.template_id:
|
|
raise HTTPException(status_code=404, detail="段落不存在")
|
|
if body.prompt_text:
|
|
paragraph.prompt_text = body.prompt_text
|
|
|
|
model = None
|
|
if body.model_id:
|
|
model = await db.get(AiModel, body.model_id)
|
|
elif paragraph.model_id:
|
|
model = await db.get(AiModel, paragraph.model_id)
|
|
|
|
if model is None or model.status != "enabled":
|
|
raise HTTPException(status_code=400, detail="当前段落未配置可用的流式模型")
|
|
if not model.supports_streaming:
|
|
raise HTTPException(status_code=400, detail="当前模型未开启流式传输")
|
|
|
|
file_name_mapping = await _build_reference_name_mapping(db, body.file_paths or [])
|
|
file_summaries = (
|
|
await asyncio.to_thread(summarize_minio_files, body.file_paths or [], file_name_mapping)
|
|
if body.file_paths
|
|
else []
|
|
)
|
|
setattr(paragraph, "enable_reasoning", bool(model.enable_reasoning))
|
|
|
|
async def event_stream():
|
|
yield {
|
|
"event": "message",
|
|
"data": json.dumps(
|
|
{
|
|
"type": "meta",
|
|
"message": f"正在通过模型 {model.name} 流式生成...",
|
|
"file_summaries": file_summaries,
|
|
},
|
|
ensure_ascii=False,
|
|
),
|
|
}
|
|
try:
|
|
async for chunk in stream_ai_preview(paragraph, model, file_summaries):
|
|
yield {
|
|
"event": "message",
|
|
"data": json.dumps({"type": "delta", "content": chunk}, ensure_ascii=False),
|
|
}
|
|
yield {
|
|
"event": "message",
|
|
"data": json.dumps({"type": "done"}, ensure_ascii=False),
|
|
}
|
|
except Exception as error:
|
|
fallback_message = str(error)
|
|
if "503" in fallback_message or "temporarily unavailable" in fallback_message.lower():
|
|
try:
|
|
result = await call_ai(paragraph, model, file_summaries)
|
|
yield {
|
|
"event": "message",
|
|
"data": json.dumps(
|
|
{
|
|
"type": "meta",
|
|
"message": "流式通道暂时不可用,已自动回退为普通返回。",
|
|
"file_summaries": file_summaries,
|
|
},
|
|
ensure_ascii=False,
|
|
),
|
|
}
|
|
yield {
|
|
"event": "message",
|
|
"data": json.dumps({"type": "delta", "content": result.raw_text}, ensure_ascii=False),
|
|
}
|
|
yield {
|
|
"event": "message",
|
|
"data": json.dumps({"type": "done"}, ensure_ascii=False),
|
|
}
|
|
return
|
|
except Exception as fallback_error:
|
|
fallback_message = f"{fallback_message};普通调用回退也失败:{fallback_error}"
|
|
yield {
|
|
"event": "message",
|
|
"data": json.dumps({"type": "error", "message": fallback_message}, ensure_ascii=False),
|
|
}
|
|
|
|
return EventSourceResponse(
|
|
event_stream(),
|
|
headers={
|
|
"Cache-Control": "no-cache",
|
|
"Connection": "keep-alive",
|
|
"X-Accel-Buffering": "no",
|
|
},
|
|
)
|
|
|
|
|
|
@router.post("/upload")
|
|
async def upload_reference_file(file: UploadFile = File(...), db: AsyncSession = Depends(get_db)):
|
|
if not file.filename:
|
|
raise HTTPException(status_code=400, detail="文件名不能为空")
|
|
|
|
ext = os.path.splitext(file.filename)[1].lower()
|
|
if ext not in settings.ALLOWED_EXTENSIONS:
|
|
allowed = " / ".join(settings.ALLOWED_EXTENSIONS)
|
|
raise HTTPException(status_code=400, detail=f"文件类型不支持:{ext or '无扩展名'}。当前支持:{allowed}")
|
|
|
|
content = await file.read()
|
|
if not content:
|
|
raise HTTPException(status_code=400, detail="上传文件不能为空")
|
|
if len(content) > settings.MAX_UPLOAD_SIZE:
|
|
raise HTTPException(status_code=400, detail="文件大小超过限制")
|
|
|
|
object_name = f"{datetime.now().strftime('%Y%m%d')}/{uuid.uuid4().hex}{ext}"
|
|
await asyncio.to_thread(
|
|
upload_bytes,
|
|
settings.MINIO_BUCKET_UPLOADS,
|
|
object_name,
|
|
content,
|
|
file.content_type or "application/octet-stream",
|
|
)
|
|
record = ReferenceFile(
|
|
file_name=file.filename,
|
|
file_path=f"{settings.MINIO_BUCKET_UPLOADS}/{object_name}",
|
|
file_size=len(content),
|
|
content_type=file.content_type or "application/octet-stream",
|
|
)
|
|
db.add(record)
|
|
await db.commit()
|
|
await db.refresh(record)
|
|
return Response(
|
|
data=_serialize_reference_file(record)
|
|
)
|
|
|
|
|
|
@router.get("/reference-files")
|
|
async def list_reference_files(
|
|
page: int = Query(1, ge=1),
|
|
page_size: int = Query(20, ge=1, le=100),
|
|
keyword: str = Query("", description="按文件名搜索"),
|
|
db: AsyncSession = Depends(get_db),
|
|
):
|
|
stmt = select(ReferenceFile)
|
|
count_stmt = select(func.count(ReferenceFile.id))
|
|
if keyword:
|
|
like_keyword = f"%{keyword.strip()}%"
|
|
stmt = stmt.where(ReferenceFile.file_name.like(like_keyword))
|
|
count_stmt = count_stmt.where(ReferenceFile.file_name.like(like_keyword))
|
|
|
|
total = (await db.execute(count_stmt)).scalar_one()
|
|
result = await db.execute(
|
|
stmt.order_by(ReferenceFile.id.desc()).offset((page - 1) * page_size).limit(page_size)
|
|
)
|
|
items = [_serialize_reference_file(item) for item in result.scalars().all()]
|
|
return Response(data={"items": items, "total": total, "page": page, "page_size": page_size})
|
|
|
|
|
|
@router.get("/reference-files/{file_id}/download")
|
|
async def download_reference_file(file_id: int, db: AsyncSession = Depends(get_db)):
|
|
record = await db.get(ReferenceFile, file_id)
|
|
if record is None:
|
|
raise HTTPException(status_code=404, detail="附件记录不存在")
|
|
bucket, object_name = split_bucket_path(record.file_path)
|
|
return Response(data={"url": get_presigned_url(bucket, object_name), "file_name": record.file_name})
|
|
|
|
|
|
@router.put("/reference-files/{file_id}")
|
|
async def update_reference_file(file_id: int, body: ReferenceFileUpdate, db: AsyncSession = Depends(get_db)):
|
|
record = await db.get(ReferenceFile, file_id)
|
|
if record is None:
|
|
raise HTTPException(status_code=404, detail="附件记录不存在")
|
|
record.file_name = body.file_name.strip()
|
|
await db.commit()
|
|
await db.refresh(record)
|
|
return Response(data=_serialize_reference_file(record))
|
|
|
|
|
|
@router.delete("/reference-files/{file_id}")
|
|
async def delete_reference_file(file_id: int, db: AsyncSession = Depends(get_db)):
|
|
record = await db.get(ReferenceFile, file_id)
|
|
if record is None:
|
|
raise HTTPException(status_code=404, detail="附件记录不存在")
|
|
try:
|
|
bucket, object_name = split_bucket_path(record.file_path)
|
|
await asyncio.to_thread(delete_object, bucket, object_name)
|
|
except Exception:
|
|
pass
|
|
await db.delete(record)
|
|
await db.commit()
|
|
return Response(data={"id": file_id})
|
|
|
|
|
|
@router.post("/full")
|
|
async def generate_full(body: GenerateFullRequest, db: AsyncSession = Depends(get_db)):
|
|
template = await db.get(Template, body.template_id)
|
|
if template is None:
|
|
raise HTTPException(status_code=404, detail="模板不存在")
|
|
|
|
result = await db.execute(
|
|
select(Paragraph)
|
|
.where(Paragraph.template_id == body.template_id)
|
|
.order_by(Paragraph.sort_index.asc(), Paragraph.id.asc())
|
|
)
|
|
paragraphs = result.scalars().all()
|
|
if not paragraphs:
|
|
raise HTTPException(status_code=400, detail="模板下暂无可生成段落")
|
|
|
|
normalized_file_map: dict[str, list[str]] = {}
|
|
for key, values in (body.file_map or {}).items():
|
|
if not isinstance(values, list):
|
|
continue
|
|
normalized_file_map[str(key)] = [item for item in values if item]
|
|
all_file_paths = [item for values in normalized_file_map.values() for item in values]
|
|
file_records = await _build_reference_records_mapping(db, all_file_paths)
|
|
paragraph_snapshot = []
|
|
for paragraph in paragraphs:
|
|
selected_paths = normalized_file_map.get(str(paragraph.id), [])
|
|
paragraph_snapshot.append(
|
|
{
|
|
"paragraph_id": paragraph.id,
|
|
"title": paragraph.title,
|
|
"sort_index": paragraph.sort_index,
|
|
"need_file": bool(paragraph.need_file),
|
|
"file_note": paragraph.file_note or "",
|
|
"selected_files": [
|
|
_serialize_reference_file(file_records[file_path])
|
|
for file_path in selected_paths
|
|
if file_path in file_records
|
|
],
|
|
}
|
|
)
|
|
request_payload = {
|
|
"template_name": template.name,
|
|
"file_map": normalized_file_map,
|
|
"paragraphs": paragraph_snapshot,
|
|
}
|
|
|
|
document = Document(
|
|
template_id=template.id,
|
|
name=f"{template.name}-{datetime.now().strftime('%Y%m%d%H%M%S')}",
|
|
para_count_done=0,
|
|
para_count_total=len(paragraphs),
|
|
status="generating",
|
|
file_path="",
|
|
error="",
|
|
request_payload_json=json.dumps(request_payload, ensure_ascii=False),
|
|
)
|
|
db.add(document)
|
|
await db.flush()
|
|
await db.commit()
|
|
await db.refresh(document)
|
|
update_progress(document.id, status="pending", percent=0, done=0, total=len(paragraphs), message="任务已创建")
|
|
asyncio.create_task(run_generation(document.id, template.id))
|
|
return Response(data=_serialize_document(document))
|
|
|
|
|
|
@router.post("/full-v2")
|
|
async def generate_full_v2(body: GenerateFullRequest, db: AsyncSession = Depends(get_db)):
|
|
"""基于块结构的全量生成(可视化编辑器使用)。"""
|
|
template = await db.get(Template, body.template_id)
|
|
if template is None:
|
|
raise HTTPException(status_code=404, detail="模板不存在")
|
|
|
|
result = await db.execute(
|
|
select(TemplateBlock)
|
|
.where(TemplateBlock.template_id == body.template_id)
|
|
.order_by(TemplateBlock.sort_index.asc(), TemplateBlock.id.asc())
|
|
)
|
|
blocks = result.scalars().all()
|
|
if not blocks:
|
|
raise HTTPException(status_code=400, detail="模板下暂无可生成块")
|
|
|
|
ai_block_count = sum(1 for b in blocks if b.edit_mode == "ai")
|
|
if ai_block_count == 0:
|
|
raise HTTPException(status_code=400, detail="模板下没有 AI 生成块")
|
|
|
|
document = Document(
|
|
template_id=template.id,
|
|
name=f"{template.name}-{datetime.now().strftime('%Y%m%d%H%M%S')}",
|
|
para_count_done=0,
|
|
para_count_total=ai_block_count,
|
|
status="generating",
|
|
file_path="",
|
|
error="",
|
|
request_payload_json=json.dumps({"template_name": template.name, "block_count": len(blocks), "ai_block_count": ai_block_count}, ensure_ascii=False),
|
|
)
|
|
db.add(document)
|
|
await db.flush()
|
|
await db.commit()
|
|
await db.refresh(document)
|
|
update_progress(document.id, status="pending", percent=0, done=0, total=ai_block_count, message="块级生成任务已创建")
|
|
asyncio.create_task(run_generation_for_blocks(document.id, template.id))
|
|
return Response(data=_serialize_document(document))
|
|
|
|
|
|
@router.get("/progress/{document_id}")
|
|
async def generate_progress(document_id: int):
|
|
async def event_generator():
|
|
while True:
|
|
state = generation_progress.get(
|
|
document_id,
|
|
{"status": "pending", "percent": 0, "message": "等待中", "done": 0, "total": 0},
|
|
)
|
|
yield {
|
|
"event": "progress",
|
|
"data": json.dumps(state, ensure_ascii=False),
|
|
}
|
|
if state.get("status") in {"completed", "failed", "cancelled"}:
|
|
break
|
|
await asyncio.sleep(1)
|
|
|
|
return EventSourceResponse(event_generator())
|
|
|
|
|
|
@router.post("/blocks/{block_id}/regenerate")
|
|
async def regenerate_block(
|
|
block_id: int,
|
|
body: GenerateTestRequest,
|
|
db: AsyncSession = Depends(get_db),
|
|
):
|
|
"""对单个块重新生成内容。"""
|
|
block = await db.get(TemplateBlock, block_id)
|
|
if block is None:
|
|
raise HTTPException(status_code=404, detail="块不存在")
|
|
|
|
document = await db.get(Document, body.template_id)
|
|
if document is None:
|
|
raise HTTPException(status_code=404, detail="生成记录不存在")
|
|
|
|
model = None
|
|
if body.model_id:
|
|
model = await db.get(AiModel, body.model_id)
|
|
elif block.model_id:
|
|
model = await db.get(AiModel, block.model_id)
|
|
|
|
if model is None or model.status != "enabled":
|
|
# 使用默认模型
|
|
result = await db.execute(
|
|
select(AiModel).where(AiModel.status == "enabled").order_by(AiModel.id.asc()).limit(1)
|
|
)
|
|
model = result.scalars().first()
|
|
|
|
content = {}
|
|
status = "success"
|
|
error_msg = ""
|
|
duration = 0.0
|
|
|
|
if model is None:
|
|
content = build_mock_content_for_block(block)
|
|
else:
|
|
start = time.perf_counter()
|
|
try:
|
|
temp_para = Paragraph(
|
|
template_id=block.template_id,
|
|
sort_index=block.sort_index,
|
|
title=block.title,
|
|
content="",
|
|
prompt_text=body.prompt_text or block.prompt_text or "",
|
|
output_format=block.output_format,
|
|
need_prompt=block.need_prompt,
|
|
edit_mode=block.edit_mode,
|
|
)
|
|
setattr(temp_para, "enable_reasoning", bool(model.enable_reasoning))
|
|
result_data = await call_ai(temp_para, model, None)
|
|
content = result_data.content
|
|
except Exception as error:
|
|
content = build_mock_content_for_block(block)
|
|
status = "failed"
|
|
error_msg = str(error)
|
|
duration = round(time.perf_counter() - start, 4)
|
|
|
|
log = GenerationLog(
|
|
document_id=document.id,
|
|
block_id=block.id,
|
|
model_id=model.id if model else None,
|
|
status=status,
|
|
content=json.dumps(content, ensure_ascii=False),
|
|
duration=duration,
|
|
error_msg=error_msg,
|
|
)
|
|
db.add(log)
|
|
await db.commit()
|
|
await db.refresh(log)
|
|
|
|
return Response(data={
|
|
"id": log.id,
|
|
"block_id": block.id,
|
|
"status": log.status,
|
|
"content": content,
|
|
"duration": log.duration,
|
|
"error_msg": log.error_msg,
|
|
})
|
|
|
|
|
|
@router.get("/documents")
|
|
async def list_documents(
|
|
page: int = Query(1, ge=1),
|
|
page_size: int = Query(20, ge=1, le=100),
|
|
db: AsyncSession = Depends(get_db),
|
|
):
|
|
total = (await db.execute(select(func.count(Document.id)))).scalar_one()
|
|
result = await db.execute(
|
|
select(Document)
|
|
.order_by(Document.id.desc())
|
|
.offset((page - 1) * page_size)
|
|
.limit(page_size)
|
|
)
|
|
items = [_serialize_document(item) for item in result.scalars().all()]
|
|
return Response(data={"items": items, "total": total, "page": page, "page_size": page_size})
|
|
|
|
|
|
@router.get("/documents/{document_id}")
|
|
async def get_document(document_id: int, db: AsyncSession = Depends(get_db)):
|
|
document = await db.get(Document, document_id)
|
|
if document is None:
|
|
raise HTTPException(status_code=404, detail="生成记录不存在")
|
|
|
|
# 查询段落级日志
|
|
log_result = await db.execute(
|
|
select(GenerationLog, Paragraph)
|
|
.join(Paragraph, Paragraph.id == GenerationLog.paragraph_id)
|
|
.where(GenerationLog.document_id == document_id)
|
|
.order_by(Paragraph.sort_index.asc(), Paragraph.id.asc())
|
|
)
|
|
items = []
|
|
for log, paragraph in log_result.all():
|
|
items.append(
|
|
{
|
|
"id": log.id,
|
|
"paragraph_id": paragraph.id,
|
|
"title": paragraph.title,
|
|
"sort_index": paragraph.sort_index,
|
|
"status": log.status,
|
|
"content": json.loads(log.content) if log.content else {"content": []},
|
|
}
|
|
)
|
|
|
|
payload = _serialize_document(document)
|
|
payload["logs"] = items
|
|
|
|
# 查询块级日志和模板块(用于可视化预览)
|
|
blocks_result = await db.execute(
|
|
select(TemplateBlock)
|
|
.where(TemplateBlock.template_id == document.template_id)
|
|
.order_by(TemplateBlock.sort_index.asc(), TemplateBlock.id.asc())
|
|
)
|
|
template_blocks = blocks_result.scalars().all()
|
|
|
|
block_logs_result = await db.execute(
|
|
select(GenerationLog)
|
|
.where(GenerationLog.document_id == document_id, GenerationLog.block_id.isnot(None))
|
|
)
|
|
block_logs = block_logs_result.scalars().all()
|
|
|
|
block_log_map: dict[int, dict] = {}
|
|
for log in block_logs:
|
|
if log.block_id is not None:
|
|
block_log_map[log.block_id] = {
|
|
"id": log.id,
|
|
"block_id": log.block_id,
|
|
"status": log.status,
|
|
"content": json.loads(log.content) if log.content else {"content": []},
|
|
"duration": log.duration,
|
|
"error_msg": log.error_msg,
|
|
}
|
|
|
|
payload["template_blocks"] = [
|
|
{
|
|
"id": b.id,
|
|
"sort_index": b.sort_index,
|
|
"block_type": b.block_type,
|
|
"anchor_start_index": b.anchor_start_index,
|
|
"anchor_end_index": b.anchor_end_index,
|
|
"html_snippet": b.html_snippet,
|
|
"title": b.title,
|
|
"edit_mode": b.edit_mode,
|
|
"output_format": b.output_format,
|
|
"model_id": b.model_id,
|
|
"prompt_text": b.prompt_text,
|
|
}
|
|
for b in template_blocks
|
|
]
|
|
payload["block_logs"] = block_log_map
|
|
payload["has_blocks"] = len(template_blocks) > 0
|
|
|
|
return Response(data=payload)
|
|
|
|
|
|
@router.post("/cancel/{document_id}")
|
|
async def cancel_document(document_id: int, db: AsyncSession = Depends(get_db)):
|
|
document = await db.get(Document, document_id)
|
|
if document is None:
|
|
raise HTTPException(status_code=404, detail="生成记录不存在")
|
|
|
|
if document.status in {"completed", "failed", "cancelled"}:
|
|
return Response(data=_serialize_document(document))
|
|
|
|
request_cancel(document_id)
|
|
return Response(data=_serialize_document(document))
|
|
|
|
|
|
@router.delete("/documents/{document_id}")
|
|
async def delete_document(document_id: int, db: AsyncSession = Depends(get_db)):
|
|
document = await db.get(Document, document_id)
|
|
if document is None:
|
|
raise HTTPException(status_code=404, detail="生成记录不存在")
|
|
|
|
result = await db.execute(select(GenerationLog).where(GenerationLog.document_id == document_id))
|
|
for log in result.scalars().all():
|
|
await db.delete(log)
|
|
|
|
await db.delete(document)
|
|
await db.commit()
|
|
return Response(data={"id": document_id})
|
|
|
|
|
|
@router.delete("/blocks/{block_id}/reset/{document_id}")
|
|
async def reset_block(block_id: int, document_id: int, db: AsyncSession = Depends(get_db)):
|
|
"""删除块的生成日志,恢复为模板默认内容。"""
|
|
block = await db.get(TemplateBlock, block_id)
|
|
if block is None:
|
|
raise HTTPException(status_code=404, detail="块不存在")
|
|
|
|
result = await db.execute(
|
|
select(GenerationLog)
|
|
.where(GenerationLog.document_id == document_id, GenerationLog.block_id == block_id)
|
|
)
|
|
logs = result.scalars().all()
|
|
for log in logs:
|
|
await db.delete(log)
|
|
await db.commit()
|
|
|
|
return Response(data={"block_id": block_id, "deleted_logs": len(logs)})
|