Team Ai
Modelpublic

thealper2/graphcodebert-code-clone-detection

sourceHugging Facemitupdated 17d agoView on Hugging Face
0likes30downloads
dfg_parser.py422 linesDownload Raw Back to code
1"""Data-flow graph (DFG) extraction for GraphCodeBERT.2 3This is a faithful port of Microsoft's GraphCodeBERT ``parser/`` package4(``utils.py`` + the ``DFG_python`` extractor from ``DFG.py``), adapted to the5modern ``py-tree-sitter`` API (>= 0.22, ``Language(tree_sitter_python.language())``)6instead of the original hand-compiled ``my-languages.so``.7 8The dataset used in this project contains **Python** snippets (verified in9``preprocess.py``), so only the Python extractor is ported; adding another10language means adding its ``DFG_<lang>`` function and grammar package here.11 12A DFG entry is the 5-tuple used throughout GraphCodeBERT::13 14    (variable_name, token_index, edge_type, source_variable_names, source_token_indices)15 16``edge_type`` is ``"comesFrom"`` (value flows from a previous definition) or17``"computedFrom"`` (value is computed from the right-hand side of an assignment).18"""19 20from __future__ import annotations21 22import io23import re24import sys25import tokenize26from typing import Any27 28from tree_sitter import Language, Node, Parser29 30__all__ = [31    "get_parser",32    "extract_dataflow",33    "remove_comments_and_docstrings",34    "DataFlowExtractionError",35]36 37#: tree-sitter recursion is mirrored by the recursive Python walkers below.38#: Competitive-programming snippets can nest deeply, so raise the ceiling but39#: keep it bounded so a pathological file raises RecursionError instead of40#: segfaulting the worker.41_RECURSION_LIMIT = 10_00042 43 44class DataFlowExtractionError(RuntimeError):45    """Raised when a snippet cannot be turned into code tokens at all."""46 47 48_PARSER_CACHE: dict[str, Parser] = {}49 50 51def get_parser(language: str = "python") -> Parser:52    """Return a cached tree-sitter parser for ``language``.53 54    Cached per process so that ``datasets.map(num_proc=...)`` workers each build55    the parser once rather than once per snippet.56    """57    if language in _PARSER_CACHE:58        return _PARSER_CACHE[language]59    if language != "python":60        raise ValueError(61            f"Only the Python grammar is wired up (requested {language!r}). "62            "Add the matching tree_sitter_<lang> package and DFG_<lang> function."63        )64    try:65        import tree_sitter_python66    except ImportError as exc:  # pragma: no cover - environment problem67        raise ImportError(68            "tree_sitter_python is required for GraphCodeBERT data-flow extraction. "69            "Install it with `pip install tree-sitter tree-sitter-python`."70        ) from exc71    parser = Parser(Language(tree_sitter_python.language()))72    _PARSER_CACHE[language] = parser73    return parser74 75 76# --------------------------------------------------------------------------- #77# parser/utils.py78# --------------------------------------------------------------------------- #79def remove_comments_and_docstrings(source: str, lang: str = "python") -> str:80    """Strip comments and docstrings, preserving token columns.81 82    Column positions are preserved because the DFG indices are ``(row, column)``83    points into the *cleaned* source.84    """85    if lang == "python":86        io_obj = io.StringIO(source)87        out = ""88        prev_toktype = tokenize.INDENT89        last_lineno = -190        last_col = 091        for tok in tokenize.generate_tokens(io_obj.readline):92            token_type, token_string = tok[0], tok[1]93            start_line, start_col = tok[2]94            end_line, end_col = tok[3]95            if start_line > last_lineno:96                last_col = 097            if start_col > last_col:98                out += " " * (start_col - last_col)99            if token_type == tokenize.COMMENT:100                pass101            elif token_type == tokenize.STRING:102                # A string that starts a logical line is a docstring -> drop it.103                if prev_toktype != tokenize.INDENT and prev_toktype != tokenize.NEWLINE:104                    if start_col > 0:105                        out += token_string106            else:107                out += token_string108            prev_toktype = token_type109            last_col = end_col110            last_lineno = end_line111        return "\n".join(x for x in out.split("\n") if x.strip() != "")112 113    def _replacer(match: re.Match[str]) -> str:114        s = match.group(0)115        return " " if s.startswith("/") else s116 117    pattern = re.compile(118        r"//.*?$|/\*.*?\*/|\'(?:\\.|[^\\\'])*\'|\"(?:\\.|[^\\\"])*\"",119        re.DOTALL | re.MULTILINE,120    )121    cleaned = re.sub(pattern, _replacer, source)122    return "\n".join(x for x in cleaned.split("\n") if x.strip() != "")123 124 125def tree_to_token_index(root_node: Node) -> list[tuple[Any, Any]]:126    """Collect ``(start_point, end_point)`` spans of every leaf token."""127    if (len(root_node.children) == 0 or root_node.type == "string") and root_node.type != "comment":128        return [(root_node.start_point, root_node.end_point)]129    spans: list[tuple[Any, Any]] = []130    for child in root_node.children:131        spans += tree_to_token_index(child)132    return spans133 134 135def tree_to_variable_index(root_node: Node, index_to_code: dict) -> list[tuple[Any, Any]]:136    """Collect spans of leaves that are *variables* (token text != node type)."""137    if (len(root_node.children) == 0 or root_node.type == "string") and root_node.type != "comment":138        index = (root_node.start_point, root_node.end_point)139        _, code = index_to_code[index]140        return [] if root_node.type == code else [index]141    spans: list[tuple[Any, Any]] = []142    for child in root_node.children:143        spans += tree_to_variable_index(child, index_to_code)144    return spans145 146 147def index_to_code_token(index: tuple[Any, Any], code: list[str]) -> str:148    """Slice the source text covered by a ``(start_point, end_point)`` span."""149    start_point, end_point = index150    if start_point[0] == end_point[0]:151        return code[start_point[0]][start_point[1] : end_point[1]]152    s = code[start_point[0]][start_point[1] :]153    for i in range(start_point[0] + 1, end_point[0]):154        s += code[i]155    s += code[end_point[0]][: end_point[1]]156    return s157 158 159# --------------------------------------------------------------------------- #160# parser/DFG.py :: DFG_python161# --------------------------------------------------------------------------- #162_ASSIGNMENT = ("assignment", "augmented_assignment", "for_in_clause")163_IF_STATEMENT = ("if_statement",)164_FOR_STATEMENT = ("for_statement",)165_WHILE_STATEMENT = ("while_statement",)166_DO_FIRST_STATEMENT = ("for_in_clause",)167_DEF_STATEMENT = ("default_parameter",)168 169 170def DFG_python(root_node: Node, index_to_code: dict, states: dict) -> tuple[list, dict]:171    """Build the data-flow graph of a Python AST subtree.172 173    Returns ``(dfg_edges, variable_states)`` where ``variable_states`` maps a174    variable name to the token indices that currently define it.175    """176    states = states.copy()177 178    if (len(root_node.children) == 0 or root_node.type == "string") and root_node.type != "comment":179        idx, code = index_to_code[(root_node.start_point, root_node.end_point)]180        if root_node.type == code:  # a keyword/operator, not a variable181            return [], states182        if code in states:183            return [(code, idx, "comesFrom", [code], states[code].copy())], states184        if root_node.type == "identifier":185            states[code] = [idx]186        return [(code, idx, "comesFrom", [], [])], states187 188    if root_node.type in _DEF_STATEMENT:189        name = root_node.child_by_field_name("name")190        value = root_node.child_by_field_name("value")191        dfg: list = []192        if value is None:193            for index in tree_to_variable_index(name, index_to_code):194                idx, code = index_to_code[index]195                dfg.append((code, idx, "comesFrom", [], []))196                states[code] = [idx]197            return sorted(dfg, key=lambda x: x[1]), states198        name_indexs = tree_to_variable_index(name, index_to_code)199        value_indexs = tree_to_variable_index(value, index_to_code)200        temp, states = DFG_python(value, index_to_code, states)201        dfg += temp202        for index1 in name_indexs:203            idx1, code1 = index_to_code[index1]204            for index2 in value_indexs:205                idx2, code2 = index_to_code[index2]206                dfg.append((code1, idx1, "comesFrom", [code2], [idx2]))207            states[code1] = [idx1]208        return sorted(dfg, key=lambda x: x[1]), states209 210    if root_node.type in _ASSIGNMENT:211        if root_node.type == "for_in_clause":212            right_nodes = [root_node.children[-1]]213            left_nodes = [root_node.child_by_field_name("left")]214        else:215            if root_node.child_by_field_name("right") is None:216                return [], states217            left_nodes = [x for x in root_node.child_by_field_name("left").children if x.type != ","]218            right_nodes = [219                x for x in root_node.child_by_field_name("right").children if x.type != ","220            ]221            if len(right_nodes) != len(left_nodes):222                left_nodes = [root_node.child_by_field_name("left")]223                right_nodes = [root_node.child_by_field_name("right")]224            if len(left_nodes) == 0:225                left_nodes = [root_node.child_by_field_name("left")]226            if len(right_nodes) == 0:227                right_nodes = [root_node.child_by_field_name("right")]228        dfg = []229        for node in right_nodes:230            temp, states = DFG_python(node, index_to_code, states)231            dfg += temp232        for left_node, right_node in zip(left_nodes, right_nodes):233            left_tokens_index = tree_to_variable_index(left_node, index_to_code)234            right_tokens_index = tree_to_variable_index(right_node, index_to_code)235            for token1_index in left_tokens_index:236                idx1, code1 = index_to_code[token1_index]237                dfg.append(238                    (239                        code1,240                        idx1,241                        "computedFrom",242                        [index_to_code[x][1] for x in right_tokens_index],243                        [index_to_code[x][0] for x in right_tokens_index],244                    )245                )246                states[code1] = [idx1]247        return sorted(dfg, key=lambda x: x[1]), states248 249    if root_node.type in _IF_STATEMENT:250        dfg = []251        current_states = states.copy()252        others_states = []253        tag = "else" in root_node.type254        for child in root_node.children:255            if "else" in child.type:256                tag = True257            if child.type not in ("elif_clause", "else_clause"):258                temp, current_states = DFG_python(child, index_to_code, current_states)259                dfg += temp260            else:261                temp, new_states = DFG_python(child, index_to_code, states)262                dfg += temp263                others_states.append(new_states)264        others_states.append(current_states)265        if tag is False:266            others_states.append(states)267        merged: dict = {}268        for dic in others_states:269            for key in dic:270                merged.setdefault(key, [])271                merged[key] += dic[key]272        for key in merged:273            merged[key] = sorted(set(merged[key]))274        return sorted(dfg, key=lambda x: x[1]), merged275 276    if root_node.type in _FOR_STATEMENT:277        dfg = []278        # Two passes: loop bodies can consume values defined later in the loop.279        for _ in range(2):280            right_nodes = [x for x in root_node.child_by_field_name("right").children if x.type != ","]281            left_nodes = [x for x in root_node.child_by_field_name("left").children if x.type != ","]282            if len(right_nodes) != len(left_nodes):283                left_nodes = [root_node.child_by_field_name("left")]284                right_nodes = [root_node.child_by_field_name("right")]285            if len(left_nodes) == 0:286                left_nodes = [root_node.child_by_field_name("left")]287            if len(right_nodes) == 0:288                right_nodes = [root_node.child_by_field_name("right")]289            for node in right_nodes:290                temp, states = DFG_python(node, index_to_code, states)291                dfg += temp292            for left_node, right_node in zip(left_nodes, right_nodes):293                left_tokens_index = tree_to_variable_index(left_node, index_to_code)294                right_tokens_index = tree_to_variable_index(right_node, index_to_code)295                for token1_index in left_tokens_index:296                    idx1, code1 = index_to_code[token1_index]297                    dfg.append(298                        (299                            code1,300                            idx1,301                            "computedFrom",302                            [index_to_code[x][1] for x in right_tokens_index],303                            [index_to_code[x][0] for x in right_tokens_index],304                        )305                    )306                    states[code1] = [idx1]307            if root_node.children[-1].type == "block":308                temp, states = DFG_python(root_node.children[-1], index_to_code, states)309                dfg += temp310        return _merge_duplicate_edges(dfg), states311 312    if root_node.type in _WHILE_STATEMENT:313        dfg = []314        for _ in range(2):315            for child in root_node.children:316                temp, states = DFG_python(child, index_to_code, states)317                dfg += temp318        return _merge_duplicate_edges(dfg), states319 320    dfg = []321    for child in root_node.children:322        if child.type in _DO_FIRST_STATEMENT:323            temp, states = DFG_python(child, index_to_code, states)324            dfg += temp325    for child in root_node.children:326        if child.type not in _DO_FIRST_STATEMENT:327            temp, states = DFG_python(child, index_to_code, states)328            dfg += temp329    return sorted(dfg, key=lambda x: x[1]), states330 331 332def _merge_duplicate_edges(dfg: list) -> list:333    """Collapse the duplicate edges produced by the two-pass loop handling."""334    dic: dict = {}335    for x in dfg:336        key = (x[0], x[1], x[2])337        if key not in dic:338            dic[key] = [x[3], x[4]]339        else:340            dic[key][0] = list(set(dic[key][0] + x[3]))341            dic[key][1] = sorted(set(dic[key][1] + x[4]))342    merged = [(k[0], k[1], k[2], v[0], v[1]) for k, v in sorted(dic.items(), key=lambda t: t[0][1])]343    return sorted(merged, key=lambda x: x[1])344 345 346# --------------------------------------------------------------------------- #347# Public entry point (GraphCodeBERT's `extract_dataflow`)348# --------------------------------------------------------------------------- #349def extract_dataflow(code: str, language: str = "python") -> tuple[list[str], list, dict]:350    """Tokenise ``code`` and extract its data-flow graph.351 352    Returns ``(code_tokens, dfg, status)``. ``status`` records *why* a stage353    degraded so callers can report it instead of hiding it:354 355    ``comment_strip`` : ``"ok"`` | ``"failed"``356    ``parse``         : ``"ok"`` | ``"failed"``357    ``dfg``           : ``"ok"`` | ``"failed"`` | ``"recursion_limit"``358    ``error``         : ``None`` or ``"<ExcType>: <message>"``359 360    A degraded DFG yields an **empty** data-flow component -- the snippet is361    still trained on (GraphCodeBERT tolerates zero nodes), it is never dropped.362    """363    status: dict[str, Any] = {"comment_strip": "ok", "parse": "ok", "dfg": "ok", "error": None}364 365    try:366        cleaned = remove_comments_and_docstrings(code, language)367    except Exception as exc:368        # Syntactically broken snippets are common in the wild; fall back to the369        # raw source rather than discarding the example.370        status["comment_strip"] = "failed"371        status["error"] = f"{type(exc).__name__}: {exc}"372        cleaned = code373 374    parser = get_parser(language)375    try:376        tree = parser.parse(bytes(cleaned, "utf8"))377        root_node = tree.root_node378    except Exception as exc:379        raise DataFlowExtractionError(f"tree-sitter failed to parse snippet: {exc}") from exc380 381    old_limit = sys.getrecursionlimit()382    sys.setrecursionlimit(_RECURSION_LIMIT)383    try:384        try:385            tokens_index = tree_to_token_index(root_node)386        except RecursionError as exc:387            status["parse"] = "failed"388            status["dfg"] = "recursion_limit"389            status["error"] = f"{type(exc).__name__}: token index recursion limit"390            raise DataFlowExtractionError("snippet nests deeper than the recursion limit") from exc391 392        lines = cleaned.split("\n")393        code_tokens = [index_to_code_token(x, lines) for x in tokens_index]394        index_to_code = {395            index: (idx, token) for idx, (index, token) in enumerate(zip(tokens_index, code_tokens))396        }397 398        try:399            dfg, _ = DFG_python(root_node, index_to_code, {})400        except RecursionError as exc:401            status["dfg"] = "recursion_limit"402            status["error"] = f"{type(exc).__name__}: DFG recursion limit"403            dfg = []404        except Exception as exc:405            status["dfg"] = "failed"406            status["error"] = f"{type(exc).__name__}: {exc}"407            dfg = []408    finally:409        sys.setrecursionlimit(old_limit)410 411    # Keep only nodes that participate in at least one edge (GraphCodeBERT does412    # the same: isolated nodes carry no data-flow signal).413    dfg = sorted(dfg, key=lambda x: x[1])414    keep: set[int] = set()415    for d in dfg:416        if len(d[-1]) != 0:417            keep.add(d[1])418        keep.update(d[-1])419    dfg = [d for d in dfg if d[1] in keep]420 421    return code_tokens, dfg, status422