"""
Corporate access detector: deep-parses RDP files, PuTTY/WinSCP registry
exports, vault credentials, and FileZilla XMLs. Aggregates per host,
scores corporate vs personal access, and produces a prioritized report.

Produces:
  - data/corporate_access.json
  - docs/corporate_access_report.md
"""

import base64
import csv
import ipaddress
import json
import re
import socket
import time
import xml.etree.ElementTree as ET
from collections import defaultdict
from pathlib import Path
from urllib.parse import unquote

ROOT_DIR = Path(__file__).resolve().parents[2]
LEGACY_ROOT = ROOT_DIR / "findings" / "legacy_dump_assessment"
DUMPS_DIR = LEGACY_ROOT / "dumps"
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"

OUTPUT_JSON = DATA_DIR / "legacy_corporate_access.json"
OUTPUT_REPORT = DOCS_DIR / "corporate_access_report.md"

GAME_NOISE = re.compile(
    r'apexhosting|akliz\.net|minehut|shockbyte|bisect\.?hosting'
    r'|\.?minecraft|\.?terraria|\.?valheim|nodecraft'
    r'|pebblehost|exaroton|server\.pro|craftserve'
    r'|\.?vultam\.net|nitrado',
    re.I,
)

COMMERCIAL_VPN = re.compile(
    r'vpntype\.(com|dev)|nordvpn|expressvpn|surfshark|protonvpn'
    r'|mullvad|privateinternetaccess|cyberghost',
    re.I,
)

OVPN_TEMPLATE_HOSTS = {"my-server-1", "my-server-2"}


# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------

def resolve_dump_path(dump_id: int, relative_path: str) -> Path:
    folder = "dump" if dump_id == 0 else f"dump ({dump_id})"
    return DUMPS_DIR / folder / relative_path


def read_text_file(path: Path, max_size: int = 5 * 1024 * 1024) -> str | None:
    if not path.exists():
        return None
    try:
        size = path.stat().st_size
        if size > max_size or size == 0:
            return None
        raw = path.read_bytes()
        if raw[:2] in (b'\xff\xfe', b'\xfe\xff'):
            return raw.decode("utf-16", errors="replace")
        if b'\x00' in raw[:512]:
            return None
        return raw.decode("utf-8", errors="replace")
    except Exception:
        return None


def classify_ip(ip_str: str) -> str:
    try:
        ip = ipaddress.ip_address(ip_str)
    except ValueError:
        return "hostname"
    if ip.is_loopback:
        return "loopback"
    if ip in ipaddress.ip_network('10.0.0.0/8'):
        return "corporate_10x"
    if ip in ipaddress.ip_network('172.16.0.0/12'):
        return "corporate_172x"
    if ip in ipaddress.ip_network('192.168.0.0/16'):
        return "private_192x"
    if ip.is_private:
        return "private_other"
    return "public"


def is_game_host(hostname: str) -> bool:
    return bool(GAME_NOISE.search(hostname))


def is_commercial_vpn(endpoint: str) -> bool:
    return bool(COMMERCIAL_VPN.search(endpoint))


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_triage_rows(*csv_paths: Path) -> list[dict]:
    rows = []
    for p in csv_paths:
        if not p.exists():
            continue
        with open(p, encoding="utf-8") as f:
            rows.extend(csv.DictReader(f))
    return rows


# ---------------------------------------------------------------------------
# WHOIS Enrichment
# ---------------------------------------------------------------------------

HOSTING_PATTERNS = re.compile(
    r'hetzner|ovh\b|ovhcloud|digitalocean|vultr|linode|akamai.connected'
    r'|amazon|aws\b|google.cloud|microsoft.azure|azure\b'
    r'|oracle.cloud|oracle-bmc|contabo|scaleway|online\.s\.a\.s'
    r'|online\.net|online.sas|leaseweb|hostgator|godaddy|hostinger|ionos'
    r'|choopa|the.constant.company|m247|psychz|colocrossing'
    r'|hostwinds|interserver|liquidweb|rackspace|softlayer'
    r'|packet\.net|equinix|clouvider|datacamp|quickpacket'
    r'|phoenixnap|servermania|kamatera|upcloud|cherry.?servers'
    r'|heficed|stark.industries|combahton|netcup|strato'
    r'|1\&1|aruba\.it|fasthosts|host ?europe|pulsant|ukfast'
    r'|limestone|singlehop|fdcservers|quadranet|cogent'
    r'|zenlayer|buyvm|frantech|privatelayer|namecheap'
    r'|dreamhost|bluehost|siteground|a2.hosting|inmotion'
    r'|greengeeks|hostpapa|ipage|fatcow|justhost'
    r'|smartape|internap|unitas.global|sparked'
    r'|datapacket|serverius|i3d\.net|path\.net|cogeco.peer'
    r'|tzulo|100tb|serverpoint|liteserver|trabia'
    r'|ihor\.net|maxihost|sharktech|ramnode|virtuo\.?za',
    re.I,
)

ISP_PATTERNS = re.compile(
    r'comcast|xfinity|at&t|att\b|at\.t\b|verizon|spectrum|charter'
    r'|cox\b|centurylink|lumen|frontier|t-mobile|sprint'
    r'|deutsche.telekom|telekom\.de|vodafone|orange\b|bouygues'
    r'|free\.fr|sfr\b|bt\b|british.telecom|sky.uk|virgin.media'
    r'|talktalk|plusnet|telefonica|movistar|o2\b'
    r'|rostelecom|ростелеком|mts\b|мтс|beeline|билайн|megafon|мегафон'
    r'|telia\b|swisscom|proximus|kpn\b|ziggo|telenor|tele2'
    r'|elisa\b|dna\b|sonera|rogers\b|bell.canada|telus|shaw\b'
    r'|videotron|optus|telstra|tpg\b|nbn\b|singtel|starhub'
    r'|bharti|airtel|jio\b|reliance|bsnl|idea.cellular'
    r'|suddenlink|mediacom|windstream|consolidated.comm'
    r'|altice|optimum|cablevision|rcn\b|wow\b|wide.open.west'
    r'|buckeye.broadband|atlantic.broadband|midco\b|grande.comm',
    re.I,
)

