Investment Plans workspace
Open raw ↗
"""User authentication and management for HAD Digital MVP.

Password hashing uses hashlib.scrypt with per-user salts.
Implements 5-failure lockout mechanism.
Roles: oncologist, had_nurse, community_nurse, gp, pharmacist, patient, caregiver, admin
"""

import hashlib
import os
import secrets
import time
from datetime import datetime, timezone, timedelta
from database import get_db

VALID_ROLES = (
    "oncologist", "had_nurse", "community_nurse", "gp",
    "pharmacist", "patient", "caregiver", "admin"
)

MAX_FAILED_ATTEMPTS = 5
LOCKOUT_DURATION_MINUTES = 30


def _generate_salt() -> str:
    """Generate a cryptographically secure salt."""
    return secrets.token_hex(32)


def _hash_password(password: str, salt: str) -> str:
    """Hash password using scrypt with the given salt."""
    salt_bytes = bytes.fromhex(salt)
    hash_bytes = hashlib.scrypt(
        password.encode("utf-8"),
        salt=salt_bytes,
        n=2**14,
        r=8,
        p=1,
        dklen=64,
    )
    return hash_bytes.hex()


def create_user(
    username: str,
    password: str,
    role: str,
    display_name: str,
    email: str = None,
    patient_id: int = None,
) -> int:
    """Create a new user. Returns user ID."""
    if role not in VALID_ROLES:
        raise ValueError(f"Invalid role: {role}. Must be one of {VALID_ROLES}")

    db = get_db()
    salt = _generate_salt()
    password_hash = _hash_password(password, salt)

    cursor = db.execute(
        """INSERT INTO users (username, password_hash, salt, role, display_name, email, patient_id)
           VALUES (?, ?, ?, ?, ?, ?, ?)""",
        (username, password_hash, salt, role, display_name, email, patient_id),
    )
    db.commit()
    return cursor.lastrowid


def authenticate(username: str, password: str, ip_address: str = None) -> dict | None:
    """Authenticate a user. Returns user dict on success, None on failure.
    
    Implements 5-fail lockout: after 5 failed attempts, account is locked for 30 minutes.
    """
    db = get_db()
    user = db.fetchone("SELECT * FROM users WHERE username = ?", (username,))

    if user is None:
        # Log failed attempt for non-existent user (no lockout needed)
        return None

    # Check lockout
    if user["locked_until"]:
        lock_until = datetime.fromisoformat(user["locked_until"])
        if datetime.now(timezone.utc) < lock_until.replace(tzinfo=timezone.utc):
            return None  # Still locked
        else:
            # Lock expired, reset
            db.execute(
                "UPDATE users SET failed_attempts = 0, locked_until = NULL WHERE id = ?",
                (user["id"],),
            )
            db.commit()

    # Check if active
    if not user["active"]:
        return None

    # Verify password
    password_hash = _hash_password(password, user["salt"])
    if password_hash != user["password_hash"]:
        # Increment failed attempts
        new_attempts = user["failed_attempts"] + 1
        locked_until = None
        if new_attempts >= MAX_FAILED_ATTEMPTS:
            locked_until = (
                datetime.now(timezone.utc) + timedelta(minutes=LOCKOUT_DURATION_MINUTES)
            ).strftime("%Y-%m-%d %H:%M:%S")

        db.execute(
            "UPDATE users SET failed_attempts = ?, locked_until = ? WHERE id = ?",
            (new_attempts, locked_until, user["id"]),
        )
        db.commit()
        return None

    # Success - reset failed attempts
    db.execute(
        "UPDATE users SET failed_attempts = 0, locked_until = NULL WHERE id = ?",
        (user["id"],),
    )
    db.commit()

    return {
        "id": user["id"],
        "username": user["username"],
        "role": user["role"],
        "display_name": user["display_name"],
        "email": user["email"],
        "patient_id": user["patient_id"],
    }


def get_user(user_id: int) -> dict | None:
    """Get user by ID."""
    db = get_db()
    user = db.fetchone(
        "SELECT id, username, role, display_name, email, patient_id, active FROM users WHERE id = ?",
        (user_id,),
    )
    if user is None:
        return None
    return dict(user)


def get_user_by_username(username: str) -> dict | None:
    """Get user by username."""
    db = get_db()
    user = db.fetchone(
        "SELECT id, username, role, display_name, email, patient_id, active FROM users WHERE username = ?",
        (username,),
    )
    if user is None:
        return None
    return dict(user)


def list_users() -> list[dict]:
    """List all active users."""
    db = get_db()
    rows = db.fetchall(
        "SELECT id, username, role, display_name, email, patient_id, active FROM users WHERE active = 1"
    )
    return [dict(r) for r in rows]


def list_users_by_role(role: str) -> list[dict]:
    """List all active users with a specific role."""
    db = get_db()
    rows = db.fetchall(
        "SELECT id, username, role, display_name, email, patient_id, active FROM users WHERE role = ? AND active = 1",
        (role,),
    )
    return [dict(r) for r in rows]


def update_user(user_id: int, **kwargs) -> bool:
    """Update user fields. Returns True on success."""
    allowed = {"display_name", "email", "role", "active"}
    updates = {k: v for k, v in kwargs.items() if k in allowed}
    if not updates:
        return False

    db = get_db()
    set_clause = ", ".join(f"{k} = ?" for k in updates)
    values = list(updates.values()) + [user_id]
    db.execute(
        f"UPDATE users SET {set_clause}, updated_at = datetime('now') WHERE id = ?",
        tuple(values),
    )
    db.commit()
    return True


def change_password(user_id: int, new_password: str) -> bool:
    """Change a user's password."""
    db = get_db()
    salt = _generate_salt()
    password_hash = _hash_password(new_password, salt)
    db.execute(
        "UPDATE users SET password_hash = ?, salt = ?, updated_at = datetime('now') WHERE id = ?",
        (password_hash, salt, user_id),
    )
    db.commit()
    return True


def get_patient_users(patient_id: int) -> list[dict]:
    """Get all users associated with a patient (patient, caregiver, etc.)."""
    db = get_db()
    rows = db.fetchall(
        "SELECT id, username, role, display_name FROM users WHERE patient_id = ? AND active = 1",
        (patient_id,),
    )
    return [dict(r) for r in rows]