Spaces:
Sleeping
Sleeping
| """Fetch, normalize, and split Layer 2 evaluation/training datasets. | |
| Downloads datasets for two classification tasks: | |
| - Prompt Injection (Model A): Gandalf, DeepSet, HackAPrompt, BIPIA, Enron | |
| - Malicious Intent (Model B): SpamAssassin Ham/Spam, Nazario Phishing, | |
| Fraudulent Email, Enron | |
| Local-First Strategy: | |
| For each dataset, the script checks ``data/raw/`` for manually downloaded | |
| files BEFORE attempting any network fetch. This is essential for gated | |
| HuggingFace datasets (HackAPrompt, BIPIA) that require authentication, | |
| and for very large datasets (Enron) that can OOM if loaded eagerly. | |
| All records are normalized into a common JSONL schema: | |
| {id, source, task, label, text_body, html_body} | |
| Stratified train/val/test splits (70/15/15) are generated per task, | |
| stratified by (source, label) to ensure proportional representation. | |
| This script never synthesizes content. All data comes from real sources. | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import email | |
| import email.policy | |
| import hashlib | |
| import io | |
| import json | |
| import logging | |
| import random | |
| import tarfile | |
| from collections import Counter | |
| from dataclasses import dataclass, field | |
| from pathlib import Path | |
| from typing import Any, Iterable, Iterator | |
| import pandas as pd | |
| import requests | |
| from sklearn.model_selection import train_test_split | |
| from tqdm import tqdm | |
| logger = logging.getLogger("fetch_datasets_l2") | |
| DEFAULT_OUTPUT_DIR = Path("data/l2_datasets") | |
| DEFAULT_RAW_DIR = Path("data/raw") | |
| DEFAULT_SEED = 42 | |
| # --- HuggingFace dataset identifiers --- | |
| GANDALF_ID = "Lakera/gandalf_ignore_instructions" | |
| DEEPSET_PI_ID = "deepset/prompt-injections" | |
| HACKAPROMPT_ID = "hackaprompt/hackaprompt-dataset" | |
| BIPIA_ID = "MAlmasabi/Indirect-Prompt-Injection-BIPIA-GPT" | |
| ENRON_ID = "SuccessfulCrab/enron" | |
| # --- SpamAssassin public corpus --- | |
| SPAMASSASSIN_BASE = "https://spamassassin.apache.org/old/publiccorpus" | |
| SPAMASSASSIN_FILES = { | |
| "ham": [ | |
| "20030228_easy_ham.tar.bz2", | |
| "20030228_easy_ham_2.tar.bz2", | |
| "20030228_hard_ham.tar.bz2", | |
| ], | |
| "spam": [ | |
| "20030228_spam.tar.bz2", | |
| "20030228_spam_2.tar.bz2", | |
| ], | |
| } | |
| # --- Nazario phishing corpus --- | |
| NAZARIO_URL = "https://monkey.org/~jose/phishing/phishing3.mbox" | |
| # Split ratios | |
| TRAIN_RATIO = 0.70 | |
| VAL_RATIO = 0.15 | |
| TEST_RATIO = 0.15 | |
| # ============================================================ | |
| # Local-First File Registry | |
| # ============================================================ | |
| # Maps each source to a list of candidate filenames (checked in order) | |
| # in the data/raw/ directory. The first matching file wins. | |
| LOCAL_FILE_REGISTRY: dict[str, list[str]] = { | |
| "gandalf": [ | |
| "gandalf.parquet", | |
| "gandalf.csv", | |
| "gandalf.jsonl", | |
| ], | |
| "deepset": [ | |
| # HuggingFace auto-download names | |
| "deepset_train.parquet", | |
| "deepset_test.parquet", | |
| "train-00000-of-00001-9564e8b05b4757ab.parquet", | |
| "test-00000-of-00001-701d16158af87368.parquet", | |
| # Manual download names | |
| "deepset.parquet", | |
| "deepset.csv", | |
| "deepset.jsonl", | |
| ], | |
| "hackaprompt": [ | |
| "hackaprompt.parquet", | |
| "hackaprompt.csv", | |
| "hackaprompt.jsonl", | |
| ], | |
| "bipia": [ | |
| "bipia.jsonl", | |
| "dataset_for_huggingface.jsonl", | |
| "bipia.parquet", | |
| "bipia.csv", | |
| ], | |
| "enron": [ | |
| "enron.parquet", | |
| "enron.csv", | |
| "enron.jsonl", | |
| ], | |
| "nazario": [ | |
| "nazario.mbox", | |
| "phishing3.mbox", | |
| ], | |
| "fraudulent": [ | |
| "fraudulent_email.csv", | |
| "fraudulent_email.jsonl", | |
| "fraudulent_email.parquet", | |
| ], | |
| } | |
| # Chunk size for reading large CSVs to prevent OOM | |
| CSV_CHUNK_SIZE = 10_000 | |
| class FetchConfig: | |
| """Runtime configuration for L2 dataset fetching.""" | |
| output_dir: Path | |
| raw_dir: Path | |
| seed: int | |
| timeout: int | |
| hf_token: str | None | |
| strict: bool | |
| # Per-source caps | |
| gandalf_count: int | |
| deepset_count: int | |
| hackaprompt_count: int | |
| bipia_count: int | |
| enron_pi_count: int | |
| enron_mi_count: int | |
| spamassassin_ham_count: int | |
| spamassassin_spam_count: int | |
| nazario_count: int | |
| fraudulent_count: int | |
| skip: set[str] = field(default_factory=set) | |
| # ============================================================ | |
| # Shared utilities | |
| # ============================================================ | |
| def write_jsonl(path: Path, records: Iterable[dict[str, Any]]) -> int: | |
| """Write records to JSONL, creating parent directories. | |
| Args: | |
| path: Output JSONL path. | |
| records: Records to write. | |
| Returns: | |
| Number of records written. | |
| """ | |
| count = 0 | |
| path.parent.mkdir(parents=True, exist_ok=True) | |
| with path.open("w", encoding="utf-8") as handle: | |
| for record in records: | |
| handle.write( | |
| json.dumps(record, ensure_ascii=False, sort_keys=True) + "\n" | |
| ) | |
| count += 1 | |
| logger.info("Wrote %s records to %s", count, path) | |
| return count | |
| def load_jsonl(path: Path) -> list[dict[str, Any]]: | |
| """Read all records from a JSONL file. | |
| Args: | |
| path: JSONL file path. | |
| Returns: | |
| List of parsed JSON objects. | |
| """ | |
| records = [] | |
| with path.open("r", encoding="utf-8", errors="replace") as handle: | |
| for line in handle: | |
| line = line.strip() | |
| if line: | |
| records.append(json.loads(line)) | |
| return records | |
| def reservoir_sample( | |
| rows: Iterable[dict[str, Any]], | |
| sample_count: int, | |
| rng: random.Random, | |
| description: str, | |
| ) -> list[dict[str, Any]]: | |
| """Uniformly sample from an iterable without loading all rows first. | |
| Uses Algorithm R (Vitter, 1985) for O(k) memory regardless of stream | |
| size. This is critical for large HF streaming datasets like Enron | |
| (500K+ rows) to avoid OOM. | |
| Args: | |
| rows: Iterable of records. | |
| sample_count: Maximum number of records to sample. | |
| rng: Random number generator. | |
| description: Label for progress bar. | |
| Returns: | |
| Sampled records, shuffled. | |
| """ | |
| sample: list[dict[str, Any]] = [] | |
| for seen, row in enumerate(tqdm(rows, desc=description, unit="rows"), start=1): | |
| if len(sample) < sample_count: | |
| sample.append(row) | |
| continue | |
| j = rng.randint(1, seen) | |
| if j <= sample_count: | |
| sample[j - 1] = row | |
| if len(sample) < sample_count: | |
| logger.warning( | |
| "%s only yielded %s/%s records", description, len(sample), sample_count | |
| ) | |
| rng.shuffle(sample) | |
| return sample | |
| def load_hf_dataset( | |
| dataset_id: str, config: FetchConfig, split: str = "train" | |
| ) -> Iterable[dict[str, Any]]: | |
| """Load a HuggingFace dataset as a stream. | |
| Always uses streaming=True to prevent loading the entire dataset | |
| into RAM. This is the ONLY way to safely handle large datasets | |
| like Enron (500K+ rows, 1.8GB+ in memory). | |
| Args: | |
| dataset_id: HuggingFace dataset identifier. | |
| config: Fetch configuration. | |
| split: Dataset split to load. | |
| Returns: | |
| Iterable of dataset rows (streamed, not materialized). | |
| """ | |
| from datasets import load_dataset | |
| return load_dataset( | |
| dataset_id, | |
| split=split, | |
| streaming=True, | |
| token=config.hf_token, | |
| ) | |
| def first_present(row: dict[str, Any], *keys: str) -> Any: | |
| """Return the first non-empty value from a row. | |
| Args: | |
| row: Data row. | |
| *keys: Keys to search in order. | |
| Returns: | |
| First non-empty value, or None. | |
| """ | |
| for key in keys: | |
| value = row.get(key) | |
| if value not in (None, ""): | |
| return value | |
| return None | |
| def looks_like_html(value: str) -> bool: | |
| """Return whether a body appears to contain HTML markup. | |
| Args: | |
| value: Text to check. | |
| Returns: | |
| True if the text contains HTML tags. | |
| """ | |
| lowered = value.lower() | |
| return "<html" in lowered or "<body" in lowered or "<table" in lowered | |
| def handle_source_error(source: str, exc: Exception, strict: bool) -> None: | |
| """Handle source download/load failures. | |
| Args: | |
| source: Source name. | |
| exc: Exception that occurred. | |
| strict: Whether to raise on errors. | |
| """ | |
| message = f"Could not fetch {source}: {exc}" | |
| if strict: | |
| raise RuntimeError(message) from exc | |
| logger.error(message) | |
| def stable_id(source: str, text: str) -> str: | |
| """Generate a stable deterministic ID from source and content. | |
| Args: | |
| source: Source name. | |
| text: Content to hash. | |
| Returns: | |
| ID string in format "source-hash8". | |
| """ | |
| digest = hashlib.sha256(text.encode("utf-8", errors="replace")).hexdigest()[:8] | |
| return f"{source}-{digest}" | |
| def build_record( | |
| record_id: str, | |
| source: str, | |
| task: str, | |
| label: int, | |
| text_body: str, | |
| html_body: str, | |
| ) -> dict[str, Any]: | |
| """Build a normalized L2 record. | |
| Args: | |
| record_id: Unique identifier. | |
| source: Dataset source name. | |
| task: "prompt_injection" or "malicious_intent". | |
| label: 0 (benign) or 1 (malicious/injection). | |
| text_body: Plain-text body. | |
| html_body: HTML body (empty if not available). | |
| Returns: | |
| Normalized record dict. | |
| """ | |
| return { | |
| "id": record_id, | |
| "source": source, | |
| "task": task, | |
| "label": label, | |
| "text_body": text_body[:50000], # Cap at 50KB to prevent bloat | |
| "html_body": html_body[:100000], # Cap at 100KB | |
| } | |
| # ============================================================ | |
| # Local-First file resolution | |
| # ============================================================ | |
| def find_local_files(source: str, raw_dir: Path) -> list[Path]: | |
| """Find all locally available files for a given dataset source. | |
| Searches the raw_dir for files matching the LOCAL_FILE_REGISTRY | |
| entries for the given source. Returns all matches (some datasets | |
| like DeepSet have separate train/test files). | |
| Args: | |
| source: Dataset source name (key in LOCAL_FILE_REGISTRY). | |
| raw_dir: Directory to search for local files. | |
| Returns: | |
| List of existing file paths, in registry order. | |
| """ | |
| candidates = LOCAL_FILE_REGISTRY.get(source, []) | |
| found = [] | |
| for filename in candidates: | |
| path = raw_dir / filename | |
| if path.exists(): | |
| found.append(path) | |
| return found | |
| def load_local_tabular( | |
| path: Path, max_rows: int | None = None | |
| ) -> Iterator[dict[str, Any]]: | |
| """Load rows from a local parquet, CSV, or JSONL file as an iterator. | |
| Uses chunked reading for CSVs to prevent OOM on large files. | |
| Parquet files are read in full (they're columnar and memory-mapped). | |
| Args: | |
| path: Path to the file. | |
| max_rows: Optional maximum number of rows to yield. | |
| Yields: | |
| Row dicts from the file. | |
| """ | |
| suffix = path.suffix.lower() | |
| yielded = 0 | |
| if suffix == ".parquet": | |
| df = pd.read_parquet(path) | |
| for _, row in df.iterrows(): | |
| if max_rows and yielded >= max_rows: | |
| return | |
| yield row.to_dict() | |
| yielded += 1 | |
| elif suffix == ".csv": | |
| # Chunked reading to prevent OOM on large CSVs | |
| for chunk in pd.read_csv( | |
| path, chunksize=CSV_CHUNK_SIZE, encoding="utf-8", | |
| on_bad_lines="skip", | |
| ): | |
| for _, row in chunk.iterrows(): | |
| if max_rows and yielded >= max_rows: | |
| return | |
| yield row.to_dict() | |
| yielded += 1 | |
| elif suffix == ".jsonl": | |
| with path.open("r", encoding="utf-8", errors="replace") as handle: | |
| for line in handle: | |
| if max_rows and yielded >= max_rows: | |
| return | |
| line = line.strip() | |
| if line: | |
| yield json.loads(line) | |
| yielded += 1 | |
| else: | |
| raise ValueError(f"Unsupported file format: {suffix}") | |
| # ============================================================ | |
| # Prompt Injection dataset fetchers | |
| # ============================================================ | |
| def fetch_gandalf(config: FetchConfig) -> list[dict[str, Any]]: | |
| """Fetch Gandalf Ignore Instructions dataset. | |
| Contains layered, multi-step injection bypass attempts | |
| from the Lakera Gandalf challenge. | |
| Strategy: local file first, then HuggingFace streaming fallback. | |
| Args: | |
| config: Fetch configuration. | |
| Returns: | |
| List of normalized records. | |
| """ | |
| rng = random.Random(config.seed) | |
| def normalize_row(idx: int, row: dict[str, Any]) -> dict[str, Any] | None: | |
| """Normalize a single Gandalf row into L2 schema.""" | |
| text = first_present(row, "text", "prompt", "instruction", "input") | |
| if not isinstance(text, str) or not text.strip(): | |
| return None | |
| return build_record( | |
| record_id=f"gandalf-{idx}", | |
| source="gandalf", | |
| task="prompt_injection", | |
| label=1, | |
| text_body=text.strip(), | |
| html_body="", | |
| ) | |
| # Local-first: check data/raw/ for pre-downloaded files | |
| local_files = find_local_files("gandalf", config.raw_dir) | |
| if local_files: | |
| logger.info("Loading Gandalf from local file: %s", local_files[0]) | |
| rows = load_local_tabular(local_files[0]) | |
| records = _normalize_stream(rows, normalize_row, config.gandalf_count, rng) | |
| if records: | |
| return records | |
| logger.warning("Local Gandalf file yielded 0 records; falling back to HF") | |
| # Network fallback: HuggingFace streaming | |
| try: | |
| dataset = load_hf_dataset(GANDALF_ID, config) | |
| except Exception as exc: | |
| handle_source_error("Gandalf", exc, config.strict) | |
| return [] | |
| def iter_records() -> Iterator[dict[str, Any]]: | |
| for idx, row in enumerate(dataset): | |
| record = normalize_row(idx, row) | |
| if record: | |
| yield record | |
| return reservoir_sample(iter_records(), config.gandalf_count, rng, "Gandalf") | |
| def fetch_deepset(config: FetchConfig) -> list[dict[str, Any]]: | |
| """Fetch DeepSet prompt-injections dataset. | |
| Contains curated injection vs. benign text pairs with labels. | |
| This is one of the few PI datasets with both positive AND negative samples. | |
| Strategy: local parquet files first (supports separate train/test files), | |
| then HuggingFace streaming fallback. | |
| Args: | |
| config: Fetch configuration. | |
| Returns: | |
| List of normalized records. | |
| """ | |
| rng = random.Random(config.seed) | |
| def normalize_row(idx: int, row: dict[str, Any]) -> dict[str, Any] | None: | |
| """Normalize a single DeepSet row into L2 schema.""" | |
| text = first_present(row, "text", "prompt", "input") | |
| if not isinstance(text, str) or not text.strip(): | |
| return None | |
| label_raw = row.get("label", None) | |
| if label_raw is None: | |
| return None | |
| label = int(label_raw) | |
| if label not in (0, 1): | |
| return None | |
| return build_record( | |
| record_id=f"deepset-{idx}", | |
| source="deepset", | |
| task="prompt_injection", | |
| label=label, | |
| text_body=text.strip(), | |
| html_body="", | |
| ) | |
| # Local-first: DeepSet may have separate train/test parquet files | |
| local_files = find_local_files("deepset", config.raw_dir) | |
| if local_files: | |
| logger.info( | |
| "Loading DeepSet from %d local file(s): %s", | |
| len(local_files), | |
| [f.name for f in local_files], | |
| ) | |
| all_rows: list[dict[str, Any]] = [] | |
| for lf in local_files: | |
| for row in load_local_tabular(lf): | |
| all_rows.append(row) | |
| records = _normalize_stream( | |
| iter(all_rows), normalize_row, config.deepset_count, rng, | |
| ) | |
| if records: | |
| return records | |
| logger.warning("Local DeepSet files yielded 0 records; falling back to HF") | |
| # Network fallback | |
| try: | |
| dataset = load_hf_dataset(DEEPSET_PI_ID, config) | |
| except Exception as exc: | |
| handle_source_error("DeepSet", exc, config.strict) | |
| return [] | |
| def iter_records() -> Iterator[dict[str, Any]]: | |
| for idx, row in enumerate(dataset): | |
| record = normalize_row(idx, row) | |
| if record: | |
| yield record | |
| return reservoir_sample(iter_records(), config.deepset_count, rng, "DeepSet") | |
| def fetch_hackaprompt_l2(config: FetchConfig) -> list[dict[str, Any]]: | |
| """Fetch HackAPrompt adversarial payloads for L2 evaluation. | |
| This is a GATED dataset on HuggingFace requiring HF_TOKEN. | |
| The local-first path is critical here: if hackaprompt.parquet | |
| exists in data/raw/, it is used directly without network access. | |
| Strategy: local parquet first, then HuggingFace streaming fallback. | |
| Args: | |
| config: Fetch configuration. | |
| Returns: | |
| List of normalized records. | |
| """ | |
| rng = random.Random(config.seed) | |
| def normalize_row(idx: int, row: dict[str, Any]) -> dict[str, Any] | None: | |
| """Normalize a single HackAPrompt row into L2 schema.""" | |
| payload = first_present(row, "user_input", "prompt") | |
| if not isinstance(payload, str) or not payload.strip(): | |
| return None | |
| return build_record( | |
| record_id=f"hackaprompt-{idx}", | |
| source="hackaprompt", | |
| task="prompt_injection", | |
| label=1, | |
| text_body=payload.strip(), | |
| html_body="", | |
| ) | |
| # Local-first: hackaprompt.parquet is 601K rows — use reservoir sampling | |
| local_files = find_local_files("hackaprompt", config.raw_dir) | |
| if local_files: | |
| logger.info("Loading HackAPrompt from local file: %s", local_files[0]) | |
| rows = load_local_tabular(local_files[0]) | |
| records = _normalize_stream( | |
| rows, normalize_row, config.hackaprompt_count, rng, | |
| ) | |
| if records: | |
| return records | |
| logger.warning( | |
| "Local HackAPrompt file yielded 0 records; falling back to HF" | |
| ) | |
| # Network fallback (requires HF_TOKEN for this gated dataset) | |
| try: | |
| dataset = load_hf_dataset(HACKAPROMPT_ID, config) | |
| except Exception as exc: | |
| handle_source_error("HackAPrompt", exc, config.strict) | |
| return [] | |
| def iter_records() -> Iterator[dict[str, Any]]: | |
| for idx, row in enumerate(dataset): | |
| record = normalize_row(idx, row) | |
| if record: | |
| yield record | |
| return reservoir_sample( | |
| iter_records(), config.hackaprompt_count, rng, "HackAPrompt" | |
| ) | |
| def fetch_bipia_l2(config: FetchConfig) -> list[dict[str, Any]]: | |
| """Fetch BIPIA email-context indirect injection payloads. | |
| These are specifically crafted for email context, making them | |
| highly relevant for our use case. This is a GATED dataset. | |
| Strategy: local JSONL/parquet first (dataset_for_huggingface.jsonl | |
| is the expected name), then HuggingFace streaming fallback. | |
| Args: | |
| config: Fetch configuration. | |
| Returns: | |
| List of normalized records. | |
| """ | |
| rng = random.Random(config.seed) | |
| def normalize_row(idx: int, row: dict[str, Any]) -> dict[str, Any] | None: | |
| """Normalize a single BIPIA row into L2 schema.""" | |
| payload = first_present( | |
| row, "context", "email", "text", "prompt", "payload" | |
| ) | |
| if not isinstance(payload, str) or not payload.strip(): | |
| return None | |
| # Skip explicitly benign-labeled rows | |
| if row.get("label") in (0, "0", False, "benign"): | |
| return None | |
| return build_record( | |
| record_id=f"bipia-{idx}", | |
| source="bipia", | |
| task="prompt_injection", | |
| label=1, | |
| text_body=payload.strip(), | |
| html_body="", | |
| ) | |
| # Local-first: BIPIA JSONL (70K rows) — use reservoir sampling | |
| local_files = find_local_files("bipia", config.raw_dir) | |
| if local_files: | |
| logger.info("Loading BIPIA from local file: %s", local_files[0]) | |
| rows = load_local_tabular(local_files[0]) | |
| records = _normalize_stream(rows, normalize_row, config.bipia_count, rng) | |
| if records: | |
| return records | |
| logger.warning("Local BIPIA file yielded 0 records; falling back to HF") | |
| # Network fallback (requires HF_TOKEN for this gated dataset) | |
| try: | |
| dataset = load_hf_dataset(BIPIA_ID, config) | |
| except Exception as exc: | |
| handle_source_error("BIPIA", exc, config.strict) | |
| return [] | |
| def iter_records() -> Iterator[dict[str, Any]]: | |
| for idx, row in enumerate(dataset): | |
| record = normalize_row(idx, row) | |
| if record: | |
| yield record | |
| return reservoir_sample(iter_records(), config.bipia_count, rng, "BIPIA") | |
| def fetch_enron_negative( | |
| config: FetchConfig, task: str, count: int | |
| ) -> list[dict[str, Any]]: | |
| """Fetch Enron emails as negative (benign) samples for a given task. | |
| CRITICAL: The Enron dataset is ~500K rows and ~1.8GB in RAM if loaded | |
| eagerly. This function MUST use either: | |
| - Local file with streaming iteration (load_local_tabular) | |
| - HuggingFace streaming=True with reservoir sampling | |
| Never call load_dataset() without streaming=True for Enron. | |
| Args: | |
| config: Fetch configuration. | |
| task: Task name ("prompt_injection" or "malicious_intent"). | |
| count: Number of samples to fetch. | |
| Returns: | |
| List of normalized negative records. | |
| """ | |
| rng = random.Random(config.seed + hash(task)) | |
| source_tag = f"enron_{task[:2]}" | |
| def normalize_row(idx: int, row: dict[str, Any]) -> dict[str, Any] | None: | |
| """Normalize a single Enron row into L2 schema.""" | |
| message = first_present(row, "message", "text", "body", "email") | |
| if not isinstance(message, str) or not message.strip(): | |
| return None | |
| html_body = message if looks_like_html(message) else "" | |
| return build_record( | |
| record_id=f"{source_tag}-{idx}", | |
| source=source_tag, | |
| task=task, | |
| label=0, | |
| text_body=message.strip(), | |
| html_body=html_body, | |
| ) | |
| # Local-first: check for pre-downloaded Enron file | |
| local_files = find_local_files("enron", config.raw_dir) | |
| if local_files: | |
| logger.info( | |
| "Loading Enron (%s) from local file: %s", task, local_files[0] | |
| ) | |
| rows = load_local_tabular(local_files[0]) | |
| records = _normalize_stream(rows, normalize_row, count, rng) | |
| if records: | |
| return records | |
| logger.warning( | |
| "Local Enron file yielded 0 records for %s; falling back to HF", | |
| task, | |
| ) | |
| # Network fallback: MUST use streaming=True to prevent OOM | |
| try: | |
| dataset = load_hf_dataset(ENRON_ID, config) | |
| except Exception as exc: | |
| handle_source_error(f"Enron ({task})", exc, config.strict) | |
| return [] | |
| def iter_records() -> Iterator[dict[str, Any]]: | |
| for idx, row in enumerate(dataset): | |
| record = normalize_row(idx, row) | |
| if record: | |
| yield record | |
| return reservoir_sample(iter_records(), count, rng, f"Enron ({task})") | |
| # ============================================================ | |
| # Shared normalization helper | |
| # ============================================================ | |
| def _normalize_stream( | |
| rows: Iterable[dict[str, Any]], | |
| normalize_fn: Any, | |
| cap: int, | |
| rng: random.Random, | |
| ) -> list[dict[str, Any]]: | |
| """Normalize rows through a function and reservoir-sample to cap. | |
| Combines normalization and sampling in a single streaming pass | |
| to avoid materializing the entire dataset in memory. | |
| Args: | |
| rows: Raw row iterator. | |
| normalize_fn: Callable(idx, row) -> normalized record or None. | |
| cap: Maximum number of records to keep. | |
| rng: Random number generator for sampling. | |
| Returns: | |
| List of normalized, sampled records. | |
| """ | |
| def iter_normalized() -> Iterator[dict[str, Any]]: | |
| for idx, row in enumerate(rows): | |
| record = normalize_fn(idx, row) | |
| if record is not None: | |
| yield record | |
| return reservoir_sample(iter_normalized(), cap, rng, "local-file") | |
| # ============================================================ | |
| # Malicious Intent dataset fetchers | |
| # ============================================================ | |
| def _parse_mbox_email(raw_bytes: bytes) -> tuple[str, str]: | |
| """Parse a raw email into text_body and html_body. | |
| Preserves HTML body when present in multipart MIME messages. | |
| This is critical for Layer 1 evaluation, which needs the original | |
| HTML to detect hidden content, CSS hiding, and other structural attacks. | |
| Args: | |
| raw_bytes: Raw email bytes. | |
| Returns: | |
| Tuple of (text_body, html_body). | |
| """ | |
| try: | |
| msg = email.message_from_bytes(raw_bytes, policy=email.policy.default) | |
| except Exception: | |
| # Fall back to string decoding if email parsing fails | |
| text = raw_bytes.decode("utf-8", errors="replace") | |
| return text, "" | |
| text_body = "" | |
| html_body = "" | |
| if msg.is_multipart(): | |
| for part in msg.walk(): | |
| content_type = part.get_content_type() | |
| try: | |
| payload = part.get_content() | |
| except Exception: | |
| continue | |
| if not isinstance(payload, str): | |
| continue | |
| if content_type == "text/plain" and not text_body: | |
| text_body = payload | |
| elif content_type == "text/html" and not html_body: | |
| html_body = payload | |
| else: | |
| try: | |
| payload = msg.get_content() | |
| except Exception: | |
| payload = raw_bytes.decode("utf-8", errors="replace") | |
| if isinstance(payload, str): | |
| content_type = msg.get_content_type() | |
| if content_type == "text/html": | |
| html_body = payload | |
| else: | |
| text_body = payload | |
| # If we only have HTML, derive text as fallback (keep html_body intact) | |
| if html_body and not text_body: | |
| text_body = html_body | |
| return text_body, html_body | |
| def fetch_spamassassin(config: FetchConfig) -> list[dict[str, Any]]: | |
| """Fetch SpamAssassin public corpus (ham + spam). | |
| Downloads tar.bz2 archives from Apache and parses individual | |
| email files. Ham emails are labeled 0 (benign), spam as 1 (malicious). | |
| HTML body is preserved from MIME multipart messages. | |
| Args: | |
| config: Fetch configuration. | |
| Returns: | |
| List of normalized records. | |
| """ | |
| records: list[dict[str, Any]] = [] | |
| for label_name, archives in SPAMASSASSIN_FILES.items(): | |
| label = 0 if label_name == "ham" else 1 | |
| cap = ( | |
| config.spamassassin_ham_count | |
| if label == 0 | |
| else config.spamassassin_spam_count | |
| ) | |
| archive_records: list[dict[str, Any]] = [] | |
| for archive_name in archives: | |
| url = f"{SPAMASSASSIN_BASE}/{archive_name}" | |
| try: | |
| logger.info("Downloading SpamAssassin %s ...", archive_name) | |
| response = requests.get(url, timeout=config.timeout) | |
| response.raise_for_status() | |
| except requests.RequestException as exc: | |
| handle_source_error( | |
| f"SpamAssassin {archive_name}", exc, config.strict | |
| ) | |
| continue | |
| try: | |
| with tarfile.open( | |
| fileobj=io.BytesIO(response.content), mode="r:bz2" | |
| ) as tar: | |
| for member in tar.getmembers(): | |
| if not member.isfile(): | |
| continue | |
| name = member.name.split("/")[-1] | |
| if name.startswith(".") or name in ("cmds", "README"): | |
| continue | |
| try: | |
| raw = tar.extractfile(member) | |
| if raw is None: | |
| continue | |
| raw_bytes = raw.read() | |
| except Exception: | |
| continue | |
| text_body, html_body = _parse_mbox_email(raw_bytes) | |
| if not text_body.strip(): | |
| continue | |
| source_tag = f"spamassassin_{label_name}" | |
| archive_records.append( | |
| build_record( | |
| record_id=stable_id(source_tag, text_body), | |
| source=source_tag, | |
| task="malicious_intent", | |
| label=label, | |
| text_body=text_body.strip(), | |
| html_body=html_body, | |
| ) | |
| ) | |
| except Exception as exc: | |
| handle_source_error( | |
| f"SpamAssassin archive {archive_name}", exc, config.strict | |
| ) | |
| # Sample down to cap | |
| rng = random.Random(config.seed) | |
| if len(archive_records) > cap: | |
| rng.shuffle(archive_records) | |
| archive_records = archive_records[:cap] | |
| logger.info( | |
| "SpamAssassin %s: %s records (cap=%s)", | |
| label_name, | |
| len(archive_records), | |
| cap, | |
| ) | |
| records.extend(archive_records) | |
| return records | |
| def fetch_nazario(config: FetchConfig) -> list[dict[str, Any]]: | |
| """Fetch Nazario phishing corpus. | |
| Downloads phishing emails from Jose Nazario's public collection. | |
| These are real phishing emails, labeled as malicious intent positive. | |
| HTML body is preserved from MIME multipart messages. | |
| Strategy: local mbox file first, then HTTP download fallback. | |
| Args: | |
| config: Fetch configuration. | |
| Returns: | |
| List of normalized records. | |
| """ | |
| raw_content: bytes | None = None | |
| # Local-first: check for pre-downloaded mbox | |
| local_files = find_local_files("nazario", config.raw_dir) | |
| if local_files: | |
| logger.info("Loading Nazario from local file: %s", local_files[0]) | |
| raw_content = local_files[0].read_bytes() | |
| # Network fallback | |
| if raw_content is None: | |
| try: | |
| logger.info("Downloading Nazario phishing corpus...") | |
| response = requests.get(NAZARIO_URL, timeout=config.timeout) | |
| response.raise_for_status() | |
| raw_content = response.content | |
| except requests.RequestException as exc: | |
| handle_source_error("Nazario phishing corpus", exc, config.strict) | |
| return [] | |
| records: list[dict[str, Any]] = [] | |
| raw_emails = raw_content.split(b"\nFrom ") | |
| for idx, raw in enumerate(raw_emails): | |
| if idx > 0: | |
| raw = b"From " + raw | |
| text_body, html_body = _parse_mbox_email(raw) | |
| if not text_body.strip() or len(text_body.strip()) < 20: | |
| continue | |
| records.append( | |
| build_record( | |
| record_id=f"nazario-{idx}", | |
| source="nazario", | |
| task="malicious_intent", | |
| label=1, | |
| text_body=text_body.strip(), | |
| html_body=html_body, | |
| ) | |
| ) | |
| rng = random.Random(config.seed) | |
| if len(records) > config.nazario_count: | |
| rng.shuffle(records) | |
| records = records[: config.nazario_count] | |
| logger.info("Nazario: %s records", len(records)) | |
| return records | |
| def fetch_fraudulent_email(config: FetchConfig) -> list[dict[str, Any]]: | |
| """Fetch fraudulent email (419 scam) dataset. | |
| Strategy: local CSV/JSONL first, then HuggingFace mirror fallback. | |
| Args: | |
| config: Fetch configuration. | |
| Returns: | |
| List of normalized records. | |
| """ | |
| rng = random.Random(config.seed) | |
| def normalize_row(idx: int, row: dict[str, Any]) -> dict[str, Any] | None: | |
| """Normalize a single fraudulent email row.""" | |
| text = first_present( | |
| row, "text", "body", "email", "message", "content" | |
| ) | |
| if not isinstance(text, str) or not text.strip(): | |
| return None | |
| label_raw = row.get("label", None) | |
| if label_raw is not None: | |
| label = int(label_raw) | |
| if label not in (0, 1): | |
| label = 1 | |
| else: | |
| label = 1 # Entire dataset is fraudulent | |
| return build_record( | |
| record_id=f"fraudulent-{idx}", | |
| source="fraudulent_email", | |
| task="malicious_intent", | |
| label=label, | |
| text_body=text.strip(), | |
| html_body="", | |
| ) | |
| # Local-first: check for pre-downloaded CSV/JSONL | |
| local_files = find_local_files("fraudulent", config.raw_dir) | |
| if local_files: | |
| logger.info( | |
| "Loading Fraudulent Email from local file: %s", local_files[0] | |
| ) | |
| rows = load_local_tabular(local_files[0]) | |
| records = _normalize_stream( | |
| rows, normalize_row, config.fraudulent_count, rng, | |
| ) | |
| if records: | |
| return records | |
| logger.warning( | |
| "Local fraudulent email file yielded 0 records; falling back to HF" | |
| ) | |
| # Network fallback: try known HuggingFace mirrors | |
| hf_ids = [ | |
| "ealvaradob/phishing-dataset", | |
| "talby/fraudulent-email", | |
| ] | |
| for hf_id in hf_ids: | |
| try: | |
| dataset = load_hf_dataset(hf_id, config) | |
| def iter_records() -> Iterator[dict[str, Any]]: | |
| for idx, row in enumerate(dataset): | |
| record = normalize_row(idx, row) | |
| if record: | |
| yield record | |
| records = reservoir_sample( | |
| iter_records(), | |
| config.fraudulent_count, | |
| rng, | |
| f"Fraudulent ({hf_id})", | |
| ) | |
| if records: | |
| logger.info( | |
| "Loaded %s fraudulent email records from %s", | |
| len(records), | |
| hf_id, | |
| ) | |
| return records | |
| except Exception as exc: | |
| logger.warning("Could not load %s: %s", hf_id, exc) | |
| continue | |
| handle_source_error( | |
| "Fraudulent email corpus", | |
| RuntimeError("No accessible local file or HF mirror found"), | |
| config.strict, | |
| ) | |
| return [] | |
| # ============================================================ | |
| # Split generation | |
| # ============================================================ | |
| def create_stratified_splits( | |
| records: list[dict[str, Any]], | |
| task: str, | |
| output_dir: Path, | |
| seed: int, | |
| ) -> dict[str, int]: | |
| """Create stratified train/val/test splits for one task. | |
| Stratifies by (source, label) to ensure each source contributes | |
| proportionally to all three splits. | |
| Args: | |
| records: All records for this task. | |
| task: Task name for file naming. | |
| output_dir: Directory for output JSONL files. | |
| seed: Random seed for reproducibility. | |
| Returns: | |
| Dict with split names and record counts. | |
| """ | |
| if not records: | |
| logger.warning("No records for task %s; skipping split creation", task) | |
| return {} | |
| # Create stratification key | |
| strat_keys = [f"{r['source']}_{r['label']}" for r in records] | |
| # Check minimum samples per stratum for stratification | |
| key_counts = Counter(strat_keys) | |
| min_count = min(key_counts.values()) | |
| if min_count < 3: | |
| logger.warning( | |
| "Task %s: Some strata have < 3 samples. " | |
| "Falling back to label-only stratification.", | |
| task, | |
| ) | |
| strat_keys = [str(r["label"]) for r in records] | |
| try: | |
| # First split: train vs. (val+test) | |
| train_records, valtest_records, _, valtest_strat = train_test_split( | |
| records, | |
| strat_keys, | |
| test_size=(VAL_RATIO + TEST_RATIO), | |
| random_state=seed, | |
| stratify=strat_keys, | |
| ) | |
| # Second split: val vs. test (50/50 of the remaining) | |
| relative_test_size = TEST_RATIO / (VAL_RATIO + TEST_RATIO) | |
| val_records, test_records = train_test_split( | |
| valtest_records, | |
| test_size=relative_test_size, | |
| random_state=seed, | |
| stratify=valtest_strat, | |
| ) | |
| except ValueError: | |
| logger.warning( | |
| "Stratified split failed for task %s; using random split", task | |
| ) | |
| rng = random.Random(seed) | |
| shuffled = list(records) | |
| rng.shuffle(shuffled) | |
| n = len(shuffled) | |
| n_train = int(n * TRAIN_RATIO) | |
| n_val = int(n * VAL_RATIO) | |
| train_records = shuffled[:n_train] | |
| val_records = shuffled[n_train : n_train + n_val] | |
| test_records = shuffled[n_train + n_val :] | |
| task_dir = output_dir / task | |
| counts = {} | |
| for split_name, split_records in [ | |
| ("train", train_records), | |
| ("val", val_records), | |
| ("test", test_records), | |
| ]: | |
| n = write_jsonl(task_dir / f"{split_name}.jsonl", split_records) | |
| counts[split_name] = n | |
| for split_name, split_records in [ | |
| ("train", train_records), | |
| ("val", val_records), | |
| ("test", test_records), | |
| ]: | |
| source_dist = Counter(r["source"] for r in split_records) | |
| label_dist = Counter(r["label"] for r in split_records) | |
| logger.info( | |
| "%s/%s: %s records | labels=%s | sources=%s", | |
| task, | |
| split_name, | |
| len(split_records), | |
| dict(label_dist), | |
| dict(source_dist), | |
| ) | |
| return counts | |
| # ============================================================ | |
| # Manifest and reporting | |
| # ============================================================ | |
| def write_manifest( | |
| config: FetchConfig, | |
| raw_counts: dict[str, int], | |
| split_counts: dict[str, dict[str, int]], | |
| load_methods: dict[str, str], | |
| ) -> None: | |
| """Write fetch and split provenance manifest. | |
| Args: | |
| config: Fetch configuration. | |
| raw_counts: Per-source record counts before splitting. | |
| split_counts: Per-task split counts. | |
| load_methods: Per-source load method ("local" or "network"). | |
| """ | |
| manifest = { | |
| "zero_synthesis": True, | |
| "seed": config.seed, | |
| "raw_dir": str(config.raw_dir), | |
| "sources": { | |
| "gandalf": GANDALF_ID, | |
| "deepset": DEEPSET_PI_ID, | |
| "hackaprompt": HACKAPROMPT_ID, | |
| "bipia": BIPIA_ID, | |
| "enron": ENRON_ID, | |
| "spamassassin": SPAMASSASSIN_BASE, | |
| "nazario": NAZARIO_URL, | |
| }, | |
| "raw_counts": raw_counts, | |
| "load_methods": load_methods, | |
| "split_counts": split_counts, | |
| "split_ratios": { | |
| "train": TRAIN_RATIO, | |
| "val": VAL_RATIO, | |
| "test": TEST_RATIO, | |
| }, | |
| } | |
| manifest_path = config.output_dir / "dataset_manifest.json" | |
| manifest_path.write_text( | |
| json.dumps(manifest, indent=2, sort_keys=True), encoding="utf-8" | |
| ) | |
| logger.info("Wrote manifest to %s", manifest_path) | |
| # ============================================================ | |
| # CLI | |
| # ============================================================ | |
| def parse_args() -> argparse.Namespace: | |
| """Parse command-line arguments.""" | |
| parser = argparse.ArgumentParser(description=__doc__) | |
| parser.add_argument("--output-dir", type=Path, default=DEFAULT_OUTPUT_DIR) | |
| parser.add_argument( | |
| "--raw-dir", | |
| type=Path, | |
| default=DEFAULT_RAW_DIR, | |
| help=( | |
| "Directory containing manually downloaded dataset files. " | |
| "Checked BEFORE any network fetch. Default: data/raw/" | |
| ), | |
| ) | |
| parser.add_argument("--seed", type=int, default=DEFAULT_SEED) | |
| parser.add_argument("--timeout", type=int, default=120) | |
| parser.add_argument("--hf-token", default=None) | |
| parser.add_argument("--strict", action="store_true") | |
| parser.add_argument("--skip", action="append", default=[]) | |
| # Per-source caps | |
| parser.add_argument("--gandalf-count", type=int, default=2000) | |
| parser.add_argument("--deepset-count", type=int, default=1000) | |
| parser.add_argument("--hackaprompt-count", type=int, default=3000) | |
| parser.add_argument("--bipia-count", type=int, default=2000) | |
| parser.add_argument("--enron-pi-count", type=int, default=5000) | |
| parser.add_argument("--enron-mi-count", type=int, default=5000) | |
| parser.add_argument("--spamassassin-ham-count", type=int, default=4000) | |
| parser.add_argument("--spamassassin-spam-count", type=int, default=3000) | |
| parser.add_argument("--nazario-count", type=int, default=2000) | |
| parser.add_argument("--fraudulent-count", type=int, default=2000) | |
| return parser.parse_args() | |
| def build_config(args: argparse.Namespace) -> FetchConfig: | |
| """Build immutable fetch configuration from CLI args. | |
| Args: | |
| args: Parsed CLI arguments. | |
| Returns: | |
| Frozen FetchConfig dataclass. | |
| """ | |
| return FetchConfig( | |
| output_dir=args.output_dir, | |
| raw_dir=args.raw_dir, | |
| seed=args.seed, | |
| timeout=args.timeout, | |
| hf_token=args.hf_token, | |
| strict=args.strict, | |
| gandalf_count=args.gandalf_count, | |
| deepset_count=args.deepset_count, | |
| hackaprompt_count=args.hackaprompt_count, | |
| bipia_count=args.bipia_count, | |
| enron_pi_count=args.enron_pi_count, | |
| enron_mi_count=args.enron_mi_count, | |
| spamassassin_ham_count=args.spamassassin_ham_count, | |
| spamassassin_spam_count=args.spamassassin_spam_count, | |
| nazario_count=args.nazario_count, | |
| fraudulent_count=args.fraudulent_count, | |
| skip=set(args.skip), | |
| ) | |
| def main() -> None: | |
| """Fetch all L2 datasets, normalize, and create splits.""" | |
| logging.basicConfig(level=logging.INFO, format="%(levelname)s %(message)s") | |
| args = parse_args() | |
| config = build_config(args) | |
| config.output_dir.mkdir(parents=True, exist_ok=True) | |
| config.raw_dir.mkdir(parents=True, exist_ok=True) | |
| # Log local-first status | |
| logger.info("Raw directory: %s", config.raw_dir.resolve()) | |
| for source, filenames in LOCAL_FILE_REGISTRY.items(): | |
| found = find_local_files(source, config.raw_dir) | |
| if found: | |
| logger.info(" [LOCAL] %s: %s", source, [f.name for f in found]) | |
| else: | |
| logger.info(" [NETWORK] %s: no local files found", source) | |
| # Track which method was used per source | |
| load_methods: dict[str, str] = {} | |
| # --- Prompt Injection datasets --- | |
| pi_records: list[dict[str, Any]] = [] | |
| pi_raw_counts: dict[str, int] = {} | |
| pi_fetchers: dict[str, Any] = { | |
| "gandalf": lambda: fetch_gandalf(config), | |
| "deepset": lambda: fetch_deepset(config), | |
| "hackaprompt": lambda: fetch_hackaprompt_l2(config), | |
| "bipia": lambda: fetch_bipia_l2(config), | |
| "enron_pi": lambda: fetch_enron_negative( | |
| config, "prompt_injection", config.enron_pi_count | |
| ), | |
| } | |
| for name, fetcher in pi_fetchers.items(): | |
| if name in config.skip: | |
| logger.info("Skipping %s", name) | |
| continue | |
| try: | |
| # Check if local files exist for load_method tracking | |
| local_key = name.replace("_pi", "").replace("_mi", "") | |
| has_local = bool(find_local_files(local_key, config.raw_dir)) | |
| records = fetcher() | |
| pi_raw_counts[name] = len(records) | |
| pi_records.extend(records) | |
| load_methods[name] = "local" if has_local and records else "network" | |
| logger.info("Fetched %s: %s records", name, len(records)) | |
| except Exception as exc: | |
| handle_source_error(name, exc, config.strict) | |
| pi_raw_counts[name] = 0 | |
| load_methods[name] = "failed" | |
| # --- Malicious Intent datasets --- | |
| mi_records: list[dict[str, Any]] = [] | |
| mi_raw_counts: dict[str, int] = {} | |
| mi_fetchers: dict[str, Any] = { | |
| "spamassassin": lambda: fetch_spamassassin(config), | |
| "nazario": lambda: fetch_nazario(config), | |
| "fraudulent_email": lambda: fetch_fraudulent_email(config), | |
| "enron_mi": lambda: fetch_enron_negative( | |
| config, "malicious_intent", config.enron_mi_count | |
| ), | |
| } | |
| for name, fetcher in mi_fetchers.items(): | |
| if name in config.skip: | |
| logger.info("Skipping %s", name) | |
| continue | |
| try: | |
| local_key = name.replace("_mi", "").replace("_email", "") | |
| has_local = bool(find_local_files(local_key, config.raw_dir)) | |
| records = fetcher() | |
| mi_raw_counts[name] = len(records) | |
| mi_records.extend(records) | |
| load_methods[name] = "local" if has_local and records else "network" | |
| logger.info("Fetched %s: %s records", name, len(records)) | |
| except Exception as exc: | |
| handle_source_error(name, exc, config.strict) | |
| mi_raw_counts[name] = 0 | |
| load_methods[name] = "failed" | |
| # --- Write raw combined JSONL (before splitting) --- | |
| raw_out = config.output_dir / "raw" | |
| if pi_records: | |
| write_jsonl(raw_out / "prompt_injection_all.jsonl", pi_records) | |
| if mi_records: | |
| write_jsonl(raw_out / "malicious_intent_all.jsonl", mi_records) | |
| # --- Create stratified splits --- | |
| split_counts = {} | |
| if pi_records: | |
| logger.info( | |
| "Creating prompt_injection splits (%s total records)...", | |
| len(pi_records), | |
| ) | |
| split_counts["prompt_injection"] = create_stratified_splits( | |
| pi_records, "prompt_injection", config.output_dir, config.seed | |
| ) | |
| if mi_records: | |
| logger.info( | |
| "Creating malicious_intent splits (%s total records)...", | |
| len(mi_records), | |
| ) | |
| split_counts["malicious_intent"] = create_stratified_splits( | |
| mi_records, "malicious_intent", config.output_dir, config.seed | |
| ) | |
| # --- Summary --- | |
| raw_counts = {**pi_raw_counts, **mi_raw_counts} | |
| write_manifest(config, raw_counts, split_counts, load_methods) | |
| print("\n=== L2 Dataset Fetch Summary ===") | |
| print(f"\nRaw directory: {config.raw_dir.resolve()}") | |
| print(f"\nPrompt Injection: {len(pi_records)} total records") | |
| for source, count in pi_raw_counts.items(): | |
| method = load_methods.get(source, "?") | |
| print(f" {source}: {count} [{method}]") | |
| print(f"\nMalicious Intent: {len(mi_records)} total records") | |
| for source, count in mi_raw_counts.items(): | |
| method = load_methods.get(source, "?") | |
| print(f" {source}: {count} [{method}]") | |
| print("\nSplit counts:") | |
| print(json.dumps(split_counts, indent=2)) | |
| if __name__ == "__main__": | |
| main() | |