Coverage for core/security.py: 100.00%

165 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-09-03 15:30 +0000

1import ast 

2import logging 

3from datetime import datetime, timedelta 

4from typing import Any 

5 

6import bcrypt 

7import redis 

8from fastapi import Depends, HTTPException, status 

9from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer 

10from jose import JWTError, jwt 

11from sqlalchemy import select, update 

12from sqlalchemy.ext.asyncio import AsyncSession 

13 

14from core.config import settings 

15from core.dependencies import get_db 

16from core.redis import get_redis 

17from models.user_sessions import UserSessions 

18from models.users import Users 

19from utils.custom_exception import ServerException 

20 

21logger = logging.getLogger("security") 

22 

23 

24async def hash_password(password: str) -> str: 

25 return bcrypt.hashpw(password.encode("utf-8"), bcrypt.gensalt()).decode("utf-8") 

26 

27 

28async def verify_password(plain_password: str, hashed_password: str) -> bool: 

29 return bcrypt.checkpw( 

30 plain_password.encode("utf-8"), 

31 hashed_password.encode("utf-8"), 

32 ) 

33 

34 

35async def create_access_token(data: dict[str, Any]) -> str: 

36 to_encode = data.copy() 

37 now = datetime.now().astimezone() 

38 expire = now + timedelta(minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES) 

39 to_encode.update({"exp": expire, "iat": now}) 

40 encoded_jwt = jwt.encode(to_encode, settings.SECRET_KEY, algorithm=settings.ALGORITHM) 

41 return encoded_jwt 

42 

43 

44async def create_password_reset_token(user_id: str, email: str) -> str: 

45 """Create password reset token""" 

46 try: 

47 now = datetime.now().astimezone() 

48 expire = now + timedelta(minutes=settings.PASSWORD_RESET_TOKEN_EXPIRE_MINUTES) 

49 

50 payload = { 

51 "sub": user_id, 

52 "email": email, 

53 "token_type": "password_reset", 

54 "force_change_password": True, 

55 "iat": now, 

56 "exp": expire, 

57 } 

58 

59 token = jwt.encode(payload, settings.SECRET_KEY, algorithm=settings.ALGORITHM) 

60 

61 return token 

62 

63 except Exception as e: 

64 raise ServerException(f"Failed to create password reset token: {str(e)}") 

65 

66 

67async def create_csrf_token(session_id: str) -> str: 

68 """Create CSRF token bound to a session""" 

69 try: 

70 now = datetime.now().astimezone() 

71 expire = now + timedelta(minutes=settings.CSRF_TOKEN_EXPIRE_MINUTES) 

72 

73 payload = { 

74 "sid": session_id, 

75 "token_type": "csrf", 

76 "iat": now, 

77 "exp": expire, 

78 } 

79 

80 return jwt.encode(payload, settings.SECRET_KEY, algorithm=settings.ALGORITHM) 

81 

82 except Exception as e: 

83 raise ServerException(f"Failed to create CSRF token: {str(e)}") 

84 

85 

86async def create_email_verification_token(user_id: str, email: str, token_type: str) -> str: 

87 """Create email verification token""" 

88 try: 

89 now = datetime.now().astimezone() 

90 expire = now + timedelta(minutes=settings.EMAIL_VERIFICATION_TOKEN_EXPIRE_MINUTES) 

91 

92 payload = { 

93 "sub": user_id, 

94 "email": email, 

95 "token_type": "email_verification", 

96 "verification_type": token_type, # 'registration' or 'email_change' 

97 "iat": now, 

98 "exp": expire, 

99 } 

100 

101 token = jwt.encode(payload, settings.SECRET_KEY, algorithm=settings.ALGORITHM) 

102 

103 return token 

104 

105 except Exception as e: 

106 raise ServerException(f"Failed to create email verification token: {str(e)}") 

107 

108 

109async def get_token( 

110 credentials: HTTPAuthorizationCredentials | None = Depends(HTTPBearer(auto_error=False)), 

111) -> str: 

112 if not credentials: 

113 raise HTTPException( 

114 status_code=status.HTTP_401_UNAUTHORIZED, 

115 detail="Invalid or expired token", 

116 headers={"WWW-Authenticate": "Bearer"}, 

117 ) 

