Team Ai
Apppublic

Tom-Dev-space/code_graph

sourceHugging Facemitupdated 1y agoView on Hugging Face
0likes
utility.py167 linesDownload Raw Back to root
1import ast2 3import ast4 5class CallCollector(ast.NodeVisitor):6    def __init__(self, defined_funcs):7        self.calls = []8        self.defined_funcs = defined_funcs9 10    def visit_Call(self, node):11        if isinstance(node.func, ast.Name) and node.func.id in self.defined_funcs:12            self.calls.append(node.func.id)13        self.generic_visit(node)14 15def parse_functions_from_files(file_dict):16    functions = {}17    defined_funcs = set()18 19    def infer_type_from_value(value_node):20        if isinstance(value_node, ast.Call) and isinstance(value_node.func, ast.Attribute):21            if value_node.func.attr in ("read_csv", "DataFrame"): return "pd.DataFrame"22            if value_node.func.attr == "array": return "np.ndarray"23        elif isinstance(value_node, ast.List): return "list"24        elif isinstance(value_node, ast.Dict): return "dict"25        elif isinstance(value_node, ast.Set): return "set"26        elif isinstance(value_node, ast.Constant):27            if isinstance(value_node.value, str): return "str"28            if isinstance(value_node.value, bool): return "bool"29            if isinstance(value_node.value, int): return "int"30            if isinstance(value_node.value, float): return "float"31        return "?"32 33    for fname, code in file_dict.items():34        tree = ast.parse(code)35        for node in ast.walk(tree):36            if isinstance(node, ast.FunctionDef): defined_funcs.add(node.name)37 38    for fname, code in file_dict.items():39        tree = ast.parse(code)40        for node in ast.walk(tree):41            if isinstance(node, ast.FunctionDef):42                func_name = node.name43                args, returns = [], []44                local_assignments = {}45                reads_state, writes_state = set(), set()46 47                for arg in node.args.args:48                    arg_type = ast.unparse(arg.annotation) if arg.annotation else "?"49                    args.append(f"{arg.arg}: {arg_type}")50 51                collector = CallCollector(defined_funcs)52                collector.visit(node)53                calls = collector.calls54 55                for sub in ast.walk(node):56                    if isinstance(sub, ast.Assign):57                        for target in sub.targets:58                            if isinstance(target, ast.Name):59                                local_assignments[target.id] = infer_type_from_value(sub.value)60                    elif isinstance(sub, ast.Return):61                        if sub.value is None: continue62                        if isinstance(sub.value, ast.Tuple):63                            for elt in sub.value.elts:64                                label = ast.unparse(elt)65                                returns.append(f"{label}: {local_assignments.get(label, infer_type_from_value(elt))}")66                        else:67                            label = ast.unparse(sub.value)68                            returns.append(f"{label}: {local_assignments.get(label, infer_type_from_value(sub.value))}")69 70                functions[func_name] = {71                    "args": args,72                    "returns": returns,73                    "calls": calls,74                    "filename": fname,75                    "reads_state": sorted(reads_state),76                    "writes_state": sorted(writes_state)77                }78 79    return functions80 81 82 83def get_reachable_functions(start, graph):84    visited, stack = set(), [start]85    while stack:86        node = stack.pop()87        if node not in visited:88            visited.add(node)89            stack.extend(graph.get(node, []))90    return visited91 92def get_backtrace_functions(target, graph):93    reverse_graph = {}94    for caller, callees in graph.items():95        for callee in callees:96            reverse_graph.setdefault(callee, []).append(caller)97    visited, stack = set(), [target]98    while stack:99        node = stack.pop()100        if node not in visited:101            visited.add(node)102            stack.extend(reverse_graph.get(node, []))103    return visited104 105from collections import deque, defaultdict106 107def build_nodes_and_edges(parsed, root_func, reachable, reverse=False, max_depth=10):108    depth_map = {}109    x_offset_map = defaultdict(int)110    positions = {}111    visited = set()112    queue = deque([(root_func, 0)])113 114    while queue:115        current, depth = queue.popleft()116        if current in visited or depth > max_depth:117            continue118        visited.add(current)119 120        adjusted_depth = -depth if reverse else depth121        x = x_offset_map[adjusted_depth] * 300122        y = adjusted_depth * 150123        positions[current] = {"x": x, "y": y}124        x_offset_map[adjusted_depth] += 1125 126        if reverse:127            # Find callers of this function128            next_funcs = [129                caller for caller, meta in parsed.items()130                if current in meta["calls"] and caller in reachable131            ]132        else:133            # Find callees134            next_funcs = [135                callee for callee in parsed[current]["calls"]136                if callee in reachable137            ]138 139        for nxt in next_funcs:140            queue.append((nxt, depth + 1))141 142    nodes = [{143        "data": {"id": name, "label": name},144        "position": positions.get(name, {"x": 0, "y": 0}),145        "classes": "main" if name == root_func else ""146    } for name in visited]147 148    edges = []149    for src in visited:150        call_sequence = parsed[src]["calls"]151        call_index = 1152        for tgt in call_sequence:153            if tgt in visited:154                edges.append({155                    "data": {156                        "source": src,157                        "target": tgt,158                        "label": str(call_index)159                    },160                    "style": {161                        "line-width": 4 if call_sequence.count(tgt) > 1 else 2162                    }163                })164                call_index += 1165 166    return nodes, edges167