"""
High-Performance Two-Tier Chat Moderation & Anti-RMT Sentinel.
Tier 1: Synchronous Trie Profanity Filter with Unicode NFKC & Leetspeak folding (< 0.1ms).
Tier 2: Asynchronous Anti-RMT & Black Market Pattern Detection.
"""

from __future__ import annotations

import re
import unicodedata
from dataclasses import dataclass
from typing import Dict, List, Set, Tuple


class TrieNode:
    """Node structure for high-throughput string matching Trie."""
    __slots__ = ("children", "is_end_of_word", "original_word")

    def __init__(self) -> None:
        self.children: Dict[str, TrieNode] = {}
        self.is_end_of_word: bool = False
        self.original_word: str = ""


class SynchronousTrieFilter:
    """
    Tier 1 Native Filter: Aho-Corasick inspired Trie with Leetspeak & Homoglyph normalization.
    Executes in < 0.1ms per message.
    """

    DEFAULT_BANNED_WORDS = {
        "dm", "dkm", "cl", "vcl", "lon", "cac", "dit", "buoi",
        "fuck", "shit", "bitch", "scam", "cnm", "sb", "nmd",
        "dcm", "cc", "vkl", "đm", "đkm", "địt", "lồn", "cặc", "buồi"
    }

    HOMOGLYPHS: Dict[str, str] = {
        "@": "a",
        "0": "o",
        "1": "i",
        "!": "i",
        "3": "e",
        "$": "s",
        "4": "a",
        "5": "s",
        "7": "t",
        "8": "b",
    }

    CHAR_VARIANTS: Dict[str, str] = {
        "a": r"[aáàảãạăắằẳẵặâấầẩẫậ@4]",
        "b": r"[b8]",
        "c": r"[c]",
        "d": r"[dđ]",
        "e": r"[eéèẻẽẹêếềểễệ3]",
        "i": r"[iíìỉĩị!1|]",
        "k": r"[k]",
        "l": r"[l1|]",
        "m": r"[m]",
        "n": r"[n]",
        "o": r"[oóòỏõọôốồổỗộơớờởỡợ0]",
        "p": r"[p]",
        "r": r"[r]",
        "s": r"[s$5]",
        "t": r"[t7]",
        "u": r"[uúùủũụưứừửữự]",
        "v": r"[v]",
        "x": r"[x]",
        "y": r"[yýỳỷỹỵ]",
    }

    IGNORE_SEPARATORS = set(". _-/\\+*~`'\",;:()[]{}|")

    def __init__(self, banned_words: Set[str] | None = None) -> None:
        self.root = TrieNode()
        words = banned_words if banned_words is not None else self.DEFAULT_BANNED_WORDS
        for word in words:
            self._insert(self.normalize_fold(word))

    def _insert(self, word: str) -> None:
        node = self.root
        for char in word:
            if char not in node.children:
                node.children[char] = TrieNode()
            node = node.children[char]
        node.is_end_of_word = True
        node.original_word = word

    def normalize_fold(self, text: str) -> str:
        """Converts Unicode NFKC, strips zero-width chars, folds homoglyphs and diacritics."""
        normalized = unicodedata.normalize("NFKC", text).lower()
        # Strip zero-width spaces
        normalized = normalized.replace("\u200b", "").replace("\ufeff", "")
        # Map Vietnamese 'đ' to 'd'
        normalized = normalized.replace("đ", "d")
        # Fold homoglyphs
        folded = "".join(self.HOMOGLYPHS.get(c, c) for c in normalized)
        # Strip diacritics for uniform check
        nfd = unicodedata.normalize("NFD", folded)
        return "".join(c for c in nfd if unicodedata.category(c) != "Mn")

    def strip_noise(self, text: str) -> str:
        """Removes punctuation and separators used to obfuscate words."""
        return "".join(c for c in text if c not in self.IGNORE_SEPARATORS)

    def scan_and_censor(self, raw_content: str) -> Tuple[bool, str]:
        """
        Scans message content for profane or abusive keywords.
        Returns: (has_profanity, censored_content)
        """
        if not raw_content:
            return False, raw_content

        normalized_folded = self.normalize_fold(raw_content)
        condensed = self.strip_noise(normalized_folded)

        words_in_condensed: Set[str] = set()

        # Check in condensed form for bypassed words like d.k.m or c_a_c
        for start_idx in range(len(condensed)):
            node = self.root
            for curr_idx in range(start_idx, len(condensed)):
                char = condensed[curr_idx]
                if char not in node.children:
                    break
                node = node.children[char]
                if node.is_end_of_word:
                    words_in_condensed.add(node.original_word)

        if not words_in_condensed:
            return False, raw_content

        # Mask found words in original content with word boundary preservation
        result = raw_content
        censored_any = False
        # Sort by length descending to mask longer compounds first
        sorted_words = sorted(words_in_condensed, key=len, reverse=True)

        for word in sorted_words:
            pattern_parts = [self.CHAR_VARIANTS.get(ch, re.escape(ch)) for ch in word]
            pattern_body = r"[\W_]*".join(pattern_parts)
            pattern = rf"(?<![a-zA-Z0-9]){pattern_body}(?![a-zA-Z0-9])"
            new_result, count = re.subn(pattern, "***", result, flags=re.IGNORECASE)
            if count > 0:
                censored_any = True
                result = new_result

        return censored_any, result