EDUCATION_PATTERNS = re.compile(
    r'universit|college|\.edu\b|academ|school.of|institute.of'
    r'|polytechnic|\.ac\.|research.network|\.ren\b|geant\b|internet2'
    r'|jisc\b|surf\.net|garr\.it|dfn\.de|renater',
    re.I,
)

GOVERNMENT_PATTERNS = re.compile(
    r'\.gov\b|government|federal|military|\.mil\b|department.of'
    r'|ministry|state.of|city.of|county.of',
    re.I,
)


def collect_public_endpoints(rdp_data: dict, reg_data: dict,
                             vault_data: dict, fz_data: dict,
                             vpn_profiles: dict) -> tuple[set, set]:
    """Collect unique public IPs and unresolved hostnames from all parsed data."""
    public_ips = set()
    hostnames = set()

    def _triage(host_str: str):
        host_str = host_str.split(":")[0].strip()
        if not host_str:
            return
        cl = classify_ip(host_str)
        if cl == "public":
            public_ips.add(host_str)
        elif cl == "hostname":
            if not is_game_host(host_str) and "." in host_str:
                hostnames.add(host_str)

    for entries in rdp_data.values():
        for e in entries:
            _triage(e.get("target", ""))
            _triage(e.get("gateway", ""))

    for data in reg_data.values():
        for s in data.get("putty_sessions", []):
            _triage(s["hostname"])
        for s in data.get("winscp_sessions", []):
            _triage(s["hostname"])

    for entries in vault_data.values():
        for c in entries:
            _triage(c.get("target_host", ""))

    for entries in fz_data.values():
        for fz in entries:
            _triage(fz["host"])

    for dump_id, profile in vpn_profiles.items():
        for ep in profile:
            _triage(ep.split(":")[0])

    return public_ips, hostnames


def dns_resolve_hostnames(hostnames: set) -> dict[str, str | None]:
    """Resolve hostnames to IPs via socket.getaddrinfo(). Returns {hostname: ip|None}."""
    resolved = {}
    for h in sorted(hostnames):
        try:
            results = socket.getaddrinfo(h, None, socket.AF_INET,
                                         socket.SOCK_STREAM)
            if results:
                ip = results[0][4][0]
                resolved[h] = ip
            else:
                resolved[h] = None
        except (socket.gaierror, socket.timeout, OSError):
            resolved[h] = None
    return resolved


def bulk_cymru_whois(ips: list[str]) -> dict[str, dict]:
    """Query Team Cymru bulk WHOIS for ASN/org info. Single TCP connection."""
    if not ips:
        return {}
    query = "begin\nverbose\n" + "\n".join(ips) + "\nend\n"
    try:
        with socket.create_connection(("whois.cymru.com", 43), timeout=30) as s:
            s.sendall(query.encode())
            data = b""
            while True:
                chunk = s.recv(8192)
                if not chunk:
                    break
                data += chunk
    except (socket.timeout, OSError) as exc:
        print(f"  WARNING: Cymru WHOIS failed: {exc}")
        return {}

    results = {}
    for line in data.decode("utf-8", errors="replace").splitlines():
        line = line.strip()
        if not line or line.startswith("Bulk") or line.startswith("AS"):
            continue
        parts = [p.strip() for p in line.split("|")]
        if len(parts) >= 7:
            results[parts[1]] = {
                "asn": parts[0],
                "prefix": parts[2],
                "country": parts[3],
                "registry": parts[4],
                "allocated": parts[5],
                "org_name": parts[6],
            }
    return results


def classify_by_org(org_name: str) -> str:
    """Heuristic classification of ASN organization into ip_type."""
    if not org_name or org_name == "NA":
        return "unknown"
    if EDUCATION_PATTERNS.search(org_name):
        return "education"
    if GOVERNMENT_PATTERNS.search(org_name):
        return "government"
    if HOSTING_PATTERNS.search(org_name):
        return "hosting"
    if ISP_PATTERNS.search(org_name):
        return "residential"
    return "corporate"


def enrich_whois(rdp_data: dict, reg_data: dict, vault_data: dict,
                 fz_data: dict, extracted: list[dict]
                 ) -> tuple[dict[str, dict], dict[str, str | None]]:
    """Run full WHOIS enrichment pipeline. Returns (whois_db, dns_map)."""

    vpn_profiles: dict[int, list[str]] = defaultdict(list)
    for entry in extracted:
        for f in entry.get("findings", []):
            if f["type"] == "vpn_endpoint":
                vpn_profiles[entry["dump_id"]].append(f.get("detail", ""))

    print("Collecting public endpoints...")
    public_ips, hostnames = collect_public_endpoints(
        rdp_data, reg_data, vault_data, fz_data, dict(vpn_profiles))
    print(f"  {len(public_ips)} unique public IPs, "
          f"{len(hostnames)} hostnames to resolve")

    print("Resolving hostnames via DNS...")
    dns_map = dns_resolve_hostnames(hostnames)
    resolved_count = sum(1 for v in dns_map.values() if v)
    print(f"  {resolved_count}/{len(hostnames)} resolved")

    for ip in dns_map.values():
        if ip and classify_ip(ip) == "public":
            public_ips.add(ip)

    all_ips = sorted(public_ips)
    print(f"Querying Team Cymru WHOIS for {len(all_ips)} IPs...")
    whois_db = bulk_cymru_whois(all_ips)
    print(f"  {len(whois_db)} results received")

    for ip, info in whois_db.items():
        info["ip_type"] = classify_by_org(info.get("org_name", ""))

    type_counts: dict[str, int] = defaultdict(int)
    for info in whois_db.values():
        type_counts[info["ip_type"]] += 1
    for t, c in sorted(type_counts.items()):
        print(f"    {t}: {c}")

    return whois_db, dns_map


def get_whois_for_host(host_str: str, whois_db: dict,
                       dns_map: dict) -> dict | None:
    """Look up WHOIS info for a host string (IP or hostname)."""
    host_str = host_str.split(":")[0].strip()
    if not host_str:
        return None

    cl = classify_ip(host_str)
    if cl == "public":
        return whois_db.get(host_str)
    elif cl == "hostname":
        resolved_ip = dns_map.get(host_str)
        if resolved_ip:
            info = whois_db.get(resolved_ip)
            if info:
                return {**info, "resolved_ip": resolved_ip}
    return None


