__init__.py 3.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119
  1. import urllib
  2. from urllib.parse import urlencode
  3. import jwt
  4. from fastapi import FastAPI, Depends, HTTPException
  5. from fastapi.security import OAuth2PasswordBearer
  6. from passlib.context import CryptContext
  7. from pydantic import BaseModel
  8. from starlette import status
  9. from starlette.websockets import WebSocket, WebSocketDisconnect
  10. from app.models.user_model import UserModel
  11. from app.service.auth import SECRET_KEY, ALGORITHM
  12. app = FastAPI()
  13. pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto")
  14. oauth2_scheme = OAuth2PasswordBearer(tokenUrl="token")
  15. class Response(BaseModel):
  16. code: int = 200
  17. msg: str = ""
  18. data: dict = {}
  19. class ResponseList(BaseModel):
  20. code: int = 200
  21. msg: str = ""
  22. data: list[dict] = []
  23. def get_current_user(token: str = Depends(oauth2_scheme)):
  24. try:
  25. payload = jwt.decode(token, SECRET_KEY, algorithms=[ALGORITHM])
  26. username: str = payload.get("sub")
  27. if username is None:
  28. raise HTTPException(
  29. status_code=status.HTTP_401_UNAUTHORIZED,
  30. detail="无法验证凭证",
  31. headers={"WWW-Authenticate": "Bearer"},
  32. )
  33. user = UserModel(username=username, id=payload.get("user_id"))
  34. if user.id == 0:
  35. raise HTTPException(
  36. status_code=status.HTTP_401_UNAUTHORIZED,
  37. detail="用户不存在",
  38. headers={"WWW-Authenticate": "Bearer"},
  39. )
  40. return user
  41. except jwt.PyJWTError:
  42. raise HTTPException(
  43. status_code=status.HTTP_401_UNAUTHORIZED,
  44. detail="令牌无效或已过期",
  45. headers={"WWW-Authenticate": "Bearer"},
  46. )
  47. async def get_current_user_websocket(websocket: WebSocket):
  48. token = websocket.query_params.get('token')
  49. if token is None:
  50. await websocket.close(code=1008)
  51. raise WebSocketDisconnect(code=status.WS_1008_POLICY_VIOLATION)
  52. try:
  53. payload = jwt.decode(token, SECRET_KEY, algorithms=[ALGORITHM])
  54. username: str = payload.get("sub")
  55. if username is None:
  56. await websocket.close(code=1008)
  57. raise WebSocketDisconnect(code=status.WS_1008_POLICY_VIOLATION)
  58. user = UserModel(username=username, id=payload.get("user_id"))
  59. if user is None:
  60. await websocket.close(code=1008)
  61. raise WebSocketDisconnect(code=status.WS_1008_POLICY_VIOLATION)
  62. return user
  63. except jwt.PyJWTError as e:
  64. print(e)
  65. await websocket.close(code=1008)
  66. raise WebSocketDisconnect(code=status.WS_1008_POLICY_VIOLATION)
  67. def format_file_url(agent_id: str, file_url: str, doc_id: str = None, doc_name: str = None) -> str:
  68. if file_url:
  69. # 对 file_url 进行 URL 编码
  70. encoded_file_url = urllib.parse.quote(file_url, safe=':/')
  71. return f"./api/files/download/?url={encoded_file_url}&agent_id={agent_id}"
  72. if doc_id:
  73. # 对 doc_id 和 doc_name 进行 URL 编码
  74. encoded_doc_id = urllib.parse.quote(doc_id, safe='')
  75. encoded_doc_name = urllib.parse.quote(doc_name, safe='')
  76. return f"./api/files/download/?doc_id={encoded_doc_id}&doc_name={encoded_doc_name}&agent_id={agent_id}"
  77. return file_url
  78. def process_files(files, agent_id):
  79. """
  80. 处理文件列表,格式化每个文件的 URL。
  81. :param files: 文件列表,每个文件是一个字典
  82. :param agent_id: 代理 ID
  83. """
  84. if not files:
  85. return # 如果文件列表为空,直接返回
  86. for file in files:
  87. if "file_url" in file and file["file_url"]:
  88. try:
  89. file["file_url"] = format_file_url(agent_id, file["file_url"])
  90. except Exception as e:
  91. # 记录异常信息,但继续处理其他文件
  92. print(f"Error processing file URL: {e}")
  93. if __name__=="__main__":
  94. files1 = [{"file_url": "aaa.com"}, {"file_url":"bbb.com"}]
  95. print(files1)
  96. process_files(files1,11111)
  97. print(files1)