from __future__ import annotations

import re
from pathlib import Path
from typing import Any

ROOT = Path(__file__).resolve().parents[3]
FORBIDDEN_FIELDS_FILE = ROOT / "docs" / "security" / "forbidden-fields.md"

RUNTIME_LEAK_PATTERNS = [
    "/home/agent/.hermes/assets/Gesundheit/",
    "/home/agent/jarvis_runtime/",
    "~/jarvis_runtime/",
    "health_data.db",
    "finance.sqlite",
    "finance.sqlite3",
    "drive_web_url",
    "drive_file_id",
    "local_original_path",
    "extrahierte_inhalte",
]
HEALTH_CONTENT_PATTERNS = [
    "diagnosis",
    "diagnose",
    "befund",
    "labor",
    "lab",
    "medication",
    "medikament",
    "arztbericht",
    "report_text",
    "ocr",
    "patient",
    "symptom",
    "yazio",
    "apple_health",
    "blood",
    "crp",
    "cholesterol",
    "faktor",
    "eliquis",
    "colchicin",
    "adalimumab",
    "hyrimoz",
    "pdf_text",
]
SECRET_PATTERNS = [
    re.compile(r"github_pat_[A-Za-z0-9_]+", re.IGNORECASE),
    re.compile(r"ghp_[A-Za-z0-9]+", re.IGNORECASE),
    re.compile(r"[A-Z0-9_]*(TOKEN|SECRET|PASSWORD|API_KEY)[A-Z0-9_]*\s*=\s*[^\s]+", re.IGNORECASE),
    re.compile(r"passphrase\s*=\s*[^\s]+", re.IGNORECASE),
    re.compile(r"private_key\s*=\s*[^\s]+", re.IGNORECASE),
]
STACKTRACE_PATTERNS = [
    re.compile(r"Traceback.*", re.IGNORECASE | re.DOTALL),
    re.compile(r"File \"[^\"]+\", line \d+"),
]


def load_forbidden_fields() -> set[str]:
    fields: set[str] = set()
    for line in FORBIDDEN_FIELDS_FILE.read_text(encoding="utf-8").splitlines():
        stripped = line.strip()
        if not stripped.startswith("-"):
            continue
        value = stripped[1:].strip().strip("`").strip()
        if value and re.fullmatch(r"[A-Za-z0-9_]+", value):
            fields.add(value)
    return fields


def contains_forbidden_field(payload: Any) -> bool:
    forbidden = load_forbidden_fields()
    forbidden_lower = {field.lower() for field in forbidden}

    def walk(value: Any) -> bool:
        if isinstance(value, dict):
            for key, nested in value.items():
                if str(key).lower() in forbidden_lower:
                    return True
                if walk(nested):
                    return True
        elif isinstance(value, list):
            return any(walk(item) for item in value)
        elif isinstance(value, str):
            lower = value.lower()
            if any(field in lower for field in forbidden_lower):
                return True
            if any(pattern.lower() in lower for pattern in RUNTIME_LEAK_PATTERNS):
                return True
            if any(pattern.search(value) for pattern in SECRET_PATTERNS):
                return True
        return False

    return walk(payload)


def redact_error_message(message: str) -> str:
    redacted = message
    for pattern in STACKTRACE_PATTERNS:
        redacted = pattern.sub("[redacted_error]", redacted)
    for pattern in RUNTIME_LEAK_PATTERNS:
        redacted = redacted.replace(pattern, "[redacted_runtime]")
    for pattern in HEALTH_CONTENT_PATTERNS:
        redacted = re.sub(re.escape(pattern), "[redacted_health]", redacted, flags=re.IGNORECASE)
    redacted = re.sub(r"/home/agent/[^\s]+", "[redacted_path]", redacted)
    redacted = re.sub(r"/(?:tmp|var|opt|synthetic|private|root)/[^\s]+", "[redacted_path]", redacted)
    for pattern in SECRET_PATTERNS:
        redacted = pattern.sub("[redacted_secret]", redacted)
    for field in load_forbidden_fields():
        redacted = re.sub(re.escape(field), "[redacted_field]", redacted, flags=re.IGNORECASE)
    return redacted[:240]


def assert_safe_overview(payload: Any) -> None:
    if contains_forbidden_field(payload):
        raise ValueError("overview contains forbidden fields or unsafe values")


def assert_safe_module_snapshot(payload: Any) -> None:
    if contains_forbidden_field(payload):
        raise ValueError("module snapshot contains forbidden fields or unsafe values")
