#!/usr/bin/env python3
"""Harvest Telegram bot tokens from stealer malware samples (P0 pilot, docs/log-source-discovery.md §5.2).

Low/mid-tier stealers (Snake, AgentTesla/OriginLogger, PXA, Prynt, Phemedrone,
Worldwind, VipKeylogger) exfiltrate via Telegram Bot API. The bot token +
chat_id are hardcoded in each sample config. Anyone holding the token can poll
the Bot API (getUpdates) and passively receive the same exfiltrated per-victim
archives the operator receives.

Pipeline per sample:
  1. Download from MalwareBazaar (get_taginfo → get_file, zip pw "infected")
  2. Extract ASCII + UTF-16LE strings in-memory
  3. Regex-extract bot tokens + chat_ids (telegram context only)
  4. Validate via Bot API: getMe (alive?) + getWebhookInfo (collection strategy)
  5. Append to findings/data/tg_bots_validated.tsv (idempotent by sha256+token)

Usage:
  python3 scripts/tg_bot_harvest.py                          # default tags, 5 samples/tag
  python3 scripts/tg_bot_harvest.py --tags Snake,Phemedrone --limit 10
  python3 scripts/tg_bot_harvest.py --file sample.exe        # local sample, no download
  python3 scripts/tg_bot_harvest.py --no-validate            # extract only, no Bot API calls

Auth: MalwareBazaar API requires an Auth-Key. Pass --auth-key, or set
MB_AUTH_KEY in environment / .env. Free account: https://bazaar.abuse.ch/api/

Rules: no-masking (tokens written in full); validation = getMe executed + observed.
"""
import argparse
import csv
import io
import json
import os
import re
import sys
import time
import urllib.parse
import urllib.request
import zipfile
from datetime import datetime, timezone
from pathlib import Path

SCRIPT_DIR = Path(__file__).parent
PROJECT_DIR = SCRIPT_DIR.parent
DATA_DIR = PROJECT_DIR / "findings" / "data"
OUT_TSV = DATA_DIR / "tg_bots_validated.tsv"
PROCESSED_TXT = DATA_DIR / "tg_bots_processed_sha256.txt"

MB_API = "https://mb-api.abuse.ch/api/v1/"
TF_API = "https://threatfox-api.abuse.ch/api/v1/"
UH_API = "https://urlhaus-api.abuse.ch/v1/host/"
TG_API = "https://api.telegram.org/bot{token}/{method}"
MB_ZIP_PASSWORD = b"infected"

# ThreatFox family tags known to exfiltrate via Telegram Bot API
TF_TAGS = [
    "AgentTesla",
    "AsyncRAT",
    "OriginLogger",
    "SnakeKeylogger",
    "Phemedrone",
    "Stealc",
    "Vidar",
    "XWorm",
    "DcRat",
    "VenomRAT",
    "QuasarRAT",
    "RevengeRAT",
    "Remcos",
    "NanoCore",
]

# TG-exfil stealer families (docs/log-source-discovery.md §3.1)
DEFAULT_TAGS = [
    "Phemedrone",
    "Snake Keylogger",
    "AgentTesla",
    "OriginLogger",
    "Prynt Stealer",
    "WorldWind Stealer",
    "VipKeylogger",
    "PXA Stealer",
]

TOKEN_RE = re.compile(r"(?<![0-9])[0-9]{8,10}:[A-Za-z0-9_-]{35}(?![A-Za-z0-9_-])")
# chat_id only when seen in a telegram context (URL or key-value near it)
CHAT_ID_RE = re.compile(r"chat_id[\"'=:\s]+(-?[0-9]{5,15})", re.I)
ASCII_STR_RE = re.compile(rb"[\x20-\x7e]{6,}")
UTF16_STR_RE = re.compile(rb"(?:[\x20-\x7e]\x00){6,}")