118 return credentials.credentials 

119 

120 

121async def verify_session(sid: str, token: str, redis_client) -> dict[str, Any]: 

122 try: 

123 redis_key = f"session:{sid}" 

124 raw = await redis_client.get(redis_key) 

125 if not raw: 

126 raise ValueError("Invalid or expired session") 

127 try: 

128 session_data = ast.literal_eval(raw) 

129 except ValueError, SyntaxError: 

130 logger.error(f"Invalid session data: {raw}") 

131 raise ValueError("Invalid session data") 

132 

133 if session_data.get("access_token") and session_data.get("access_token") != token: 

134 logger.error(f"Token mismatch: {session_data.get('access_token')} != {token}") 

135 raise JWTError("Token mismatch") 

136 return session_data 

137 except JWTError as e: 

138 logger.warning(f"JWT validation failed: {type(e).__name__}: {str(e)}") 

139 raise HTTPException( 

140 status_code=status.HTTP_401_UNAUTHORIZED, 

141 detail="Invalid or expired token", 

142 headers={"WWW-Authenticate": "Bearer"}, 

143 ) 

144 except Exception as e: 

145 logger.error(f"Failed to verify session: {e}") 

146 raise HTTPException( 

147 status_code=status.HTTP_401_UNAUTHORIZED, 

148 detail="Invalid or expired session", 

149 headers={"WWW-Authenticate": "Bearer"}, 

150 ) 

151 

152 

153async def verify_token( 

154 token: str = Depends(get_token), 

155 redis_client=Depends(get_redis), 

156 db: AsyncSession = Depends(get_db), 

157) -> dict[str, Any]: 

158 try: 

159 payload = jwt.decode(token, settings.SECRET_KEY, algorithms=[settings.ALGORITHM]) 

160 sid = payload.get("sid") 

161 if not sid: 

162 raise ValueError("Missing session ID") 

163 await verify_session(sid, token, redis_client) 

164 

165 user = await db.execute(select(Users.status).where(Users.id == payload.get("sub"))) 

166 user_status = user.scalar_one_or_none() 

167 if user_status is not None and not user_status: 

168 raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Account is disabled") 

169 

170 return payload 

171 

172 except JWTError as e: 

173 logger.warning(f"JWT validation failed: {type(e).__name__}: {str(e)}") 

174 raise HTTPException( 

175 status_code=status.HTTP_401_UNAUTHORIZED, 

176 detail="Invalid or expired token", 

177 headers={"WWW-Authenticate": "Bearer"}, 

178 ) 

179 except ValueError as e: 

180 logger.warning(f"Token validation error: {str(e)}") 

181 raise HTTPException( 

182 status_code=status.HTTP_401_UNAUTHORIZED, 

183 detail="Invalid or expired token", 

184 headers={"WWW-Authenticate": "Bearer"}, 

185 ) 

186 

187 

188async def verify_password_reset_token(token: str = Depends(get_token)) -> dict[str, Any]: 

189 """Verify password reset token""" 

190 try: 

191 payload = jwt.decode(token, settings.SECRET_KEY, algorithms=[settings.ALGORITHM]) 

192 

193 if payload.get("token_type") != "password_reset": 

194 raise ValueError("Invalid token type") 

195 

196 # Check if force change password 

197 if not payload.get("force_change_password"): 

198 raise ValueError("Token not authorized for password reset") 

199 

200 if payload.get("exp") < datetime.now().astimezone().timestamp(): 

201 raise ValueError("Token expired") 

202 

203 user_id = payload.get("sub") 

204 email = payload.get("email") 

205 

206 if not user_id or not email: 

207 raise ValueError("Invalid token payload") 

208 

209 return { 

210 "token": token, 

211 "sub": user_id, 

212 "email": email, 

213 "exp": payload.get("exp"), 

214 "iat": payload.get("iat"), 

215 } 

216 

217 except JWTError as e: 

218 logger.warning(f"JWT validation failed: {type(e).__name__}: {str(e)}") 

