Maaac/CodeLLaMA-Linux-BugFix
08
1from pydriller import Repository2import os3import json4from tqdm import tqdm5import re6from multiprocessing import Pool7 8REPO_PATH = '../linux'9OUTPUT_FILE = './output/linux_bugfix_dataset.jsonl'10 11TEST_MODE = False # Set to False to process the full repository12MAX_COMMITS_TEST = 50 # Set a limit if TEST_MODE is True13NUM_WORKERS = 16 # Adjust to your actual core count14 15BUGFIX_KEYWORDS = [16 'fix', 'bug', 'leak', 'null', 'overflow', 'error', 'failure',17 'crash', 'panic', 'memory', 'race', 'deadlock', 'corruption',18 'security', 'vulnerability', 'exploit', 'buffer', 'stack'19]20 21def is_bugfix_commit(msg):22 msg_lower = msg.lower()23 return any(keyword in msg_lower for keyword in BUGFIX_KEYWORDS)24 25def extract_instruction_from_commit_msg(msg):26 lines = msg.strip().splitlines()27 for line in lines:28 line = line.strip()29 if len(line) < 5 or not any(c.isalpha() for c in line):30 continue31 if line.lower().startswith((32 '[patch]', 'signed-off-by', 'reviewed-by', 'tested-by', 'ack',33 'reported-by', 'cc:', 'co-authored-by', 'patchwork-id',34 'suggested-by', 'fixes:', 'link:', 'cherry picked from commit'35 )):36 continue37 return line38 return msg.strip().splitlines()[0] if msg.strip() else "fix"39 40def extract_code_context(code, line_number, context_lines=10):41 if not code:42 return ""43 lines = code.split('\n')44 start = max(0, line_number - context_lines)45 end = min(len(lines), line_number + context_lines)46 return '\n'.join(lines[start:end])47 48def extract_diff_context(diff_text, context_lines=5):49 if not diff_text:50 return ""51 lines = diff_text.split('\n')52 change_lines = [i for i, line in enumerate(lines) if line.startswith('+') or line.startswith('-')]53 if not change_lines:54 return diff_text55 start = max(0, change_lines[0] - context_lines)56 end = min(len(lines), change_lines[-1] + context_lines + 1)57 return '\n'.join(lines[start:end])58 59def create_dataset_entry(original_code, commit_msg, diff_code):60 return {61 "input": {62 "original code": original_code.strip(),63 "instruction": extract_instruction_from_commit_msg(commit_msg)64 },65 "output": {66 "diff codes": diff_code.strip()67 }68 }69 70def process_commit(commit):71 entries = []72 if not is_bugfix_commit(commit.msg):73 return entries74 75 for mod in commit.modified_files:76 if not mod.new_path or not mod.new_path.endswith(('.c', '.h')):77 continue78 if mod.change_type.name != "MODIFY":79 continue80 if not mod.diff or not mod.source_code_before:81 continue82 83 focused_diff = extract_diff_context(mod.diff)84 85 diff_lines = mod.diff.split('\n')86 line_numbers = []87 for line in diff_lines:88 if line.startswith('@@'):89 match = re.search(r'@@ -(\d+),?\d* \+\d+,?\d* @@', line)90 if match:91 line_numbers.append(int(match.group(1)))92 93 if line_numbers:94 focused_code = extract_code_context(mod.source_code_before, line_numbers[0])95 else:96 focused_code = '\n'.join(mod.source_code_before.split('\n')[:50])97 98 entry = create_dataset_entry(99 original_code=focused_code,100 commit_msg=commit.msg,101 diff_code=focused_diff102 )103 entries.append(entry)104 105 return entries106 107def collect_entries_from_hash(commit_hash):108 try:109 commit = next(Repository(REPO_PATH, only_commits=[commit_hash]).traverse_commits())110 return process_commit(commit)111 except Exception:112 return []113 114def main():115 if not os.path.exists(REPO_PATH):116 print("[ERROR] Repository not found at:", REPO_PATH)117 return118 119 os.makedirs('./output', exist_ok=True)120 121 print("[INFO] Building Linux kernel bug-fix dataset...")122 print("[INFO] Repository:", REPO_PATH)123 print("[INFO] Output file:", OUTPUT_FILE)124 125 output_file = OUTPUT_FILE.replace('.jsonl', '_test.jsonl') if TEST_MODE else OUTPUT_FILE126 127 all_hashes = [c.hash for c in Repository(REPO_PATH).traverse_commits()]128 if TEST_MODE and MAX_COMMITS_TEST:129 all_hashes = all_hashes[:MAX_COMMITS_TEST]130 131 dataset_entries = []132 with Pool(NUM_WORKERS) as pool:133 results = list(tqdm(pool.imap_unordered(collect_entries_from_hash, all_hashes), total=len(all_hashes)))134 135 for entries in results:136 dataset_entries.extend(entries)137 138 with open(output_file, 'w', encoding='utf-8') as f:139 for entry in dataset_entries:140 f.write(json.dumps(entry, ensure_ascii=False) + '\n')141 142 print("[DONE] Dataset creation completed!")143 print("[INFO] Total commits processed:", len(all_hashes))144 print("[INFO] Total dataset entries:", len(dataset_entries))145 print("[INFO] Saved to:", output_file)146 147 if dataset_entries:148 print("[INFO] Sample dataset entry:")149 sample = dataset_entries[0]150 print(json.dumps(sample, indent=2, ensure_ascii=False)[:800] + "...")151 print("[INFO] Dataset structure:")152 print(" - Input: original code + instruction")153 print(" - Output: diff codes")154 print(" - Format: JSONL (one JSON object per line)")155 156if __name__ == "__main__":157 main()158 