Spaces:
Paused
Paused
| import uuid | |
| import asyncio | |
| import logging | |
| import json | |
| import secrets | |
| import re | |
| from typing import Optional | |
| from datetime import datetime, timedelta | |
| from fastapi import Depends, Request, HTTPException | |
| from fastapi_users import BaseUserManager, UUIDIDMixin | |
| from sqlalchemy.ext.asyncio import AsyncSession | |
| from sqlalchemy import select | |
| from app.auth.models import User, RefreshToken | |
| from app.auth.db import get_user_db | |
| from app.database import get_db, AsyncSessionLocal | |
| from app.config import get_settings | |
| from app.utils.security import mask_email | |
| from app.utils.encryption import encrypt_backup_codes, decrypt_backup_codes | |
| logger = logging.getLogger(__name__) | |
| settings = get_settings() | |
| class UserManager(UUIDIDMixin, BaseUserManager[User, uuid.UUID]): | |
| reset_password_token_secret = settings.reset_password_token_secret or settings.secret_key | |
| verification_token_secret = settings.verification_token_secret or settings.secret_key | |
| async def on_after_register(self, user: User, request: Optional[Request] = None): | |
| logger.info(f"User {user.id} ({mask_email(user.email)}) has registered.") | |
| try: | |
| from app.utils.email_service import send_welcome_email | |
| asyncio.create_task(send_welcome_email( | |
| to_email=user.email, | |
| full_name=user.full_name, | |
| )) | |
| except Exception as e: | |
| logger.warning(f"Failed to send welcome email to {mask_email(user.email)}: {e}") | |
| async def on_after_forgot_password( | |
| self, user: User, token: str, request: Optional[Request] = None | |
| ): | |
| logger.info(f"User {user.id} has forgot their password. Reset token generated.") | |
| async def on_after_request_verify( | |
| self, user: User, token: str, request: Optional[Request] = None | |
| ): | |
| logger.info(f"Verification requested for user {user.id}.") | |
| # ─── Password Validation ─── | |
| async def validate_password(self, password: str, user: Optional[User] = None) -> None: | |
| """Validate password strength using zxcvbn and check against HIBP.""" | |
| # Minimum length check | |
| if len(password) < 12: | |
| raise ValueError("La contraseña debe tener al menos 12 caracteres") | |
| # Check for common patterns - only flag long sequences (5+ chars) | |
| if re.search(r'(.)\1{2,}', password): # 3+ repeated characters | |
| raise ValueError("La contraseña no debe contener caracteres repetidos (3 o más)") | |
| # Check for sequential patterns (5+ chars) - e.g., abcde, 12345, etc. | |
| sequential_patterns = [ | |
| 'abcde', 'bcdef', 'cdefg', 'defgh', 'efghi', 'fghij', 'ghijk', 'hijkl', 'ijklm', 'jklmn', | |
| 'klmno', 'lmnop', 'mnopq', 'nopqr', 'opqrs', 'pqrst', 'qrstu', 'rstuv', 'stuvw', 'tuvwx', | |
| 'uvwxy', 'vwxy', 'wxyz', | |
| '01234', '12345', '23456', '34567', '45678', '56789' | |
| ] | |
| password_lower = password.lower() | |
| for seq in sequential_patterns: | |
| if seq in password_lower: | |
| raise ValueError("La contraseña no debe contener secuencias comunes (5+ caracteres)") | |
| # zxcvbn score check (minimum score 3 = good) | |
| try: | |
| from zxcvbn import zxcvbn | |
| result = zxcvbn(password) | |
| if result['score'] < 3: | |
| raise ValueError(f"Contraseña muy débil. Mejora: {'; '.join(result['feedback']['suggestions'])}") | |
| # Check against HIBP (k-anonymity) | |
| import httpx | |
| import hashlib | |
| sha1 = hashlib.sha1(password.encode()).hexdigest().upper() | |
| prefix, suffix = sha1[:5], sha1[5:] | |
| try: | |
| with httpx.Client(timeout=5.0) as client: | |
| resp = client.get(f"https://api.pwnedpasswords.com/range/{prefix}") | |
| if resp.status_code == 200: | |
| if suffix in resp.text: | |
| raise ValueError("Esta contraseña ha sido filtrada en brechas de seguridad conocidas. Usa otra.") | |
| except httpx.TimeoutException: | |
| logger.warning("HIBP check timeout, skipping") | |
| except Exception as e: | |
| logger.warning(f"HIBP check failed: {e}") | |
| except ImportError: | |
| # zxcvbn not installed, skip advanced checks | |
| logger.warning("zxcvbn not installed, skipping advanced password checks") | |
| pass | |
| async def authenticate(self, credentials): | |
| """Override authenticate to add account lockout and password validation.""" | |
| from fastapi_users.exceptions import InvalidPasswordException, UserNotExists, UserInactive | |
| # Get user by email | |
| try: | |
| user = await self.user_db.get_by_email(credentials.username) | |
| except Exception: | |
| user = None | |
| if not user: | |
| raise UserNotExists() | |
| if not user.is_active: | |
| raise UserInactive() | |
| # Check account lockout | |
| if user.locked_until and user.locked_until > datetime.utcnow(): | |
| remaining = int((user.locked_until - datetime.utcnow()).total_seconds() / 60) | |
| raise HTTPException( | |
| status_code=403, | |
| detail=f"Cuenta bloqueada temporalmente. Intente en {remaining} minutos." | |
| ) | |
| # Verify password | |
| is_valid, new_hash = self.password_helper.verify_and_update(credentials.password, user.hashed_password) | |
| if not is_valid: | |
| # Increment failed attempts | |
| user.failed_login_attempts += 1 | |
| if user.failed_login_attempts >= 5: | |
| user.locked_until = datetime.utcnow() + timedelta(minutes=15) | |
| logger.warning(f"Account locked for user {user.id} ({mask_email(user.email)}) after 5 failed attempts") | |
| await self.user_db.update(user) | |
| raise InvalidPasswordException() | |
| # Update hash if it was upgraded (e.g., bcrypt rounds increased) | |
| if new_hash: | |
| user.hashed_password = new_hash | |
| await self.user_db.update(user) | |
| # Successful login - reset failed attempts and lock | |
| if user.failed_login_attempts > 0 or user.locked_until: | |
| user.failed_login_attempts = 0 | |
| user.locked_until = None | |
| await self.user_db.update(user) | |
| # Verify email is verified | |
| if not user.is_verified: | |
| raise HTTPException( | |
| status_code=403, | |
| detail="Tu cuenta no ha sido verificada. Revisa tu email para verificar tu cuenta." | |
| ) | |
| return user | |
| # ─── Refresh Token Methods ─── | |
| async def create_refresh_token( | |
| self, | |
| user: User, | |
| request: Optional[Request] = None, | |
| db: Optional[AsyncSession] = None, | |
| ) -> str: | |
| """Crear nuevo refresh token y almacenar hash en BD.""" | |
| if db is None: | |
| async for session in get_db(): | |
| return await self._create_refresh_token_internal(user, request, session) | |
| return await self._create_refresh_token_internal(user, request, db) | |
| async def _create_refresh_token_internal( | |
| self, | |
| user: User, | |
| request: Optional[Request], | |
| db: AsyncSession, | |
| ) -> str: | |
| # Revocar tokens anteriores del usuario (rotación) | |
| await self.revoke_user_refresh_tokens(user.id, db) | |
| # Generar nuevo token | |
| raw_token = RefreshToken.generate_token() | |
| token_hash = RefreshToken.hash_token(raw_token) | |
| expires_at = datetime.utcnow() + timedelta(days=settings.refresh_token_expire_days) | |
| # Extraer info de request | |
| user_agent = request.headers.get("user-agent") if request else None | |
| ip_address = request.client.host if request and request.client else None | |
| refresh_token = RefreshToken( | |
| user_id=user.id, | |
| token_hash=token_hash, | |
| expires_at=expires_at, | |
| user_agent=user_agent, | |
| ip_address=ip_address, | |
| ) | |
| db.add(refresh_token) | |
| await db.commit() | |
| logger.info(f"Refresh token created for user {user.id}") | |
| return raw_token | |
| async def verify_refresh_token( | |
| self, | |
| token: str, | |
| db: Optional[AsyncSession] = None, | |
| ) -> Optional[User]: | |
| """Verificar refresh token y retornar usuario si válido.""" | |
| if db is None: | |
| async for session in get_db(): | |
| return await self._verify_refresh_token_internal(token, session) | |
| return await self._verify_refresh_token_internal(token, db) | |
| async def _verify_refresh_token_internal( | |
| self, | |
| token: str, | |
| db: AsyncSession, | |
| ) -> Optional[User]: | |
| token_hash = RefreshToken.hash_token(token) | |
| result = await db.execute( | |
| select(RefreshToken).where( | |
| RefreshToken.token_hash == token_hash, | |
| RefreshToken.revoked == False, | |
| RefreshToken.expires_at > datetime.utcnow(), | |
| ) | |
| ) | |
| refresh_token = result.scalar_one_or_none() | |
| if not refresh_token: | |
| return None | |
| # Obtener usuario | |
| user = await self.user_db.get(refresh_token.user_id) | |
| if not user or not user.is_active: | |
| return None | |
| return user | |
| async def revoke_refresh_token( | |
| self, | |
| token: str, | |
| db: Optional[AsyncSession] = None, | |
| ) -> bool: | |
| """Revocar un refresh token específico.""" | |
| if db is None: | |
| async for session in get_db(): | |
| return await self._revoke_refresh_token_internal(token, session) | |
| return await self._revoke_refresh_token_internal(token, db) | |
| async def _revoke_refresh_token_internal( | |
| self, | |
| token: str, | |
| db: AsyncSession, | |
| ) -> bool: | |
| token_hash = RefreshToken.hash_token(token) | |
| result = await db.execute( | |
| select(RefreshToken).where(RefreshToken.token_hash == token_hash) | |
| ) | |
| refresh_token = result.scalar_one_or_none() | |
| if not refresh_token: | |
| return False | |
| refresh_token.revoked = True | |
| await db.commit() | |
| return True | |
| async def revoke_user_refresh_tokens( | |
| self, | |
| user_id: uuid.UUID, | |
| db: Optional[AsyncSession] = None, | |
| ) -> int: | |
| """Revocar TODOS los refresh tokens de un usuario (logout everywhere).""" | |
| if db is None: | |
| async for session in get_db(): | |
| return await self._revoke_user_refresh_tokens_internal(user_id, session) | |
| return await self._revoke_user_refresh_tokens_internal(user_id, db) | |
| async def _revoke_user_refresh_tokens_internal( | |
| self, | |
| user_id: uuid.UUID, | |
| db: AsyncSession, | |
| ) -> int: | |
| result = await db.execute( | |
| select(RefreshToken).where( | |
| RefreshToken.user_id == user_id, | |
| RefreshToken.revoked == False, | |
| ) | |
| ) | |
| tokens = result.scalars().all() | |
| for token in tokens: | |
| token.revoked = True | |
| await db.commit() | |
| return len(tokens) | |
| async def create_user(self, user_create, safe: bool = False, request: Request | None = None): | |
| """Override to validate password strength on registration.""" | |
| # Validate password strength | |
| self.validate_password(user_create.password) | |
| # Call parent create_user | |
| return await super().create_user(user_create, safe, request) | |
| async def update_user(self, user_update, user, safe: bool = False, request: Request | None = None): | |
| """Override to validate password on update.""" | |
| if hasattr(user_update, 'password') and user_update.password: | |
| self.validate_password(user_update.password) | |
| return await super().update_user(user_update, user, safe, request) | |
| async def get_user_manager(user_db=Depends(get_user_db)): | |
| yield UserManager(user_db) |