@dataclass(slots=True, frozen=True)
class ModerationAnalysisResult:
    """Output summary of two-tier chat moderation."""
    is_censored: bool
    filtered_content: str
    is_rmt: bool
    risk_score: int
    should_auto_mute: bool
    violation_reasons: Tuple[str, ...]


class AsyncSentinelRMTDetector:
    """
    Tier 2 Asynchronous Sentinel: Real-Money Trading (RMT) and scam detector.
    Identifies phone numbers, social media handles, and bank transfer keywords.
    """

    RMT_PATTERNS = [
        ("PHONE_VN", re.compile(r"(?:0|84|\+84)(?:3|5|7|8|9)\d{8}")),
        ("SOCIAL_HANDLE", re.compile(r"(zalo|telegram|tele|fb\.com|facebook\.com)[\s\:\@\.\-_]*\w+", re.IGNORECASE)),
        ("PAYMENT_KEYWORD", re.compile(r"(chuyen khoan|atm|momo|ban vang|thu mua ngoc|ban acc|gdtg|shopacc|zalopay)", re.IGNORECASE)),
        ("CRYPTO_PAYMENT", re.compile(r"(usdt|binance|crypto|vi dien tu|bitcoin|eth)", re.IGNORECASE)),
    ]

    def evaluate(self, raw_content: str, normalized_content: str) -> Tuple[bool, int, List[str]]:
        reasons: List[str] = []
        score = 0

        for category, pattern in self.RMT_PATTERNS:
            if pattern.search(raw_content) or pattern.search(normalized_content):
                reasons.append(category)
                if category == "PHONE_VN":
                    score += 40
                elif category == "SOCIAL_HANDLE":
                    score += 35
                elif category == "PAYMENT_KEYWORD":
                    score += 50
                elif category == "CRYPTO_PAYMENT":
                    score += 45

        is_rmt = len(reasons) > 0
        return is_rmt, min(score, 100), reasons


class ChatModerationPipeline:
    """Coordinates Tier 1 Synchronous and Tier 2 Asynchronous moderation."""

    def __init__(self) -> None:
        self.trie_filter = SynchronousTrieFilter()
        self.rmt_detector = AsyncSentinelRMTDetector()

    def process(self, raw_content: str) -> ModerationAnalysisResult:
        """Processes a chat message through both moderation tiers."""
        is_profane, censored_content = self.trie_filter.scan_and_censor(raw_content)
        normalized = self.trie_filter.strip_noise(self.trie_filter.normalize_fold(raw_content))

        is_rmt, rmt_score, reasons = self.rmt_detector.evaluate(raw_content, normalized)

        total_risk = rmt_score
        if is_profane:
            total_risk += 25
            reasons.append("PROFANITY")

        final_risk = min(total_risk, 100)
        should_mute = final_risk >= 60

        return ModerationAnalysisResult(
            is_censored=is_profane,
            filtered_content=censored_content,
            is_rmt=is_rmt,
            risk_score=final_risk,
            should_auto_mute=should_mute,
            violation_reasons=tuple(reasons)
        )
