Team Ai
Datasetpublic

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.

sourceHugging Faceupdated 2y agoView on Hugging Face
3likes104downloads
extract_function.py175 linesDownload Raw Back to Dataset_Construction
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 
SYSUSELab/RustRepoTrans · Team Ai