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
1from tree_sitter import Language, Parser2import tree_sitter_rust as tsrust3import sys4import os5import re6 7RS_LANGUAGE = Language(tsrust.language(), "rust")8parser = Parser()9parser.set_language(RS_LANGUAGE)10# call_functions = []11 12# 获取impl定义13query_impl_defin_text = """14(15 (impl_item) @impl.defin16)17"""18 19# 获取函数定义20query_function_defin_text = """21(22 (function_item) @function.defin23)24"""25 26query_function_name_text = """27(function_item28 (identifier) @function.name)29"""30 31# 获取macro定义32query_macro_defin_text = """33(34 (macro_definition) @macro.defin35)36"""37 38query_macro_name_text = """39(macro_definition40 (identifier) @macro.name)41"""42 43# 获取struct定义44query_struct_defin_text = """45(46 (struct_item) @struct.defin47)48"""49query_struct_name_text = """50(struct_item51 (type_identifier) @struct.name52)53"""54 55# 调用的函数56query_call_function_text = """57(58 (call_expression) @function.call59)60"""61query_call_function_name_text = """62(63 (field_identifier) @function.call_name64)65"""66 67# 调用的macro68query_call_macro_text = """69(70 (macro_invocation) @macro.call71)72"""73query_call_macro_name_text = """74(macro_invocation75 (identifier) @macro.call_name)76"""77 78 79 80# 调用的数据类型81query_call_vars_type_text = """82(83 (type_identifier) @call_vars.type84)85"""86 87# 调用的变量88query_call_vars_text = """89(90 (expression_statement) @call_vars.exp91)92(93 (let_declaration) @call_vars.let94)95"""96 97query_call_vars_name_text = """98(99 (field_identifier) @call_vars.let100)101"""102 103query_import_text = """104(105 (use_declaration) @use.name106)107"""108 109# Create a query object110query_impl_defin = RS_LANGUAGE.query(query_impl_defin_text)111query_function_defin = RS_LANGUAGE.query(query_function_defin_text)112query_function_name = RS_LANGUAGE.query(query_function_name_text)113query_macro_defin = RS_LANGUAGE.query(query_macro_defin_text)114query_macro_name = RS_LANGUAGE.query(query_macro_name_text)115query_struct_defin = RS_LANGUAGE.query(query_struct_defin_text)116query_struct_name = RS_LANGUAGE.query(query_struct_name_text)117query_call_function = RS_LANGUAGE.query(query_call_function_text)118query_call_function_name = RS_LANGUAGE.query(query_call_function_name_text)119query_call_macro = RS_LANGUAGE.query(query_call_macro_text)120query_call_macro_name = RS_LANGUAGE.query(query_call_macro_name_text)121query_call_vars = RS_LANGUAGE.query(query_call_vars_text)122query_call_vars_name = RS_LANGUAGE.query(query_call_vars_name_text)123query_call_vars_type = RS_LANGUAGE.query(query_call_vars_type_text)124query_import = RS_LANGUAGE.query(query_import_text)125 126 127 128def traverse_call(node, source_code, call_functions):129 if node.type == "call":130 call_functions.append(node)131 function_call_code = source_code[node.start_byte:node.end_byte].decode('utf-8')132 for child in node.children:133 traverse_call(child, source_code, call_functions)134 135 136def traverse(node, source_code, depth=0):137 # Get the node text138 node_text = source_code[node.start_byte:node.end_byte].decode('utf-8')139 for child in node.children:140 traverse(child, source_code, depth + 1)141 142# Execute the query to get the captures143 144 145def get_source_code(target_file_path):146 147 with open(target_file_path, 'r', encoding='utf-8', errors='ignore') as input_file:148 source_code = input_file.read()149 150 # 转化成bytes!!否则在出现中文注释时,根据偏移获得对应内容会出错151 source_code = bytes(source_code, "utf-8")152 153 return source_code154 155def get_source_code_and_path(target_file_path):156 with open(target_file_path, 'r', encoding='utf-8', errors='ignore') as input_file:157 content = input_file.read()158 159 content = content.split("------")[0]160 # print(content)161 162 pattern = r'<path>(.*?)</path>'163 function_path = re.findall(pattern, content, re.DOTALL)[0].strip()164 165 pattern = r'<function>(.*?)</function>'166 source_code = re.findall(pattern, content, re.DOTALL)[0].strip()167 168 # 转化成bytes!!否则在出现中文注释时,根据偏移获得对应内容会出错169 source_code = bytes(source_code, "utf-8")170 171 if "deltachat-core" in function_path:172 function_path = function_path.replace("rust/", "rust/src/")173 174 return source_code, function_path175 176def get_call_macro(node, source_code):177 call_macro_names = set()178 # 获取依赖数据类型179 call_macro_captures = query_call_macro.captures(node)180 for call_macro_capture in call_macro_captures:181 call_macro_node , _ = call_macro_capture182 call_macro_name_capture = query_call_macro_name.captures(call_macro_node)183 try:184 call_macro_name_node , _ = call_macro_name_capture[0]185 call_macro_code = source_code[call_macro_name_node.start_byte:call_macro_name_node.end_byte].decode()186 call_macro_names.add(call_macro_code)187 except:188 pass189 return call_macro_names190 191def get_call_function(node, source_code):192 call_function_names = set()193 # print(f"Function : {function_name}")194 195 # 获取call function196 call_function_captures = query_call_function.captures(node)197 198 199 200 for call_function_capture in call_function_captures:201 call_function_node, _ = call_function_capture202 # call function203 call_function_var_code = source_code[call_function_node.start_byte:call_function_node.end_byte].decode()204 call_function_var_code = call_function_var_code.split("=")205 for call_function_var in call_function_var_code:206 if "(" in call_function_var and ")" in call_function_var:207 call_function_var = call_function_var.replace("\n", "")208 call_function_var = call_function_var.split("self.")[-1]209 pattern = r'([a-zA-Z_][a-zA-Z0-9_:]*)\.|([a-zA-Z_][a-zA-Z0-9_:]*\([^\)]*\))'210 call_function_var = re.findall(pattern, call_function_var)211 212 # call_vars = [match[0] for match in call_function_var if match[0] ]213 call_function = [match[1].split("(")[0] for match in call_function_var if match[1]]214 215 # call_vars_name.update(call_vars)216 call_function_names.update(call_function)217 218 return call_function_names219 220def get_call_vars_type(node, source_code):221 call_vars_type_name = set()222 # 获取依赖数据类型223 vars_type_captures = query_call_vars_type.captures(node)224 for vars_type_capture in vars_type_captures:225 vars_type_node , _ = vars_type_capture226 vars_type_code = source_code[vars_type_node.start_byte:vars_type_node.end_byte].decode()227 call_vars_type_name.add(vars_type_code)228 return call_vars_type_name229 230def get_file_function_dependency(target_file_path):231 232 Dependency_func = {}233 Dependency_vars = {}234 235 function_source_code, function_path = get_source_code_and_path(target_file_path)236 tree = parser.parse(function_source_code)237 238 239 240 # 以下是从已经提取好的函数文件(只有单个目标函数)中获取依赖241 # 先直接按照该函数进行获取242 captures = query_function_defin.captures(tree.root_node)243 target_function_name = ""244 for capture in captures:245 node, _ = capture246 function_name_captures = query_function_name.captures(node)247 function_name_node, _ = function_name_captures[0]248 function_name = function_source_code[function_name_node.start_byte:function_name_node.end_byte].decode()249 target_function_name = function_name250 # 先获取全部函数再减去impl函数251 # if function_name in Dependency_func.keys():252 # continue253 254 call_function_names = get_call_function(node, function_source_code)255 call_macro_names = get_call_macro(node, function_source_code)256 call_vars_type_name = get_call_vars_type(node, function_source_code)257 258 call_function_names.update(call_macro_names)259 Dependency_func[function_name] = call_function_names260 Dependency_vars[function_name] = call_vars_type_name261 262 # 从该函数的function_path来查看该函数是否在某个impl内部,如果是的话则将该impl的内容添加给该函数的Dependency_vars263 source_code = get_source_code(function_path)264 tree = parser.parse(source_code)265 query_impl_defin_captures = query_impl_defin.captures(tree.root_node)266 for query_impl_defin_capture in query_impl_defin_captures:267 query_impl_defin_node , _ = query_impl_defin_capture268 269 struct_name = source_code[query_impl_defin_node.start_byte:query_impl_defin_node.end_byte].decode().split("{")[0].split("for")[-1].split("impl")[-1].split(" ")[-2].split("<")[0].strip()270 271 captures = query_function_defin.captures(query_impl_defin_node)272 for capture in captures:273 node, _ = capture274 function_name_captures = query_function_name.captures(node)275 function_name_node, _ = function_name_captures[0]276 function_name = source_code[function_name_node.start_byte:function_name_node.end_byte].decode()277 278 if function_name == target_function_name:279 # impl里的函数添加相关struct的定义280 Dependency_vars[function_name].add(struct_name)281 282 return Dependency_func, Dependency_vars, function_path, function_source_code283 284 285def filtered_os_walk(top):286 for root, dirs, files in os.walk(top):287 # 过滤掉名字以"."开头的目录和名字包含"test"的目录288 dirs[:] = [d for d in dirs if not d.startswith('.') and 'test' not in d]289 yield root, dirs, files290 291def get_function_defin(node, source_code, function_name_to_code, project_functions, file_path):292 function_defin_captures = query_function_defin.captures(node)293 function_names = set()294 for function_defin_capture in function_defin_captures:295 function_defin_node , _ = function_defin_capture296 function_code = source_code[function_defin_node.start_byte:function_defin_node.end_byte].decode()297 function_name_captures = query_function_name.captures(function_defin_node)298 function_name_node, _ = function_name_captures[0]299 function_name = source_code[function_name_node.start_byte:function_name_node.end_byte].decode()300 301 function_names.add(function_name)302 # 以@为分隔符,将函数的文件路径加入,防止在单个项目中存在多个同名函数303 function_name_to_code[file_path + "@" + function_name] = function_code304 project_functions[file_path] = function_names305 306def get_struct_defin(node, source_code, struct_name_to_code, project_structs, file_path):307 struct_names = set()308 struct_defin_captures = query_struct_defin.captures(node)309 for struct_defin_capture in struct_defin_captures:310 struct_defin_node , _ = struct_defin_capture311 struct_defin_code = source_code[struct_defin_node.start_byte:struct_defin_node.end_byte].decode()312 struct_name_captures = query_struct_name.captures(struct_defin_node)313 sturct_name_node, _ = struct_name_captures[0]314 sturct_name = source_code[sturct_name_node.start_byte:sturct_name_node.end_byte].decode()315 struct_names.add(sturct_name)316 struct_name_to_code[file_path + "@" + sturct_name] = struct_defin_code317 318 project_structs[file_path] = struct_names319 320def get_macro_defin(node, source_code, macro_name_to_code, project_macros, file_path):321 macro_names = set()322 macro_defin_captures = query_macro_defin.captures(node)323 for macro_defin_capture in macro_defin_captures:324 macro_defin_node , _ = macro_defin_capture325 macro_defin_code = source_code[macro_defin_node.start_byte:macro_defin_node.end_byte].decode()326 macro_name_captures = query_macro_name.captures(macro_defin_node)327 macro_name_node, _ = macro_name_captures[0]328 macro_name = source_code[macro_name_node.start_byte:macro_name_node.end_byte].decode()329 330 macro_names.add(macro_name)331 332 macro_name_to_code[file_path + "@" + macro_name] = macro_defin_code333 project_macros[file_path] = macro_names334 335# 读取项目336def get_project_functions(project_path):337 338 project_functions = {}339 project_structs = {}340 project_macros = {}341 342 function_name_to_code = {}343 struct_name_to_code = {}344 macro_name_to_code = {}345 346 project_imports = {}347 project_vars = {}348 349 # 手动添加350 project_structs["IString"] = ["IString"]351 struct_name_to_code["IString@IString"] = "pub type IString = ::string_cache::Atom<IStringStaticSet>;"352 353 for current_path, _, files in os.walk(project_path):354 for file in files:355 if "test" in file:356 continue357 if file.endswith(".rs"):358 359 file_path = os.path.join(current_path, file)360 source_code = get_source_code(file_path)361 tree = parser.parse(source_code)362 363 # get function defin364 get_function_defin(tree.root_node, source_code, function_name_to_code, project_functions, file_path)365 366 # get struct defin367 get_struct_defin(tree.root_node, source_code, struct_name_to_code, project_structs, file_path)368 369 # get macro defin370 get_macro_defin(tree.root_node, source_code, macro_name_to_code, project_macros, file_path)371 372 # get import373 import_codes = []374 import_captures = query_import.captures(tree.root_node)375 for import_capture in import_captures:376 import_node, _ = import_capture377 import_code = source_code[import_node.start_byte:import_node.end_byte].decode().split("use")[-1].strip()378 import_codes.append(import_code)379 project_imports[file_path] = import_codes380 381 # 将macro归入function382 for file_path, macros in project_macros.items():383 if file_path in project_functions.keys():384 project_functions[file_path].update(macros) 385 function_name_to_code = {**macro_name_to_code, **function_name_to_code}386 387 return project_imports, project_functions, function_name_to_code, project_structs, struct_name_to_code388 389def match(project_functions, dependency_funcs, function_path, project_imports):390 # 依次进行匹配391 project_dependency_function = {}392 393 for target_function , call_functions in dependency_funcs.items():394 dependencies = []395 for call_function in call_functions:396 key = False397 # 先从同个文件中找398 for file_path, potential_functions in project_functions.items():399 if file_path == function_path:400 if call_function in potential_functions:401 dependencies.append(function_path + "@" + call_function)402 key = True403 break404 if key:405 continue406 407 # 从import中找408 for project_file_path, file_imports in project_imports.items():409 # 获得目标文件的对应import410 if project_file_path == function_path:411 for file_path, potential_functions in project_functions.items():412 # 获取文件名simple_path,判断该文件是否在目标文件的import中413 # 如果在file_import中,那么目标文件可以使用该文件内的函数414 simple_path = file_path.split("/")[-1].split(".")[0].strip()415 416 for file_import in file_imports:417 if simple_path in file_import:418 if call_function in potential_functions:419 dependencies.append(file_path + "@" + call_function)420 key = True421 break422 if key:423 break424 break425 if key:426 continue427 428 project_dependency_function[target_function] = dependencies429 430 431 return project_dependency_function432 433 434project_name = sys.argv[1]435target_lang_pair = sys.argv[2]436target_lang = sys.argv[3]437target_files_path = os.path.join("function_pair_with_identical_functionality", project_name, target_lang_pair)438project_path = os.path.join("projects", project_name, target_lang)439 440project_imports, project_functions, function_name_to_code, project_structs, struct_name_to_code = get_project_functions(project_path)441 442unit_test_function = set()443for current_path, _, target_files in os.walk(target_files_path):444 for target_file in target_files:445 446 dependency_funcs, dependency_vars, function_path, source_code = get_file_function_dependency(os.path.join(current_path, target_file))447 if not dependency_funcs:448 continue449 result_function = match(project_functions, dependency_funcs, function_path, project_imports)450 result_vars = match(project_structs, dependency_vars, function_path, project_imports)451 output_file_path = os.path.join("related_functions_and_datatypes_and_import", project_name, target_lang_pair, target_lang)452 if not os.path.exists(output_file_path):453 os.makedirs(output_file_path)454 455 with open(os.path.join(output_file_path, target_file), "w") as output_file:456 for dependencies in result_function.values():457 for dependency in dependencies:458 if function_name_to_code[dependency].strip() == source_code.decode("utf-8").strip():459 continue460 output_file.write(function_name_to_code[dependency])461 output_file.write("\n\n")462 463 for dependencies in result_vars.values():464 for dependency in dependencies:465 output_file.write(struct_name_to_code[dependency])466 output_file.write("\n\n")467 468 output_file.write("------\n")469 470 tmp = [f"use {use}\n" for use in project_imports[function_path]]471 output_file.writelines(tmp)472 473 474 475 476 