dialog.py 2.9 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364656667686970717273747576
  1. from datetime import datetime
  2. from app.models import KnowledgeModel, GroupModel, DialogModel, ConversationModel, group_dialog_table
  3. from app.models.user_model import UserModel
  4. from Log import logger
  5. async def get_dialog_list(db, user_id, keyword, page_size, page_index):
  6. user = db.query(UserModel).filter(UserModel.id == user_id).first()
  7. if user is None:
  8. return {"rows": []}
  9. if user.permission == "admin":
  10. query = db.query(DialogModel)
  11. else:
  12. group_list = [i.id for i in user.groups]
  13. query = db.query(DialogModel)
  14. query = query.filter(DialogModel.tenant_id == user_id)
  15. query = query.union(
  16. db.query(DialogModel).join(
  17. group_dialog_table,
  18. DialogModel.id == group_dialog_table.c.dialog_id
  19. ).filter(
  20. group_dialog_table.c.group_id.in_(group_list)
  21. )
  22. )
  23. if keyword:
  24. query = query.filter(DialogModel.name.like('%{}%'.format(keyword)))
  25. total = query.count()
  26. if page_size:
  27. query = query.limit(page_size).offset((page_index - 1) * page_size)
  28. rows = []
  29. user_id_set = set()
  30. for kld in query.all():
  31. user_id_set.add(kld.tenant_id)
  32. rows.append(kld.to_json())
  33. print(rows)
  34. user_dict = {i.id: i.to_dict() for i in db.query(UserModel).filter(UserModel.id.in_(user_id_set)).all()}
  35. for r in rows:
  36. r["user"] = user_dict.get(r["user_id"], {})
  37. return {"total": total, "rows": rows}
  38. async def update_session_history(db, data: dict, user_id):
  39. session_id = data.get("id")
  40. if not session_id:
  41. logger.error("更新回话记录失败!{}".format(data))
  42. return
  43. data["create_date"] = datetime.strptime(data["create_date"], '%a, %d %b %Y %H:%M:%S %Z')
  44. data["update_date"] = datetime.strptime(data["update_date"], '%a, %d %b %Y %H:%M:%S %Z')
  45. conversation = db.query(ConversationModel).filter(ConversationModel.id == session_id).first()
  46. if not conversation:
  47. try:
  48. data["tenant_id"] = user_id
  49. conversation_model = ConversationModel(**data)
  50. db.add(conversation_model)
  51. db.commit()
  52. except Exception as e:
  53. logger.error(e)
  54. db.rollback()
  55. else:
  56. try:
  57. # data["tenant_id"] = user_id
  58. del data["id"]
  59. db.query(ConversationModel).filter(ConversationModel.id == session_id).update(data)
  60. db.commit()
  61. except Exception as e:
  62. logger.error(e)
  63. db.rollback()
  64. async def get_session_history(db, user_id, dialog_id):
  65. session_list = db.query(ConversationModel).filter(ConversationModel.tenant_id.__eq__(user_id),
  66. ConversationModel.dialog_id.__eq__(dialog_id)).order_by(
  67. ConversationModel.update_time.desc()).all()
  68. return [i.to_json() for i in session_list]