"""
Cryptographic helpers — single source of truth for all security primitives.

Design decisions
----------------
Passwords / bcrypt
  passlib CryptContext with bcrypt, cost factor from settings (min 12).
  verify_password() is constant-time by passlib design.

API key secrets — two storage strategies
  Caller Type 2 (HMAC-signed external): the server must *recompute* the HMAC
  to verify each request. bcrypt is one-way and cannot be used here.
  Solution: Fernet symmetric encryption (AES-128-CBC + HMAC-SHA256 envelope)
  so the server can recover the plain secret at verify time.

  Caller Type 3 (cron, simple): bcrypt is sufficient — the server just
  verifies the raw key against the stored hash without needing the plain value.

  The secret_hash column stores either:
    "$2b$..."       → bcrypt  (cron keys)
    "fernet:..."    → Fernet  (HMAC keys)

JWT
  HS256 with SECRET_KEY. Access tokens are short-lived (15 min) and never
  hit the DB on verification. Refresh tokens embed a jti that is checked in
  the DB for revocation status.

HMAC request signing
  HMAC-SHA256(secret, f"{timestamp}.{body_bytes}").hexdigest()
  Compared with hmac.compare_digest() — never == to prevent timing attacks.
"""

from __future__ import annotations

import hashlib
import hmac as _hmac
import secrets
import time
from datetime import datetime, timezone
from typing import Any

import bcrypt as _bcrypt_lib
from cryptography.fernet import Fernet, InvalidToken
from jose import jwt
from passlib.context import CryptContext

# passlib ≤ 1.7.4 reads bcrypt.__about__.__version__ which was removed in
# bcrypt 4.x.  Injecting a minimal shim silences the "(trapped)" warning
# without downgrading either package.
if not hasattr(_bcrypt_lib, "__about__"):
    class _About:
        __version__ = _bcrypt_lib.__version__
    _bcrypt_lib.__about__ = _About()  # type: ignore[attr-defined]

from app.config import get_settings

settings = get_settings()

ALGORITHM = "HS256"

# bcrypt context — cost factor validated at startup (min 12)
_pwd_ctx = CryptContext(
    schemes=["bcrypt"],
    deprecated="auto",
    bcrypt__rounds=settings.bcrypt_rounds,
)

# Fernet instance for symmetric encryption of HMAC API secrets
_fernet = Fernet(settings.fernet_key)

# Column value prefix to distinguish storage type without a separate DB column
_FERNET_PREFIX = "fernet:"


# ---------------------------------------------------------------------------
# Passwords  (users table)
# ---------------------------------------------------------------------------

def hash_password(plain: str) -> str:
    """Return a bcrypt hash of a plaintext password. Never store plain."""
    return _pwd_ctx.hash(plain)


def verify_password(plain: str, hashed: str) -> bool:
    """
    Constant-time bcrypt verification.
    Called even when user is not found (dummy hash) to prevent timing leaks.
    """
    return _pwd_ctx.verify(plain, hashed)


# ---------------------------------------------------------------------------
# API key secrets
# ---------------------------------------------------------------------------

def generate_api_secret() -> str:
    """Generate a 32-byte hex secret (256 bits of entropy)."""
    return secrets.token_hex(32)


def encrypt_api_secret(plain: str) -> str:
    """
    Fernet-encrypt an API key secret for HMAC callers (Type 2).
    Fernet = AES-128-CBC + HMAC-SHA256 envelope using the derived fernet_key.
    The plain secret is never stored on disk or in the DB.
    """
    return _FERNET_PREFIX + _fernet.encrypt(plain.encode()).decode()


def decrypt_api_secret(stored: str) -> str | None:
    """
    Decrypt a Fernet-stored API secret.
    Returns None if stored value is not Fernet (e.g. a bcrypt hash for cron keys).
    Returns None on decryption failure (tampered ciphertext).
    """
    if not stored.startswith(_FERNET_PREFIX):
        return None
    try:
        return _fernet.decrypt(stored[len(_FERNET_PREFIX):].encode()).decode()
    except InvalidToken:
        return None


def hash_api_secret_bcrypt(plain: str) -> str:
    """
    bcrypt-hash an API key secret for cron keys (Type 3).
    Use this when HMAC recomputation is not needed — verify with compare_bcrypt.
    """
    return _pwd_ctx.hash(plain)


