diff --git a/docs/index.md b/docs/index.md index 920edd23..05f182b6 100644 --- a/docs/index.md +++ b/docs/index.md @@ -69,6 +69,7 @@ user-guide/maps user-guide/structs user-guide/compilation user-guide/helpers +user-guide/integers ``` ```{toctree} diff --git a/docs/user-guide/index.md b/docs/user-guide/index.md index dce1b5eb..a7504f93 100644 --- a/docs/user-guide/index.md +++ b/docs/user-guide/index.md @@ -44,6 +44,9 @@ PythonBPF uses Python's `ctypes` module for type definitions: * `c_void_p` - Void pointers * `str(N)` - Fixed-length strings (e.g., `str(16)` for 16-byte string) +Integers follow C's rules for width, sign, conversion and arithmetic; see +{doc}`integers` for the details and the places where this differs from Python. + ## Example Structure A typical PythonBPF program follows this structure: diff --git a/docs/user-guide/integers.md b/docs/user-guide/integers.md new file mode 100644 index 00000000..3fbe6cac --- /dev/null +++ b/docs/user-guide/integers.md @@ -0,0 +1,153 @@ +# Integer Semantics + +PythonBPF programs are Python syntax, but the integers in them behave as C integers: the +program runs in the kernel as BPF bytecode, where every value is a fixed-width machine +word. This page describes the rules the compiler applies. They are C's rules, applied to +the `ctypes` types you declare, so a program's arithmetic matches what the equivalent C +program compiled with clang would compute. + +```{note} +This is one of the few places where PythonBPF deliberately differs from Python. Python +integers have arbitrary precision and no unsigned types; BPF has neither. The +[divergences from Python](#divergences-from-python) are listed at the end of this page. +``` + +## Types + +An integer's type is the `ctypes` type it was declared with, and the type carries both +a width and a sign: + +| Signed | Unsigned | Width | +|---|---|---| +| `c_int8` | `c_uint8` | 8 | +| `c_int16` | `c_uint16` | 16 | +| `c_int32` | `c_uint32` | 32 | +| `c_int64` | `c_uint64` | 64 | + +Every declaration site uses these types: local variables initialised with a constructor +call, `@bpfglobal` variables, `@struct` fields, map keys and values, and fields read from +`vmlinux` structures. Helper functions return the type of the kernel's signature, so +`pid()` and `ktime()` are unsigned while `probe_read`-style helpers return a signed +`long`. + +```python +count = c_uint32(0) # a 32-bit unsigned local +delta = c_int64(-1) # a 64-bit signed local +``` + +A local assigned without a constructor takes its type from the expression: + +```python +now = ktime() # c_uint64, the helper's return type +total = count + 1 # the type of the addition (see below), held in a 64-bit slot +``` + +Undeclared locals are always 64 bits wide; the inferred type only decides their sign. +Declare the local with a constructor when a narrower width matters. + +### Literals + +A literal has the type a C compiler gives it: `int` (32-bit signed) if the value fits, +`long long` (64-bit signed) otherwise. This matters for mixed arithmetic: in +`count / -2` with `count` a `c_uint32`, the literal `-2` is a 32-bit `int`, so the +division happens in `c_uint32` exactly as it would in C. + +## Assignment and conversion + +Assigning a value to a variable of a different integer type converts it, and the +variable's declared type is what the stored value *is* afterwards: + +* **Widening preserves the value.** The conversion looks at the *source*'s sign: an + unsigned source is zero-extended, a signed source is sign-extended. So a `c_uint32` + holding `0xFFFFFFFF` stored into a `c_int64` gives `4294967295`, and a `c_int32` + holding `-1` stored into a `c_uint64` gives `0xFFFFFFFFFFFFFFFF`. This is what C and + `ctypes` both do. +* **Narrowing truncates.** Only the low bits survive. +* **Same width reinterprets.** A `c_uint32` `0xFFFFFFFF` stored into a `c_int32` reads + as `-1`. + +The same rules apply to explicit conversions written as constructor calls +(`c_int64(count)`), to struct field stores and to `return`. + +## Arithmetic + +Each binary operation is typed on its own, from its two operands, following C's usual +arithmetic conversions: + +1. Operands narrower than 32 bits are promoted to `c_int32`. +2. If both operands have the same sign, the result has the wider width and that sign. +3. If the signs differ, the unsigned type wins when it is at least as wide as the signed + one; otherwise the signed type wins. + +The operation is then performed in that type, and its result has that type. The variable +receiving the result plays no part until the final store. Two consequences worth knowing: + +* **Intermediate results wrap at their own width.** `c_uint32(0x80000000) * c_uint32(2)` + is a `c_uint32` multiplication, so it wraps to `0` before being stored, even if the + destination is a `c_uint64`. Widen an operand first if you want a 64-bit product. +* **Mixed signs go unsigned.** `c_uint32(10) / c_int32(-2)` is an unsigned division by + `0xFFFFFFFE`, giving `0`, not `-5`. + +The sign of the operation's type selects the instruction for the operations where it +matters: + +| Operator | Signed type | Unsigned type | +|---|---|---| +| `/`, `//` | truncating signed division | unsigned division | +| `%` | remainder with the dividend's sign | unsigned remainder | +| `>>` | arithmetic shift (sign bit shifts in) | logical shift (zeros shift in) | +| `<`, `<=`, `>`, `>=` | signed comparison | unsigned comparison | + +`+`, `-`, `*`, `<<`, `&`, `|`, `^`, `==` and `!=` produce the same bits for either sign. + +Unary minus on an unsigned value follows C too: `-x` is `2^N - x` in the value's type. + +## Comparisons + +A comparison converts both operands with the same usual arithmetic conversions and then +compares in the resulting type. `c_uint64(10) > c_int64(-1)` is therefore an unsigned +comparison in which `-1` is the largest possible value, and the result is false. The +result of a comparison is `1` or `0`, as in C. + +## A verifier gotcha: packet pointer fields + +A few context fields are declared as 32-bit integers but are pointers as far as the +kernel verifier is concerned: `data`, `data_end` and `data_meta` on `xdp_md`, and `data` +and `data_end` on `__sk_buff`. Because they are `c_uint32`, arithmetic on them directly +is a 32-bit operation, exactly as in C, and the verifier rejects 32-bit arithmetic on a +pointer: + +``` +R0 32-bit pointer arithmetic prohibited +``` + +C programs cast these fields through `(void *)(long)` before using them for the same +reason. Until PythonBPF does this for you, copy the field into a local first, which is a +64-bit slot, or cast it with `c_void_p`: + +```python +data = ctx.data # 64-bit local +end = ctx.data_end +if data + 34 < end: # 64-bit pointer arithmetic, accepted + ... +``` + +```{note} +This is a known gap. The plan is to give these fields pointer rank automatically so that +no cast or copy is needed; this section will go away when that lands. +``` + +## Divergences from Python + +Because the semantics are C's, some Python behaviour does not carry over: + +* `/` is integer division; there is no floating-point result. +* `//` and `/` are the same operation, and both truncate toward zero: `-7 // 2` is `-3`, + where Python gives `-4`. +* `%` takes the sign of the dividend: `-7 % 2` is `-1`, where Python gives `1`. +* Integers have a fixed width and wrap on overflow; there is no arbitrary precision. +* Unsigned types exist, and mixing them with signed values follows C's conversions + rather than Python's mathematical integers. + +The test programs under `tests/passing_tests/signedness/` show each rule with its +expected value, and `tests/c-form/signedness.bpf.c` is the equivalent C program. diff --git a/pythonbpf/allocation_pass.py b/pythonbpf/allocation_pass.py index b2098e0f..aaf25391 100644 --- a/pythonbpf/allocation_pass.py +++ b/pythonbpf/allocation_pass.py @@ -6,7 +6,8 @@ from pythonbpf.helper import HelperHandlerRegistry from pythonbpf.vmlinux_parser.dependency_node import Field from .expr import VmlinuxHandlerRegistry -from pythonbpf.type_deducer import ctypes_to_ir +from pythonbpf.type_deducer import ctypes_to_ir, IntTy, signedness +from pythonbpf.expr.type_inference import infer_int_type from pythonbpf.maps import BPFMapType logger = logging.getLogger(__name__) @@ -74,7 +75,9 @@ def handle_assign_allocation(compilation_context, builder, stmt, local_sym_tab): elif isinstance(rval, ast.Constant): _allocate_for_constant(builder, var_name, rval, local_sym_tab) elif isinstance(rval, ast.BinOp): - _allocate_for_binop(builder, var_name, local_sym_tab) + _allocate_for_binop( + builder, var_name, rval, local_sym_tab, compilation_context + ) elif isinstance(rval, ast.Name): # Variable-to-variable assignment (b = a) _allocate_for_name( @@ -116,7 +119,11 @@ def _allocate_for_call(builder, var_name, rval, local_sym_tab, compilation_conte # Helper functions elif HelperHandlerRegistry.has_handler(call_type): - ir_type = ir.IntType(64) # Assume i64 return type + # Undeclared locals are 64-bit; the sign comes from the helper. + ret = HelperHandlerRegistry.get_return_type(call_type) + ir_type = IntTy( + 64, signedness(ret) if isinstance(ret, ir.IntType) else True + ) var = builder.alloca(ir_type, name=var_name) var.align = 8 local_sym_tab[var_name] = LocalSymbol(var, ir_type) @@ -256,7 +263,7 @@ def _allocate_for_constant(builder, var_name, rval, local_sym_tab): """Allocate memory for variable assigned from a constant.""" if isinstance(rval.value, bool): - ir_type = ir.IntType(1) + ir_type = IntTy(1, False) # a bool widens to 0 or 1, never sign-extends var = builder.alloca(ir_type, name=var_name) var.align = 1 local_sym_tab[var_name] = LocalSymbol(var, ir_type) @@ -282,9 +289,17 @@ def _allocate_for_constant(builder, var_name, rval, local_sym_tab): ) -def _allocate_for_binop(builder, var_name, local_sym_tab): - """Allocate memory for variable assigned from a binary operation.""" - ir_type = ir.IntType(64) # Assume i64 result +def _allocate_for_binop(builder, var_name, rval, local_sym_tab, compilation_context): + """Allocate memory for variable assigned from a binary operation. + + Undeclared locals are 64-bit; the sign is that of the expression's C type, + inferred statically. Falls back to signed when the expression involves + something the inference does not know. + """ + inferred = infer_int_type(rval, local_sym_tab, compilation_context) + if inferred is None: + logger.debug(f"Could not infer a type for {var_name}, assuming signed i64") + ir_type = IntTy(64, signedness(inferred) if inferred is not None else True) var = builder.alloca(ir_type, name=var_name) var.align = 8 local_sym_tab[var_name] = LocalSymbol(var, ir_type) diff --git a/pythonbpf/assign_pass.py b/pythonbpf/assign_pass.py index 5d931f70..f243d8ae 100644 --- a/pythonbpf/assign_pass.py +++ b/pythonbpf/assign_pass.py @@ -3,7 +3,7 @@ from inspect import isclass from llvmlite import ir -from pythonbpf.expr import eval_expr +from pythonbpf.expr import eval_expr, convert from pythonbpf.helper import emit_probe_read_kernel_str_call from pythonbpf.type_deducer import ctypes_to_ir from pythonbpf.vmlinux_parser.dependency_node import Field @@ -57,10 +57,7 @@ def handle_struct_field_assignment( # Same implicit widening/truncation as assignment to a local: expressions # evaluate in i64, but a field may be narrower. if isinstance(val_type, ir.IntType) and isinstance(field_type, ir.IntType): - if val_type.width < field_type.width: - val = builder.sext(val, field_type) - elif val_type.width > field_type.width: - val = builder.trunc(val, field_type) + val = convert(builder, val, val_type, field_type) # Regular assignment builder.store(val, field_ptr) @@ -154,7 +151,12 @@ def handle_variable_assignment( f"Evaluated value for {var_name}: {val} of type {val_type}, expected {var_type}" ) - if val_type != var_type: + if isinstance(val_type, ir.IntType) and isinstance(var_type, ir.IntType): + # The descriptor may be narrower than the constant carrying the value + # (a literal is a 64-bit constant typed as C int), so never decide + # from descriptor equality: convert is a no-op when widths match. + val = convert(builder, val, val_type, var_type) + elif val_type != var_type: # Handle vmlinux struct pointers - they're represented as Python classes but are i64 pointers if isclass(val_type) and (val_type.__module__ == "vmlinux"): logger.info("Handling vmlinux struct pointer assignment") @@ -219,14 +221,6 @@ def handle_variable_assignment( f"Failed to assign ctype struct field to {var_name}: {val_type} != {var_type}" ) return False - elif isinstance(val_type, ir.IntType) and isinstance(var_type, ir.IntType): - # Allow implicit int widening - if val_type.width < var_type.width: - val = builder.sext(val, var_type) - logger.info(f"Implicitly widened int for variable {var_name}") - elif val_type.width > var_type.width: - val = builder.trunc(val, var_type) - logger.info(f"Implicitly truncated int for variable {var_name}") elif isinstance(val_type, ir.IntType) and isinstance(var_type, ir.PointerType): # NOTE: This is assignment to a PTR_TO_MAP_VALUE_OR_NULL logger.info( diff --git a/pythonbpf/expr/__init__.py b/pythonbpf/expr/__init__.py index 5002113e..040c0359 100644 --- a/pythonbpf/expr/__init__.py +++ b/pythonbpf/expr/__init__.py @@ -1,5 +1,12 @@ -from .expr_pass import eval_expr, handle_expr, get_operand_value -from .type_normalization import convert_to_bool, get_base_type_and_depth +from .expr_pass import eval_expr, handle_expr, get_typed_operand +from .type_normalization import ( + convert_to_bool, + get_base_type_and_depth, + convert, + canonicalise, + to_promoted, +) +from .operators import usual_arithmetic_conversions from .ir_ops import deref_to_depth, access_struct_field from .operators import apply_binop from .call_registry import CallHandlerRegistry @@ -9,11 +16,15 @@ "eval_expr", "handle_expr", "convert_to_bool", + "convert", + "canonicalise", + "to_promoted", + "get_typed_operand", + "usual_arithmetic_conversions", "get_base_type_and_depth", "deref_to_depth", "apply_binop", "access_struct_field", - "get_operand_value", "CallHandlerRegistry", "VmlinuxHandlerRegistry", ] diff --git a/pythonbpf/expr/expr_pass.py b/pythonbpf/expr/expr_pass.py index 8d905746..c454e704 100644 --- a/pythonbpf/expr/expr_pass.py +++ b/pythonbpf/expr/expr_pass.py @@ -4,11 +4,21 @@ import logging from typing import Dict -from pythonbpf.type_deducer import ctypes_to_ir, is_ctypes +from pythonbpf.type_deducer import ( + ctypes_to_ir, + field_int_type, + is_ctypes, + IntTy, + int_literal_type, + signedness, +) from .call_registry import CallHandlerRegistry from .ir_ops import deref_to_depth, access_struct_field -from .operators import apply_binop, UNARY_OPS, BOOL_OPS +from .operators import apply_binop, usual_arithmetic_conversions, UNARY_OPS, BOOL_OPS from .type_normalization import ( + convert, + to_promoted, + canonicalise, convert_to_bool, handle_comparator, get_base_type_and_depth, @@ -48,10 +58,19 @@ def _handle_name_expr( raise SyntaxError(f"Undefined variable {expr.id}") +def _int_literal(v: int): + """An integer literal: a 64-bit constant with C's literal rank as its + descriptor, `int` if the value fits and `long long` otherwise.""" + return ir.Constant(ir.IntType(64), v), int_literal_type(v) + + def _handle_constant_expr(compilation_context, builder, expr: ast.Constant): """Handle ast.Constant expressions.""" if isinstance(expr.value, int) or isinstance(expr.value, bool): - return ir.Constant(ir.IntType(64), int(expr.value)), ir.IntType(64) + # C gives a literal the type int if it fits, otherwise long long. That + # rank is what makes `u32 / -2` an unsigned 32-bit division as in C. + v = int(expr.value) + return _int_literal(v) elif isinstance(expr.value, str): str_name = f".str.{id(expr)}" str_bytes = expr.value.encode("utf-8") + b"\x00" @@ -175,73 +194,111 @@ def _handle_deref_call(expr: ast.Call, local_sym_tab: Dict, builder: ir.IRBuilde # ============================================================================ -def get_operand_value(func, compilation_context, operand, builder, local_sym_tab): - """Extract the value from an operand, handling variables and constants.""" - logger.info(f"Getting operand value for: {ast.dump(operand)}") +def _descriptor(val, ty): + """IntTy descriptor for an evaluated integer value: width from the physical + value unless the descriptor is itself an integer type, sign from the + descriptor (an IntTy, a vmlinux Field, or plain -> signed). None when the + value is not an integer at all.""" + if isinstance(ty, ir.IntType): + return IntTy(ty.width, signedness(ty)) + field = field_int_type(ty) + if field is not None: + # A vmlinux field: load_ctx_field already widened the value, but C + # ranks it by its declared width (a c_uint32 field is unsigned int). + # + # TODO(gotcha): some u32 context fields are packet pointers to the + # verifier, not numbers: xdp_md.data / data_end / data_meta and + # __sk_buff.data / data_end. Ranking them as u32 is what C does, and + # the verifier then rejects any arithmetic on them ("32-bit pointer + # arithmetic prohibited"), so C code casts them through + # (void *)(long) first. The plan is to spare users that: give these + # fields 64-bit pointer rank here, so `ctx.data + 34 < ctx.data_end` + # lowers to 64-bit pointer arithmetic without a cast. Until then, + # copy the field into a local (a 64-bit slot) or cast it via + # c_void_p before using it. + return field + if val is not None and isinstance(val.type, ir.IntType): + return IntTy(val.type.width, signedness(ty)) + return None + + +def get_typed_operand(func, compilation_context, operand, builder, local_sym_tab): + """Evaluate an operand to (value, IntTy). Pointers (map-lookup results) are + dereferenced to the scalar they point at.""" + logger.info(f"Getting typed operand for: {ast.dump(operand)}") if isinstance(operand, ast.Name): if operand.id in local_sym_tab: - var = local_sym_tab[operand.id].var - var_type = var.type - base_type, depth = get_base_type_and_depth(var_type) - logger.info(f"var is {var}, base_type is {base_type}, depth is {depth}") + sym = local_sym_tab[operand.id] + var = sym.var + base_type, depth = get_base_type_and_depth(var.type) if depth == 1: val = builder.load(var) - return val - else: - val = deref_to_depth(func, builder, var, depth) - return val + return val, _descriptor(val, sym.ir_type) + val = deref_to_depth(func, builder, var, depth) + # A map-lookup local: the slot points at the value, and the map's + # declared value ctype is the symbol's metadata. That, not the + # physical pointee, carries the sign (a c_uint64 value is unsigned). + declared = ( + ctypes_to_ir(sym.metadata) + if isinstance(sym.metadata, str) and is_ctypes(sym.metadata) + else None + ) + return val, _descriptor( + val, declared if declared is not None else base_type + ) elif operand.id in compilation_context.bpf_globals: - # A @bpfglobal: plain load off the global symbol. - return builder.load(compilation_context.bpf_globals[operand.id].var) + sym = compilation_context.bpf_globals[operand.id] + val = builder.load(sym.var) + return val, _descriptor(val, sym.ir_type) else: - # Check if it's a vmlinux enum/constant vmlinux_result = VmlinuxHandlerRegistry.handle_name(operand.id) if vmlinux_result is not None: - val, _ = vmlinux_result - return val + return vmlinux_result # (i64 constant, its C rank) elif isinstance(operand, ast.Constant): - if isinstance(operand.value, int): - cst = ir.Constant(ir.IntType(64), int(operand.value)) - return cst + if isinstance(operand.value, (int, bool)): + v = int(operand.value) + lit_ty = IntTy(32, True) if -(1 << 31) <= v < (1 << 31) else IntTy(64, True) + return ir.Constant(ir.IntType(64), v), lit_ty raise TypeError(f"Unsupported constant type: {type(operand.value)}") elif isinstance(operand, ast.BinOp): - res = _handle_binary_op_impl( + return _handle_binary_op_impl( func, compilation_context, operand, builder, local_sym_tab ) - return res else: res = eval_expr(func, compilation_context, builder, operand, local_sym_tab) if res is None: raise ValueError(f"Failed to evaluate call expression: {operand}") - val, _ = res + val, ty = res logger.info(f"Evaluated expr to {val} of type {val.type}") base_type, depth = get_base_type_and_depth(val.type) if depth > 0: val = deref_to_depth(func, builder, val, depth) - return val + return val, _descriptor(val, ty) raise TypeError(f"Unsupported operand type: {type(operand)}") def _handle_binary_op_impl(func, compilation_context, rval, builder, local_sym_tab): + """A binary operation, typed per node the way C types it: the operation is + performed in the type given by the usual arithmetic conversions of its two + operands, each operand converted to that type first, and the result + narrowed to it -- so u32 * u32 wraps at 32 bits even though the arithmetic + itself runs in an i64 register. Returns (value, IntTy).""" op = rval.op - left = get_operand_value( + left, left_ty = get_typed_operand( func, compilation_context, rval.left, builder, local_sym_tab ) - right = get_operand_value( + right, right_ty = get_typed_operand( func, compilation_context, rval.right, builder, local_sym_tab ) - logger.info(f"left is {left}, right is {right}, op is {op}") - - # NOTE: Before doing the operation, if the operands are integers - # we always extend them to i64. The assignment to LHS will take - # care of truncation if needed. - if isinstance(left.type, ir.IntType) and left.type.width < 64: - left = builder.sext(left, ir.IntType(64)) - if isinstance(right.type, ir.IntType) and right.type.width < 64: - right = builder.sext(right, ir.IntType(64)) - - # Map AST operation nodes to LLVM IR builder methods - return apply_binop(builder, op, left, right) + result_ty = usual_arithmetic_conversions(left_ty, right_ty) + logger.info( + f"binop {type(op).__name__}: {left_ty.describe()} x {right_ty.describe()} " + f"-> {result_ty.describe()}" + ) + left = to_promoted(builder, left, left_ty, result_ty) + right = to_promoted(builder, right, right_ty, result_ty) + result = apply_binop(builder, op, left, right, signedness(result_ty)) + return canonicalise(builder, result, result_ty), result_ty def _handle_binary_op( @@ -252,15 +309,14 @@ def _handle_binary_op( var_name, local_sym_tab, ): - result = _handle_binary_op_impl( + result, result_ty = _handle_binary_op_impl( func, compilation_context, rval, builder, local_sym_tab ) if var_name and var_name in local_sym_tab: - logger.info( - f"Storing result {result} into variable {local_sym_tab[var_name].var}" - ) - builder.store(result, local_sym_tab[var_name].var) - return result, result.type + slot = local_sym_tab[var_name] + logger.info(f"Storing result {result} into variable {slot.var}") + builder.store(convert(builder, result, result_ty, slot.ir_type), slot.var) + return result, result_ty # ============================================================================ @@ -311,32 +367,17 @@ def _handle_ctypes_call( else: actual_ir_type = val_type - if actual_ir_type != expected_type: - # NOTE: We are only considering casting to and from int types for now - if isinstance(actual_ir_type, ir.IntType) and isinstance( - expected_type, ir.IntType - ): - if actual_ir_type.width < expected_type.width: - value = builder.sext(value, expected_type) - logger.info( - f"Sign-extended from i{actual_ir_type.width} to i{ - expected_type.width - }" - ) - elif actual_ir_type.width > expected_type.width: - value = builder.trunc(value, expected_type) - logger.info( - f"Truncated from i{actual_ir_type.width} to i{expected_type.width}" - ) - else: - # Same width, just use as-is (e.g., both i64) - pass - else: - raise ValueError( - f"Type mismatch: expected {expected_type}, got { - actual_ir_type - } (original type: {val_type})" - ) + if isinstance(actual_ir_type, ir.IntType) and isinstance(expected_type, ir.IntType): + # A cast is truncate-or-extend per the source's sign; the result then + # takes the cast's type. Decide from the value's physical width (as + # convert does), never from descriptor equality: a literal is a 64-bit + # constant whose descriptor may already read as C int. + value = convert(builder, value, actual_ir_type, expected_type) + elif actual_ir_type != expected_type: + raise ValueError( + f"Type mismatch: expected {expected_type}, got {actual_ir_type} " + f"(original type: {val_type})" + ) return value, expected_type @@ -366,8 +407,19 @@ def _handle_compare(func, compilation_context, builder, cond, local_sym_tab): logger.error("Failed to evaluate comparison operands") return None - lhs, _ = lhs - rhs, _ = rhs + lhs, lhs_ty = lhs + rhs, rhs_ty = rhs + lhs_desc, rhs_desc = _descriptor(lhs, lhs_ty), _descriptor(rhs, rhs_ty) + if lhs_desc is not None and rhs_desc is not None: + # Both integers: compare in the promoted type, which also picks the + # signed or unsigned predicate (u64 > s64 is an unsigned compare in C). + cmp_ty = usual_arithmetic_conversions(lhs_desc, rhs_desc) + lhs = to_promoted(builder, lhs, lhs_desc, cmp_ty) + rhs = to_promoted(builder, rhs, rhs_desc, cmp_ty) + return handle_comparator( + func, builder, cond.ops[0], lhs, rhs, signed=signedness(cmp_ty) + ) + # Pointers and struct values: the depth-normalising path return handle_comparator(func, builder, cond.ops[0], lhs, rhs) @@ -383,7 +435,7 @@ def _handle_unary_op( logger.error("Only 'not' and '-' unary operators are supported") return None - operand = get_operand_value( + operand, operand_ty = get_typed_operand( func, compilation_context, expr.operand, builder, local_sym_tab ) if operand is None: @@ -393,12 +445,18 @@ def _handle_unary_op( if isinstance(expr.op, ast.Not): true_const = ir.Constant(ir.IntType(1), 1) result = builder.xor(convert_to_bool(builder, operand), true_const) - return result, ir.IntType(1) + return result, IntTy(1, False) elif isinstance(expr.op, ast.USub): - # Multiply by -1 - neg_one = ir.Constant(ir.IntType(64), -1) - result = builder.mul(operand, neg_one) - return result, ir.IntType(64) + if isinstance(operand, ir.Constant) and isinstance(operand.constant, int): + # -2 parses as USub(Constant 2); fold it so it is a literal like + # any other, with a literal's C rank. + return _int_literal(-operand.constant) + # Negation happens in the operand's promoted type; for an unsigned + # operand that is C's 2^N - x, which the narrowing produces. + result_ty = usual_arithmetic_conversions(operand_ty, operand_ty) + operand = to_promoted(builder, operand, operand_ty, result_ty) + result = builder.mul(operand, ir.Constant(ir.IntType(64), -1)) + return canonicalise(builder, result, result_ty), result_ty return None @@ -481,7 +539,7 @@ def _handle_and_op(func, builder, expr, local_sym_tab, compilation_context): phi.add_incoming(val, block) logger.debug(f"Generated 'and' with {len(incoming_values)} incoming values") - return phi, ir.IntType(1) + return phi, IntTy(1, False) def _handle_or_op(func, builder, expr, local_sym_tab, compilation_context): @@ -536,7 +594,7 @@ def _handle_or_op(func, builder, expr, local_sym_tab, compilation_context): phi.add_incoming(val, block) logger.debug(f"Generated 'or' with {len(incoming_values)} incoming values") - return phi, ir.IntType(1) + return phi, IntTy(1, False) def _handle_boolean_op( diff --git a/pythonbpf/expr/operators.py b/pythonbpf/expr/operators.py index fb4cadd1..7a4fbea4 100644 --- a/pythonbpf/expr/operators.py +++ b/pythonbpf/expr/operators.py @@ -9,20 +9,26 @@ import ast -# ast.BinOp.op class -> llvmlite IRBuilder method name. -# Shared by binary-op evaluation and augmented assignment. +from pythonbpf.type_deducer import IntTy, signedness + +# ast.BinOp.op class -> (signed IRBuilder method, unsigned IRBuilder method). +# Shared by binary-op evaluation and augmented assignment. The ring operations +# are sign-blind (two's complement gives identical low bits); division, +# remainder and right shift are not, and the operation's type decides. `/` and +# `//` are the same C truncating division -- Python's floor semantics for `//` +# and `%` on negatives are a documented divergence. BINOP_METHODS = { - ast.Add: "add", - ast.Sub: "sub", - ast.Mult: "mul", - ast.Div: "sdiv", - ast.Mod: "srem", - ast.LShift: "shl", - ast.RShift: "lshr", - ast.BitOr: "or_", - ast.BitXor: "xor", - ast.BitAnd: "and_", - ast.FloorDiv: "udiv", + ast.Add: ("add", "add"), + ast.Sub: ("sub", "sub"), + ast.Mult: ("mul", "mul"), + ast.Div: ("sdiv", "udiv"), + ast.FloorDiv: ("sdiv", "udiv"), + ast.Mod: ("srem", "urem"), + ast.LShift: ("shl", "shl"), + ast.RShift: ("ashr", "lshr"), + ast.BitOr: ("or_", "or_"), + ast.BitXor: ("xor", "xor"), + ast.BitAnd: ("and_", "and_"), } # ast.Compare op class -> icmp predicate string. @@ -42,14 +48,39 @@ BOOL_OPS = (ast.And, ast.Or) -def apply_binop(builder, op, left, right): - """Emit the LLVM instruction for a Python binary operator.""" - method = BINOP_METHODS.get(type(op)) - if method is None: +def apply_binop(builder, op, left, right, signed=True): + """Emit the LLVM instruction for a Python binary operator, in the signed or + unsigned form the operation's type calls for.""" + methods = BINOP_METHODS.get(type(op)) + if methods is None: raise SyntaxError(f"Unsupported binary operation: {type(op).__name__}") - return getattr(builder, method)(left, right) + return getattr(builder, methods[0] if signed else methods[1])(left, right) def comparison_predicate(op): """icmp predicate for a Python comparison operator, or None if unsupported.""" return COMPARISON_OPS.get(type(op)) + + +def usual_arithmetic_conversions(left, right) -> IntTy: + """The type a C binary operation on `left` and `right` is performed in. + + Integer promotion first: anything narrower than int becomes a signed 32-bit + int (int can represent every value of the narrower type, signed or not). + Then, if the signs agree, the wider type wins; if they differ, the unsigned + operand wins at equal or greater width, otherwise the signed one -- because + it can then represent every value of the unsigned one. + """ + + def promote(ty): + if ty.width < 32: + return IntTy(32, True) + return IntTy(ty.width, signedness(ty)) + + left, right = promote(left), promote(right) + if left.signed == right.signed: + return IntTy(max(left.width, right.width), left.signed) + unsigned, signed = (left, right) if not left.signed else (right, left) + if unsigned.width >= signed.width: + return IntTy(unsigned.width, False) + return IntTy(signed.width, True) diff --git a/pythonbpf/expr/type_inference.py b/pythonbpf/expr/type_inference.py new file mode 100644 index 00000000..9a0c56b8 --- /dev/null +++ b/pythonbpf/expr/type_inference.py @@ -0,0 +1,106 @@ +"""Static integer typing of an expression, for the allocation pass. + +The allocation pass runs before code generation and must size and type a slot +for `x = ` without evaluating . Undeclared locals are always 64-bit +(declare with a ctypes constructor for a narrower type); what this decides is +their sign, by walking the expression with the same usual-arithmetic-conversion +rule the code generator applies. +""" + +import ast +import ctypes + +from llvmlite import ir + +from pythonbpf.type_deducer import ( + IntTy, + ctypes_to_ir, + is_ctypes, + is_signed_ctype, + signedness, +) +from .operators import usual_arithmetic_conversions +from .vmlinux_registry import VmlinuxHandlerRegistry + + +def _as_intty(ty): + if isinstance(ty, ir.IntType): + return IntTy(ty.width, signedness(ty)) + return None + + +def infer_int_type(expr, local_sym_tab, compilation_context): + """Best static integer type of `expr`, or None when it cannot be determined.""" + if isinstance(expr, ast.Constant) and isinstance(expr.value, (int, bool)): + v = int(expr.value) + return IntTy(32, True) if -(1 << 31) <= v < (1 << 31) else IntTy(64, True) + + if isinstance(expr, ast.Name): + if expr.id in local_sym_tab: + return _as_intty(local_sym_tab[expr.id].ir_type) + if expr.id in compilation_context.bpf_globals: + return _as_intty(compilation_context.bpf_globals[expr.id].ir_type) + enum = VmlinuxHandlerRegistry.handle_name(expr.id) + if enum is not None: + return _as_intty(enum[1]) + return None + + if isinstance(expr, ast.BinOp): + left = infer_int_type(expr.left, local_sym_tab, compilation_context) + right = infer_int_type(expr.right, local_sym_tab, compilation_context) + if left is None or right is None: + return None + return usual_arithmetic_conversions(left, right) + + if isinstance(expr, ast.UnaryOp): + inner = infer_int_type(expr.operand, local_sym_tab, compilation_context) + return None if inner is None else usual_arithmetic_conversions(inner, inner) + + if isinstance(expr, ast.Call) and isinstance(expr.func, ast.Name): + from pythonbpf.helper import HelperHandlerRegistry # avoid an import cycle + + name = expr.func.id + if is_ctypes(name): + return _as_intty(ctypes_to_ir(name)) + if HelperHandlerRegistry.has_handler(name): + return _as_intty(HelperHandlerRegistry.get_return_type(name)) + return None + + if ( + isinstance(expr, ast.Call) + and isinstance(expr.func, ast.Attribute) + and expr.func.attr == "lookup" + ): + # m.lookup(...) used as a value has the type the map declares for its + # values, whatever the key argument looks like. The declaration is a + # ctypes name for a scalar map and a struct name otherwise, which is + # not an integer, so None. + map_name = getattr(expr.func.value, "id", None) + sym = compilation_context.map_sym_tab.get(map_name) + value_ctype = (sym.params or {}).get("value") if sym else None + if isinstance(value_ctype, str) and is_ctypes(value_ctype): + return _as_intty(ctypes_to_ir(value_ctype)) + return None + + if isinstance(expr, ast.Attribute) and isinstance(expr.value, ast.Name): + base = local_sym_tab.get(expr.value.id) + if base is None: + return None + meta = base.metadata + if meta in compilation_context.structs_sym_tab: + return _as_intty( + compilation_context.structs_sym_tab[meta].field_type(expr.attr) + ) + if getattr(meta, "__module__", None) == "vmlinux": + try: + _, field = VmlinuxHandlerRegistry.get_field_type( + meta.__name__, expr.attr + ) + cname = field.type.__name__ + if is_ctypes(cname): + return IntTy(ctypes.sizeof(field.type) * 8, is_signed_ctype(cname)) + except Exception: + return None + return None + + return None diff --git a/pythonbpf/expr/type_normalization.py b/pythonbpf/expr/type_normalization.py index edea4228..e89815f8 100644 --- a/pythonbpf/expr/type_normalization.py +++ b/pythonbpf/expr/type_normalization.py @@ -1,6 +1,7 @@ import logging from llvmlite import ir from .ir_ops import deref_to_depth +from pythonbpf.type_deducer import IntTy, signedness from .operators import COMPARISON_OPS logger = logging.getLogger(__name__) @@ -22,9 +23,9 @@ def _normalize_types(func, builder, lhs, rhs): logger.info(f"Normalizing types: {lhs.type} vs {rhs.type}") if isinstance(lhs.type, ir.IntType) and isinstance(rhs.type, ir.IntType): if lhs.type.width < rhs.type.width: - lhs = builder.sext(lhs, rhs.type) + lhs = convert(builder, lhs, lhs.type, rhs.type) else: - rhs = builder.sext(rhs, lhs.type) + rhs = convert(builder, rhs, rhs.type, lhs.type) return lhs, rhs elif not isinstance(lhs.type, ir.PointerType) and not isinstance( rhs.type, ir.PointerType @@ -42,6 +43,72 @@ def _normalize_types(func, builder, lhs, rhs): return _normalize_types(func, builder, lhs, rhs) +def convert(builder, val, from_ty, to_ty): + """Convert an integer value between types the way C does. + + Widening is driven by the *source* sign (zext for unsigned, sext for + signed) so the mathematical value is preserved; narrowing truncates; equal + width is a reinterpretation and emits nothing. `from_ty` and `to_ty` are + descriptors (see type_deducer.IntTy); the physical width comes from the + value itself, which may already be wider than its descriptor says. + """ + if not (isinstance(to_ty, ir.IntType) and isinstance(val.type, ir.IntType)): + # Every caller either checks both sides are integers or sits on a path + # that only carries integers, so reaching here is a type error that + # would otherwise surface as an llc rejection with no Python line. + raise TypeError( + f"integer conversion requested for a {val.type} value to {to_ty}" + ) + if val.type.width > to_ty.width: + if to_ty.width == 1: + # C's rule for bool: nonzero is true. Truncation would keep the + # low bit and turn 2 into false. + return builder.icmp_unsigned("!=", val, ir.Constant(val.type, 0)) + return builder.trunc(val, to_ty) + if val.type.width < to_ty.width: + ext = builder.zext if not signedness(from_ty) else builder.sext + return ext(val, to_ty) + return val + + +def _fold_int_constant(val, ty, width): + """A literal re-expressed at the working width holding type ty's value: + wrap to ty's width, take the representative ty's sign implies.""" + v = val.constant % (1 << ty.width) + if signedness(ty) and v >= 1 << (ty.width - 1): + v -= 1 << ty.width + return ir.Constant(ir.IntType(width), v) + + +def to_promoted(builder, val, from_ty, to_ty, width=64): + """Bring an operand to the promoted type of its operation, C-style. + + First convert it to to_ty per its *own* sign (that is C's conversion of an + operand to the common type), then widen to the working width per to_ty's + sign so the i64 register holds exactly a to_ty value. Literals are folded. + """ + if isinstance(val, ir.Constant) and isinstance(val.constant, int): + return _fold_int_constant(val, to_ty, width) + val = convert(builder, val, from_ty, ir.IntType(to_ty.width)) + return canonicalise(builder, val, to_ty, width) + + +def canonicalise(builder, val, ty, width=64): + """Bring `val` to the working width holding exactly the value of type `ty`: + truncate to ty's width if the register is wider (so the operation wraps at + ty's width, as C does), then extend per ty's sign.""" + if not isinstance(val.type, ir.IntType): + return val + if isinstance(val, ir.Constant) and isinstance(val.constant, int): + return _fold_int_constant(val, ty, width) + if val.type.width > ty.width: + val = builder.trunc(val, ir.IntType(ty.width)) + if val.type.width < width: + ext = builder.zext if not signedness(ty) else builder.sext + val = ext(val, ir.IntType(width)) + return val + + def convert_to_bool(builder, val): """Convert a value to boolean.""" if val.type == ir.IntType(1): @@ -53,8 +120,8 @@ def convert_to_bool(builder, val): return builder.icmp_signed("!=", val, zero) -def handle_comparator(func, builder, op, lhs, rhs): - """Handle comparison operations.""" +def handle_comparator(func, builder, op, lhs, rhs, signed=True): + """Handle comparison operations, signed or unsigned per the compared type.""" if lhs.type != rhs.type: lhs, rhs = _normalize_types(func, builder, lhs, rhs) @@ -67,6 +134,7 @@ def handle_comparator(func, builder, op, lhs, rhs): return None predicate = COMPARISON_OPS[type(op)] - result = builder.icmp_signed(predicate, lhs, rhs) + icmp = builder.icmp_signed if signed else builder.icmp_unsigned + result = icmp(predicate, lhs, rhs) logger.debug(f"Comparison result: {result}") - return result, ir.IntType(1) + return result, IntTy(1, False) diff --git a/pythonbpf/functions/functions_pass.py b/pythonbpf/functions/functions_pass.py index 8aee77a6..382ef38e 100644 --- a/pythonbpf/functions/functions_pass.py +++ b/pythonbpf/functions/functions_pass.py @@ -6,13 +6,17 @@ from pythonbpf.helper import ( HelperHandlerRegistry, ) -from pythonbpf.type_deducer import ctypes_to_ir, is_ctypes +from pythonbpf.type_deducer import ctypes_to_ir, is_ctypes, signedness from pythonbpf.expr import ( eval_expr, handle_expr, convert_to_bool, - get_operand_value, + get_typed_operand, apply_binop, + convert, + to_promoted, + canonicalise, + usual_arithmetic_conversions, VmlinuxHandlerRegistry, ) from pythonbpf.assign_pass import ( @@ -199,7 +203,7 @@ def handle_aug_assign(func, compilation_context, builder, stmt, local_sym_tab): compiler walks the tree the user wrote, and nodes invented mid-codegen are invisible to the passes that already ran and carry no source locations. Semantic agreement with `x = x op v` comes from sharing the value-level - helpers instead — the RHS goes through get_operand_value like any other + helpers instead — the RHS goes through get_typed_operand like any other read, and the operator table is apply_binop, the same one binary-op evaluation uses. """ @@ -266,23 +270,24 @@ def handle_aug_assign(func, compilation_context, builder, stmt, local_sym_tab): # Python evaluates the target's current value before the right-hand side. current = builder.load(slot) - rhs = get_operand_value( + rhs, rhs_ty = get_typed_operand( func, compilation_context, stmt.value, builder, local_sym_tab ) if rhs is None: raise SyntaxError( f"Failed to evaluate augmented-assignment value: {ast.dump(stmt.value)}" ) - # Same width discipline as binary-op evaluation: compute in i64, narrow - # back to the slot's width on the way out. - if current.type.width < 64: - current = builder.sext(current, ir.IntType(64)) - if isinstance(rhs.type, ir.IntType) and rhs.type.width < 64: - rhs = builder.sext(rhs, ir.IntType(64)) - result = apply_binop(builder, stmt.op, current, rhs) - if result.type.width > slot_type.width: - result = builder.trunc(result, slot_type) - builder.store(result, slot) + # x op= v is typed exactly like x = x op v: operate in the promoted type, + # then convert the result to the target's type on the way back in. + result_ty = usual_arithmetic_conversions(slot_type, rhs_ty) + current = to_promoted(builder, current, slot_type, result_ty) + rhs = to_promoted(builder, rhs, rhs_ty, result_ty) + result = canonicalise( + builder, + apply_binop(builder, stmt.op, current, rhs, signedness(result_ty)), + result_ty, + ) + builder.store(convert(builder, result, result_ty, slot_type), slot) def handle_cond(func, compilation_context, builder, cond, local_sym_tab): @@ -330,7 +335,9 @@ def handle_if(func, compilation_context, builder, stmt, local_sym_tab): builder.position_at_end(merge_block) -def handle_return(builder, stmt, local_sym_tab, ret_type, compilation_context=None): +def handle_return( + func, builder, stmt, local_sym_tab, ret_type, compilation_context=None +): logger.info(f"Handling return statement: {ast.dump(stmt)}") if stmt.value is None: return handle_none_return(builder) @@ -355,15 +362,15 @@ def handle_return(builder, stmt, local_sym_tab, ret_type, compilation_context=No "CompilationContext required for return statement evaluation" ) - val = eval_expr( - func=None, - compilation_context=compilation_context, - builder=builder, - expr=stmt.value, - local_sym_tab=local_sym_tab, + # A pointer to a value is dereferenced to it (null-checked), the way + # every other consumer of a value does; get_typed_operand is that path. + val = get_typed_operand( + func, compilation_context, stmt.value, builder, local_sym_tab ) logger.info(f"Evaluated return expression to {val}") - builder.ret(val[0]) + # The declared return type is the LHS of an implicit assignment: + # widen per the value's sign, truncate if narrower. + builder.ret(convert(builder, val[0], val[1], ret_type)) return True @@ -398,7 +405,7 @@ def process_stmt( handle_if(func, compilation_context, builder, stmt, local_sym_tab) elif isinstance(stmt, ast.Return): did_return = handle_return( - builder, stmt, local_sym_tab, ret_type, compilation_context + func, builder, stmt, local_sym_tab, ret_type, compilation_context ) else: # Dropping a statement makes the program mean something other than what diff --git a/pythonbpf/globals_pass.py b/pythonbpf/globals_pass.py index aabe8fde..c5dc0ba5 100644 --- a/pythonbpf/globals_pass.py +++ b/pythonbpf/globals_pass.py @@ -3,7 +3,7 @@ from logging import Logger import logging -from .type_deducer import ctypes_to_ir +from .type_deducer import ctypes_to_ir, is_signed_ctype from .symbols import BpfGlobalSymbol from .debuginfo import DebugInfoGenerator from .expr import VmlinuxHandlerRegistry @@ -11,18 +11,6 @@ logger: Logger = logging.getLogger(__name__) -_SIGNED_CTYPES = { - "c_int8", - "c_int16", - "c_int32", - "c_int64", - "c_int", - "c_short", - "c_long", - "c_longlong", - "c_byte", -} - _C_NAME_BY_WIDTH = {8: "char", 16: "short", 32: "int", 64: "long long"} @@ -106,7 +94,7 @@ def _emit_global_debug_info(compilation_context, gvar, name, ctype_name): """ generator = DebugInfoGenerator(compilation_context.module) width = gvar.value_type.width - signed = ctype_name in _SIGNED_CTYPES + signed = is_signed_ctype(ctype_name) base = _C_NAME_BY_WIDTH[width] if width == 8: encoding = dc.DW_ATE_signed_char if signed else dc.DW_ATE_unsigned_char diff --git a/pythonbpf/helper/bpf_helper_handler.py b/pythonbpf/helper/bpf_helper_handler.py index 3b3e61a9..92bb1307 100644 --- a/pythonbpf/helper/bpf_helper_handler.py +++ b/pythonbpf/helper/bpf_helper_handler.py @@ -3,6 +3,7 @@ from enum import Enum from .helper_registry import HelperHandlerRegistry +from pythonbpf.type_deducer import IntTy from .helper_utils import ( get_or_create_ptr_from_arg, get_flags_val, @@ -45,7 +46,7 @@ class BPFHelperID(Enum): @HelperHandlerRegistry.register( "ktime", param_types=[], - return_type=ir.IntType(64), + return_type=IntTy(64, False), ) def bpf_ktime_get_ns_emitter( call, @@ -70,7 +71,7 @@ def bpf_ktime_get_ns_emitter( @HelperHandlerRegistry.register( "get_current_cgroup_id", param_types=[], - return_type=ir.IntType(64), + return_type=IntTy(64, False), ) def bpf_get_current_cgroup_id( call, @@ -310,7 +311,7 @@ def bpf_map_delete_elem_emitter( @HelperHandlerRegistry.register( "comm", param_types=[ir.PointerType(ir.IntType(8))], - return_type=ir.IntType(64), + return_type=IntTy(64, True), ) def bpf_get_current_comm_emitter( call, @@ -369,7 +370,7 @@ def bpf_get_current_comm_emitter( @HelperHandlerRegistry.register( "pid", param_types=[], - return_type=ir.IntType(64), + return_type=IntTy(64, False), ) def bpf_get_current_pid_tgid_emitter( call, @@ -496,7 +497,7 @@ def bpf_ringbuf_output_emitter( @HelperHandlerRegistry.register( "output", param_types=[ir.PointerType(ir.IntType(8))], - return_type=ir.IntType(64), + return_type=IntTy(64, True), ) def handle_output_helper( call, @@ -566,7 +567,7 @@ def emit_probe_read_kernel_str_call(builder, dst_ptr, dst_size, src_ptr): ir.PointerType(ir.IntType(8)), ir.PointerType(ir.IntType(8)), ], - return_type=ir.IntType(64), + return_type=IntTy(64, True), ) def bpf_probe_read_kernel_str_emitter( call, @@ -633,7 +634,7 @@ def emit_probe_read_kernel_call(builder, dst_ptr, dst_size, src_ptr): ir.PointerType(ir.IntType(8)), ir.PointerType(ir.IntType(8)), ], - return_type=ir.IntType(64), + return_type=IntTy(64, True), ) def bpf_probe_read_kernel_emitter( call, @@ -670,7 +671,7 @@ def bpf_probe_read_kernel_emitter( @HelperHandlerRegistry.register( "random", param_types=[], - return_type=ir.IntType(32), + return_type=IntTy(32, False), ) def bpf_get_prandom_u32_emitter( call, @@ -698,7 +699,7 @@ def bpf_get_prandom_u32_emitter( ir.IntType(32), ir.PointerType(ir.IntType(8)), ], - return_type=ir.IntType(64), + return_type=IntTy(64, True), ) def bpf_probe_read_emitter( call, @@ -763,7 +764,7 @@ def bpf_probe_read_emitter( @HelperHandlerRegistry.register( "smp_processor_id", param_types=[], - return_type=ir.IntType(32), + return_type=IntTy(32, False), ) def bpf_get_smp_processor_id_emitter( call, @@ -788,7 +789,7 @@ def bpf_get_smp_processor_id_emitter( @HelperHandlerRegistry.register( "uid", param_types=[], - return_type=ir.IntType(64), + return_type=IntTy(64, False), ) def bpf_get_current_uid_gid_emitter( call, @@ -822,7 +823,7 @@ def bpf_get_current_uid_gid_emitter( ir.IntType(32), ir.IntType(64), ], - return_type=ir.IntType(64), + return_type=IntTy(64, True), ) def bpf_skb_store_bytes_emitter( call, @@ -1010,7 +1011,7 @@ def bpf_ringbuf_submit_emitter( @HelperHandlerRegistry.register( "get_stack", param_types=[ir.PointerType(ir.IntType(8)), ir.IntType(64)], - return_type=ir.IntType(64), + return_type=IntTy(64, True), ) def bpf_get_stack_emitter( call, @@ -1075,7 +1076,7 @@ def invoke_helper(method_name, map_ptr=None): raise NotImplementedError( f"Helper function '{method_name}' is not implemented." ) - return handler( + result = handler( call, map_ptr, compilation_context, @@ -1083,6 +1084,19 @@ def invoke_helper(method_name, map_ptr=None): func, local_sym_tab, ) + # Emitters return (value, plain LLVM type); the registry entry is the + # descriptor that knows the sign. Substitute it once, here, rather + # than in every emitter. + declared = HelperHandlerRegistry.get_return_type(method_name) + if ( + isinstance(result, tuple) + and len(result) == 2 + and isinstance(declared, ir.IntType) + and isinstance(result[1], ir.IntType) + and declared.width == result[1].width + ): + return result[0], declared + return result map_sym_tab = compilation_context.map_sym_tab diff --git a/pythonbpf/helper/helper_utils.py b/pythonbpf/helper/helper_utils.py index 57251175..28b60b34 100644 --- a/pythonbpf/helper/helper_utils.py +++ b/pythonbpf/helper/helper_utils.py @@ -3,6 +3,7 @@ from llvmlite import ir from pythonbpf.expr import ( + convert, eval_expr, access_struct_field, ) @@ -134,12 +135,8 @@ def get_or_create_ptr_from_arg( local_sym_tab, expected_type ) logger.info(f"Using temp variable '{temp_name}' for expression result") - if ( - isinstance(val.type, ir.IntType) - and expected_type - and val.type.width > expected_type.width - ): - val = builder.trunc(val, expected_type) + if expected_type is not None and isinstance(expected_type, ir.IntType): + val = convert(builder, val, val.type, expected_type) builder.store(val, ptr) # NOTE: For char arrays, also return size diff --git a/pythonbpf/helper/printk_formatter.py b/pythonbpf/helper/printk_formatter.py index f9d35370..47262c99 100644 --- a/pythonbpf/helper/printk_formatter.py +++ b/pythonbpf/helper/printk_formatter.py @@ -2,7 +2,7 @@ import logging from llvmlite import ir -from pythonbpf.expr import eval_expr, get_base_type_and_depth, deref_to_depth +from pythonbpf.expr import eval_expr, get_base_type_and_depth, deref_to_depth, convert from pythonbpf.expr.vmlinux_registry import VmlinuxHandlerRegistry from pythonbpf.helper.helper_utils import get_char_array_ptr_and_size @@ -267,8 +267,6 @@ def _handle_pointer_arg(val, func, builder): return ir.Constant(ir.IntType(64), 0) -def _handle_int_arg(val, builder): - """Convert integer type for bpf_printk (sign-extend to i64).""" - if val.type.width < 64: - return builder.sext(val, ir.IntType(64)) - return val +def _handle_int_arg(val, builder, ty=None): + """Widen an integer for bpf_printk to i64, per the value's sign.""" + return convert(builder, val, ty if ty is not None else val.type, ir.IntType(64)) diff --git a/pythonbpf/type_deducer.py b/pythonbpf/type_deducer.py index 2e4c77f4..f6f412c4 100644 --- a/pythonbpf/type_deducer.py +++ b/pythonbpf/type_deducer.py @@ -1,30 +1,132 @@ from llvmlite import ir + +class IntTy(ir.IntType): + """An LLVM integer type that also remembers its signedness. + + LLVM integer types are sign-agnostic by design: `i32` is just 32 bits, and + the sign lives in the operations (sdiv/udiv, sext/zext, icmp s*/u*). The + frontend therefore has to carry it. IntTy is a plain ir.IntType for every + purpose LLVM cares about -- it renders as `i32`, compares and hashes equal + to ir.IntType(32), and passes every isinstance check -- with one extra + attribute the compiler reads when choosing between signed and unsigned + forms of an operation. + + Invariant: the sign is read only from a *descriptor* -- a Symbol.ir_type + or the type half of an eval_expr result -- never from `value.type`. Values + produced by the IRBuilder (loads, arithmetic, extensions) come back with a + plain ir.IntType, so a sign on a value's own type is lost at the first + operation. Descriptors are constructed by the compiler; that is where the + sign lives. + """ + + def __new__(cls, bits: int, signed: bool = True): + # ir.IntType.__new__ memoises one instance per width in a cache shared + # with subclasses. Going through it would (a) merge the signed and + # unsigned flavours of a width into one object and (b) plant an IntTy + # in the cache so that ir.IntType(32) itself started returning one. + # Construct directly instead; equality and hashing are inherited and + # depend only on the width, so an IntTy still compares equal to i32. + self = object.__new__(cls) + self.width = bits + return self + + def __init__(self, bits: int, signed: bool = True): + self.signed = signed + + def __getnewargs__(self): + return self.width, self.signed + + def describe(self) -> str: + return f"{'i' if self.signed else 'u'}{self.width}" + + +def int_literal_type(value: int) -> IntTy: + """C's type for an integer constant: `int` if the value fits, else + `long long`. Literals and enum constants alike; an enum constant is an + `int` whatever the enum's underlying type is (that type belongs to + variables of the enum type, such as struct fields).""" + return IntTy(32, True) if -(1 << 31) <= value < (1 << 31) else IntTy(64, True) + + +def field_int_type(ty) -> "IntTy | None": + """The declared integer type behind a vmlinux Field descriptor (its ctypes + class), or None if the descriptor is not a Field with an integer ctype. + The loaded value may already be wider; C ranks it by the declared width.""" + ctype = getattr(getattr(ty, "type", None), "__name__", None) + if ctype in _INT_CTYPE_WIDTHS: + return IntTy(_INT_CTYPE_WIDTHS[ctype], is_signed_ctype(ctype)) + return None + + +def signedness(ty) -> bool: + """Sign of a descriptor. An IntTy carries it directly; a vmlinux Field + carries a ctypes class in .type, whose name decides; a 1-bit integer is a + bool and never negative, whatever it is wrapped in; a plain ir.IntType, a + site not yet taught to carry a sign, reads as signed, the compiler's + historical behaviour.""" + if isinstance(ty, ir.IntType) and ty.width == 1: + return False + if hasattr(ty, "signed"): + return ty.signed + ctype = getattr(getattr(ty, "type", None), "__name__", None) + if ctype in _INT_CTYPE_WIDTHS: + return is_signed_ctype(ctype) + return True + + +_SIGNED_CTYPES = { + "c_int8", + "c_int16", + "c_int32", + "c_int64", + "c_int", + "c_short", + "c_long", + "c_longlong", + "c_byte", +} + +_INT_CTYPE_WIDTHS = { + "c_int8": 8, + "c_uint8": 8, + "c_byte": 8, + "c_ubyte": 8, + "c_int16": 16, + "c_uint16": 16, + "c_short": 16, + "c_ushort": 16, + "c_int32": 32, + "c_uint32": 32, + "c_int": 32, + "c_uint": 32, + "c_int64": 64, + "c_uint64": 64, + "c_long": 64, + "c_ulong": 64, + "c_longlong": 64, + # A pointer-sized integer; treated as unsigned like uintptr_t. + "c_void_p": 64, +} + + +def is_signed_ctype(ctype: str) -> bool: + return ctype in _SIGNED_CTYPES + + # TODO: THIS IS NOT SUPPOSED TO MATCH STRINGS :skull: mapping = { - "c_int8": ir.IntType(8), - "c_uint8": ir.IntType(8), - "c_int16": ir.IntType(16), - "c_uint16": ir.IntType(16), - "c_int32": ir.IntType(32), - "c_uint32": ir.IntType(32), - "c_int64": ir.IntType(64), - "c_uint64": ir.IntType(64), - "c_float": ir.FloatType(), - "c_double": ir.DoubleType(), - "c_void_p": ir.IntType(64), - "c_long": ir.IntType(64), - "c_ulong": ir.IntType(64), - "c_longlong": ir.IntType(64), - "c_uint": ir.IntType(32), - "c_int": ir.IntType(32), - "c_ushort": ir.IntType(16), - "c_short": ir.IntType(16), - "c_ubyte": ir.IntType(8), - "c_byte": ir.IntType(8), - # Not so sure about this one - "str": ir.PointerType(ir.IntType(8)), + name: IntTy(width, is_signed_ctype(name)) + for name, width in _INT_CTYPE_WIDTHS.items() } +mapping.update( + { + "c_float": ir.FloatType(), + "c_double": ir.DoubleType(), + # Not so sure about this one + "str": ir.PointerType(ir.IntType(8)), + } +) def ctypes_to_ir(ctype: str): diff --git a/pythonbpf/vmlinux_parser/vmlinux_exports_handler.py b/pythonbpf/vmlinux_parser/vmlinux_exports_handler.py index 97a84c12..7305d111 100644 --- a/pythonbpf/vmlinux_parser/vmlinux_exports_handler.py +++ b/pythonbpf/vmlinux_parser/vmlinux_exports_handler.py @@ -4,6 +4,7 @@ from llvmlite import ir from pythonbpf.symbols import LocalSymbol +from pythonbpf.type_deducer import int_literal_type, is_signed_ctype from pythonbpf.vmlinux_parser.assignment_info import AssignmentType logger = logging.getLogger(__name__) @@ -73,7 +74,7 @@ def handle_vmlinux_enum(self, name): if self.is_vmlinux_enum(name): value = self.vmlinux_symtab[name].value logger.info(f"Resolving vmlinux enum {name} = {value}") - return ir.Constant(ir.IntType(64), value), ir.IntType(64) + return ir.Constant(ir.IntType(64), value), int_literal_type(value) return None def get_vmlinux_enum_value(self, name): @@ -367,8 +368,12 @@ def load_ctx_field(builder, ctx_arg, offset_global, field_data, struct_name=None # Widen sub-register-width context fields to i64 if needs_zext: - value = builder.zext(value, ir.IntType(64)) - logger.info(f"Zero-extended i{int_width} context field value to i64") + if is_signed_ctype(getattr(field_data.type, "__name__", "")): + value = builder.sext(value, ir.IntType(64)) + logger.info(f"Sign-extended i{int_width} context field value to i64") + else: + value = builder.zext(value, ir.IntType(64)) + logger.info(f"Zero-extended i{int_width} context field value to i64") return value diff --git a/tests/c-form/signedness.bpf.c b/tests/c-form/signedness.bpf.c new file mode 100644 index 00000000..86a5edf7 --- /dev/null +++ b/tests/c-form/signedness.bpf.c @@ -0,0 +1,47 @@ +/* Reference for integer signedness. Each case reads its operands from .bss + * globals (so clang cannot constant-fold) and writes the result to a .bss + * global (so the runtime proof is a `bpftool map dump`). The IR clang emits + * for this file is the specification: zext vs sext on widening, udiv vs sdiv, + * icmp ugt vs sgt, lshr vs ashr, and the trunc/zext pair that makes + * u32 * u32 wrap at 32 bits. */ +#define SEC(name) __attribute__((section(name), used)) +typedef unsigned int __u32; +typedef int __s32; +typedef unsigned long long __u64; +typedef long long __s64; + +/* inputs */ +__u32 u32_max = 0xFFFFFFFFu; +__s32 s32_neg = -1; +__u32 u32_ten = 10; +__s32 s32_neg2 = -2; +__u64 u64_ten = 10; +__s64 s64_neg = -1; +__u32 u32_half = 0x80000000u; +__u32 u32_two = 2; + +/* results */ +__s64 widen_unsigned; /* s64 = u32 -> 4294967295 zext */ +__u64 widen_signed; /* u64 = s32 -> 0xffffffffffffffff sext */ +__s32 reinterpret; /* s32 = u32 -> -1 no-op */ +__u32 mixed_div; /* u32 / s32 -> 0 udiv i32 */ +__u64 unsigned_cmp; /* u64 > s64 -> 0 icmp ugt */ +__u64 narrow_wrap; /* u32 * u32 -> 0 (wraps at 32) trunc */ +__u32 shr_unsigned; /* u32 >> 4 -> 0x0FFFFFFF lshr */ +__s32 shr_signed; /* s32 >> 4 -> -1 ashr */ + +SEC("tracepoint/raw_syscalls/sys_enter") +int prog(void *ctx) +{ + widen_unsigned = u32_max; + widen_signed = s32_neg; + reinterpret = u32_max; + mixed_div = u32_ten / s32_neg2; + unsigned_cmp = u64_ten > s64_neg; + narrow_wrap = u32_half * u32_two; + shr_unsigned = u32_max >> 4; + shr_signed = s32_neg >> 4; + return 0; +} + +char _license[] SEC("license") = "GPL"; diff --git a/tests/failing_tests/return_struct.py b/tests/failing_tests/return_struct.py new file mode 100644 index 00000000..4d29a0b9 --- /dev/null +++ b/tests/failing_tests/return_struct.py @@ -0,0 +1,28 @@ +# A struct value has no integer to reach by dereferencing, so returning one +# from an integer function is a type error, raised by convert() with both +# types named rather than left for llc to reject with an IR line number. +from pythonbpf import bpf, struct, section, bpfglobal, compile +from ctypes import c_void_p, c_int64, c_uint64 + + +@bpf +@struct +class task_info: + pid: c_uint64 + + +@bpf +@section("tracepoint/raw_syscalls/sys_enter") +def prog(ctx: c_void_p) -> c_int64: + t = task_info() + t.pid = 1 + return t + + +@bpf +@bpfglobal +def LICENSE() -> str: + return "GPL" + + +compile() diff --git a/tests/passing_tests/return/map_value.py b/tests/passing_tests/return/map_value.py new file mode 100644 index 00000000..b866e442 --- /dev/null +++ b/tests/passing_tests/return/map_value.py @@ -0,0 +1,31 @@ +# `return` is a consumer of a value like any other: a pointer to one is +# dereferenced (null-checked) to reach it, so returning a map lookup result +# works the way `p + 0` already did. C shape: `if (p) return *p; return 0;` +from pythonbpf import bpf, map, section, bpfglobal, compile +from pythonbpf.maps import HashMap +from ctypes import c_void_p, c_int64, c_uint32, c_uint64 + + +@bpf +@map +def m() -> HashMap: + return HashMap(key=c_uint32, value=c_uint64, max_entries=4) + + +@bpf +@section("tracepoint/raw_syscalls/sys_enter") +def prog(ctx: c_void_p) -> c_int64: + k = c_uint32(1) + p = m.lookup(k) + if p: + return p + return c_int64(0) + + +@bpf +@bpfglobal +def LICENSE() -> str: + return "GPL" + + +compile() diff --git a/tests/passing_tests/signedness/augassign_unsigned.py b/tests/passing_tests/signedness/augassign_unsigned.py new file mode 100644 index 00000000..1ba2bfa4 --- /dev/null +++ b/tests/passing_tests/signedness/augassign_unsigned.py @@ -0,0 +1,24 @@ +# Augmented assignment picks its operator from the promoted type like a +# binary operation does: on a c_uint32, >>= is a logical shift and //= and +# %= are unsigned division and remainder. +from pythonbpf import bpf, section, bpfglobal, compile +from ctypes import c_void_p, c_int64, c_uint32 + + +@bpf +@section("tracepoint/raw_syscalls/sys_enter") +def prog(ctx: c_void_p) -> c_int64: + u = c_uint32(0xF0000000) + u >>= 4 # lshr: 0x0F000000, not sign-filled + u //= 3 # udiv + u %= 7 # urem + return c_int64(u) + + +@bpf +@bpfglobal +def LICENSE() -> str: + return "GPL" + + +compile() diff --git a/tests/passing_tests/signedness/bool_int.py b/tests/passing_tests/signedness/bool_int.py new file mode 100644 index 00000000..3103ec6b --- /dev/null +++ b/tests/passing_tests/signedness/bool_int.py @@ -0,0 +1,24 @@ +# A bool is a 1-bit integer to LLVM, and C's rules for it are not the integer +# ones: it widens to 0 or 1 (never sign-extends to -1), and an integer narrows +# to it by comparing with zero (2 is true), not by keeping the low bit. +from pythonbpf import bpf, section, bpfglobal, compile +from ctypes import c_void_p, c_int64 + + +@bpf +@section("tracepoint/raw_syscalls/sys_enter") +def prog(ctx: c_void_p) -> c_int64: + t = True + n = True + n = 2 # bool = 2 is true in C; truncation would make it false + s = t + 1 # int + bool promotes the bool to int 1: 2, not 0 + return c_int64(t + n + s) # 1 + 1 + 2 = 4, not -1 + 0 + 0 + + +@bpf +@bpfglobal +def LICENSE() -> str: + return "GPL" + + +compile() diff --git a/tests/passing_tests/signedness/helper_results.py b/tests/passing_tests/signedness/helper_results.py new file mode 100644 index 00000000..b0033d4f --- /dev/null +++ b/tests/passing_tests/signedness/helper_results.py @@ -0,0 +1,23 @@ +# A helper's value carries the sign its registry entry declares: ktime() +# and pid() are unsigned, so a right shift is logical and a division is +# unsigned, the same as for a c_uint64 local. +from pythonbpf import bpf, section, bpfglobal, compile +from pythonbpf.helper import ktime, pid +from ctypes import c_void_p, c_int64 + + +@bpf +@section("tracepoint/raw_syscalls/sys_enter") +def prog(ctx: c_void_p) -> c_int64: + t = ktime() >> 1 # lshr + q = pid() // 3 # udiv + return c_int64(t + q) + + +@bpf +@bpfglobal +def LICENSE() -> str: + return "GPL" + + +compile() diff --git a/tests/passing_tests/signedness/literal_rank.py b/tests/passing_tests/signedness/literal_rank.py new file mode 100644 index 00000000..b4bfe557 --- /dev/null +++ b/tests/passing_tests/signedness/literal_rank.py @@ -0,0 +1,45 @@ +# Mirrors tests/c-form/signedness.bpf.c; the IR clang emits for that file is +# the specification (see tests/test_signedness_ir.py for the ops asserted). +from pythonbpf import bpf, section, bpfglobal, compile +from ctypes import c_void_p, c_int64, c_uint32 + + +@bpf +@bpfglobal +def u32_ten() -> c_uint32: + return c_uint32(10) + + +@bpf +@bpfglobal +def by_literal() -> c_uint32: + return c_uint32(0) + + +@bpf +@bpfglobal +def wide() -> c_int64: + return c_int64(0) + + +@bpf +@section("tracepoint/raw_syscalls/sys_enter") +def prog(ctx: c_void_p) -> c_int64: + global by_literal, wide + # A literal has C's `int` type, so u32 / -2 promotes to u32: an unsigned + # division by 0xFFFFFFFE (udiv), exactly as `u32_ten / -2` in C + by_literal = u32_ten / -2 + # A literal that does not fit in int is a `long long`; the sum is 64-bit + wide = 5000000000 + 1 + # i32-ranked literals are still 64-bit slots for an undeclared local + small = 7 * 6 + return c_int64(small) + + +@bpf +@bpfglobal +def LICENSE() -> str: + return "GPL" + + +compile() diff --git a/tests/passing_tests/signedness/map_value_sign.py b/tests/passing_tests/signedness/map_value_sign.py new file mode 100644 index 00000000..1c5c19d1 --- /dev/null +++ b/tests/passing_tests/signedness/map_value_sign.py @@ -0,0 +1,31 @@ +# A map-lookup local dereferenced to its value carries the map's declared +# value type: on a c_uint64 map, p >> 63 is a logical shift. +from pythonbpf import bpf, map, section, bpfglobal, compile +from pythonbpf.maps import HashMap +from ctypes import c_void_p, c_int64, c_uint32, c_uint64 + + +@bpf +@map +def m() -> HashMap: + return HashMap(key=c_uint32, value=c_uint64, max_entries=4) + + +@bpf +@section("tracepoint/raw_syscalls/sys_enter") +def prog(ctx: c_void_p) -> c_int64: + k = c_uint32(1) + p = m.lookup(k) + if p: + top = p >> 63 # lshr: 0 or 1, never -1 + return c_int64(top) + return c_int64(0) + + +@bpf +@bpfglobal +def LICENSE() -> str: + return "GPL" + + +compile() diff --git a/tests/passing_tests/signedness/mixed_division.py b/tests/passing_tests/signedness/mixed_division.py new file mode 100644 index 00000000..75bf8711 --- /dev/null +++ b/tests/passing_tests/signedness/mixed_division.py @@ -0,0 +1,52 @@ +# Mirrors tests/c-form/signedness.bpf.c; the IR clang emits for that file is +# the specification (see tests/test_signedness_ir.py for the ops asserted). +from pythonbpf import bpf, section, bpfglobal, compile +from ctypes import c_void_p, c_int64, c_int32, c_uint32 + + +@bpf +@bpfglobal +def u32_ten() -> c_uint32: + return c_uint32(10) + + +@bpf +@bpfglobal +def s32_neg2() -> c_int32: + return c_int32(-2) + + +@bpf +@bpfglobal +def mixed_div() -> c_uint32: + return c_uint32(0) + + +@bpf +@bpfglobal +def mixed_mod() -> c_uint32: + return c_uint32(0) + + +@bpf +@section("tracepoint/raw_syscalls/sys_enter") +def prog(ctx: c_void_p) -> c_int64: + global mixed_div, mixed_mod + # u32 / s32: the usual arithmetic conversions make both u32, so this is an + # unsigned division: 10 / 0xFFFFFFFE == 0, not -5 + mixed_div = u32_ten / s32_neg2 + mixed_mod = u32_ten % s32_neg2 + # both signed: an ordinary signed division + a = c_int32(-7) + b = c_int32(2) + q = a / b + return c_int64(q) + + +@bpf +@bpfglobal +def LICENSE() -> str: + return "GPL" + + +compile() diff --git a/tests/passing_tests/signedness/narrow_wrap.py b/tests/passing_tests/signedness/narrow_wrap.py new file mode 100644 index 00000000..51b5b844 --- /dev/null +++ b/tests/passing_tests/signedness/narrow_wrap.py @@ -0,0 +1,43 @@ +# Mirrors tests/c-form/signedness.bpf.c; the IR clang emits for that file is +# the specification (see tests/test_signedness_ir.py for the ops asserted). +from pythonbpf import bpf, section, bpfglobal, compile +from ctypes import c_void_p, c_int64, c_uint32, c_uint64 + + +@bpf +@bpfglobal +def u32_half() -> c_uint32: + return c_uint32(0x80000000) + + +@bpf +@bpfglobal +def u32_two() -> c_uint32: + return c_uint32(2) + + +@bpf +@bpfglobal +def narrow_wrap() -> c_uint64: + return c_uint64(0) + + +@bpf +@section("tracepoint/raw_syscalls/sys_enter") +def prog(ctx: c_void_p) -> c_int64: + global narrow_wrap + # u32 * u32 is a u32 multiplication: 0x80000000 * 2 wraps to 0 before the + # widening to u64 (the LHS never widens the operands) + narrow_wrap = u32_half * u32_two + # narrowing truncates: only the low 32 bits of the u64 survive + small = c_uint32(narrow_wrap) + return c_int64(small) + + +@bpf +@bpfglobal +def LICENSE() -> str: + return "GPL" + + +compile() diff --git a/tests/passing_tests/signedness/right_shift.py b/tests/passing_tests/signedness/right_shift.py new file mode 100644 index 00000000..a97bf745 --- /dev/null +++ b/tests/passing_tests/signedness/right_shift.py @@ -0,0 +1,48 @@ +# Mirrors tests/c-form/signedness.bpf.c; the IR clang emits for that file is +# the specification (see tests/test_signedness_ir.py for the ops asserted). +from pythonbpf import bpf, section, bpfglobal, compile +from ctypes import c_void_p, c_int64, c_int32, c_uint32 + + +@bpf +@bpfglobal +def u32_max() -> c_uint32: + return c_uint32(0xFFFFFFFF) + + +@bpf +@bpfglobal +def s32_neg() -> c_int32: + return c_int32(-1) + + +@bpf +@bpfglobal +def shr_unsigned() -> c_uint32: + return c_uint32(0) + + +@bpf +@bpfglobal +def shr_signed() -> c_int32: + return c_int32(0) + + +@bpf +@section("tracepoint/raw_syscalls/sys_enter") +def prog(ctx: c_void_p) -> c_int64: + global shr_unsigned, shr_signed + # u32 >> 4 shifts zeros in (lshr): 0x0FFFFFFF + shr_unsigned = u32_max >> 4 + # s32 >> 4 shifts the sign in (ashr): -1 stays -1 + shr_signed = s32_neg >> 4 + return c_int64(0) + + +@bpf +@bpfglobal +def LICENSE() -> str: + return "GPL" + + +compile() diff --git a/tests/passing_tests/signedness/unsigned_compare.py b/tests/passing_tests/signedness/unsigned_compare.py new file mode 100644 index 00000000..08a5e611 --- /dev/null +++ b/tests/passing_tests/signedness/unsigned_compare.py @@ -0,0 +1,52 @@ +# Mirrors tests/c-form/signedness.bpf.c; the IR clang emits for that file is +# the specification (see tests/test_signedness_ir.py for the ops asserted). +from pythonbpf import bpf, section, bpfglobal, compile +from ctypes import c_void_p, c_int64, c_uint64 + + +@bpf +@bpfglobal +def u64_ten() -> c_uint64: + return c_uint64(10) + + +@bpf +@bpfglobal +def s64_neg() -> c_int64: + return c_int64(-1) + + +@bpf +@bpfglobal +def unsigned_cmp() -> c_uint64: + return c_uint64(0) + + +@bpf +@bpfglobal +def signed_cmp() -> c_uint64: + return c_uint64(0) + + +@bpf +@section("tracepoint/raw_syscalls/sys_enter") +def prog(ctx: c_void_p) -> c_int64: + global unsigned_cmp, signed_cmp + # u64 > s64: the comparison happens in u64, so -1 is the largest value and + # 10 > -1 is false (icmp ugt) + if u64_ten > s64_neg: + unsigned_cmp = 1 + # s64 > s64 stays a signed comparison (icmp sgt) + a = c_int64(10) + if a > s64_neg: + signed_cmp = 1 + return c_int64(0) + + +@bpf +@bpfglobal +def LICENSE() -> str: + return "GPL" + + +compile() diff --git a/tests/passing_tests/signedness/widen_signed.py b/tests/passing_tests/signedness/widen_signed.py new file mode 100644 index 00000000..5fc7159d --- /dev/null +++ b/tests/passing_tests/signedness/widen_signed.py @@ -0,0 +1,36 @@ +# Mirrors tests/c-form/signedness.bpf.c; the IR clang emits for that file is +# the specification (see tests/test_signedness_ir.py for the ops asserted). +from pythonbpf import bpf, section, bpfglobal, compile +from ctypes import c_void_p, c_int64, c_int32, c_uint64 + + +@bpf +@bpfglobal +def s32_neg() -> c_int32: + return c_int32(-1) + + +@bpf +@bpfglobal +def widen_signed() -> c_uint64: + return c_uint64(0) + + +@bpf +@section("tracepoint/raw_syscalls/sys_enter") +def prog(ctx: c_void_p) -> c_int64: + global widen_signed + # u64 = s32: -1 sign-extends to 0xffffffffffffffff (sext), as in C and ctypes + widen_signed = s32_neg + x = c_int32(-5) + y = c_int64(x) + return y + + +@bpf +@bpfglobal +def LICENSE() -> str: + return "GPL" + + +compile() diff --git a/tests/passing_tests/signedness/widen_unsigned.py b/tests/passing_tests/signedness/widen_unsigned.py new file mode 100644 index 00000000..9ad1561b --- /dev/null +++ b/tests/passing_tests/signedness/widen_unsigned.py @@ -0,0 +1,44 @@ +# Mirrors tests/c-form/signedness.bpf.c; the IR clang emits for that file is +# the specification (see tests/test_signedness_ir.py for the ops asserted). +from pythonbpf import bpf, section, bpfglobal, compile +from ctypes import c_void_p, c_int64, c_int32, c_uint32 + + +@bpf +@bpfglobal +def u32_max() -> c_uint32: + return c_uint32(0xFFFFFFFF) + + +@bpf +@bpfglobal +def widen_unsigned() -> c_int64: + return c_int64(0) + + +@bpf +@bpfglobal +def reinterpret() -> c_int32: + return c_int32(0) + + +@bpf +@section("tracepoint/raw_syscalls/sys_enter") +def prog(ctx: c_void_p) -> c_int64: + global widen_unsigned, reinterpret + # s64 = u32: value-preserving, so 0xFFFFFFFF stays 4294967295 (zext, not sext) + widen_unsigned = u32_max + # s32 = u32: same width, the bits are reinterpreted (-1) + reinterpret = u32_max + # a local copied from a global takes the global's type: a c_uint32 slot + also = u32_max + return c_int64(also) + + +@bpf +@bpfglobal +def LICENSE() -> str: + return "GPL" + + +compile() diff --git a/tests/passing_tests/vmlinux/ctx_field_rank.py b/tests/passing_tests/vmlinux/ctx_field_rank.py new file mode 100644 index 00000000..1699858e --- /dev/null +++ b/tests/passing_tests/vmlinux/ctx_field_rank.py @@ -0,0 +1,29 @@ +# A sub-register context field is loaded widened to i64, but C ranks it by +# its declared width: xdp_md.ingress_ifindex is a u32, so ifindex - k with a +# c_int32 k is a u32 subtraction (the result is cut to 32 bits and +# zero-extended), not a 64-bit one. +# +# Not xdp_md.data: it is a u32 in C too, but the verifier tracks it as a +# packet pointer and rejects 32-bit arithmetic on it ("R0 32-bit pointer +# arithmetic prohibited"), for C programs as much as for this one. That is +# why C casts it through (void *)(long) before doing anything with it. +from ctypes import c_int32, c_int64 +from pythonbpf import bpf, section, bpfglobal, compile +from vmlinux import struct_xdp_md + + +@bpf +@section("xdp") +def prog(ctx: struct_xdp_md) -> c_int64: + k = c_int32(1) + d = ctx.ingress_ifindex - k + return c_int64(d) + + +@bpf +@bpfglobal +def LICENSE() -> str: + return "GPL" + + +compile() diff --git a/tests/passing_tests/vmlinux/enum_rank.py b/tests/passing_tests/vmlinux/enum_rank.py new file mode 100644 index 00000000..ffab2148 --- /dev/null +++ b/tests/passing_tests/vmlinux/enum_rank.py @@ -0,0 +1,24 @@ +# An enum constant has C's `int` type, whatever the enum's underlying type is +# (clang: `XDP_PASS - k` is `sub nsw i32` + sext; a *variable* of the enum +# type would be u32). So `XDP_PASS - k` with an unsigned k is a u32 operation, +# the same rank a literal that fits in int gets. +from ctypes import c_int64, c_uint32, c_void_p +from pythonbpf import bpf, section, bpfglobal, compile +from vmlinux import XDP_PASS + + +@bpf +@section("tracepoint/raw_syscalls/sys_enter") +def prog(ctx: c_void_p) -> c_int64: + k = c_uint32(3) + d = XDP_PASS - k + return c_int64(d) + + +@bpf +@bpfglobal +def LICENSE() -> str: + return "GPL" + + +compile() diff --git a/tests/test_config.toml b/tests/test_config.toml index 0697a6c4..6f9bec08 100644 --- a/tests/test_config.toml +++ b/tests/test_config.toml @@ -15,6 +15,8 @@ "failing_tests/conditionals/struct_ptr.py" = {reason = "Struct pointer used directly as boolean condition not supported", level = "ir"} +"failing_tests/return_struct.py" = {reason = "Returning a struct value from an integer function is a type error (nothing to dereference to an integer)", level = "ir"} + "failing_tests/license.py" = {reason = "Missing LICENSE global produces IR that llc rejects — should be caught earlier with a clear error message", level = "llc"} "failing_tests/undeclared_values.py" = {reason = "Undeclared variable used in f-string — should raise SyntaxError (correct behaviour, test documents it)", level = "ir"} diff --git a/tests/test_signedness_ir.py b/tests/test_signedness_ir.py new file mode 100644 index 00000000..2166ed88 --- /dev/null +++ b/tests/test_signedness_ir.py @@ -0,0 +1,105 @@ +""" +Integer signedness: the shape of the IR. + +The passing_tests/signedness cases mirror tests/c-form/signedness.bpf.c, and +the IR clang emits for that C file is the specification. Levels 1 and 2 only +prove these files compile; this test checks the operations themselves, which +is the whole point of the cases: zext vs sext on widening, udiv vs sdiv, +icmp ugt vs sgt, lshr vs ashr, and the trunc/zext pair that wraps u32 * u32. +""" + +import importlib.util +import re +from pathlib import Path + +import pytest + +from tests.framework.compiler import run_ir_generation + +PASSING_DIR = Path(__file__).parent / "passing_tests" +HAVE_VMLINUX = importlib.util.find_spec("vmlinux") is not None + +# path under passing_tests -> (patterns that must appear, patterns that must not) +CASES = { + "signedness/widen_unsigned.py": ( + [r"zext i32 .* to i64"], + [r"sext i32 .* to i64"], + ), + "signedness/widen_signed.py": ( + [r"sext i32 .* to i64"], + [r"zext i32 .* to i64"], + ), + "signedness/mixed_division.py": ( + [r"\budiv i64", r"\burem i64", r"\bsdiv i64"], + [r"\bsrem i64"], + ), + "signedness/unsigned_compare.py": ( + [r"icmp ugt i64", r"icmp sgt i64"], + [], + ), + "signedness/narrow_wrap.py": ( + # the product is computed at 64 bits, cut to 32, then zero-extended + [r"\bmul i64", r"trunc i64 .* to i32", r"zext i32 .* to i64"], + [r"sext i32 .* to i64"], + ), + "signedness/right_shift.py": ( + [r"\blshr i64 .*, 4", r"\bashr i64 .*, 4"], + [], + ), + "signedness/literal_rank.py": ( + [r"\budiv i64 .*, 4294967294"], + [r"\bsdiv i64"], + ), + "signedness/augassign_unsigned.py": ( + [r"\blshr i64", r"\budiv i64", r"\burem i64"], + [r"\bashr i64", r"\bsdiv i64", r"\bsrem i64"], + ), + "signedness/helper_results.py": ( + [r"\blshr i64", r"\budiv i64"], + [r"\bashr i64", r"\bsdiv i64"], + ), + "signedness/map_value_sign.py": ( + [r"\blshr i64 [^,]*, 63"], + [r"\bashr i64"], + ), + # u32 - int is a u32 operation: the result is cut to 32 bits and + # zero-extended; ranked as 64-bit there would be no trunc at all. + "vmlinux/ctx_field_rank.py": ( + [r"\bsub i64", r"trunc i64 .* to i32", r"zext i32 .* to i64"], + [], + ), + # A bool widens with zext and an integer narrows to it by != 0, never by + # trunc; sext of an i1 would return -1 for True. + "signedness/bool_int.py": ( + [r"zext i1 .* to i(32|64)", r"icmp ne i64 .*, 0"], + [r"sext i1 ", r"trunc i64 .* to i1"], + ), + # An enum constant is a C `int`, so `XDP_PASS - k` with k a c_uint32 is a + # u32 operation: the result is cut to 32 bits and zero-extended. Ranked + # as i64 it would be a signed 64-bit subtraction with no trunc at all. + "vmlinux/enum_rank.py": ( + [r"trunc i64 .* to i32", r"zext i32 .* to i64"], + [r"sext i32 .* to i64"], + ), + # `return p` on a map lookup dereferences through a null check and returns + # the i64, never the pointer. + "return/map_value.py": ( + [r"deref_0_not_null", r"ret i64 %"], + [r"ret i64\*"], + ), +} + + +@pytest.mark.parametrize("name", list(CASES)) +def test_signedness_ir_shape(name, tmp_path): + if name.startswith("vmlinux/") and not HAVE_VMLINUX: + pytest.skip("vmlinux.py not importable") + ll_path = tmp_path / Path(name).name.replace(".py", ".ll") + run_ir_generation(PASSING_DIR / name, ll_path) + ir_text = ll_path.read_text() + + expected, forbidden = CASES[name] + for pattern in expected: + assert re.search(pattern, ir_text), f"{name}: expected /{pattern}/ in the IR" + for pattern in forbidden: + assert not re.search(pattern, ir_text), f"{name}: /{pattern}/ must not appear"