agent.py 6.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137
  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
  8. from app.config.config import settings
  9. from app.models.agent_model import AgentType, AgentModel
  10. from app.models.base_model import get_db
  11. from app.models.user_model import UserModel
  12. from app.service.bisheng import BishengService
  13. from app.service.ragflow import RagflowService
  14. from app.service.token import get_ragflow_token, get_bisheng_token
  15. router = APIRouter()
  16. @router.get("/list", response_model=ResponseList)
  17. async def agent_list(db: Session = Depends(get_db)):
  18. agents = db.query(AgentModel).order_by(AgentModel.sort.asc()).all()
  19. result = [item.to_dict() for item in agents]
  20. return ResponseList(code=200, msg="", data=result)
  21. @router.get("/{agent_id}/sessions", response_model=ResponseList)
  22. async def chat_list(agent_id: str, db: Session = Depends(get_db), current_user: UserModel = Depends(get_current_user)):
  23. agent = db.query(AgentModel).filter(AgentModel.id == agent_id).first()
  24. if not agent:
  25. return ResponseList(code=404, msg="Agent not found")
  26. if agent.agent_type == AgentType.RAGFLOW:
  27. ragflow_service = RagflowService(base_url=settings.fwr_base_url)
  28. try:
  29. token = get_ragflow_token(db, current_user.id)
  30. result = await ragflow_service.get_chat_sessions(token, agent_id)
  31. except Exception as e:
  32. raise HTTPException(status_code=500, detail=str(e))
  33. return ResponseList(code=200, msg="", data=result)
  34. elif agent.agent_type == AgentType.BISHENG:
  35. bisheng_service = BishengService(base_url=settings.sgb_base_url)
  36. try:
  37. token = get_bisheng_token(db, current_user.id)
  38. result = await bisheng_service.get_chat_sessions(token)
  39. except Exception as e:
  40. raise HTTPException(status_code=500, detail=str(e))
  41. return ResponseList(code=200, msg="", data=result)
  42. else:
  43. return ResponseList(code=200, msg="Unsupported agent type")
  44. @router.get("/{agent_id}/{conversation_id}/session_log")
  45. async def session_log(agent_id: str, conversation_id: str, db: Session = Depends(get_db), current_user: UserModel = Depends(get_current_user)):
  46. agent = db.query(AgentModel).filter(AgentModel.id == agent_id).first()
  47. if not agent:
  48. return Response(code=404, msg="Agent not found")
  49. if agent.agent_type == AgentType.RAGFLOW:
  50. ragflow_service = RagflowService(base_url=settings.fwr_base_url)
  51. try:
  52. token = get_ragflow_token(db, current_user.id)
  53. result = await ragflow_service.get_session_log(token, conversation_id)
  54. if 'session_log' in result and 'reference' in result:
  55. combined_logs = []
  56. last_question = None
  57. references = result['reference']
  58. reference_index = 0
  59. for session in result['session_log']:
  60. if session['role'] == 'user':
  61. last_question = session['message']
  62. elif session['role'] == 'assistant' and last_question:
  63. if reference_index < len(references):
  64. reference = references[reference_index]
  65. else:
  66. reference = None
  67. combined_logs.append({
  68. 'question': last_question,
  69. 'answer': session['message'],
  70. 'reference': reference
  71. })
  72. last_question = None
  73. reference_index += 1
  74. return JSONResponse(status_code=200, content={"code": 200, "data": combined_logs})
  75. else:
  76. return JSONResponse(status_code=200, content={"code": 400, "message": "Invalid result structure"})
  77. except Exception as e:
  78. raise HTTPException(status_code=500, detail=str(e))
  79. if agent.agent_type == AgentType.BISHENG:
  80. bisheng_service = BishengService(base_url=settings.sgb_base_url)
  81. try:
  82. token = get_bisheng_token(db, current_user.id)
  83. result = await bisheng_service.get_session_log(token, agent_id, conversation_id)
  84. combined_logs = []
  85. last_question = None
  86. for session in result:
  87. message = session['message']
  88. # 判断message是字符串还是json 对象,如果是json取其中的question字段,或者report_name字段赋值给message
  89. try:
  90. message_json = json.loads(message)
  91. if 'question' in message_json:
  92. message = message_json['question']
  93. elif 'query' in message_json:
  94. message = message_json['query']
  95. elif 'report_name' in message_json:
  96. message = message_json['report_name']
  97. except json.JSONDecodeError:
  98. pass
  99. if session['role'] == 'question':
  100. last_question = message
  101. elif session['role'] == 'answer' and last_question:
  102. combined_logs.append({
  103. 'question': last_question,
  104. 'answer': message
  105. })
  106. last_question = None
  107. return JSONResponse(status_code=200, content={"code": 200, "data": combined_logs})
  108. except Exception as e:
  109. raise HTTPException(status_code=500, detail=str(e))
  110. else:
  111. return JSONResponse(status_code=200, content={"code": 200, "log": "Unsupported agent type"})
  112. @router.get("/get-chat-id/{agent_id}", response_model=Response)
  113. async def get_chat_id(agent_id: str, db: Session = Depends(get_db)):
  114. agent = db.query(AgentModel).filter(AgentModel.id == agent_id).first()
  115. if not agent:
  116. return Response(code=404, msg="Agent not found")
  117. return Response(code=200, msg="", data={"chat_id": uuid.uuid4().hex})