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