Tom-Dev-space/code_graph
0
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 