Jack-ki1's picture
Upload 31 files
f62a6a7 verified
Raw
History Blame Contribute Delete
9.15 kB
"""
code_parser.py
---------------
Turns raw Python source into a NETVIS "graph" dict.
IMPORTANT: this module only calls `ast.parse`. It never calls `exec`,
`eval`, or imports the user's module. Static analysis only — uploaded
code is data, never executed.
"""
import ast
from .detectors import (
LAYER_CATEGORY,
SKLEARN_ESTIMATORS,
PLOT_SIGNALS,
THREE_D_MARKERS,
)
FRAMEWORK_IMPORT_MAP = {
"torch": "pytorch",
"tensorflow": "keras",
"keras": "keras",
"sklearn": "sklearn",
"xgboost": "xgboost",
"lightgbm": "lightgbm",
"catboost": "catboost",
}
PLOT_IMPORT_MAP = {"matplotlib": "matplotlib", "seaborn": "seaborn", "plotly": "plotly"}
def _dotted(node):
if isinstance(node, ast.Name):
return node.id
if isinstance(node, ast.Attribute):
base = _dotted(node.value)
return f"{base}.{node.attr}" if base else node.attr
return ""
def _leaf(node):
if isinstance(node, ast.Name):
return node.id
if isinstance(node, ast.Attribute):
return node.attr
return None
def _call_leaf(call):
return _leaf(call.func) if isinstance(call, ast.Call) else None
def _args_repr(call, max_items=4):
parts = []
try:
for a in call.args[:max_items]:
parts.append(ast.unparse(a))
for kw in call.keywords[:max_items]:
if kw.arg is None:
continue
parts.append(f"{kw.arg}={ast.unparse(kw.value)}")
except Exception:
pass
return ", ".join(parts) if parts else "—"
def _collect_imports(tree):
roots = set()
for node in ast.walk(tree):
if isinstance(node, ast.Import):
for alias in node.names:
roots.add(alias.name.split(".")[0])
elif isinstance(node, ast.ImportFrom) and node.module:
roots.add(node.module.split(".")[0])
return roots
def _layers_from_init(class_node):
"""Find self.x = Layer(...) assignments inside a class's __init__."""
init_fn = next(
(n for n in class_node.body if isinstance(n, ast.FunctionDef) and n.name == "__init__"),
None,
)
if init_fn is None:
return []
layers = []
for node in ast.walk(init_fn):
if not isinstance(node, ast.Assign) or not isinstance(node.value, ast.Call):
continue
target = node.targets[0]
if not (isinstance(target, ast.Attribute) and isinstance(target.value, ast.Name) and target.value.id == "self"):
continue
leaf = _call_leaf(node.value)
if leaf is None:
continue
layers.append({
"attr": target.attr,
"leaf": leaf,
"full": _dotted(node.value.func),
"params": _args_repr(node.value),
"lineno": node.lineno,
})
return layers
def _sequential_blocks(tree):
"""Find nn.Sequential([...]) / keras.Sequential([...]) list literals anywhere."""
blocks = []
for node in ast.walk(tree):
if isinstance(node, ast.Call) and _call_leaf(node) == "Sequential":
container = node.args[0] if node.args else None
if isinstance(container, (ast.List, ast.Tuple)):
layers = []
for el in container.elts:
if isinstance(el, ast.Call):
layers.append({
"attr": None,
"leaf": _call_leaf(el),
"full": _dotted(el.func),
"params": _args_repr(el),
"lineno": el.lineno,
})
if layers:
blocks.append(layers)
return blocks
def _pipeline_blocks(tree):
"""Find sklearn Pipeline([('name', Estimator(...)), ...]) literals."""
blocks = []
for node in ast.walk(tree):
if isinstance(node, ast.Call) and _call_leaf(node) == "Pipeline":
container = node.args[0] if node.args else None
if isinstance(container, (ast.List, ast.Tuple)):
steps = []
for el in container.elts:
if isinstance(el, ast.Tuple) and len(el.elts) == 2 and isinstance(el.elts[1], ast.Call):
call = el.elts[1]
steps.append({
"attr": None,
"leaf": _call_leaf(call),
"full": _dotted(call.func),
"params": _args_repr(call),
"lineno": call.lineno,
})
if steps:
blocks.append(steps)
return blocks
def _top_level_estimators(tree):
found = []
seen_lines = set()
for node in ast.walk(tree):
if isinstance(node, ast.Call):
leaf = _call_leaf(node)
if leaf in SKLEARN_ESTIMATORS and node.lineno not in seen_lines:
seen_lines.add(node.lineno)
found.append({
"attr": None, "leaf": leaf, "full": _dotted(node.func),
"params": _args_repr(node), "lineno": node.lineno,
})
return found
def _plot_calls(tree, source):
calls = []
is_3d = any(marker in source for marker in THREE_D_MARKERS)
for node in ast.walk(tree):
if isinstance(node, ast.Call):
leaf = _call_leaf(node)
if leaf in PLOT_SIGNALS:
calls.append({"call": _dotted(node.func) or leaf, "tag": PLOT_SIGNALS[leaf], "lineno": node.lineno})
return calls, is_3d
def _best_class(tree):
"""Pick the class definition that looks most like a model (most layers found)."""
best, best_len = None, -1
for node in ast.walk(tree):
if isinstance(node, ast.ClassDef):
bases = [_dotted(b) for b in node.bases]
layers = _layers_from_init(node)
is_model_like = any("Module" in b or "Model" in b or "Layer" in b for b in bases) or len(layers) > 0
if is_model_like and len(layers) > best_len:
best, best_len = {"name": node.name, "bases": bases, "layers": layers}, len(layers)
return best
def _normalize_nodes(raw_layers):
nodes, edges = [], []
for i, l in enumerate(raw_layers):
category = LAYER_CATEGORY.get(l["leaf"], "other")
nodes.append({
"id": f"L{i+1:02d}",
"label": l["leaf"] or "Layer",
"attr": l.get("attr"),
"category": category,
"params": l["params"],
"order": i,
})
if i > 0:
edges.append([f"L{i:02d}", f"L{i+1:02d}"])
return nodes, edges
def analyze_python_source(source, filename="uploaded.py"):
try:
tree = ast.parse(source)
except SyntaxError as e:
return {"error": f"Syntax error at line {e.lineno}: {e.msg}", "filename": filename}
imports = _collect_imports(tree)
framework = next((FRAMEWORK_IMPORT_MAP[r] for r in imports if r in FRAMEWORK_IMPORT_MAP), None)
raw_layers = None
source_kind = None # 'class' | 'sequential' | 'pipeline' | 'estimator' | 'plot'
cls = _best_class(tree)
if cls and cls["layers"]:
raw_layers = cls["layers"]
source_kind = "class"
seq_blocks = _sequential_blocks(tree)
if raw_layers is None and seq_blocks:
raw_layers = max(seq_blocks, key=len)
source_kind = "sequential"
pipe_blocks = _pipeline_blocks(tree)
estimators = _top_level_estimators(tree)
plot_calls, is_3d = _plot_calls(tree, source)
nodes, edges, kind = [], [], "unknown"
if raw_layers:
nodes, edges = _normalize_nodes(raw_layers)
kind = "neural_network"
if framework is None:
framework = "pytorch" if "torch" in imports else ("keras" if ({"tensorflow", "keras"} & imports) else None)
elif pipe_blocks:
nodes, edges = _normalize_nodes(max(pipe_blocks, key=len))
kind = "sklearn_pipeline"
framework = framework or "sklearn"
elif estimators:
nodes, edges = _normalize_nodes(estimators[:1])
kind = "sklearn_model"
framework = framework or "sklearn"
elif plot_calls:
kind = "plot_3d" if is_3d else "plot"
framework = framework or next((PLOT_IMPORT_MAP[r] for r in imports if r in PLOT_IMPORT_MAP), "matplotlib")
plot_tally = {}
for p in plot_calls:
plot_tally[p["tag"]] = plot_tally.get(p["tag"], 0) + 1
class_names = [n.name for n in ast.walk(tree) if isinstance(n, ast.ClassDef)]
fn_names = [n.name for n in tree.body if isinstance(n, ast.FunctionDef)]
return {
"filename": filename,
"kind": kind,
"framework": framework,
"source_kind": source_kind,
"nodes": nodes,
"edges": edges,
"plot_signals": [{"tag": k, "count": v} for k, v in plot_tally.items()],
"is_3d": is_3d,
"summary": {
"lines": source.count("\n") + 1,
"classes": class_names,
"functions": fn_names,
"imports": sorted(imports),
},
}