from __future__ import annotations

import re
from collections import Counter
from dataclasses import dataclass, field
from enum import Enum

from app.models.context import BuiltContext, Citation, ContextChunk

REDACTED_SECRET = "[REDACTED_SECRET]"
REDACTED_INJECTION = "[REDACTED_INJECTION]"

SAFE_REFUSAL_MESSAGE = (
    "I can't reveal secrets, credentials, or exploit-ready instructions from the indexed context. "
    "I can help explain the risk and suggest remediation."
)

DEFAULT_MAX_SECRET_FINDINGS = 3


class GuardrailCategory(str, Enum):
    SECRET_LEAKAGE = "secret_leakage"
    PROMPT_INJECTION = "prompt_injection"
    UNSAFE_VULNERABILITY_GUIDANCE = "unsafe_vulnerability_guidance"


class GuardrailAction(str, Enum):
    ALLOW = "allow"
    REDACT = "redact"
    BLOCK = "block"
    REFUSE = "refuse"


class GuardrailSeverity(str, Enum):
    LOW = "low"
    MEDIUM = "medium"
    HIGH = "high"


@dataclass(frozen=True)
class GuardrailFinding:
    category: GuardrailCategory
    severity: GuardrailSeverity
    count: int = 1


@dataclass
class GuardrailResult:
    action: GuardrailAction = GuardrailAction.ALLOW
    findings: list[GuardrailFinding] = field(default_factory=list)
    context_redacted_count: int = 0
    context_blocked_count: int = 0
    output_refused: bool = False

    def summary(self) -> dict[str, int | str | bool]:
        counts: Counter[str] = Counter()
        severities: Counter[str] = Counter()
        for finding in self.findings:
            counts[finding.category.value] += finding.count
            severities[finding.severity.value] += finding.count
        return {
            "action": self.action.value,
            "context_redacted_count": self.context_redacted_count,
            "context_blocked_count": self.context_blocked_count,
            "output_refused": self.output_refused,
            "finding_counts": dict(counts),
            "severity_counts": dict(severities),
        }


_SECRET_PATTERNS: list[tuple[re.Pattern[str], GuardrailSeverity]] = [
    (re.compile(r"-----BEGIN (?:RSA |EC |OPENSSH )?PRIVATE KEY-----[\s\S]*?-----END (?:RSA |EC |OPENSSH )?PRIVATE KEY-----", re.IGNORECASE), GuardrailSeverity.HIGH),
    (re.compile(r"\bAKIA[0-9A-Z]{16}\b"), GuardrailSeverity.HIGH),
    (re.compile(r"\b(?:sk|pk)_(?:live|test)_[A-Za-z0-9]{16,}\b"), GuardrailSeverity.HIGH),
    (re.compile(r"\bghp_[A-Za-z0-9]{20,}\b"), GuardrailSeverity.HIGH),
    (re.compile(r"\bgithub_pat_[A-Za-z0-9_]{20,}\b"), GuardrailSeverity.HIGH),
    (re.compile(r"\bxox[baprs]-[A-Za-z0-9-]{10,}\b"), GuardrailSeverity.HIGH),
    (re.compile(r"\beyJ[A-Za-z0-9_-]{8,}\.[A-Za-z0-9_-]{8,}\.[A-Za-z0-9_-]{8,}\b"), GuardrailSeverity.HIGH),
    (re.compile(r"\bBearer\s+[A-Za-z0-9._-]{12,}\b", re.IGNORECASE), GuardrailSeverity.HIGH),
    (re.compile(r"\b(?:postgres|postgresql|mysql|mongodb)(?:\+srv)?://[^\s:@/]+:[^\s@/]+@", re.IGNORECASE), GuardrailSeverity.HIGH),
    (re.compile(r"(?i)(\b(?:api[_-]?key|client[_-]?secret|access[_-]?token|refresh[_-]?token|password|passwd|secret|private[_-]?key)\s*=\s*['\"]?)([^\s'\"#]{8,})(['\"]?)", re.IGNORECASE), GuardrailSeverity.MEDIUM),
    (re.compile(r"(?i)(\b(?:aws_secret_access_key|openai_api_key|anthropic_api_key)\s*[:=]\s*['\"]?)([^\s'\"]{8,})(['\"]?)", re.IGNORECASE), GuardrailSeverity.HIGH),
]

