files.py 4.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111
  1. from typing import Optional
  2. import requests
  3. from fastapi import Depends, APIRouter, HTTPException, UploadFile, File, Query, Form
  4. from pydantic import BaseModel
  5. from sqlalchemy.orm import Session
  6. from starlette.responses import StreamingResponse
  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.basic import BasicService
  13. from app.service.bisheng import BishengService
  14. from app.service.ragflow import RagflowService
  15. from app.service.service_token import get_ragflow_token, get_bisheng_token
  16. import urllib.parse
  17. router = APIRouter()
  18. @router.post("/upload/{agent_id}", response_model=Response)
  19. async def upload_file(agent_id: str,
  20. file: UploadFile = File(...),
  21. chat_id: str = Query(None, description="The ID of the chat"),
  22. db: Session = Depends(get_db),
  23. current_user: UserModel = Depends(get_current_user)
  24. ):
  25. agent = db.query(AgentModel).filter(AgentModel.id == agent_id).first()
  26. if not agent:
  27. return Response(code=404, msg="Agent not found")
  28. # 读取上传的文件内容
  29. try:
  30. file_content = await file.read()
  31. except Exception as e:
  32. return Response(code=400, msg=str(e))
  33. if agent.agent_type == AgentType.RAGFLOW:
  34. token = get_ragflow_token(db, current_user.id)
  35. ragflow_service = RagflowService(base_url=settings.fwr_base_url)
  36. # 查询会话是否存在,不存在先创建会话
  37. history = await ragflow_service.get_session_history(token, chat_id)
  38. if len(history) == 0:
  39. message = {"role": "user", "message": file.filename}
  40. await ragflow_service.set_session(token, agent_id, message, chat_id, True)
  41. ragflow_service = RagflowService(base_url=settings.fwr_base_url)
  42. token = get_ragflow_token(db, current_user.id)
  43. doc_ids = await ragflow_service.upload_and_parse(token, chat_id, file.filename, file_content)
  44. return Response(code=200, msg="", data={"doc_ids": doc_ids, "file_name": file.filename})
  45. elif agent.agent_type == AgentType.BISHENG:
  46. bisheng_service = BishengService(base_url=settings.sgb_base_url)
  47. try:
  48. token = get_bisheng_token(db, current_user.id)
  49. result = await bisheng_service.upload(token, file.filename, file_content)
  50. except Exception as e:
  51. raise HTTPException(status_code=500, detail=str(e))
  52. result["file_name"] = file.filename
  53. return Response(code=200, msg="", data=result)
  54. elif agent.agent_type == AgentType.BASIC:
  55. if agent_id == "basic_excel_talk":
  56. service = BasicService(base_url=settings.basic_base_url)
  57. result = await service.excel_talk_upload(chat_id, file.filename, file_content)
  58. return Response(code=200, msg="", data=result)
  59. else:
  60. return Response(code=200, msg="Unsupported agent type")
  61. @router.get("/download/", response_model=Response)
  62. async def download_file(
  63. url: Optional[str] = Query(None, description="URL of the file to download for bisheng"),
  64. agent_id: str = Query(..., description="Agent ID"),
  65. doc_id: Optional[str] = Query(None, description="Optional doc id for ragflow agents"),
  66. doc_name: Optional[str] = Query(None, description="Optional doc name for ragflow agents"),
  67. db: Session = Depends(get_db)
  68. ):
  69. agent = db.query(AgentModel).filter(AgentModel.id == agent_id).first()
  70. if not agent:
  71. return Response(code=404, msg="Agent not found")
  72. if agent.agent_type == AgentType.BISHENG:
  73. url = urllib.parse.unquote(url)
  74. # 从 URL 中提取文件名
  75. parsed_url = urllib.parse.urlparse(url)
  76. filename = urllib.parse.unquote(parsed_url.path.split('/')[-1])
  77. url = url.replace("http://minio:9000", settings.sgb_base_url)
  78. elif agent.agent_type == AgentType.RAGFLOW:
  79. if not doc_id:
  80. return Response(code=400, msg="doc_id is required")
  81. url = f"{settings.fwr_base_url}/v1/document/get/{doc_id}"
  82. filename = doc_name
  83. else:
  84. return Response(code=400, msg="Unsupported agent type")
  85. try:
  86. # 发送GET请求获取文件内容
  87. response = requests.get(url, stream=True)
  88. response.raise_for_status() # 检查请求是否成功
  89. # 返回流式响应
  90. return StreamingResponse(
  91. response.iter_content(chunk_size=1024),
  92. media_type="application/octet-stream",
  93. headers={"Content-Disposition": f"attachment; filename*=utf-8''{urllib.parse.quote(filename)}"}
  94. )
  95. except Exception as e:
  96. raise HTTPException(status_code=400, detail=f"Error downloading file: {e}")