TSV_FIELDS = [
    "tag", "sha256", "signature", "token", "chat_id",
    "bot_id", "bot_username", "webhook_url", "pending_updates",
    "status", "first_seen", "last_checked",
]


def load_auth_key(cli_key: str | None) -> str | None:
    if cli_key:
        return cli_key
    if os.environ.get("MB_AUTH_KEY"):
        return os.environ["MB_AUTH_KEY"]
    env_file = PROJECT_DIR / ".env"
    if env_file.exists():
        for line in env_file.read_text(errors="replace").splitlines():
            if line.startswith("MB_AUTH_KEY="):
                return line.split("=", 1)[1].strip().strip('"').strip("'")
    return None


def mb_query(payload: dict, auth_key: str) -> dict:
    data = urllib.parse.urlencode(payload).encode()
    req = urllib.request.Request(MB_API, data=data, headers={"Auth-Key": auth_key})
    with urllib.request.urlopen(req, timeout=60) as r:
        return json.loads(r.read())


def tf_query(payload: dict, auth_key: str) -> dict:
    data = json.dumps(payload).encode()
    req = urllib.request.Request(TF_API, data=data,
                                 headers={"Auth-Key": auth_key,
                                          "Content-Type": "application/json"})
    with urllib.request.urlopen(req, timeout=60) as r:
        return json.loads(r.read())


def threatfox_tokens(auth_key: str, tags: list[str]) -> list[dict]:
    """Collect bot tokens from ThreatFox IOCs (url IOCs pointing at the Bot API).

    Sources: global search for api.telegram.org + per-family taginfo for
    TG-exfil families. Returns [{ioc_id, tag, token, first_seen}]."""
    iocs: dict[str, dict] = {}

    def ingest(resp: dict, tag: str):
        for i in resp.get("data") or []:
            if not isinstance(i, dict):
                continue
            url = i.get("ioc", "")
            if "api.telegram.org" not in url:
                continue
            # chat_id often present as query param in sendMessage/sendDocument IOCs
            chat_id = ""
            qm = re.search(r"[?&]chat_id=(-?\d{5,15})", url)
            if qm:
                chat_id = qm.group(1)
            for tok in TOKEN_RE.findall(url):
                key = tok
                prev = iocs.get(key, {})
                iocs[key] = {
                    "ioc_id": str(i.get("id", "")),
                    "tag": i.get("malware_printable") or tag,
                    "token": tok,
                    "chat_id": prev.get("chat_id") or chat_id,
                    "first_seen": i.get("first_seen", ""),
                }

    print("[*] ThreatFox search_ioc: api.telegram.org")
    ingest(tf_query({"query": "search_ioc", "search_term": "api.telegram.org"},
                    auth_key), "unknown")
    # recent feed: freshest IOCs, catches TG URLs before search indexing
    for days in (1, 3):
        print(f"[*] ThreatFox get_iocs: last {days}d")
        try:
            ingest(tf_query({"query": "get_iocs", "days": days}, auth_key), "recent")
        except Exception as e:
            print(f"    failed: {e}")
    for tag in tags:
        print(f"[*] ThreatFox taginfo: {tag}")
        try:
            ingest(tf_query({"query": "taginfo", "tag": tag, "limit": 1000},
                            auth_key), tag)
        except Exception as e:
            print(f"    failed: {e}")

    # URLhaus: host query for api.telegram.org (bot API URLs as malware_download)
    print("[*] URLhaus host: api.telegram.org")
    try:
        data = urllib.parse.urlencode({"host": "api.telegram.org"}).encode()
        req = urllib.request.Request(UH_API, data=data,
                                     headers={"Auth-Key": auth_key})
        with urllib.request.urlopen(req, timeout=60) as r:
            uh = json.loads(r.read())
        for u in uh.get("urls") or []:
            u_url = u.get("url", "")
            qm = re.search(r"[?&]chat_id=(-?\d{5,15})", u_url)
            for tok in TOKEN_RE.findall(u_url):
                prev = iocs.get(tok, {})
                iocs[tok] = {
                    "ioc_id": str(u.get("id", "")),
                    "tag": "urlhaus:" + (u.get("threat") or ""),
                    "token": tok,
                    "chat_id": prev.get("chat_id") or (qm.group(1) if qm else ""),
                    "first_seen": u.get("date_added", ""),
                }
    except Exception as e:
        print(f"    failed: {e}")
    return sorted(iocs.values(), key=lambda x: x["first_seen"], reverse=True)