_PROMPT_INJECTION_PATTERNS: list[tuple[re.Pattern[str], GuardrailSeverity]] = [
    (re.compile(r"(?i)ignore (?:all )?(?:previous|above|prior|earlier) instructions"), GuardrailSeverity.HIGH),
    (re.compile(r"(?i)disregard (?:your |the )?(?:system )?instructions"), GuardrailSeverity.HIGH),
    (re.compile(r"(?i)you are now (?:a |an )?(?:unrestricted|jailbroken|DAN)"), GuardrailSeverity.HIGH),
    (re.compile(r"(?i)reveal (?:your |the )?(?:system )?prompt"), GuardrailSeverity.HIGH),
    (re.compile(r"(?i)(?:dump|exfiltrate|leak) (?:all )?(?:secrets|credentials|tokens|keys)"), GuardrailSeverity.HIGH),
    (re.compile(r"(?i)bypass (?:security|guardrails|policy|restrictions)"), GuardrailSeverity.HIGH),
    (re.compile(r"(?i)(?:new|updated) instructions\s*:"), GuardrailSeverity.MEDIUM),
    (re.compile(r"(?i)override (?:system |safety )?instructions"), GuardrailSeverity.HIGH),
    (re.compile(r"(?i)pretend (?:you are|to be) (?:not bound|unrestricted)"), GuardrailSeverity.MEDIUM),
]

_UNSAFE_VULN_PATTERNS: list[tuple[re.Pattern[str], GuardrailSeverity]] = [
    (re.compile(r"(?i)\b(?:run|execute) (?:this |the following )?(?:curl|command|payload) (?:to )?(?:exploit|bypass|attack)"), GuardrailSeverity.HIGH),
    (re.compile(r"(?i)\b(?:sqlmap|metasploit|hydra|hashcat)\b.*\b(?:against|target|exploit)\b"), GuardrailSeverity.HIGH),
    (re.compile(r"(?i)\b(?:step[- ]by[- ]step|detailed) (?:exploit|attack) (?:instructions|steps|guide)\b"), GuardrailSeverity.HIGH),
    (re.compile(r"(?i)'\s*OR\s+1\s*=\s*1"), GuardrailSeverity.MEDIUM),
    (re.compile(r"(?i)<script[^>]*>.*?</script>"), GuardrailSeverity.MEDIUM),
    (re.compile(r"(?i)\bunion\s+select\b.*\bfrom\b"), GuardrailSeverity.MEDIUM),
    (re.compile(r"(?i)\bcurl\b[^\n]{0,120}\b(?:exploit|bypass|inject|attack)\b"), GuardrailSeverity.HIGH),
]


def _count_matches(text: str, patterns: list[tuple[re.Pattern[str], GuardrailSeverity]]) -> tuple[int, GuardrailSeverity]:
    total = 0
    max_severity = GuardrailSeverity.LOW
    severity_rank = {
        GuardrailSeverity.LOW: 0,
        GuardrailSeverity.MEDIUM: 1,
        GuardrailSeverity.HIGH: 2,
    }
    for pattern, severity in patterns:
        matches = pattern.findall(text)
        if matches:
            total += len(matches)
            if severity_rank[severity] > severity_rank[max_severity]:
                max_severity = severity
    return total, max_severity


def _redact_patterns(
    text: str,
    patterns: list[tuple[re.Pattern[str], GuardrailSeverity]],
    placeholder: str,
) -> tuple[str, int]:
    redacted_count = 0

    def _replace(match: re.Match[str]) -> str:
        nonlocal redacted_count
        redacted_count += 1
        if match.lastindex and match.lastindex >= 2:
            return f"{match.group(1)}{placeholder}{match.group(3) if match.lastindex >= 3 else ''}"
        return placeholder

    sanitized = text
    for pattern, _severity in patterns:
        sanitized = pattern.sub(_replace, sanitized)
    return sanitized, redacted_count


def _injection_density(text: str) -> float:
    if not text.strip():
        return 0.0
    injection_count, _ = _count_matches(text, _PROMPT_INJECTION_PATTERNS)
    return injection_count / max(len(text.split()), 1)


