diff --git a/.jules/bolt.md b/.jules/bolt.md new file mode 100644 index 00000000..146383f3 --- /dev/null +++ b/.jules/bolt.md @@ -0,0 +1,3 @@ +## 2026-06-11 - [Fast AST Iteration] +**Learning:** `ast.iter_child_nodes` is a major performance bottleneck for AST walking in Python because it relies on `ast.iter_fields`, which uses relatively slow string-based `getattr()` calls dynamically for every field. +**Action:** When walking massive AST trees, use a custom inline `fast_iter_child_nodes` generator that loops over `node._fields` directly, handling `AttributeError` instead of using the slower `iter_fields`. diff --git a/src/wardline/scanner/analyzer.py b/src/wardline/scanner/analyzer.py index 3fc88fc4..46f7d177 100644 --- a/src/wardline/scanner/analyzer.py +++ b/src/wardline/scanner/analyzer.py @@ -18,6 +18,7 @@ from wardline.core.finding import ENGINE_PATH, Finding, Kind, Location, Severity from wardline.core.taints import TaintState, combine +from wardline.scanner.ast_primitives import fast_iter_child_nodes from wardline.scanner.context import AnalysisContext, RuleRegistry from wardline.scanner.diagnostics import ( build_diagnostic_findings, @@ -281,7 +282,7 @@ def _bind_call_site_arguments_to_parameters( def _iter_l2_body_nodes(node: ast.FunctionDef | ast.AsyncFunctionDef) -> Iterator[ast.AST]: def walk(current: ast.AST) -> Iterator[ast.AST]: - for child in ast.iter_child_nodes(current): + for child in fast_iter_child_nodes(current): if isinstance(child, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef, ast.Lambda)): continue yield child diff --git a/src/wardline/scanner/ast_primitives.py b/src/wardline/scanner/ast_primitives.py index 70f565b3..ac538b3e 100644 --- a/src/wardline/scanner/ast_primitives.py +++ b/src/wardline/scanner/ast_primitives.py @@ -46,7 +46,7 @@ def build_import_alias_map( """ alias_map: dict[str, str] = {} - for node in ast.iter_child_nodes(tree): + for node in fast_iter_child_nodes(tree): if isinstance(node, ast.Import): for alias in node.names: local_name = alias.asname if alias.asname else alias.name.split(".")[0] @@ -93,6 +93,21 @@ def build_import_alias_map( return alias_map + +def fast_iter_child_nodes(node: ast.AST) -> Iterator[ast.AST]: + """Faster alternative to fast_iter_child_nodes() that avoids slow hasattr() checks.""" + for field in node._fields: + try: + value = getattr(node, field) + except AttributeError: + continue + if isinstance(value, list): + for item in value: + if isinstance(item, ast.AST): + yield item + elif isinstance(value, ast.AST): + yield value + def iter_calls_in_function_body( node: ast.FunctionDef | ast.AsyncFunctionDef, ) -> Iterator[ast.Call]: @@ -124,7 +139,7 @@ def walk_node(current: ast.AST) -> Iterator[ast.Call]: return if isinstance(current, ast.Call): yield current - for child in ast.iter_child_nodes(current): + for child in fast_iter_child_nodes(current): yield from walk_node(child) def _walk_argument_defaults(args: ast.arguments) -> Iterator[ast.Call]: diff --git a/src/wardline/scanner/flow_trace.py b/src/wardline/scanner/flow_trace.py index 30bd744a..3bc32094 100644 --- a/src/wardline/scanner/flow_trace.py +++ b/src/wardline/scanner/flow_trace.py @@ -9,6 +9,7 @@ from wardline.core.finding import Finding, Location from wardline.core.qualname import module_dotted_name from wardline.core.taints import RAW_ZONE, TRUST_RANK, TaintState +from wardline.scanner.ast_primitives import fast_iter_child_nodes from wardline.scanner.context import AnalysisContext from wardline.scanner.rules._sink_helpers import dotted_name @@ -66,7 +67,7 @@ def _find_assignment_callee(nodes: Sequence[ast.AST], name: str, entity_node: as callee = _return_callee(node.value) if callee is not None and any(isinstance(t, ast.Name) and t.id == name for t in node.targets): result = callee - for child in ast.iter_child_nodes(node): + for child in fast_iter_child_nodes(node): nested = _find_assignment_callee([child] if isinstance(child, ast.stmt) else [], name, entity_node) if nested is not None: result = nested @@ -98,7 +99,7 @@ def visit(node: ast.AST) -> None: return if isinstance(node, ast.Call) and getattr(node, "lineno", None) == line: calls_at_line.append(node) - for child in ast.iter_child_nodes(node): + for child in fast_iter_child_nodes(node): visit(child) visit(entity.node) @@ -123,7 +124,7 @@ def find_stmt(node: ast.AST, cur_stmt: ast.stmt | None = None) -> None: if node is call: stmt_at_line = new_stmt return - for child in ast.iter_child_nodes(node): + for child in fast_iter_child_nodes(node): find_stmt(child, new_stmt) find_stmt(entity.node) diff --git a/src/wardline/scanner/index.py b/src/wardline/scanner/index.py index de38cd51..0b4c4d9c 100644 --- a/src/wardline/scanner/index.py +++ b/src/wardline/scanner/index.py @@ -15,6 +15,7 @@ from wardline.core.finding import Location from wardline.core.qualname import is_overload_stub, reconstruct_qualname +from wardline.scanner.ast_primitives import fast_iter_child_nodes @dataclass(frozen=True, slots=True) @@ -40,7 +41,7 @@ def discover_class_qualnames(tree: ast.Module, *, module: str) -> set[str]: classes: set[str] = set() def visit(node: ast.AST, scope: list[ast.AST]) -> None: - for child in ast.iter_child_nodes(node): + for child in fast_iter_child_nodes(node): if isinstance(child, ast.ClassDef): local = reconstruct_qualname(child.name, list(reversed(scope))) classes.add(f"{module}.{local}") @@ -123,7 +124,7 @@ def add_or_replace_entity(entity: Entity) -> None: entities.append(entity) def visit(node: ast.AST, scope: list[ast.AST], *, parent_is_class: bool) -> None: - for child in ast.iter_child_nodes(node): + for child in fast_iter_child_nodes(node): if isinstance(child, (ast.FunctionDef, ast.AsyncFunctionDef)): if not is_overload_stub(child): # ``scope`` is outermost->innermost; reconstruct wants diff --git a/src/wardline/scanner/rules/_ast_helpers.py b/src/wardline/scanner/rules/_ast_helpers.py index 88335c67..4b118732 100644 --- a/src/wardline/scanner/rules/_ast_helpers.py +++ b/src/wardline/scanner/rules/_ast_helpers.py @@ -12,6 +12,8 @@ import ast from typing import TYPE_CHECKING +from wardline.scanner.ast_primitives import fast_iter_child_nodes + if TYPE_CHECKING: from collections.abc import Iterator @@ -21,7 +23,7 @@ def _own_statements(node: ast.AST) -> Iterator[ast.stmt]: """Yield every statement in *node*'s own scope, not descending into nested def/class bodies. Includes the bodies of if/for/while/try/with at any depth.""" - for child in ast.iter_child_nodes(node): + for child in fast_iter_child_nodes(node): if isinstance(child, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef)): continue if isinstance(child, ast.stmt): @@ -150,7 +152,7 @@ def own_nodes(node: ast.AST) -> Iterator[ast.AST]: def _walk_own(node: ast.AST) -> Iterator[ast.AST]: - for child in ast.iter_child_nodes(node): + for child in fast_iter_child_nodes(node): if isinstance(child, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef, ast.Lambda)): yield child else: diff --git a/src/wardline/scanner/rules/_sink_helpers.py b/src/wardline/scanner/rules/_sink_helpers.py index f830eac7..56ceb97b 100644 --- a/src/wardline/scanner/rules/_sink_helpers.py +++ b/src/wardline/scanner/rules/_sink_helpers.py @@ -24,6 +24,7 @@ from wardline.core.finding import Finding, Kind, Location, Severity from wardline.core.finding import compute_finding_fingerprint as _fp from wardline.core.taints import RAW_ZONE, TRUST_RANK, TaintState +from wardline.scanner.ast_primitives import fast_iter_child_nodes from wardline.scanner.rules.severity_model import modulate if TYPE_CHECKING: @@ -79,7 +80,7 @@ def _own_calls(node: ast.AST) -> Iterator[ast.Call]: the entity index does not emit separate lambda entities; skipping them would hide dangerous calls from sink rules. """ - for child in ast.iter_child_nodes(node): + for child in fast_iter_child_nodes(node): if isinstance(child, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef)): continue if isinstance(child, ast.Call): diff --git a/src/wardline/scanner/rules/none_leak.py b/src/wardline/scanner/rules/none_leak.py index 90ef4335..6241c664 100644 --- a/src/wardline/scanner/rules/none_leak.py +++ b/src/wardline/scanner/rules/none_leak.py @@ -27,6 +27,7 @@ from wardline.core.finding import Finding, Kind, Severity from wardline.core.finding import compute_finding_fingerprint as _fp from wardline.core.taints import RAW_ZONE, TRUST_RANK +from wardline.scanner.ast_primitives import fast_iter_child_nodes from wardline.scanner.rules._ast_helpers import _own_statements from wardline.scanner.rules.metadata import RuleMetadata @@ -129,7 +130,7 @@ def _promises_non_none( def _is_generator(node: ast.AST) -> bool: """True if *node*'s own scope contains a ``yield``/``yield from`` (does not descend into nested def/class/lambda — those are separate scopes).""" - for child in ast.iter_child_nodes(node): + for child in fast_iter_child_nodes(node): if isinstance(child, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef, ast.Lambda)): continue if isinstance(child, (ast.Yield, ast.YieldFrom)) or _is_generator(child): diff --git a/src/wardline/scanner/taint/callgraph.py b/src/wardline/scanner/taint/callgraph.py index 76ea742a..5f5fdc77 100644 --- a/src/wardline/scanner/taint/callgraph.py +++ b/src/wardline/scanner/taint/callgraph.py @@ -24,6 +24,7 @@ from collections.abc import Iterator, Sequence from wardline.scanner.ast_primitives import ( + fast_iter_child_nodes, iter_calls_in_function_body, resolve_call_fqn, resolve_self_method_fqn, @@ -32,7 +33,7 @@ def _own_scope_nodes(node: ast.AST) -> Iterator[ast.AST]: - for child in ast.iter_child_nodes(node): + for child in fast_iter_child_nodes(node): if isinstance(child, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef, ast.Lambda)): continue yield child diff --git a/src/wardline/scanner/taint/variable_level.py b/src/wardline/scanner/taint/variable_level.py index 25b8dbd2..0e2a6553 100644 --- a/src/wardline/scanner/taint/variable_level.py +++ b/src/wardline/scanner/taint/variable_level.py @@ -27,6 +27,7 @@ from typing import TYPE_CHECKING from wardline.core.taints import _PROVENANCE_CLASH, TRUST_RANK, TaintState, combine +from wardline.scanner.ast_primitives import fast_iter_child_nodes if TYPE_CHECKING: from collections.abc import Iterator @@ -171,7 +172,7 @@ def _own_scope_lambdas(node: ast.AST) -> Iterator[ast.Lambda]: """Yield every ``ast.Lambda`` in *node*'s own scope (descends into lambdas, which are not separate entities, but NOT into nested ``def``/``class`` — those are analyzed as their own entities).""" - for child in ast.iter_child_nodes(node): + for child in fast_iter_child_nodes(node): if isinstance(child, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef)): continue if isinstance(child, ast.Lambda): @@ -603,7 +604,7 @@ def _resolve_comprehension( def _name_bound_by_walrus(node: ast.AST, name: str) -> bool: """True if *name* is the target of a NamedExpr anywhere in *node* (not inside a nested Lambda — those bind the lambda's scope).""" - for child in ast.iter_child_nodes(node): + for child in fast_iter_child_nodes(node): if isinstance(child, ast.Lambda): continue if isinstance(child, ast.NamedExpr) and isinstance(child.target, ast.Name) and child.target.id == name: @@ -878,7 +879,7 @@ def _walk_exprs_for_walrus( function's, so it must not leak into ``var_taints``. Comprehension walruses DO bind the enclosing scope (PEP 572) and are intentionally still captured. """ - for child in ast.iter_child_nodes(node): + for child in fast_iter_child_nodes(node): if isinstance(child, ast.Lambda): continue # separate scope — its walruses don't bind here if isinstance(child, ast.NamedExpr): @@ -1526,7 +1527,7 @@ def collect_attribute_writes( var_types: dict[str, str] = {} def _walk(node: ast.AST) -> None: - for child in ast.iter_child_nodes(node): + for child in fast_iter_child_nodes(node): if isinstance(child, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef, ast.Lambda)): continue @@ -1725,7 +1726,7 @@ def _assignment_callee( ): result = callee nested = _assignment_callee( - list(ast.iter_child_nodes(node)), name, worst, function_taint, taint_map, var_taints + list(fast_iter_child_nodes(node)), name, worst, function_taint, taint_map, var_taints ) if nested is not None: result = nested @@ -1771,4 +1772,4 @@ def _collect_return_paths( if isinstance(node, (ast.Return, ast.Yield, ast.YieldFrom)) and node.value is not None: taint = _resolve_expr(node.value, function_taint, taint_map, var_taints) out.append((taint, _return_callee(node.value), node.value)) - _collect_return_paths(list(ast.iter_child_nodes(node)), function_taint, taint_map, var_taints, out) + _collect_return_paths(list(fast_iter_child_nodes(node)), function_taint, taint_map, var_taints, out)