chat.py 6.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139
  1. import json
  2. import uuid
  3. from fastapi import WebSocket, WebSocketDisconnect, APIRouter, Depends
  4. import asyncio
  5. import websockets
  6. from sqlalchemy.orm import Session
  7. from app.api import get_current_user_websocket
  8. from app.config.config import settings
  9. from app.models.base_model import get_db
  10. from app.models.user_model import UserModel
  11. from app.service.ragflow import RagflowService
  12. from app.service.token import get_bisheng_token, get_ragflow_token
  13. router = APIRouter()
  14. # 存储客户端 WebSocket 连接
  15. client_websockets = {}
  16. # 中间层WebSocket 服务器,接收客户端的连接
  17. @router.websocket("/ws/{agent_id}/{chat_id}")
  18. async def handle_client(websocket: WebSocket,
  19. agent_id: str,
  20. chat_id: str,
  21. current_user: UserModel = Depends(get_current_user_websocket),
  22. db: Session = Depends(get_db)):
  23. await websocket.accept()
  24. print(f"Client {agent_id} connected")
  25. if agent_id == "0":
  26. agent_id = settings.bisheng_agent_id
  27. elif agent_id == "1":
  28. agent_id = settings.ragflow_agent_id
  29. chat_id = settings.ragflow_chat_id
  30. if chat_id == "0":
  31. chat_id = uuid.uuid4().hex
  32. client_websockets[chat_id] = websocket
  33. if agent_id == settings.ragflow_agent_id:
  34. ragflow_service = RagflowService(settings.ragflow_base_url)
  35. token = get_ragflow_token(db, current_user.id)
  36. try:
  37. async def forward_to_ragflow():
  38. while True:
  39. message = await websocket.receive_json()
  40. print(f"Received from client {chat_id}: {message}")
  41. async for rag_response in ragflow_service.chat(token, chat_id, message["chatHistory"]):
  42. try:
  43. print(f"Received from ragflow: {rag_response}")
  44. json_str = rag_response[5:].strip()
  45. json_data = json.loads(json_str)
  46. data = json_data.get("data")
  47. if data is True: # 完成输出
  48. result = {"message": "", "type": "close"}
  49. elif data is None: # 发生错误
  50. answer = json_data.get("retmsg", json_data.get("retcode"))
  51. result = {"message": "内部错误:" + answer, "type": "stream"}
  52. else: # 正常输出
  53. answer = json_data.get("data", {}).get("answer", "")
  54. result = {"message": answer, "type": "stream"}
  55. await websocket.send_json(result)
  56. print(f"Forwarded to client {chat_id}: {result}")
  57. except Exception as e:
  58. result = {"message": f"内部错误: {e}", "type": "close"}
  59. await websocket.send_json(result)
  60. print(f"Error process message of ragflow: {e}")
  61. # 启动任务处理客户端消息
  62. tasks = [
  63. asyncio.create_task(forward_to_ragflow())
  64. ]
  65. await asyncio.wait(tasks, return_when=asyncio.FIRST_COMPLETED)
  66. except WebSocketDisconnect:
  67. print(f"Client {chat_id} disconnected")
  68. finally:
  69. del client_websockets[chat_id]
  70. else:
  71. token = get_bisheng_token(db, current_user.id)
  72. service_uri = f"{settings.bisheng_websocket_url}/api/v1/assistant/chat/{agent_id}?t=&chat_id={chat_id}"
  73. headers = {'cookie': f"access_token_cookie={token};"}
  74. async with websockets.connect(service_uri, extra_headers=headers) as service_websocket:
  75. try:
  76. # 处理客户端发来的消息
  77. async def forward_to_service():
  78. while True:
  79. message = await websocket.receive_json()
  80. print(f"Received from client, {chat_id}: {message}")
  81. # 添加 'agent_id' 和 'chat_id' 字段
  82. message['flow_id'] = agent_id
  83. message['chat_id'] = chat_id
  84. msg = message["message"]
  85. del message["message"]
  86. message['inputs'] = {
  87. "data": {"chatId": chat_id, "id": agent_id, "type": "assistant"},
  88. "input": msg
  89. }
  90. await service_websocket.send(json.dumps(message))
  91. print(f"Forwarded to bisheng: {message}")
  92. # 监听毕昇发来的消息并转发给客户端
  93. async def forward_to_client():
  94. while True:
  95. message = await service_websocket.recv()
  96. print(f"Received from bisheng: {message}")
  97. data = json.loads(message)
  98. if data["type"] == "close" or data["type"] == "stream" or data["type"] == "end_cover":
  99. if data["type"] == "close":
  100. t = "close"
  101. else:
  102. t = "stream"
  103. result = {"message": data["message"], "type": t}
  104. await websocket.send_json(result)
  105. print(f"Forwarded to client, {chat_id}: {result}")
  106. # 启动两个任务,分别处理客户端和服务端的消息
  107. tasks = [
  108. asyncio.create_task(forward_to_service()),
  109. asyncio.create_task(forward_to_client())
  110. ]
  111. done, pending = await asyncio.wait(tasks, return_when=asyncio.FIRST_COMPLETED)
  112. # 取消未完成的任务
  113. for task in pending:
  114. task.cancel()
  115. try:
  116. await task
  117. except asyncio.CancelledError:
  118. pass
  119. except WebSocketDisconnect:
  120. print(f"Client {chat_id} disconnected")
  121. finally:
  122. del client_websockets[chat_id]