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

88 lines
2.6 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, File, Form, UploadFile
from sse_starlette.sse import EventSourceResponse
from app.agent.loop import run_agent_loop
from app.config import get_image_models_list, get_default_image_model_id
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(...),
):
"""
独立的参考图上传端点。
前端选图后立即调用,返回服务端路径,供后续发消息时引用。
"""
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():
"""返回可用的图像生成模型列表。"""
return {
"models": get_image_models_list(),
"default": get_default_image_model_id(),
}
@router.post("/chat")
async def chat(
messages: str = Form(...),
ref_image: Optional[UploadFile] = File(None),
ref_image_url: Optional[str] = Form(None),
image_model: Optional[str] = Form(None),
session_id: Optional[str] = Form(None),
):
"""
主对话端点。
参数:
- messages: JSON 字符串,对话历史 [{role, content}]
- ref_image: 可选的参考图文件(兼容旧方式)
- ref_image_url: 可选,已通过 /upload-ref-image 上传后的服务端路径
- image_model: 可选,指定本次使用的生图模型短 ID
- session_id: 可选,前端会话 ID用于 Mem0 记忆作用域
"""
parsed_messages = json.loads(messages)
resolved_ref_url: Optional[str] = None
if ref_image_url:
resolved_ref_url = ref_image_url
elif ref_image and ref_image.filename:
resolved_ref_url = await _save_upload(ref_image)
async def event_generator():
async for event in run_agent_loop(
parsed_messages,
resolved_ref_url,
image_model=image_model,
session_id=session_id,
):
yield {
"event": event["type"],
"data": json.dumps(event["data"], ensure_ascii=False),
}
return EventSourceResponse(event_generator())