Files
EPEEAIKit/art-agent/backend/app/api/chat.py

114 lines
3.5 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
import json
import uuid
from pathlib import Path
from typing import Optional
from fastapi import APIRouter, Depends, File, Form, UploadFile
from sse_starlette.sse import EventSourceResponse
from app.agent.loop import run_agent_loop
from app.auth import get_current_user
from app.config import (
get_image_models_list,
get_default_image_model_id,
get_llm_models_list,
get_default_llm_model_id,
)
from app.db import User
router = APIRouter()
UPLOADS_DIR = Path(__file__).parent.parent.parent / "uploads"
async def _save_upload(file: UploadFile) -> str:
"""保存上传的参考图,返回可访问的 URL 路径。"""
ext = Path(file.filename).suffix or ".png"
filename = f"{uuid.uuid4().hex}{ext}"
filepath = UPLOADS_DIR / filename
content = await file.read()
filepath.write_bytes(content)
return f"/uploads/{filename}"
@router.post("/upload-ref-image")
async def upload_ref_image(
file: UploadFile = File(...),
current_user: User = Depends(get_current_user),
):
"""
独立的参考图上传端点。
前端选图后立即调用,返回服务端路径,供后续发消息时引用。
"""
UPLOADS_DIR.mkdir(parents=True, exist_ok=True)
url = await _save_upload(file)
return {"url": url, "filename": file.filename}
@router.get("/models")
async def list_models(current_user: User = Depends(get_current_user)):
"""返回可用的图像生成模型列表。"""
return {
"models": get_image_models_list(),
"default": get_default_image_model_id(),
}
@router.get("/llm-models")
async def list_llm_models(current_user: User = Depends(get_current_user)):
"""返回可用的 LLM 对话模型列表。"""
return {
"models": get_llm_models_list(),
"default": get_default_llm_model_id(),
}
@router.post("/chat")
async def chat(
messages: str = Form(...),
ref_image: Optional[UploadFile] = File(None),
ref_image_url: Optional[str] = Form(None),
ref_image_urls: Optional[str] = Form(None),
image_model: Optional[str] = Form(None),
session_id: Optional[str] = Form(None),
llm_model: Optional[str] = Form(None),
current_user: User = Depends(get_current_user),
):
"""
主对话端点。
参数:
- messages: JSON 字符串,对话历史 [{role, content}]
- ref_image: 可选的参考图文件(兼容旧方式)
- ref_image_url: 兼容旧方式,单张参考图服务端路径
- ref_image_urls: JSON 数组字符串,多张参考图服务端路径列表
- image_model: 可选,指定本次使用的生图模型短 ID
- session_id: 可选,前端会话 ID用于 Mem0 记忆作用域
- llm_model: 可选,指定本次使用的 LLM 模型短 ID
"""
parsed_messages = json.loads(messages)
resolved_ref_urls: list[str] = []
if ref_image_urls:
resolved_ref_urls = json.loads(ref_image_urls)
elif ref_image_url:
resolved_ref_urls = [ref_image_url]
elif ref_image and ref_image.filename:
resolved_ref_urls = [await _save_upload(ref_image)]
async def event_generator():
async for event in run_agent_loop(
parsed_messages,
resolved_ref_urls or None,
image_model=image_model,
session_id=session_id,
user_id=current_user.id,
llm_model=llm_model,
):
yield {
"event": event["type"],
"data": json.dumps(event["data"], ensure_ascii=False),
}
return EventSourceResponse(event_generator())