接入真实模型调用与参考文件上传
This commit is contained in:
+86
-20
@@ -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)
|
||||
|
||||
@@ -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},
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user