__init__.py 6.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170
  1. import urllib
  2. from datetime import datetime
  3. from typing import Callable, Any
  4. from urllib.parse import urlencode
  5. import jwt
  6. # from cryptography.fernet import Fernet
  7. from fastapi import FastAPI, Depends, HTTPException, Header, Request
  8. from fastapi.security import OAuth2PasswordBearer
  9. from passlib.context import CryptContext
  10. from pydantic import BaseModel
  11. from starlette import status
  12. from starlette.websockets import WebSocket, WebSocketDisconnect
  13. from Log import logger
  14. from app.models.base_model import SessionLocal
  15. # from app.models.app_model import AppRegisterModel
  16. from app.models.user_model import UserModel, UserApiTokenModel
  17. from app.service.auth import SECRET_KEY, ALGORITHM
  18. from app.config.config import settings
  19. app = FastAPI()
  20. pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto")
  21. oauth2_scheme = OAuth2PasswordBearer(tokenUrl="token")
  22. # cipher_suite = Fernet(settings.HASH_SUB_KEY)
  23. class Response(BaseModel):
  24. code: int = 200
  25. msg: str = ""
  26. data: dict = {}
  27. class ResponseList(BaseModel):
  28. code: int = 200
  29. msg: str = ""
  30. data: list[dict] = []
  31. def verify_token(token: str) -> Any:
  32. """
  33. 验证 Token 是否有效
  34. """
  35. db = SessionLocal()
  36. try:
  37. db_token = db.query(UserApiTokenModel).filter(UserApiTokenModel.token == token, UserApiTokenModel.is_active == 1).first()
  38. return db_token is not None and (db_token.expires_at is None or db_token.expires_at > datetime.now())
  39. finally:
  40. db.close()
  41. def token_required()-> Callable:
  42. def decorated_function(request: Request)-> Any:
  43. authorization_str = request.headers.get("Authorization")
  44. if not authorization_str:
  45. raise HTTPException(status_code=401, detail="Authorization` can't be empty")
  46. authorization_list = authorization_str.split()
  47. if len(authorization_list) < 2:
  48. raise HTTPException(status_code=401, detail="Invalid token")
  49. token = authorization_list[1]
  50. objs = verify_token(token)
  51. if not objs:
  52. raise HTTPException(status_code=401, detail="Invalid token")
  53. user = UserModel(username="", id=objs.user_id)
  54. return user
  55. return decorated_function
  56. def get_current_user(token: str = Depends(oauth2_scheme)):
  57. try:
  58. payload = jwt.decode(token, SECRET_KEY, algorithms=[ALGORITHM])
  59. expired_time = payload.get("lex")
  60. if not expired_time:
  61. raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="令牌无效或已过期",
  62. headers={"WWW-Authenticate": "Bearer"})
  63. if datetime.strptime(expired_time, "%Y-%m-%d %H:%M:%S") < datetime.now():
  64. raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="系统授权已过期!",
  65. headers={"WWW-Authenticate": "Bearer"})
  66. username: str = payload.get("sub")
  67. if username is None:
  68. raise HTTPException(
  69. status_code=status.HTTP_401_UNAUTHORIZED,
  70. detail="无法验证凭证",
  71. headers={"WWW-Authenticate": "Bearer"},
  72. )
  73. user = UserModel(username=username, id=payload.get("user_id"))
  74. if user.id == 0:
  75. raise HTTPException(
  76. status_code=status.HTTP_401_UNAUTHORIZED,
  77. detail="用户不存在",
  78. headers={"WWW-Authenticate": "Bearer"},
  79. )
  80. return user
  81. except jwt.PyJWTError:
  82. raise HTTPException(
  83. status_code=status.HTTP_401_UNAUTHORIZED,
  84. detail="令牌无效或已过期",
  85. headers={"WWW-Authenticate": "Bearer"},
  86. )
  87. async def get_current_user_websocket(websocket: WebSocket):
  88. token = websocket.query_params.get('token')
  89. if token is None:
  90. await websocket.close(code=1008)
  91. raise WebSocketDisconnect(code=status.WS_1008_POLICY_VIOLATION)
  92. try:
  93. payload = jwt.decode(token, SECRET_KEY, algorithms=[ALGORITHM])
  94. username: str = payload.get("sub")
  95. if username is None:
  96. await websocket.close(code=1008)
  97. raise WebSocketDisconnect(code=status.WS_1008_POLICY_VIOLATION)
  98. user = UserModel(username=username, id=payload.get("user_id"))
  99. if user is None:
  100. await websocket.close(code=1008)
  101. raise WebSocketDisconnect(code=status.WS_1008_POLICY_VIOLATION)
  102. return user
  103. except jwt.PyJWTError as e:
  104. print(e)
  105. await websocket.close(code=1008)
  106. raise WebSocketDisconnect(code=status.WS_1008_POLICY_VIOLATION)
  107. def format_file_url(agent_id: str, file_url: str, doc_id: str = None, doc_name: str = None) -> str:
  108. if file_url:
  109. # 对 file_url 进行 URL 编码
  110. encoded_file_url = urllib.parse.quote(file_url, safe=':/')
  111. return f"./api/files/download/?url={encoded_file_url}&agent_id={agent_id}"
  112. if doc_id:
  113. # 对 doc_id 和 doc_name 进行 URL 编码
  114. encoded_doc_id = urllib.parse.quote(doc_id, safe='')
  115. encoded_doc_name = urllib.parse.quote(doc_name, safe='')
  116. return f"./api/files/download/?doc_id={encoded_doc_id}&doc_name={encoded_doc_name}&agent_id={agent_id}"
  117. return file_url
  118. def process_files(files, agent_id):
  119. """
  120. 处理文件列表,格式化每个文件的 URL。
  121. :param files: 文件列表,每个文件是一个字典
  122. :param agent_id: 代理 ID
  123. """
  124. if not files:
  125. return # 如果文件列表为空,直接返回
  126. for file in files:
  127. if "file_url" in file and file["file_url"]:
  128. try:
  129. file["file_url"] = format_file_url(agent_id, file["file_url"])
  130. except Exception as e:
  131. # 记录异常信息,但继续处理其他文件
  132. print(f"Error processing file URL: {e}")
  133. def get_api_key(authorization: str = Header(...)):
  134. if not authorization.startswith("Bearer "):
  135. raise HTTPException(status_code=401, detail="Invalid Authorization header format.")
  136. return authorization.split(" ")[1]
  137. if __name__=="__main__":
  138. files1 = [{"file_url": "aaa.com"}, {"file_url":"bbb.com"}]
  139. print(files1)
  140. process_files(files1,11111)
  141. print(files1)