"""
Bank-Grade Security Engine for FreeExile.
Implements Hardware/Device Fingerprint Binding, High-Entropy Remember-Me Tokens,
Progressive Exponential Lockout Delays, and Immutable Audit Logging.
"""

from __future__ import annotations
import enum
import hashlib
import hmac
import secrets
import time
import uuid
from typing import Any, Dict, List, Optional, Tuple
from dataclasses import dataclass, field

from server.auth.database import DatabaseAccountRepository


class SecurityEventType(enum.Enum):
    LOGIN_SUCCESS = "LOGIN_SUCCESS"
    LOGIN_FAIL = "LOGIN_FAIL"
    REGISTER = "REGISTER"
    OTP_VERIFY = "OTP_VERIFY"
    LOCKOUT = "LOCKOUT"
    REMEMBER_LOGIN = "REMEMBER_LOGIN"
    HIJACK_ATTEMPT = "HIJACK_ATTEMPT"
    LOGOUT = "LOGOUT"


@dataclass(slots=True, frozen=True)
class DeviceFingerprint:
    """Hardware and browser/client telemetry fingerprint."""
    device_id: str
    device_name: str
    hardware_hash: str
    client_ip: str


@dataclass(slots=True, frozen=True)
class RememberTokenResult:
    """Result of remember-token issuance or authentication."""
    success: bool
    message: str
    raw_token: Optional[str] = None
    account_id: Optional[str] = None


