thealper2/graphcodebert-code-clone-detection
030
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 