"""
Triage: classify inventory rows by security severity.
Reads data/inventory.csv + data/hosts.csv, produces:
  - data/critical_findings.csv
  - data/high_findings.csv
  - data/medium_findings.csv
  - docs/triage_summary.md
"""

import csv
import re
import sys
from collections import Counter, defaultdict
from pathlib import Path

ROOT_DIR = Path(__file__).resolve().parents[2]
DATA_DIR = ROOT_DIR / "findings" / "data"
DOCS_DIR = ROOT_DIR / "docs" / "processing" / "legacy_dump_assessment"

INVENTORY_CSV = DATA_DIR / "legacy_dump_inventory.csv"
HOSTS_CSV = DATA_DIR / "legacy_dump_hosts.csv"

NOISE_PATTERNS = [
    re.compile(r"journeymap", re.I),
    re.compile(r"minecraft", re.I),
    re.compile(r"minecells", re.I),
    re.compile(r"sgjourney", re.I),
    re.compile(r"twilightforest", re.I),
    re.compile(r"deeperdarker", re.I),
    re.compile(r"music_disc", re.I),
    re.compile(r"qt\.conf$", re.I),
    re.compile(r"xboxservices\.config$", re.I),
    re.compile(r"acrox_", re.I),
    re.compile(r"BedrockLauncher", re.I),
    re.compile(r"nlog\.config$", re.I),
    re.compile(r"updatetime\.config$", re.I),
    re.compile(r"\.theme\.config$", re.I),
    re.compile(r"emulator", re.I),
    re.compile(r"dolphin", re.I),
    re.compile(r"pcsx2", re.I),
    re.compile(r"PrismLauncher", re.I),
    re.compile(r"SummonersWar", re.I),
    re.compile(r"Super Animal Royale", re.I),
    re.compile(r"Clicker Heroes", re.I),
    re.compile(r"Kerbal Space Program", re.I),
    re.compile(r"PngTuber", re.I),
    re.compile(r"WorldPainter", re.I),
    re.compile(r"Azahar", re.I),
    re.compile(r"curseforge", re.I),
    re.compile(r"ModrinthApp", re.I),
    re.compile(r"fluency[/\\]lm", re.I),
    re.compile(r"scan_config\.json$", re.I),
]

CRITICAL_SSH_NAMES = {
    "id_rsa", "id_dsa", "id_ecdsa", "id_ed25519",
    "id_rsa_", "id_ed25519_",
    "authorized_keys", "known_hosts",
}

CRITICAL_EXTENSIONS = {
    ".kdbx", ".kdb",
    ".ppk",
    ".ovpn",
    ".rdp",
    ".pfx", ".p12",
}

HIGH_EXTENSIONS = {".reg", ".csv"}

MEDIUM_EXTENSIONS = {
    ".crt", ".cer", ".der", ".crl",
    ".pub",
    ".p7b", ".p7c", ".spc",
    ".jks", ".keystore", ".truststore",
    ".csr",
    ".pcf",
}

OUTPUT_FIELDS = [
    "dump_id", "computer_name", "username",
    "severity", "reason",
    "relative_path", "filename", "extension",
    "category", "size_bytes",
]


def is_noise(row: dict) -> bool:
    rel = row["relative_path"]
    fname = row["filename"]
    for pat in NOISE_PATTERNS:
        if pat.search(rel) or pat.search(fname):
            return True
    return False


def classify(row: dict) -> tuple[str, str] | None:
    """Return (severity, reason) or None for noise/irrelevant."""
    if is_noise(row):
        return None

    ext = row["extension"]
    fname = row["filename"].lower()
    cat = row["category"]
    rel = row["relative_path"].lower()

    if cat == "root" and fname in ("system_info.txt", "scan_config.json"):
        return None
    if cat == "results":
        return None

    if ext in CRITICAL_EXTENSIONS:
        return ("CRITICAL", f"critical_extension:{ext}")

    stem = Path(fname).stem.lower()
    if cat.startswith("files/ssh"):
        if ext == ".pub":
            return ("MEDIUM", f"ssh_public_key:{fname}")
        if stem in ("id_rsa", "id_dsa", "id_ecdsa", "id_ed25519"):
            return ("CRITICAL", f"ssh_private_key:{fname}")
        if stem == "authorized_keys":
            return ("HIGH", f"ssh_authorized_keys:{fname}")
        if stem == "known_hosts":
            return ("MEDIUM", f"ssh_known_hosts:{fname}")
        if stem == "config":
            return ("HIGH", f"ssh_possible_config:{fname}")
        if ext == "" and "config" not in fname:
            return ("HIGH", f"ssh_file_unknown:{fname}")

    if ext == ".key":
        return ("CRITICAL", "private_key_extension:.key")

    if ext == ".pem":
        key_indicators = ("privkey", "private", "key", "id_rsa", "id_ed25519", "server.pem", "client.pem")
        if any(kw in fname.lower() for kw in key_indicators):
            return ("CRITICAL", f"pem_private_key:{fname}")
        cert_indicators = ("cert", "ca-bundle", "cacert", "fullchain", "chain", "root", "intermediate")
        if any(kw in fname.lower() for kw in cert_indicators):
            return ("MEDIUM", f"pem_certificate:{fname}")
        return ("HIGH", f"pem_unknown:{fname}")

    if ext in (".xml",):
        if any(n in fname for n in ("sitemanager", "recentservers", "filezilla")):
            return ("CRITICAL", f"filezilla_xml:{fname}")

    if ext in (".ini",):
        if "winscp" in fname:
            return ("CRITICAL", f"winscp_ini:{fname}")

    if ext in (".conf", ".config"):
        if any(kw in rel for kw in ("vpn", "openvpn", "wireguard")):
            if any(kw in fname for kw in ("openvpn", "wireguard", "wg0", "wg1", "vpn")):
                return ("CRITICAL", f"vpn_config:{fname}")
        if "ssh" in cat and "config" in fname:
            return ("CRITICAL", f"ssh_config:{fname}")
        if fname in ("sshd_config", "ssh_config"):
            return ("CRITICAL", f"ssh_config:{fname}")
        if fname == "rasphone.pbk":
            return ("MEDIUM", "ras_phonebook")
        return None

    if cat == "vault":
        return ("HIGH", "windows_vault")
    if cat == "credentials":
        return ("HIGH", f"credential_file:{fname}")
    if ext in HIGH_EXTENSIONS:
        if ext == ".reg":
            return ("HIGH", "registry_export:.reg")
        if ext == ".csv" and "vault" in rel:
            return ("HIGH", "vault_csv")

    if ext in MEDIUM_EXTENSIONS:
        return ("MEDIUM", f"medium_extension:{ext}")

    if ext == ".pbk":
        return ("MEDIUM", "phonebook:.pbk")

    return None


