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