2626from 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+
356522def 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
440612def 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