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
rust_extract_dependency.py476 linesDownload Raw Back to Dataset_Construction
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 
SYSUSELab/RustRepoTrans · Team Ai