# ---------------------------------------------------------------------------
# Phase 1: Deep Parsers
# ---------------------------------------------------------------------------

def parse_rdp_files(triage_rows: list[dict]) -> dict[int, list[dict]]:
    """Parse .rdp files, extract target address/username/gateway."""
    results = defaultdict(list)
    seen = set()

    for row in triage_rows:
        if row.get("extension", "") != ".rdp":
            continue
        dump_id = int(row["dump_id"])
        rel = row["relative_path"]
        key = (dump_id, rel)
        if key in seen:
            continue
        seen.add(key)

        path = resolve_dump_path(dump_id, rel)
        text = read_text_file(path)
        if not text:
            continue

        entry = {
            "file": rel,
            "filename": row.get("filename", ""),
            "target": "",
            "username": "",
            "gateway": "",
        }

        for line in text.splitlines():
            line = line.strip()
            lower = line.lower()
            # RDP format: "key name:type_char:value"
            parts = line.split(":", 2)
            if len(parts) < 3:
                continue
            key = parts[0].strip().lower()
            val = parts[2].strip()
            if key == "full address" and val:
                entry["target"] = val
            elif key == "username" and val:
                entry["username"] = val
            elif key == "gatewayhostname" and val:
                entry["gateway"] = val

        if entry["target"]:
            target_for_ip = entry["target"].split(":")[0]
            entry["ip_class"] = classify_ip(target_for_ip)
            dedup_key = (dump_id, entry["filename"], entry["target"])
            if dedup_key not in seen:
                seen.add(dedup_key)
            results[dump_id].append(entry)

    for dump_id in results:
        seen_rdp = set()
        unique = []
        for r in results[dump_id]:
            key = (r["filename"], r["target"], r["username"])
            if key not in seen_rdp:
                seen_rdp.add(key)
                unique.append(r)
        results[dump_id] = unique

    return dict(results)


def parse_registry_files(triage_rows: list[dict]) -> dict[int, dict]:
    """Parse .reg files for PuTTY and WinSCP sessions."""
    results = defaultdict(lambda: {"putty_sessions": [], "winscp_sessions": []})
    seen = set()

    for row in triage_rows:
        if row.get("extension", "") != ".reg":
            continue
        dump_id = int(row["dump_id"])
        rel = row["relative_path"]
        key = (dump_id, rel)
        if key in seen:
            continue
        seen.add(key)

        path = resolve_dump_path(dump_id, rel)
        text = read_text_file(path, max_size=2 * 1024 * 1024)
        if not text:
            continue

        cat = row.get("category", "")
        fname_lower = row.get("filename", "").lower()

        if "putty" in cat or "putty" in fname_lower:
            _parse_putty_reg(text, dump_id, rel, results)
        if "winscp" in cat or "winscp" in fname_lower:
            _parse_winscp_reg(text, dump_id, rel, results)

    return {k: dict(v) for k, v in results.items()}


def _parse_putty_reg(text: str, dump_id: int, rel: str,
                     results: dict):
    current_session = None
    session_data = {}

    for line in text.splitlines():
        line = line.strip()
        section_match = re.match(
            r'\[HKEY_CURRENT_USER\\Software\\SimonTatham\\PuTTY\\Sessions\\(.+?)\]',
            line, re.I,
        )
        if section_match:
            if current_session and session_data.get("HostName"):
                _emit_putty_session(current_session, session_data,
                                    dump_id, rel, results)
            current_session = unquote(section_match.group(1))
            session_data = {}
            continue

        m = re.match(r'"(\w+)"="(.*)"', line)
        if m:
            session_data[m.group(1)] = m.group(2)
            continue
        m = re.match(r'"(\w+)"=dword:([0-9a-fA-F]+)', line)
        if m:
            session_data[m.group(1)] = int(m.group(2), 16)

    if current_session and session_data.get("HostName"):
        _emit_putty_session(current_session, session_data,
                            dump_id, rel, results)


def _emit_putty_session(name: str, data: dict, dump_id: int,
                        rel: str, results: dict):
    hostname = data.get("HostName", "")
    if not hostname:
        return
    port = data.get("PortNumber", 22)
    if isinstance(port, str):
        port = 22
    username = data.get("UserName", "")
    protocol = data.get("Protocol", "ssh")
    pubkey = data.get("PublicKeyFile", "")

    results[dump_id]["putty_sessions"].append({
        "name": name,
        "hostname": hostname,
        "port": port,
        "username": username,
        "protocol": protocol,
        "has_pubkey": bool(pubkey),
        "file": rel,
        "ip_class": classify_ip(hostname),
        "is_game": is_game_host(hostname),
    })


def _parse_winscp_reg(text: str, dump_id: int, rel: str,
                      results: dict):
    current_session = None
    session_data = {}

    for line in text.splitlines():
        line = line.strip()
        section_match = re.match(
            r'\[HKEY_CURRENT_USER\\Software\\Martin Prikryl\\WinSCP 2\\Sessions\\(.+?)\]',
            line, re.I,
        )
        if section_match:
            if current_session and session_data.get("HostName"):
                _emit_winscp_session(current_session, session_data,
                                     dump_id, rel, results)
            current_session = unquote(section_match.group(1))
            session_data = {}
            continue

        m = re.match(r'"(\w+)"="(.*)"', line)
        if m:
            session_data[m.group(1)] = m.group(2)

    if current_session and session_data.get("HostName"):
        _emit_winscp_session(current_session, session_data,
                             dump_id, rel, results)


def _emit_winscp_session(name: str, data: dict, dump_id: int,
                         rel: str, results: dict):
    hostname = data.get("HostName", "")
    if not hostname:
        return
    username = data.get("UserName", "")
    has_password = bool(data.get("Password", ""))

    results[dump_id]["winscp_sessions"].append({
        "name": name,
        "hostname": hostname,
        "username": username,
        "has_password": has_password,
        "file": rel,
        "ip_class": classify_ip(hostname),
        "is_game": is_game_host(hostname),
    })


