"""
Anti-Bot Captcha Engine & Rate Limiter for FreeExile.
Provides stateless HMAC-signed visual/math puzzles, Proof-of-Work (Hashcash), and sliding-window rate limiting.
"""

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


class CaptchaType(enum.Enum):
    MATH = "math"
    POW = "pow"


@dataclass(slots=True, frozen=True)
class CaptchaChallenge:
    """Cryptographic Captcha or Proof-of-Work challenge."""
    challenge_id: str
    captcha_type: CaptchaType
    question: str
    seed: str
    difficulty: int
    expires_at_ms: int
    signature: str


class CaptchaEngine:
    """Stateless HMAC challenge generation and replay-resistant verification."""

    def __init__(self, hmac_secret: str) -> None:
        self.hmac_secret = hmac_secret.encode("utf-8")
        self._used_challenges: Set[str] = set()

    def _sign(self, data: str) -> str:
        return hmac.new(self.hmac_secret, data.encode("utf-8"), hashlib.sha256).hexdigest()

    def generate_math_challenge(self, ttl_seconds: int = 180) -> CaptchaChallenge:
        """Generate a simple arithmetic question with cryptographic signature."""
        challenge_id = f"cap_{uuid.uuid4().hex[:16]}"
        a = random.randint(10, 50)
        b = random.randint(1, 30)
        op = random.choice(["+", "-"])
        
        answer = a + b if op == "+" else a - b
        expires_at_ms = int((time.time() + ttl_seconds) * 1000)
        
        # Payload for signature binds ID, expected answer, and expiry
        sig_payload = f"{challenge_id}:{answer}:{expires_at_ms}"
        signature = self._sign(sig_payload)

        return CaptchaChallenge(
            challenge_id=challenge_id,
            captcha_type=CaptchaType.MATH,
            question=f"{a} {op} {b}",
            seed="",
            difficulty=0,
            expires_at_ms=expires_at_ms,
            signature=signature,
        )

    def verify_math_challenge(
        self,
        challenge_id: str,
        user_answer: str,
        signature: str,
        expires_at: int,
    ) -> Tuple[bool, str]:
        """Verify user solution with replay attack prevention."""
        now_ms = int(time.time() * 1000)
        if now_ms > expires_at:
            return False, "Captcha challenge expired"

        if challenge_id in self._used_challenges:
            return False, "Captcha challenge already used (Replay blocked)"

        cleaned_answer = user_answer.strip()
        expected_sig = self._sign(f"{challenge_id}:{cleaned_answer}:{expires_at}")

        if not secrets.compare_digest(expected_sig, signature):
            return False, "Incorrect captcha solution or invalid signature"

        # Record challenge as consumed
        self._used_challenges.add(challenge_id)
        return True, "OK"

    def generate_pow_challenge(self, difficulty: int = 2, ttl_seconds: int = 180) -> CaptchaChallenge:
        """Generate a Proof-of-Work (Hashcash) challenge."""
        challenge_id = f"pow_{uuid.uuid4().hex[:16]}"
        seed = secrets.token_hex(16)
        expires_at_ms = int((time.time() + ttl_seconds) * 1000)
        
        signature = self._sign(f"{seed}:{difficulty}:{expires_at_ms}")

        return CaptchaChallenge(
            challenge_id=challenge_id,
            captcha_type=CaptchaType.POW,
            question=f"Compute SHA256({seed} + nonce) with {difficulty} leading zeros",
            seed=seed,
            difficulty=difficulty,
            expires_at_ms=expires_at_ms,
            signature=signature,
        )

    def solve_pow(self, seed: str, difficulty: int) -> Optional[str]:
        """Utility method to solve PoW (used by genuine clients or test suites)."""
        target_prefix = "0" * difficulty
        for nonce in range(1_000_000):
            nonce_str = str(nonce)
            digest = hashlib.sha256(f"{seed}:{nonce_str}".encode("utf-8")).hexdigest()
            if digest.startswith(target_prefix):
                return nonce_str
        return None

    def verify_pow_challenge(
        self,
        seed: str,
        nonce: str,
        difficulty: int,
        signature: str,
        expires_at: int,
    ) -> Tuple[bool, str]:
        """Verify Proof-of-Work solution with single-use replay protection."""
        now_ms = int(time.time() * 1000)
        if now_ms > expires_at:
            return False, "Proof-of-Work challenge expired"

        expected_sig = self._sign(f"{seed}:{difficulty}:{expires_at}")
        if not secrets.compare_digest(expected_sig, signature):
            return False, "Invalid Proof-of-Work challenge signature"

        cache_key = f"pow_used_{seed}"
        if cache_key in self._used_challenges:
            return False, "Proof-of-Work challenge already used"

        # Check hash condition
        digest = hashlib.sha256(f"{seed}:{nonce}".encode("utf-8")).hexdigest()
        target_prefix = "0" * difficulty
        if not digest.startswith(target_prefix):
            return False, f"Proof-of-Work difficulty {difficulty} unsatisfied"

        self._used_challenges.add(cache_key)
        return True, "OK"


class RateLimiter:
    """Sliding-window rate limiter per client IP or key."""

    def __init__(self, max_requests: int = 5, window_seconds: float = 60.0) -> None:
        self.max_requests = max_requests
        self.window_seconds = window_seconds
        self._history: Dict[str, List[float]] = {}

    def is_allowed(self, key: str) -> bool:
        now = time.time()
        timestamps = self._history.setdefault(key, [])
        
        # Evict timestamps older than sliding window
        cutoff = now - self.window_seconds
        self._history[key] = [t for t in timestamps if t > cutoff]

        if len(self._history[key]) < self.max_requests:
            self._history[key].append(now)
            return True
        return False

    def get_remaining(self, key: str) -> int:
        now = time.time()
        cutoff = now - self.window_seconds
        valid = [t for t in self._history.get(key, []) if t > cutoff]
        return max(0, self.max_requests - len(valid))