def load_hosts() -> dict[int, dict]:
    hosts = {}
    with open(HOSTS_CSV, encoding="utf-8") as f:
        for row in csv.DictReader(f):
            hosts[int(row["dump_id"])] = row
    return hosts


def run():
    DATA_DIR.mkdir(parents=True, exist_ok=True)
    DOCS_DIR.mkdir(parents=True, exist_ok=True)

    hosts = load_hosts()

    severity_files = {
        "CRITICAL": open(DATA_DIR / "legacy_critical_findings.csv", "w", newline="", encoding="utf-8"),
        "HIGH": open(DATA_DIR / "legacy_high_findings.csv", "w", newline="", encoding="utf-8"),
        "MEDIUM": open(DATA_DIR / "legacy_medium_findings.csv", "w", newline="", encoding="utf-8"),
    }
    writers = {}
    for sev, fh in severity_files.items():
        w = csv.DictWriter(fh, fieldnames=OUTPUT_FIELDS)
        w.writeheader()
        writers[sev] = w

    stats = Counter()
    by_host = defaultdict(lambda: defaultdict(list))
    ext_counter = defaultdict(lambda: Counter())
    total = 0
    classified = 0
    noise = 0

    with open(INVENTORY_CSV, encoding="utf-8") as f:
        for row in csv.DictReader(f):
            total += 1
            result = classify(row)
            if result is None:
                noise += 1
                continue

            severity, reason = result
            classified += 1
            stats[severity] += 1

            dump_id = int(row["dump_id"])
            host = hosts.get(dump_id, {})
            comp = host.get("computer_name", "")
            user = host.get("username", "")

            out = {
                "dump_id": dump_id,
                "computer_name": comp,
                "username": user,
                "severity": severity,
                "reason": reason,
                "relative_path": row["relative_path"],
                "filename": row["filename"],
                "extension": row["extension"],
                "category": row["category"],
                "size_bytes": row["size_bytes"],
            }

            writers[severity].writerow(out)
            by_host[dump_id][severity].append(out)
            ext_counter[severity][row["extension"]] += 1

    for fh in severity_files.values():
        fh.close()

    print(f"Total files: {total}")
    print(f"Noise/skipped: {noise}")
    print(f"Classified: {classified}")
    for sev in ("CRITICAL", "HIGH", "MEDIUM"):
        print(f"  {sev}: {stats[sev]}")

    write_summary(stats, by_host, ext_counter, hosts, total, noise, classified)


def write_summary(stats, by_host, ext_counter, hosts, total, noise, classified):
    out = DOCS_DIR / "triage_summary.md"
    with open(out, "w", encoding="utf-8") as f:
        f.write("# Triage Summary\n\n")
        f.write("## Overview\n\n")
        f.write(f"- Total files scanned: **{total:,}**\n")
        f.write(f"- Noise / skipped: **{noise:,}**\n")
        f.write(f"- Classified findings: **{classified:,}**\n\n")

        f.write("| Severity | Count |\n")
        f.write("|----------|-------|\n")
        for sev in ("CRITICAL", "HIGH", "MEDIUM"):
            f.write(f"| {sev} | {stats[sev]:,} |\n")
        f.write("\n")

        f.write("## Findings by Extension\n\n")
        for sev in ("CRITICAL", "HIGH", "MEDIUM"):
            f.write(f"### {sev}\n\n")
            f.write("| Extension | Count |\n")
            f.write("|-----------|-------|\n")
            for ext, cnt in ext_counter[sev].most_common():
                f.write(f"| `{ext}` | {cnt:,} |\n")
            f.write("\n")

        f.write("## Top Hosts by CRITICAL Findings\n\n")
        f.write("| Dump ID | Computer | User | CRITICAL | HIGH | MEDIUM |\n")
        f.write("|---------|----------|------|----------|------|--------|\n")

        host_scores = []
        for dump_id, sevs in by_host.items():
            c = len(sevs.get("CRITICAL", []))
            h = len(sevs.get("HIGH", []))
            m = len(sevs.get("MEDIUM", []))
            host_scores.append((dump_id, c, h, m))

        host_scores.sort(key=lambda x: (x[1], x[2], x[3]), reverse=True)
        for dump_id, c, h, m in host_scores[:50]:
            host = hosts.get(dump_id, {})
            comp = host.get("computer_name", "?")
            user = host.get("username", "?")
            f.write(f"| {dump_id} | {comp} | {user} | {c} | {h} | {m} |\n")
        f.write("\n")

    print(f"\nSummary written to {out}")


if __name__ == "__main__":
    run()