def parse_vault_credentials(extracted: list[dict]) -> dict[int, list[dict]]:
    """Deep-parse vault_credential detail strings into structured records."""
    results = defaultdict(list)

    domain_re = re.compile(
        r'target=(Domain):target=(?:TERMSRV/)?(.+?)\s+user=(.+)',
        re.I,
    )
    legacy_re = re.compile(
        r'target=(LegacyGeneric):target=(.+?)\s+user=(.+)',
        re.I,
    )

    for entry in extracted:
        dump_id = entry["dump_id"]
        for finding in entry.get("findings", []):
            if finding.get("type") != "vault_credential":
                continue
            detail = finding.get("detail", "")

            m = domain_re.match(detail)
            if m:
                cred_type = m.group(1)
                target_raw = m.group(2)
                user_raw = m.group(3)

                is_termsrv = "TERMSRV/" in detail
                domain_name = ""
                username = user_raw
                if "\\" in user_raw:
                    domain_name, username = user_raw.split("\\", 1)

                rec = {
                    "cred_type": cred_type,
                    "target_host": target_raw,
                    "is_termsrv": is_termsrv,
                    "domain_name": domain_name,
                    "username": username,
                    "ip_class": classify_ip(target_raw),
                    "raw": detail,
                }
                results[dump_id].append(rec)
                continue

            m = legacy_re.match(detail)
            if m:
                target_info = m.group(2)
                user_raw = m.group(3)
                indicators = []
                if "SSMS" in target_info or "Microsoft:SSMS" in target_info:
                    indicators.append("sql_server")
                if "autodiscover" in target_info.lower():
                    indicators.append("exchange")
                if ".edu" in target_info.lower():
                    indicators.append("academic")
                if indicators:
                    results[dump_id].append({
                        "cred_type": "LegacyGeneric",
                        "target_host": target_info,
                        "is_termsrv": False,
                        "domain_name": "",
                        "username": user_raw,
                        "ip_class": "service",
                        "indicators": indicators,
                        "raw": detail,
                    })

    return dict(results)


def parse_filezilla_xmls(triage_rows: list[dict]) -> dict[int, list[dict]]:
    """Parse FileZilla sitemanager/recentservers XML for structured site data."""
    results = defaultdict(list)
    seen = set()

    for row in triage_rows:
        fname = row.get("filename", "").lower()
        if not (fname.endswith(".xml") and
                ("filezilla" in row.get("category", "") or
                 "filezilla" in fname or
                 "sitemanager" in fname or
                 "recentservers" in fname)):
            continue

        dump_id = int(row["dump_id"])
        rel = row["relative_path"]
        key = (dump_id, rel)
        if key in seen:
            continue
        seen.add(key)

        path = resolve_dump_path(dump_id, rel)
        text = read_text_file(path)
        if not text:
            continue

        try:
            root = ET.fromstring(text)
        except ET.ParseError:
            continue

        for server in root.iter("Server"):
            host_el = server.find("Host")
            if host_el is None or not host_el.text:
                continue
            host = host_el.text.strip()
            port = (server.findtext("Port") or "21").strip()
            user = (server.findtext("User") or "").strip()
            protocol_num = (server.findtext("Protocol") or "0").strip()
            protocol = "SFTP" if protocol_num == "1" else "FTP"
            site_name = (server.findtext("Name") or "").strip()

            pass_el = server.find("Pass")
            has_password = False
            password_b64 = ""
            if pass_el is not None and pass_el.text:
                password_b64 = pass_el.text.strip()
                has_password = bool(password_b64)

            results[dump_id].append({
                "host": host,
                "port": port,
                "user": user,
                "protocol": protocol,
                "site_name": site_name,
                "has_password": has_password,
                "file": rel,
                "ip_class": classify_ip(host),
                "is_game": is_game_host(host),
            })

    return dict(results)


# ---------------------------------------------------------------------------
# Phase 2: Aggregation
# ---------------------------------------------------------------------------

def aggregate(hosts: dict[int, dict],
              extracted: list[dict],
              rdp_data: dict,
              reg_data: dict,
              vault_data: dict,
              fz_data: dict,
              triage_rows: list[dict],
              whois_db: dict | None = None,
              dns_map: dict | None = None) -> dict[int, dict]:
    """Build per-host aggregated profile."""
    whois_db = whois_db or {}
    dns_map = dns_map or {}
    profiles = {}

    all_dump_ids = set(hosts.keys())
    for src in (rdp_data, reg_data, vault_data, fz_data):
        all_dump_ids.update(src.keys())
    for entry in extracted:
        all_dump_ids.add(entry["dump_id"])

    def _attach_whois(entry: dict, host_key: str):
        w = get_whois_for_host(entry.get(host_key, ""), whois_db, dns_map)
        if w:
            entry["whois"] = w

    for dump_id in sorted(all_dump_ids):
        host = hosts.get(dump_id, {})
        rdp_conns = [dict(r) for r in rdp_data.get(dump_id, [])]
        putty_sess = [dict(s) for s in reg_data.get(dump_id, {}).get("putty_sessions", [])]
        winscp_sess = [dict(s) for s in reg_data.get(dump_id, {}).get("winscp_sessions", [])]
        vault_creds = [dict(c) for c in vault_data.get(dump_id, [])]
        fz_sites = [dict(fz) for fz in fz_data.get(dump_id, [])]

        for r in rdp_conns:
            _attach_whois(r, "target")
        for s in putty_sess:
            _attach_whois(s, "hostname")
        for s in winscp_sess:
            _attach_whois(s, "hostname")
        for c in vault_creds:
            _attach_whois(c, "target_host")
        for fz in fz_sites:
            _attach_whois(fz, "host")

        profile = {
            "dump_id": dump_id,
            "computer_name": host.get("computer_name", "?"),
            "username": host.get("username", "?"),
            "scan_date": host.get("scan_date", ""),
            "rdp_connections": rdp_conns,
            "putty_sessions": putty_sess,
            "winscp_sessions": winscp_sess,
            "vault_domain_creds": vault_creds,
            "filezilla_sites": fz_sites,
            "vpn_endpoints": [],
            "keepass_dbs": [],
            "ssh_keys_unencrypted": 0,
            "ssh_keys_encrypted": 0,
            "server_certs": [],
            "has_ssms": False,
            "has_exchange": False,
            "has_edu": False,
            "has_server_ovpn": False,
            "has_jira_cert": False,
        }

        for entry in extracted:
            if entry["dump_id"] != dump_id:
                continue
            fname = entry.get("filename", "").lower()
            for f in entry.get("findings", []):
                ft = f["type"]
                det = f.get("detail", "")
                if ft == "vpn_endpoint":
                    if det.split(":")[0] not in OVPN_TEMPLATE_HOSTS:
                        vpn_entry = {
                            "endpoint": det,
                            "file": entry["file"],
                            "is_commercial": is_commercial_vpn(det),
                        }
                        w = get_whois_for_host(det, whois_db, dns_map)
                        if w:
                            vpn_entry["whois"] = w
                        profile["vpn_endpoints"].append(vpn_entry)
                elif ft == "vpn_auth_user_pass":
                    for vpn in profile["vpn_endpoints"]:
                        if vpn["file"] == entry["file"]:
                            vpn["has_auth_user_pass"] = True
                elif ft == "private_key":
                    if det == "UNENCRYPTED":
                        profile["ssh_keys_unencrypted"] += 1
                    else:
                        profile["ssh_keys_encrypted"] += 1

            if "server.ovpn" in fname:
                profile["has_server_ovpn"] = True
            if "jira" in fname or "confluence" in fname:
                profile["has_jira_cert"] = True
            if "server" in fname and any(
                fname.endswith(ext) for ext in (".key", ".pem")
            ):
                profile["server_certs"].append(entry["file"])

        for cred in profile["vault_domain_creds"]:
            for ind in cred.get("indicators", []):
                if ind == "sql_server":
                    profile["has_ssms"] = True
                elif ind == "exchange":
                    profile["has_exchange"] = True
                elif ind == "academic":
                    profile["has_edu"] = True
            target = cred.get("target_host", "").lower()
            if ".edu" in target:
                profile["has_edu"] = True
            if "autodiscover" in target:
                profile["has_exchange"] = True

        for row in triage_rows:
            if int(row["dump_id"]) != dump_id:
                continue
            ext = row.get("extension", "")
            if ext == ".kdbx":
                profile["keepass_dbs"].append({
                    "file": row["relative_path"],
                    "size": int(row.get("size_bytes", 0)),
                })

        dedupe_vpn(profile)
        dedupe_filezilla(profile)
        profiles[dump_id] = profile

    return profiles


