KChad/Prompt-Injection-RL-environment
1
1from __future__ import annotations2 3import argparse4import sys5from pathlib import Path6 7ROOT = Path(__file__).resolve().parents[1]8sys.path.insert(0, str(ROOT))9 10from env.llmail_curator import (11 DIFFICULTIES,12 curate_rows,13 load_llmail_contexts,14 load_rows_from_hf,15 load_rows_from_jsonl,16 write_curated_scenarios,17)18 19DEFAULT_OUTPUT_DIR = ROOT / "env" / "data" / "scenarios" / "curated"20 21 22def parse_args() -> argparse.Namespace:23 parser = argparse.ArgumentParser(24 description="Curate LLMail submissions into local easy/medium/hard scenario files for the IDPI mail benchmark."25 )26 parser.add_argument("--phase", action="append", choices=["Phase1", "Phase2"], help="Repeat to choose one or both Hugging Face splits. Defaults to both.")27 parser.add_argument("--input-jsonl-phase1", type=Path, help="Optional local Phase1 JSONL file.")28 parser.add_argument("--input-jsonl-phase2", type=Path, help="Optional local Phase2 JSONL file.")29 parser.add_argument("--scenarios-json", type=Path, help="Optional local copy of LLMail scenarios.json.")30 parser.add_argument("--output-dir", type=Path, default=DEFAULT_OUTPUT_DIR)31 parser.add_argument("--prefix", default="")32 parser.add_argument("--max-per-difficulty", type=int, default=50)33 parser.add_argument("--sample-limit", type=int, default=None, help="Optional cap per split for quick dry runs.")34 parser.add_argument("--seed", type=int, default=42)35 parser.add_argument("--overwrite", action="store_true")36 return parser.parse_args()37 38 39def rows_for_phase(phase: str, args: argparse.Namespace):40 if phase == "Phase1" and args.input_jsonl_phase1:41 return load_rows_from_jsonl(args.input_jsonl_phase1, sample_limit=args.sample_limit)42 if phase == "Phase2" and args.input_jsonl_phase2:43 return load_rows_from_jsonl(args.input_jsonl_phase2, sample_limit=args.sample_limit)44 return load_rows_from_hf(phase, sample_limit=args.sample_limit)45 46 47def main() -> None:48 args = parse_args()49 phases = args.phase or ["Phase1", "Phase2"]50 contexts = load_llmail_contexts(local_path=args.scenarios_json)51 52 merged = {difficulty: [] for difficulty in DIFFICULTIES}53 for phase in phases:54 curated = curate_rows(55 rows_for_phase(phase, args),56 phase=phase,57 contexts=contexts,58 max_per_difficulty=args.max_per_difficulty,59 seed=args.seed,60 )61 for difficulty in DIFFICULTIES:62 merged[difficulty].extend(curated[difficulty])63 64 written = write_curated_scenarios(65 merged,66 output_dir=args.output_dir,67 prefix=args.prefix,68 overwrite=args.overwrite,69 )70 71 for path in written:72 with path.open("r", encoding="utf-8-sig") as handle:73 count = sum(1 for _ in handle)74 print(f"wrote {count:>3} scenarios -> {path}")75 76 77if __name__ == "__main__":78 main()79 