__init__.py 4.4 KB

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