def dedupe_vpn(profile: dict):
    seen = set()
    unique = []
    for v in profile["vpn_endpoints"]:
        ep = v["endpoint"]
        if ep not in seen:
            seen.add(ep)
            unique.append(v)
    profile["vpn_endpoints"] = unique


def dedupe_filezilla(profile: dict):
    seen = set()
    unique = []
    for fz in profile["filezilla_sites"]:
        key = (fz["host"], fz["port"], fz["user"], fz["protocol"])
        if key not in seen:
            seen.add(key)
            unique.append(fz)
    profile["filezilla_sites"] = unique


# ---------------------------------------------------------------------------
# Phase 3: Corporate Scoring
# ---------------------------------------------------------------------------

def _whois_tag(entry: dict) -> str:
    """Format a short WHOIS annotation for reason strings."""
    w = entry.get("whois")
    if not w:
        return ""
    ip_type = w.get("ip_type", "")
    org = w.get("org_name", "")
    cc = w.get("country", "")
    parts = []
    if ip_type:
        parts.append(ip_type)
    if org:
        short = org[:40] + "..." if len(org) > 40 else org
        parts.append(short)
    if cc:
        parts.append(cc)
    return f" [{', '.join(parts)}]" if parts else ""


def _ip_type_of(entry: dict) -> str:
    w = entry.get("whois")
    return w.get("ip_type", "") if w else ""


