dialog.py 9.4 KB

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