dialog.py 7.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187
  1. from datetime import datetime
  2. from sqlalchemy import or_
  3. from app.config.const import Dialog_STATSU_DELETE
  4. from app.models import KnowledgeModel, GroupModel, DialogModel, ConversationModel, group_dialog_table, LabelWorkerModel, \
  5. LabelModel
  6. from app.models.user_model import UserModel
  7. from Log import logger
  8. async def get_dialog_list(db, user_id, keyword, label, status, page_size, page_index):
  9. user = db.query(UserModel).filter(UserModel.id == user_id).first()
  10. if user is None:
  11. return {"rows": []}
  12. query = db.query(DialogModel)
  13. if status:
  14. query = query.filter(DialogModel.status == status)
  15. else:
  16. query = query.filter(DialogModel.status != Dialog_STATSU_DELETE)
  17. id_list = []
  18. if label:
  19. id_list = [i.object_id for i in db.query(LabelWorkerModel).filter(LabelWorkerModel.label_id==label).all()]
  20. if user.permission != "admin":
  21. dia_list = [j.id for i in user.groups for j in i.dialogs if not label or j.id in id_list]
  22. query = query.filter(or_(DialogModel.tenant_id == user_id, DialogModel.id.in_(dia_list)))
  23. else:
  24. if label:
  25. query = query.filter(or_(DialogModel.id.in_(id_list)))
  26. if keyword:
  27. query = query.filter(DialogModel.name.like('%{}%'.format(keyword)))
  28. query = query.order_by(DialogModel.update_date.desc())
  29. total = query.count()
  30. if page_size:
  31. query = query.limit(page_size).offset((page_index - 1) * page_size)
  32. rows = []
  33. user_id_set = set()
  34. dialog_id_set = set()
  35. label_dict = {}
  36. for kld in query.all():
  37. user_id_set.add(kld.tenant_id)
  38. dialog_id_set.add(kld.id)
  39. rows.append(kld.to_json())
  40. user_dict = {str(i.id): i.to_dict() for i in db.query(UserModel).filter(UserModel.id.in_(user_id_set)).all()}
  41. for i in db.query(LabelModel.id, LabelModel.name, LabelWorkerModel.object_id).outerjoin(LabelWorkerModel,
  42. LabelModel.id == LabelWorkerModel.label_id).filter(
  43. LabelWorkerModel.object_id.in_(dialog_id_set)).all():
  44. label_dict[i.object_id] = label_dict.get(i.object_id, []) +[{"labelId": i.id, "labelName": i.name}]
  45. for r in rows:
  46. r["user"] = user_dict.get(r["user_id"], {})
  47. r["label"] = label_dict.get(r["id"], [])
  48. return {"total": total, "rows": rows}
  49. async def update_session_history(db, data: dict, user_id):
  50. session_id = data.get("id")
  51. if not session_id:
  52. logger.error("更新回话记录失败!{}".format(data))
  53. return
  54. data["create_date"] = datetime.strptime(data["create_date"], '%a, %d %b %Y %H:%M:%S %Z')
  55. data["update_date"] = datetime.strptime(data["update_date"], '%a, %d %b %Y %H:%M:%S %Z')
  56. conversation = db.query(ConversationModel).filter(ConversationModel.id == session_id).first()
  57. if not conversation:
  58. try:
  59. data["tenant_id"] = user_id
  60. conversation_model = ConversationModel(**data)
  61. db.add(conversation_model)
  62. db.commit()
  63. except Exception as e:
  64. logger.error(e)
  65. db.rollback()
  66. else:
  67. try:
  68. # data["tenant_id"] = user_id
  69. del data["id"]
  70. db.query(ConversationModel).filter(ConversationModel.id == session_id).update(data)
  71. db.commit()
  72. except Exception as e:
  73. logger.error(e)
  74. db.rollback()
  75. async def get_session_history(db, user_id, dialog_id, page, limit):
  76. session_list = db.query(ConversationModel).filter(ConversationModel.tenant_id.__eq__(user_id),
  77. ConversationModel.dialog_id.__eq__(dialog_id)).order_by(
  78. ConversationModel.update_time.desc()).limit(limit).offset((page - 1) * limit).all()
  79. return [i.to_json() for i in session_list]
  80. async def create_dialog_service(db, dialog_id, dialog_name, description, icon, dialog_type, mode, user_id):
  81. try:
  82. dialog_model = DialogModel(id=dialog_id,name=dialog_name, description=description,icon=icon, dialog_type=dialog_type, tenant_id=user_id, mode=mode,update_date=datetime.now(),create_date=datetime.now())
  83. db.add(dialog_model)
  84. db.commit()
  85. db.refresh(dialog_model)
  86. except Exception as e:
  87. logger.error(e)
  88. db.rollback()
  89. return False
  90. return True
  91. async def update_dialog_status_service(db, dialog_id, status):
  92. try:
  93. db.query(DialogModel).filter_by(id=dialog_id).update({"status":status, "update_date": datetime.now()})
  94. db.commit()
  95. except Exception as e:
  96. logger.error(e)
  97. db.rollback()
  98. return False
  99. return True
  100. async def delete_dialog_service(db, dialog_id):
  101. try:
  102. db.query(DialogModel).filter_by(id=dialog_id).update({"status":Dialog_STATSU_DELETE, "update_date": datetime.now()})
  103. db.commit()
  104. except Exception as e:
  105. logger.error(e)
  106. db.rollback()
  107. return False
  108. return True
  109. async def update_dialog_icon_service(db, dialog_id, icon):
  110. try:
  111. db.query(DialogModel).filter_by(id=dialog_id).update({"icon":icon, "update_date": datetime.now()})
  112. db.commit()
  113. except Exception as e:
  114. logger.error(e)
  115. db.rollback()
  116. return False
  117. return True
  118. async def get_dialog_manage_list(db, user_id, keyword, label, status, page_size, page_index, mode):
  119. user = db.query(UserModel).filter(UserModel.id == user_id).first()
  120. if user is None:
  121. return {"rows": []}
  122. query = db.query(DialogModel).filter(DialogModel.status != Dialog_STATSU_DELETE)
  123. id_list = []
  124. if label:
  125. id_list = set([i.object_id for i in db.query(LabelWorkerModel).filter(LabelWorkerModel.label_id.in_(label)).all()])
  126. if user.permission != "admin":
  127. dia_list = [j.id for i in user.groups for j in i.dialogs if not label or j.id in id_list]
  128. query = query.filter(or_(DialogModel.tenant_id == user_id, DialogModel.id.in_(dia_list)))
  129. else:
  130. if label:
  131. query = query.filter(or_(DialogModel.id.in_(id_list)))
  132. if keyword:
  133. query = query.filter(DialogModel.name.like('%{}%'.format(keyword)))
  134. if status:
  135. # print(status)
  136. query = query.filter(DialogModel.status == status)
  137. if mode:
  138. query = query.filter(DialogModel.mode == mode)
  139. query = query.order_by(DialogModel.update_date.desc())
  140. total = query.count()
  141. if page_size:
  142. query = query.limit(page_size).offset((page_index - 1) * page_size)
  143. rows = []
  144. user_id_set = set()
  145. dialog_id_set = set()
  146. label_dict = {}
  147. for kld in query.all():
  148. user_id_set.add(kld.tenant_id)
  149. dialog_id_set.add(kld.id)
  150. rows.append(kld.to_json())
  151. user_dict = {str(i.id): i.to_dict() for i in db.query(UserModel).filter(UserModel.id.in_(user_id_set)).all()}
  152. for i in db.query(LabelModel.id, LabelModel.name, LabelWorkerModel.object_id).outerjoin(LabelWorkerModel,
  153. LabelModel.id == LabelWorkerModel.label_id).filter(
  154. LabelWorkerModel.object_id.in_(dialog_id_set)).all():
  155. label_dict[i.object_id] = label_dict.get(i.object_id, []) +[{"labelId": i.id, "labelName": i.name}]
  156. for r in rows:
  157. r["user"] = user_dict.get(r["user_id"], {})
  158. r["label"] = label_dict.get(r["id"], [])
  159. return {"total": total, "rows": rows}