"""
Generate assessment reports from extracted_credentials.json and triage data.

Produces:
  - docs/critical_report.md   -- CRITICAL findings grouped by host
  - docs/access_vectors.md    -- access vectors (SSH, VPN, RDP, FTP)
"""

import csv
import json
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"

HOSTS_CSV = DATA_DIR / "legacy_dump_hosts.csv"
EXTRACTED_JSON = DATA_DIR / "legacy_extracted_credentials.json"
CRITICAL_CSV = DATA_DIR / "legacy_critical_findings.csv"
HIGH_CSV = DATA_DIR / "legacy_high_findings.csv"


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 load_findings() -> list[dict]:
    with open(EXTRACTED_JSON, encoding="utf-8") as f:
        return json.load(f)


def load_triage(path: Path) -> list[dict]:
    if not path.exists():
        return []
    with open(path, encoding="utf-8") as f:
        return list(csv.DictReader(f))


def generate_critical_report(findings: list[dict], hosts: dict):
    by_host = defaultdict(list)
    for entry in findings:
        by_host[entry["dump_id"]].append(entry)

    host_scores = []
    for dump_id, entries in by_host.items():
        score = 0
        crit_count = 0
        for e in entries:
            for f in e["findings"]:
                t = f["type"]
                if t == "private_key" and f.get("detail") == "UNENCRYPTED":
                    score += 100
                elif t == "private_key":
                    score += 50
                elif t == "password":
                    score += 80
                elif t == "filezilla_password":
                    score += 90
                elif t == "winscp_password":
                    score += 90
                elif t == "rdp_password_hash":
                    score += 95
                elif t == "vpn_auth_user_pass":
                    score += 70
                elif t == "vpn_endpoint":
                    score += 40
                elif t == "ftp_endpoint":
                    score += 30
                elif t == "aws_key":
                    score += 100
                elif t == "mnemonic_phrase":
                    score += 100
                elif t == "vault_credential":
                    score += 10
            if e["severity"] == "CRITICAL":
                crit_count += 1
        host_scores.append((dump_id, score, crit_count, entries))

    host_scores.sort(key=lambda x: x[1], reverse=True)

    out = DOCS_DIR / "critical_report.md"
    with open(out, "w", encoding="utf-8") as f:
        f.write("# Critical Findings Report\n\n")
        f.write("Hosts ranked by risk score (higher = more exploitable).\n\n")
        f.write("---\n\n")

        for dump_id, score, crit_count, entries in host_scores:
            if score == 0:
                continue
            host = hosts.get(dump_id, {})
            comp = host.get("computer_name", "?")
            user = host.get("username", "?")

            f.write(f"## [{dump_id}] {comp} / {user}  (score: {score})\n\n")

            type_groups = defaultdict(list)
            for e in entries:
                for finding in e["findings"]:
                    type_groups[finding["type"]].append({
                        "file": e["file"],
                        "detail": finding.get("detail", ""),
                        "severity": e["severity"],
                    })

            high_value_types = [
                "private_key", "password", "filezilla_password", "winscp_password",
                "rdp_password_hash", "vpn_auth_user_pass", "aws_key", "mnemonic_phrase",
                "ftp_endpoint", "vpn_endpoint",
            ]
            for ftype in high_value_types:
                if ftype not in type_groups:
                    continue
                items = type_groups[ftype]
                f.write(f"### {ftype} ({len(items)})\n\n")
                for item in items[:20]:
                    f.write(f"- `{item['file']}` -- {item['detail']}\n")
                if len(items) > 20:
                    f.write(f"- ... and {len(items) - 20} more\n")
                f.write("\n")

            other_types = [t for t in type_groups if t not in high_value_types
                           and t not in ("binary_file", "url", "ip_addresses", "vault_credential")]
            for ftype in other_types:
                items = type_groups[ftype]
                f.write(f"### {ftype} ({len(items)})\n\n")
                for item in items[:5]:
                    f.write(f"- `{item['file']}` -- {item['detail']}\n")
                if len(items) > 5:
                    f.write(f"- ... and {len(items) - 5} more\n")
                f.write("\n")

            f.write("---\n\n")

    print(f"Written: {out}")


