完善模型测试与参考文件历史能力

This commit is contained in:
zwt13703
2026-07-02 17:37:35 +08:00
parent 8743a02110
commit d3530ab0b7
16 changed files with 697 additions and 32 deletions
+136 -8
View File
@@ -15,9 +15,10 @@ 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 schemas.schemas import GenerateFullRequest, GenerateTestRequest, Response
from services.ai_service import call_ai
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,
@@ -45,6 +46,17 @@ def _serialize_document(document: Document) -> dict:
"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,
}
@router.post("/test")
async def generate_test(body: GenerateTestRequest, db: AsyncSession = Depends(get_db)):
paragraph = await db.get(Paragraph, body.paragraph_id)
@@ -64,6 +76,7 @@ async def generate_test(body: GenerateTestRequest, db: AsyncSession = Depends(ge
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} 生成。"
@@ -77,14 +90,101 @@ async def generate_test(body: GenerateTestRequest, db: AsyncSession = Depends(ge
)
@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_summaries = await asyncio.to_thread(summarize_minio_files, body.file_paths or []) 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(...)):
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:
raise HTTPException(status_code=400, detail="文件类型不支持")
allowed = " / ".join(settings.ALLOWED_EXTENSIONS)
raise HTTPException(status_code=400, detail=f"文件类型不支持:{ext or '无扩展名'}。当前支持:{allowed}")
content = await file.read()
if not content:
@@ -100,12 +200,40 @@ async def upload_reference_file(file: UploadFile = File(...)):
content,
file.content_type or "application/octet-stream",
)
return Response(
data={
"file_name": file.filename,
"file_path": f"{settings.MINIO_BUCKET_UPLOADS}/{object_name}",
}
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.post("/full")