agent.py 15 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300
  1. import json
  2. import uuid
  3. from fastapi import Depends, APIRouter, Query, HTTPException
  4. from fastapi.responses import JSONResponse
  5. from pydantic import BaseModel
  6. from sqlalchemy.orm import Session
  7. from app.api import Response, get_current_user, ResponseList, process_files
  8. from app.api.user import reset_user_pwd
  9. from app.config.config import settings
  10. from app.models import DialogModel, MenuCapacityModel
  11. from app.models.agent_model import AgentType, AgentModel
  12. from app.models.base_model import get_db
  13. from app.models.session_model import SessionModel
  14. from app.models.user_model import UserModel
  15. from app.service.bisheng import BishengService
  16. from app.service.dialog import get_session_history
  17. from app.service.ragflow import RagflowService
  18. from app.service.service_token import get_ragflow_token, get_bisheng_token
  19. # from app.task.fetch_agent import initialize_agents
  20. router = APIRouter()
  21. @router.get("/list", response_model=ResponseList)
  22. async def agent_list(db: Session = Depends(get_db)):
  23. agents = db.query(AgentModel).order_by(AgentModel.sort.asc()).all()
  24. result = [item.to_dict() for item in agents]
  25. return ResponseList(code=200, msg="", data=result)
  26. @router.get("/{agent_id}/sessions", response_model=ResponseList)
  27. async def chat_list(
  28. agent_id: str,
  29. page: int = Query(1, ge=1),
  30. limit: int = Query(1000, ge=1, le=1000),
  31. db: Session = Depends(get_db),
  32. current_user: UserModel = Depends(get_current_user)):
  33. # agent = db.query(AgentModel).filter(AgentModel.id == agent_id).first()
  34. agent = db.query(MenuCapacityModel).filter(MenuCapacityModel.chat_id == agent_id).first()
  35. if not agent:
  36. return ResponseList(code=404, msg="Agent not found")
  37. agent_type = int(agent.capacity_type)
  38. if agent_type == AgentType.RAGFLOW:
  39. ragflow_service = RagflowService(base_url=settings.fwr_base_url)
  40. try:
  41. result = await get_session_history(db, current_user.id, agent_id, page, limit)
  42. if not result:
  43. token = await get_ragflow_token(db, current_user.id)
  44. result = await ragflow_service.get_chat_sessions(token, agent_id)
  45. except Exception as e:
  46. print(e)
  47. raise HTTPException(status_code=500, detail=str(e))
  48. return ResponseList(code=200, msg="", data=result)
  49. elif agent_type == AgentType.BISHENG:
  50. bisheng_service = BishengService(base_url=settings.sgb_base_url)
  51. try:
  52. token = await get_bisheng_token(db, current_user.id)
  53. result = await bisheng_service.get_chat_sessions(token, agent_id, page, limit)
  54. except Exception as e:
  55. raise HTTPException(status_code=500, detail=str(e))
  56. return ResponseList(code=200, msg="", data=result)
  57. elif agent_type == AgentType.BASIC:
  58. offset = (page - 1) * limit
  59. records = db.query(SessionModel).filter(SessionModel.agent_id == agent_id, SessionModel.tenant_id==current_user.id).order_by(SessionModel.create_date.desc()).offset(offset).limit(limit).all()
  60. result = [item.to_dict() for item in records]
  61. return ResponseList(code=200, msg="", data=result)
  62. elif agent_type == AgentType.DIFY:
  63. offset = (page - 1) * limit
  64. records = db.query(SessionModel).filter(SessionModel.agent_id == agent_id, SessionModel.tenant_id==current_user.id).order_by(SessionModel.create_date.desc()).offset(offset).limit(limit).all()
  65. result = [item.to_dict() for item in records]
  66. return ResponseList(code=200, msg="", data=result)
  67. else:
  68. return ResponseList(code=200, msg="Unsupported agent type")
  69. @router.get("/{agent_id}/{conversation_id}/session_log")
  70. async def session_log(agent_id: str, conversation_id: str, db: Session = Depends(get_db), current_user: UserModel = Depends(get_current_user)):
  71. # agent = db.query(AgentModel).filter(AgentModel.id == agent_id).first()
  72. # if not agent:
  73. # return Response(code=404, msg="Agent not found")
  74. agent = db.query(MenuCapacityModel).filter(MenuCapacityModel.chat_id == agent_id).first()
  75. if not agent:
  76. return ResponseList(code=404, msg="Agent not found")
  77. agent_type = int(agent.capacity_type)
  78. if agent_type == AgentType.RAGFLOW:
  79. ragflow_service = RagflowService(base_url=settings.fwr_base_url)
  80. try:
  81. token = await get_ragflow_token(db, current_user.id)
  82. result = await ragflow_service.get_session_log(token, conversation_id)
  83. if 'session_log' in result and 'reference' in result:
  84. combined_logs = []
  85. last_question = None
  86. references = result['reference']
  87. reference_index = 0
  88. for session in result['session_log']:
  89. if session['role'] == 'user':
  90. last_question = session['message']
  91. elif session['role'] == 'assistant' and last_question:
  92. if reference_index < len(references):
  93. reference = references[reference_index]
  94. else:
  95. reference = None
  96. combined_logs.append({
  97. 'question': last_question,
  98. 'answer': session['message'],
  99. 'reference': reference
  100. })
  101. last_question = None
  102. reference_index += 1
  103. return JSONResponse(status_code=200, content={"code": 200, "data": combined_logs})
  104. else:
  105. return JSONResponse(status_code=200, content={"code": 400, "message": "Invalid result structure"})
  106. except Exception as e:
  107. raise HTTPException(status_code=500, detail=str(e))
  108. elif agent_type == AgentType.BISHENG:
  109. is_join = False
  110. if agent.name == "报告生成":
  111. is_join = True
  112. bisheng_service = BishengService(base_url=settings.sgb_base_url)
  113. try:
  114. token = await get_bisheng_token(db, current_user.id)
  115. result = await bisheng_service.get_session_log(token, agent_id, conversation_id)
  116. combined_logs = []
  117. last_question = None
  118. answer_str = ""
  119. files = []
  120. for session in result:
  121. # print(session)
  122. # 检查 session 是否为 None
  123. if session is None:
  124. continue
  125. message = session.get('message')
  126. # 判断 message 是否是字符串,然后尝试解析为 JSON 对象
  127. if isinstance(message, str):
  128. try:
  129. message_json = json.loads(message)
  130. if 'question' in message_json:
  131. message = message_json['question']
  132. elif 'query' in message_json:
  133. message = message_json['query']
  134. elif 'report_name' in message_json:
  135. message = message_json['report_name']
  136. except Exception as e:
  137. pass # 非 JSON 字符串,继续使用原始 message
  138. if session.get('files') and isinstance(session.get('files'), str):
  139. try:
  140. files = json.loads(session.get('files'))
  141. process_files(files, agent_id)
  142. except Exception as e:
  143. pass # 非 JSON 字符串,继续使用原始 message
  144. # 检查 message 是否为 None
  145. if message is None:
  146. continue
  147. if is_join:
  148. ...
  149. if session.get('role') == 'question':
  150. last_question = message
  151. elif session.get('role') == 'answer':
  152. answer_str += message
  153. else:
  154. if session.get('role') == 'question':
  155. last_question = message
  156. elif session.get('role') == 'answer' and last_question:
  157. combined_logs.append({
  158. 'question': last_question,
  159. 'answer': message
  160. })
  161. last_question = None
  162. return JSONResponse(status_code=200, content={"code": 200, "data": combined_logs if combined_logs else [{'question': last_question,
  163. 'answer': answer_str, 'files': files}]})
  164. except Exception as e:
  165. raise HTTPException(status_code=500, detail=str(e))
  166. elif agent_type == AgentType.BASIC:
  167. data = []
  168. session = db.query(SessionModel).filter(SessionModel.id == conversation_id).first()
  169. if session:
  170. tmp_data = {}
  171. for i in session.log_to_json().get("message", []):
  172. if i.get("role") == "user":
  173. tmp_data["question"]=i.get("content")
  174. if "upload_filenames" in i:
  175. tmp_data["upload_filenames"] = i.get("upload_filenames")
  176. elif i.get("role") == "assistant":
  177. if isinstance(i.get("content"), dict):
  178. tmp_data["answer"] = i.get("content", {}).get("message")
  179. if "file_name" in i.get("content", {}):
  180. tmp_data["files"] = [{"file_name":i.get("content", {}).get("file_name"), "file_url":i.get("content", {}).get("file_url")}]
  181. else:
  182. tmp_data["answer"] = i.get("content")
  183. if "excel_url" in i:
  184. tmp_data["excel_url"] = i.get("excel_url")
  185. if "image_url" in i:
  186. tmp_data["image_url"] = i.get("image_url")
  187. if "sql" in i:
  188. tmp_data["sql"] = i.get("sql")
  189. if "code" in i:
  190. tmp_data["code"] = i.get("code")
  191. if "e" in i:
  192. tmp_data["e"] = i.get("e")
  193. if "image_name" in i:
  194. tmp_data["image_name"] = i.get("image_name")
  195. if "excel_name" in i:
  196. tmp_data["excel_name"] = i.get("excel_name")
  197. data.append(tmp_data)
  198. tmp_data = {}
  199. if tmp_data:
  200. data.append(tmp_data)
  201. return JSONResponse(status_code=200, content={"code": 200, "data": data})
  202. elif agent_type == AgentType.DIFY:
  203. data = []
  204. session = db.query(SessionModel).filter(SessionModel.id == conversation_id).first()
  205. if session:
  206. tmp_data = {}
  207. for i in session.log_to_json().get("message", []):
  208. if i.get("role") == "user":
  209. tmp_data["question"] = i.get("content")
  210. if "upload_filenames" in i:
  211. tmp_data["upload_filenames"] = i.get("upload_filenames")
  212. elif i.get("role") == "assistant":
  213. if isinstance(i.get("content"), dict):
  214. content = i.get("content", {})
  215. tmp_data["answer"] = i.get("content", {}).get("answer")
  216. if "file_name" in i.get("content", {}):
  217. tmp_data["files"] = [{"file_name": i.get("content", {}).get("file_name"),
  218. "file_url": i.get("content", {}).get("file_url")}]
  219. if "images" in i.get("content", {}):
  220. tmp_data["images"] = i.get("content", {}).get("images")
  221. if "download_url" in i.get("content", {}):
  222. tmp_data["download_url"] = i.get("content", {}).get("download_url")
  223. if "node_list" in content:
  224. node_dict = {
  225. "node_data": [],
  226. # {"title": "去除冗余", # 节点名称 "status": "succeeded", # 节点状态"created_at": 1735817337, # 开始时间"finished_at": 1735817337, # 结束时间"error": "" # 错误日志}
  227. "total_tokens": 0, # 花费token数
  228. "created_at": 0, # 开始时间
  229. "finished_at": 0, # 结束时间
  230. "elapsed_time": 0, # 结束时间
  231. "status": "succeeded", # 工作流状态
  232. "error": "", # 错误日志
  233. }
  234. for node in content["node_list"]:
  235. if node.get("event") == "node_finished":
  236. node_dict["node_data"].append({
  237. "title": node.get("data", {}).get("title", ""),
  238. "status": node.get("data", {}).get("status", ""),
  239. "created_at": node.get("data", {}).get("created_at", 0),
  240. "finished_at": node.get("data", {}).get("finished_at", 0),
  241. "node_type": node.get("data", {}).get("node_type", 0),
  242. "elapsed_time": node.get("data", {}).get("elapsed_time", 0),
  243. "error": node.get("data", {}).get("error", ""),
  244. })
  245. elif node.get("event") == "workflow_finished":
  246. node_dict["total_tokens"] = node.get("data", {}).get("total_tokens", 0)
  247. node_dict["created_at"] = node.get("data", {}).get("created_at", 0)
  248. node_dict["finished_at"] = node.get("data", {}).get("finished_at", 0)
  249. node_dict["status"] = node.get("data", {}).get("status", "")
  250. node_dict["error"] = node.get("data", {}).get("error", "")
  251. node_dict["elapsed_time"] = node.get("data", {}).get("elapsed_time", 0)
  252. tmp_data["workflow"] = node_dict
  253. else:
  254. tmp_data["answer"] = i.get("content")
  255. data.append(tmp_data)
  256. tmp_data = {}
  257. if tmp_data:
  258. data.append(tmp_data)
  259. return JSONResponse(status_code=200, content={"code": 200, "data": data})
  260. else:
  261. return JSONResponse(status_code=200, content={"code": 200, "log": "Unsupported agent type"})
  262. @router.get("/get-chat-id/{agent_id}", response_model=Response)
  263. async def get_chat_id(agent_id: str, db: Session = Depends(get_db)):
  264. # agent = db.query(AgentModel).filter(AgentModel.id == agent_id).first()
  265. # if not agent:
  266. # return Response(code=404, msg="Agent not found")
  267. return Response(code=200, msg="", data={"chat_id": uuid.uuid4().hex})