chat.py 8.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226
  1. import json
  2. from Log import logger
  3. from app.config.agent_base_url import RG_CHAT_DIALOG, DF_CHAT_AGENT, DF_CHAT_PARAMETERS, RG_CHAT_SESSIONS, \
  4. DF_CHAT_WORKFLOW
  5. from app.config.config import settings
  6. from app.config.const import *
  7. from app.models import DialogModel, ApiTokenModel
  8. from app.models.v2.session_model import ChatSessionDao, ChatData
  9. from app.service.v2.app_driver.chat_agent import ChatAgent
  10. from app.service.v2.app_driver.chat_data import ChatBaseApply
  11. from app.service.v2.app_driver.chat_dialog import ChatDialog
  12. from app.service.v2.app_driver.chat_workflow import ChatWorkflow
  13. async def update_session_log(db, session_id: str, message: dict, conversation_id: str):
  14. await ChatSessionDao(db).update_session_by_id(
  15. session_id=session_id,
  16. session=None,
  17. message=message,
  18. conversation_id=conversation_id
  19. )
  20. async def add_session_log(db, session_id: str, question: str, chat_id: str, user_id, event_type: str,
  21. conversation_id: str):
  22. try:
  23. session = await ChatSessionDao(db).update_or_insert_by_id(
  24. session_id=session_id,
  25. name=question[:255],
  26. agent_id=chat_id,
  27. agent_type=1,
  28. tenant_id=user_id,
  29. message={"role": "user", "content": question},
  30. conversation_id=conversation_id,
  31. event_type=event_type
  32. )
  33. return session
  34. except Exception as e:
  35. logger.error(e)
  36. return None
  37. async def get_chat_token(db, app_id):
  38. app_token = db.query(ApiTokenModel).filter_by(app_id=app_id).first()
  39. if app_token:
  40. return app_token.token
  41. return ""
  42. async def get_chat_info(db, chat_id: str):
  43. return db.query(DialogModel).filter_by(id=chat_id, status=Dialog_STATSU_ON).first()
  44. async def get_chat_object(mode):
  45. if mode == workflow_chat:
  46. url = settings.dify_base_url + DF_CHAT_WORKFLOW
  47. return ChatWorkflow(), url
  48. else:
  49. url = settings.dify_base_url + DF_CHAT_AGENT
  50. return ChatAgent(), url
  51. async def service_chat_dialog(db, chat_id: str, question: str, session_id: str, user_id, mode: str):
  52. conversation_id = ""
  53. token = await get_chat_token(db, rg_api_token)
  54. url = settings.fwr_base_url + RG_CHAT_DIALOG.format(chat_id)
  55. chat = ChatDialog()
  56. session = await add_session_log(db, session_id, question, chat_id, user_id, mode, session_id)
  57. if session:
  58. conversation_id = session.conversation_id
  59. message = {"role": "assistant", "answer": "", "reference": {}}
  60. try:
  61. async for ans in chat.chat_completions(url, await chat.request_data(question, conversation_id),
  62. await chat.get_headers(token)):
  63. data = {}
  64. error = ""
  65. status = http_200
  66. if ans.get("code", None) == 102:
  67. error = ans.get("message", "error!")
  68. status = http_400
  69. event = smart_message_error
  70. else:
  71. if isinstance(ans.get("data"), bool) and ans.get("data") is True:
  72. event = smart_message_end
  73. else:
  74. data = ans.get("data", {})
  75. # conversation_id = data.get("session_id", "")
  76. if "session_id" in data:
  77. del data["session_id"]
  78. message = data
  79. event = smart_message_cover
  80. message_str = "data: " + json.dumps(
  81. {"event": event, "data": data, "error": error, "status": status, "session_id": session_id},
  82. ensure_ascii=False) + "\n\n"
  83. for i in range(0, len(message_str), max_chunk_size):
  84. chunk = message_str[i:i + max_chunk_size]
  85. # print(chunk)
  86. yield chunk # 发送分块消息
  87. except Exception as e:
  88. logger.error(e)
  89. try:
  90. yield "data: " + json.dumps({"message": smart_message_error,
  91. "error": "**ERROR**: " + str(e), "status": http_500},
  92. ensure_ascii=False) + "\n\n"
  93. except:
  94. ...
  95. finally:
  96. await update_session_log(db, session_id, message, conversation_id)
  97. async def service_chat_workflow(db, chat_id: str, chat_data: ChatData, session_id: str, user_id, mode: str):
  98. conversation_id = ""
  99. answer_event = ""
  100. answer_agent = ""
  101. message_id = ""
  102. task_id = ""
  103. error = ""
  104. files = []
  105. node_list = []
  106. token = await get_chat_token(db, chat_id)
  107. chat, url = await get_chat_object(mode)
  108. if hasattr(chat_data, "query"):
  109. query = chat_data.query
  110. else:
  111. query = "start new workflow"
  112. session = await add_session_log(db, session_id, query, chat_id, user_id, mode, conversation_id)
  113. if session:
  114. conversation_id = session.conversation_id
  115. try:
  116. async for ans in chat.chat_completions(url,
  117. await chat.request_data(query, conversation_id, str(user_id), chat_data),
  118. await chat.get_headers(token)):
  119. data = {}
  120. status = http_200
  121. conversation_id = ans.get("conversation_id")
  122. task_id = ans.get("task_id")
  123. if ans.get("event") == message_error:
  124. error = ans.get("message", "参数异常!")
  125. status = http_400
  126. event = smart_message_error
  127. elif ans.get("event") == message_agent:
  128. data = {"answer": ans.get("answer", ""), "id": ans.get("message_id", "")}
  129. answer_agent += ans.get("answer", "")
  130. message_id = ans.get("message_id", "")
  131. event = smart_message_stream
  132. elif ans.get("event") == message_event:
  133. data = {"answer": ans.get("answer", ""), "id": ans.get("message_id", "")}
  134. answer_event += ans.get("answer", "")
  135. message_id = ans.get("message_id", "")
  136. event = smart_message_stream
  137. elif ans.get("event") == message_file:
  138. data = {"url": ans.get("url", ""), "id": ans.get("id", ""),
  139. "type": ans.get("type", "")}
  140. files.append(data)
  141. event = smart_message_file
  142. elif ans.get("event") in [workflow_started, node_started, node_finished]:
  143. data = ans.get("data", {})
  144. data["inputs"] = []
  145. data["outputs"] = []
  146. data["process_data"] = ""
  147. node_list.append(ans)
  148. event = [smart_workflow_started, smart_node_started, smart_node_finished][
  149. [workflow_started, node_started, node_finished].index(ans.get("event"))]
  150. elif ans.get("event") == workflow_finished:
  151. data = ans.get("data", {})
  152. event = smart_workflow_finished
  153. node_list.append(ans)
  154. elif ans.get("event") == message_end:
  155. event = smart_message_end
  156. else:
  157. continue
  158. yield "data: " + json.dumps(
  159. {"event": event, "data": data, "error": error, "status": status, "task_id": task_id,
  160. "session_id": session_id},
  161. ensure_ascii=False) + "\n\n"
  162. except Exception as e:
  163. logger.error(e)
  164. try:
  165. yield "data: " + json.dumps({"message": smart_message_error,
  166. "error": "**ERROR**: " + str(e), "status": http_500},
  167. ensure_ascii=False) + "\n\n"
  168. except:
  169. ...
  170. finally:
  171. await update_session_log(db, session_id, {"role": "assistant", "answer": answer_event or answer_agent,
  172. "node_list": node_list, "task_id": task_id, "id": message_id,
  173. "error": error}, conversation_id)
  174. async def service_chat_basic(db, chat_id: str, question: str, session_id: str, user_id):
  175. ...
  176. async def service_chat_parameters(db, chat_id, user_id):
  177. chat_info = db.query(DialogModel).filter_by(id=chat_id).first()
  178. if not chat_info:
  179. return {}
  180. if chat_info.dialog_type == RG_TYPE:
  181. return {"retriever_resource":
  182. {
  183. "enabled": True
  184. }
  185. }
  186. elif chat_info.dialog_type == BASIC_TYPE:
  187. ...
  188. elif chat_info.dialog_type == DF_TYPE:
  189. token = await get_chat_token(db, chat_id)
  190. if not token:
  191. return {}
  192. url = settings.dify_base_url + DF_CHAT_PARAMETERS
  193. chat = ChatBaseApply()
  194. return await chat.chat_parameters(url, {"user": str(user_id)}, await chat.get_headers(token))
  195. async def service_chat_sessions(db, chat_id, name):
  196. token = await get_chat_token(db, rg_api_token)
  197. if not token:
  198. return {}
  199. url = settings.fwr_base_url + RG_CHAT_SESSIONS.format(chat_id)
  200. chat = ChatDialog()
  201. return await chat.chat_sessions(url, {"name": name}, await chat.get_headers(token))