class GuardrailService:
    """Deterministic guardrails for retrieved context and LLM output."""

    def __init__(
        self,
        *,
        enabled: bool | None = None,
        max_secret_findings: int | None = None,
        block_on_context_injection: bool | None = None,
    ):
        self._enabled = True if enabled is None else enabled
        self._max_secret_findings = (
            DEFAULT_MAX_SECRET_FINDINGS
            if max_secret_findings is None
            else max_secret_findings
        )
        self._block_on_context_injection = (
            True if block_on_context_injection is None else block_on_context_injection
        )

    @property
    def enabled(self) -> bool:
        return self._enabled

    def sanitize_context(self, built: BuiltContext) -> tuple[BuiltContext, GuardrailResult]:
        if not self._enabled:
            return built, GuardrailResult(action=GuardrailAction.ALLOW)

        findings: list[GuardrailFinding] = []
        sanitized_chunks: list[ContextChunk] = []
        citations: list[Citation] = []
        parts: list[str] = []
        redacted_count = 0
        blocked_count = 0

        for chunk in built.chunks:
            text = chunk.text
            secret_count, secret_severity = _count_matches(text, _SECRET_PATTERNS)
            injection_count, injection_severity = _count_matches(text, _PROMPT_INJECTION_PATTERNS)

            if secret_count:
                findings.append(
                    GuardrailFinding(
                        category=GuardrailCategory.SECRET_LEAKAGE,
                        severity=secret_severity,
                        count=secret_count,
                    )
                )
                text, secret_redactions = _redact_patterns(text, _SECRET_PATTERNS, REDACTED_SECRET)
                redacted_count += secret_redactions

            if injection_count:
                findings.append(
                    GuardrailFinding(
                        category=GuardrailCategory.PROMPT_INJECTION,
                        severity=injection_severity,
                        count=injection_count,
                    )
                )
                should_block = (
                    self._block_on_context_injection
                    and (
                        injection_severity == GuardrailSeverity.HIGH
                        or _injection_density(chunk.text) >= 0.05
                    )
                )
                if should_block:
                    blocked_count += 1
                    continue
                text, injection_redactions = _redact_patterns(
                    text,
                    _PROMPT_INJECTION_PATTERNS,
                    REDACTED_INJECTION,
                )
                redacted_count += injection_redactions

            if not text.strip() or text.strip() in {REDACTED_SECRET, REDACTED_INJECTION}:
                blocked_count += 1
                continue

            sanitized_chunk = chunk.model_copy(update={"text": text})
            sanitized_chunks.append(sanitized_chunk)
            header = f"[{chunk.index}] {chunk.label}\n"
            parts.append(f"{header}{text.strip()}\n")
            citation = next((c for c in built.citations if c.index == chunk.index), None)
            if citation:
                citations.append(citation)

        action = GuardrailAction.ALLOW
        if blocked_count or redacted_count:
            action = GuardrailAction.REDACT if sanitized_chunks else GuardrailAction.BLOCK
        if secret_count_over_limit(findings, self._max_secret_findings):
            action = GuardrailAction.BLOCK
            sanitized_chunks = []
            citations = []
            parts = []

        sanitized_built = BuiltContext(
            chunks=sanitized_chunks,
            prompt_text="".join(parts),
            citations=citations,
            total_tokens=built.total_tokens,
            trimmed_count=built.trimmed_count + blocked_count,
        )
        result = GuardrailResult(
            action=action,
            findings=findings,
            context_redacted_count=redacted_count,
            context_blocked_count=blocked_count,
        )
        return sanitized_built, result

    def validate_output(self, answer: str) -> tuple[str, GuardrailResult]:
        if not self._enabled or not answer.strip():
            return answer, GuardrailResult(action=GuardrailAction.ALLOW)

        findings: list[GuardrailFinding] = []
        secret_count, secret_severity = _count_matches(answer, _SECRET_PATTERNS)
        injection_count, injection_severity = _count_matches(answer, _PROMPT_INJECTION_PATTERNS)
        vuln_count, vuln_severity = _count_matches(answer, _UNSAFE_VULN_PATTERNS)

        if secret_count:
            findings.append(
                GuardrailFinding(
                    category=GuardrailCategory.SECRET_LEAKAGE,
                    severity=secret_severity,
                    count=secret_count,
                )
            )
        if injection_count:
            findings.append(
                GuardrailFinding(
                    category=GuardrailCategory.PROMPT_INJECTION,
                    severity=injection_severity,
                    count=injection_count,
                )
            )
        if vuln_count:
            findings.append(
                GuardrailFinding(
                    category=GuardrailCategory.UNSAFE_VULNERABILITY_GUIDANCE,
                    severity=vuln_severity,
                    count=vuln_count,
                )
            )

        should_refuse = bool(findings) and any(
            finding.severity in {GuardrailSeverity.MEDIUM, GuardrailSeverity.HIGH}
            for finding in findings
        )
        if should_refuse:
            return SAFE_REFUSAL_MESSAGE, GuardrailResult(
                action=GuardrailAction.REFUSE,
                findings=findings,
                output_refused=True,
            )

        sanitized, _ = _redact_patterns(answer, _SECRET_PATTERNS, REDACTED_SECRET)
        action = GuardrailAction.REDACT if sanitized != answer else GuardrailAction.ALLOW
        return sanitized, GuardrailResult(action=action, findings=findings)

    def redact_stream_delta(self, delta: str) -> str:
        if not self._enabled or not delta:
            return delta
        sanitized, _ = _redact_patterns(delta, _SECRET_PATTERNS, REDACTED_SECRET)
        sanitized, _ = _redact_patterns(sanitized, _PROMPT_INJECTION_PATTERNS, REDACTED_INJECTION)
        return sanitized


def secret_count_over_limit(findings: list[GuardrailFinding], max_secret_findings: int) -> bool:
    total = sum(
        finding.count
        for finding in findings
        if finding.category == GuardrailCategory.SECRET_LEAKAGE
    )
    return total > max_secret_findings
