chat.py 68 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364656667686970717273747576777879808182838485868788899091929394959697989910010110210310410510610710810911011111211311411511611711811912012112212312412512612712812913013113213313413513613713813914014114214314414514614714814915015115215315415515615715815916016116216316416516616716816917017117217317417517617717817918018118218318418518618718818919019119219319419519619719819920020120220320420520620720820921021121221321421521621721821922022122222322422522622722822923023123223323423523623723823924024124224324424524624724824925025125225325425525625725825926026126226326426526626726826927027127227327427527627727827928028128228328428528628728828929029129229329429529629729829930030130230330430530630730830931031131231331431531631731831932032132232332432532632732832933033133233333433533633733833934034134234334434534634734834935035135235335435535635735835936036136236336436536636736836937037137237337437537637737837938038138238338438538638738838939039139239339439539639739839940040140240340440540640740840941041141241341441541641741841942042142242342442542642742842943043143243343443543643743843944044144244344444544644744844945045145245345445545645745845946046146246346446546646746846947047147247347447547647747847948048148248348448548648748848949049149249349449549649749849950050150250350450550650750850951051151251351451551651751851952052152252352452552652752852953053153253353453553653753853954054154254354454554654754854955055155255355455555655755855956056156256356456556656756856957057157257357457557657757857958058158258358458558658758858959059159259359459559659759859960060160260360460560660760860961061161261361461561661761861962062162262362462562662762862963063163263363463563663763863964064164264364464564664764864965065165265365465565665765865966066166266366466566666766866967067167267367467567667767867968068168268368468568668768868969069169269369469569669769869970070170270370470570670770870971071171271371471571671771871972072172272372472572672772872973073173273373473573673773873974074174274374474574674774874975075175275375475575675775875976076176276376476576676776876977077177277377477577677777877978078178278378478578678778878979079179279379479579679779879980080180280380480580680780880981081181281381481581681781881982082182282382482582682782882983083183283383483583683783883984084184284384484584684784884985085185285385485585685785885986086186286386486586686786886987087187287387487587687787887988088188288388488588688788888989089189289389489589689789889990090190290390490590690790890991091191291391491591691791891992092192292392492592692792892993093193293393493593693793893994094194294394494594694794894995095195295395495595695795895996096196296396496596696796896997097197297397497597697797897998098198298398498598698798898999099199299399499599699799899910001001100210031004100510061007100810091010101110121013101410151016101710181019102010211022102310241025102610271028102910301031103210331034103510361037103810391040104110421043104410451046104710481049105010511052105310541055105610571058105910601061106210631064106510661067106810691070107110721073107410751076107710781079108010811082108310841085108610871088108910901091109210931094109510961097109810991100110111021103110411051106110711081109111011111112111311141115
  1. import json
  2. import re
  3. import uuid
  4. from copy import deepcopy
  5. from fastapi import WebSocket, WebSocketDisconnect, APIRouter, Depends
  6. import asyncio
  7. import websockets
  8. from sqlalchemy.orm import Session
  9. from Log import logger
  10. from app.api import get_current_user_websocket
  11. from app.config.config import settings
  12. from app.config.const import IMAGE_TO_TEXT, DOCUMENT_TO_REPORT, DOCUMENT_TO_CLEANING, DOCUMENT_IA_QUESTIONS, \
  13. DOCUMENT_TO_REPORT_TITLE, DOCUMENT_TO_TITLE, DOCUMENT_TO_PAPER
  14. from app.models import MenuCapacityModel
  15. from app.models.agent_model import AgentModel, AgentType
  16. from app.models.base_model import get_db
  17. from app.models.user_model import UserModel
  18. from app.service.v2.api_token import DfTokenDao
  19. from app.service.dialog import update_session_history
  20. from app.service.basic import BasicService
  21. from app.service.difyService import DifyService
  22. from app.service.ragflow import RagflowService
  23. from app.service.service_token import get_bisheng_token, get_ragflow_token
  24. from app.service.session import SessionService
  25. router = APIRouter()
  26. # 中间层WebSocket 服务器,接收客户端的连接
  27. @router.websocket("/ws/{agent_id}/{chat_id}")
  28. async def handle_client(websocket: WebSocket,
  29. agent_id: str,
  30. chat_id: str,
  31. current_user: UserModel = Depends(get_current_user_websocket),
  32. db: Session = Depends(get_db)):
  33. tasks = []
  34. await websocket.accept()
  35. print(f"Client {agent_id} connected")
  36. agent = db.query(MenuCapacityModel).filter(MenuCapacityModel.chat_id == agent_id).first()
  37. if not agent:
  38. agent = db.query(AgentModel).filter(AgentModel.id == agent_id).first()
  39. agent_type = agent.agent_type
  40. chat_type = agent.type
  41. else:
  42. agent_type = agent.capacity_type
  43. chat_type = agent.chat_type
  44. # print(agent_type)
  45. # print(chat_type)
  46. if not agent:
  47. ret = {"message": "Agent not found", "type": "close"}
  48. await websocket.send_json(ret)
  49. return
  50. if chat_id == "" or chat_id == "0":
  51. ret = {"message": "Chat ID not found", "type": "close"}
  52. await websocket.send_json(ret)
  53. return
  54. # print(agent_type)
  55. # print(chat_type)
  56. if agent_type == AgentType.RAGFLOW:
  57. ragflow_service = RagflowService(settings.fwr_base_url)
  58. token = await get_ragflow_token(db, current_user.id)
  59. try:
  60. async def forward_to_ragflow():
  61. while True:
  62. message = await websocket.receive_json()
  63. print(f"Received from client {chat_id}: {message}")
  64. chat_history = message.get('chatHistory', [])
  65. message["role"] = "user"
  66. if len(chat_history) == 0:
  67. chat_history = await ragflow_service.get_session_history(token, chat_id)
  68. if len(chat_history) == 0:
  69. chat_history = await ragflow_service.set_session(token, agent_id,
  70. message, chat_id, True)
  71. # print("chat_history------------------------", chat_history)
  72. if len(chat_history) == 0:
  73. result = {"message": "内部错误:创建会话失败", "type": "close"}
  74. await websocket.send_json(result)
  75. await websocket.close()
  76. return
  77. else:
  78. chat_history.append({
  79. "content": message["message"],
  80. "doc_ids": message.get("doc_ids", []),
  81. "role": "user"
  82. })
  83. complete_response = ""
  84. async for rag_response in ragflow_service.chat(token, chat_id, chat_history):
  85. try:
  86. if rag_response[:5] == "data:":
  87. # 如果是,则截取掉前5个字符,并去除首尾空白符
  88. text = rag_response[5:].strip()
  89. else:
  90. # 否则,保持原样
  91. text = rag_response
  92. complete_response += text
  93. try:
  94. json_data = json.loads(complete_response)
  95. data = json_data.get("data")
  96. if data is True: # 完成输出
  97. result = {"message": "", "type": "close"}
  98. elif data is None: # 发生错误
  99. answer = json_data.get("retmsg", json_data.get("retcode"))
  100. result = {"message": "内部错误:" + answer, "type": "message"}
  101. else: # 正常输出
  102. answer = data.get("answer", "")
  103. reference = data.get("reference", {})
  104. result = {"message": answer, "type": "message", "reference": reference}
  105. await websocket.send_json(result)
  106. complete_response = ""
  107. except json.JSONDecodeError as e:
  108. print(f"Error decoding JSON: {e}")
  109. # print(f"Response text: {text}")
  110. except Exception as e2:
  111. result = {"message": f"内部错误: {e2}", "type": "close"}
  112. await websocket.send_json(result)
  113. print(f"Error process message of ragflow: {e2}")
  114. try:
  115. dialog_chat_history = await ragflow_service.get_session_history(token, chat_id, 1)
  116. await update_session_history(db, dialog_chat_history, current_user.id)
  117. except Exception as e:
  118. logger.error(e)
  119. logger.error("-----------------保存ragflow的历史会话异常-----------------")
  120. # 启动任务处理客户端消息
  121. tasks = [
  122. asyncio.create_task(forward_to_ragflow())
  123. ]
  124. await asyncio.wait(tasks, return_when=asyncio.FIRST_COMPLETED)
  125. except WebSocketDisconnect as e1:
  126. print(f"Client {chat_id} disconnected: {e1}")
  127. await websocket.close()
  128. except Exception as e:
  129. print(f"Exception occurred: {e}")
  130. finally:
  131. print("Cleaning up resources of ragflow")
  132. # 取消所有任务
  133. for task in tasks:
  134. if not task.done():
  135. task.cancel()
  136. try:
  137. await task
  138. except asyncio.CancelledError:
  139. pass
  140. elif agent_type == AgentType.BISHENG:
  141. token = await get_bisheng_token(db, current_user.id)
  142. service_uri = f"{settings.sgb_websocket_url}/api/v1/assistant/chat/{agent_id}?t=&chat_id={chat_id}"
  143. headers = {'cookie': f"access_token_cookie={token};"}
  144. async with websockets.connect(service_uri, extra_headers=headers) as service_websocket:
  145. try:
  146. # 处理客户端发来的消息
  147. async def forward_to_service():
  148. while True:
  149. message = await websocket.receive_json()
  150. print(f"Received from client, {chat_id}: {message}")
  151. # 添加 'agent_id' 和 'chat_id' 字段
  152. message['flow_id'] = agent_id
  153. message['chat_id'] = chat_id
  154. msg = message["message"]
  155. del message["message"]
  156. message['inputs'] = {
  157. "data": {"chatId": chat_id, "id": agent_id, "type": "assistant"},
  158. "input": msg
  159. }
  160. await service_websocket.send(json.dumps(message))
  161. print(f"Forwarded to bisheng: {message}")
  162. # 监听毕昇发来的消息并转发给客户端
  163. async def forward_to_client():
  164. while True:
  165. message = await service_websocket.recv()
  166. print(f"Received from bisheng: {message}")
  167. data = json.loads(message)
  168. if data["type"] == "close" or data["type"] == "stream" or data["type"] == "end_cover":
  169. if data["type"] == "close":
  170. t = "close"
  171. else:
  172. t = "stream"
  173. result = {"message": data["message"], "type": t}
  174. await websocket.send_json(result)
  175. print(f"Forwarded to client, {chat_id}: {result}")
  176. # 启动两个任务,分别处理客户端和服务端的消息
  177. tasks = [
  178. asyncio.create_task(forward_to_service()),
  179. asyncio.create_task(forward_to_client())
  180. ]
  181. done, pending = await asyncio.wait(tasks, return_when=asyncio.FIRST_COMPLETED)
  182. # 取消未完成的任务
  183. for task in pending:
  184. task.cancel()
  185. try:
  186. await task
  187. except asyncio.CancelledError:
  188. pass
  189. except WebSocketDisconnect as e:
  190. print(f"WebSocket connection closed with code {e.code}: {e.reason}")
  191. await websocket.close()
  192. await service_websocket.close()
  193. except Exception as e:
  194. print(f"Exception occurred: {e}")
  195. finally:
  196. print("Cleaning up resources of bisheng")
  197. # 取消所有任务
  198. for task in tasks:
  199. if not task.done():
  200. task.cancel()
  201. try:
  202. await task
  203. except asyncio.CancelledError:
  204. pass
  205. elif agent_type == AgentType.BASIC:
  206. try:
  207. service = BasicService(base_url=settings.basic_base_url)
  208. while True:
  209. # 接收前端消息
  210. message = await websocket.receive_json()
  211. question = message.get("message")
  212. try:
  213. SessionService(db).create_session(
  214. chat_id,
  215. question,
  216. agent_id,
  217. AgentType.BASIC,
  218. current_user.id
  219. )
  220. except Exception as e:
  221. logger.error(e)
  222. if not question:
  223. await websocket.send_json({"message": "Invalid request", "type": "error"})
  224. continue
  225. # logger.error(agent.type)
  226. if chat_type == "questionTalk":
  227. try:
  228. data = await service.questions_talk(question, chat_id)
  229. output = data.get("output", "")
  230. file_name = data.get("filename", "")
  231. excel_url = None
  232. if file_name:
  233. excel_url = f"/api/files/download/?agent_id=basic_question_talk&file_id={file_name}&file_type=word"
  234. result = {"message": output, "type": "message", "file_url": excel_url, "file_name": file_name}
  235. try:
  236. SessionService(db).update_session(chat_id,
  237. message={"role": "assistant", "content": result})
  238. except Exception as e:
  239. logger.error(e)
  240. logger.error("-----------------返回数据--------------------")
  241. await websocket.send_json(result)
  242. except Exception as e2:
  243. result = {"message": f"内部错误: {e2}", "type": "close"}
  244. logger.error(str(e2))
  245. logger.error(f"Error process message of basic chuti agent: {e2}")
  246. await websocket.send_json(result)
  247. else:
  248. message_data = {}
  249. logger.error("---------------------excel_talk-----------------------------")
  250. excel_url = ""
  251. image_url = ""
  252. image_name = ""
  253. excel_name = ""
  254. async for data in service.excel_talk(question, chat_id):
  255. # logger.error(data)
  256. output = data.get("output", "")
  257. e_name = data.get("excel_name", "")
  258. i_name = data.get("image_name", "")
  259. def build_file_url(name, file_type):
  260. if not name:
  261. return None
  262. return (f"/api/files/download/?agent_id={agent_id}&file_id={name}"
  263. f"&file_type={file_type}")
  264. if e_name:
  265. excel_url = build_file_url(e_name, 'excel')
  266. excel_name = e_name
  267. if i_name:
  268. image_url = build_file_url(i_name, 'image')
  269. image_name = i_name
  270. if data["type"] == "message":
  271. message_data = {
  272. "content": output,
  273. "excel_url": excel_url,
  274. "image_url": image_url,
  275. "image_name": image_name,
  276. "excel_name": excel_name,
  277. "sql": data.get("sql", ""),
  278. "code": data.get("code", ""),
  279. "e": data.get("e", ""),
  280. "role": "assistant"}
  281. # 发送结果给客户端
  282. # data["type"] = "message"
  283. data["message"] = output
  284. data["excel_url"] = excel_url
  285. data["image_url"] = image_url
  286. await websocket.send_json(data)
  287. if message_data:
  288. try:
  289. SessionService(db).update_session(chat_id, message=message_data)
  290. except Exception as e:
  291. logger.error(f"Unexpected error when update_session: {e}")
  292. except Exception as e:
  293. logger.error(e)
  294. await websocket.send_json({"message": "出现错误!", "type": "error"})
  295. finally:
  296. await websocket.close()
  297. print(f"Client {agent_id} disconnected")
  298. if agent_type == AgentType.DIFY:
  299. dify_service = DifyService(settings.dify_base_url)
  300. # token = get_dify_token(db, current_user.id)
  301. try:
  302. async def forward_to_dify():
  303. if chat_type == "imageTalk":
  304. token = DfTokenDao(db).get_token_by_id(IMAGE_TO_TEXT)
  305. if not token:
  306. await websocket.send_json({"message": "Invalid token", "type": "error"})
  307. while True:
  308. image_list = []
  309. is_image = False
  310. conversation_id = ""
  311. receive_message = await websocket.receive_json()
  312. print(f"Received from client {chat_id}: {receive_message}")
  313. upload_file_id = receive_message.get('upload_file_id', "")
  314. question = receive_message.get('message', "")
  315. if not question and not image_url:
  316. await websocket.send_json({"message": "Invalid request", "type": "error"})
  317. continue
  318. try:
  319. session = SessionService(db).create_session(
  320. chat_id,
  321. question,
  322. agent_id,
  323. AgentType.DIFY,
  324. current_user.id
  325. )
  326. conversation_id = session.conversation_id
  327. except Exception as e:
  328. logger.error(e)
  329. # complete_response = ""
  330. answer_str = ""
  331. files = []
  332. if upload_file_id:
  333. files.append({
  334. "type": "image",
  335. "transfer_method": "local_file",
  336. "url": "",
  337. "upload_file_id": upload_file_id
  338. })
  339. async for rag_response in dify_service.chat(token, current_user.id, question, files,
  340. conversation_id, {}):
  341. # print(rag_response)
  342. try:
  343. if rag_response[:5] == "data:":
  344. # 如果是,则截取掉前5个字符,并去除首尾空白符
  345. complete_response = rag_response[5:].strip()
  346. else:
  347. # 否则,保持原样
  348. complete_response = rag_response
  349. try:
  350. data = json.loads(complete_response)
  351. if data.get("event") == "agent_message": # "event": "message_end"
  352. if "answer" not in data or not data["answer"]: # 信息过滤
  353. logger.error("非法数据--------------------")
  354. # logger.error(data)
  355. continue
  356. else: # 正常输出
  357. answer = data.get("answer", "")
  358. if isinstance(answer, str):
  359. if "![](https://res.stepfun.com/" in answer and image_list:
  360. is_image = True
  361. pattern = r'!\[\] *\(https://res\.stepfun\.com/image_gen/[^)]+\)'
  362. url_image = image_list.pop()
  363. new_answer = re.sub(pattern, url_image, answer)
  364. answer_str += new_answer
  365. else:
  366. answer_str += answer
  367. elif isinstance(answer, dict):
  368. logger.error("未知数据体:0---------------------------------")
  369. logger.error(answer)
  370. answer_str += answer.get("action_input", "")
  371. result = {"message": answer_str, "type": "message"}
  372. elif data.get("event") == "message_end":
  373. images_url = []
  374. if image_list and not is_image:
  375. answer_str += image_list[-1]
  376. result = {"message": answer_str,
  377. "type": "close"} # , "message_files": images_url
  378. try:
  379. SessionService(db).update_session(chat_id,
  380. message={"role": "assistant",
  381. "content": {"answer": answer_str,
  382. "images": images_url}},
  383. conversation_id=data.get(
  384. "conversation_id"))
  385. except Exception as e:
  386. logger.error("保存dify的会话异常!")
  387. logger.error(e)
  388. elif data.get("event") == "message_file":
  389. await dify_service.save_images(data.get("url"), data.get("id") + ".png")
  390. image_list.append(f"![](/api/files/image/{data.get('id')})")
  391. # result = {"message": answer_str, "type": "message"}
  392. continue
  393. else:
  394. continue
  395. await websocket.send_json(result)
  396. complete_response = ""
  397. except json.JSONDecodeError as e:
  398. print(f"Error decoding JSON: {e}")
  399. # print(f"Response text: {text}")
  400. except Exception as e2:
  401. result = {"message": f"内部错误: {e2}", "type": "close"}
  402. await websocket.send_json(result)
  403. print(f"Error process message of ragflow: {e2}")
  404. elif chat_type == "reportWorkflow":
  405. token = DfTokenDao(db).get_token_by_id(DOCUMENT_TO_CLEANING)
  406. if not token:
  407. await websocket.send_json({"message": "Invalid token document_to_cleaning", "type": "error"})
  408. while True:
  409. receive_message = await websocket.receive_json()
  410. print(f"Received from client {chat_id}: {receive_message}")
  411. upload_files = receive_message.get('upload_files', [])
  412. title = receive_message.get('title', "")
  413. workflow_type = receive_message.get('workflow', 1)
  414. sub_titles = receive_message.get('sub_titles', "")
  415. title_number = receive_message.get('title_number', 8)
  416. title_style = receive_message.get('title_style', "")
  417. title_query = receive_message.get('title_query', "")
  418. is_clean = receive_message.get('is_clean', 0)
  419. file_type = receive_message.get('file_type', 1)
  420. max_token = receive_message.get('max_tokens', 100000)
  421. tokens = receive_message.get('tokens', 0)
  422. if upload_files:
  423. title_query = "start"
  424. try:
  425. session = SessionService(db).create_session(
  426. chat_id,
  427. title,
  428. agent_id,
  429. AgentType.DIFY,
  430. current_user.id
  431. )
  432. conversation_id = session.conversation_id
  433. except Exception as e:
  434. logger.error(e)
  435. inputs = {
  436. }
  437. files = []
  438. for file in upload_files:
  439. if file_type == 1:
  440. files.append({
  441. "type": "document",
  442. "transfer_method": "local_file",
  443. "url": "",
  444. "upload_file_id": file
  445. })
  446. else:
  447. files.append({
  448. "type": "document",
  449. "transfer_method": "remote_url",
  450. "url": file,
  451. "upload_file_id": ""
  452. })
  453. inputs_list = []
  454. is_next = 0
  455. if workflow_type == 1:
  456. token = DfTokenDao(db).get_token_by_id(DOCUMENT_TO_CLEANING)
  457. if not token:
  458. await websocket.send_json(
  459. {"message": "Invalid token document_to_cleaning", "type": "error"})
  460. inputs["input_files"] = files
  461. inputs["Completion_of_main_indicators"] = title
  462. inputs_list.append({"inputs": inputs, "token": token, "workflow_type": workflow_type})
  463. if workflow_type == 2:
  464. inputs["file_list"] = files
  465. inputs["Completion_of_main_indicators"] = title
  466. inputs["sub_titles"] = sub_titles
  467. token = DfTokenDao(db).get_token_by_id(DOCUMENT_TO_REPORT_TITLE)
  468. if not token:
  469. await websocket.send_json(
  470. {"message": "Invalid token document_to_cleaning", "type": "error"})
  471. inputs_list.append({"inputs": inputs, "token": token, "workflow_type": workflow_type})
  472. elif workflow_type == 3 and is_clean == 0 and tokens < max_token:
  473. inputs["file_list"] = files
  474. inputs["number_of_title"] = title_number
  475. inputs["title_style"] = title_style
  476. token = DfTokenDao(db).get_token_by_id(DOCUMENT_TO_TITLE)
  477. if not token:
  478. await websocket.send_json(
  479. {"message": "Invalid token document_to_title", "type": "error"})
  480. inputs_list.append({"inputs": inputs, "token": token, "workflow_type": workflow_type})
  481. elif workflow_type == 3 and is_clean == 1 or tokens >= max_token:
  482. token = DfTokenDao(db).get_token_by_id(DOCUMENT_TO_CLEANING)
  483. if not token:
  484. await websocket.send_json(
  485. {"message": "Invalid token document_to_cleaning", "type": "error"})
  486. inputs["input_files"] = files
  487. inputs["Completion_of_main_indicators"] = title
  488. inputs_list.append({"inputs": inputs, "token": token, "workflow_type": 1})
  489. inputs1 = {}
  490. inputs1["file_list"] = files
  491. inputs1["number_of_title"] = title_number
  492. inputs1["title_style"] = title_style
  493. token = DfTokenDao(db).get_token_by_id(DOCUMENT_TO_TITLE)
  494. if not token:
  495. await websocket.send_json(
  496. {"message": "Invalid token document_to_report", "type": "error"})
  497. inputs_list.append({"inputs": inputs1, "token": token, "workflow_type": 3})
  498. complete_response = ""
  499. for idx, input in enumerate(inputs_list):
  500. # print(input)
  501. if idx < len(inputs_list) - 1:
  502. is_next = 1
  503. else:
  504. is_next = 0
  505. i = input["inputs"]
  506. if "file_list" in i:
  507. i["file_list"] = files
  508. # print(i)
  509. node_list = []
  510. complete_response = ""
  511. workflow_list = []
  512. workflow_dict = {}
  513. if input["workflow_type"] == 1 or input["workflow_type"] == 2:
  514. async for rag_response in dify_service.workflow(input["token"], current_user.id, i):
  515. # print(rag_response)
  516. try:
  517. if rag_response[:5] == "data:":
  518. # 如果是,则截取掉前5个字符,并去除首尾空白符
  519. complete_response = rag_response[5:].strip()
  520. elif "event: ping" in rag_response:
  521. continue
  522. else:
  523. # 否则,保持原样
  524. complete_response += rag_response
  525. try:
  526. data = json.loads(complete_response)
  527. # print(data)
  528. node_data = deepcopy(data)
  529. if "data" in node_data:
  530. if "outputs" in node_data["data"]:
  531. node_data["data"]["outputs"] = {}
  532. if "inputs" in node_data["data"]:
  533. node_data["data"]["inputs"] = {}
  534. # print(node_data)
  535. node_list.append(node_data)
  536. complete_response = ""
  537. if data.get("event") == "node_started": # "event": "message_end"
  538. if "data" not in data or not data["data"]: # 信息过滤
  539. logger.error("非法数据--------------------")
  540. logger.error(data)
  541. continue
  542. else: # 正常输出
  543. answer = data.get("data", "")
  544. if isinstance(answer, str):
  545. logger.error("----------------未知数据--------------------")
  546. logger.error(data)
  547. continue
  548. elif isinstance(answer, dict):
  549. message = answer.get("title", "")
  550. result = {"message": message, "type": "system",
  551. "workflow": {"node_data": workflow_list}}
  552. elif data.get("event") == "node_finished":
  553. workflow_list.append({
  554. "title": data.get("data", {}).get("title", ""),
  555. "status": data.get("data", {}).get("status", ""),
  556. "created_at": data.get("data", {}).get("created_at", 0),
  557. "finished_at": data.get("data", {}).get("finished_at", 0),
  558. "node_type": data.get("data", {}).get("node_type", 0),
  559. "elapsed_time": data.get("data", {}).get("elapsed_time", 0),
  560. "error": data.get("data", {}).get("error", ""),
  561. })
  562. answer = data.get("data", "")
  563. if isinstance(answer, str):
  564. logger.error("----------------未知数据--------------------")
  565. logger.error(data)
  566. continue
  567. elif isinstance(answer, dict):
  568. message = answer.get("title", "")
  569. if answer.get("status") == "failed":
  570. message = answer.get("error", "")
  571. result = {"message": message, "type": "system",
  572. "workflow": {"node_data": workflow_list}}
  573. elif data.get("event") == "workflow_finished":
  574. answer = data.get("data", "")
  575. if isinstance(answer, str):
  576. logger.error("----------------未知数据--------------------")
  577. logger.error(data)
  578. result = {"message": "", "type": "close", "download_url": "",
  579. "is_next": is_next}
  580. elif isinstance(answer, dict):
  581. download_url = ""
  582. outputs = answer.get("outputs", {})
  583. if outputs:
  584. message = outputs.get("output", "")
  585. download_url = outputs.get("download_url", "")
  586. else:
  587. message = answer.get("error", "")
  588. if download_url:
  589. files = [{
  590. "type": "document",
  591. "transfer_method": "remote_url",
  592. "url": download_url,
  593. "upload_file_id": ""
  594. }]
  595. workflow_dict = {
  596. "node_data": workflow_list,
  597. "total_tokens": answer.get("total_tokens", 0),
  598. "created_at": answer.get("created_at", 0),
  599. "finished_at": answer.get("finished_at", 0),
  600. "status": answer.get("status", ""),
  601. "error": answer.get("error", ""),
  602. "elapsed_time": answer.get("elapsed_time", 0)
  603. }
  604. result = {"message": message, "type": "message",
  605. "download_url": download_url, "workflow": workflow_dict}
  606. try:
  607. SessionService(db).update_session(chat_id,
  608. message={"role": "assistant",
  609. "content": {
  610. "answer": message,
  611. "node_list": node_list,
  612. "download_url": download_url}},
  613. conversation_id=data.get(
  614. "conversation_id"))
  615. node_list = []
  616. except Exception as e:
  617. logger.error("保存dify的会话异常!")
  618. logger.error(e)
  619. try:
  620. await websocket.send_json(result)
  621. except Exception as e:
  622. logger.error(e)
  623. logger.error("返回客户端消息异常!")
  624. result = {"message": "", "type": "close", "workflow": workflow_dict,
  625. "is_next": is_next, "download_url": download_url}
  626. else:
  627. continue
  628. try:
  629. await websocket.send_json(result)
  630. except Exception as e:
  631. logger.error(e)
  632. logger.error("返回客户端消息异常!")
  633. complete_response = ""
  634. except json.JSONDecodeError as e:
  635. print(f"Error decoding JSON: {e}")
  636. # print(f"Response text: {text}")
  637. except Exception as e2:
  638. result = {"message": f"内部错误: {e2}", "type": "close"}
  639. await websocket.send_json(result)
  640. print(f"Error process message of ragflow: {e2}")
  641. elif input["workflow_type"] == 3:
  642. image_list = []
  643. # print(inputs)
  644. complete_response = ""
  645. answer_str = ""
  646. async for rag_response in dify_service.chat(input["token"], current_user.id,
  647. title_query, [],
  648. conversation_id, i):
  649. # print(rag_response)
  650. try:
  651. if rag_response[:5] == "data:":
  652. # 如果是,则截取掉前5个字符,并去除首尾空白符
  653. complete_response = rag_response[5:].strip()
  654. elif "event: ping" in rag_response:
  655. continue
  656. else:
  657. # 否则,保持原样
  658. complete_response += rag_response
  659. try:
  660. data = json.loads(complete_response)
  661. node_data = deepcopy(data)
  662. if "data" in node_data:
  663. if "outputs" in node_data["data"]:
  664. node_data["data"]["outputs"] = {}
  665. if "inputs" in node_data["data"]:
  666. node_data["data"]["inputs"] = {}
  667. # print(node_data)
  668. node_list.append(node_data)
  669. complete_response = ""
  670. if data.get("event") == "node_started": # "event": "message_end"
  671. if "data" not in data or not data["data"]: # 信息过滤
  672. logger.error("非法数据--------------------")
  673. logger.error(data)
  674. continue
  675. else: # 正常输出
  676. answer = data.get("data", "")
  677. if isinstance(answer, str):
  678. logger.error("----------------未知数据--------------------")
  679. logger.error(data)
  680. continue
  681. elif isinstance(answer, dict):
  682. message = answer.get("title", "")
  683. result = {"message": message, "type": "system",
  684. "workflow": {"node_data": workflow_list}}
  685. elif data.get("event") == "node_finished":
  686. workflow_list.append({
  687. "title": data.get("data", {}).get("title", ""),
  688. "status": data.get("data", {}).get("status", ""),
  689. "created_at": data.get("data", {}).get("created_at", 0),
  690. "finished_at": data.get("data", {}).get("finished_at", 0),
  691. "node_type": data.get("data", {}).get("node_type", 0),
  692. "elapsed_time": data.get("data", {}).get("elapsed_time", 0),
  693. "error": data.get("data", {}).get("error", ""),
  694. })
  695. answer = data.get("data", "")
  696. if isinstance(answer, str):
  697. logger.error("----------------未知数据--------------------")
  698. logger.error(data)
  699. continue
  700. elif isinstance(answer, dict):
  701. message = answer.get("title", "")
  702. if answer.get("status") == "failed":
  703. message = answer.get("error", "")
  704. result = {"message": message, "type": "system",
  705. "workflow": {"node_data": workflow_list}}
  706. elif data.get("event") == "message":
  707. answer_str = data.get("answer", "")
  708. # try:
  709. # msg_dict = json.loads(answer)
  710. # message = msg_dict.get("output", "")
  711. # except Exception as e:
  712. # print(e)
  713. # continue
  714. result = {"message": answer_str, "type": "message",
  715. "download_url": "", "workflow": {"node_data": workflow_list}}
  716. # try:
  717. # await websocket.send_json(result)
  718. # except Exception as e:
  719. # logger.error(e)
  720. # logger.error("返回客户端消息异常!")
  721. elif data.get("event") == "workflow_finished":
  722. workflow_dict = {
  723. "node_data": workflow_list,
  724. "total_tokens": data.get("data", {}).get("total_tokens", 0),
  725. "created_at": data.get("data", {}).get("created_at", 0),
  726. "finished_at": data.get("data", {}).get("finished_at", 0),
  727. "status": data.get("data", {}).get("status", ""),
  728. "error": data.get("data", {}).get("error", ""),
  729. "elapsed_time": data.get("data", {}).get("elapsed_time", 0)
  730. }
  731. try:
  732. SessionService(db).update_session(chat_id,
  733. message={"role": "assistant",
  734. "content": {
  735. "answer": answer_str,
  736. "node_list": node_list,
  737. "download_url": ""}},
  738. conversation_id=data.get(
  739. "conversation_id"))
  740. node_list = []
  741. except Exception as e:
  742. logger.error("保存dify的会话异常!")
  743. logger.error(e)
  744. elif data.get("event") == "message_end":
  745. result = {"message": "", "type": "close", "workflow": workflow_dict,
  746. "is_next": is_next}
  747. else:
  748. continue
  749. try:
  750. await websocket.send_json(result)
  751. except Exception as e:
  752. logger.error(e)
  753. logger.error("dify返回客户端消息异常!")
  754. complete_response = ""
  755. except json.JSONDecodeError as e:
  756. print(f"Error decoding JSON: {e}")
  757. # print(f"Response text: {text}")
  758. except Exception as e2:
  759. result = {"message": f"内部错误: {e2}", "type": "close"}
  760. await websocket.send_json(result)
  761. print(f"Error process message of ragflow: {e2}")
  762. elif chat_type == "documentIa":
  763. # print(122112)
  764. token = DfTokenDao(db).get_token_by_id(DOCUMENT_IA_QUESTIONS)
  765. # print(token)
  766. if not token:
  767. await websocket.send_json({"message": "Invalid token", "type": "error"})
  768. while True:
  769. conversation_id = ""
  770. # print(4343)
  771. receive_message = await websocket.receive_json()
  772. print(f"Received from client {chat_id}: {receive_message}")
  773. upload_file_id = receive_message.get('upload_file_id', [])
  774. question = receive_message.get('message', "")
  775. if not question and not image_url:
  776. await websocket.send_json({"message": "Invalid request", "type": "error"})
  777. continue
  778. try:
  779. session = SessionService(db).create_session(
  780. chat_id,
  781. question,
  782. agent_id,
  783. AgentType.DIFY,
  784. current_user.id
  785. )
  786. conversation_id = session.conversation_id
  787. except Exception as e:
  788. logger.error(e)
  789. # complete_response = ""
  790. files = []
  791. for fileId in upload_file_id:
  792. files.append({
  793. "type": "document",
  794. "transfer_method": "local_file",
  795. "url": "",
  796. "upload_file_id": fileId
  797. })
  798. answer_str = ""
  799. complete_response = ""
  800. async for rag_response in dify_service.chat(token, current_user.id, question, files,
  801. conversation_id, {}):
  802. print(rag_response)
  803. try:
  804. if rag_response[:5] == "data:":
  805. # 如果是,则截取掉前5个字符,并去除首尾空白符
  806. complete_response = rag_response[5:].strip()
  807. elif "event: ping" in rag_response:
  808. continue
  809. else:
  810. # 否则,保持原样
  811. complete_response += rag_response
  812. try:
  813. data = json.loads(complete_response)
  814. if data.get("event") == "node_started" or data.get(
  815. "event") == "node_finished": # "event": "message_end"
  816. if "data" not in data or not data["data"]: # 信息过滤
  817. logger.error("非法数据--------------------")
  818. logger.error(data)
  819. continue
  820. else: # 正常输出
  821. answer = data.get("data", "")
  822. if isinstance(answer, str):
  823. logger.error("----------------未知数据--------------------")
  824. logger.error(data)
  825. continue
  826. elif isinstance(answer, dict):
  827. message = answer.get("title", "")
  828. if answer.get("status") == "failed":
  829. message = answer.get("error")
  830. result = {"message": message, "type": "system"}
  831. # continue
  832. elif data.get("event") == "message": # "event": "message_end"
  833. # 正常输出
  834. answer = data.get("answer", "")
  835. result = {"message": answer, "type": "stream"}
  836. elif data.get("event") == "error":
  837. answer = data.get("message", "")
  838. result = {"message": answer, "type": "system"}
  839. elif data.get("event") == "workflow_finished":
  840. answer = data.get("data", "")
  841. if isinstance(answer, str):
  842. logger.error("----------------未知数据--------------------")
  843. logger.error(data)
  844. # result = {"message": "", "type": "close", "download_url": ""}
  845. elif isinstance(answer, dict):
  846. download_url = ""
  847. outputs = answer.get("outputs", {})
  848. if outputs:
  849. message = outputs.get("answer", "")
  850. # download_url = outputs.get("download_url", "")
  851. else:
  852. message = answer.get("error", "")
  853. result = {"message": message, "type": "system",
  854. "download_url": download_url}
  855. try:
  856. SessionService(db).update_session(chat_id,
  857. message={"role": "assistant",
  858. "content": {
  859. "answer": message,
  860. "download_url": download_url}},
  861. conversation_id=data.get(
  862. "conversation_id"))
  863. except Exception as e:
  864. logger.error("保存dify的会话异常!")
  865. logger.error(e)
  866. # await websocket.send_json(result)
  867. # continue
  868. elif data.get("event") == "message_end":
  869. result = {"message": "", "type": "close"}
  870. else:
  871. continue
  872. try:
  873. await websocket.send_json(result)
  874. except Exception as e:
  875. logger.error(e)
  876. logger.error("返回客户端消息异常!")
  877. complete_response = ""
  878. except json.JSONDecodeError as e:
  879. print(f"Error decoding JSON: {e}")
  880. # print(f"Response text: {text}")
  881. except Exception as e2:
  882. result = {"message": f"内部错误: {e2}", "type": "close"}
  883. await websocket.send_json(result)
  884. print(f"Error process message of ragflow: {e2}")
  885. elif chat_type == "paperTalk":
  886. token = DfTokenDao(db).get_token_by_id(DOCUMENT_TO_PAPER)
  887. # print(token)
  888. if not token:
  889. await websocket.send_json({"message": "Invalid token", "type": "error"})
  890. while True:
  891. conversation_id = ""
  892. inputs = {}
  893. # print(4343)
  894. receive_message = await websocket.receive_json()
  895. print(f"Received from client {chat_id}: {receive_message}")
  896. if "difficulty" in receive_message:
  897. inputs["Question_Difficulty"] = receive_message["difficulty"]
  898. if "is_paper" in receive_message:
  899. inputs["Generate_test_paper"] = receive_message["is_paper"]
  900. if "single_choice" in receive_message:
  901. inputs["Multiple_choice_questions"] = receive_message["single_choice"]
  902. if "gap_filling" in receive_message:
  903. inputs["Fill_in_blank"] = receive_message["gap_filling"]
  904. if "true_or_false" in receive_message:
  905. inputs["true_or_false"] = receive_message["true_or_false"]
  906. if "multiple_choice" in receive_message:
  907. inputs["Multiple_Choice"] = receive_message["multiple_choice"]
  908. if "easy_question" in receive_message:
  909. inputs["Short_Answer_Questions"] = receive_message["easy_question"]
  910. if "case_questions" in receive_message:
  911. inputs["Case_Questions"] = receive_message["case_questions"]
  912. if "key_words" in receive_message:
  913. inputs["key_words"] = receive_message["key_words"]
  914. upload_files = receive_message.get('upload_files', [])
  915. question = receive_message.get('message', "")
  916. session_log = SessionService(db).get_session_by_id(chat_id)
  917. if not session_log and not upload_files:
  918. await websocket.send_json({"message": "需要上传文档!", "type": "error"})
  919. continue
  920. try:
  921. session = SessionService(db).create_session(
  922. chat_id,
  923. question if question else "开始出题",
  924. agent_id,
  925. AgentType.DIFY,
  926. current_user.id
  927. )
  928. conversation_id = session.conversation_id
  929. except Exception as e:
  930. logger.error(e)
  931. # complete_response = ""
  932. files = []
  933. for fileId in upload_files:
  934. files.append({
  935. "type": "document",
  936. "transfer_method": "local_file",
  937. "url": "",
  938. "upload_file_id": fileId
  939. })
  940. if files:
  941. inputs["upload_files"] = files
  942. # print(inputs)
  943. if not question and not inputs:
  944. await websocket.send_json({"message": "Invalid request", "type": "error"})
  945. continue
  946. if not question:
  947. question = "开始出题"
  948. complete_response = ""
  949. async for rag_response in dify_service.chat(token, current_user.id, question, files,
  950. conversation_id, inputs):
  951. # print(rag_response)
  952. try:
  953. if rag_response[:5] == "data:":
  954. # 如果是,则截取掉前5个字符,并去除首尾空白符
  955. complete_response = rag_response[5:].strip()
  956. elif "event: ping" in rag_response:
  957. continue
  958. else:
  959. # 否则,保持原样
  960. complete_response += rag_response
  961. try:
  962. data = json.loads(complete_response)
  963. # print(data)
  964. if data.get("event") == "node_started" or data.get(
  965. "event") == "node_finished": # "event": "message_end"
  966. if "data" not in data or not data["data"]: # 信息过滤
  967. logger.error("非法数据--------------------")
  968. logger.error(data)
  969. continue
  970. else: # 正常输出
  971. answer = data.get("data", "")
  972. if isinstance(answer, str):
  973. logger.error("----------------未知数据--------------------")
  974. logger.error(data)
  975. continue
  976. elif isinstance(answer, dict):
  977. message = answer.get("title", "")
  978. result = {"message": message, "type": "system"}
  979. # continue
  980. elif data.get("event") == "message": # "event": "message_end"
  981. # 正常输出
  982. answer = data.get("answer", "")
  983. result = {"message": answer, "type": "stream"}
  984. elif data.get("event") == "error":
  985. answer = data.get("message", "")
  986. result = {"message": answer, "type": "system"}
  987. elif data.get("event") == "workflow_finished":
  988. answer = data.get("data", "")
  989. if isinstance(answer, str):
  990. logger.error("----------------未知数据--------------------")
  991. logger.error(data)
  992. result = {"message": "", "type": "close", "download_url": ""}
  993. elif isinstance(answer, dict):
  994. download_url = ""
  995. outputs = answer.get("outputs", {})
  996. if outputs:
  997. message = outputs.get("answer", "")
  998. download_url = outputs.get("download_url", "")
  999. else:
  1000. message = answer.get("error", "")
  1001. result = {"message": message, "type": "system",
  1002. "download_url": download_url}
  1003. try:
  1004. SessionService(db).update_session(chat_id,
  1005. message={"role": "assistant",
  1006. "content": {
  1007. "answer": message,
  1008. "download_url": download_url}},
  1009. conversation_id=data.get(
  1010. "conversation_id"))
  1011. except Exception as e:
  1012. logger.error("保存dify的会话异常!")
  1013. logger.error(e)
  1014. # await websocket.send_json(result)
  1015. # continue
  1016. elif data.get("event") == "message_end":
  1017. result = {"message": "", "type": "close"}
  1018. else:
  1019. continue
  1020. try:
  1021. await websocket.send_json(result)
  1022. except Exception as e:
  1023. logger.error(e)
  1024. logger.error("返回客户端消息异常!")
  1025. complete_response = ""
  1026. except json.JSONDecodeError as e:
  1027. print(f"Error decoding JSON: {e}")
  1028. # print(f"Response text: {text}")
  1029. except Exception as e2:
  1030. result = {"message": f"内部错误: {e2}", "type": "close"}
  1031. await websocket.send_json(result)
  1032. print(f"Error process message of ragflow: {e2}")
  1033. # 启动任务处理客户端消息
  1034. tasks = [
  1035. asyncio.create_task(forward_to_dify())
  1036. ]
  1037. await asyncio.wait(tasks, return_when=asyncio.FIRST_COMPLETED)
  1038. except WebSocketDisconnect as e1:
  1039. print(f"Client {chat_id} disconnected: {e1}")
  1040. await websocket.close()
  1041. except Exception as e:
  1042. print(f"Exception occurred: {e}")
  1043. finally:
  1044. print("Cleaning up resources of ragflow")
  1045. # 取消所有任务
  1046. for task in tasks:
  1047. if not task.done():
  1048. task.cancel()
  1049. try:
  1050. await task
  1051. except asyncio.CancelledError:
  1052. pass
  1053. else:
  1054. ret = {"message": "Agent not found", "type": "close"}
  1055. await websocket.send_json(ret)