SYSUSELab/RustRepoTrans
Evaluating Large Language Models in Repository-level Code Translation RustRepoTrans is the first repository-level code translation benchmark described in the paper "RustRepoTrans: Repository-level Code Translation Benchmark Targeting Rust". Feel free to contact us to submit new results. Benchmark Dataset RustRepoTrans, the first repository-level code translation benchmark comprising 375 tasks targeting Rust, consists of 122 java-rust function pairs, 145 c-rust… See the full description on the dataset page: https://huggingface.co/datasets/SYSUSELab/RustRepoTrans.
3104
1import re
2import os
3import sys
4
5
6total_functions = set()
7
8def extract_functions_from_code(code, pattern):
9
10 function_head_pattern = re.compile(pattern)
11
12 lines = code.split('\n')
13 functions = []
14 brace_count = 0
15 function_code = []
16 inside_function = False
17 key = False
18 for i, line in enumerate(lines):
19 if not inside_function and function_head_pattern.search(line):
20 inside_function = True
21
22 if inside_function:
23 function_code.append(line)
24 if not line.lstrip(" ").startswith("//"):
25 brace_count += line.count('{')
26 brace_count -= line.count('}')
27 if brace_count == 0 :
28 if line.strip().endswith('}'):
29 inside_function = False
30 functions.append('\n'.join(function_code))
31 function_code = []
32 elif line.strip().endswith(';'):
33 inside_function = False
34 function_code = []
35 return functions
36
37def extract_functions_from_code_py(code):
38 lines = code.split('\n')
39 functions = []
40 function_code = []
41 inside_function = False
42
43 for line in lines:
44 if not inside_function and line.lstrip().startswith("def "):
45 # 获取函数起始的缩进长度
46 pre_cnt = len(line) - len(line.lstrip())
47 function_code.append(line)
48 inside_function = True
49 # 跳过def那一行
50 continue
51
52 if inside_function:
53 # 空行和缩进比def要小表示还在函数内
54 if len(line) == 0 or len(line) - len(line.lstrip()) >= pre_cnt + 4:
55 function_code.append(line)
56 else:
57 functions.append('\n'.join(function_code))
58 function_code = []
59 # 当前行有可能是下一个函数的声明行,不处理会跳过该函数
60 if line.lstrip().startswith("def "):
61 pre_cnt = len(line) - len(line.lstrip())
62 function_code.append(line)
63 else:
64 inside_function = False
65
66 # 处理在文件末尾声明定义的function
67 if function_code:
68 functions.append('\n'.join(function_code))
69
70 return functions
71
72def extract_functions_from_code_rb(code):
73 lines = code.split('\n')
74 functions = []
75 function_code = []
76 inside_function = False
77
78 for line in lines:
79 if not inside_function and line.lstrip().startswith("def "):
80 # 获取函数起始的缩进长度
81 pre_cnt = len(line) - len(line.lstrip())
82 function_code.append(line)
83 inside_function = True
84 # 跳过def那一行
85 continue
86
87 if inside_function:
88 if len(line) - len(line.lstrip()) == pre_cnt and line.lstrip().startswith("end"):
89 inside_function = False
90 functions.append('\n'.join(function_code))
91 function_code = []
92 else:
93 function_code.append(line)
94
95 return functions
96
97def save_functions_to_files(functions, output_dir, output_file_name):
98 if not os.path.exists(output_dir):
99 os.makedirs(output_dir)
100 try:
101 for i, func in enumerate(functions):
102 # influxdb-1.8\\client\\influxdb_test.go -> influxdb-1.8__client__influxdb_test
103 output_file = os.path.splitext(output_file_name)
104 output_file = output_file[0].replace("/", "__") + "__" + output_file[1]
105 file_path = os.path.join(output_dir, f'{output_file}__function__{i + 1}.txt')
106 with open(file_path, 'w', encoding='utf-8') as file:
107 file.write(f"<path>\n{output_file_name}\n</path>\n")
108 file.write(f"<function>\n{func}\n</function>")
109 except Exception as e:
110 print(e)
111 pass
112
113def process_file(input_file, lang, output_dir, pattern):
114
115 with open(input_file, 'r', encoding='utf-8', errors='ignore') as file:
116 code = file.read()
117
118 if lang == "py":
119 functions = extract_functions_from_code_py(code)
120 elif lang == "rb":
121 functions = extract_functions_from_code_rb(code)
122 else:
123 functions = extract_functions_from_code(code, pattern)
124
125 save_functions_to_files(functions, output_dir, input_file)
126
127def main():
128 project_dir = "projects"
129 target_project = sys.argv[1]
130
131
132 patterns = {
133 'cpp': r'^\s*[\w\s\*\[\]\<\>\:]+\s+[\w\s\*\[\]\<\>\:]+\s*\(',
134 'cxx': r'^\s*[\w\s\*\[\]\<\>\:]+\s+[\w\s\*\[\]\<\>\:]+\s*\(',
135 'h': r'^\s*[\w\s\*\[\]\<\>\:]+\s+[\w\s\*\[\]\<\>\:]+\s*\(',
136 'java': r'^\s*(public|protected|private|static|final|synchronized|native|abstract|strictfp|default)?\s*(public|protected|private|static|final|synchronized|native|abstract|strictfp|default)?\s*[\w\<\>\[\] ]+\s+[\w\<\>\[\]]+\s*\(',
137 'rs': r'^\s*(unsafe)?\s*(pub(\(crate\))?)?\s*(async)?\s*fn\s',
138 'c': r'^\s*[\w\s\*\[\]]*\s*\w+\s*\(',
139 'py': r''
140 }
141 lang_to_fileType = {
142 'cpp' : ['cpp', 'cxx', 'h'],
143 'c' : ['c', 'h'],
144 'java' : ['java'],
145 'rust' : ['rs'],
146 'python' : ['py']
147 }
148
149 projects = os.listdir(project_dir)
150 for project in projects:
151 if project != target_project:
152 continue
153 project_pair_path = os.path.join(project_dir, project)
154 langs = os.listdir(project_pair_path)
155 for lang in langs:
156 root_dir = os.path.join(project_pair_path, lang)
157 # 对项目进行遍历
158 for current_path, dirs, files in os.walk(root_dir):
159 dirs[:] = [d for d in dirs if not d.startswith('.')]
160
161 for file in files:
162 try :
163 file_lang = file.split('.')[-1]
164 except:
165 continue
166 if file_lang in lang_to_fileType[lang]:
167 file_path = os.path.join(current_path, file)
168 if "test" in file_path or "Test" in file_path:
169 continue
170 process_file(file_path, file_lang, root_dir.replace("projects", "functions"), patterns[file_lang])
171
172
173if __name__ == '__main__':
174 main()
175 