Skip to content

Commit b3d79af

Browse files
r41k0uclaude
andcommitted
Core: Count scratch temps needed by if conditions
The per-statement temp count returned early for an if after recursing into its body and else, and never walked the condition. A helper call whose only occurrence was in a condition, `if m.lookup(5):`, therefore found no pool and failed with "Scratch pool exhausted" while the same call on a line of its own compiled. The condition is evaluated before any nested statement and the pool resets between statements, so it is counted as a piece of its own and the bodies take the max, as before. Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
1 parent a949a90 commit b3d79af

1 file changed

Lines changed: 14 additions & 10 deletions

File tree

‎pythonbpf/functions/functions_pass.py‎

Lines changed: 14 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -109,23 +109,27 @@ def merge_type_counts(count_dict):
109109
for typ, cnt in count_dict.items():
110110
max_temps_needed[typ] = max(max_temps_needed.get(typ, 0), cnt)
111111

112-
def update_max_temps_for_stmt(stmt):
113-
nonlocal max_temps_needed
112+
def count_temps_in(node):
113+
"""Temps one statement-sized piece of AST needs, merged into the max."""
114+
temps = {}
115+
for sub in ast.walk(node):
116+
if isinstance(sub, ast.Call):
117+
for typ, cnt in count_temps_in_call(sub, local_sym_tab).items():
118+
temps[typ] = temps.get(typ, 0) + cnt
119+
merge_type_counts(temps)
114120

121+
def update_max_temps_for_stmt(stmt):
115122
if isinstance(stmt, ast.If):
123+
# The condition is evaluated before any nested statement and the
124+
# pool resets between statements, so it counts as a piece of its
125+
# own; the bodies then take the max. `elif` is an If in orelse.
126+
count_temps_in(stmt.test)
116127
for s in stmt.body:
117128
update_max_temps_for_stmt(s)
118129
for s in stmt.orelse:
119130
update_max_temps_for_stmt(s)
120131
return
121-
122-
stmt_temps = {}
123-
for node in ast.walk(stmt):
124-
if isinstance(node, ast.Call):
125-
call_temps = count_temps_in_call(node, local_sym_tab)
126-
for typ, cnt in call_temps.items():
127-
stmt_temps[typ] = stmt_temps.get(typ, 0) + cnt
128-
merge_type_counts(stmt_temps)
132+
count_temps_in(stmt)
129133

130134
for stmt in body:
131135
update_max_temps_for_stmt(stmt)

0 commit comments

Comments
 (0)