File size: 9,048 Bytes
994182c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
#!/usr/bin/env python3
"""Probe or download Hugging Face dataset splits for the CyberGym SFT mix.

Two modes:

  Probe (no heavy deps, just stdlib + network) -- inspect a dataset's real
  schema via the public datasets-server before trusting the manifest:

      python training/scripts/hf_download.py --probe colin/PrimeVul
      python training/scripts/hf_download.py --probe-manifest training/configs/datasets.yaml

  Download (requires `datasets`; runs on the GPU host) -- stream a split to
  raw JSONL under data/download/<key>/raw.jsonl:

      python training/scripts/hf_download.py --key primevul
      python training/scripts/hf_download.py --all --profile pilot
      python training/scripts/hf_download.py --hf-id colin/PrimeVul --split train --out data/download/primevul/raw.jsonl

Auth: set HF_TOKEN (or HUGGINGFACE_HUB_TOKEN) for gated datasets. Probe uses the
anonymous datasets-server; gated sets return HTTP 401 there (recorded in the
manifest `auth:` field).
"""

from __future__ import annotations

import argparse
import json
import os
import sys
import urllib.parse
import urllib.request
from pathlib import Path
from typing import Any

import yaml

DATASETS_SERVER = "https://datasets-server.huggingface.co"


def read_yaml(path: str | Path) -> dict[str, Any]:
    with Path(path).open("r", encoding="utf-8") as fh:
        payload = yaml.safe_load(fh) or {}
    if not isinstance(payload, dict):
        raise TypeError(f"Expected a YAML mapping in {path}")
    return payload


def _http_get(url: str, timeout: int = 20) -> dict[str, Any]:
    req = urllib.request.Request(url, headers={"User-Agent": "infosec-hf-probe"})
    token = os.getenv("HF_TOKEN") or os.getenv("HUGGINGFACE_HUB_TOKEN")
    if token:
        req.add_header("Authorization", f"Bearer {token}")
    with urllib.request.urlopen(req, timeout=timeout) as resp:
        return json.loads(resp.read())


def probe(hf_id: str) -> dict[str, Any]:
    """Return {splits, features, sample_row} or {error} for a dataset."""
    enc = urllib.parse.quote(hf_id)
    try:
        splits_payload = _http_get(f"{DATASETS_SERVER}/splits?dataset={enc}")
    except Exception as exc:  # noqa: BLE001 - surfaced to the caller
        return {"hf_id": hf_id, "error": repr(exc)}
    splits = splits_payload.get("splits", [])
    if not splits:
        return {"hf_id": hf_id, "error": "no splits", "raw": splits_payload}
    first = splits[0]
    cfg, sp = first["config"], first["split"]
    try:
        rows_payload = _http_get(
            f"{DATASETS_SERVER}/first-rows?dataset={enc}"
            f"&config={urllib.parse.quote(cfg)}&split={urllib.parse.quote(sp)}"
        )
    except Exception as exc:  # noqa: BLE001
        return {"hf_id": hf_id, "splits": splits, "error": f"first-rows: {exc!r}"}
    features = [
        {"name": f["name"], "type": f["type"].get("dtype", f["type"].get("_type"))}
        for f in rows_payload.get("features", [])
    ]
    sample = rows_payload.get("rows", [{}])[0].get("row", {}) if rows_payload.get("rows") else {}
    return {
        "hf_id": hf_id,
        "splits": [{"config": s["config"], "split": s["split"]} for s in splits],
        "config_used": cfg,
        "split_used": sp,
        "features": features,
        "sample_row": sample,
    }


def print_probe(result: dict[str, Any]) -> None:
    print("=" * 72)
    print(result["hf_id"])
    if "error" in result and "features" not in result:
        print(f"  ERROR: {result['error']}")
        return
    print(f"  splits: {result.get('splits')}")
    print(f"  used:   {result.get('config_used')}/{result.get('split_used')}")
    print("  features:")
    for feat in result.get("features", []):
        print(f"    - {feat['name']}: {feat['type']}")


def resolve_sources(manifest: dict[str, Any]) -> dict[str, dict[str, Any]]:
    return manifest.get("sources", {})


