接入真实模型调用与参考文件上传

This commit is contained in:
zwt13703
2026-07-02 15:16:34 +08:00
parent b313083766
commit 5a43cc70d5
9 changed files with 420 additions and 53 deletions
+86 -20
View File
@@ -1,17 +1,23 @@
import json
import time
import os
import uuid
from datetime import datetime
from fastapi import APIRouter, Depends, HTTPException, Query
from fastapi import APIRouter, Depends, File, HTTPException, Query, UploadFile
from sqlalchemy import func, select
from sqlalchemy.ext.asyncio import AsyncSession
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.template import Template
from schemas.schemas import GenerateFullRequest, GenerateTestRequest, Response
from services.ai_service import call_ai
from services.minio_client import upload_bytes
router = APIRouter()
@@ -60,18 +66,72 @@ def _build_mock_content(paragraph: Paragraph) -> dict:
return {"content": blocks}
async def _get_effective_model(db: AsyncSession, paragraph: Paragraph) -> AiModel | None:
if paragraph.model_id:
model = await db.get(AiModel, paragraph.model_id)
if model is not None and model.status == "enabled":
return model
result = await db.execute(
select(AiModel).where(AiModel.status == "enabled").order_by(AiModel.id.asc()).limit(1)
)
return result.scalars().first()
@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="段落不存在")
content = _build_mock_content(paragraph)
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":
content = _build_mock_content(paragraph)
message = "当前未找到可用模型,返回本地模拟生成结果。"
else:
result = await call_ai(paragraph, model)
content = result.content
message = f"已通过模型 {result.used_model} 生成。"
return Response(
data={
"paragraph_id": paragraph.id,
"content": content,
"message": "当前返回本地模拟生成结果,便于前端联调。",
"message": message,
}
)
@router.post("/upload")
async def upload_reference_file(file: UploadFile = File(...)):
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="文件类型不支持")
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",
)
return Response(
data={
"file_name": file.filename,
"file_path": f"{settings.MINIO_BUCKET_UPLOADS}/{object_name}",
}
)
@@ -104,40 +164,46 @@ async def generate_full(body: GenerateFullRequest, db: AsyncSession = Depends(ge
await db.flush()
done_count = 0
failed_count = 0
for paragraph in paragraphs:
if paragraph.edit_mode == "manual":
content = {"content": [{"type": "text", "text": paragraph.content or "该段落为人工编辑模式。"}]}
status = "success"
duration = 0
error_message = ""
else:
start = time.perf_counter()
content = _build_mock_content(paragraph)
model = await _get_effective_model(db, paragraph)
try:
if model is None:
content = _build_mock_content(paragraph)
else:
result = await call_ai(paragraph, model)
content = result.content
status = "success"
error_message = ""
except Exception as error:
content = _build_mock_content(paragraph)
status = "failed"
error_message = str(error)
failed_count += 1
duration = round(time.perf_counter() - start, 4)
log = GenerationLog(
document_id=document.id,
paragraph_id=paragraph.id,
model_id=paragraph.model_id,
status="success",
content=json.dumps(content, ensure_ascii=False),
duration=duration,
error_msg="",
)
db.add(log)
done_count += 1
continue
log = GenerationLog(
document_id=document.id,
paragraph_id=paragraph.id,
model_id=paragraph.model_id,
status="success",
status=status,
content=json.dumps(content, ensure_ascii=False),
duration=0,
error_msg="",
duration=duration,
error_msg=error_message,
)
db.add(log)
done_count += 1
document.para_count_done = done_count
document.status = "completed"
document.status = "completed" if failed_count == 0 else "failed"
document.error = "" if failed_count == 0 else f"{failed_count} 个段落生成失败,已回退为模拟结果。"
document.file_path = f"mock://document/{document.id}"
await db.commit()
await db.refresh(document)
+24 -7
View File
@@ -5,6 +5,7 @@ from sqlalchemy.ext.asyncio import AsyncSession
from database import get_db
from models.ai_model import AiModel
from schemas.schemas import AiModelCreate, AiModelUpdate, Response
from services.ai_service import call_ai
from services.security import decrypt_text, encrypt_text, mask_secret
router = APIRouter()
@@ -88,10 +89,26 @@ async def test_model(model_id: int, db: AsyncSession = Depends(get_db)):
if model is None:
raise HTTPException(status_code=404, detail="模型不存在")
return Response(
data={
"id": model.id,
"success": True,
"message": f"模型 {model.name} 配置校验通过(当前为本地模拟测试)",
}
)
class FakeParagraph:
title = "连接测试"
content = "请返回一段非常简短的测试文本。"
need_prompt = False
prompt_text = ""
output_format = "text"
try:
result = await call_ai(FakeParagraph(), model)
return Response(
data={
"id": model.id,
"success": True,
"message": f"模型 {model.name} 连接测试成功",
"preview": result.content,
}
)
except Exception as error:
return Response(
code=-1,
message=str(error),
data={"id": model.id, "success": False},
)