def mb_download(sha256: str, auth_key: str) -> bytes:
    data = urllib.parse.urlencode({"query": "get_file", "sha256_hash": sha256}).encode()
    req = urllib.request.Request(MB_API, data=data, headers={"Auth-Key": auth_key})
    with urllib.request.urlopen(req, timeout=120) as r:
        body = r.read()
    try:
        err = json.loads(body)
        raise RuntimeError(f"MB download error: {err}")
    except (json.JSONDecodeError, UnicodeDecodeError):
        pass
    return body


def extract_payload(zip_bytes: bytes) -> bytes:
    """Extract first (largest) file from MB password-protected zip.

    MB zips are AES-encrypted (method 99) — Python zipfile can't read them,
    so fall back to system 7z (project prerequisite, see SETUP.md)."""
    try:
        with zipfile.ZipFile(io.BytesIO(zip_bytes)) as zf:
            infos = sorted(zf.infolist(), key=lambda i: i.file_size, reverse=True)
            if not infos:
                raise RuntimeError("empty zip")
            return zf.read(infos[0], pwd=MB_ZIP_PASSWORD)
    except RuntimeError as e:
        if "compression method" not in str(e):
            raise
    import subprocess
    import tempfile
    with tempfile.TemporaryDirectory() as td:
        zpath = Path(td) / "sample.zip"
        zpath.write_bytes(zip_bytes)
        outdir = Path(td) / "out"
        r = subprocess.run(
            ["7z", "x", "-p" + MB_ZIP_PASSWORD.decode(), "-y", f"-o{outdir}", str(zpath)],
            capture_output=True, timeout=120)
        if r.returncode != 0:
            raise RuntimeError(f"7z failed: {r.stderr.decode(errors='replace')[:200]}")
        files = [p for p in outdir.rglob("*") if p.is_file()]
        if not files:
            raise RuntimeError("7z extracted nothing")
        return max(files, key=lambda p: p.stat().st_size).read_bytes()


def iter_strings(blob: bytes):
    for m in ASCII_STR_RE.finditer(blob):
        yield m.group().decode("ascii", errors="replace")
    for m in UTF16_STR_RE.finditer(blob):
        yield m.group().decode("utf-16-le", errors="replace")


import zlib

ZLIB_HEADER_RE = re.compile(rb"\x78[\x01\x9c\xda]")


def iter_decompressed(blob: bytes, max_blobs: int = 2000, max_out: int = 20_000_000):
    """Yield zlib-decompressed payloads found at any offset (PyInstaller PYZ
    entries, embedded resources). Best-effort: partial streams are fine."""
    n = 0
    for m in ZLIB_HEADER_RE.finditer(blob):
        if n >= max_blobs:
            break
        try:
            d = zlib.decompressobj()
            out = d.decompress(blob[m.start():], max_out)
        except Exception:
            continue
        if len(out) > 64:
            n += 1
            yield out


def extract_config(blob: bytes) -> list[dict]:
    """Return [{token, chat_id}] pairs found in the sample (raw + zlib layers)."""
    found: dict[str, str] = {}

    def scan(strings):
        for s in strings:
            for tok in TOKEN_RE.findall(s):
                found.setdefault(tok, "")
            m = CHAT_ID_RE.search(s)
            if m and ("telegram" in s.lower() or TOKEN_RE.search(s)):
                tok_in_str = TOKEN_RE.search(s)
                if tok_in_str:
                    found[tok_in_str.group()] = m.group(1)

    scan(iter_strings(blob))
    for layer in iter_decompressed(blob):
        scan(iter_strings(layer))
    return [{"token": t, "chat_id": c} for t, c in sorted(found.items())]