def download_split(
    hf_id: str,
    config: str | None,
    split: str,
    out_path: Path,
    max_rows: int | None,
) -> dict[str, Any]:
    try:
        from datasets import load_dataset
    except Exception as exc:  # pragma: no cover - exercised only on GPU host
        raise RuntimeError(
            "The `datasets` package is required for download mode. "
            "Install it on the training host (see requirements-cu126.txt)."
        ) from exc

    out_path.parent.mkdir(parents=True, exist_ok=True)
    token = os.getenv("HF_TOKEN") or os.getenv("HUGGINGFACE_HUB_TOKEN")
    stream = load_dataset(hf_id, config or None, split=split, streaming=True, token=token)

    written = 0
    with out_path.open("w", encoding="utf-8") as out:
        for row in stream:
            out.write(json.dumps(row, ensure_ascii=False, default=str) + "\n")
            written += 1
            if max_rows is not None and written >= max_rows:
                break
    return {"hf_id": hf_id, "split": split, "out": str(out_path), "rows": written}


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
    parser.add_argument("--probe", metavar="HF_ID", help="Probe one dataset schema via the datasets-server.")
    parser.add_argument("--probe-manifest", metavar="YAML", help="Probe every source in a manifest.")
    parser.add_argument("--manifest", default="training/configs/datasets.yaml")
    parser.add_argument("--key", help="Manifest source key to download.")
    parser.add_argument("--all", action="store_true", help="Download every enabled manifest source.")
    parser.add_argument("--eval", action="store_true", help="Download the eval_split instead of the train_split.")
    parser.add_argument("--profile", choices=["full", "pilot"], default="full")
    parser.add_argument("--hf-id", help="Ad-hoc download: dataset id (bypasses the manifest).")
    parser.add_argument("--config", help="Ad-hoc download: config name.")
    parser.add_argument("--split", help="Ad-hoc download: split name.")
    parser.add_argument("--out", help="Ad-hoc download: output JSONL path.")
    parser.add_argument("--max-rows", type=int, default=None)
    parser.add_argument("--include-disabled", action="store_true", help="With --all, also fetch disabled sources.")
    return parser.parse_args()


def _cap_for(source: dict[str, Any], defaults: dict[str, Any], profile: str) -> int | None:
    if profile == "pilot":
        return source.get("pilot_sample_cap", defaults.get("pilot_sample_cap"))
    cap = source.get("sample_cap", defaults.get("sample_cap"))
    return cap


def main() -> int:
    args = parse_args()

    if args.probe:
        print_probe(probe(args.probe))
        return 0

    if args.probe_manifest:
        manifest = read_yaml(args.probe_manifest)
        for key, source in resolve_sources(manifest).items():
            print(f"\n### {key} ({source.get('auth', 'public')}, enabled={source.get('enabled')})")
            print_probe(probe(source["hf_id"]))
        return 0

    if args.hf_id:
        out = Path(args.out or f"data/download/_adhoc/{args.hf_id.replace('/', '__')}.jsonl")
        result = download_split(args.hf_id, args.config, args.split or "train", out, args.max_rows)
        print(json.dumps(result, indent=2))
        return 0

    manifest = read_yaml(args.manifest)
    defaults = manifest.get("defaults", {})
    raw_dir = Path(defaults.get("raw_dir", "data/download"))
    sources = resolve_sources(manifest)

    keys: list[str]
    if args.all:
        keys = [
            k
            for k, s in sources.items()
            if (s.get("enabled", False) or args.include_disabled)
        ]
    elif args.key:
        keys = [args.key]
    else:
        print("Nothing to do. Use --probe, --probe-manifest, --key, --all, or --hf-id.", file=sys.stderr)
        return 2

    results = []
    for key in keys:
        if key not in sources:
            print(f"Unknown source key: {key}", file=sys.stderr)
            return 2
        source = sources[key]
        if source.get("auth") == "gated" and not (os.getenv("HF_TOKEN") or os.getenv("HUGGINGFACE_HUB_TOKEN")):
            print(f"[skip] {key}: gated dataset and no HF_TOKEN set.", file=sys.stderr)
            continue
        split_key = "eval_split" if args.eval else "train_split"
        split = source.get(split_key)
        if not split:
            print(f"[skip] {key}: no {split_key} configured.", file=sys.stderr)
            continue
        suffix = "eval" if args.eval else "raw"
        out_path = raw_dir / key / f"{suffix}.jsonl"
        max_rows = args.max_rows if args.max_rows is not None else _cap_for(source, defaults, args.profile)
        print(f"[download] {key}: {source['hf_id']} [{split}] -> {out_path} (max_rows={max_rows})")
        result = download_split(source["hf_id"], source.get("config"), split, out_path, max_rows)
        result["key"] = key
        results.append(result)
        print(json.dumps(result))

    print(json.dumps({"downloaded": results}, indent=2))
    return 0


if __name__ == "__main__":
    raise SystemExit(main())