| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189 |
- """旅游AI对话API路由(SSE流式输出)"""
- import json
- import asyncio
- from fastapi import APIRouter, HTTPException, Request
- from fastapi.responses import StreamingResponse
- from ...models.schemas import ChatSessionResponse, ChatSessionListResponse, ChatMessagesResponse, ChatSendMessageRequest, ChatDeleteResponse
- from ...database import (
- create_chat_session, list_chat_sessions, get_chat_session,
- delete_chat_session, add_chat_message, get_chat_messages,
- update_chat_session_title
- )
- from ...services.travel_chat_service import get_travel_chat_service
- from ...services.user_profile_service import load_profile_text, extract_and_update_profile, get_cross_session_context
- from .auth import require_auth
- router = APIRouter(prefix="/chat", tags=["旅游AI对话"])
- def _require_auth(request: Request) -> dict:
- """统一鉴权"""
- try:
- return require_auth(request)
- except HTTPException:
- raise HTTPException(status_code=401, detail="请先登录后再使用AI对话")
- @router.post("/sessions", summary="创建新会话")
- async def create_session(request: Request):
- """创建一个新的聊天会话"""
- user = _require_auth(request)
- session = create_chat_session(user["id"])
- return ChatSessionResponse(success=True, session=session)
- @router.get("/sessions", summary="获取会话列表")
- async def list_sessions(request: Request):
- """获取当前用户的所有会话"""
- user = _require_auth(request)
- sessions = list_chat_sessions(user["id"])
- return ChatSessionListResponse(success=True, sessions=sessions)
- @router.get("/sessions/{session_id}", summary="获取会话详情")
- async def get_session(session_id: int, request: Request):
- """获取单个会话信息"""
- user = _require_auth(request)
- session = get_chat_session(session_id, user["id"])
- if not session:
- raise HTTPException(status_code=404, detail="会话不存在")
- return ChatSessionResponse(success=True, session=session)
- @router.delete("/sessions/{session_id}", summary="删除会话")
- async def delete_session(session_id: int, request: Request):
- """删除会话及其所有消息"""
- user = _require_auth(request)
- deleted = delete_chat_session(session_id, user["id"])
- if not deleted:
- raise HTTPException(status_code=404, detail="会话不存在")
- return ChatDeleteResponse(success=True, message="会话已删除")
- @router.get("/sessions/{session_id}/messages", summary="获取会话消息")
- async def get_messages(session_id: int, request: Request):
- """获取会话的所有聊天消息"""
- user = _require_auth(request)
- session = get_chat_session(session_id, user["id"])
- if not session:
- raise HTTPException(status_code=404, detail="会话不存在")
- messages = get_chat_messages(session_id)
- return ChatMessagesResponse(success=True, messages=messages)
- @router.post("/sessions/{session_id}/messages", summary="发送消息(流式)")
- async def send_message(session_id: int, req: ChatSendMessageRequest, request: Request):
- """
- 发送消息并流式获取AI回复(SSE格式)
- 流式返回 SSE 事件:
- - data: {"type": "token", "content": "文本片段"}
- - data: {"type": "error", "content": "错误信息"}
- - data: {"type": "done", "title": "更新后的会话标题"}
- """
- user = _require_auth(request)
- session = get_chat_session(session_id, user["id"])
- if not session:
- raise HTTPException(status_code=404, detail="会话不存在")
- content = req.content.strip()
- if not content:
- raise HTTPException(status_code=400, detail="消息不能为空")
- # 1. 保存用户消息
- add_chat_message(session_id, "user", content)
- # 2. 获取历史消息(作为上下文)
- history = get_chat_messages(session_id)
- # 3. 如果是会话首条消息,加载用户画像并包装为 XML 标签用户消息
- profile_message = ""
- if len(history) <= 1:
- profile_text = load_profile_text(user["id"])
- if profile_text:
- profile_message = (
- f"<user_profile>\n{profile_text}\n"
- f"(注意:如果我现在说的与上述偏好不一致,请以我当前说的为准。)\n"
- f"</user_profile>"
- )
- print(f" 👤 用户 {user['id']} 已加载画像上下文")
- # 4. 返回流式响应
- return StreamingResponse(
- _stream_ai_response(user["id"], session_id, content, history, profile_message),
- media_type="text/event-stream",
- headers={
- "Cache-Control": "no-cache",
- "Connection": "keep-alive",
- "X-Accel-Buffering": "no",
- }
- )
- async def _stream_ai_response(user_id: int, session_id: int, content: str, history: list, profile_message: str = ""):
- """流式生成AI回复的SSE事件"""
- travel_chat = get_travel_chat_service()
- full_response = ""
- try:
- # 获取流式生成器(携带用户画像消息)
- stream = travel_chat.chat_stream(
- user_message=content,
- history=history[:-1], # 排除刚保存的最后一条
- profile_message=profile_message,
- )
- for chunk in stream:
- if chunk:
- full_response += chunk
- # 发送 token 事件
- yield f"data: {json.dumps({'type': 'token', 'content': chunk}, ensure_ascii=False)}\n\n"
- # 流式完成 - 保存AI回复到数据库
- add_chat_message(session_id, "assistant", full_response)
- # 如果是第一条消息,自动生成会话标题
- title = None
- if len(history) <= 1:
- title = _generate_title(content)
- update_chat_session_title(session_id, title)
- # 发送完成事件(先发送,不阻塞)
- done_event = {"type": "done"}
- if title:
- done_event["title"] = title
- yield f"data: {json.dumps(done_event, ensure_ascii=False)}\n\n"
- # 后台异步执行画像提取,不阻塞主进程(SSE 流已关闭)
- asyncio.create_task(
- asyncio.to_thread(_run_profile_extraction, user_id, content, history)
- )
- except Exception as e:
- error_msg = f"抱歉,AI暂时无法回答您的问题,请稍后重试。"
- # 尝试发送错误事件
- yield f"data: {json.dumps({'type': 'error', 'content': error_msg}, ensure_ascii=False)}\n\n"
- yield f"data: {json.dumps({'type': 'done'}, ensure_ascii=False)}\n\n"
- def _generate_title(user_message: str) -> str:
- """根据用户第一条消息生成会话标题"""
- title = user_message.strip()[:20]
- if len(user_message) > 20:
- title += "..."
- return title
- def _run_profile_extraction(user_id: int, content: str, history: list):
- """在后台线程中同步执行画像提取(不阻塞 SSE 流主进程)
- 由 asyncio.to_thread 调度到线程池执行,避免阻塞事件循环。
- """
- try:
- cross_ctx = get_cross_session_context(user_id, max_sessions=5, max_messages=6)
- extract_and_update_profile(
- user_id, content, history,
- cross_session_context=cross_ctx,
- )
- except Exception:
- pass # 画像提取失败不影响主流程
|