"""
Content scanner: reads CRITICAL and HIGH findings, inspects file contents,
and extracts passwords, keys, tokens, mnemonics, endpoints.

Produces:
  - data/extracted_credentials.json
"""

import csv
import json
import re
import sys
import time
from pathlib import Path

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"

CRITICAL_CSV = DATA_DIR / "legacy_critical_findings.csv"
HIGH_CSV = DATA_DIR / "legacy_high_findings.csv"
OUTPUT_JSON = DATA_DIR / "legacy_extracted_credentials.json"

MAX_FILE_SIZE = 10 * 1024 * 1024  # 10MB

# ---------------------------------------------------------------------------
# BIP-39 mnemonic detection (English)
# ---------------------------------------------------------------------------
BIP39_WORDS = None
BIP39_PATH = Path(__file__).with_name("bip39_english.txt")


def load_bip39():
    global BIP39_WORDS
    if BIP39_PATH.exists():
        BIP39_WORDS = set(BIP39_PATH.read_text().strip().splitlines())
    else:
        BIP39_WORDS = set()


# ---------------------------------------------------------------------------
# Regex patterns
# ---------------------------------------------------------------------------
PASSWORD_PATTERNS = [
    re.compile(r'(?:password|passwd|pass|pwd)\s*[=:]\s*["\']?(.+?)["\']?\s*$', re.I | re.M),
    re.compile(r'(?:secret|token|api[_-]?key)\s*[=:]\s*["\']?(.+?)["\']?\s*$', re.I | re.M),
    re.compile(r'auth[_-]?user[_-]?pass', re.I),
]

PRIVATE_KEY_RE = re.compile(
    r'-----BEGIN (?:RSA |DSA |EC |OPENSSH |ENCRYPTED )?PRIVATE KEY-----'
)
ENCRYPTED_KEY_RE = re.compile(r'Proc-Type:\s*4,ENCRYPTED|ENCRYPTED', re.I)

IP_RE = re.compile(r'\b(?:\d{1,3}\.){3}\d{1,3}\b')
DOMAIN_RE = re.compile(
    r'\b(?:[a-zA-Z0-9](?:[a-zA-Z0-9-]{0,61}[a-zA-Z0-9])?\.)'
    r'+(?:com|net|org|io|dev|ru|de|uk|fr|info|biz|co|me|us|ca|au|nl|se|no|fi|ch|at)\b',
    re.I,
)
URL_RE = re.compile(r'https?://[^\s"\'<>]+', re.I)

RDP_PASSWORD_RE = re.compile(r'password\s+51:b:(.+)', re.I)
FILEZILLA_PASS_RE = re.compile(r'<Pass[^>]*>([^<]+)</Pass>', re.I)
WINSCP_PASS_RE = re.compile(r'Password\s*=\s*(.+)', re.I)

AWS_KEY_RE = re.compile(r'(?:AKIA|ASIA)[A-Z0-9]{16,}')
BEARER_RE = re.compile(r'Bearer\s+[A-Za-z0-9\-_.~+/]+=*', re.I)

VAULT_INTERESTING_RE = re.compile(
    r'(?:password|passwd|credential|secret|token|api.?key)',
    re.I,
)


def detect_mnemonic(text: str) -> list[str]:
    """Find potential BIP-39 mnemonic phrases (12 or 24 words)."""
    if not BIP39_WORDS:
        return []
    findings = []
    words = text.lower().split()
    for window in (12, 24):
        if len(words) < window:
            continue
        for i in range(len(words) - window + 1):
            chunk = words[i : i + window]
            if all(w in BIP39_WORDS for w in chunk):
                phrase = " ".join(chunk)
                if phrase not in findings:
                    findings.append(phrase)
    return findings


def scan_text_file(text: str, row: dict) -> list[dict]:
    """Scan text content and return list of finding dicts."""
    findings = []
    reason = row.get("reason", "")
    ext = row.get("extension", "")
    fname = row.get("filename", "").lower()

    if PRIVATE_KEY_RE.search(text):
        encrypted = bool(ENCRYPTED_KEY_RE.search(text))
        findings.append({
            "type": "private_key",
            "encrypted": encrypted,
            "detail": "encrypted" if encrypted else "UNENCRYPTED",
        })

    for pat in PASSWORD_PATTERNS:
        for m in pat.finditer(text):
            val = m.group(1) if pat.groups else m.group(0)
            val = val.strip()
            if len(val) > 2 and val not in ("none", "null", "empty", "N/A", "''", '""'):
                findings.append({"type": "password", "detail": val})

    if ext == ".rdp":
        for m in RDP_PASSWORD_RE.finditer(text):
            findings.append({"type": "rdp_password_hash", "detail": m.group(1).strip()})

    if "filezilla" in fname or "sitemanager" in fname or "recentservers" in fname:
        for m in FILEZILLA_PASS_RE.finditer(text):
            findings.append({"type": "filezilla_password", "detail": m.group(1).strip()})
        hosts = re.findall(r'<Host>([^<]+)</Host>', text, re.I)
        users = re.findall(r'<User>([^<]+)</User>', text, re.I)
        ports = re.findall(r'<Port>([^<]+)</Port>', text, re.I)
        for i, h in enumerate(hosts):
            u = users[i] if i < len(users) else ""
            p = ports[i] if i < len(ports) else ""
            findings.append({"type": "ftp_endpoint", "detail": f"{u}@{h}:{p}"})

    if "winscp" in fname:
        for m in WINSCP_PASS_RE.finditer(text):
            findings.append({"type": "winscp_password", "detail": m.group(1).strip()})

    if ext == ".ovpn" or "openvpn" in fname:
        if "auth-user-pass" in text:
            findings.append({"type": "vpn_auth_user_pass", "detail": "auth-user-pass directive found"})
        for m in re.finditer(r'remote\s+(\S+)\s+(\d+)', text):
            findings.append({"type": "vpn_endpoint", "detail": f"{m.group(1)}:{m.group(2)}"})

    for m in AWS_KEY_RE.finditer(text):
        findings.append({"type": "aws_key", "detail": m.group(0)[:30]})

    mnemonics = detect_mnemonic(text)
    for phrase in mnemonics:
        findings.append({"type": "mnemonic_phrase", "detail": phrase})

    ips = set(IP_RE.findall(text))
    private_ips = {ip for ip in ips if not ip.startswith(("127.", "0.", "255."))}
    if private_ips and len(private_ips) < 50:
        findings.append({"type": "ip_addresses", "detail": ", ".join(sorted(private_ips)[:20])})

    urls = set(URL_RE.findall(text))
    interesting_urls = {u for u in urls if not any(
        skip in u.lower() for skip in
        ("microsoft.com", "xbox", "windows.net", "aka.ms", "bing.com",
         "google.com", "googleapis", "gstatic", "mozilla.org",
         "w3.org", "xmlsoap.org", "xml.org", "schemas.microsoft")
    )}
    if interesting_urls and len(interesting_urls) < 50:
        for url in sorted(interesting_urls)[:10]:
            findings.append({"type": "url", "detail": url[:200]})

    return findings


