"""
Cryptographic services: Password Hashing, JWT Token Issuance & Revocation.
Implements OWASP 2026 guidelines with PBKDF2-HMAC-SHA256 and constant-time comparisons.
"""

from __future__ import annotations
import base64
import hashlib
import hmac
import secrets
import time
from typing import Any, Dict, Optional, Set
import jwt

from server.auth.models import AuthTokenPair


class PasswordHasher:
    """OWASP-compliant password hashing and constant-time verification."""
    
    def __init__(self, iterations: int = 600_000, salt_bytes: int = 16) -> None:
        self.iterations = iterations
        self.salt_bytes = salt_bytes

    def hash_password(self, password: str) -> str:
        """Hash a password using PBKDF2-HMAC-SHA256 with random salt."""
        salt = secrets.token_bytes(self.salt_bytes)
        dk = hashlib.pbkdf2_hmac(
            hash_name="sha256",
            password=password.encode("utf-8"),
            salt=salt,
            iterations=self.iterations,
            dklen=32,
        )
        salt_b64 = base64.b64encode(salt).decode("ascii")
        hash_b64 = base64.b64encode(dk).decode("ascii")
        return f"pbkdf2_sha256${self.iterations}${salt_b64}${hash_b64}"

    def verify_password(self, password: str, hashed_password: str) -> bool:
        """Verify password against stored hash using constant-time comparison."""
        try:
            algorithm, iter_str, salt_b64, hash_b64 = hashed_password.split("$")
            if algorithm != "pbkdf2_sha256":
                return False
            iterations = int(iter_str)
            salt = base64.b64decode(salt_b64.encode("ascii"))
            expected_hash = base64.b64decode(hash_b64.encode("ascii"))

            actual_hash = hashlib.pbkdf2_hmac(
                hash_name="sha256",
                password=password.encode("utf-8"),
                salt=salt,
                iterations=iterations,
                dklen=len(expected_hash),
            )
            return secrets.compare_digest(expected_hash, actual_hash)
        except Exception:
            # Prevent timing side-channels and crash resilience
            return False


class TokenService:
    """JWT Access and Refresh token generation, validation, and revocation."""

    def __init__(
        self,
        secret_key: str,
        algorithm: str = "HS256",
        access_token_ttl_seconds: int = 900,       # 15 minutes
        refresh_token_ttl_seconds: int = 604_800,  # 7 days
    ) -> None:
        self.secret_key = secret_key
        self.algorithm = algorithm
        self.access_token_ttl_seconds = access_token_ttl_seconds
        self.refresh_token_ttl_seconds = refresh_token_ttl_seconds
        self._revoked_tokens: Set[str] = set()

    def generate_token_pair(
        self,
        account_id: str,
        roles: list[str],
        email: str,
    ) -> AuthTokenPair:
        """Create signed access token and high-entropy refresh token."""
        now = int(time.time())
        jti = secrets.token_hex(16)
        access_payload = {
            "sub": account_id,
            "email": email,
            "roles": roles,
            "iat": now,
            "exp": now + self.access_token_ttl_seconds,
            "jti": jti,
            "iss": "freeexile-auth-gateway",
        }
        access_token = jwt.encode(access_payload, self.secret_key, algorithm=self.algorithm)
        refresh_token = secrets.token_urlsafe(32)

        return AuthTokenPair(
            access_token=access_token,
            refresh_token=refresh_token,
            expires_in_seconds=self.access_token_ttl_seconds,
            token_type="Bearer",
        )

    def verify_access_token(self, token: str) -> Optional[Dict[str, Any]]:
        """Verify token signature and ensure it has not been revoked."""
        if self.is_token_revoked(token):
            return None
        try:
            payload: Dict[str, Any] = jwt.decode(
                token,
                self.secret_key,
                algorithms=[self.algorithm],
                issuer="freeexile-auth-gateway",
            )
            return payload
        except (jwt.PyJWTError, Exception):
            return None

    def revoke_token(self, token: str) -> None:
        """Revoke a token immediately."""
        self._revoked_tokens.add(token)

    def is_token_revoked(self, token: str) -> bool:
        """Check if token is in blacklist."""
        return token in self._revoked_tokens
