dialog.py 10 KB

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