def generate_access_vectors(findings: list[dict], hosts: dict, critical_rows: list, high_rows: list):
    ssh_keys = []
    vpn_configs = []
    rdp_files = []
    ftp_entries = []
    passwords = []
    kdbx_files = []

    for entry in findings:
        dump_id = entry["dump_id"]
        host = hosts.get(dump_id, {})
        comp = host.get("computer_name", "?")
        user = host.get("username", "?")
        label = f"[{dump_id}] {comp}/{user}"

        for finding in entry["findings"]:
            t = finding["type"]
            d = finding.get("detail", "")
            rec = {"dump_id": dump_id, "host": label, "file": entry["file"], "detail": d}
            if t == "private_key":
                rec["encrypted"] = finding.get("encrypted", True)
                ssh_keys.append(rec)
            elif t == "vpn_endpoint":
                vpn_configs.append(rec)
            elif t == "vpn_auth_user_pass":
                vpn_configs.append(rec)
            elif t == "rdp_password_hash":
                rdp_files.append(rec)
            elif t in ("ftp_endpoint", "filezilla_password"):
                ftp_entries.append(rec)
            elif t in ("password", "winscp_password"):
                passwords.append(rec)

    for row in critical_rows + high_rows:
        ext = row.get("extension", "")
        if ext == ".kdbx":
            dump_id = int(row["dump_id"])
            host = hosts.get(dump_id, {})
            label = f"[{dump_id}] {host.get('computer_name', '?')}/{host.get('username', '?')}"
            kdbx_files.append({
                "dump_id": dump_id,
                "host": label,
                "file": row["relative_path"],
                "detail": f"{row.get('size_bytes', '?')} bytes",
            })
        if ext == ".rdp" and not any(r["file"] == row["relative_path"] and r["dump_id"] == int(row["dump_id"]) for r in rdp_files):
            dump_id = int(row["dump_id"])
            host = hosts.get(dump_id, {})
            label = f"[{dump_id}] {host.get('computer_name', '?')}/{host.get('username', '?')}"
            rdp_files.append({
                "dump_id": dump_id,
                "host": label,
                "file": row["relative_path"],
                "detail": "RDP file found",
            })

    out = DOCS_DIR / "access_vectors.md"
    with open(out, "w", encoding="utf-8") as f:
        f.write("# Access Vectors\n\n")
        f.write("Potential ways to gain access to systems based on dump analysis.\n\n")

        # SSH
        f.write("## 1. SSH Private Keys\n\n")
        unencrypted = [k for k in ssh_keys if not k.get("encrypted")]
        encrypted = [k for k in ssh_keys if k.get("encrypted")]
        f.write(f"Total: **{len(ssh_keys)}** ({len(unencrypted)} unencrypted, {len(encrypted)} encrypted)\n\n")
        if unencrypted:
            f.write("### UNENCRYPTED (immediate access)\n\n")
            f.write("| Dump | Host | File |\n")
            f.write("|------|------|------|\n")
            for k in unencrypted[:100]:
                f.write(f"| {k['dump_id']} | {k['host']} | `{k['file']}` |\n")
            f.write("\n")
        if encrypted:
            f.write(f"### Encrypted ({len(encrypted)} keys, need passphrase)\n\n")
            seen = set()
            for k in encrypted[:50]:
                key = (k["dump_id"], k["file"])
                if key not in seen:
                    seen.add(key)
                    f.write(f"- {k['host']}: `{k['file']}`\n")
            f.write("\n")

        # VPN
        f.write("## 2. VPN Configurations\n\n")
        f.write(f"Total: **{len(vpn_configs)}** entries\n\n")
        if vpn_configs:
            f.write("| Dump | Host | File | Detail |\n")
            f.write("|------|------|------|--------|\n")
            for v in vpn_configs:
                f.write(f"| {v['dump_id']} | {v['host']} | `{v['file']}` | {v['detail']} |\n")
            f.write("\n")

        # RDP
        f.write("## 3. RDP Files\n\n")
        f.write(f"Total: **{len(rdp_files)}** files\n\n")
        if rdp_files:
            f.write("| Dump | Host | File | Detail |\n")
            f.write("|------|------|------|--------|\n")
            for r in rdp_files[:100]:
                f.write(f"| {r['dump_id']} | {r['host']} | `{r['file']}` | {r['detail']} |\n")
            f.write("\n")

        # FTP
        f.write("## 4. FTP / FileZilla\n\n")
        f.write(f"Total: **{len(ftp_entries)}** entries\n\n")
        if ftp_entries:
            f.write("| Dump | Host | File | Detail |\n")
            f.write("|------|------|------|--------|\n")
            for e in ftp_entries[:100]:
                f.write(f"| {e['dump_id']} | {e['host']} | `{e['file']}` | {e['detail'][:80]} |\n")
            f.write("\n")

        # KeePass
        f.write("## 5. KeePass Databases\n\n")
        f.write(f"Total: **{len(kdbx_files)}** files\n\n")
        if kdbx_files:
            f.write("| Dump | Host | File | Size |\n")
            f.write("|------|------|------|------|\n")
            for k in kdbx_files[:100]:
                f.write(f"| {k['dump_id']} | {k['host']} | `{k['file']}` | {k['detail']} |\n")
            f.write("\n")

        # Plaintext passwords
        f.write("## 6. Plaintext Passwords\n\n")
        f.write(f"Total: **{len(passwords)}** entries\n\n")
        if passwords:
            f.write("| Dump | Host | File | Value |\n")
            f.write("|------|------|------|-------|\n")
            for p in passwords:
                f.write(f"| {p['dump_id']} | {p['host']} | `{p['file']}` | `{p['detail']}` |\n")
            f.write("\n")

        # Summary
        f.write("## Summary\n\n")
        f.write("| Vector | Count | Immediate Access |\n")
        f.write("|--------|-------|------------------|\n")
        f.write(f"| SSH Keys | {len(ssh_keys)} | {len(unencrypted)} unencrypted |\n")
        f.write(f"| VPN | {len(vpn_configs)} | check configs |\n")
        f.write(f"| RDP | {len(rdp_files)} | check for saved passwords |\n")
        f.write(f"| FTP/SFTP | {len(ftp_entries)} | {len([e for e in ftp_entries if 'password' in e.get('detail','').lower() or e['detail'].count('@') > 0])} with credentials |\n")
        f.write(f"| KeePass | {len(kdbx_files)} | need master password |\n")
        f.write(f"| Passwords | {len(passwords)} | plaintext |\n")
        f.write("\n")

    print(f"Written: {out}")


def run():
    hosts = load_hosts()
    findings = load_findings()
    critical_rows = load_triage(CRITICAL_CSV)
    high_rows = load_triage(HIGH_CSV)

    generate_critical_report(findings, hosts)
    generate_access_vectors(findings, hosts, critical_rows, high_rows)


if __name__ == "__main__":
    run()
