dialog.py 11 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259
  1. import json
  2. from datetime import datetime
  3. from sqlalchemy import or_
  4. from app.config.agent_base_url import DF_CHAT_PARAMETERS, DF_CHAT_API_KEY
  5. from app.config.config import settings
  6. from app.config.const import Dialog_STATSU_DELETE, DF_TYPE, Dialog_STATSU_ON, workflow_server, RG_TYPE, basic_chat
  7. from app.models import KnowledgeModel, GroupModel, DialogModel, ConversationModel, group_dialog_table, LabelWorkerModel, \
  8. LabelModel, ApiTokenModel
  9. from app.models.user_model import UserModel, UserTokenModel
  10. from Log import logger
  11. from app.service.v2.app_driver.chat_data import ChatBaseApply
  12. from app.service.v2.chat import get_chat_token, add_chat_token, get_app_token
  13. from app.task.fetch_agent import get_one_from_ragflow_dialog
  14. async def get_dialog_list(db, user_id, keyword, label, status, page_size, page_index, mode):
  15. user = db.query(UserModel).filter(UserModel.id == user_id).first()
  16. if user is None:
  17. return {"rows": []}
  18. query = db.query(DialogModel)
  19. if status:
  20. query = query.filter(DialogModel.status == status)
  21. else:
  22. query = query.filter(DialogModel.status != Dialog_STATSU_DELETE)
  23. if mode == 1:
  24. query = query.filter(DialogModel.mode != basic_chat)
  25. id_list = []
  26. # if label:
  27. # id_list = [i.object_id for i in db.query(LabelWorkerModel).filter(LabelWorkerModel.label_id==label).all()]
  28. if user.permission != "admin":
  29. dia_list = [j.id for i in user.groups for j in i.dialogs]
  30. query = query.filter(or_(DialogModel.tenant_id == user_id, DialogModel.id.in_(dia_list)))
  31. # else:
  32. if label:
  33. id_list = [i.object_id for i in db.query(LabelWorkerModel).filter(LabelWorkerModel.label_id == label).all()]
  34. query = query.filter(DialogModel.id.in_(id_list))
  35. if keyword:
  36. query = query.filter(DialogModel.name.like('%{}%'.format(keyword)))
  37. query = query.order_by(DialogModel.update_date.desc())
  38. total = query.count()
  39. if page_size:
  40. query = query.limit(page_size).offset((page_index - 1) * page_size)
  41. rows = []
  42. user_id_set = set()
  43. dialog_id_set = set()
  44. label_dict = {}
  45. for kld in query.all():
  46. user_id_set.add(kld.tenant_id)
  47. dialog_id_set.add(kld.id)
  48. rows.append(kld.to_json())
  49. user_dict = {str(i.id): i.to_dict() for i in db.query(UserModel).filter(UserModel.id.in_(user_id_set)).all()}
  50. for i in db.query(LabelModel.id, LabelModel.name, LabelWorkerModel.object_id).outerjoin(LabelWorkerModel,
  51. LabelModel.id == LabelWorkerModel.label_id).filter(
  52. LabelWorkerModel.object_id.in_(dialog_id_set)).all():
  53. label_dict[i.object_id] = label_dict.get(i.object_id, []) + [{"labelId": i.id, "labelName": i.name}]
  54. for r in rows:
  55. r["user"] = user_dict.get(r["user_id"], {})
  56. r["label"] = label_dict.get(r["id"], [])
  57. return {"total": total, "rows": rows}
  58. async def update_session_history(db, data: dict, user_id):
  59. session_id = data.get("id")
  60. if not session_id:
  61. logger.error("更新回话记录失败!{}".format(data))
  62. return
  63. data["create_date"] = datetime.strptime(data["create_date"], '%a, %d %b %Y %H:%M:%S %Z')
  64. data["update_date"] = datetime.strptime(data["update_date"], '%a, %d %b %Y %H:%M:%S %Z')
  65. conversation = db.query(ConversationModel).filter(ConversationModel.id == session_id).first()
  66. if not conversation:
  67. try:
  68. data["tenant_id"] = user_id
  69. conversation_model = ConversationModel(**data)
  70. db.add(conversation_model)
  71. db.commit()
  72. except Exception as e:
  73. logger.error(e)
  74. db.rollback()
  75. else:
  76. try:
  77. # data["tenant_id"] = user_id
  78. del data["id"]
  79. db.query(ConversationModel).filter(ConversationModel.id == session_id).update(data)
  80. db.commit()
  81. except Exception as e:
  82. logger.error(e)
  83. db.rollback()
  84. async def get_session_history(db, user_id, dialog_id, page, limit):
  85. session_list = db.query(ConversationModel).filter(ConversationModel.tenant_id.__eq__(user_id),
  86. ConversationModel.dialog_id.__eq__(dialog_id)).order_by(
  87. ConversationModel.update_time.desc()).limit(limit).offset((page - 1) * limit).all()
  88. return [i.to_json() for i in session_list]
  89. async def create_dialog_service(db, dialog_id, dialog_name, description, icon, dialog_type, mode, user_id):
  90. para = {
  91. "user_input_form": [],
  92. "retriever_resource": {
  93. "enabled": True
  94. },
  95. "file_upload": {
  96. "enabled": False
  97. }
  98. }
  99. try:
  100. dialog_model = DialogModel(id=dialog_id, name=dialog_name, description=description, icon=icon,
  101. dialog_type=dialog_type, tenant_id=user_id, mode=mode, update_date=datetime.now(),
  102. create_date=datetime.now(), parameters=json.dumps(para))
  103. db.add(dialog_model)
  104. db.commit()
  105. db.refresh(dialog_model)
  106. except Exception as e:
  107. logger.error(e)
  108. db.rollback()
  109. return False
  110. return True
  111. async def update_dialog_status_service(db, dialog_id, status, user_id):
  112. try:
  113. dialog = db.query(DialogModel).filter_by(id=dialog_id).first()
  114. dialog.status = status
  115. dialog.update_date = datetime.now()
  116. # db.query(DialogModel).filter_by(id=dialog_id).update({"status":status, "update_date": datetime.now()})
  117. if dialog.dialog_type == DF_TYPE and status == Dialog_STATSU_ON:
  118. chat = ChatBaseApply()
  119. token = await get_chat_token(db, dialog_id)
  120. if not token:
  121. access_token = await get_app_token(db, workflow_server)
  122. # print(workflow)
  123. if access_token:
  124. url = settings.dify_base_url + DF_CHAT_API_KEY.format(dialog_id)
  125. param = await chat.chat_get(url, {}, await chat.get_headers(access_token))
  126. if param and param.get("data"):
  127. token = param.get("data", [{}])[0].get("token")
  128. token_id = param.get("data", [{}])[0].get("id")
  129. await add_chat_token(db, {"id":token_id, "app_id": dialog_id, "type":"app", "token": token})
  130. # dialog.parameters = json.dumps(param)
  131. else:
  132. param = await chat.chat_post(url, {}, await chat.get_headers(access_token))
  133. if param:
  134. token = param.get("token")
  135. token_id = param.get("id")
  136. await add_chat_token(db, {"id": token_id, "app_id": dialog_id, "type": "app", "token": token})
  137. if token:
  138. url = settings.dify_base_url + DF_CHAT_PARAMETERS
  139. param = await chat.chat_get(url, {"user": str(user_id)}, await chat.get_headers(token))
  140. if param:
  141. dialog.parameters = json.dumps(param)
  142. db.commit()
  143. except Exception as e:
  144. logger.error(e)
  145. db.rollback()
  146. return False
  147. return True
  148. async def delete_dialog_service(db, dialog_id):
  149. try:
  150. db.query(DialogModel).filter_by(id=dialog_id).update(
  151. {"status": Dialog_STATSU_DELETE, "update_date": datetime.now()})
  152. db.commit()
  153. except Exception as e:
  154. logger.error(e)
  155. db.rollback()
  156. return False
  157. return True
  158. async def update_dialog_icon_service(db, dialog_id, icon, name, description):
  159. update = {"icon": icon, "update_date": datetime.now()}
  160. if name:
  161. update["name"] = name
  162. if description or description == "":
  163. update["description"] = description
  164. try:
  165. db.query(DialogModel).filter_by(id=dialog_id).update(update)
  166. db.commit()
  167. except Exception as e:
  168. logger.error(e)
  169. db.rollback()
  170. return False
  171. return True
  172. async def get_dialog_manage_list(db, user_id, keyword, label, status, page_size, page_index, mode):
  173. user = db.query(UserModel).filter(UserModel.id == user_id).first()
  174. if user is None:
  175. return {"rows": []}
  176. query = db.query(DialogModel).filter(DialogModel.status != Dialog_STATSU_DELETE)
  177. if user.permission != "admin":
  178. dia_list = [j.id for i in user.groups for j in i.dialogs]
  179. query = query.filter(or_(DialogModel.tenant_id == user_id, DialogModel.id.in_(dia_list)))
  180. if label:
  181. id_list = set(
  182. [i.object_id for i in db.query(LabelWorkerModel).filter(LabelWorkerModel.label_id.in_(label)).all()])
  183. query = query.filter(DialogModel.id.in_(id_list))
  184. if keyword:
  185. query = query.filter(DialogModel.name.like('%{}%'.format(keyword)))
  186. if status:
  187. # print(status)
  188. query = query.filter(DialogModel.status == status)
  189. if mode:
  190. query = query.filter(DialogModel.mode == mode)
  191. query = query.order_by(DialogModel.update_date.desc())
  192. total = query.count()
  193. if page_size:
  194. query = query.limit(page_size).offset((page_index - 1) * page_size)
  195. rows = []
  196. user_id_set = set()
  197. dialog_id_set = set()
  198. label_dict = {}
  199. for kld in query.all():
  200. user_id_set.add(kld.tenant_id)
  201. dialog_id_set.add(kld.id)
  202. rows.append(kld.to_json())
  203. user_dict = {str(i.id): i.to_dict() for i in db.query(UserModel).filter(UserModel.id.in_(user_id_set)).all()}
  204. for i in db.query(LabelModel.id, LabelModel.name, LabelWorkerModel.object_id).outerjoin(LabelWorkerModel,
  205. LabelModel.id == LabelWorkerModel.label_id).filter(
  206. LabelWorkerModel.object_id.in_(dialog_id_set)).all():
  207. label_dict[i.object_id] = label_dict.get(i.object_id, []) + [{"labelId": i.id, "labelName": i.name}]
  208. for r in rows:
  209. r["user"] = user_dict.get(r["user_id"], {})
  210. r["label"] = label_dict.get(r["id"], [])
  211. return {"total": total, "rows": rows}
  212. async def sync_dialog_service(db, dialog_id):
  213. dialog = db.query(DialogModel).filter(DialogModel.id == dialog_id).first()
  214. if dialog and dialog.dialog_type == RG_TYPE:
  215. try:
  216. app_dialog = get_one_from_ragflow_dialog(dialog_id)
  217. if app_dialog:
  218. dialog.name = app_dialog["name"]
  219. dialog.description = app_dialog["description"]
  220. dialog.kb_ids = app_dialog["kb_ids"]
  221. dialog.update_date = datetime.now()
  222. db.add(dialog)
  223. db.commit()
  224. db.refresh(dialog)
  225. except Exception as e:
  226. logger.error(e)
  227. db.rollback()
  228. return False
  229. return True