chat.py 22 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424
  1. import json
  2. import re
  3. import uuid
  4. from fastapi import WebSocket, WebSocketDisconnect, APIRouter, Depends
  5. import asyncio
  6. import websockets
  7. from sqlalchemy.orm import Session
  8. from Log import logger
  9. from app.api import get_current_user_websocket
  10. from app.config.config import settings
  11. from app.models.agent_model import AgentModel, AgentType
  12. from app.models.base_model import get_db
  13. from app.models.user_model import UserModel
  14. from app.service.dialog import update_session_history
  15. from app.service.basic import BasicService
  16. from app.service.difyService import DifyService
  17. from app.service.ragflow import RagflowService
  18. from app.service.service_token import get_bisheng_token, get_ragflow_token
  19. from app.service.session import SessionService
  20. router = APIRouter()
  21. # 中间层WebSocket 服务器,接收客户端的连接
  22. @router.websocket("/ws/{agent_id}/{chat_id}")
  23. async def handle_client(websocket: WebSocket,
  24. agent_id: str,
  25. chat_id: str,
  26. current_user: UserModel = Depends(get_current_user_websocket),
  27. db: Session = Depends(get_db)):
  28. tasks = []
  29. await websocket.accept()
  30. print(f"Client {agent_id} connected")
  31. agent = db.query(AgentModel).filter(AgentModel.id == agent_id).first()
  32. if not agent:
  33. ret = {"message": "Agent not found", "type": "close"}
  34. await websocket.send_json(ret)
  35. return
  36. agent_type = agent.agent_type
  37. if chat_id == "" or chat_id == "0":
  38. ret = {"message": "Chat ID not found", "type": "close"}
  39. await websocket.send_json(ret)
  40. return
  41. if agent_type == AgentType.RAGFLOW:
  42. ragflow_service = RagflowService(settings.fwr_base_url)
  43. token = get_ragflow_token(db, current_user.id)
  44. try:
  45. async def forward_to_ragflow():
  46. while True:
  47. message = await websocket.receive_json()
  48. print(f"Received from client {chat_id}: {message}")
  49. chat_history = message.get('chatHistory', [])
  50. message["role"] = "user"
  51. if len(chat_history) == 0:
  52. chat_history = await ragflow_service.get_session_history(token, chat_id)
  53. if len(chat_history) == 0:
  54. chat_history = await ragflow_service.set_session(token, agent_id,
  55. message, chat_id, True)
  56. # print("chat_history------------------------", chat_history)
  57. if len(chat_history) == 0:
  58. result = {"message": "内部错误:创建会话失败", "type": "close"}
  59. await websocket.send_json(result)
  60. await websocket.close()
  61. return
  62. else:
  63. chat_history.append({
  64. "content": message["message"],
  65. "doc_ids": message.get("doc_ids", []),
  66. "role": "user"
  67. })
  68. complete_response = ""
  69. async for rag_response in ragflow_service.chat(token, chat_id, chat_history):
  70. try:
  71. if rag_response[:5] == "data:":
  72. # 如果是,则截取掉前5个字符,并去除首尾空白符
  73. text = rag_response[5:].strip()
  74. else:
  75. # 否则,保持原样
  76. text = rag_response
  77. complete_response += text
  78. try:
  79. json_data = json.loads(complete_response)
  80. data = json_data.get("data")
  81. if data is True: # 完成输出
  82. result = {"message": "", "type": "close"}
  83. elif data is None: # 发生错误
  84. answer = json_data.get("retmsg", json_data.get("retcode"))
  85. result = {"message": "内部错误:" + answer, "type": "message"}
  86. else: # 正常输出
  87. answer = data.get("answer", "")
  88. reference = data.get("reference", {})
  89. result = {"message": answer, "type": "message", "reference": reference}
  90. await websocket.send_json(result)
  91. complete_response = ""
  92. except json.JSONDecodeError as e:
  93. print(f"Error decoding JSON: {e}")
  94. # print(f"Response text: {text}")
  95. except Exception as e2:
  96. result = {"message": f"内部错误: {e2}", "type": "close"}
  97. await websocket.send_json(result)
  98. print(f"Error process message of ragflow: {e2}")
  99. try:
  100. dialog_chat_history = await ragflow_service.get_session_history(token, chat_id, 1)
  101. await update_session_history(db, dialog_chat_history, current_user.id)
  102. except Exception as e:
  103. logger.error(e)
  104. logger.error("-----------------保存ragflow的历史会话异常-----------------")
  105. # 启动任务处理客户端消息
  106. tasks = [
  107. asyncio.create_task(forward_to_ragflow())
  108. ]
  109. await asyncio.wait(tasks, return_when=asyncio.FIRST_COMPLETED)
  110. except WebSocketDisconnect as e1:
  111. print(f"Client {chat_id} disconnected: {e1}")
  112. await websocket.close()
  113. except Exception as e:
  114. print(f"Exception occurred: {e}")
  115. finally:
  116. print("Cleaning up resources of ragflow")
  117. # 取消所有任务
  118. for task in tasks:
  119. if not task.done():
  120. task.cancel()
  121. try:
  122. await task
  123. except asyncio.CancelledError:
  124. pass
  125. elif agent_type == AgentType.BISHENG:
  126. token = get_bisheng_token(db, current_user.id)
  127. service_uri = f"{settings.sgb_websocket_url}/api/v1/assistant/chat/{agent_id}?t=&chat_id={chat_id}"
  128. headers = {'cookie': f"access_token_cookie={token};"}
  129. async with websockets.connect(service_uri, extra_headers=headers) as service_websocket:
  130. try:
  131. # 处理客户端发来的消息
  132. async def forward_to_service():
  133. while True:
  134. message = await websocket.receive_json()
  135. print(f"Received from client, {chat_id}: {message}")
  136. # 添加 'agent_id' 和 'chat_id' 字段
  137. message['flow_id'] = agent_id
  138. message['chat_id'] = chat_id
  139. msg = message["message"]
  140. del message["message"]
  141. message['inputs'] = {
  142. "data": {"chatId": chat_id, "id": agent_id, "type": "assistant"},
  143. "input": msg
  144. }
  145. await service_websocket.send(json.dumps(message))
  146. print(f"Forwarded to bisheng: {message}")
  147. # 监听毕昇发来的消息并转发给客户端
  148. async def forward_to_client():
  149. while True:
  150. message = await service_websocket.recv()
  151. print(f"Received from bisheng: {message}")
  152. data = json.loads(message)
  153. if data["type"] == "close" or data["type"] == "stream" or data["type"] == "end_cover":
  154. if data["type"] == "close":
  155. t = "close"
  156. else:
  157. t = "stream"
  158. result = {"message": data["message"], "type": t}
  159. await websocket.send_json(result)
  160. print(f"Forwarded to client, {chat_id}: {result}")
  161. # 启动两个任务,分别处理客户端和服务端的消息
  162. tasks = [
  163. asyncio.create_task(forward_to_service()),
  164. asyncio.create_task(forward_to_client())
  165. ]
  166. done, pending = await asyncio.wait(tasks, return_when=asyncio.FIRST_COMPLETED)
  167. # 取消未完成的任务
  168. for task in pending:
  169. task.cancel()
  170. try:
  171. await task
  172. except asyncio.CancelledError:
  173. pass
  174. except WebSocketDisconnect as e:
  175. print(f"WebSocket connection closed with code {e.code}: {e.reason}")
  176. await websocket.close()
  177. await service_websocket.close()
  178. except Exception as e:
  179. print(f"Exception occurred: {e}")
  180. finally:
  181. print("Cleaning up resources of bisheng")
  182. # 取消所有任务
  183. for task in tasks:
  184. if not task.done():
  185. task.cancel()
  186. try:
  187. await task
  188. except asyncio.CancelledError:
  189. pass
  190. elif agent_type == AgentType.BASIC:
  191. try:
  192. service = BasicService(base_url=settings.basic_base_url)
  193. while True:
  194. # 接收前端消息
  195. message = await websocket.receive_json()
  196. question = message.get("message")
  197. try:
  198. SessionService(db).create_session(
  199. chat_id,
  200. question,
  201. agent_id,
  202. AgentType.BASIC,
  203. current_user.id
  204. )
  205. except Exception as e:
  206. logger.error(e)
  207. if not question:
  208. await websocket.send_json({"message": "Invalid request", "type": "error"})
  209. continue
  210. logger.error(agent.type)
  211. if agent.type == "questionTalk":
  212. try:
  213. data = await service.questions_talk(question, chat_id)
  214. output = data.get("output", "")
  215. file_name = data.get("filename", "")
  216. excel_url = None
  217. if file_name:
  218. excel_url = f"/api/files/download/?agent_id=basic_question_talk&file_id={file_name}&file_type=word"
  219. result = {"message": output, "type": "message", "file_url": excel_url, "file_name":file_name}
  220. try:
  221. SessionService(db).update_session(chat_id,
  222. message={"role": "assistant", "content": result})
  223. except Exception as e:
  224. logger.error(e)
  225. logger.error("-----------------返回数据--------------------")
  226. await websocket.send_json(result)
  227. except Exception as e2:
  228. result = {"message": f"内部错误: {e2}", "type": "close"}
  229. logger.error(str(e2))
  230. logger.error(f"Error process message of basic chuti agent: {e2}")
  231. await websocket.send_json(result)
  232. else:
  233. logger.error("---------------------excel_talk-----------------------------")
  234. async for data in service.excel_talk(question, chat_id):
  235. logger.error(data)
  236. output = data.get("output", "")
  237. excel_name = data.get("excel_name", "")
  238. image_name = data.get("image_name", "")
  239. def build_file_url(name, file_type):
  240. if not name:
  241. return None
  242. return (f"/api/files/download/?agent_id={agent_id}&file_id={name}"
  243. f"&file_type={file_type}")
  244. excel_url = build_file_url(excel_name, 'excel')
  245. image_url = build_file_url(image_name, 'image')
  246. if excel_url or data.get("e", ""):
  247. try:
  248. SessionService(db).update_session(chat_id,
  249. message={
  250. "content": output,
  251. "excel_url": excel_url,
  252. "image_url": image_url,
  253. "sql": data.get("sql", ""),
  254. "code": data.get("code", ""),
  255. "e": data.get("e", ""),
  256. "role": "assistant"})
  257. except Exception as e:
  258. logger.error(f"Unexpected error when update_session: {e}")
  259. # 发送结果给客户端
  260. data["type"] = "message"
  261. data["message"] = output
  262. data["excel_url"] = excel_url
  263. data["image_url"] = image_url
  264. await websocket.send_json(data)
  265. except Exception as e:
  266. logger.error(e)
  267. await websocket.send_json({"message": "出现错误!", "type": "error"})
  268. finally:
  269. await websocket.close()
  270. print(f"Client {agent_id} disconnected")
  271. if agent_type == AgentType.DIFY:
  272. dify_service = DifyService(settings.dify_base_url)
  273. # token = get_dify_token(db, current_user.id)
  274. token = settings.dify_api_token
  275. try:
  276. async def forward_to_dify():
  277. while True:
  278. image_list = []
  279. is_image = False
  280. conversation_id = ""
  281. receive_message = await websocket.receive_json()
  282. print(f"Received from client {chat_id}: {receive_message}")
  283. upload_file_id = receive_message.get('upload_file_id', "")
  284. question = receive_message.get('message', "")
  285. if not question and not image_url:
  286. await websocket.send_json({"message": "Invalid request", "type": "error"})
  287. continue
  288. try:
  289. session = SessionService(db).create_session(
  290. chat_id,
  291. question,
  292. agent_id,
  293. AgentType.DIFY,
  294. current_user.id
  295. )
  296. conversation_id = session.conversation_id
  297. except Exception as e:
  298. logger.error(e)
  299. # complete_response = ""
  300. answer_str = ""
  301. async for rag_response in dify_service.chat(token, current_user.id, question, upload_file_id, conversation_id):
  302. # print("=============================================")
  303. # print(rag_response)
  304. try:
  305. if rag_response[:5] == "data:":
  306. # 如果是,则截取掉前5个字符,并去除首尾空白符
  307. complete_response = rag_response[5:].strip()
  308. else:
  309. # 否则,保持原样
  310. complete_response = rag_response
  311. # complete_response += text
  312. try:
  313. data = json.loads(complete_response)
  314. complete_response = ""
  315. # data = json_data.get("data")
  316. if data.get("event") == "agent_message":# "event": "message_end"
  317. if "answer" not in data or not data["answer"]: # 信息过滤
  318. logger.error("非法数据--------------------")
  319. # logger.error(data)
  320. continue
  321. else: # 正常输出
  322. answer = data.get("answer", "")
  323. if isinstance(answer, str):
  324. if "![](https://res.stepfun.com/" in answer and image_list:
  325. is_image = True
  326. pattern = r'!\[\] *\(https://res\.stepfun\.com/image_gen/[^)]+\)'
  327. url_image = image_list.pop()
  328. new_answer = re.sub(pattern, url_image, answer)
  329. answer_str += new_answer
  330. else:
  331. answer_str += answer
  332. elif isinstance(answer, dict):
  333. logger.error("未知数据体:0---------------------------------")
  334. logger.error(answer)
  335. answer_str += answer.get("action_input", "")
  336. result = {"message": answer_str, "type": "message"}
  337. elif data.get("event") == "message_end":
  338. images_url = []
  339. # res_msg = await dify_service.get_session_history(token, data.get("conversation_id"), str(current_user.id))
  340. # if len(res_msg) > 0:
  341. # message_files = res_msg[-1].get("message_files")
  342. # for msg_file in message_files:
  343. # await dify_service.save_images(msg_file.get("url"), msg_file.get("id")+".png")
  344. # images_url.append(msg_file.get("id"))
  345. # result = {"message": answer_str, "type": "close"} # , "message_files": images_url
  346. if image_list and not is_image:
  347. answer_str += image_list[-1]
  348. result = {"message": answer_str,
  349. "type": "close"} # , "message_files": images_url
  350. try:
  351. SessionService(db).update_session(chat_id,
  352. message={"role": "assistant", "content": {"answer":answer_str, "images":images_url}},conversation_id=data.get("conversation_id"))
  353. except Exception as e:
  354. logger.error("保存dify的会话异常!")
  355. logger.error(e)
  356. elif data.get("event") == "message_file":
  357. await dify_service.save_images(data.get("url"), data.get("id") + ".png")
  358. image_list.append(f"![](/api/files/image/{data.get('id')})")
  359. # result = {"message": answer_str, "type": "message"}
  360. continue
  361. else:
  362. continue
  363. await websocket.send_json(result)
  364. complete_response = ""
  365. except json.JSONDecodeError as e:
  366. print(f"Error decoding JSON: {e}")
  367. # print(f"Response text: {text}")
  368. except Exception as e2:
  369. result = {"message": f"内部错误: {e2}", "type": "close"}
  370. await websocket.send_json(result)
  371. print(f"Error process message of ragflow: {e2}")
  372. # 启动任务处理客户端消息
  373. tasks = [
  374. asyncio.create_task(forward_to_dify())
  375. ]
  376. await asyncio.wait(tasks, return_when=asyncio.FIRST_COMPLETED)
  377. except WebSocketDisconnect as e1:
  378. print(f"Client {chat_id} disconnected: {e1}")
  379. await websocket.close()
  380. except Exception as e:
  381. print(f"Exception occurred: {e}")
  382. finally:
  383. print("Cleaning up resources of ragflow")
  384. # 取消所有任务
  385. for task in tasks:
  386. if not task.done():
  387. task.cancel()
  388. try:
  389. await task
  390. except asyncio.CancelledError:
  391. pass
  392. else:
  393. ret = {"message": "Agent not found", "type": "close"}
  394. await websocket.send_json(ret)