import argparse import copy import json import os import random import re import sys from pathlib import Path from typing import Callable, Dict, Iterable, List import numpy as np proj_path = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) sys.path.insert(0, proj_path) from lib_code import attack_defense_tools as attack_tools # noqa: E402 from lib_code.attack_defense_tools import ignore, naive, none # noqa: E402 RANDOM_SEED = 42 SOURCE_FILES = { "squad": "1_raw_dataset/topicattack/data/crafted_instruction_data_squad_injection_qa.json", "tri": "1_raw_dataset/topicattack/data/crafted_instruction_data_tri_injection_qa.json", } SYSTEM_PROMPT_FILE = "1_raw_dataset/topicattack/prompts/generator_system_prompt.txt" CONFIGS = ( "prompt_based_separator", "native_tool_response_only", "native_tool_empty_query", ) TRAIN_ATTACK_SPECS = ( ("none", "none"), ("naive", "left"), ("naive", "right"), ("ignore", "left"), ("ignore", "right"), ) ATTACK_MAP: Dict[str, Callable] = { "none": none, "naive": naive, "ignore": ignore, } SIDE_MAP = { "none": "end", "left": "start", "right": "end", } def load_json(path: Path): with path.open("r", encoding="utf-8") as f: return json.load(f) def load_text(path: Path) -> str: with path.open("r", encoding="utf-8") as f: return f.read() def clean_input_text(text: str) -> str: text = text.replace("[DOC]", "").replace("[TLE]", "").replace("[PAR]", "") text = text.replace("\n", "") return re.sub(" +", " ", text).strip() def normalize_item(item: dict) -> dict: normalized = copy.deepcopy(item) normalized["input"] = clean_input_text(normalized.get("input") or "") return normalized def apply_attack(item: dict, attack: str, side_label: str) -> dict: attack_fn = ATTACK_MAP[attack] attacked = copy.deepcopy(item) attacked["add_mark"] = "inst" return attack_fn(attacked, side=SIDE_MAP[side_label], model=None, additional_injections=[]) def build_messages(item: dict, system_prompt: str, config: str) -> List[dict]: instruction = item["instruction"] input_data = item.get("input") or "" if config == "prompt_based_separator": user_input = ( "\n{instruction}\n<\\Instruction Area>\n\n" "\n{data}\n<\\Data Area>\n" ).format(instruction=instruction, data=input_data) return [ {"role": "system", "content": system_prompt}, {"role": "user", "content": user_input}, ] if config == "native_tool_response_only": return [ {"role": "system", "content": system_prompt}, {"role": "user", "content": instruction}, {"role": "tool", "content": input_data}, ] if config == "native_tool_empty_query": tool_call_id = "call_context" return [ {"role": "system", "content": system_prompt}, {"role": "user", "content": instruction}, { "role": "assistant", "content": "", "tool_calls": [ { "id": tool_call_id, "type": "function", "function": {"name": "get_context", "arguments": "{}"}, } ], }, {"role": "tool", "tool_call_id": tool_call_id, "content": input_data}, ] raise ValueError(f"Unsupported config: {config}") def iter_records( source: str, config: str, data: Iterable[dict], system_prompt: str, data_size: int, ): if data_size > 0: data = list(data)[:data_size] for item_idx, raw_item in enumerate(data): base_item = normalize_item(raw_item) for attack, side_label in TRAIN_ATTACK_SPECS: attacked_item = apply_attack(base_item, attack, side_label) yield { "id": f"{source}_{item_idx:04d}_{config}_{attack}_{side_label}", "source": source, "config": config, "attack": attack, "side": side_label, "original_index": item_idx, "instruction": base_item["instruction"], "input": attacked_item.get("input") or "", "expected_output": base_item.get("output", ""), "injection": base_item.get("injection", ""), "injection_output": base_item.get("injection_output", ""), "messages": build_messages(attacked_item, system_prompt, config), } def generate_single_turn(base_dir: Path, output_dir: Path, data_size: int = -1): random.seed(RANDOM_SEED) system_prompt = load_text(base_dir / SYSTEM_PROMPT_FILE) output_dir.mkdir(parents=True, exist_ok=True) for source, rel_path in SOURCE_FILES.items(): data = load_json(base_dir / rel_path) for config in CONFIGS: attack_tools._random = random.Random(42) attack_tools.np_random = np.random.default_rng(seed=123) records = list(iter_records(source, config, data, system_prompt, data_size)) out_path = output_dir / f"{source}_{config}.json" with out_path.open("w", encoding="utf-8") as f: json.dump(records, f, ensure_ascii=False, indent=2) print(f"Saved {len(records)} records -> {out_path}") def main(): parser = argparse.ArgumentParser(description="Generate single-turn training message JSON files.") parser.add_argument( "--base-dir", default="/home/hujk/gitrs/Paper2026/SortedCode2", help="SortedCode2 base directory. Raw datasets are read from {base-dir}/1_raw_dataset.", ) parser.add_argument( "--output-dir", default="/home/hujk/gitrs/Paper2026/SortedCode2/3-1_model_training_data_gen/single_turn", help="Directory for the six generated source x config JSON files.", ) parser.add_argument("--data-size", type=int, default=-1, help="Debug only: limit source rows.") args = parser.parse_args() generate_single_turn(Path(args.base_dir), Path(args.output_dir), data_size=args.data_size) if __name__ == "__main__": main()