astscribe 0.8.2__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- astscribe/__init__.py +37 -0
- astscribe/__main__.py +3 -0
- astscribe/api.py +52 -0
- astscribe/cli.py +84 -0
- astscribe/dependency.py +351 -0
- astscribe/diagnostics.py +157 -0
- astscribe/impact.py +220 -0
- astscribe/ipython/__init__.py +3 -0
- astscribe/ipython/magic.py +18 -0
- astscribe/methodology.py +244 -0
- astscribe/notebook.py +233 -0
- astscribe/parser/__init__.py +13 -0
- astscribe/parser/ast_parser.py +24 -0
- astscribe/parser/imports.py +32 -0
- astscribe/parser/symbols.py +155 -0
- astscribe/patterns/__init__.py +4 -0
- astscribe/patterns/evaluation.py +7 -0
- astscribe/patterns/inference.py +19 -0
- astscribe/patterns/training.py +23 -0
- astscribe/pipeline.py +144 -0
- astscribe/py.typed +1 -0
- astscribe/renderers/__init__.py +20 -0
- astscribe/renderers/concise.py +13 -0
- astscribe/renderers/educational.py +11 -0
- astscribe/renderers/scientific.py +43 -0
- astscribe/semantics/__init__.py +3 -0
- astscribe/semantics/datasets.py +260 -0
- astscribe/semantics/peft.py +372 -0
- astscribe/semantics/pytorch.py +615 -0
- astscribe/semantics/pytorch_experiment.py +418 -0
- astscribe/semantics/pytorch_reproducibility.py +149 -0
- astscribe/semantics/registry.py +54 -0
- astscribe/semantics/transformers.py +520 -0
- astscribe/semantics/transformers_models.py +197 -0
- astscribe/semantics/transformers_quantization.py +212 -0
- astscribe/sir/__init__.py +19 -0
- astscribe/sir/nodes.py +90 -0
- astscribe/techniques.py +110 -0
- astscribe-0.8.2.dist-info/METADATA +134 -0
- astscribe-0.8.2.dist-info/RECORD +45 -0
- astscribe-0.8.2.dist-info/WHEEL +5 -0
- astscribe-0.8.2.dist-info/entry_points.txt +2 -0
- astscribe-0.8.2.dist-info/licenses/LICENSE +201 -0
- astscribe-0.8.2.dist-info/licenses/NOTICE +5 -0
- astscribe-0.8.2.dist-info/top_level.txt +1 -0
astscribe/__init__.py
ADDED
|
@@ -0,0 +1,37 @@
|
|
|
1
|
+
from .api import analyze, explain
|
|
2
|
+
from .dependency import (
|
|
3
|
+
CellDependencyEdge,
|
|
4
|
+
CellDependencyNode,
|
|
5
|
+
NotebookDependencyGraph,
|
|
6
|
+
SymbolRedefinition,
|
|
7
|
+
)
|
|
8
|
+
from .diagnostics import NotebookDiagnostic, NotebookDiagnostics
|
|
9
|
+
from .impact import CellImpactSummary, ImpactHop, ImpactPath, NotebookImpactReport
|
|
10
|
+
from .methodology import MethodologyReport, MethodologySection
|
|
11
|
+
from .notebook import NotebookAnalyzer, SkippedCell
|
|
12
|
+
from .pipeline import ExperimentPipeline, PipelineStage
|
|
13
|
+
from .sir import AnalysisResult
|
|
14
|
+
|
|
15
|
+
__version__ = "0.8.2"
|
|
16
|
+
|
|
17
|
+
__all__ = [
|
|
18
|
+
"AnalysisResult",
|
|
19
|
+
"CellDependencyEdge",
|
|
20
|
+
"CellDependencyNode",
|
|
21
|
+
"CellImpactSummary",
|
|
22
|
+
"ExperimentPipeline",
|
|
23
|
+
"ImpactHop",
|
|
24
|
+
"ImpactPath",
|
|
25
|
+
"MethodologyReport",
|
|
26
|
+
"MethodologySection",
|
|
27
|
+
"NotebookAnalyzer",
|
|
28
|
+
"NotebookDependencyGraph",
|
|
29
|
+
"NotebookDiagnostic",
|
|
30
|
+
"NotebookDiagnostics",
|
|
31
|
+
"NotebookImpactReport",
|
|
32
|
+
"PipelineStage",
|
|
33
|
+
"SkippedCell",
|
|
34
|
+
"SymbolRedefinition",
|
|
35
|
+
"analyze",
|
|
36
|
+
"explain",
|
|
37
|
+
]
|
astscribe/__main__.py
ADDED
astscribe/api.py
ADDED
|
@@ -0,0 +1,52 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from astscribe.parser import (
|
|
4
|
+
ImportTable,
|
|
5
|
+
ParsedSource,
|
|
6
|
+
SymbolTable,
|
|
7
|
+
build_import_table,
|
|
8
|
+
build_symbol_table,
|
|
9
|
+
parse_source,
|
|
10
|
+
)
|
|
11
|
+
from astscribe.patterns import detect_inference, detect_training_step
|
|
12
|
+
from astscribe.semantics import DEFAULT_REGISTRY
|
|
13
|
+
from astscribe.sir import AnalysisResult
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
def _analyze_parsed(
|
|
17
|
+
parsed: ParsedSource,
|
|
18
|
+
imports: ImportTable,
|
|
19
|
+
symbols: SymbolTable,
|
|
20
|
+
) -> AnalysisResult:
|
|
21
|
+
result = AnalysisResult(source=parsed.source)
|
|
22
|
+
for framework, analyzers in DEFAULT_REGISTRY.analyzers().items():
|
|
23
|
+
framework_detected = False
|
|
24
|
+
for analyzer in analyzers:
|
|
25
|
+
semantic = analyzer(parsed, imports, symbols)
|
|
26
|
+
if semantic.operations or semantic.claims:
|
|
27
|
+
framework_detected = True
|
|
28
|
+
result.operations.extend(semantic.operations)
|
|
29
|
+
result.claims.extend(semantic.claims)
|
|
30
|
+
if framework_detected:
|
|
31
|
+
result.frameworks.append(framework)
|
|
32
|
+
|
|
33
|
+
result.training_step = detect_training_step(result.operations)
|
|
34
|
+
result.inference = detect_inference(result.operations)
|
|
35
|
+
return result
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
def analyze(source: str) -> AnalysisResult:
|
|
39
|
+
parsed = parse_source(source)
|
|
40
|
+
imports = build_import_table(parsed.tree)
|
|
41
|
+
symbols = build_symbol_table(parsed.tree, imports, source=source)
|
|
42
|
+
return _analyze_parsed(parsed, imports, symbols)
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
def explain(
|
|
46
|
+
source: str,
|
|
47
|
+
style: str = "scientific",
|
|
48
|
+
*,
|
|
49
|
+
return_result: bool = False,
|
|
50
|
+
) -> str | AnalysisResult:
|
|
51
|
+
result = analyze(source)
|
|
52
|
+
return result if return_result else result.render(style)
|
astscribe/cli.py
ADDED
|
@@ -0,0 +1,84 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import argparse
|
|
4
|
+
import sys
|
|
5
|
+
from collections.abc import Sequence
|
|
6
|
+
from pathlib import Path
|
|
7
|
+
|
|
8
|
+
from astscribe import NotebookAnalyzer, __version__, explain
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
def _parser() -> argparse.ArgumentParser:
|
|
12
|
+
parser = argparse.ArgumentParser(
|
|
13
|
+
prog="astscribe",
|
|
14
|
+
description="Explain Python ML code and inspect Jupyter notebooks without executing them.",
|
|
15
|
+
)
|
|
16
|
+
parser.add_argument(
|
|
17
|
+
"path",
|
|
18
|
+
type=Path,
|
|
19
|
+
help="Python source file, Jupyter notebook, or '-' to read Python source from stdin",
|
|
20
|
+
)
|
|
21
|
+
parser.add_argument("--version", action="version", version=f"%(prog)s {__version__}")
|
|
22
|
+
parser.add_argument(
|
|
23
|
+
"--style",
|
|
24
|
+
choices=("scientific", "educational", "concise"),
|
|
25
|
+
default="scientific",
|
|
26
|
+
help="rendering style for Python source (default: scientific)",
|
|
27
|
+
)
|
|
28
|
+
parser.add_argument(
|
|
29
|
+
"--report",
|
|
30
|
+
choices=("methodology", "pipeline", "techniques", "dependencies", "diagnostics", "impact"),
|
|
31
|
+
default="methodology",
|
|
32
|
+
help="notebook report to render (default: methodology)",
|
|
33
|
+
)
|
|
34
|
+
parser.add_argument("--cell", type=int, help="notebook cell for an impact report")
|
|
35
|
+
parser.add_argument(
|
|
36
|
+
"--evidence",
|
|
37
|
+
action="store_true",
|
|
38
|
+
help="include source evidence in methodology reports",
|
|
39
|
+
)
|
|
40
|
+
parser.add_argument(
|
|
41
|
+
"--strict",
|
|
42
|
+
action="store_true",
|
|
43
|
+
help="fail if a notebook code cell is not valid Python instead of skipping it",
|
|
44
|
+
)
|
|
45
|
+
return parser
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
def _render_notebook(args: argparse.Namespace) -> str:
|
|
49
|
+
notebook = NotebookAnalyzer.from_ipynb(
|
|
50
|
+
args.path,
|
|
51
|
+
skip_invalid_python=not args.strict,
|
|
52
|
+
)
|
|
53
|
+
if args.report == "methodology":
|
|
54
|
+
return notebook.render_methodology(include_evidence=args.evidence)
|
|
55
|
+
if args.report == "pipeline":
|
|
56
|
+
return notebook.render_pipeline()
|
|
57
|
+
if args.report == "techniques":
|
|
58
|
+
return notebook.render_techniques()
|
|
59
|
+
if args.report == "dependencies":
|
|
60
|
+
return notebook.render_dependency_graph()
|
|
61
|
+
if args.report == "diagnostics":
|
|
62
|
+
return notebook.render_diagnostics()
|
|
63
|
+
if args.cell is None:
|
|
64
|
+
raise ValueError("--cell is required when --report impact is selected")
|
|
65
|
+
return notebook.render_impact(args.cell)
|
|
66
|
+
|
|
67
|
+
|
|
68
|
+
def main(argv: Sequence[str] | None = None) -> int:
|
|
69
|
+
parser = _parser()
|
|
70
|
+
args = parser.parse_args(argv)
|
|
71
|
+
|
|
72
|
+
try:
|
|
73
|
+
if args.path.suffix.lower() == ".ipynb":
|
|
74
|
+
output = _render_notebook(args)
|
|
75
|
+
else:
|
|
76
|
+
source = sys.stdin.read() if str(args.path) == "-" else args.path.read_text(encoding="utf-8")
|
|
77
|
+
output = str(explain(source, style=args.style))
|
|
78
|
+
except (OSError, ValueError, SyntaxError) as exc:
|
|
79
|
+
parser.error(str(exc))
|
|
80
|
+
|
|
81
|
+
sys.stdout.write(output)
|
|
82
|
+
if output and not output.endswith("\n"):
|
|
83
|
+
sys.stdout.write("\n")
|
|
84
|
+
return 0
|
astscribe/dependency.py
ADDED
|
@@ -0,0 +1,351 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import ast
|
|
4
|
+
from dataclasses import asdict, dataclass
|
|
5
|
+
from typing import Any, Literal
|
|
6
|
+
|
|
7
|
+
EventKind = Literal["read", "write", "delete"]
|
|
8
|
+
_BUILTIN_NAMES = frozenset(dir(__import__("builtins"))) | {
|
|
9
|
+
"__name__",
|
|
10
|
+
"__file__",
|
|
11
|
+
"__package__",
|
|
12
|
+
}
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
@dataclass(frozen=True)
|
|
16
|
+
class SymbolEvent:
|
|
17
|
+
kind: EventKind
|
|
18
|
+
symbol: str
|
|
19
|
+
line: int | None
|
|
20
|
+
column: int | None
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
@dataclass(frozen=True)
|
|
24
|
+
class CellDependencyNode:
|
|
25
|
+
cell: int
|
|
26
|
+
defines: tuple[str, ...]
|
|
27
|
+
reads: tuple[str, ...]
|
|
28
|
+
unresolved_reads: tuple[str, ...]
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
@dataclass(frozen=True)
|
|
32
|
+
class CellDependencyEdge:
|
|
33
|
+
producer_cell: int
|
|
34
|
+
consumer_cell: int
|
|
35
|
+
symbols: tuple[str, ...]
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
@dataclass(frozen=True)
|
|
39
|
+
class SymbolRedefinition:
|
|
40
|
+
symbol: str
|
|
41
|
+
previous_cell: int
|
|
42
|
+
new_cell: int
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
@dataclass(frozen=True)
|
|
46
|
+
class NotebookDependencyGraph:
|
|
47
|
+
nodes: tuple[CellDependencyNode, ...]
|
|
48
|
+
edges: tuple[CellDependencyEdge, ...]
|
|
49
|
+
redefinitions: tuple[SymbolRedefinition, ...]
|
|
50
|
+
|
|
51
|
+
def to_dict(self) -> dict[str, Any]:
|
|
52
|
+
return asdict(self)
|
|
53
|
+
|
|
54
|
+
def parents(self, cell: int) -> tuple[int, ...]:
|
|
55
|
+
return tuple(
|
|
56
|
+
sorted(
|
|
57
|
+
{
|
|
58
|
+
edge.producer_cell
|
|
59
|
+
for edge in self.edges
|
|
60
|
+
if edge.consumer_cell == cell
|
|
61
|
+
}
|
|
62
|
+
)
|
|
63
|
+
)
|
|
64
|
+
|
|
65
|
+
def children(self, cell: int) -> tuple[int, ...]:
|
|
66
|
+
return tuple(
|
|
67
|
+
sorted(
|
|
68
|
+
{
|
|
69
|
+
edge.consumer_cell
|
|
70
|
+
for edge in self.edges
|
|
71
|
+
if edge.producer_cell == cell
|
|
72
|
+
}
|
|
73
|
+
)
|
|
74
|
+
)
|
|
75
|
+
|
|
76
|
+
def render(self) -> str:
|
|
77
|
+
if not self.nodes:
|
|
78
|
+
return "No analyzable notebook cells are available."
|
|
79
|
+
if not self.edges:
|
|
80
|
+
return "No cross-cell symbol dependencies were detected."
|
|
81
|
+
|
|
82
|
+
lines = ["# Cell dependency graph", ""]
|
|
83
|
+
for edge in self.edges:
|
|
84
|
+
symbols = ", ".join(edge.symbols)
|
|
85
|
+
lines.append(
|
|
86
|
+
f"Cell {edge.producer_cell} -> Cell {edge.consumer_cell} [{symbols}]"
|
|
87
|
+
)
|
|
88
|
+
return "\n".join(lines)
|
|
89
|
+
|
|
90
|
+
def to_dot(self) -> str:
|
|
91
|
+
lines = ["digraph ASTScribeNotebook {", " rankdir=LR;"]
|
|
92
|
+
for node in self.nodes:
|
|
93
|
+
lines.append(f' c{node.cell} [label="Cell {node.cell}"];')
|
|
94
|
+
for edge in self.edges:
|
|
95
|
+
symbols = ", ".join(edge.symbols).replace('"', '\\"')
|
|
96
|
+
lines.append(
|
|
97
|
+
f' c{edge.producer_cell} -> c{edge.consumer_cell} [label="{symbols}"];'
|
|
98
|
+
)
|
|
99
|
+
lines.append("}")
|
|
100
|
+
return "\n".join(lines)
|
|
101
|
+
|
|
102
|
+
|
|
103
|
+
class _EventCollector(ast.NodeVisitor):
|
|
104
|
+
def __init__(self) -> None:
|
|
105
|
+
self.events: list[SymbolEvent] = []
|
|
106
|
+
self._locals: list[set[str]] = []
|
|
107
|
+
|
|
108
|
+
def _emit(self, kind: EventKind, symbol: str, node: ast.AST) -> None:
|
|
109
|
+
self.events.append(
|
|
110
|
+
SymbolEvent(
|
|
111
|
+
kind=kind,
|
|
112
|
+
symbol=symbol,
|
|
113
|
+
line=getattr(node, "lineno", None),
|
|
114
|
+
column=getattr(node, "col_offset", None),
|
|
115
|
+
)
|
|
116
|
+
)
|
|
117
|
+
|
|
118
|
+
def _is_local(self, name: str) -> bool:
|
|
119
|
+
return any(name in scope for scope in reversed(self._locals))
|
|
120
|
+
|
|
121
|
+
def _write_target(self, node: ast.AST) -> None:
|
|
122
|
+
if isinstance(node, ast.Name):
|
|
123
|
+
if not self._locals:
|
|
124
|
+
self._emit("write", node.id, node)
|
|
125
|
+
else:
|
|
126
|
+
self._locals[-1].add(node.id)
|
|
127
|
+
return
|
|
128
|
+
if isinstance(node, ast.Tuple | ast.List):
|
|
129
|
+
for element in node.elts:
|
|
130
|
+
self._write_target(element)
|
|
131
|
+
return
|
|
132
|
+
if isinstance(node, ast.Starred):
|
|
133
|
+
self._write_target(node.value)
|
|
134
|
+
return
|
|
135
|
+
# Attribute/subscript assignment reads the base/index expressions but
|
|
136
|
+
# does not create a new notebook-global symbol.
|
|
137
|
+
self.visit(node)
|
|
138
|
+
|
|
139
|
+
def _delete_target(self, node: ast.AST) -> None:
|
|
140
|
+
if isinstance(node, ast.Name):
|
|
141
|
+
if not self._is_local(node.id):
|
|
142
|
+
self._emit("delete", node.id, node)
|
|
143
|
+
return
|
|
144
|
+
if isinstance(node, ast.Tuple | ast.List):
|
|
145
|
+
for element in node.elts:
|
|
146
|
+
self._delete_target(element)
|
|
147
|
+
return
|
|
148
|
+
self.visit(node)
|
|
149
|
+
|
|
150
|
+
def _visit_function_definition(
|
|
151
|
+
self,
|
|
152
|
+
node: ast.FunctionDef | ast.AsyncFunctionDef,
|
|
153
|
+
) -> None:
|
|
154
|
+
for decorator in node.decorator_list:
|
|
155
|
+
self.visit(decorator)
|
|
156
|
+
for default in (*node.args.defaults, *node.args.kw_defaults):
|
|
157
|
+
if default is not None:
|
|
158
|
+
self.visit(default)
|
|
159
|
+
self._emit("write", node.name, node)
|
|
160
|
+
|
|
161
|
+
def _visit_for(self, node: ast.For | ast.AsyncFor) -> None:
|
|
162
|
+
self.visit(node.iter)
|
|
163
|
+
self._write_target(node.target)
|
|
164
|
+
for statement in node.body:
|
|
165
|
+
self.visit(statement)
|
|
166
|
+
for statement in node.orelse:
|
|
167
|
+
self.visit(statement)
|
|
168
|
+
|
|
169
|
+
def _visit_with(self, node: ast.With | ast.AsyncWith) -> None:
|
|
170
|
+
for item in node.items:
|
|
171
|
+
self.visit(item.context_expr)
|
|
172
|
+
if item.optional_vars is not None:
|
|
173
|
+
self._write_target(item.optional_vars)
|
|
174
|
+
for statement in node.body:
|
|
175
|
+
self.visit(statement)
|
|
176
|
+
|
|
177
|
+
def visit_Name(self, node: ast.Name) -> None:
|
|
178
|
+
if isinstance(node.ctx, ast.Load) and not self._is_local(node.id):
|
|
179
|
+
self._emit("read", node.id, node)
|
|
180
|
+
elif isinstance(node.ctx, ast.Store):
|
|
181
|
+
self._write_target(node)
|
|
182
|
+
|
|
183
|
+
def visit_Assign(self, node: ast.Assign) -> None:
|
|
184
|
+
self.visit(node.value)
|
|
185
|
+
for target in node.targets:
|
|
186
|
+
self._write_target(target)
|
|
187
|
+
|
|
188
|
+
def visit_AnnAssign(self, node: ast.AnnAssign) -> None:
|
|
189
|
+
if node.value is not None:
|
|
190
|
+
self.visit(node.value)
|
|
191
|
+
self._write_target(node.target)
|
|
192
|
+
|
|
193
|
+
def visit_AugAssign(self, node: ast.AugAssign) -> None:
|
|
194
|
+
if isinstance(node.target, ast.Name) and not self._is_local(node.target.id):
|
|
195
|
+
self._emit("read", node.target.id, node.target)
|
|
196
|
+
else:
|
|
197
|
+
self.visit(node.target)
|
|
198
|
+
self.visit(node.value)
|
|
199
|
+
self._write_target(node.target)
|
|
200
|
+
|
|
201
|
+
def visit_NamedExpr(self, node: ast.NamedExpr) -> None:
|
|
202
|
+
self.visit(node.value)
|
|
203
|
+
self._write_target(node.target)
|
|
204
|
+
|
|
205
|
+
def visit_Delete(self, node: ast.Delete) -> None:
|
|
206
|
+
for target in node.targets:
|
|
207
|
+
self._delete_target(target)
|
|
208
|
+
|
|
209
|
+
def visit_Import(self, node: ast.Import) -> None:
|
|
210
|
+
for alias in node.names:
|
|
211
|
+
name = alias.asname or alias.name.split(".", 1)[0]
|
|
212
|
+
self._emit("write", name, node)
|
|
213
|
+
|
|
214
|
+
def visit_ImportFrom(self, node: ast.ImportFrom) -> None:
|
|
215
|
+
for alias in node.names:
|
|
216
|
+
if alias.name == "*":
|
|
217
|
+
continue
|
|
218
|
+
self._emit("write", alias.asname or alias.name, node)
|
|
219
|
+
|
|
220
|
+
def visit_FunctionDef(self, node: ast.FunctionDef) -> None:
|
|
221
|
+
self._visit_function_definition(node)
|
|
222
|
+
|
|
223
|
+
def visit_AsyncFunctionDef(self, node: ast.AsyncFunctionDef) -> None:
|
|
224
|
+
self._visit_function_definition(node)
|
|
225
|
+
|
|
226
|
+
def visit_Lambda(self, node: ast.Lambda) -> None:
|
|
227
|
+
for default in (*node.args.defaults, *node.args.kw_defaults):
|
|
228
|
+
if default is not None:
|
|
229
|
+
self.visit(default)
|
|
230
|
+
|
|
231
|
+
def visit_ClassDef(self, node: ast.ClassDef) -> None:
|
|
232
|
+
for decorator in node.decorator_list:
|
|
233
|
+
self.visit(decorator)
|
|
234
|
+
for base in node.bases:
|
|
235
|
+
self.visit(base)
|
|
236
|
+
for keyword in node.keywords:
|
|
237
|
+
self.visit(keyword.value)
|
|
238
|
+
self._emit("write", node.name, node)
|
|
239
|
+
|
|
240
|
+
def visit_For(self, node: ast.For) -> None:
|
|
241
|
+
self._visit_for(node)
|
|
242
|
+
|
|
243
|
+
def visit_AsyncFor(self, node: ast.AsyncFor) -> None:
|
|
244
|
+
self._visit_for(node)
|
|
245
|
+
|
|
246
|
+
def visit_With(self, node: ast.With) -> None:
|
|
247
|
+
self._visit_with(node)
|
|
248
|
+
|
|
249
|
+
def visit_AsyncWith(self, node: ast.AsyncWith) -> None:
|
|
250
|
+
self._visit_with(node)
|
|
251
|
+
|
|
252
|
+
def _visit_comprehension(
|
|
253
|
+
self,
|
|
254
|
+
generators: list[ast.comprehension],
|
|
255
|
+
final_nodes: tuple[ast.AST, ...],
|
|
256
|
+
) -> None:
|
|
257
|
+
self._locals.append(set())
|
|
258
|
+
try:
|
|
259
|
+
for generator in generators:
|
|
260
|
+
self.visit(generator.iter)
|
|
261
|
+
self._write_target(generator.target)
|
|
262
|
+
for condition in generator.ifs:
|
|
263
|
+
self.visit(condition)
|
|
264
|
+
for node in final_nodes:
|
|
265
|
+
self.visit(node)
|
|
266
|
+
finally:
|
|
267
|
+
self._locals.pop()
|
|
268
|
+
|
|
269
|
+
def visit_ListComp(self, node: ast.ListComp) -> None:
|
|
270
|
+
self._visit_comprehension(node.generators, (node.elt,))
|
|
271
|
+
|
|
272
|
+
def visit_SetComp(self, node: ast.SetComp) -> None:
|
|
273
|
+
self._visit_comprehension(node.generators, (node.elt,))
|
|
274
|
+
|
|
275
|
+
def visit_GeneratorExp(self, node: ast.GeneratorExp) -> None:
|
|
276
|
+
self._visit_comprehension(node.generators, (node.elt,))
|
|
277
|
+
|
|
278
|
+
def visit_DictComp(self, node: ast.DictComp) -> None:
|
|
279
|
+
self._visit_comprehension(node.generators, (node.key, node.value))
|
|
280
|
+
|
|
281
|
+
|
|
282
|
+
def collect_symbol_events(source: str) -> tuple[SymbolEvent, ...]:
|
|
283
|
+
tree = ast.parse(source)
|
|
284
|
+
collector = _EventCollector()
|
|
285
|
+
collector.visit(tree)
|
|
286
|
+
return tuple(collector.events)
|
|
287
|
+
|
|
288
|
+
|
|
289
|
+
def build_dependency_graph(
|
|
290
|
+
cells: tuple[tuple[int, str], ...],
|
|
291
|
+
) -> NotebookDependencyGraph:
|
|
292
|
+
producer: dict[str, int] = {}
|
|
293
|
+
nodes: list[CellDependencyNode] = []
|
|
294
|
+
edge_symbols: dict[tuple[int, int], set[str]] = {}
|
|
295
|
+
redefinitions: list[SymbolRedefinition] = []
|
|
296
|
+
|
|
297
|
+
for cell, source in cells:
|
|
298
|
+
events = collect_symbol_events(source)
|
|
299
|
+
defines: set[str] = set()
|
|
300
|
+
reads: set[str] = set()
|
|
301
|
+
unresolved: set[str] = set()
|
|
302
|
+
|
|
303
|
+
for event in events:
|
|
304
|
+
if event.kind == "read":
|
|
305
|
+
reads.add(event.symbol)
|
|
306
|
+
source_cell = producer.get(event.symbol)
|
|
307
|
+
if source_cell is None:
|
|
308
|
+
if event.symbol not in _BUILTIN_NAMES:
|
|
309
|
+
unresolved.add(event.symbol)
|
|
310
|
+
elif source_cell != cell:
|
|
311
|
+
edge_symbols.setdefault((source_cell, cell), set()).add(event.symbol)
|
|
312
|
+
continue
|
|
313
|
+
|
|
314
|
+
if event.kind == "delete":
|
|
315
|
+
producer.pop(event.symbol, None)
|
|
316
|
+
continue
|
|
317
|
+
|
|
318
|
+
defines.add(event.symbol)
|
|
319
|
+
previous_cell = producer.get(event.symbol)
|
|
320
|
+
if previous_cell is not None and previous_cell != cell:
|
|
321
|
+
redefinitions.append(
|
|
322
|
+
SymbolRedefinition(
|
|
323
|
+
symbol=event.symbol,
|
|
324
|
+
previous_cell=previous_cell,
|
|
325
|
+
new_cell=cell,
|
|
326
|
+
)
|
|
327
|
+
)
|
|
328
|
+
producer[event.symbol] = cell
|
|
329
|
+
|
|
330
|
+
nodes.append(
|
|
331
|
+
CellDependencyNode(
|
|
332
|
+
cell=cell,
|
|
333
|
+
defines=tuple(sorted(defines)),
|
|
334
|
+
reads=tuple(sorted(reads)),
|
|
335
|
+
unresolved_reads=tuple(sorted(unresolved)),
|
|
336
|
+
)
|
|
337
|
+
)
|
|
338
|
+
|
|
339
|
+
edges = tuple(
|
|
340
|
+
CellDependencyEdge(
|
|
341
|
+
producer_cell=producer_cell,
|
|
342
|
+
consumer_cell=consumer_cell,
|
|
343
|
+
symbols=tuple(sorted(symbols)),
|
|
344
|
+
)
|
|
345
|
+
for (producer_cell, consumer_cell), symbols in sorted(edge_symbols.items())
|
|
346
|
+
)
|
|
347
|
+
return NotebookDependencyGraph(
|
|
348
|
+
nodes=tuple(nodes),
|
|
349
|
+
edges=edges,
|
|
350
|
+
redefinitions=tuple(redefinitions),
|
|
351
|
+
)
|
astscribe/diagnostics.py
ADDED
|
@@ -0,0 +1,157 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from dataclasses import asdict, dataclass
|
|
4
|
+
from typing import Any, Literal
|
|
5
|
+
|
|
6
|
+
from astscribe.dependency import NotebookDependencyGraph
|
|
7
|
+
|
|
8
|
+
DiagnosticSeverity = Literal["info", "warning"]
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
@dataclass(frozen=True)
|
|
12
|
+
class NotebookDiagnostic:
|
|
13
|
+
code: str
|
|
14
|
+
severity: DiagnosticSeverity
|
|
15
|
+
message: str
|
|
16
|
+
cell: int
|
|
17
|
+
symbol: str
|
|
18
|
+
related_cell: int | None = None
|
|
19
|
+
|
|
20
|
+
def to_dict(self) -> dict[str, Any]:
|
|
21
|
+
return asdict(self)
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
@dataclass(frozen=True)
|
|
25
|
+
class NotebookDiagnostics:
|
|
26
|
+
items: tuple[NotebookDiagnostic, ...]
|
|
27
|
+
|
|
28
|
+
def to_dict(self) -> dict[str, Any]:
|
|
29
|
+
return {"items": [item.to_dict() for item in self.items]}
|
|
30
|
+
|
|
31
|
+
def by_code(self, code: str) -> tuple[NotebookDiagnostic, ...]:
|
|
32
|
+
return tuple(item for item in self.items if item.code == code)
|
|
33
|
+
|
|
34
|
+
def render(self) -> str:
|
|
35
|
+
if not self.items:
|
|
36
|
+
return "No supported notebook dependency diagnostics were detected."
|
|
37
|
+
|
|
38
|
+
lines = ["# Notebook diagnostics", ""]
|
|
39
|
+
for item in self.items:
|
|
40
|
+
label = item.severity.upper()
|
|
41
|
+
lines.append(f"- [{label}] cell {item.cell}: {item.message} (`{item.code}`)")
|
|
42
|
+
return "\n".join(lines)
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
def _definitions_by_symbol(graph: NotebookDependencyGraph) -> dict[str, tuple[int, ...]]:
|
|
46
|
+
cells: dict[str, list[int]] = {}
|
|
47
|
+
for node in graph.nodes:
|
|
48
|
+
for symbol in node.defines:
|
|
49
|
+
cells.setdefault(symbol, []).append(node.cell)
|
|
50
|
+
return {symbol: tuple(values) for symbol, values in cells.items()}
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
def _next_definition(
|
|
54
|
+
definitions: dict[str, tuple[int, ...]],
|
|
55
|
+
symbol: str,
|
|
56
|
+
cell: int,
|
|
57
|
+
) -> int | None:
|
|
58
|
+
return next(
|
|
59
|
+
(definition_cell for definition_cell in definitions.get(symbol, ()) if definition_cell > cell),
|
|
60
|
+
None,
|
|
61
|
+
)
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
def _definition_consumed_before_redefinition(
|
|
65
|
+
graph: NotebookDependencyGraph,
|
|
66
|
+
symbol: str,
|
|
67
|
+
producer_cell: int,
|
|
68
|
+
redefinition_cell: int,
|
|
69
|
+
) -> bool:
|
|
70
|
+
return any(
|
|
71
|
+
edge.producer_cell == producer_cell
|
|
72
|
+
and edge.consumer_cell <= redefinition_cell
|
|
73
|
+
and symbol in edge.symbols
|
|
74
|
+
for edge in graph.edges
|
|
75
|
+
)
|
|
76
|
+
|
|
77
|
+
|
|
78
|
+
def build_notebook_diagnostics(graph: NotebookDependencyGraph) -> NotebookDiagnostics:
|
|
79
|
+
definitions = _definitions_by_symbol(graph)
|
|
80
|
+
diagnostics: list[NotebookDiagnostic] = []
|
|
81
|
+
|
|
82
|
+
for node in graph.nodes:
|
|
83
|
+
for symbol in node.unresolved_reads:
|
|
84
|
+
later_cell = _next_definition(definitions, symbol, node.cell)
|
|
85
|
+
if later_cell is not None:
|
|
86
|
+
diagnostics.append(
|
|
87
|
+
NotebookDiagnostic(
|
|
88
|
+
code="dependency.forward_reference",
|
|
89
|
+
severity="warning",
|
|
90
|
+
message=(
|
|
91
|
+
f"`{symbol}` is read before its next supported definition in "
|
|
92
|
+
f"cell {later_cell}; source-order analysis cannot resolve this read."
|
|
93
|
+
),
|
|
94
|
+
cell=node.cell,
|
|
95
|
+
symbol=symbol,
|
|
96
|
+
related_cell=later_cell,
|
|
97
|
+
)
|
|
98
|
+
)
|
|
99
|
+
else:
|
|
100
|
+
diagnostics.append(
|
|
101
|
+
NotebookDiagnostic(
|
|
102
|
+
code="dependency.unresolved_symbol",
|
|
103
|
+
severity="warning",
|
|
104
|
+
message=(
|
|
105
|
+
f"`{symbol}` is read without a supported prior notebook definition; "
|
|
106
|
+
"it may come from external or hidden kernel state."
|
|
107
|
+
),
|
|
108
|
+
cell=node.cell,
|
|
109
|
+
symbol=symbol,
|
|
110
|
+
)
|
|
111
|
+
)
|
|
112
|
+
|
|
113
|
+
for redefinition in graph.redefinitions:
|
|
114
|
+
diagnostics.append(
|
|
115
|
+
NotebookDiagnostic(
|
|
116
|
+
code="dependency.symbol_redefinition",
|
|
117
|
+
severity="info",
|
|
118
|
+
message=(
|
|
119
|
+
f"`{redefinition.symbol}` replaces the definition from cell "
|
|
120
|
+
f"{redefinition.previous_cell}."
|
|
121
|
+
),
|
|
122
|
+
cell=redefinition.new_cell,
|
|
123
|
+
symbol=redefinition.symbol,
|
|
124
|
+
related_cell=redefinition.previous_cell,
|
|
125
|
+
)
|
|
126
|
+
)
|
|
127
|
+
|
|
128
|
+
if not _definition_consumed_before_redefinition(
|
|
129
|
+
graph,
|
|
130
|
+
redefinition.symbol,
|
|
131
|
+
redefinition.previous_cell,
|
|
132
|
+
redefinition.new_cell,
|
|
133
|
+
):
|
|
134
|
+
diagnostics.append(
|
|
135
|
+
NotebookDiagnostic(
|
|
136
|
+
code="dependency.overwritten_before_cross_cell_use",
|
|
137
|
+
severity="info",
|
|
138
|
+
message=(
|
|
139
|
+
f"The definition of `{redefinition.symbol}` in cell "
|
|
140
|
+
f"{redefinition.previous_cell} is replaced in cell "
|
|
141
|
+
f"{redefinition.new_cell} before any later analyzed cell consumes it."
|
|
142
|
+
),
|
|
143
|
+
cell=redefinition.new_cell,
|
|
144
|
+
symbol=redefinition.symbol,
|
|
145
|
+
related_cell=redefinition.previous_cell,
|
|
146
|
+
)
|
|
147
|
+
)
|
|
148
|
+
|
|
149
|
+
diagnostics.sort(
|
|
150
|
+
key=lambda item: (
|
|
151
|
+
item.cell,
|
|
152
|
+
item.code,
|
|
153
|
+
item.symbol,
|
|
154
|
+
item.related_cell if item.related_cell is not None else -1,
|
|
155
|
+
)
|
|
156
|
+
)
|
|
157
|
+
return NotebookDiagnostics(items=tuple(diagnostics))
|