93 lines
2.8 KiB
Python
93 lines
2.8 KiB
Python
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
|
||
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.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),
|
||
current_user: User = Depends(get_current_user),
|
||
):
|
||
"""
|
||
主对话端点。
|
||
|
||
参数:
|
||
- 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,
|
||
user_id=current_user.id,
|
||
):
|
||
yield {
|
||
"event": event["type"],
|
||
"data": json.dumps(event["data"], ensure_ascii=False),
|
||
}
|
||
|
||
return EventSourceResponse(event_generator())
|