Skip to content

Commit 24e7feb

Browse files
Core: Lower for-over-range and while loops, with break and continue
`for i in range(...)` steps a hidden induction counter, allocated in the entry block next to the loop variable, and copies it into the variable each iteration, so rebinding the variable in the body leaves the trip count alone as in Python (and the loop stays visibly bounded for the verifier). The counter is 64-bit, signed unless the bounds make C's arithmetic unsigned. Bounds are evaluated once, before the loop; the step must be a nonzero integer literal, because its sign picks the loop test. `while` re-evaluates its test at the top of each iteration. `break` and `continue` branch through a loop-target stack on the compilation context; outside a loop they are a SyntaxError, as in Python. A loop's else-branch runs on normal exit only. Anything but range() as the iterable is rejected with NotImplementedError. The allocation pass now descends into loop bodies and counts helper temps in loop headers. Checked against clang -O2 on tests/c-form/loops.bpf.c: the constant-bound cases fold to the same ret, and a helper-bounded loop gets the same rotated loop.
1 parent 919d994 commit 24e7feb

3 files changed

Lines changed: 292 additions & 10 deletions

File tree

‎pythonbpf/allocation_pass.py‎

Lines changed: 105 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,7 @@
88
from .expr import VmlinuxHandlerRegistry
99
from pythonbpf.type_deducer import ctypes_to_ir, is_ctypes, IntTy, signedness
1010
from pythonbpf.expr.type_inference import infer_int_type
11+
from pythonbpf.expr.operators import usual_arithmetic_conversions
1112
from pythonbpf.maps import BPFMapType
1213

1314
logger = logging.getLogger(__name__)
@@ -89,6 +90,110 @@ def handle_ann_assign_allocation(compilation_context, builder, stmt, local_sym_t
8990
)
9091

9192

93+
def handle_for_allocation(compilation_context, builder, stmt, local_sym_tab):
94+
"""Handle memory allocation for `for <name> in range(...)`.
95+
96+
Two slots, both in the entry block so a loop never grows the stack: the
97+
loop variable, and a hidden induction counter the loop actually steps.
98+
Keeping them apart is what makes rebinding the loop variable in the body
99+
leave the trip count alone, as in Python -- and a counter the body cannot
100+
touch is what keeps the loop visibly bounded for the verifier.
101+
"""
102+
start, stop, step = parse_range(stmt)
103+
104+
# range() yields Python ints; like any undeclared local they are 64-bit,
105+
# signed unless the bounds make C's arithmetic unsigned.
106+
bound_types = [
107+
infer_int_type(bound, local_sym_tab, compilation_context)
108+
for bound in (start, stop)
109+
if bound is not None
110+
]
111+
signed = True
112+
if all(ty is not None for ty in bound_types):
113+
common = bound_types[0]
114+
for ty in bound_types[1:]:
115+
common = usual_arithmetic_conversions(common, ty)
116+
signed = signedness(common)
117+
if not signed and step < 0:
118+
raise SyntaxError(
119+
f"range() on line {stmt.lineno} counts down over unsigned bounds; "
120+
f"the counter would wrap instead of stopping"
121+
)
122+
loop_ty = IntTy(64, signed)
123+
124+
counter = range_counter_name(stmt)
125+
local_sym_tab[counter] = LocalSymbol(
126+
_allocate_with_type(builder, counter, loop_ty), loop_ty
127+
)
128+
_bind_name(
129+
compilation_context,
130+
stmt,
131+
stmt.target,
132+
local_sym_tab,
133+
lambda var_name: local_sym_tab.__setitem__(
134+
var_name,
135+
LocalSymbol(_allocate_with_type(builder, var_name, loop_ty), loop_ty),
136+
),
137+
)
138+
139+
140+
def range_counter_name(stmt):
141+
"""Symbol-table name of a range loop's hidden induction counter, unique per
142+
loop so nested loops each get their own."""
143+
return f"__range_idx_{stmt.lineno}_{stmt.col_offset}"
144+
145+
146+
def parse_range(stmt):
147+
"""Split `for <name> in range(...)` into (start, stop, step): start is an
148+
expression or None (meaning 0), stop an expression, step a nonzero int.
149+
150+
step has to be known at compile time, because its sign decides whether the
151+
loop runs while the counter is below stop or above it.
152+
"""
153+
it = stmt.iter
154+
if not (
155+
isinstance(it, ast.Call)
156+
and isinstance(it.func, ast.Name)
157+
and it.func.id == "range"
158+
):
159+
raise NotImplementedError(
160+
f"for loop on line {stmt.lineno}: only range(...) can be iterated "
161+
f"so far, got {ast.unparse(it)}"
162+
)
163+
if not isinstance(stmt.target, ast.Name):
164+
raise NotImplementedError(
165+
f"for loop on line {stmt.lineno}: only a plain name can be the loop "
166+
f"variable, got {ast.unparse(stmt.target)}"
167+
)
168+
if it.keywords or not 1 <= len(it.args) <= 3:
169+
raise SyntaxError(
170+
f"range() on line {stmt.lineno} takes 1 to 3 positional arguments"
171+
)
172+
173+
if len(it.args) == 1:
174+
return None, it.args[0], 1
175+
start, stop = it.args[0], it.args[1]
176+
if len(it.args) == 2:
177+
return start, stop, 1
178+
179+
step_node = it.args[2]
180+
negate = isinstance(step_node, ast.UnaryOp) and isinstance(step_node.op, ast.USub)
181+
literal = step_node.operand if negate else step_node
182+
if not (
183+
isinstance(literal, ast.Constant)
184+
and isinstance(literal.value, int)
185+
and not isinstance(literal.value, bool)
186+
):
187+
raise SyntaxError(
188+
f"range() step on line {stmt.lineno} must be an integer literal, "
189+
f"got {ast.unparse(step_node)}"
190+
)
191+
step = -literal.value if negate else literal.value
192+
if step == 0:
193+
raise ValueError(f"range() arg 3 must not be zero (line {stmt.lineno})")
194+
return start, stop, step
195+
196+
92197
def annotation_to_ir(annotation, lineno):
93198
"""IR type for a ctypes annotation, written `c_int32` or `ctypes.c_int32`."""
94199
if isinstance(annotation, ast.Name):

