完善模型测试与参考文件历史能力
This commit is contained in:
+136
-8
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user