auth.py 8.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222
  1. import os.path
  2. import re
  3. import uuid
  4. import base64
  5. from datetime import datetime, timedelta
  6. from typing import Type
  7. from jwt import encode, decode, exceptions
  8. from passlib.context import CryptContext
  9. from fastapi import HTTPException, status
  10. from sqlalchemy.orm import Session
  11. from Log import logger
  12. from app.config.config import settings
  13. from app.config.const import USER_STATSU_DELETE, APP_SERVICE_PATH
  14. from app.models import RoleModel, GroupModel, TokenModel
  15. from app.models.user_model import UserModel, UserAppModel
  16. from cryptography.hazmat.backends import default_backend
  17. from cryptography.hazmat.primitives import serialization
  18. from cryptography.hazmat.primitives.asymmetric import padding
  19. SECRET_KEY = settings.secret_key
  20. ALGORITHM = "HS256"
  21. ACCESS_TOKEN_EXPIRE_MINUTES = 24*60
  22. pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto")
  23. def verify_password(plain_password, hashed_password):
  24. return pwd_context.verify(plain_password, hashed_password)
  25. def get_password_hash(password):
  26. return pwd_context.hash(password)
  27. def authenticate_user(db, username: str, password: str):
  28. user = db.query(UserModel).filter(UserModel.username == username, UserModel.status != USER_STATSU_DELETE).first()
  29. if not user:
  30. return False
  31. if not verify_password(password, user.hashed_password):
  32. return False
  33. return user
  34. def create_access_token(data: dict, expires_delta: timedelta = None):
  35. to_encode = data.copy()
  36. if expires_delta:
  37. expire = datetime.utcnow() + expires_delta
  38. else:
  39. expire = datetime.utcnow() + timedelta(minutes=ACCESS_TOKEN_EXPIRE_MINUTES)
  40. to_encode.update({"exp": expire})
  41. encoded_jwt = encode(to_encode, SECRET_KEY, algorithm=ALGORITHM)
  42. return encoded_jwt
  43. def decode_access_token(token: str):
  44. try:
  45. payload = decode(token, SECRET_KEY, algorithms=[ALGORITHM])
  46. return payload
  47. except exceptions.DecodeError:
  48. raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Could not validate credentials")
  49. def is_valid_password(password: str) -> bool:
  50. if len(password) < 8:
  51. return False
  52. has_digit = re.search(r'[0-9]', password)
  53. has_letter = re.search(r'[A-Za-z]', password)
  54. # 如果密码包含数字和字母,则返回True,否则返回None
  55. return has_digit is not None and has_letter is not None
  56. async def save_register_user(db, username, password, email, app_password, register_dict):
  57. user_id = ""
  58. sync_flag = str(uuid.uuid4())
  59. try:
  60. hashed_password = pwd_context.hash(password)
  61. db_user = UserModel(username=username, hashed_password=hashed_password, email=email, sync_flag=sync_flag)
  62. # pwd = db_user.encrypted_password(app_password)
  63. # db_user.password = pwd
  64. db_user.roles = [db.query(RoleModel).filter(RoleModel.role_type == 2).first()]
  65. db_user.groups = [db.query(GroupModel).filter(GroupModel.group_type == 2).first()]
  66. db.add(db_user)
  67. db.commit()
  68. db.refresh(db_user)
  69. '''
  70. user_id = db_user.id
  71. for k, v in register_dict.items():
  72. await UserAppDao(db).update_and_insert_data(v.get("name"), pwd, v.get("email"), user_id, str(v.get("id")), k)
  73. '''
  74. except Exception as e:
  75. logger.error(e)
  76. db.rollback()
  77. return False
  78. return sync_flag
  79. async def update_user_token(db, user_id, token_dict):
  80. try:
  81. for k, v in token_dict.items():
  82. await UserAppDao(db).update_user_app_data({"user_id": user_id, "app_type": k},
  83. {"access_token": v, "token_at": datetime.now()})
  84. except Exception as e:
  85. logger.error(e)
  86. return False
  87. return True
  88. """
  89. async def update_user_info(db, user_id):
  90. app_register = AppRegisterDao(db).get_apps()
  91. register_dict = {}
  92. user = db.query(UserModel).filter(UserModel.id==user_id).first()
  93. for app in app_register:
  94. if app["id"] == RAGFLOW:
  95. register_dict[app['id']] = {"id": user.ragflow_id, "name": user.username, "email": f"{user.username}@example.com"}
  96. elif app["id"] == BISHENG:
  97. register_dict[app['id']] = {"id": user.bisheng_id, "name": user.username, "email": ""}
  98. elif app["id"] == DIFY:
  99. register_dict[app['id']] = {"id": "", "name": user.username, "email": ""}
  100. else:
  101. logger.error("未知注册应用---")
  102. continue
  103. try:
  104. for k, v in register_dict.items():
  105. await UserAppDao(db).update_and_insert_data(v.get("name"), user.password, v.get("email"), user_id,
  106. str(v.get("id")), k)
  107. except Exception as e:
  108. logger.error(e)
  109. # 存储用户信息
  110. # hashed_password = pwd_context.hash(user.password)
  111. # db_user = UserModel(username=user.username, hashed_password=hashed_password, email=user.email)
  112. # db_user.password = db_user.encrypted_password(user.password)
  113. # for k, v in register_dict.items():
  114. # setattr(db_user, k.replace("app", "id"), v)
  115. # db.add(db_user)
  116. # db.commit()
  117. # db.refresh(db_user)
  118. # is_sava = await save_register_user(db, user.username, user.password, user.email, register_dict)
  119. """
  120. class UserAppDao:
  121. def __init__(self, db: Session):
  122. self.db = db
  123. async def get_data_by_id(self, user_id: int, app_type: int) -> Type[UserAppModel] | None:
  124. session = self.db.query(UserAppModel).filter_by(user_id=user_id, app_type=app_type).first()
  125. return session
  126. async def update_user_app_data(self, query: dict, update_data: dict):
  127. logger.error("更新数据df update_app_data---------------------------")
  128. try:
  129. self.db.query(UserAppModel).filter_by(**query).update(update_data)
  130. self.db.commit()
  131. except Exception as e:
  132. logger.error(e)
  133. self.db.rollback()
  134. raise Exception("更新失败!")
  135. async def insert_user_app_data(self, username: str, password: str, email: str, user_id: int, app_id: str,
  136. app_type: int):
  137. logger.error("新增数据df insert_user_app_data---------------------------")
  138. new_session = UserAppModel(
  139. username=username,
  140. password=password,
  141. email=email,
  142. user_id=user_id,
  143. app_id=app_id,
  144. app_type=app_type,
  145. )
  146. self.db.add(new_session)
  147. self.db.commit()
  148. self.db.refresh(new_session)
  149. return new_session
  150. async def update_and_insert_data(self, username: str, password: str, email: str, user_id: int, app_id: str,
  151. app_type: int):
  152. logger.error("更新或者添加数据 update_and_insert_token---------------------------")
  153. token_boj = await self.get_data_by_id(user_id, app_type)
  154. if token_boj:
  155. await self.update_user_app_data({"id": token_boj.id}, {"username": username,
  156. "password": password, "email": email,
  157. "updated_at": datetime.now(),
  158. })
  159. else:
  160. await self.insert_user_app_data(username, password, email, user_id, app_id, app_type)
  161. async def get_user_datas(self, user_id: int):
  162. return self.db.query(UserAppModel).filter_by(user_id=user_id).all()
  163. async def password_rsa(password):
  164. with open(os.path.join(APP_SERVICE_PATH, "pom/private_key.pem"), "rb") as key_file:
  165. private_key = serialization.load_pem_private_key(
  166. key_file.read(),
  167. password=None, # 如果私钥加密,请提供密码
  168. backend=default_backend()
  169. )
  170. # Base64 解码
  171. try:
  172. # 解密消息
  173. ciphertext = base64.b64decode(password)
  174. # 使用 PKCS#1 v1.5 填充解密
  175. plaintext = private_key.decrypt(
  176. ciphertext,
  177. padding.PKCS1v15() # 改为 PKCS#1 v1.5 填充
  178. )
  179. return plaintext.decode()
  180. except Exception as e:
  181. print(e)
  182. return ""