def score_host(profile: dict) -> tuple[int, str, list[str]]:
    """Return (score, classification, list_of_reasons)."""
    score = 0
    reasons = []

    IGNORED_DOMAINS = {"MicrosoftAccount", "MicrosoftAccount ", ""}
    for cred in profile["vault_domain_creds"]:
        domain = cred.get("domain_name", "")
        if domain and domain not in IGNORED_DOMAINS:
            hostname = profile.get("computer_name", "")
            if domain != hostname:
                score += 50
                reasons.append(
                    f"AD domain '{domain}' credential: "
                    f"{domain}\\{cred['username']} -> {cred['target_host']}"
                )

        if cred.get("is_termsrv"):
            cred_domain = cred.get("domain_name", "")
            if cred_domain == "MicrosoftAccount":
                continue
            target = cred["target_host"]
            ip_cl = cred.get("ip_class", "")
            user = cred.get("username", "")
            tag = _whois_tag(cred)
            ipt = _ip_type_of(cred)
            if ip_cl in ("corporate_10x", "corporate_172x"):
                score += 40
                reasons.append(f"RDP to corporate subnet {target} ({ip_cl})")
            elif ip_cl == "public":
                if ipt == "corporate":
                    score += 35
                    reasons.append(f"RDP to corporate IP {target} as {user}{tag}")
                elif ipt == "education":
                    score += 30
                    reasons.append(f"RDP to education IP {target} as {user}{tag}")
                elif user.lower() in ("administrator", "admin", "root"):
                    score += 35
                    reasons.append(f"RDP to public IP {target} as {user}{tag}")
                elif ipt == "hosting":
                    score += 15
                    reasons.append(f"RDP to hosting IP {target} as {user}{tag}")
                else:
                    score += 15
                    reasons.append(f"RDP to public IP {target} as {user}{tag}")
            elif ip_cl == "hostname":
                host_lower = target.lower()
                if "dc" in host_lower or "domain" in host_lower:
                    score += 45
                    reasons.append(
                        f"RDP to likely Domain Controller '{target}'"
                    )
                elif not is_game_host(target):
                    score += 20
                    reasons.append(f"RDP to named host '{target}'{tag}")

    if profile["has_ssms"]:
        score += 45
        reasons.append("SQL Server Management Studio (SSMS) credential found")

    if profile["has_exchange"]:
        score += 35
        reasons.append("Exchange/autodiscover credential found")

    if profile["has_edu"]:
        score += 30
        reasons.append("Academic (.edu) domain credential found")

    seen_rdp_targets = set()
    for rdp in profile["rdp_connections"]:
        rdp_key = (rdp.get("filename", ""), rdp.get("target", ""))
        if rdp_key in seen_rdp_targets:
            continue
        seen_rdp_targets.add(rdp_key)
        fname = rdp.get("filename", "")
        if fname and fname.lower() not in ("default.rdp",):
            score += 30
            reasons.append(f"Named RDP file: '{fname}' -> {rdp['target']}")
        if rdp.get("gateway"):
            score += 40
            reasons.append(f"RD Gateway configured: {rdp['gateway']}")
        target = rdp.get("target", "")
        ip_cl = rdp.get("ip_class", "")
        if target and ip_cl in ("corporate_10x", "corporate_172x"):
            score += 25
            reasons.append(
                f"RDP file targets corporate subnet: {target} ({ip_cl})"
            )

    for vpn in profile["vpn_endpoints"]:
        ep = vpn["endpoint"]
        if vpn.get("is_commercial"):
            continue
        ep_host = ep.split(":")[0]
        tag = _whois_tag(vpn)
        ipt = _ip_type_of(vpn)
        if ep_host.startswith("vpn.") or ep_host.startswith("openvpn."):
            score += 40
            reasons.append(f"Corporate VPN endpoint: {ep}{tag}")
        elif not is_game_host(ep_host):
            ip_cl = classify_ip(ep_host)
            if ip_cl == "public":
                if ipt == "corporate":
                    score += 30
                    reasons.append(f"VPN to corporate server: {ep}{tag}")
                else:
                    score += 15
                    reasons.append(f"VPN to public server: {ep}{tag}")
            elif ip_cl == "hostname":
                score += 20
                reasons.append(f"VPN to named host: {ep}{tag}")

    if profile["has_server_ovpn"]:
        score += 25
        reasons.append("Server-side OpenVPN config (server.ovpn) found")

    if profile["has_jira_cert"]:
        score += 25
        reasons.append("Jira/Confluence certificate found")

    for sess in profile["putty_sessions"]:
        if sess["is_game"]:
            score -= 5
            continue
        hostname = sess["hostname"]
        ip_cl = sess["ip_class"]
        tag = _whois_tag(sess)
        ipt = _ip_type_of(sess)
        if ip_cl in ("corporate_10x", "corporate_172x"):
            score += 25
            reasons.append(
                f"PuTTY session '{sess['name']}' to corporate subnet "
                f"{hostname}"
            )
        elif ip_cl == "public":
            if ipt == "corporate":
                score += 25
                reasons.append(
                    f"PuTTY session '{sess['name']}' to corporate IP "
                    f"{hostname}{tag}"
                )
            elif ipt == "education":
                score += 25
                reasons.append(
                    f"PuTTY session '{sess['name']}' to education IP "
                    f"{hostname}{tag}"
                )
            else:
                score += 15
                reasons.append(
                    f"PuTTY session '{sess['name']}' to {hostname}{tag}"
                )
        elif ip_cl == "hostname" and not is_game_host(hostname):
            score += 15
            reasons.append(
                f"PuTTY session '{sess['name']}' to {hostname}{tag}"
            )

    for sess in profile["winscp_sessions"]:
        if sess["is_game"]:
            score -= 5
            continue
        hostname = sess["hostname"]
        ip_cl = sess["ip_class"]
        pw_note = " (with saved password)" if sess["has_password"] else ""
        tag = _whois_tag(sess)
        ipt = _ip_type_of(sess)
        if ip_cl in ("corporate_10x", "corporate_172x"):
            score += 25
            reasons.append(
                f"WinSCP session to corporate subnet {hostname}{pw_note}"
            )
        elif ip_cl == "public":
            if ipt == "corporate":
                score += 25
                reasons.append(
                    f"WinSCP session to corporate IP {hostname}{pw_note}{tag}"
                )
            elif ipt == "education":
                score += 25
                reasons.append(
                    f"WinSCP session to education IP {hostname}{pw_note}{tag}"
                )
            else:
                score += 15
                reasons.append(f"WinSCP session to {hostname}{pw_note}{tag}")
        elif ip_cl == "hostname" and not is_game_host(hostname):
            score += 15
            reasons.append(f"WinSCP session to {hostname}{pw_note}{tag}")

    seen_fz_hosts = set()
    for fz in profile["filezilla_sites"]:
        if fz["is_game"]:
            continue
        fz_key = (fz["host"], fz["port"], fz["user"])
        if fz_key in seen_fz_hosts:
            continue
        seen_fz_hosts.add(fz_key)
        if fz["has_password"]:
            ip_cl = fz["ip_class"]
            tag = _whois_tag(fz)
            ipt = _ip_type_of(fz)
            if ip_cl in ("corporate_10x", "corporate_172x"):
                score += 20
                reasons.append(
                    f"FileZilla site '{fz['site_name']}' to "
                    f"{fz['host']}:{fz['port']} with password"
                )
            elif ip_cl == "public" or ip_cl == "hostname":
                if ipt == "corporate":
                    score += 15
                    reasons.append(
                        f"FileZilla site to corporate {fz['host']}:{fz['port']} "
                        f"(user: {fz['user']}){tag}"
                    )
                else:
                    score += 10
                    reasons.append(
                        f"FileZilla site to {fz['host']}:{fz['port']} "
                        f"(user: {fz['user']}){tag}"
                    )

    big_keepass = [k for k in profile["keepass_dbs"] if k["size"] > 10000]
    if big_keepass:
        score += 10
        reasons.append(
            f"{len(big_keepass)} KeePass db(s) > 10KB "
            f"(largest: {max(k['size'] for k in big_keepass)} bytes)"
        )

    if score >= 50:
        classification = "CORPORATE"
    elif score >= 25:
        classification = "INFRASTRUCTURE"
    elif score >= 10:
        classification = "POSSIBLE"
    else:
        classification = "PERSONAL"

    return score, classification, reasons


# ---------------------------------------------------------------------------
# Phase 4: Output
# ---------------------------------------------------------------------------

def generate_json(scored: list[dict]):
    with open(OUTPUT_JSON, "w", encoding="utf-8") as f:
        json.dump(scored, f, indent=2, ensure_ascii=False)
    print(f"Written: {OUTPUT_JSON}")