def scan_vault_csv(text: str, row: dict) -> list[dict]:
    """Extract interesting entries from Windows Vault CSV."""
    findings = []
    lines = text.splitlines()
    for line in lines:
        lower = line.lower()
        if any(kw in lower for kw in ("password", "credential", "secret")):
            if "xbl" not in lower and "xbox" not in lower and "minecraft" not in lower:
                parts = line.split(",")
                target = parts[1] if len(parts) > 1 else ""
                username = parts[2] if len(parts) > 2 else ""
                if target or username:
                    findings.append({
                        "type": "vault_credential",
                        "detail": f"target={target[:100]} user={username[:50]}",
                    })
    return findings


def resolve_path(row: dict) -> Path:
    dump_id = int(row["dump_id"])
    folder = "dump" if dump_id == 0 else f"dump ({dump_id})"
    return DUMPS_DIR / folder / row["relative_path"]


def scan_file(row: dict) -> list[dict]:
    """Scan a single file and return findings."""
    fpath = resolve_path(row)
    if not fpath.exists():
        return []

    size = int(row.get("size_bytes", 0))
    if size > MAX_FILE_SIZE:
        return [{"type": "skipped", "detail": f"file too large: {size} bytes"}]
    if size == 0:
        return []

    try:
        raw = fpath.read_bytes()
    except Exception as e:
        return [{"type": "error", "detail": str(e)[:200]}]

    if b'\x00' in raw[:512]:
        return [{"type": "binary_file", "detail": f"binary, {size} bytes"}]

    try:
        text = raw.decode("utf-8", errors="replace")
    except Exception:
        return [{"type": "decode_error", "detail": "could not decode"}]

    cat = row.get("category", "")
    if cat == "vault" or "vault" in row.get("filename", "").lower():
        return scan_vault_csv(text, row)

    return scan_text_file(text, row)


def run():
    load_bip39()
    if BIP39_WORDS:
        print(f"Loaded {len(BIP39_WORDS)} BIP-39 words")
    else:
        print("BIP-39 wordlist not found, mnemonic detection disabled")
        print(f"  (place wordlist at {BIP39_PATH})")

    all_findings = []
    files_scanned = 0
    files_with_findings = 0

    for csv_path, severity in [(CRITICAL_CSV, "CRITICAL"), (HIGH_CSV, "HIGH")]:
        print(f"\nScanning {severity} findings from {csv_path.name}...")
        if not csv_path.exists():
            print(f"  File not found, skipping")
            continue

        with open(csv_path, encoding="utf-8") as f:
            rows = list(csv.DictReader(f))

        t0 = time.time()
        for i, row in enumerate(rows):
            findings = scan_file(row)
            files_scanned += 1

            if findings:
                non_trivial = [f for f in findings if f["type"] not in ("skipped", "error", "decode_error")]
                if non_trivial:
                    files_with_findings += 1
                    entry = {
                        "dump_id": int(row["dump_id"]),
                        "computer_name": row.get("computer_name", ""),
                        "username": row.get("username", ""),
                        "severity": severity,
                        "file": row["relative_path"],
                        "filename": row.get("filename", ""),
                        "findings": non_trivial,
                    }
                    all_findings.append(entry)

            if (i + 1) % 2000 == 0:
                elapsed = time.time() - t0
                print(f"  [{i+1}/{len(rows)}] {elapsed:.1f}s, findings so far: {files_with_findings}")

        elapsed = time.time() - t0
        print(f"  Done: {len(rows)} files in {elapsed:.1f}s")

    all_findings.sort(key=lambda x: (
        0 if x["severity"] == "CRITICAL" else 1,
        -len(x["findings"]),
        x["dump_id"],
    ))

    with open(OUTPUT_JSON, "w", encoding="utf-8") as f:
        json.dump(all_findings, f, indent=2, ensure_ascii=False)

    print(f"\nTotal files scanned: {files_scanned}")
    print(f"Files with findings: {files_with_findings}")
    print(f"Output: {OUTPUT_JSON}")

    from collections import Counter
    type_counts = Counter()
    for entry in all_findings:
        for finding in entry["findings"]:
            type_counts[finding["type"]] += 1
    print("\nFinding types:")
    for ftype, cnt in type_counts.most_common():
        print(f"  {ftype}: {cnt}")


if __name__ == "__main__":
    run()
