Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
22 changes: 20 additions & 2 deletions pythonbpf/allocation_pass.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,15 @@
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, is_ctypes, IntTy, signedness, byte_size
from pythonbpf.type_deducer import (
ctypes_to_ir,
is_ctypes,
IntTy,
signedness,
byte_size,
PktPtrTy,
)
from pythonbpf.expr.packet_pointer import packet_type
from pythonbpf.expr.type_inference import infer_int_type
from pythonbpf.maps import BPFMapType

Expand Down Expand Up @@ -301,7 +309,10 @@ def _allocate_for_binop(builder, var_name, rval, local_sym_tab, compilation_cont
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)
if isinstance(inferred, PktPtrTy):
ir_type = inferred # data + 14 is still a packet pointer
else:
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)
Expand Down Expand Up @@ -401,6 +412,13 @@ def _allocate_for_attribute(
# Same discriminator handle_vmlinux_struct_field uses: a context
# argument has no alloca of its own.
is_context_field = local_sym_tab[struct_var].var is None
pkt_ty = packet_type(rval, local_sym_tab)
if pkt_ty is not None:
# A packet-pointer field: a 64-bit slot that carries the kind.
var = _allocate_with_type(builder, var_name, pkt_ty)
local_sym_tab[var_name] = LocalSymbol(var, pkt_ty)
logger.info(f"Pre-allocated {var_name} as {pkt_ty.describe()}")
return
if not VmlinuxHandlerRegistry.has_field(vmlinux_struct_name, field_name):
logger.error(
f"Field '{field_name}' not found in vmlinux struct '{vmlinux_struct_name}'"
Expand Down
64 changes: 64 additions & 0 deletions pythonbpf/expr/expr_pass.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@
field_int_type,
is_ctypes,
IntTy,
PktPtrTy,
int_literal_type,
signedness,
)
Expand All @@ -24,6 +25,7 @@
get_base_type_and_depth,
)
from .vmlinux_registry import VmlinuxHandlerRegistry
from .packet_pointer import MAX_PACKET_OFF, packet_type
from ..vmlinux_parser.dependency_node import Field