def generate_report(scored: list[dict], whois_db: dict | None = None):
    whois_db = whois_db or {}
    tiers = defaultdict(list)
    for item in scored:
        tiers[item["classification"]].append(item)

    DOCS_DIR.mkdir(exist_ok=True)
    with open(OUTPUT_REPORT, "w", encoding="utf-8") as f:
        f.write("# Corporate Access Detection Report\n\n")

        total = len(scored)
        f.write("## Summary\n\n")
        f.write(f"Hosts analyzed: **{total}** (out of those with any score > 0)\n\n")
        f.write("| Classification | Count |\n")
        f.write("|----------------|-------|\n")
        for tier in ("CORPORATE", "INFRASTRUCTURE", "POSSIBLE", "PERSONAL"):
            cnt = len(tiers.get(tier, []))
            if cnt:
                f.write(f"| **{tier}** | {cnt} |\n")
        f.write("\n")

        if whois_db:
            f.write("## WHOIS IP Intelligence\n\n")
            type_groups: dict[str, list] = defaultdict(list)
            for ip, info in sorted(whois_db.items()):
                type_groups[info.get("ip_type", "unknown")].append((ip, info))

            f.write(f"**{len(whois_db)}** unique public IPs enriched via "
                    f"Team Cymru WHOIS\n\n")
            f.write("| IP Type | Count |\n")
            f.write("|---------|-------|\n")
            for t in ("hosting", "residential", "corporate",
                      "education", "government", "unknown"):
                if t in type_groups:
                    f.write(f"| **{t}** | {len(type_groups[t])} |\n")
            f.write("\n")

            f.write("| IP | ASN | Organization | Country | Type |\n")
            f.write("|----|-----|--------------|---------|------|\n")
            for ip, info in sorted(whois_db.items(),
                                   key=lambda x: x[1].get("ip_type", "")):
                f.write(
                    f"| {ip} | AS{info.get('asn', '?')} | "
                    f"{info.get('org_name', '?')} | "
                    f"{info.get('country', '?')} | "
                    f"**{info.get('ip_type', '?')}** |\n"
                )
            f.write("\n")

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

        for tier in ("CORPORATE", "INFRASTRUCTURE", "POSSIBLE"):
            items = tiers.get(tier, [])
            if not items:
                continue
            f.write(f"## {tier} ({len(items)} hosts)\n\n")

            for item in items:
                p = item["profile"]
                f.write(
                    f"### [{p['dump_id']}] {p['computer_name']} / "
                    f"{p['username']}  (score: {item['score']})\n\n"
                )

                f.write("**Indicators:**\n\n")
                for reason in item["reasons"]:
                    f.write(f"- {reason}\n")
                f.write("\n")

                _write_profile_details(f, p)
                f.write("---\n\n")

    print(f"Written: {OUTPUT_REPORT}")


def _write_profile_details(f, p: dict):
    if p["rdp_connections"]:
        f.write("**RDP Connections (from .rdp files):**\n\n")
        f.write("| File | Target | IP Class | Username | Gateway |\n")
        f.write("|------|--------|----------|----------|---------|\n")
        for r in p["rdp_connections"]:
            f.write(
                f"| `{r['filename']}` | {r['target']} | "
                f"{r.get('ip_class', '')} | {r.get('username', '')} | "
                f"{r.get('gateway', '')} |\n"
            )
        f.write("\n")

    domain_creds = [c for c in p["vault_domain_creds"]
                    if c.get("is_termsrv") or c.get("domain_name")]
    if domain_creds:
        f.write("**Domain/RDP Credentials (from vault):**\n\n")
        f.write("| Target | Type | Domain | User |\n")
        f.write("|--------|------|--------|------|\n")
        for c in domain_creds:
            termsrv = "TERMSRV" if c.get("is_termsrv") else "Domain"
            f.write(
                f"| {c['target_host']} | {termsrv} | "
                f"{c.get('domain_name', '')} | {c['username']} |\n"
            )
        f.write("\n")

    service_creds = [c for c in p["vault_domain_creds"]
                     if c.get("indicators")]
    if service_creds:
        f.write("**Service Credentials:**\n\n")
        for c in service_creds:
            inds = ", ".join(c["indicators"])
            f.write(f"- [{inds}] {c['target_host']} (user: {c['username']})\n")
        f.write("\n")

    if p["putty_sessions"]:
        non_game = [s for s in p["putty_sessions"] if not s["is_game"]]
        if non_game:
            f.write("**PuTTY Sessions:**\n\n")
            f.write("| Session | Host | Port | User | IP Class | WHOIS |\n")
            f.write("|---------|------|------|------|----------|-------|\n")
            for s in non_game:
                w = s.get("whois", {})
                whois_col = (f"{w.get('org_name', '')} "
                             f"({w.get('ip_type', '')})"
                             if w else "")
                f.write(
                    f"| {s['name']} | {s['hostname']} | {s['port']} | "
                    f"{s['username']} | {s['ip_class']} | {whois_col} |\n"
                )
            f.write("\n")

    if p["winscp_sessions"]:
        non_game = [s for s in p["winscp_sessions"] if not s["is_game"]]
        if non_game:
            f.write("**WinSCP Sessions:**\n\n")
            f.write("| Session | Host | User | Password? | IP Class | WHOIS |\n")
            f.write("|---------|------|------|-----------|----------|-------|\n")
            for s in non_game:
                w = s.get("whois", {})
                whois_col = (f"{w.get('org_name', '')} "
                             f"({w.get('ip_type', '')})"
                             if w else "")
                f.write(
                    f"| {s['name']} | {s['hostname']} | {s['username']} | "
                    f"{'Yes' if s['has_password'] else 'No'} | "
                    f"{s['ip_class']} | {whois_col} |\n"
                )
            f.write("\n")

    corp_vpn = [v for v in p["vpn_endpoints"] if not v.get("is_commercial")]
    if corp_vpn:
        f.write("**VPN Endpoints:**\n\n")
        for v in corp_vpn:
            auth = " [auth-user-pass]" if v.get("has_auth_user_pass") else ""
            w = v.get("whois", {})
            whois_note = ""
            if w:
                whois_note = (f" -- {w.get('org_name', '?')} "
                              f"({w.get('ip_type', '?')}, "
                              f"{w.get('country', '?')})")
            f.write(f"- {v['endpoint']}{auth}{whois_note} (`{v['file']}`)\n")
        f.write("\n")

    non_game_fz = [fz for fz in p["filezilla_sites"] if not fz["is_game"]]
    if non_game_fz:
        seen_fz = set()
        deduped = []
        for fz in non_game_fz:
            key = (fz["host"], fz["port"], fz["user"], fz["protocol"])
            if key not in seen_fz:
                seen_fz.add(key)
                deduped.append(fz)
        if deduped:
            f.write("**FileZilla Sites:**\n\n")
            f.write("| Site | Host:Port | User | Protocol | Password? | WHOIS |\n")
            f.write("|------|-----------|------|----------|-----------|-------|\n")
            for fz in deduped:
                w = fz.get("whois", {})
                whois_col = (f"{w.get('org_name', '')} "
                             f"({w.get('ip_type', '')})"
                             if w else "")
                f.write(
                    f"| {fz['site_name']} | {fz['host']}:{fz['port']} | "
                    f"{fz['user']} | {fz['protocol']} | "
                    f"{'Yes' if fz['has_password'] else 'No'} | "
                    f"{whois_col} |\n"
                )
            f.write("\n")

    notes = []
    if p["has_server_ovpn"]:
        notes.append("Server-side OpenVPN config found")
    if p["has_jira_cert"]:
        notes.append("Jira/Confluence certificate found")
    if p["has_ssms"]:
        notes.append("SQL Server Management Studio credential")
    if p["has_exchange"]:
        notes.append("Exchange/autodiscover credential")
    if p["has_edu"]:
        notes.append("Academic (.edu) domain")
    big_kp = [k for k in p["keepass_dbs"] if k["size"] > 10000]
    if big_kp:
        notes.append(
            f"KeePass databases > 10KB: "
            + ", ".join(f"`{k['file']}` ({k['size']}B)" for k in big_kp)
        )
    if p["ssh_keys_unencrypted"]:
        notes.append(f"{p['ssh_keys_unencrypted']} unencrypted SSH keys")
    if notes:
        f.write("**Additional Notes:**\n\n")
        for n in notes:
            f.write(f"- {n}\n")
        f.write("\n")


