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
« 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
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
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
21logger = logging.getLogger("security")
24async def hash_password(password: str) -> str:
25 return bcrypt.hashpw(password.encode("utf-8"), bcrypt.gensalt()).decode("utf-8")
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 )
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
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)
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 }
59 token = jwt.encode(payload, settings.SECRET_KEY, algorithm=settings.ALGORITHM)
61 return token
63 except Exception as e:
64 raise ServerException(f"Failed to create password reset token: {str(e)}")
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)
73 payload = {
74 "sid": session_id,
75 "token_type": "csrf",
76 "iat": now,
77 "exp": expire,
78 }
80 return jwt.encode(payload, settings.SECRET_KEY, algorithm=settings.ALGORITHM)
82 except Exception as e:
83 raise ServerException(f"Failed to create CSRF token: {str(e)}")
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)
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 }
101 token = jwt.encode(payload, settings.SECRET_KEY, algorithm=settings.ALGORITHM)
103 return token
105 except Exception as e:
106 raise ServerException(f"Failed to create email verification token: {str(e)}")
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
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")
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 )
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)
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")
170 return payload
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 )
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])
193 if payload.get("token_type") != "password_reset":
194 raise ValueError("Invalid token type")
196 # Check if force change password
197 if not payload.get("force_change_password"):
198 raise ValueError("Token not authorized for password reset")
200 if payload.get("exp") < datetime.now().astimezone().timestamp():
201 raise ValueError("Token expired")
203 user_id = payload.get("sub")
204 email = payload.get("email")
206 if not user_id or not email:
207 raise ValueError("Invalid token payload")
209 return {
210 "token": token,
211 "sub": user_id,
212 "email": email,
213 "exp": payload.get("exp"),
214 "iat": payload.get("iat"),
215 }
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 )
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])
238 if payload.get("token_type") != "email_verification":
239 raise ValueError("Invalid token type")
241 verification_type = payload.get("verification_type")
242 if verification_type not in ["registration", "email_change"]:
243 raise ValueError("Invalid verification type")
245 if payload.get("exp") < datetime.now().astimezone().timestamp():
246 raise ValueError("Token expired")
248 user_id = payload.get("sub")
249 email = payload.get("email")
251 if not user_id or not email:
252 raise ValueError("Invalid token payload")
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 }
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 )
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()
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))
289 except Exception as e:
290 logger.error(f"Failed to extend session TTL: {e}")
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)
308 result = await db.execute(sessions_query)
309 session_ids = list(result.scalars().all())
311 if not session_ids:
312 return True
314 await db.execute(
315 update(UserSessions).where(UserSessions.id.in_(session_ids)).values(is_active=False)
316 )
317 await db.commit()
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)
325 return True
326 except Exception as e:
327 raise ServerException(f"Failed to logout all devices: {e}")