chat.py 7.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189
  1. """旅游AI对话API路由(SSE流式输出)"""
  2. import json
  3. import asyncio
  4. from fastapi import APIRouter, HTTPException, Request
  5. from fastapi.responses import StreamingResponse
  6. from ...models.schemas import ChatSessionResponse, ChatSessionListResponse, ChatMessagesResponse, ChatSendMessageRequest, ChatDeleteResponse
  7. from ...database import (
  8. create_chat_session, list_chat_sessions, get_chat_session,
  9. delete_chat_session, add_chat_message, get_chat_messages,
  10. update_chat_session_title
  11. )
  12. from ...services.travel_chat_service import get_travel_chat_service
  13. from ...services.user_profile_service import load_profile_text, extract_and_update_profile, get_cross_session_context
  14. from .auth import require_auth
  15. router = APIRouter(prefix="/chat", tags=["旅游AI对话"])
  16. def _require_auth(request: Request) -> dict:
  17. """统一鉴权"""
  18. try:
  19. return require_auth(request)
  20. except HTTPException:
  21. raise HTTPException(status_code=401, detail="请先登录后再使用AI对话")
  22. @router.post("/sessions", summary="创建新会话")
  23. async def create_session(request: Request):
  24. """创建一个新的聊天会话"""
  25. user = _require_auth(request)
  26. session = create_chat_session(user["id"])
  27. return ChatSessionResponse(success=True, session=session)
  28. @router.get("/sessions", summary="获取会话列表")
  29. async def list_sessions(request: Request):
  30. """获取当前用户的所有会话"""
  31. user = _require_auth(request)
  32. sessions = list_chat_sessions(user["id"])
  33. return ChatSessionListResponse(success=True, sessions=sessions)
  34. @router.get("/sessions/{session_id}", summary="获取会话详情")
  35. async def get_session(session_id: int, request: Request):
  36. """获取单个会话信息"""
  37. user = _require_auth(request)
  38. session = get_chat_session(session_id, user["id"])
  39. if not session:
  40. raise HTTPException(status_code=404, detail="会话不存在")
  41. return ChatSessionResponse(success=True, session=session)
  42. @router.delete("/sessions/{session_id}", summary="删除会话")
  43. async def delete_session(session_id: int, request: Request):
  44. """删除会话及其所有消息"""
  45. user = _require_auth(request)
  46. deleted = delete_chat_session(session_id, user["id"])
  47. if not deleted:
  48. raise HTTPException(status_code=404, detail="会话不存在")
  49. return ChatDeleteResponse(success=True, message="会话已删除")
  50. @router.get("/sessions/{session_id}/messages", summary="获取会话消息")
  51. async def get_messages(session_id: int, request: Request):
  52. """获取会话的所有聊天消息"""
  53. user = _require_auth(request)
  54. session = get_chat_session(session_id, user["id"])
  55. if not session:
  56. raise HTTPException(status_code=404, detail="会话不存在")
  57. messages = get_chat_messages(session_id)
  58. return ChatMessagesResponse(success=True, messages=messages)
  59. @router.post("/sessions/{session_id}/messages", summary="发送消息(流式)")
  60. async def send_message(session_id: int, req: ChatSendMessageRequest, request: Request):
  61. """
  62. 发送消息并流式获取AI回复(SSE格式)
  63. 流式返回 SSE 事件:
  64. - data: {"type": "token", "content": "文本片段"}
  65. - data: {"type": "error", "content": "错误信息"}
  66. - data: {"type": "done", "title": "更新后的会话标题"}
  67. """
  68. user = _require_auth(request)
  69. session = get_chat_session(session_id, user["id"])
  70. if not session:
  71. raise HTTPException(status_code=404, detail="会话不存在")
  72. content = req.content.strip()
  73. if not content:
  74. raise HTTPException(status_code=400, detail="消息不能为空")
  75. # 1. 保存用户消息
  76. add_chat_message(session_id, "user", content)
  77. # 2. 获取历史消息(作为上下文)
  78. history = get_chat_messages(session_id)
  79. # 3. 如果是会话首条消息,加载用户画像并包装为 XML 标签用户消息
  80. profile_message = ""
  81. if len(history) <= 1:
  82. profile_text = load_profile_text(user["id"])
  83. if profile_text:
  84. profile_message = (
  85. f"<user_profile>\n{profile_text}\n"
  86. f"(注意:如果我现在说的与上述偏好不一致,请以我当前说的为准。)\n"
  87. f"</user_profile>"
  88. )
  89. print(f" 👤 用户 {user['id']} 已加载画像上下文")
  90. # 4. 返回流式响应
  91. return StreamingResponse(
  92. _stream_ai_response(user["id"], session_id, content, history, profile_message),
  93. media_type="text/event-stream",
  94. headers={
  95. "Cache-Control": "no-cache",
  96. "Connection": "keep-alive",
  97. "X-Accel-Buffering": "no",
  98. }
  99. )
  100. async def _stream_ai_response(user_id: int, session_id: int, content: str, history: list, profile_message: str = ""):
  101. """流式生成AI回复的SSE事件"""
  102. travel_chat = get_travel_chat_service()
  103. full_response = ""
  104. try:
  105. # 获取流式生成器(携带用户画像消息)
  106. stream = travel_chat.chat_stream(
  107. user_message=content,
  108. history=history[:-1], # 排除刚保存的最后一条
  109. profile_message=profile_message,
  110. )
  111. for chunk in stream:
  112. if chunk:
  113. full_response += chunk
  114. # 发送 token 事件
  115. yield f"data: {json.dumps({'type': 'token', 'content': chunk}, ensure_ascii=False)}\n\n"
  116. # 流式完成 - 保存AI回复到数据库
  117. add_chat_message(session_id, "assistant", full_response)
  118. # 如果是第一条消息,自动生成会话标题
  119. title = None
  120. if len(history) <= 1:
  121. title = _generate_title(content)
  122. update_chat_session_title(session_id, title)
  123. # 发送完成事件(先发送,不阻塞)
  124. done_event = {"type": "done"}
  125. if title:
  126. done_event["title"] = title
  127. yield f"data: {json.dumps(done_event, ensure_ascii=False)}\n\n"
  128. # 后台异步执行画像提取,不阻塞主进程(SSE 流已关闭)
  129. asyncio.create_task(
  130. asyncio.to_thread(_run_profile_extraction, user_id, content, history)
  131. )
  132. except Exception as e:
  133. error_msg = f"抱歉,AI暂时无法回答您的问题,请稍后重试。"
  134. # 尝试发送错误事件
  135. yield f"data: {json.dumps({'type': 'error', 'content': error_msg}, ensure_ascii=False)}\n\n"
  136. yield f"data: {json.dumps({'type': 'done'}, ensure_ascii=False)}\n\n"
  137. def _generate_title(user_message: str) -> str:
  138. """根据用户第一条消息生成会话标题"""
  139. title = user_message.strip()[:20]
  140. if len(user_message) > 20:
  141. title += "..."
  142. return title
  143. def _run_profile_extraction(user_id: int, content: str, history: list):
  144. """在后台线程中同步执行画像提取(不阻塞 SSE 流主进程)
  145. 由 asyncio.to_thread 调度到线程池执行,避免阻塞事件循环。
  146. """
  147. try:
  148. cross_ctx = get_cross_session_context(user_id, max_sessions=5, max_messages=6)
  149. extract_and_update_profile(
  150. user_id, content, history,
  151. cross_session_context=cross_ctx,
  152. )
  153. except Exception:
  154. pass # 画像提取失败不影响主流程