class BankGradeSecurityEngine:
    """Enterprise/Bank-Grade Authentication and Device Integrity Engine."""

    def __init__(self, repository: DatabaseAccountRepository) -> None:
        self.repository = repository

    def _hash_token(self, raw_token: str) -> str:
        return hashlib.sha256(raw_token.encode("utf-8")).hexdigest()

    def log_audit_event(
        self,
        event_type: SecurityEventType,
        account_id: Optional[str],
        ip_address: str,
        device_fingerprint: Optional[str],
        status: str,
        details: str = "",
    ) -> str:
        """Record an immutable security audit event."""
        log_id = f"aud_{uuid.uuid4().hex[:16]}"
        now = time.time()
        with self.repository._get_connection() as conn:
            conn.execute(
                """
                INSERT INTO security_audit_logs (log_id, account_id, event_type, ip_address, device_fingerprint, status, details, timestamp)
                VALUES (?, ?, ?, ?, ?, ?, ?, ?)
                """,
                (log_id, account_id, event_type.value, ip_address, device_fingerprint, status, details, now),
            )
            conn.commit()
        return log_id

    def get_audit_logs_for_account(self, account_id: str, limit: int = 50) -> List[Dict[str, Any]]:
        with self.repository._get_connection() as conn:
            cur = conn.execute(
                "SELECT * FROM security_audit_logs WHERE account_id = ? ORDER BY timestamp DESC LIMIT ?",
                (account_id, limit),
            )
            return [dict(r) for r in cur.fetchall()]

    def issue_remember_token(
        self,
        account_id: str,
        device: DeviceFingerprint,
        valid_days: int = 30,
    ) -> RememberTokenResult:
        """Issue a 256-bit cryptographically secure remember token tied to device fingerprint."""
        raw_token = secrets.token_urlsafe(32)
        token_hash = self._hash_token(raw_token)
        now = time.time()
        expires_at = now + (valid_days * 86400)

        with self.repository._get_connection() as conn:
            conn.execute(
                """
                INSERT OR REPLACE INTO remember_tokens (
                    token_hash, account_id, device_id, device_name, hardware_hash, client_ip, expires_at, is_revoked, created_at
                ) VALUES (?, ?, ?, ?, ?, ?, ?, 0, ?)
                """,
                (token_hash, account_id, device.device_id, device.device_name, device.hardware_hash, device.client_ip, expires_at, now),
            )
            conn.commit()

        self.log_audit_event(
            event_type=SecurityEventType.REGISTER,
            account_id=account_id,
            ip_address=device.client_ip,
            device_fingerprint=device.hardware_hash,
            status="SUCCESS",
            details=f"Remember token issued for {device.device_name}",
        )

        return RememberTokenResult(
            success=True,
            message="Remember token issued successfully",
            raw_token=raw_token,
            account_id=account_id,
        )

    def authenticate_remember_token(
        self,
        raw_token: str,
        current_device: DeviceFingerprint,
    ) -> RememberTokenResult:
        """Authenticate remember token with strict device fingerprint and hardware hash check."""
        token_hash = self._hash_token(raw_token)
        now = time.time()

        with self.repository._get_connection() as conn:
            cur = conn.execute(
                "SELECT * FROM remember_tokens WHERE token_hash = ? AND is_revoked = 0 AND expires_at > ?",
                (token_hash, now),
            )
            row = cur.fetchone()

        if not row:
            return RememberTokenResult(success=False, message="Invalid or expired remember-me token")

        account_id = row["account_id"]
        stored_hw_hash = row["hardware_hash"]
        stored_device_id = row["device_id"]

        # Constant-time comparison on hardware fingerprint
        is_hw_match = secrets.compare_digest(stored_hw_hash, current_device.hardware_hash)
        is_dev_match = secrets.compare_digest(stored_device_id, current_device.device_id)

        if not (is_hw_match and is_dev_match):
            # Potential token theft / session hijacking across devices!
            # Revoke token immediately for user protection
            self.revoke_remember_token(raw_token)
            self.log_audit_event(
                event_type=SecurityEventType.HIJACK_ATTEMPT,
                account_id=account_id,
                ip_address=current_device.client_ip,
                device_fingerprint=current_device.hardware_hash,
                status="REVOKED",
                details=f"Fingerprint mismatch! Expected HW: {stored_hw_hash[:8]}..., Got: {current_device.hardware_hash[:8]}...",
            )
            return RememberTokenResult(
                success=False,
                message="Device fingerprint mismatch. Re-authentication required for security.",
            )

        self.log_audit_event(
            event_type=SecurityEventType.REMEMBER_LOGIN,
            account_id=account_id,
            ip_address=current_device.client_ip,
            device_fingerprint=current_device.hardware_hash,
            status="SUCCESS",
            details="1-Click login via valid remember-me token",
        )

        return RememberTokenResult(
            success=True,
            message="Remember-me authentication successful",
            raw_token=raw_token,
            account_id=account_id,
        )

    def revoke_remember_token(self, raw_token: str) -> None:
        token_hash = self._hash_token(raw_token)
        with self.repository._get_connection() as conn:
            conn.execute("UPDATE remember_tokens SET is_revoked = 1 WHERE token_hash = ?", (token_hash,))
            conn.commit()

    def record_login_failure(
        self,
        account_id: str,
        device: DeviceFingerprint,
    ) -> Tuple[float, bool]:
        """Record login failure, compute progressive exponential delay, or lock account."""
        account = self.repository.get_by_id(account_id)
        if not account:
            return 0.0, False

        account.failed_login_attempts += 1
        now = time.time()
        is_locked = False
        delay = 0.0

        if account.failed_login_attempts == 1:
            delay = 0.0
        elif account.failed_login_attempts == 2:
            delay = 1.0
        elif account.failed_login_attempts == 3:
            delay = 2.0
        elif account.failed_login_attempts == 4:
            delay = 4.0
        else:
            # 5th attempt or higher -> 15 minutes lockout (900 seconds)
            is_locked = True
            delay = 900.0
            account.lockout_until = now + 900.0
            self.log_audit_event(
                event_type=SecurityEventType.LOCKOUT,
                account_id=account_id,
                ip_address=device.client_ip,
                device_fingerprint=device.hardware_hash,
                status="LOCKED",
                details=f"Account locked for 15 minutes after {account.failed_login_attempts} failed attempts",
            )

        self.repository.update_account(account)

        if not is_locked:
            self.log_audit_event(
                event_type=SecurityEventType.LOGIN_FAIL,
                account_id=account_id,
                ip_address=device.client_ip,
                device_fingerprint=device.hardware_hash,
                status="FAIL",
                details=f"Failed attempt #{account.failed_login_attempts}",
            )

        return delay, is_locked

    def record_login_success(
        self,
        account_id: str,
        device: DeviceFingerprint,
    ) -> None:
        """Reset failed attempt counters on successful login."""
        account = self.repository.get_by_id(account_id)
        if account:
            account.failed_login_attempts = 0
            account.lockout_until = None
            account.last_login_at = time.time()
            self.repository.update_account(account)

            self.log_audit_event(
                event_type=SecurityEventType.LOGIN_SUCCESS,
                account_id=account_id,
                ip_address=device.client_ip,
                device_fingerprint=device.hardware_hash,
                status="SUCCESS",
                details=f"User authenticated successfully on {device.device_name}",
            )