219 raise HTTPException( 

220 status_code=status.HTTP_401_UNAUTHORIZED, 

221 detail="Invalid or expired token", 

222 headers={"WWW-Authenticate": "Bearer"}, 

223 ) 

224 except ValueError as e: 

225 logger.error(f"Failed to verify password reset token: {str(e)}") 

226 raise HTTPException( 

227 status_code=status.HTTP_401_UNAUTHORIZED, 

228 detail="Invalid or expired token", 

229 headers={"WWW-Authenticate": "Bearer"}, 

230 ) 

231 

232 

233async def verify_email_verification_token(token: str = Depends(get_token)) -> dict[str, Any]: 

234 """Verify email verification token""" 

235 try: 

236 payload = jwt.decode(token, settings.SECRET_KEY, algorithms=[settings.ALGORITHM]) 

237 

238 if payload.get("token_type") != "email_verification": 

239 raise ValueError("Invalid token type") 

240 

241 verification_type = payload.get("verification_type") 

242 if verification_type not in ["registration", "email_change"]: 

243 raise ValueError("Invalid verification type") 

244 

245 if payload.get("exp") < datetime.now().astimezone().timestamp(): 

246 raise ValueError("Token expired") 

247 

248 user_id = payload.get("sub") 

249 email = payload.get("email") 

250 

251 if not user_id or not email: 

252 raise ValueError("Invalid token payload") 

253 

254 return { 

255 "token": token, 

256 "sub": user_id, 

257 "email": email, 

258 "verification_type": verification_type, 

259 "exp": payload.get("exp"), 

260 "iat": payload.get("iat"), 

261 } 

262 

263 except JWTError as e: 

264 logger.warning(f"JWT validation failed: {type(e).__name__}: {str(e)}") 

265 raise HTTPException( 

266 status_code=status.HTTP_401_UNAUTHORIZED, 

267 detail="Invalid or expired token", 

268 headers={"WWW-Authenticate": "Bearer"}, 

269 ) 

270 except ValueError as e: 

271 logger.error(f"Failed to verify email verification token: {str(e)}") 

272 raise HTTPException( 

273 status_code=status.HTTP_401_UNAUTHORIZED, 

274 detail="Invalid or expired token", 

275 headers={"WWW-Authenticate": "Bearer"}, 

276 ) 

277 

278 

279async def extend_session_ttl(redis_client, session_id: str, session_data: dict[str, Any]) -> None: 

280 """Extend session TTL and update last activity time""" 

281 try: 

282 # Update last activity time using system timezone with timezone info 

283 session_data["last_activity"] = datetime.now().astimezone().isoformat() 

284 

285 # Reset TTL, start from current time 

286 ttl = settings.SESSION_EXPIRE_MINUTES * 60 

287 await redis_client.setex(f"session:{session_id}", ttl, str(session_data)) 

288 

289 except Exception as e: 

290 logger.error(f"Failed to extend session TTL: {e}") 

291 

292 

293async def clear_user_all_sessions( 

294 db: AsyncSession, 

295 redis_client: redis.Redis, 

296 user_id: str, 

297 exclude_session_id: str | None = None, 

298) -> bool: 

299 """Logout user from all devices, optionally keeping one active session.""" 

300 try: 

301 sessions_query = select(UserSessions.id).where( 

302 UserSessions.user_id == user_id, 

303 UserSessions.is_active.is_(True), 

304 ) 

305 if exclude_session_id: 

306 sessions_query = sessions_query.where(UserSessions.id != exclude_session_id) 

307 

308 result = await db.execute(sessions_query) 

309 session_ids = list(result.scalars().all()) 

310 

311 if not session_ids: 

312 return True 

313 

314 await db.execute( 

315 update(UserSessions).where(UserSessions.id.in_(session_ids)).values(is_active=False) 

316 ) 

317 await db.commit() 

318 

319 redis_keys = [] 

320 for sid in session_ids: 

321 redis_keys.append(f"session:{sid}") 

322 redis_keys.append(f"csrf:{sid}") 

323 await redis_client.delete(*redis_keys) 

324 

325 return True 

326 except Exception as e: 

327 raise ServerException(f"Failed to logout all devices: {e}")