import secrets import bcrypt import logging from datetime import datetime, timezone from fastapi import APIRouter, Depends, HTTPException, status from pydantic import BaseModel from sqlalchemy.orm import Session from database import get_db from models.user import User from models.recovery_code import RecoveryCode from schemas.user import UserOut from schemas.auth import TokenResponse from routers.deps import get_current_user, make_token router = APIRouter() _logger = logging.getLogger(__name__) CODES_PER_BATCH = 5 def _generate_code() -> str: """Return a human-readable code like XENIA-A3K9-MW2F.""" part = lambda: secrets.token_hex(2).upper() return f"XENIA-{part()}-{part()}" def _hash_code(plain: str) -> str: return bcrypt.hashpw(plain.encode(), bcrypt.gensalt()).decode() def _verify_code(plain: str, hashed: str) -> bool: return bcrypt.checkpw(plain.encode(), hashed.encode()) def try_recovery_code(plain: str, user: "User", db: "Session") -> bool: """Try to consume a recovery code for user. Returns True and burns the code if matched.""" unused = db.query(RecoveryCode).filter( RecoveryCode.user_id == user.id, RecoveryCode.used_at.is_(None), ).all() matched = next((rc for rc in unused if _verify_code(plain, rc.code_hash)), None) if not matched: return False matched.used_at = datetime.now(timezone.utc) db.commit() _logger.warning("RECOVERY CODE USED for user_id=%d username=%s code_id=%d", user.id, user.username, matched.id) return True # ─── Public: use a recovery code to log in (kept for direct API use) ───────── class UseRecoveryCodeRequest(BaseModel): username: str code: str @router.post("/use", response_model=TokenResponse) def use_recovery_code(body: UseRecoveryCodeRequest, db: Session = Depends(get_db)): user = db.query(User).filter( User.username == body.username, User.is_active == True, ).first() if not user or not try_recovery_code(body.code, user, db): raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Invalid credentials") token = make_token(user) return TokenResponse(access_token=token, user=UserOut.model_validate(user)) # ─── Authenticated: generate a new batch (burns all existing unused codes) ─── class RecoveryCodesGenerated(BaseModel): codes: list[str] remaining_after: int @router.post("/generate", response_model=RecoveryCodesGenerated) def generate_recovery_codes( db: Session = Depends(get_db), current_user: User = Depends(get_current_user), ): # Burn all existing unused codes for this user db.query(RecoveryCode).filter( RecoveryCode.user_id == current_user.id, RecoveryCode.used_at.is_(None), ).delete(synchronize_session=False) db.flush() plain_codes = [_generate_code() for _ in range(CODES_PER_BATCH)] for plain in plain_codes: db.add(RecoveryCode(user_id=current_user.id, code_hash=_hash_code(plain))) db.commit() _logger.info("Recovery codes regenerated for user_id=%d", current_user.id) return RecoveryCodesGenerated(codes=plain_codes, remaining_after=CODES_PER_BATCH) # ─── Authenticated: check how many unused codes remain ─────────────────────── class RecoveryCodeStatus(BaseModel): unused_count: int @router.get("/status", response_model=RecoveryCodeStatus) def recovery_code_status( db: Session = Depends(get_db), current_user: User = Depends(get_current_user), ): count = db.query(RecoveryCode).filter( RecoveryCode.user_id == current_user.id, RecoveryCode.used_at == None, ).count() return RecoveryCodeStatus(unused_count=count)