def verify_api_secret_bcrypt(plain: str, stored: str) -> bool:
    """Constant-time bcrypt verify for cron API keys."""
    return _pwd_ctx.verify(plain, stored)


def verify_api_secret_plain(plain: str, expected: str) -> bool:
    """
    Constant-time compare for HMAC subscriber secrets (Bearer auth mode).

    Uses hmac.compare_digest so timing does not leak secret length/content.
    """
    if not plain or not expected:
        return False
    return _hmac.compare_digest(plain.encode(), expected.encode())


# ---------------------------------------------------------------------------
# HMAC-SHA256 request signing  (Caller Type 2)
# ---------------------------------------------------------------------------

def verify_hmac_signature(
    plain_secret: str,
    timestamp: str,
    body: bytes,
    received_signature: str,
) -> bool:
    """
    Verify an HMAC-SHA256 request signature.

    Protocol (subscriber must implement the same):
        payload   = f"{timestamp}.{body_bytes}"
        signature = HMAC-SHA256(secret, payload).hexdigest()
        header    = "sha256=" + signature

    Security: hmac.compare_digest() — constant-time, prevents timing oracle.
    Never use == for secret comparison.
    """
    msg = f"{timestamp}.".encode() + body
    expected_hex = _hmac.new(
        plain_secret.encode(),
        msg,
        hashlib.sha256,
    ).hexdigest()
    expected = f"sha256={expected_hex}"
    # Constant-time comparison — both operands same type (str)
    return _hmac.compare_digest(expected, received_signature)


def is_timestamp_fresh(timestamp_str: str) -> bool:
    """
    Return True if |now() - timestamp| <= HMAC_TIMESTAMP_TOLERANCE_SECONDS.
    Rejects replayed requests outside the tolerance window (default ±5 min).
    """
    try:
        ts = int(timestamp_str)
    except (ValueError, TypeError):
        return False
    return abs(int(time.time()) - ts) <= settings.hmac_timestamp_tolerance_seconds


# ---------------------------------------------------------------------------
# JWT  (Caller Type 1 — human users)
# ---------------------------------------------------------------------------

def create_access_token(sub: str, role: str) -> tuple[str, int]:
    """
    Issue an HS256 access token.

    Payload contains only non-sensitive identity claims.
    Never include: email, PII, permissions lists, or large objects.
    The payload is BASE64-encoded and *signed*, not encrypted — anyone can
    read it. Keep it minimal.

    Returns (token_string, exp_unix_timestamp).
    """
    now = int(datetime.now(timezone.utc).timestamp())
    exp = now + settings.access_token_ttl
    payload: dict[str, Any] = {
        "sub": sub,       # user UUID string — opaque identifier
        "role": role,     # viewer | analyst | admin
        "type": "access",
        "iat": now,
        "exp": exp,
    }
    token = jwt.encode(payload, settings.secret_key, algorithm=ALGORITHM)
    return token, exp


def create_refresh_token(jti: str, sub: str) -> tuple[str, datetime]:
    """
    Issue an HS256 refresh token that embeds the DB row identifier (jti).

    The jti UUID is stored in the refresh_tokens table for revocation checks.
    The JWT signature prevents the client from crafting arbitrary jti values.

    Returns (token_string, expiry_datetime).
    """
    now = int(datetime.now(timezone.utc).timestamp())
    exp_ts = now + settings.refresh_token_ttl
    exp_dt = datetime.fromtimestamp(exp_ts, tz=timezone.utc)
    payload: dict[str, Any] = {
        "jti": jti,       # maps to refresh_tokens.jti in DB
        "sub": sub,       # user UUID
        "type": "refresh",
        "iat": now,
        "exp": exp_ts,
    }
    token = jwt.encode(payload, settings.secret_key, algorithm=ALGORITHM)
    return token, exp_dt


def decode_token(token: str) -> dict[str, Any]:
    """
    Decode + verify JWT signature and expiry.
    Raises jose.JWTError on tampered or expired tokens.
    Does NOT touch the database — pure cryptographic check.
    """
    return jwt.decode(token, settings.secret_key, algorithms=[ALGORITHM])