def tg_api(token: str, method: str, timeout: int = 20) -> dict:
    url = TG_API.format(token=token, method=method)
    try:
        with urllib.request.urlopen(url, timeout=timeout) as r:
            return json.loads(r.read())
    except urllib.error.HTTPError as e:
        return {"ok": False, "error": f"HTTP {e.code}"}
    except Exception as e:
        return {"ok": False, "error": str(e)}


def validate_token(token: str) -> dict:
    """getMe + getWebhookInfo. Returns status row fields."""
    me = tg_api(token, "getMe")
    if not me.get("ok"):
        return {"bot_id": "", "bot_username": "", "webhook_url": "",
                "pending_updates": "", "status": f"dead:{me.get('error', 'unknown')}"}
    wh = tg_api(token, "getWebhookInfo").get("result", {})
    return {
        "bot_id": str(me["result"].get("id", "")),
        "bot_username": me["result"].get("username", ""),
        "webhook_url": wh.get("url", ""),
        "pending_updates": str(wh.get("pending_update_count", "")),
        "status": "valid",
    }


def load_processed() -> set[str]:
    if PROCESSED_TXT.exists():
        return set(PROCESSED_TXT.read_text().split())
    return set()


def mark_processed(sha256: str):
    with PROCESSED_TXT.open("a") as f:
        f.write(sha256 + "\n")


def known_token_sha_pairs() -> set[tuple[str, str]]:
    if not OUT_TSV.exists():
        return set()
    pairs = set()
    with OUT_TSV.open() as f:
        for row in csv.DictReader(f, delimiter="\t"):
            pairs.add((row["sha256"], row["token"]))
    return pairs


def append_rows(rows: list[dict]):
    DATA_DIR.mkdir(parents=True, exist_ok=True)
    write_header = not OUT_TSV.exists()
    with OUT_TSV.open("a", newline="") as f:
        w = csv.DictWriter(f, fieldnames=TSV_FIELDS, delimiter="\t")
        if write_header:
            w.writeheader()
        for r in rows:
            w.writerow(r)


def process_sample(tag: str, sha256: str, signature: str, blob: bytes,
                   do_validate: bool) -> list[dict]:
    now = datetime.now(timezone.utc).strftime("%Y-%m-%dT%H:%M:%SZ")
    pairs = extract_config(blob)
    rows = []
    for p in pairs:
        row = {
            "tag": tag, "sha256": sha256, "signature": signature,
            "token": p["token"], "chat_id": p["chat_id"],
            "bot_id": "", "bot_username": "", "webhook_url": "",
            "pending_updates": "", "status": "extracted",
            "first_seen": now, "last_checked": now,
        }
        if do_validate:
            row.update(validate_token(p["token"]))
            row["last_checked"] = now
        rows.append(row)
    return rows