‎pythonbpf/context.py‎

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -82,7 +82,12 @@ def __init__(self, module: ir.Module):
8282
# Current function context (optional, if needed globally during function processing)
8383
self.current_func = None
8484

85+
# Enclosing loops of the statement being lowered, innermost last, as
86+
# (continue target, break target) blocks.
87+
self.loop_stack: list[tuple[ir.Block, ir.Block]] = []
88+
8589
def reset(self):
8690
"""Reset state between functions if necessary, though new context per compile is preferred."""
8791
self.scratch_pool.reset()
8892
self.current_func = None
93+
self.loop_stack = []

‎pythonbpf/functions/functions_pass.py‎

Lines changed: 182 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -26,6 +26,9 @@
2626
from pythonbpf.allocation_pass import (
2727
handle_assign_allocation,
2828
handle_ann_assign_allocation,
29+
handle_for_allocation,
30+
parse_range,
31+
range_counter_name,
2932
allocate_temp_pool,
3033
create_targets_and_rvals,
3134
LocalSymbol,
@@ -82,10 +85,11 @@ def count_temps_in_call(call_node, local_sym_tab):
8285
return count
8386

8487

85-
def handle_if_allocation(
88+
def handle_block_allocation(
8689
compilation_context, builder, stmt, func, ret_type, local_sym_tab
8790
):
88-
"""Recursively handle allocations in if/else branches."""
91+
"""Recursively handle allocations in the body and else-branch of an
92+
if, for or while statement."""
8993
if stmt.body:
9094
allocate_mem(
9195
compilation_context,
@@ -117,15 +121,23 @@ def merge_type_counts(count_dict):
117121
def update_max_temps_for_stmt(stmt):
118122
nonlocal max_temps_needed
119123

120-
if isinstance(stmt, ast.If):
124+
if isinstance(stmt, (ast.If, ast.For, ast.While)):
125+
# A loop header is evaluated like a statement of its own: range()
126+
# bounds once before the loop, a while test on every iteration.
127+
header = {ast.For: "iter", ast.While: "test"}.get(type(stmt))
128+
if header is not None:
129+
count_temps_in_tree(getattr(stmt, header))
121130
for s in stmt.body:
122131
update_max_temps_for_stmt(s)
123132
for s in stmt.orelse:
124133
update_max_temps_for_stmt(s)
125134
return
126135

136+
count_temps_in_tree(stmt)
137+
138+
def count_temps_in_tree(tree):
127139
stmt_temps = {}
128-
for node in ast.walk(stmt):
140+
for node in ast.walk(tree):
129141
if isinstance(node, ast.Call):
130142
call_temps = count_temps_in_call(node, local_sym_tab)
131143
for typ, cnt in call_temps.items():
@@ -136,8 +148,10 @@ def update_max_temps_for_stmt(stmt):
136148
update_max_temps_for_stmt(stmt)
137149

138150
# Handle allocations
139-
if isinstance(stmt, ast.If):
140-
handle_if_allocation(
151+
if isinstance(stmt, ast.For):
152+
handle_for_allocation(compilation_context, builder, stmt, local_sym_tab)
153+
if isinstance(stmt, (ast.If, ast.For, ast.While)):
154+
handle_block_allocation(
141155
compilation_context,
142156
builder,
143157
stmt,
@@ -353,6 +367,158 @@ def handle_if(func, compilation_context, builder, stmt, local_sym_tab, ret_type)
353367
builder.position_at_end(merge_block)
354368

355369

370+
def _lower_loop(
371+
func,
372+
compilation_context,
373+
builder,
374+
stmt,
375+
local_sym_tab,
376+
ret_type,
377+
body_block,
378+
continue_block,
379+
end_block,
380+
else_block,
381+
):
382+
"""What for and while share once their header is emitted: the body, with
383+
`continue` and `break` bound to this loop, falling through to
384+
continue_block; then the else-branch, which runs only when the loop ends
385+
without a break, so it sits between the exit test and end_block."""
386+
builder.position_at_end(body_block)
387+
compilation_context.loop_stack.append((continue_block, end_block))
388+
try:
389+
process_block(
390+
func, compilation_context, builder, stmt.body, local_sym_tab, ret_type
391+
)
392+
finally:
393+
compilation_context.loop_stack.pop()
394+
if not builder.block.is_terminated:
395+
builder.branch(continue_block)
396+
397+
if else_block is not None:
398+
# Outside this loop's scope: a break here leaves the enclosing loop.
399+
builder.position_at_end(else_block)
400+
process_block(
401+
func, compilation_context, builder, stmt.orelse, local_sym_tab, ret_type
402+
)
403+
if not builder.block.is_terminated:
404+
builder.branch(end_block)
405+
406+
builder.position_at_end(end_block)
407+
408+
409+
def handle_while(func, compilation_context, builder, stmt, local_sym_tab, ret_type):
410+
"""Handle `while test: body [else: orelse]`. The test is re-evaluated at
411+
the top of every iteration, and is where `continue` goes."""
412+
cond_block = func.append_basic_block(name="while.cond")
413+
body_block = func.append_basic_block(name="while.body")
414+
else_block = func.append_basic_block(name="while.else") if stmt.orelse else None
415+
end_block = func.append_basic_block(name="while.end")
416+
417+
builder.branch(cond_block)
418+
builder.position_at_end(cond_block)
419+
cond = handle_cond(func, compilation_context, builder, stmt.test, local_sym_tab)
420+
builder.cbranch(cond, body_block, else_block or end_block)
421+
422+
_lower_loop(
423+
func,
424+
compilation_context,
425+
builder,
426+
stmt,
427+
local_sym_tab,
428+
ret_type,
429+
body_block,
430+
cond_block,
431+
end_block,
432+
else_block,
433+
)
434+
435+
436+
def handle_for(func, compilation_context, builder, stmt, local_sym_tab, ret_type):
437+
"""Handle `for name in range(...): body [else: orelse]`.
438+
439+
The allocation pass made a hidden induction counter (typed from the
440+
bounds) next to the loop variable. The bounds are evaluated once, before
441+
the loop, as Python does; each iteration copies the counter into the loop
442+
variable, and `continue` goes to the step, not straight back to the test.
443+
"""
444+
start, stop, step = parse_range(stmt)
445+
counter = local_sym_tab[range_counter_name(stmt)]
446+
loop_ty = counter.ir_type
447+
448+
def bound(expr):
449+
val, ty = get_typed_operand(
450+
func, compilation_context, expr, builder, local_sym_tab
451+
)
452+
if val is None or not isinstance(ty, ir.IntType):
453+
raise SyntaxError(
454+
f"range() bound on line {stmt.lineno} must be an integer: "
455+
f"{ast.unparse(expr)}"
456+
)
457+
return convert(builder, val, ty, loop_ty)
458+
459+
start_val = ir.Constant(loop_ty, 0) if start is None else bound(start)
460+
stop_val = bound(stop)
461+
builder.store(start_val, counter.var)
462+
463+
target = local_sym_tab[stmt.target.id]
464+
if target.var is None:
465+
raise SyntaxError(
466+
f"cannot use '{stmt.target.id}' as a loop variable: it is the "
467+
f"context parameter"
468+
)
469+
470+
cond_block = func.append_basic_block(name="for.cond")
471+
body_block = func.append_basic_block(name="for.body")
472+
inc_block = func.append_basic_block(name="for.inc")
473+
else_block = func.append_basic_block(name="for.else") if stmt.orelse else None
474+
end_block = func.append_basic_block(name="for.end")
475+
476+
builder.branch(cond_block)
477+
builder.position_at_end(cond_block)
478+
idx = builder.load(counter.var)
479+
# Counting up runs while below stop, counting down while above it.
480+
predicate = "<" if step > 0 else ">"
481+
compare = builder.icmp_signed if signedness(loop_ty) else builder.icmp_unsigned
482+
builder.cbranch(
483+
compare(predicate, idx, stop_val), body_block, else_block or end_block
484+
)
485+
486+
# The loop variable is bound to the counter's value, per iteration.
487+
builder.position_at_end(body_block)
488+
builder.store(
489+
convert(builder, builder.load(counter.var), loop_ty, target.ir_type),
490+
target.var,
491+
)
492+
493+
builder.position_at_end(inc_block)
494+
next_idx = builder.add(builder.load(counter.var), ir.Constant(loop_ty, step))
495+
builder.store(next_idx, counter.var)
496+
builder.branch(cond_block)
497+
498+
_lower_loop(
499+
func,
500+
compilation_context,
501+
builder,
502+
stmt,
503+
local_sym_tab,
504+
ret_type,
505+
body_block,
506+
inc_block,
507+
end_block,
508+
else_block,
509+
)
510+
511+
512+
def handle_loop_jump(compilation_context, builder, stmt):
513+
"""Handle `break` and `continue`: branch to the innermost loop's exit or
514+
next-iteration block."""
515+
keyword = "break" if isinstance(stmt, ast.Break) else "continue"
516+
if not compilation_context.loop_stack:
517+
raise SyntaxError(f"'{keyword}' outside loop (line {stmt.lineno})")
518+
continue_block, break_block = compilation_context.loop_stack[-1]
519+
builder.branch(break_block if keyword == "break" else continue_block)
520+
521+
356522
def handle_return(
357523
func, builder, stmt, local_sym_tab, ret_type, compilation_context=None
358524
):
@@ -423,6 +589,12 @@ def process_stmt(
423589
logger.debug(f"global declaration of {', '.join(stmt.names)} already bound")
424590
elif isinstance(stmt, ast.If):
425591
handle_if(func, compilation_context, builder, stmt, local_sym_tab, ret_type)
592+
elif isinstance(stmt, ast.While):
593+
handle_while(func, compilation_context, builder, stmt, local_sym_tab, ret_type)
594+
elif isinstance(stmt, ast.For):
595+
handle_for(func, compilation_context, builder, stmt, local_sym_tab, ret_type)
596+
elif isinstance(stmt, (ast.Break, ast.Continue)):
597+
handle_loop_jump(compilation_context, builder, stmt)
426598
elif isinstance(stmt, ast.Return):
427599
did_return = handle_return(
428600
func, builder, stmt, local_sym_tab, ret_type, compilation_context
@@ -438,10 +610,10 @@ def process_stmt(
438610

439611

440612
def process_block(func, compilation_context, builder, stmts, local_sym_tab, ret_type):
441-
"""Process a nested statement list, such as an if-branch, in the
442-
enclosing function's return type. Stops at the first statement that ends
443-
the block (a return), because whatever follows it in the same list can
444-
never run, and would otherwise be emitted after a terminator."""
613+
"""Process a nested statement list (an if-branch or loop body). Stops at
614+
the first statement that ends the block -- break, continue or return --
615+
because whatever follows it in the same list can never run, and would
616+
otherwise be emitted after a terminator."""
445617
for s in stmts:
446618
if builder.block.is_terminated:
447619
break

0 commit comments

Comments
 (0)