# ---------------------------------------------------------------------------
# Main
# ---------------------------------------------------------------------------

def run():
    t0 = time.time()
    print("=== Corporate Access Detector ===\n")

    print("Loading hosts...")
    hosts = load_hosts()
    print(f"  {len(hosts)} hosts loaded")

    print("Loading extracted credentials...")
    with open(EXTRACTED_JSON, encoding="utf-8") as f:
        extracted = json.load(f)
    print(f"  {len(extracted)} entries loaded")

    print("Loading triage rows...")
    triage_rows = load_triage_rows(CRITICAL_CSV, HIGH_CSV)
    print(f"  {len(triage_rows)} triage rows loaded")

    # Phase 1
    print("\n--- Phase 1: Deep Parsing ---")

    print("Parsing RDP files...")
    rdp_data = parse_rdp_files(triage_rows)
    rdp_count = sum(len(v) for v in rdp_data.values())
    print(f"  {rdp_count} RDP connections from {len(rdp_data)} hosts")

    print("Parsing registry files (PuTTY/WinSCP)...")
    reg_data = parse_registry_files(triage_rows)
    putty_count = sum(len(v.get("putty_sessions", []))
                      for v in reg_data.values())
    winscp_count = sum(len(v.get("winscp_sessions", []))
                       for v in reg_data.values())
    print(f"  {putty_count} PuTTY sessions, {winscp_count} WinSCP sessions "
          f"from {len(reg_data)} hosts")

    print("Parsing vault credentials...")
    vault_data = parse_vault_credentials(extracted)
    vault_count = sum(len(v) for v in vault_data.values())
    print(f"  {vault_count} structured vault creds from {len(vault_data)} hosts")

    print("Parsing FileZilla XMLs...")
    fz_data = parse_filezilla_xmls(triage_rows)
    fz_count = sum(len(v) for v in fz_data.values())
    print(f"  {fz_count} FileZilla sites from {len(fz_data)} hosts")

    # Phase 1.5: WHOIS Enrichment
    print("\n--- Phase 1.5: WHOIS Enrichment ---")
    whois_db, dns_map = enrich_whois(rdp_data, reg_data, vault_data,
                                     fz_data, extracted)

    # Phase 2
    print("\n--- Phase 2: Aggregation ---")
    profiles = aggregate(hosts, extracted, rdp_data, reg_data,
                         vault_data, fz_data, triage_rows,
                         whois_db, dns_map)
    print(f"  {len(profiles)} host profiles built")

    # Phase 3
    print("\n--- Phase 3: Scoring ---")
    scored = []
    for dump_id, profile in sorted(profiles.items()):
        sc, classification, reasons = score_host(profile)
        if sc > 0:
            scored.append({
                "dump_id": dump_id,
                "computer_name": profile["computer_name"],
                "username": profile["username"],
                "score": sc,
                "classification": classification,
                "reasons": reasons,
                "profile": profile,
            })

    scored.sort(key=lambda x: x["score"], reverse=True)

    from collections import Counter
    tier_counts = Counter(x["classification"] for x in scored)
    print(f"  Scored {len(scored)} hosts with score > 0")
    for tier in ("CORPORATE", "INFRASTRUCTURE", "POSSIBLE", "PERSONAL"):
        if tier_counts[tier]:
            print(f"    {tier}: {tier_counts[tier]}")

    # Phase 4
    print("\n--- Phase 4: Output ---")
    generate_json(scored)
    generate_report(scored, whois_db)

    elapsed = time.time() - t0
    print(f"\nDone in {elapsed:.1f}s")

    if scored:
        print(f"\nTop 10 hosts by corporate score:")
        for item in scored[:10]:
            print(
                f"  [{item['dump_id']}] {item['computer_name']}/"
                f"{item['username']}: score={item['score']} "
                f"({item['classification']})"
            )


if __name__ == "__main__":
    run()