logger: Logger = logging.getLogger(__name__)
Expand Down Expand Up @@ -122,6 +124,11 @@ def _handle_attribute_expr(
expr, local_sym_tab, None, builder
)
if vmlinux_result is not None:
pkt_ty = packet_type(expr, local_sym_tab)
if pkt_ty is not None:
# A packet-pointer field: the value is the loaded
# pointer, and the descriptor says so from here on.
return vmlinux_result[0], pkt_ty
return vmlinux_result
else:
raise RuntimeError("Vmlinux struct did not process successfully")
Expand Down Expand Up @@ -199,6 +206,8 @@ def _descriptor(val, ty):
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, PktPtrTy):
return ty # a packet pointer keeps its kind, never re-ranked
if isinstance(ty, ir.IntType):
return IntTy(ty.width, signedness(ty))
field = field_int_type(ty)
Expand Down Expand Up @@ -291,6 +300,10 @@ def _handle_binary_op_impl(func, compilation_context, rval, builder, local_sym_t
right, right_ty = get_typed_operand(
func, compilation_context, rval.right, builder, local_sym_tab
)
if isinstance(left_ty, PktPtrTy) or isinstance(right_ty, PktPtrTy):
# Before the usual arithmetic conversions, which would re-rank the
# pointer as an ordinary integer and drop its kind.
return _packet_pointer_binop(builder, rval, left, left_ty, right, right_ty)
Comment on lines +303 to +306
result_ty = usual_arithmetic_conversions(left_ty, right_ty)
logger.info(
f"binop {type(op).__name__}: {left_ty.describe()} x {right_ty.describe()} "
Expand All @@ -302,6 +315,50 @@ def _handle_binary_op_impl(func, compilation_context, rval, builder, local_sym_t
return canonicalise(builder, result, result_ty), result_ty


def _packet_pointer_binop(builder, rval, left, left_ty, right, right_ty):
"""A binary operation with a packet-pointer operand (see packet_pointer.py):
64-bit, never ranked as an ordinary integer, with the offset taken as an
unsigned 16-bit value so the verifier can track the packet range. The
result of pointer +/- offset is a pointer of the same kind; pointer -
pointer is a plain 64-bit length."""
op = rval.op
where = f"line {rval.lineno}: {ast.unparse(rval)}"
left_pkt = isinstance(left_ty, PktPtrTy)
right_pkt = isinstance(right_ty, PktPtrTy)
if left_pkt and right_pkt:
if isinstance(op, ast.Sub):
# data_end - data: a length, an ordinary 64-bit scalar.
return builder.sub(left, right), IntTy(64, False)
raise SyntaxError(
f"only subtraction is defined between packet pointers ({where})"
)
if not isinstance(op, (ast.Add, ast.Sub)):
raise SyntaxError(
f"only + and - with an offset are allowed on a packet pointer ({where})"
)
if right_pkt and isinstance(op, ast.Sub):
raise SyntaxError(f"cannot subtract a packet pointer from an offset ({where})")

ptr, ptr_ty, off = (left, left_ty, right) if left_pkt else (right, right_ty, left)
if ptr_ty.kind == "pkt_end":
# The verifier only compares against the end of the packet.
raise SyntaxError(
f"a packet-end pointer can only be compared, not offset ({where})"
)
if isinstance(off, ir.Constant) and isinstance(off.constant, int):
if not 0 <= off.constant <= MAX_PACKET_OFF:
raise SyntaxError(
f"packet offset {off.constant} is outside 0..{MAX_PACKET_OFF} "
Comment on lines +348 to +351
f"({where}); subtract instead of adding a negative offset"
)
# The offset as u16, widened with zero-extension: a value the verifier can
# bound to [0, 0xffff].
off = builder.trunc(off, ir.IntType(16)) if off.type.width > 16 else off
off = builder.zext(off, ir.IntType(64)) if off.type.width < 64 else off
result = builder.add(ptr, off) if isinstance(op, ast.Add) else builder.sub(ptr, off)
return result, PktPtrTy(ptr_ty.kind)


def _handle_binary_op(
func,
compilation_context,
Expand Down Expand Up @@ -411,6 +468,13 @@ def _handle_compare(func, compilation_context, builder, cond, local_sym_tab):
lhs, lhs_ty = lhs
rhs, rhs_ty = rhs
lhs_desc, rhs_desc = _descriptor(lhs, lhs_ty), _descriptor(rhs, rhs_ty)
if isinstance(lhs_desc, PktPtrTy) or isinstance(rhs_desc, PktPtrTy):
# A packet bounds check (data + n > data_end): the verifier accepts
# pointer comparisons only at 64 bits, and they are unsigned.
u64 = IntTy(64, False)
lhs = convert(builder, lhs, lhs_desc, u64)
rhs = convert(builder, rhs, rhs_desc, u64)
return handle_comparator(func, builder, cond.ops[0], lhs, rhs, signed=False)
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).
Expand Down
68 changes: 68 additions & 0 deletions pythonbpf/expr/packet_pointer.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,68 @@
"""Packet-pointer context fields.

A few context fields are declared u32 in C but are packet pointers to the
verifier: it rewrites the context load into a 64-bit pointer load. C's
integer rules (which rank them as u32) therefore produce code the verifier
rejects ("32-bit pointer arithmetic prohibited"); C programs dodge that by
casting through (void *)(long). The compiler special-cases them instead, so
users never need the cast:

- the field is used at 64 bits, as the pointer it is;
- an offset added to or subtracted from it is taken as an unsigned 16-bit
value (truncated, then zero-extended), because the verifier only tracks a
packet range whose offset stays within MAX_PACKET_OFF (0xffff);
- pointer - pointer (e.g. data_end - data) is an ordinary 64-bit length.

The field read gives its value a PktPtrTy descriptor (a 64-bit unsigned
IntTy tagged with the verifier's kind), and the rules key on that descriptor,
so they follow the pointer through locals (`d = ctx.data`), through
`d + 14`, and into comparisons (`ctx.data + 34 > ctx.data_end` is a 64-bit
unsigned compare). packet-end pointers take no offset at all.

Storing a packet pointer into a narrower slot truncates it, as C does, and the
result is an ordinary integer: the verifier forbids 32-bit *arithmetic* on a
packet pointer, not a 32-bit store of one (upstream's
cgroup_skb_direct_packet_access stores data_end into a __u32 global). Keeping
the pointer when a local is declared narrower is a job for allocation, which
could give such a slot the widest type ever stored into it.
"""

import ast

from pythonbpf.type_deducer import PktPtrTy

# Context struct -> {field: packet kind}. Taken from the verifier's
# is_valid_access callbacks, restricted to the fields vmlinux declares as
# 32-bit integers.
PACKET_POINTER_FIELDS = {
"struct_xdp_md": {"data": "pkt", "data_meta": "pkt_meta", "data_end": "pkt_end"},
"struct___sk_buff": {
"data": "pkt",
"data_meta": "pkt_meta",
"data_end": "pkt_end",
},
}

MAX_PACKET_OFF = 0xFFFF


def packet_kind(expr, local_sym_tab) -> "str | None":
"""The packet kind of `ctx.<field>` when ctx is a context parameter of one
of the structs above and the field is one of its packet-pointer fields;
None otherwise. This is the only place a PktPtrTy originates: from here it
travels as a descriptor, through locals, binop results and comparisons."""
if not (isinstance(expr, ast.Attribute) and isinstance(expr.value, ast.Name)):
return None
sym = (local_sym_tab or {}).get(expr.value.id)
if sym is None:
return None
meta = sym.metadata
struct_name = meta if isinstance(meta, str) else getattr(meta, "__name__", None)
if not isinstance(struct_name, str):
return None
return PACKET_POINTER_FIELDS.get(struct_name, {}).get(expr.attr)


def packet_type(expr, local_sym_tab) -> "PktPtrTy | None":
kind = packet_kind(expr, local_sym_tab)
return PktPtrTy(kind) if kind is not None else None
11 changes: 11 additions & 0 deletions pythonbpf/expr/type_inference.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
from llvmlite import ir

from pythonbpf.type_deducer import (
PktPtrTy,
IntTy,
ctypes_to_ir,
is_ctypes,
Expand All @@ -21,9 +22,12 @@
)
from .operators import usual_arithmetic_conversions
from .vmlinux_registry import VmlinuxHandlerRegistry
from .packet_pointer import packet_type


def _as_intty(ty):
if isinstance(ty, PktPtrTy):
return ty
if isinstance(ty, ir.IntType):
return IntTy(ty.width, signedness(ty))
return None
Expand All @@ -48,6 +52,10 @@ def infer_int_type(expr, local_sym_tab, compilation_context):
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 isinstance(left, PktPtrTy) and isinstance(right, PktPtrTy):
return IntTy(64, False) # pointer - pointer: a length
if isinstance(left, PktPtrTy) or isinstance(right, PktPtrTy):
return left if isinstance(left, PktPtrTy) else right
if left is None or right is None:
return None
return usual_arithmetic_conversions(left, right)
Expand Down Expand Up @@ -83,6 +91,9 @@ def infer_int_type(expr, local_sym_tab, compilation_context):
return None

if isinstance(expr, ast.Attribute) and isinstance(expr.value, ast.Name):
pkt_ty = packet_type(expr, local_sym_tab)
if pkt_ty is not None:
return pkt_ty
base = local_sym_tab.get(expr.value.id)
if base is None:
return None
Expand Down
27 changes: 27 additions & 0 deletions pythonbpf/type_deducer.py
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,33 @@ def describe(self) -> str:
return f"{'i' if self.signed else 'u'}{self.width}"


PACKET_KINDS = ("pkt", "pkt_meta", "pkt_end")


class PktPtrTy(IntTy):
"""A packet pointer: 64-bit, unsigned, and tagged with the verifier's kind
for it (packet data, packet metadata, or packet end). It is an IntTy so
every integer path accepts it; the few sites that must treat it as a
pointer check isinstance(ty, PktPtrTy) before the integer rules apply.
Comment on lines +50 to +51
See pythonbpf/expr/packet_pointer.py for where kinds come from and the
rules applied to them."""

def __new__(cls, kind: str):
return IntTy.__new__(cls, 64, False)

def __init__(self, kind: str):
if kind not in PACKET_KINDS:
raise ValueError(f"unknown packet pointer kind {kind!r}")
IntTy.__init__(self, 64, False)
self.kind = kind

def __getnewargs__(self): # type: ignore[override]
return (self.kind,)

def describe(self) -> str:
return f"{self.kind}*"


def byte_size(ty) -> int:
"""Bytes an integer type occupies in memory: its width rounded up to whole
bytes, so a 1-bit bool still takes one byte, as C's _Bool does."""
Expand Down
21 changes: 21 additions & 0 deletions tests/failing_tests/vmlinux/ctx_data_end_offset.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,21 @@
# data_end can only be compared against; the verifier rejects arithmetic on
# it, so the compiler refuses it.
from ctypes import c_int64, c_uint32 # noqa: F401
from pythonbpf import bpf, section, bpfglobal, compile
from vmlinux import struct_xdp_md


@bpf
@section("xdp")
def prog(ctx: struct_xdp_md) -> c_int64:
e = ctx.data_end + 1
return c_int64(e)


@bpf
@bpfglobal
def LICENSE() -> str:
return "GPL"


compile()
21 changes: 21 additions & 0 deletions tests/failing_tests/vmlinux/ctx_data_offset_too_big.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,21 @@
# 70000 cannot be a packet offset: the verifier tracks packet ranges only
# up to 0xffff. Rejected at compile time rather than silently wrapped.
from ctypes import 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:
d = ctx.data + 70000
return c_int64(d)


@bpf
@bpfglobal
def LICENSE() -> str:
return "GPL"


compile()
25 changes: 25 additions & 0 deletions tests/passing_tests/vmlinux/ctx_bounds_check.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,25 @@
# The canonical XDP bounds check, through locals: the packet-pointer type
# travels with data and data_end, so data + 34 > data_end is a 64-bit
# unsigned comparison, which is what the verifier needs to prove the range.
from ctypes import c_int64, c_uint32 # noqa: F401
from pythonbpf import bpf, section, bpfglobal, compile
from vmlinux import struct_xdp_md


@bpf
@section("xdp")
def prog(ctx: struct_xdp_md) -> c_int64:
data = ctx.data
data_end = ctx.data_end
if data + 34 > data_end:
return c_int64(1)
return c_int64(2)


@bpf
@bpfglobal
def LICENSE() -> str:
return "GPL"


compile()
21 changes: 21 additions & 0 deletions tests/passing_tests/vmlinux/ctx_bounds_check_direct.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,21 @@
# The same bounds check written on the fields directly.
from ctypes import c_int64, c_uint32 # noqa: F401
from pythonbpf import bpf, section, bpfglobal, compile
from vmlinux import struct_xdp_md


@bpf
@section("xdp")
def prog(ctx: struct_xdp_md) -> c_int64:
if ctx.data + 14 > ctx.data_end:
return c_int64(1)
return c_int64(2)


@bpf
@bpfglobal
def LICENSE() -> str:
return "GPL"


compile()
Loading
Loading