def main() -> int:
    ap = argparse.ArgumentParser(description=__doc__.splitlines()[0])
    ap.add_argument("--source", choices=["threatfox", "mb"], default="threatfox",
                    help="token source: ThreatFox IOCs (default, high yield) or "
                         "MalwareBazaar static sample extraction")
    ap.add_argument("--tags", default="",
                    help="comma-separated tags (ThreatFox families / MB tags; "
                         "defaults built-in per source)")
    ap.add_argument("--limit", type=int, default=5, help="samples per tag (mb mode)")
    ap.add_argument("--auth-key", help="abuse.ch Auth-Key (or MB_AUTH_KEY env/.env)")
    ap.add_argument("--file", help="process a local sample file instead of downloading")
    ap.add_argument("--tag", default="local", help="tag label for --file mode")
    ap.add_argument("--no-validate", action="store_true", help="skip Bot API validation")
    ap.add_argument("--sleep", type=float, default=1.0, help="pause between samples (API politeness)")
    args = ap.parse_args()

    do_validate = not args.no_validate
    processed = load_processed()
    known = known_token_sha_pairs()
    all_rows: list[dict] = []
    now = datetime.now(timezone.utc).strftime("%Y-%m-%dT%H:%M:%SZ")

    if args.file:
        sha = __import__("hashlib").sha256(Path(args.file).read_bytes()).hexdigest()
        blob = Path(args.file).read_bytes()
        rows = process_sample(args.tag, sha, "local", blob, do_validate)
        all_rows.extend(rows)
    else:
        auth_key = load_auth_key(args.auth_key)
        if not auth_key:
            print("ERROR: abuse.ch Auth-Key required. Use --auth-key or MB_AUTH_KEY "
                  "(env/.env). Free account: https://bazaar.abuse.ch/api/", file=sys.stderr)
            return 2
        if args.source == "threatfox":
            tags = ([t.strip() for t in args.tags.split(",") if t.strip()] or TF_TAGS)
            for item in threatfox_tokens(auth_key, tags):
                row = {
                    "tag": item["tag"], "sha256": f"tf-{item['ioc_id']}",
                    "signature": "", "token": item["token"],
                    "chat_id": item.get("chat_id", ""),
                    "bot_id": "", "bot_username": "", "webhook_url": "",
                    "pending_updates": "", "status": "extracted",
                    "first_seen": item["first_seen"] or now, "last_checked": now,
                }
                if do_validate:
                    row.update(validate_token(item["token"]))
                    row["last_checked"] = now
                all_rows.append(row)
            n_valid = sum(1 for r in all_rows if r["status"] == "valid")
            print(f"[*] ThreatFox: {len(all_rows)} token(s), {n_valid} valid")
        else:  # mb static extraction
            import hashlib
            tags = ([t.strip() for t in args.tags.split(",") if t.strip()]
                    or DEFAULT_TAGS)
            for tag in tags:
                print(f"[*] tag={tag} limit={args.limit}")
                try:
                    resp = mb_query({"query": "get_taginfo", "tag": tag, "limit": args.limit},
                                    auth_key)
                except Exception as e:
                    print(f"    tag query failed: {e}")
                    continue
                samples = resp.get("data") or []
                if not samples:
                    print(f"    no samples ({resp.get('query_status', '?')})")
                    continue
                for s in samples:
                    sha256 = s.get("sha256_hash", "")
                    sig = s.get("signature") or ""
                    if sha256 in processed:
                        print(f"    skip {sha256[:16]}… (already processed)")
                        continue
                    try:
                        zbytes = mb_download(sha256, auth_key)
                        blob = extract_payload(zbytes)
                    except Exception as e:
                        print(f"    {sha256[:16]}… download/extract failed: {e}")
                        mark_processed(sha256)
                        continue
                    rows = process_sample(tag, sha256, sig, blob, do_validate)
                    n_valid = sum(1 for r in rows if r["status"] == "valid")
                    print(f"    {sha256[:16]}… {sig or '?'}: {len(rows)} token(s), {n_valid} valid")
                    all_rows.extend(rows)
                    mark_processed(sha256)
                    time.sleep(args.sleep)

    # dedup against known + within this run (keep first occurrence per sha+token)
    new_rows, seen = [], set(known)
    for r in all_rows:
        key = (r["sha256"], r["token"])
        if key not in seen:
            seen.add(key)
            new_rows.append(r)
    if new_rows:
        append_rows(new_rows)
    n_valid = sum(1 for r in new_rows if r["status"] == "valid")
    print(f"\n[+] {len(new_rows)} new row(s) → {OUT_TSV} ({n_valid} valid tokens)")
    return 0


if __name__ == "__main__":
    sys.exit(main())
