@@ -159,7 +159,7 @@ def verify_mfa_code(self, secret: str, otp_code: str) -> bool:
159159 return totp .verify (otp_code )
160160
161161 def generate_qr_code_svg (self , uri : str ) -> str :
162- """Generates an SVG string for a QR code from a URI."""
162+ """Generates a base64-encoded PNG for a QR code from a URI."""
163163 img = qrcode .make (uri )
164164 buffer = io .BytesIO ()
165165 img .save (buffer , format = "PNG" )
@@ -179,16 +179,21 @@ def __init__(self, redis_client: redis.Redis) -> None:
179179
180180 async def is_allowed (
181181 self , key : str , limit : int , window : int , identifier : str = "default"
182- ) -> tuple [ bool , Dict [ str , Any ]] :
182+ ) -> tuple :
183183 """
184- Check if request is allowed based on rate limit
184+ Check if request is allowed based on rate limit.
185185 Returns (is_allowed, info_dict)
186186 """
187187 current_time = int (time .time ())
188- pipe = self .redis .pipeline ()
189188 window_start = current_time - window
190- await pipe .zremrangebyscore (key , 0 , window_start )
191- current_requests = await pipe .zcard (key )
189+
190+ # Use pipeline correctly: execute cleanup and count atomically
191+ pipe = self .redis .pipeline ()
192+ pipe .zremrangebyscore (key , 0 , window_start )
193+ pipe .zcard (key )
194+ results = await pipe .execute ()
195+ current_requests = results [1 ]
196+
192197 if current_requests >= limit :
193198 oldest_request = await self .redis .zrange (key , 0 , 0 , withscores = True )
194199 if oldest_request :
@@ -205,9 +210,12 @@ async def is_allowed(
205210 "retry_after" : time_until_reset ,
206211 },
207212 )
208- await pipe .zadd (key , {f"{ current_time } :{ identifier } " : current_time })
209- await pipe .expire (key , window )
210- await pipe .execute ()
213+
214+ pipe2 = self .redis .pipeline ()
215+ pipe2 .zadd (key , {f"{ current_time } :{ identifier } " : current_time })
216+ pipe2 .expire (key , window )
217+ await pipe2 .execute ()
218+
211219 remaining = limit - current_requests - 1
212220 return (
213221 True ,
@@ -306,7 +314,9 @@ async def get_current_user_from_token(
306314 )
307315 user = (
308316 db .query (User )
309- .filter (User .id == int (user_id ), User .is_active , User .is_deleted == False )
317+ .filter (
318+ User .id == int (user_id ), User .is_active == True , User .is_deleted == False
319+ )
310320 .first ()
311321 )
312322 if not user :
@@ -324,7 +334,8 @@ async def get_current_user_from_token(
324334 raise HTTPException (
325335 status_code = status .HTTP_403_FORBIDDEN , detail = "MFA code required"
326336 )
327- if user .mfa_secret == False :
337+ # Bug fix: was `user.mfa_secret == False` which compares object to bool
338+ if not user .mfa_secret :
328339 raise HTTPException (
329340 status_code = status .HTTP_500_INTERNAL_SERVER_ERROR ,
330341 detail = "MFA enabled but secret not found" ,
@@ -348,7 +359,7 @@ async def get_current_user_from_api_key(
348359 db .query (ApiKey )
349360 .filter (
350361 ApiKey .key_hash == key_hash ,
351- ApiKey .is_active ,
362+ ApiKey .is_active == True ,
352363 ApiKey .is_deleted == False ,
353364 )
354365 .first ()
@@ -367,7 +378,7 @@ async def get_current_user_from_api_key(
367378 db .query (User )
368379 .filter (
369380 User .id == api_key_obj .user_id ,
370- User .is_active ,
381+ User .is_active == True ,
371382 User .is_deleted == False ,
372383 )
373384 .first ()
@@ -400,17 +411,26 @@ def decorator(func):
400411 @wraps (func )
401412 async def wrapper (* args , ** kwargs ):
402413 current_user = None
414+ # Check positional args
403415 for arg in args :
404416 if isinstance (arg , User ):
405417 current_user = arg
406418 break
419+ # Check keyword args (FastAPI passes dependencies as kwargs)
420+ if not current_user :
421+ for v in kwargs .values ():
422+ if isinstance (v , User ):
423+ current_user = v
424+ break
407425 if not current_user :
408426 raise HTTPException (
409427 status_code = status .HTTP_401_UNAUTHORIZED ,
410428 detail = "Authentication required" ,
411429 )
412- user_permissions = set (
413- [p .permission_name for p in current_user .role .permissions ]
430+ user_permissions = (
431+ set ([p .permission_name for p in current_user .role .permissions ])
432+ if current_user .role
433+ else set ()
414434 )
415435 if not all ((perm in user_permissions for perm in required_permissions )):
416436 raise HTTPException (
@@ -426,7 +446,8 @@ async def wrapper(*args, **kwargs):
426446
427447def require_admin (current_user : User = Depends (get_current_user )) -> Any :
428448 """Dependency to require admin role"""
429- if current_user .role .value != "admin" :
449+ role_name = current_user .role .role_name if current_user .role else ""
450+ if role_name != "admin" :
430451 raise HTTPException (
431452 status_code = status .HTTP_403_FORBIDDEN , detail = "Admin access required"
432453 )
@@ -546,10 +567,10 @@ async def create_user_session(
546567 request : Request ,
547568 max_concurrent_sessions : int = settings .security .max_concurrent_sessions ,
548569) -> UserSession :
549- """Create a new user session, handling concurrent sessions and IP/User-Agent binding """
570+ """Create a new user session, handling concurrent sessions"""
550571 active_sessions = (
551572 db .query (UserSession )
552- .filter (UserSession .user_id == user .id , UserSession .is_active )
573+ .filter (UserSession .user_id == user .id , UserSession .is_active == True )
553574 .order_by (UserSession .last_activity .asc ())
554575 .all ()
555576 )
@@ -572,7 +593,7 @@ async def create_user_session(
572593 UserSession .user_id == user .id ,
573594 UserSession .ip_address == request .client .host ,
574595 UserSession .user_agent == request .headers .get ("user-agent" ),
575- UserSession .is_active ,
596+ UserSession .is_active == True ,
576597 )
577598 .first ()
578599 )
@@ -627,11 +648,14 @@ def authenticate_user(db: Session, username: str, password: str) -> Optional[Use
627648
628649def create_tokens (user : User ) -> Token :
629650 """Create access and refresh tokens for a user"""
630- user_permissions = [p .permission_name for p in user .role .permissions ]
651+ user_permissions = []
652+ if user .role and user .role .permissions :
653+ user_permissions = [p .permission_name for p in user .role .permissions ]
654+ role_name = user .role .role_name if user .role else "user"
631655 access_token = security_manager .create_access_token (
632656 user_id = user .id ,
633657 username = user .username ,
634- role = user . role . value ,
658+ role = role_name ,
635659 permissions = user_permissions ,
636660 )
637661 refresh_token = security_manager .create_refresh_token (
@@ -655,7 +679,9 @@ def refresh_access_token(db: Session, refresh_token: str) -> Optional[Token]:
655679 return None
656680 user = (
657681 db .query (User )
658- .filter (User .id == int (user_id ), User .is_active , User .is_deleted == False )
682+ .filter (
683+ User .id == int (user_id ), User .is_active == True , User .is_deleted == False
684+ )
659685 .first ()
660686 )
661687 if not user :
@@ -665,7 +691,7 @@ def refresh_access_token(db: Session, refresh_token: str) -> Optional[Token]:
665691 .filter (
666692 UserSession .user_id == user .id ,
667693 UserSession .refresh_token == refresh_token ,
668- UserSession .is_active ,
694+ UserSession .is_active == True ,
669695 )
670696 .first ()
671697 )
@@ -683,6 +709,3 @@ def refresh_access_token(db: Session, refresh_token: str) -> Optional[Token]:
683709from fastapi import APIRouter
684710
685711router = APIRouter ()
686-
687- # Auth endpoints would be defined here if needed
688- # Currently auth functionality is provided through dependencies
0 commit comments